Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
53 changes: 19 additions & 34 deletions api/controllers/console/workspace/model_providers.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -15,6 +15,7 @@
RBACResourceScope,
account_initialization_required,
is_admin_or_owner_required,
model_validate,
rbac_permission_required,
setup_required,
with_current_tenant_id,
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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
)
Expand All @@ -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:
Expand All @@ -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:
Expand All @@ -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
Expand All @@ -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,
Expand All @@ -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()
Expand Down Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
@@ -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 (
Expand All @@ -17,6 +19,14 @@
ModelProviderPaymentCheckoutUrlApi,
ModelProviderSummaryListApi,
ModelProviderValidateApi,
ParserCredentialCreate,
ParserCredentialDelete,
ParserCredentialId,
ParserCredentialSwitch,
ParserCredentialUpdate,
ParserCredentialValidate,
ParserModelList,
ParserPreferredProviderType,
PreferredProviderTypeUpdateApi,
)
from core.entities.provider_entities import CredentialConfiguration
Expand All @@ -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,
Expand Down Expand Up @@ -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()
Expand All @@ -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()]}
Expand All @@ -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": []}
Expand Down Expand Up @@ -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
Expand All @@ -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}
Expand All @@ -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()
Expand All @@ -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",
Expand All @@ -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()
Expand All @@ -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",
Expand All @@ -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()
Expand All @@ -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
Expand All @@ -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",
Expand All @@ -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:
Expand All @@ -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"}
Expand All @@ -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"}

Expand Down Expand Up @@ -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"
Expand All @@ -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:
Expand Down Expand Up @@ -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()
Loading