refactor: replace manual model_validate with @model_validate in workspace remaining controllers (#40239)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
Likalikali
2026-08-09 08:23:51 +00:00
committed by GitHub
co-authored by autofix-ci[bot]
parent 5ebb389d26
commit b692ddd80c
10 changed files with 303 additions and 172 deletions
@@ -16,6 +16,7 @@ from controllers.console.auth.error import (
from controllers.console.error import AccountInFreezeError
from controllers.console.workspace.account import (
AccountAvatarApi,
AccountAvatarQuery,
AccountDeleteApi,
AccountDeleteVerifyApi,
AccountInitApi,
@@ -216,7 +217,7 @@ class TestAccountAvatarApiGet:
return_value="https://signed/example",
) as sign_mock,
):
result = method(api, user)
result = method(api, AccountAvatarQuery(avatar=file_id), user)
assert result == {"avatar_url": "https://signed/example"}
sign_mock.assert_called_once_with(upload_file_id=file_id)
@@ -256,7 +257,7 @@ class TestAccountAvatarApiGet:
) as sign_mock,
):
with pytest.raises(NotFound):
method(api, user)
method(api, AccountAvatarQuery(avatar=file_id), user)
sign_mock.assert_not_called()
@@ -289,7 +290,7 @@ class TestAccountAvatarApiGet:
return_value="https://signed/example",
) as sign_mock,
):
result = method(api, user)
result = method(api, AccountAvatarQuery(avatar=file_id), user)
assert result == {"avatar_url": "https://signed/example"}
sign_mock.assert_called_once_with(upload_file_id=file_id)
@@ -308,7 +309,7 @@ class TestAccountAvatarApiGet:
return_value="https://signed/should-not-use",
) as sign_mock,
):
result = method(api, user)
result = method(api, AccountAvatarQuery(avatar=external), user)
assert result == {"avatar_url": external}
sign_mock.assert_not_called()
@@ -17,7 +17,9 @@ from controllers.console.workspace.endpoint import (
EndpointIdPayload,
EndpointItemApi,
EndpointListApi,
EndpointListForPluginQuery,
EndpointListForSinglePluginApi,
EndpointListQuery,
EndpointUpdatePayload,
LegacyEndpointUpdatePayload,
)
@@ -146,7 +148,7 @@ class TestEndpointListApi:
return_value=[endpoint_entity],
),
):
result = method(api, "t1", "u1")
result = method(api, EndpointListQuery(page=1, page_size=10), "t1", "u1")
endpoint = result["endpoints"][0]
assert endpoint["id"] == "e1"
@@ -180,7 +182,7 @@ class TestEndpointListApi:
app.test_request_context("/?page=0&page_size=10"),
):
with pytest.raises(ValueError):
method(api, "t1", "u1")
method(api, EndpointListQuery(page=0, page_size=10), "t1", "u1")
class TestEndpointListForSinglePluginApi:
@@ -195,7 +197,7 @@ class TestEndpointListForSinglePluginApi:
return_value=[_endpoint_entity()],
),
):
result = method(api, "t1", "u1")
result = method(api, EndpointListForPluginQuery(page=1, page_size=10, plugin_id="p1"), "t1", "u1")
assert result["endpoints"][0]["id"] == "e1"
assert result["endpoints"][0]["settings"]["api_key"] == "pl********et"
@@ -209,7 +211,7 @@ class TestEndpointListForSinglePluginApi:
app.test_request_context("/?page=1&page_size=10"),
):
with pytest.raises(ValueError):
method(api, "t1", "u1")
method(api, EndpointListForPluginQuery(page=1, page_size=10), "t1", "u1")
class TestEndpointItemApi:
@@ -15,6 +15,16 @@ from controllers.console.workspace.models import (
ModelProviderModelEnableApi,
ModelProviderModelParameterRuleApi,
ModelProviderModelValidateApi,
ParserCreateCredential,
ParserDeleteCredential,
ParserDeleteModels,
ParserGetCredentials,
ParserGetDefault,
ParserParameter,
ParserPostDefault,
ParserPostModels,
ParserSwitch,
ParserValidate,
)
from graphon.model_runtime.entities.model_entities import ModelType
from graphon.model_runtime.errors.validate import CredentialsValidateFailedError
@@ -43,7 +53,7 @@ class TestDefaultModelApi:
},
}
result = method(api, "tenant1")
result = method(api, ParserGetDefault(model_type=ModelType.LLM), "tenant1")
assert "data" in result
@@ -65,7 +75,7 @@ class TestDefaultModelApi:
app.test_request_context("/", json=payload),
patch("controllers.console.workspace.models.ModelProviderService"),
):
result = method(api, "tenant1")
result = method(api, ParserPostDefault.model_validate(payload), "tenant1")
assert result["result"] == "success"
@@ -79,7 +89,7 @@ class TestDefaultModelApi:
):
service.return_value.get_default_model_of_model_type.return_value = None
result = method(api, "t1")
result = method(api, ParserGetDefault(model_type=ModelType.LLM), "t1")
assert "data" in result
@@ -117,7 +127,7 @@ class TestModelProviderModelApi:
patch("controllers.console.workspace.models.ModelProviderService"),
patch("controllers.console.workspace.models.ModelLoadBalancingService"),
):
result, status = method(api, "tenant1", "openai")
result, status = method(api, ParserPostModels.model_validate(payload), "tenant1", "openai")
assert status == 200
@@ -134,7 +144,7 @@ class TestModelProviderModelApi:
app.test_request_context("/", json=payload),
patch("controllers.console.workspace.models.ModelProviderService"),
):
result, status = method(api, "tenant1", "openai")
result, status = method(api, ParserDeleteModels.model_validate(payload), "tenant1", "openai")
assert status == 204
@@ -177,7 +187,13 @@ class TestModelProviderModelCredentialApi:
provider_service.return_value.provider_manager.get_provider_model_available_credentials.return_value = []
lb_service.return_value.get_load_balancing_configs.return_value = (False, [])
result = method(api, "tenant1", SimpleNamespace(id="u1"), "openai")
result = method(
api,
ParserGetCredentials(model="gpt-4", model_type=ModelType.LLM),
"tenant1",
SimpleNamespace(id="u1"),
"openai",
)
assert "credentials" in result
@@ -195,7 +211,7 @@ class TestModelProviderModelCredentialApi:
app.test_request_context("/", json=payload),
patch("controllers.console.workspace.models.ModelProviderService"),
):
result, status = method(api, "tenant1", "openai")
result, status = method(api, ParserCreateCredential.model_validate(payload), "tenant1", "openai")
assert status == 201
@@ -212,7 +228,13 @@ class TestModelProviderModelCredentialApi:
service.return_value.provider_manager.get_provider_model_available_credentials.return_value = []
lb.return_value.get_load_balancing_configs.return_value = (False, [])
result = method(api, "t1", SimpleNamespace(id="u1"), "openai")
result = method(
api,
ParserGetCredentials(model="gpt", model_type=ModelType.LLM),
"t1",
SimpleNamespace(id="u1"),
"openai",
)
assert result["credentials"] == {}
@@ -230,7 +252,7 @@ class TestModelProviderModelCredentialApi:
app.test_request_context("/", json=payload),
patch("controllers.console.workspace.models.ModelProviderService"),
):
result, status = method(api, "t1", "openai")
result, status = method(api, ParserDeleteCredential.model_validate(payload), "t1", "openai")
assert status == 204
@@ -250,7 +272,7 @@ class TestModelProviderModelCredentialSwitchApi:
app.test_request_context("/", json=payload),
patch("controllers.console.workspace.models.ModelProviderService"),
):
result = method(api, "tenant1", "openai")
result = method(api, ParserSwitch.model_validate(payload), "tenant1", "openai")
assert result["result"] == "success"
@@ -269,7 +291,7 @@ class TestModelEnableDisableApis:
app.test_request_context("/", json=payload),
patch("controllers.console.workspace.models.ModelProviderService"),
):
result = method(api, "tenant1", "openai")
result = method(api, ParserDeleteModels.model_validate(payload), "tenant1", "openai")
assert result["result"] == "success"
@@ -286,7 +308,7 @@ class TestModelEnableDisableApis:
app.test_request_context("/", json=payload),
patch("controllers.console.workspace.models.ModelProviderService"),
):
result = method(api, "tenant1", "openai")
result = method(api, ParserDeleteModels.model_validate(payload), "tenant1", "openai")
assert result["result"] == "success"
@@ -306,7 +328,7 @@ class TestModelProviderModelValidateApi:
app.test_request_context("/", json=payload),
patch("controllers.console.workspace.models.ModelProviderService"),
):
result = method(api, "tenant1", "openai")
result = method(api, ParserValidate.model_validate(payload), "tenant1", "openai")
assert result["result"] == "success"
@@ -327,7 +349,7 @@ class TestModelProviderModelValidateApi:
):
service_mock.return_value.validate_model_credentials.side_effect = CredentialsValidateFailedError("invalid")
result = method(api, "tenant1", "openai")
result = method(api, ParserValidate.model_validate(payload), "tenant1", "openai")
assert result["result"] == "error"
@@ -343,7 +365,7 @@ class TestParameterAndAvailableModels:
):
service_mock.return_value.get_model_parameter_rules.return_value = []
result = method(api, "tenant1", "openai")
result = method(api, ParserParameter(model="gpt-4"), "tenant1", "openai")
assert "data" in result
@@ -371,7 +393,7 @@ class TestParameterAndAvailableModels:
):
service.return_value.get_model_parameter_rules.return_value = []
result = method(api, "t1", "openai")
result = method(api, ParserParameter(model="gpt"), "t1", "openai")
assert result["data"] == []
@@ -26,6 +26,7 @@ from werkzeug.exceptions import Forbidden, NotFound
from configs import dify_config
from controllers.console.workspace import rbac as rbac_mod
from controllers.console.workspace.rbac import _RolesListQuery
@pytest.fixture
@@ -175,7 +176,10 @@ class TestPaginationMapping:
patch("controllers.console.workspace.rbac._current_ids", return_value=("tenant-1", "acct-1")),
patch("controllers.console.workspace.rbac.svc.RBACService.Roles.list") as mock_list,
):
response = inspect.unwrap(rbac_mod.RBACRolesApi.get)(rbac_mod.RBACRolesApi())
response = inspect.unwrap(rbac_mod.RBACRolesApi.get)(
rbac_mod.RBACRolesApi(),
_RolesListQuery.model_validate({"page": 1, "limit": 2, "include_owner": 1}),
)
owner_permission_keys = rbac_mod._LEGACY_ROLE_PERMISSION_KEYS["owner"]
valid_owner_permission_keys = []
@@ -230,7 +234,7 @@ class TestPaginationMapping:
patch("controllers.console.workspace.rbac._current_ids", return_value=("tenant-1", "acct-1")),
patch("controllers.console.workspace.rbac.svc.RBACService.Roles.list"),
):
response = inspect.unwrap(rbac_mod.RBACRolesApi.get)(rbac_mod.RBACRolesApi())
response = inspect.unwrap(rbac_mod.RBACRolesApi.get)(rbac_mod.RBACRolesApi(), _RolesListQuery())
names = [r["name"] for r in response["data"]]
assert "owner" not in names
@@ -242,7 +246,10 @@ class TestPaginationMapping:
patch("controllers.console.workspace.rbac._current_ids", return_value=("tenant-1", "acct-1")),
patch("controllers.console.workspace.rbac.svc.RBACService.Roles.list"),
):
response = inspect.unwrap(rbac_mod.RBACRolesApi.get)(rbac_mod.RBACRolesApi())
response = inspect.unwrap(rbac_mod.RBACRolesApi.get)(
rbac_mod.RBACRolesApi(),
_RolesListQuery.model_validate({"include_owner": 1}),
)
names = [r["name"] for r in response["data"]]
assert "owner" in names
@@ -254,7 +261,7 @@ class TestPaginationMapping:
patch("controllers.console.workspace.rbac._current_ids", return_value=("tenant-1", "acct-1")),
patch("controllers.console.workspace.rbac.svc.RBACService.Roles.list"),
):
response = inspect.unwrap(rbac_mod.RBACRolesApi.get)(rbac_mod.RBACRolesApi())
response = inspect.unwrap(rbac_mod.RBACRolesApi.get)(rbac_mod.RBACRolesApi(), _RolesListQuery())
names = [r["name"] for r in response["data"]]
assert "owner" not in names
@@ -267,7 +274,10 @@ class TestPaginationMapping:
patch("controllers.console.workspace.rbac.svc.RBACService.Roles.list") as mock_list,
patch("controllers.console.workspace.rbac._dump", return_value={}),
):
inspect.unwrap(rbac_mod.RBACRolesApi.get)(rbac_mod.RBACRolesApi())
inspect.unwrap(rbac_mod.RBACRolesApi.get)(
rbac_mod.RBACRolesApi(),
_RolesListQuery.model_validate({"page": 2, "limit": 50, "reverse": True, "include_owner": 1}),
)
_, kwargs = mock_list.call_args
options = kwargs["options"]
@@ -15,15 +15,19 @@ from controllers.console.workspace.trigger_providers import (
TriggerOAuthAuthorizeApi,
TriggerOAuthCallbackApi,
TriggerOAuthClientManageApi,
TriggerOAuthClientPayload,
TriggerProviderIconApi,
TriggerProviderInfoApi,
TriggerProviderListApi,
TriggerSubscriptionBuilderBuildApi,
TriggerSubscriptionBuilderCreateApi,
TriggerSubscriptionBuilderCreatePayload,
TriggerSubscriptionBuilderGetApi,
TriggerSubscriptionBuilderLogsApi,
TriggerSubscriptionBuilderUpdateApi,
TriggerSubscriptionBuilderUpdatePayload,
TriggerSubscriptionBuilderVerifyApi,
TriggerSubscriptionBuilderVerifyPayload,
TriggerSubscriptionListApi,
TriggerSubscriptionUpdateApi,
TriggerSubscriptionVerifyApi,
@@ -163,7 +167,13 @@ class TestTriggerSubscriptionBuilderApis:
return_value=subscription_builder(),
),
):
result = method(api, "t1", mock_user(), "github")
result = method(
api,
TriggerSubscriptionBuilderCreatePayload(credential_type="UNAUTHORIZED"),
"t1",
mock_user(),
"github",
)
assert result["subscription_builder"]["id"] == "b1"
def test_get_builder(self, app: Flask) -> None:
@@ -196,7 +206,14 @@ class TestTriggerSubscriptionBuilderApis:
return_value={"verified": True},
),
):
assert method(api, "t1", mock_user(), "github", "b1") == {"verified": True}
assert method(
api,
TriggerSubscriptionBuilderVerifyPayload(credentials={"a": 1}),
"t1",
mock_user(),
"github",
"b1",
) == {"verified": True}
def test_verify_builder_error(self, app: Flask) -> None:
api = TriggerSubscriptionBuilderVerifyApi()
@@ -210,7 +227,14 @@ class TestTriggerSubscriptionBuilderApis:
),
):
with pytest.raises(ValueError):
method(api, "t1", mock_user(), "github", "b1")
method(
api,
TriggerSubscriptionBuilderVerifyPayload(credentials={}),
"t1",
mock_user(),
"github",
"b1",
)
def test_update_builder(self, app: Flask) -> None:
api = TriggerSubscriptionBuilderUpdateApi()
@@ -223,7 +247,17 @@ class TestTriggerSubscriptionBuilderApis:
return_value=subscription_builder(),
) as mock_update_builder,
):
assert method(api, "t1", mock_user(), "github", "b1")["id"] == "b1"
assert (
method(
api,
TriggerSubscriptionBuilderUpdatePayload(name="n"),
"t1",
mock_user(),
"github",
"b1",
)["id"]
== "b1"
)
mock_update_builder.assert_called_once_with(
tenant_id="t1",
user_id="u1",
@@ -263,7 +297,14 @@ class TestTriggerSubscriptionBuilderApis:
return_value=None,
),
):
assert method(api, "t1", mock_user(), "github", "b1") == {"result": "success"}
assert method(
api,
TriggerSubscriptionBuilderUpdatePayload(name="x"),
"t1",
mock_user(),
"github",
"b1",
) == {"result": "success"}
class TestTriggerSubscriptionCrud:
@@ -283,7 +324,12 @@ class TestTriggerSubscriptionCrud:
),
patch("controllers.console.workspace.trigger_providers.TriggerProviderService.update_trigger_subscription"),
):
assert method(api, "t1", "s1") == {"result": "success"}
assert method(
api,
TriggerSubscriptionBuilderUpdatePayload(name="x"),
"t1",
"s1",
) == {"result": "success"}
def test_update_not_found(self, app: Flask) -> None:
api = TriggerSubscriptionUpdateApi()
@@ -297,7 +343,7 @@ class TestTriggerSubscriptionCrud:
),
):
with pytest.raises(NotFoundError):
method(api, "t1", "x")
method(api, TriggerSubscriptionBuilderUpdatePayload(name="x"), "t1", "x")
def test_update_rebuild(self, app: Flask) -> None:
api = TriggerSubscriptionUpdateApi()
@@ -319,7 +365,12 @@ class TestTriggerSubscriptionCrud:
"controllers.console.workspace.trigger_providers.TriggerProviderService.rebuild_trigger_subscription"
),
):
assert method(api, "t1", "s1") == {"result": "success"}
assert method(
api,
TriggerSubscriptionBuilderUpdatePayload(credentials={}),
"t1",
"s1",
) == {"result": "success"}
class TestTriggerOAuthApis:
@@ -499,7 +550,12 @@ class TestTriggerOAuthClientManageApi:
return_value={"result": "success"},
),
):
assert method(api, "t1", "github") == {"result": "success"}
assert method(
api,
TriggerOAuthClientPayload(enabled=True),
"t1",
"github",
) == {"result": "success"}
def test_delete_client(self, app: Flask) -> None:
api = TriggerOAuthClientManageApi()
@@ -526,7 +582,7 @@ class TestTriggerOAuthClientManageApi:
),
):
with pytest.raises(BadRequest):
method(api, "t1", "github")
method(api, TriggerOAuthClientPayload(enabled=True), "t1", "github")
class TestTriggerSubscriptionVerifyApi:
@@ -541,7 +597,14 @@ class TestTriggerSubscriptionVerifyApi:
return_value={"verified": True},
),
):
assert method(api, "t1", mock_user(), "github", "s1") == {"verified": True}
assert method(
api,
TriggerSubscriptionBuilderVerifyPayload(credentials={}),
"t1",
mock_user(),
"github",
"s1",
) == {"verified": True}
@pytest.mark.parametrize("raised_exception", [ValueError("bad"), Exception("boom")])
def test_verify_errors(self, app: Flask, raised_exception: Exception) -> None:
@@ -556,4 +619,11 @@ class TestTriggerSubscriptionVerifyApi:
),
):
with pytest.raises(BadRequest):
method(api, "t1", mock_user(), "github", "s1")
method(
api,
TriggerSubscriptionBuilderVerifyPayload(credentials={}),
"t1",
mock_user(),
"github",
"s1",
)