diff --git a/api/controllers/console/workspace/model_providers.py b/api/controllers/console/workspace/model_providers.py index ca0b536f85ed50..fe844204d276bd 100644 --- a/api/controllers/console/workspace/model_providers.py +++ b/api/controllers/console/workspace/model_providers.py @@ -1,7 +1,7 @@ import io from typing import Any, Literal -from flask import request, send_file +from flask import send_file from flask_restx import Resource from pydantic import BaseModel, Field, field_validator from sqlalchemy.orm import Session @@ -15,6 +15,7 @@ RBACResourceScope, account_initialization_required, is_admin_or_owner_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -154,10 +155,8 @@ class ModelProviderListApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, tenant_id: str): - payload = request.args.to_dict(flat=True) - args = ParserModelList.model_validate(payload) - + @model_validate(ParserModelList) + def get(self, args: ParserModelList, tenant_id: str): model_provider_service = ModelProviderService() provider_list = model_provider_service.get_provider_list(tenant_id=tenant_id, model_type=args.model_type) @@ -212,12 +211,10 @@ class ModelProviderCredentialApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.CREDENTIAL_MANAGE, resource_required=False) @account_initialization_required @with_current_tenant_id - def get(self, tenant_id: str, provider: str): - # if credential_id is not provided, return current used credential - payload = request.args.to_dict(flat=True) - args = ParserCredentialId.model_validate(payload) - + @model_validate(ParserCredentialId) + def get(self, args: ParserCredentialId, tenant_id: str, provider: str): model_provider_service = ModelProviderService() + # if credential_id is not provided, return current used credential credentials = model_provider_service.get_provider_credential( tenant_id=tenant_id, provider=provider, credential_id=args.credential_id ) @@ -232,10 +229,8 @@ def get(self, tenant_id: str, provider: str): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.CREDENTIAL_CREATE, resource_required=False) @account_initialization_required @with_current_tenant_id - def post(self, current_tenant_id: str, provider: str): - payload = console_ns.payload or {} - args = ParserCredentialCreate.model_validate(payload) - + @model_validate(ParserCredentialCreate) + def post(self, args: ParserCredentialCreate, current_tenant_id: str, provider: str): model_provider_service = ModelProviderService() try: @@ -258,10 +253,8 @@ def post(self, current_tenant_id: str, provider: str): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.CREDENTIAL_MANAGE, resource_required=False) @account_initialization_required @with_current_tenant_id - def put(self, current_tenant_id: str, provider: str): - payload = console_ns.payload or {} - args = ParserCredentialUpdate.model_validate(payload) - + @model_validate(ParserCredentialUpdate) + def put(self, args: ParserCredentialUpdate, current_tenant_id: str, provider: str): model_provider_service = ModelProviderService() try: @@ -285,10 +278,8 @@ def put(self, current_tenant_id: str, provider: str): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.CREDENTIAL_MANAGE, resource_required=False) @account_initialization_required @with_current_tenant_id - def delete(self, current_tenant_id: str, provider: str): - payload = console_ns.payload or {} - args = ParserCredentialDelete.model_validate(payload) - + @model_validate(ParserCredentialDelete) + def delete(self, args: ParserCredentialDelete, current_tenant_id: str, provider: str): model_provider_service = ModelProviderService() model_provider_service.remove_provider_credential( tenant_id=current_tenant_id, provider=provider, credential_id=args.credential_id @@ -307,10 +298,8 @@ class ModelProviderCredentialSwitchApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.CREDENTIAL_USE, resource_required=False) @account_initialization_required @with_current_tenant_id - def post(self, current_tenant_id: str, provider: str): - payload = console_ns.payload or {} - args = ParserCredentialSwitch.model_validate(payload) - + @model_validate(ParserCredentialSwitch) + def post(self, args: ParserCredentialSwitch, current_tenant_id: str, provider: str): service = ModelProviderService() service.switch_active_provider_credential( tenant_id=current_tenant_id, @@ -332,10 +321,8 @@ class ModelProviderValidateApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def post(self, current_tenant_id: str, provider: str): - payload = console_ns.payload or {} - args = ParserCredentialValidate.model_validate(payload) - + @model_validate(ParserCredentialValidate) + def post(self, args: ParserCredentialValidate, current_tenant_id: str, provider: str): tenant_id = current_tenant_id model_provider_service = ModelProviderService() @@ -388,10 +375,8 @@ class PreferredProviderTypeUpdateApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.CREDENTIAL_USE, resource_required=False) @account_initialization_required @with_current_tenant_id - def post(self, tenant_id: str, provider: str): - payload = console_ns.payload or {} - args = ParserPreferredProviderType.model_validate(payload) - + @model_validate(ParserPreferredProviderType) + def post(self, args: ParserPreferredProviderType, tenant_id: str, provider: str): model_provider_service = ModelProviderService() model_provider_service.switch_preferred_provider( tenant_id=tenant_id, provider=provider, preferred_provider_type=args.preferred_provider_type diff --git a/api/tests/unit_tests/controllers/console/workspace/test_model_providers.py b/api/tests/unit_tests/controllers/console/workspace/test_model_providers.py index a4f14e5255ddcb..2af0ce3db7abf1 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_model_providers.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_model_providers.py @@ -1,11 +1,13 @@ +from collections.abc import Callable from inspect import unwrap +from types import SimpleNamespace from unittest.mock import patch import pytest -from flask import Flask, g +from flask import Flask, g, request from pydantic_core import ValidationError from sqlalchemy.orm import Session -from werkzeug.exceptions import Forbidden +from werkzeug.exceptions import Forbidden, UnprocessableEntity from configs import dify_config from controllers.console.workspace.model_providers import ( @@ -17,6 +19,14 @@ ModelProviderPaymentCheckoutUrlApi, ModelProviderSummaryListApi, ModelProviderValidateApi, + ParserCredentialCreate, + ParserCredentialDelete, + ParserCredentialId, + ParserCredentialSwitch, + ParserCredentialUpdate, + ParserCredentialValidate, + ParserModelList, + ParserPreferredProviderType, PreferredProviderTypeUpdateApi, ) from core.entities.provider_entities import CredentialConfiguration @@ -26,6 +36,7 @@ from graphon.model_runtime.entities.provider_entities import ConfigurateMethod from graphon.model_runtime.errors.validate import CredentialsValidateFailedError from models import Account +from models.account import TenantAccountRole from models.provider import ProviderType from services.entities.model_provider_entities import ( CustomConfigurationResponse, @@ -116,6 +127,19 @@ def expected_provider_payload() -> dict[str, object]: } +def _payload() -> dict[str, object]: + """Mirror the source the ``@model_validate`` decorator reads for the request in scope. + + The tests build their contexts with ``test_request_context``, which defaults to GET even for + handlers mounted on POST, so fall back to the JSON body whenever the query string is empty. + """ + args: dict[str, object] = dict(request.args.to_dict(flat=True)) + if args: + return args + body = request.get_json(silent=True) + return dict(body) if isinstance(body, dict) else {} + + class TestModelProviderListApi: def test_get_success(self, app: Flask): api = ModelProviderListApi() @@ -129,7 +153,7 @@ def test_get_success(self, app: Flask): return_value=[provider], ) as get_provider_list, ): - result = method(api, "tenant1") + result = method(api, ParserModelList.model_validate(_payload()), "tenant1") get_provider_list.assert_called_once_with(tenant_id="tenant1", model_type=ModelType.LLM) assert result == {"data": [expected_provider_payload()]} @@ -145,7 +169,7 @@ def test_get_without_model_type_passes_none(self, app: Flask): return_value=[], ) as get_provider_list, ): - result = method(api, "tenant1") + result = method(api, ParserModelList.model_validate(_payload()), "tenant1") get_provider_list.assert_called_once_with(tenant_id="tenant1", model_type=None) assert result == {"data": []} @@ -297,7 +321,7 @@ def test_get_success(self, app: Flask): }, ) as get_provider_credential, ): - result = method(api, "tenant1", provider="openai") + result = method(api, ParserCredentialId.model_validate(_payload()), "tenant1", provider="openai") get_provider_credential.assert_called_once_with( tenant_id="tenant1", provider="openai", credential_id=VALID_UUID @@ -321,7 +345,7 @@ def test_get_current_credential_without_id(self, app: Flask): return_value=None, ) as get_provider_credential, ): - result = method(api, "tenant1", provider="openai") + result = method(api, ParserCredentialId.model_validate(_payload()), "tenant1", provider="openai") get_provider_credential.assert_called_once_with(tenant_id="tenant1", provider="openai", credential_id=None) assert result == {"credentials": None} @@ -332,7 +356,7 @@ def test_get_invalid_uuid(self, app: Flask): with app.test_request_context(f"/?credential_id={INVALID_UUID}"): with pytest.raises(ValidationError): - method(api, "tenant1", provider="openai") + method(api, ParserCredentialId.model_validate(_payload()), "tenant1", provider="openai") def test_post_create_success(self, app: Flask): api = ModelProviderCredentialApi() @@ -347,7 +371,9 @@ def test_post_create_success(self, app: Flask): return_value=None, ) as create_provider_credential, ): - result, status = method(api, "tenant1", provider="openai") + result, status = method( + api, ParserCredentialCreate.model_validate(_payload()), "tenant1", provider="openai" + ) create_provider_credential.assert_called_once_with( tenant_id="tenant1", @@ -372,7 +398,7 @@ def test_post_create_validation_error(self, app: Flask): ), ): with pytest.raises(ValueError): - method(api, "tenant1", provider="openai") + method(api, ParserCredentialCreate.model_validate(_payload()), "tenant1", provider="openai") def test_put_update_success(self, app: Flask): api = ModelProviderCredentialApi() @@ -387,7 +413,7 @@ def test_put_update_success(self, app: Flask): return_value=None, ) as update_provider_credential, ): - result = method(api, "tenant1", provider="openai") + result = method(api, ParserCredentialUpdate.model_validate(_payload()), "tenant1", provider="openai") update_provider_credential.assert_called_once_with( tenant_id="tenant1", @@ -406,7 +432,7 @@ def test_put_invalid_uuid(self, app: Flask): with app.test_request_context("/", json=payload): with pytest.raises(ValidationError): - method(api, "tenant1", provider="openai") + method(api, ParserCredentialUpdate.model_validate(_payload()), "tenant1", provider="openai") def test_delete_success(self, app: Flask): api = ModelProviderCredentialApi() @@ -421,7 +447,9 @@ def test_delete_success(self, app: Flask): return_value=None, ) as remove_provider_credential, ): - result, status = method(api, "tenant1", provider="openai") + result, status = method( + api, ParserCredentialDelete.model_validate(_payload()), "tenant1", provider="openai" + ) remove_provider_credential.assert_called_once_with( tenant_id="tenant1", provider="openai", credential_id=VALID_UUID @@ -444,7 +472,7 @@ def test_switch_success(self, app: Flask): return_value=None, ) as switch_active_provider_credential, ): - result = method(api, "tenant1", provider="openai") + result = method(api, ParserCredentialSwitch.model_validate(_payload()), "tenant1", provider="openai") switch_active_provider_credential.assert_called_once_with( tenant_id="tenant1", @@ -461,7 +489,7 @@ def test_switch_invalid_uuid(self, app: Flask): with app.test_request_context("/", json=payload): with pytest.raises(ValidationError): - method(api, "tenant1", provider="openai") + method(api, ParserCredentialSwitch.model_validate(_payload()), "tenant1", provider="openai") class TestModelProviderValidateApi: @@ -478,7 +506,7 @@ def test_validate_success(self, app: Flask): return_value=None, ) as validate_provider_credentials, ): - result = method(api, "tenant1", provider="openai") + result = method(api, ParserCredentialValidate.model_validate(_payload()), "tenant1", provider="openai") validate_provider_credentials.assert_called_once_with( tenant_id="tenant1", provider="openai", credentials={"a": "b"} @@ -498,7 +526,7 @@ def test_validate_failure(self, app: Flask): side_effect=CredentialsValidateFailedError("bad"), ), ): - result = method(api, "tenant1", provider="openai") + result = method(api, ParserCredentialValidate.model_validate(_payload()), "tenant1", provider="openai") assert result == {"result": "error", "error": "bad"} @@ -549,7 +577,7 @@ def test_update_success(self, app: Flask): return_value=None, ) as switch_preferred_provider, ): - result = method(api, "tenant1", provider="openai") + result = method(api, ParserPreferredProviderType.model_validate(_payload()), "tenant1", provider="openai") switch_preferred_provider.assert_called_once_with( tenant_id="tenant1", provider="openai", preferred_provider_type="custom" @@ -564,7 +592,7 @@ def test_invalid_enum(self, app: Flask): with app.test_request_context("/", json=payload): with pytest.raises(ValidationError): - method(api, "tenant1", provider="openai") + method(api, ParserPreferredProviderType.model_validate(_payload()), "tenant1", provider="openai") class TestModelProviderPaymentCheckoutUrlApi: @@ -618,3 +646,57 @@ def test_checkout_rejects_non_privileged_role(self, app: Flask): api.get(provider="anthropic") get_model_provider_payment_link.assert_not_called() + + +class TestModelValidateDecorator: + """The tests above unwrap the view, so this is what covers the decorators themselves.""" + + EMPTY_BODY: dict[str, object] = {} + EMPTY_KWARGS: dict[str, object] = {} + + @pytest.mark.parametrize( + ("verb", "url", "body", "call"), + [ + ("GET", "/?model_type=not-a-model-type", None, lambda: ModelProviderListApi().get()), + ( + "GET", + "/?credential_id=not-a-uuid", + None, + lambda: ModelProviderCredentialApi().get(provider="openai"), + ), + ("POST", "/", EMPTY_BODY, lambda: ModelProviderCredentialApi().post(provider="openai")), + ("PUT", "/", EMPTY_BODY, lambda: ModelProviderCredentialApi().put(provider="openai")), + ("DELETE", "/", EMPTY_BODY, lambda: ModelProviderCredentialApi().delete(provider="openai")), + ("POST", "/", EMPTY_BODY, lambda: ModelProviderCredentialSwitchApi().post(provider="openai")), + ("POST", "/", EMPTY_BODY, lambda: ModelProviderValidateApi().post(provider="openai")), + ("POST", "/", EMPTY_BODY, lambda: PreferredProviderTypeUpdateApi().post(provider="openai")), + ], + ) + def test_invalid_input_is_rejected_before_the_handler_runs( + self, + app: Flask, + verb: str, + url: str, + body: dict[str, object] | None, + call: Callable[[], object], + ) -> None: + account = make_account() + account.role = TenantAccountRole.OWNER + + with ( + app.test_request_context(url, method=verb, json=body), + config_overrides_context(LOGIN_DISABLED=True, RBAC_ENABLED=False), + patch("controllers.console.wraps._is_setup_completed", return_value=True), + patch( + "controllers.console.wraps.current_account_with_tenant", + return_value=(account, "tenant1"), + ), + patch("libs.login.current_user", SimpleNamespace(_get_current_object=lambda: account)), + patch("controllers.console.workspace.model_providers.ModelProviderService") as service, + ): + g._login_user = account + with pytest.raises(UnprocessableEntity) as exc_info: + call() + + assert exc_info.value.code == 422 + service.assert_not_called()