mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
refactor(api): migrate dataset rag pipeline endpoints to BaseModel (#37958)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> Co-authored-by: Asuka Minato <i@asukaminato.eu.org>
This commit is contained in:
co-authored by
autofix-ci[bot]
Asuka Minato
parent
c3b1508712
commit
77ae583b44
+6
-2
@@ -607,7 +607,11 @@ class TestMiscApis:
|
||||
method = unwrap(api.get)
|
||||
|
||||
service = MagicMock()
|
||||
service.get_recommended_plugins.return_value = [{"id": "p1"}]
|
||||
recommended_plugins = {
|
||||
"installed_recommended_plugins": [{"id": "p1"}],
|
||||
"uninstalled_recommended_plugins": [{"id": "p2"}],
|
||||
}
|
||||
service.get_recommended_plugins.return_value = recommended_plugins
|
||||
user = make_account()
|
||||
tenant_id = "tenant-1"
|
||||
|
||||
@@ -619,7 +623,7 @@ class TestMiscApis:
|
||||
),
|
||||
):
|
||||
result = method(api, tenant_id, user)
|
||||
assert result == [{"id": "p1"}]
|
||||
assert result == recommended_plugins
|
||||
service.get_recommended_plugins.assert_called_once_with("all", user, tenant_id)
|
||||
|
||||
|
||||
|
||||
+224
-46
@@ -1,4 +1,5 @@
|
||||
import inspect
|
||||
from datetime import UTC, datetime
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
@@ -23,6 +24,76 @@ from graphon.model_runtime.errors.validate import CredentialsValidateFailedError
|
||||
from services.datasource_provider_service import DatasourceProviderService
|
||||
from services.plugin.oauth_service import OAuthProxyService
|
||||
|
||||
_PROVIDER_ID = "langgenius/notion_datasource/notion"
|
||||
|
||||
|
||||
def _i18n(text: str) -> dict[str, str]:
|
||||
return {"en_US": text, "zh_Hans": text, "pt_BR": text, "ja_JP": text}
|
||||
|
||||
|
||||
def _provider_config(name: str, type_: str, label: str, *, required: bool = True) -> dict:
|
||||
return {
|
||||
"type": type_,
|
||||
"name": name,
|
||||
"scope": None,
|
||||
"required": required,
|
||||
"default": None,
|
||||
"options": None,
|
||||
"multiple": False,
|
||||
"label": _i18n(label),
|
||||
"help": None,
|
||||
"url": None,
|
||||
"placeholder": None,
|
||||
}
|
||||
|
||||
|
||||
def _datasource_credential(credential_id: str = "cred-1", *, is_default: bool = True) -> dict:
|
||||
return {
|
||||
"credential": {
|
||||
"api_key": "******",
|
||||
"workspace": "engineering",
|
||||
"database_id": "db-123",
|
||||
},
|
||||
"type": "api-key",
|
||||
"name": "API Key",
|
||||
"avatar_url": "https://cdn.example.com/notion.png",
|
||||
"id": credential_id,
|
||||
"is_default": is_default,
|
||||
}
|
||||
|
||||
|
||||
def _datasource_auth() -> dict:
|
||||
return {
|
||||
"author": "Dify",
|
||||
"provider": "notion",
|
||||
"plugin_id": "langgenius/notion_datasource",
|
||||
"plugin_unique_identifier": "langgenius/notion_datasource:0.0.1",
|
||||
"icon": "icon.svg",
|
||||
"name": "notion",
|
||||
"label": _i18n("Notion"),
|
||||
"description": _i18n("Notion datasource"),
|
||||
"credential_schema": [
|
||||
_provider_config("api_key", "secret-input", "API key"),
|
||||
],
|
||||
"oauth_schema": {
|
||||
"client_schema": [
|
||||
_provider_config("client_id", "text-input", "Client ID"),
|
||||
],
|
||||
"credentials_schema": [
|
||||
_provider_config("access_token", "secret-input", "Access token"),
|
||||
],
|
||||
"oauth_custom_client_params": {"client_id": "masked-client", "client_secret": "********"},
|
||||
"is_oauth_custom_client_enabled": True,
|
||||
"is_system_oauth_params_exists": True,
|
||||
"redirect_uri": "https://api.example.com/oauth/callback",
|
||||
},
|
||||
"credentials_list": [_datasource_credential(), _datasource_credential("cred-2", is_default=False)],
|
||||
}
|
||||
|
||||
|
||||
def _success_response() -> dict[str, str]:
|
||||
return {"result": "success"}
|
||||
|
||||
|
||||
class TestDatasourcePluginOAuthAuthorizationUrl:
|
||||
def test_get_success(self, app: Flask):
|
||||
@@ -30,28 +101,50 @@ class TestDatasourcePluginOAuthAuthorizationUrl:
|
||||
method = inspect.unwrap(api.get)
|
||||
|
||||
user = MagicMock(id="user-1")
|
||||
oauth_client = {"client_id": "abc", "client_secret": "shh", "scopes": ["read", "write"]}
|
||||
auth_url_payload = {
|
||||
"authorization_url": "https://auth.example.com/oauth?client_id=abc&state=xyz",
|
||||
}
|
||||
|
||||
with (
|
||||
app.test_request_context("/?credential_id=cred-1"),
|
||||
patch.object(
|
||||
DatasourceProviderService,
|
||||
"get_oauth_client",
|
||||
return_value={"client_id": "abc"},
|
||||
),
|
||||
return_value=oauth_client,
|
||||
) as get_oauth_client,
|
||||
patch.object(
|
||||
OAuthProxyService,
|
||||
"create_proxy_context",
|
||||
return_value="ctx-1",
|
||||
),
|
||||
) as create_proxy_context,
|
||||
patch.object(
|
||||
OAuthHandler,
|
||||
"get_authorization_url",
|
||||
return_value={"url": "http://auth"},
|
||||
),
|
||||
return_value=auth_url_payload,
|
||||
) as get_authorization_url,
|
||||
):
|
||||
response = method(api, "tenant-1", user, "notion")
|
||||
response = method(api, "tenant-1", user, _PROVIDER_ID)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.get_json() == auth_url_payload
|
||||
assert "context_id=ctx-1" in response.headers.get("Set-Cookie")
|
||||
provider_id = get_oauth_client.call_args.kwargs["datasource_provider_id"]
|
||||
assert str(provider_id) == _PROVIDER_ID
|
||||
get_oauth_client.assert_called_once()
|
||||
create_proxy_context.assert_called_once_with(
|
||||
user_id="user-1",
|
||||
tenant_id="tenant-1",
|
||||
plugin_id="langgenius/notion_datasource",
|
||||
provider="notion",
|
||||
credential_id="cred-1",
|
||||
)
|
||||
get_authorization_url.assert_called_once()
|
||||
assert get_authorization_url.call_args.kwargs["tenant_id"] == "tenant-1"
|
||||
assert get_authorization_url.call_args.kwargs["user_id"] == "user-1"
|
||||
assert get_authorization_url.call_args.kwargs["plugin_id"] == "langgenius/notion_datasource"
|
||||
assert get_authorization_url.call_args.kwargs["provider"] == "notion"
|
||||
assert get_authorization_url.call_args.kwargs["system_credentials"] == oauth_client
|
||||
|
||||
def test_get_no_oauth_config(self, app: Flask):
|
||||
api = DatasourcePluginOAuthAuthorizationUrl()
|
||||
@@ -90,10 +183,10 @@ class TestDatasourcePluginOAuthAuthorizationUrl:
|
||||
patch.object(
|
||||
OAuthHandler,
|
||||
"get_authorization_url",
|
||||
return_value={"url": "http://auth"},
|
||||
return_value={"authorization_url": "http://auth"},
|
||||
),
|
||||
):
|
||||
response = method(api, "tenant-1", user, "notion")
|
||||
response = method(api, "tenant-1", user, _PROVIDER_ID)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert "context_id" in response.headers.get("Set-Cookie")
|
||||
@@ -106,8 +199,9 @@ class TestDatasourceOAuthCallback:
|
||||
|
||||
oauth_response = MagicMock()
|
||||
oauth_response.credentials = {"token": "abc"}
|
||||
oauth_response.expires_at = None
|
||||
oauth_response.metadata = {"name": "test"}
|
||||
expires_at = datetime(2024, 1, 2, 3, 4, 5, tzinfo=UTC)
|
||||
oauth_response.expires_at = expires_at
|
||||
oauth_response.metadata = {"name": "Workspace Bot", "avatar_url": "https://avatar.example.com/bot.png"}
|
||||
|
||||
context = {
|
||||
"user_id": "user-1",
|
||||
@@ -125,7 +219,7 @@ class TestDatasourceOAuthCallback:
|
||||
patch.object(
|
||||
DatasourceProviderService,
|
||||
"get_oauth_client",
|
||||
return_value={"client_id": "abc"},
|
||||
return_value={"client_id": "abc", "client_secret": "secret"},
|
||||
),
|
||||
patch.object(
|
||||
OAuthHandler,
|
||||
@@ -136,11 +230,22 @@ class TestDatasourceOAuthCallback:
|
||||
DatasourceProviderService,
|
||||
"add_datasource_oauth_provider",
|
||||
return_value=None,
|
||||
),
|
||||
) as add_oauth_provider,
|
||||
):
|
||||
response = method(api, "notion")
|
||||
response = method(api, _PROVIDER_ID)
|
||||
|
||||
assert response.status_code == 302
|
||||
assert "/oauth-callback" in response.location
|
||||
add_oauth_provider.assert_called_once()
|
||||
assert add_oauth_provider.call_args.kwargs == {
|
||||
"tenant_id": "tenant-1",
|
||||
"provider_id": add_oauth_provider.call_args.kwargs["provider_id"],
|
||||
"avatar_url": "https://avatar.example.com/bot.png",
|
||||
"name": "Workspace Bot",
|
||||
"expire_at": expires_at,
|
||||
"credentials": {"token": "abc"},
|
||||
}
|
||||
assert str(add_oauth_provider.call_args.kwargs["provider_id"]) == _PROVIDER_ID
|
||||
|
||||
def test_callback_missing_context(self, app: Flask):
|
||||
api = DatasourceOAuthCallback()
|
||||
@@ -223,12 +328,16 @@ class TestDatasourceOAuthCallback:
|
||||
DatasourceProviderService,
|
||||
"reauthorize_datasource_oauth_provider",
|
||||
return_value=None,
|
||||
),
|
||||
) as reauthorize_provider,
|
||||
):
|
||||
response = method(api, "notion")
|
||||
response = method(api, _PROVIDER_ID)
|
||||
|
||||
assert response.status_code == 302
|
||||
assert "/oauth-callback" in response.location
|
||||
reauthorize_provider.assert_called_once()
|
||||
assert str(reauthorize_provider.call_args.kwargs["provider_id"]) == _PROVIDER_ID
|
||||
assert reauthorize_provider.call_args.kwargs["credential_id"] == "cred-1"
|
||||
assert reauthorize_provider.call_args.kwargs["credentials"] == {"token": "abc"}
|
||||
|
||||
def test_callback_context_id_from_cookie(self, app: Flask):
|
||||
api = DatasourceOAuthCallback()
|
||||
@@ -278,7 +387,14 @@ class TestDatasourceAuth:
|
||||
api = DatasourceAuth()
|
||||
method = inspect.unwrap(api.post)
|
||||
|
||||
payload = {"credentials": {"key": "val"}}
|
||||
payload = {
|
||||
"name": "Engineering Notion",
|
||||
"credentials": {
|
||||
"api_key": "secret-token",
|
||||
"workspace": "engineering",
|
||||
"database_id": "db-123",
|
||||
},
|
||||
}
|
||||
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
@@ -287,11 +403,17 @@ class TestDatasourceAuth:
|
||||
DatasourceProviderService,
|
||||
"add_datasource_api_key_provider",
|
||||
return_value=None,
|
||||
),
|
||||
) as add_api_key_provider,
|
||||
):
|
||||
response, status = method(api, "tenant-1", "notion")
|
||||
response, status = method(api, "tenant-1", _PROVIDER_ID)
|
||||
|
||||
assert response == _success_response()
|
||||
assert status == 200
|
||||
add_api_key_provider.assert_called_once()
|
||||
assert add_api_key_provider.call_args.kwargs["tenant_id"] == "tenant-1"
|
||||
assert str(add_api_key_provider.call_args.kwargs["provider_id"]) == _PROVIDER_ID
|
||||
assert add_api_key_provider.call_args.kwargs["credentials"] == payload["credentials"]
|
||||
assert add_api_key_provider.call_args.kwargs["name"] == "Engineering Notion"
|
||||
|
||||
def test_post_invalid_credentials(self, app: Flask):
|
||||
api = DatasourceAuth()
|
||||
@@ -321,19 +443,19 @@ class TestDatasourceAuth:
|
||||
patch.object(
|
||||
DatasourceProviderService,
|
||||
"list_datasource_credentials",
|
||||
return_value=[{"id": "1"}],
|
||||
return_value=[_datasource_credential()],
|
||||
),
|
||||
):
|
||||
response, status = method(api, "tenant-1", user, "notion")
|
||||
response, status = method(api, "tenant-1", user, _PROVIDER_ID)
|
||||
|
||||
assert status == 200
|
||||
assert response["result"]
|
||||
assert response == {"result": [_datasource_credential()]}
|
||||
|
||||
def test_post_missing_credentials(self, app: Flask):
|
||||
api = DatasourceAuth()
|
||||
method = inspect.unwrap(api.post)
|
||||
|
||||
payload = {}
|
||||
payload: dict[str, object] = {}
|
||||
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
@@ -375,17 +497,24 @@ class TestDatasourceAuthDeleteApi:
|
||||
DatasourceProviderService,
|
||||
"remove_datasource_credentials",
|
||||
return_value=None,
|
||||
),
|
||||
) as remove_datasource_credentials,
|
||||
):
|
||||
response, status = method(api, "tenant-1", "notion")
|
||||
response, status = method(api, "tenant-1", _PROVIDER_ID)
|
||||
|
||||
assert response == _success_response()
|
||||
assert status == 200
|
||||
remove_datasource_credentials.assert_called_once_with(
|
||||
tenant_id="tenant-1",
|
||||
auth_id="cred-1",
|
||||
provider="notion",
|
||||
plugin_id="langgenius/notion_datasource",
|
||||
)
|
||||
|
||||
def test_delete_missing_credential_id(self, app: Flask):
|
||||
api = DatasourceAuthDeleteApi()
|
||||
method = inspect.unwrap(api.post)
|
||||
|
||||
payload = {}
|
||||
payload: dict[str, object] = {}
|
||||
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
@@ -400,7 +529,11 @@ class TestDatasourceAuthUpdateApi:
|
||||
api = DatasourceAuthUpdateApi()
|
||||
method = inspect.unwrap(api.post)
|
||||
|
||||
payload = {"credential_id": "id", "credentials": {"k": "v"}}
|
||||
payload = {
|
||||
"credential_id": "cred-1",
|
||||
"name": "Updated Notion",
|
||||
"credentials": {"api_key": "new-secret", "database_id": "db-456"},
|
||||
}
|
||||
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
@@ -409,11 +542,20 @@ class TestDatasourceAuthUpdateApi:
|
||||
DatasourceProviderService,
|
||||
"update_datasource_credentials",
|
||||
return_value=None,
|
||||
),
|
||||
) as update_datasource_credentials,
|
||||
):
|
||||
response, status = method(api, "tenant-1", "notion")
|
||||
response, status = method(api, "tenant-1", _PROVIDER_ID)
|
||||
|
||||
assert response == _success_response()
|
||||
assert status == 201
|
||||
update_datasource_credentials.assert_called_once_with(
|
||||
tenant_id="tenant-1",
|
||||
auth_id="cred-1",
|
||||
provider="notion",
|
||||
plugin_id="langgenius/notion_datasource",
|
||||
credentials=payload["credentials"],
|
||||
name="Updated Notion",
|
||||
)
|
||||
|
||||
def test_update_with_credentials_none(self, app: Flask):
|
||||
api = DatasourceAuthUpdateApi()
|
||||
@@ -432,7 +574,9 @@ class TestDatasourceAuthUpdateApi:
|
||||
):
|
||||
response, status = method(api, "tenant-1", "notion")
|
||||
|
||||
assert response == _success_response()
|
||||
update_mock.assert_called_once()
|
||||
assert update_mock.call_args.kwargs["credentials"] == {}
|
||||
assert status == 201
|
||||
|
||||
def test_update_name_only(self, app: Flask):
|
||||
@@ -450,8 +594,9 @@ class TestDatasourceAuthUpdateApi:
|
||||
return_value=None,
|
||||
),
|
||||
):
|
||||
_, status = method(api, "tenant-1", "notion")
|
||||
response, status = method(api, "tenant-1", "notion")
|
||||
|
||||
assert response == _success_response()
|
||||
assert status == 201
|
||||
|
||||
def test_update_with_empty_credentials_dict(self, app: Flask):
|
||||
@@ -469,8 +614,9 @@ class TestDatasourceAuthUpdateApi:
|
||||
return_value=None,
|
||||
) as update_mock,
|
||||
):
|
||||
_, status = method(api, "tenant-1", "notion")
|
||||
response, status = method(api, "tenant-1", "notion")
|
||||
|
||||
assert response == _success_response()
|
||||
update_mock.assert_called_once()
|
||||
assert status == 201
|
||||
|
||||
@@ -485,12 +631,13 @@ class TestDatasourceAuthListApi:
|
||||
patch.object(
|
||||
DatasourceProviderService,
|
||||
"get_all_datasource_credentials",
|
||||
return_value=[{"id": "1"}],
|
||||
return_value=[_datasource_auth()],
|
||||
),
|
||||
):
|
||||
response, status = method(api, "tenant-1")
|
||||
|
||||
assert status == 200
|
||||
assert response == {"result": [_datasource_auth()]}
|
||||
|
||||
def test_auth_list_empty(self, app: Flask):
|
||||
api = DatasourceAuthListApi()
|
||||
@@ -537,7 +684,7 @@ class TestDatasourceHardCodeAuthListApi:
|
||||
patch.object(
|
||||
DatasourceProviderService,
|
||||
"get_hard_code_datasource_credentials",
|
||||
return_value=[{"id": "1"}],
|
||||
return_value=[_datasource_auth()],
|
||||
),
|
||||
):
|
||||
response, status = method(api, "tenant-1")
|
||||
@@ -550,7 +697,14 @@ class TestDatasourceAuthOauthCustomClient:
|
||||
api = DatasourceAuthOauthCustomClient()
|
||||
method = inspect.unwrap(api.post)
|
||||
|
||||
payload = {"client_params": {}, "enable_oauth_custom_client": True}
|
||||
payload = {
|
||||
"client_params": {
|
||||
"client_id": "custom-client",
|
||||
"client_secret": "custom-secret",
|
||||
"authorize_url": "https://auth.example.com/authorize",
|
||||
},
|
||||
"enable_oauth_custom_client": True,
|
||||
}
|
||||
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
@@ -559,11 +713,17 @@ class TestDatasourceAuthOauthCustomClient:
|
||||
DatasourceProviderService,
|
||||
"setup_oauth_custom_client_params",
|
||||
return_value=None,
|
||||
),
|
||||
) as setup_custom_client,
|
||||
):
|
||||
response, status = method(api, "tenant-1", "notion")
|
||||
response, status = method(api, "tenant-1", _PROVIDER_ID)
|
||||
|
||||
assert response == _success_response()
|
||||
assert status == 200
|
||||
setup_custom_client.assert_called_once()
|
||||
assert setup_custom_client.call_args.kwargs["tenant_id"] == "tenant-1"
|
||||
assert str(setup_custom_client.call_args.kwargs["datasource_provider_id"]) == _PROVIDER_ID
|
||||
assert setup_custom_client.call_args.kwargs["client_params"] == payload["client_params"]
|
||||
assert setup_custom_client.call_args.kwargs["enabled"] is True
|
||||
|
||||
def test_delete_success(self, app: Flask):
|
||||
api = DatasourceAuthOauthCustomClient()
|
||||
@@ -575,17 +735,20 @@ class TestDatasourceAuthOauthCustomClient:
|
||||
DatasourceProviderService,
|
||||
"remove_oauth_custom_client_params",
|
||||
return_value=None,
|
||||
),
|
||||
) as remove_custom_client,
|
||||
):
|
||||
response, status = method(api, "tenant-1", "notion")
|
||||
response, status = method(api, "tenant-1", _PROVIDER_ID)
|
||||
|
||||
assert response == _success_response()
|
||||
assert status == 200
|
||||
remove_custom_client.assert_called_once()
|
||||
assert str(remove_custom_client.call_args.kwargs["datasource_provider_id"]) == _PROVIDER_ID
|
||||
|
||||
def test_post_empty_payload(self, app: Flask):
|
||||
api = DatasourceAuthOauthCustomClient()
|
||||
method = inspect.unwrap(api.post)
|
||||
|
||||
payload = {}
|
||||
payload: dict[str, object] = {}
|
||||
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
@@ -596,8 +759,9 @@ class TestDatasourceAuthOauthCustomClient:
|
||||
return_value=None,
|
||||
),
|
||||
):
|
||||
_, status = method(api, "tenant-1", "notion")
|
||||
response, status = method(api, "tenant-1", "notion")
|
||||
|
||||
assert response == _success_response()
|
||||
assert status == 200
|
||||
|
||||
def test_post_disabled_flag(self, app: Flask):
|
||||
@@ -618,9 +782,12 @@ class TestDatasourceAuthOauthCustomClient:
|
||||
return_value=None,
|
||||
) as setup_mock,
|
||||
):
|
||||
_, status = method(api, "tenant-1", "notion")
|
||||
response, status = method(api, "tenant-1", "notion")
|
||||
|
||||
assert response == _success_response()
|
||||
setup_mock.assert_called_once()
|
||||
assert setup_mock.call_args.kwargs["client_params"] == {"a": 1}
|
||||
assert setup_mock.call_args.kwargs["enabled"] is False
|
||||
assert status == 200
|
||||
|
||||
|
||||
@@ -638,17 +805,22 @@ class TestDatasourceAuthDefaultApi:
|
||||
DatasourceProviderService,
|
||||
"set_default_datasource_provider",
|
||||
return_value=None,
|
||||
),
|
||||
) as set_default_datasource_provider,
|
||||
):
|
||||
response, status = method(api, "tenant-1", "notion")
|
||||
response, status = method(api, "tenant-1", _PROVIDER_ID)
|
||||
|
||||
assert response == _success_response()
|
||||
assert status == 200
|
||||
set_default_datasource_provider.assert_called_once()
|
||||
assert set_default_datasource_provider.call_args.kwargs["tenant_id"] == "tenant-1"
|
||||
assert str(set_default_datasource_provider.call_args.kwargs["datasource_provider_id"]) == _PROVIDER_ID
|
||||
assert set_default_datasource_provider.call_args.kwargs["credential_id"] == "cred-1"
|
||||
|
||||
def test_default_missing_id(self, app: Flask):
|
||||
api = DatasourceAuthDefaultApi()
|
||||
method = inspect.unwrap(api.post)
|
||||
|
||||
payload = {}
|
||||
payload: dict[str, object] = {}
|
||||
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
@@ -663,7 +835,7 @@ class TestDatasourceUpdateProviderNameApi:
|
||||
api = DatasourceUpdateProviderNameApi()
|
||||
method = inspect.unwrap(api.post)
|
||||
|
||||
payload = {"credential_id": "id", "name": "New Name"}
|
||||
payload = {"credential_id": "cred-1", "name": "New Name"}
|
||||
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
@@ -672,11 +844,17 @@ class TestDatasourceUpdateProviderNameApi:
|
||||
DatasourceProviderService,
|
||||
"update_datasource_provider_name",
|
||||
return_value=None,
|
||||
),
|
||||
) as update_datasource_provider_name,
|
||||
):
|
||||
response, status = method(api, "tenant-1", "notion")
|
||||
response, status = method(api, "tenant-1", _PROVIDER_ID)
|
||||
|
||||
assert response == _success_response()
|
||||
assert status == 200
|
||||
update_datasource_provider_name.assert_called_once()
|
||||
assert update_datasource_provider_name.call_args.kwargs["tenant_id"] == "tenant-1"
|
||||
assert str(update_datasource_provider_name.call_args.kwargs["datasource_provider_id"]) == _PROVIDER_ID
|
||||
assert update_datasource_provider_name.call_args.kwargs["name"] == "New Name"
|
||||
assert update_datasource_provider_name.call_args.kwargs["credential_id"] == "cred-1"
|
||||
|
||||
def test_update_name_too_long(self, app: Flask):
|
||||
api = DatasourceUpdateProviderNameApi()
|
||||
|
||||
+62
@@ -158,3 +158,65 @@ def test_rag_pipeline_workflow_patch_serializes_response_model(app: Flask, monke
|
||||
assert response["id"] == "workflow-1"
|
||||
assert response["marked_name"] == "Updated release"
|
||||
assert response["hash"] == "hash-1"
|
||||
|
||||
|
||||
def test_default_rag_pipeline_block_configs_serializes_root_response(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
block_configs = [{"type": "start", "config": {"title": "Start"}}]
|
||||
monkeypatch.setattr(
|
||||
module,
|
||||
"RagPipelineService",
|
||||
lambda: SimpleNamespace(get_default_block_configs=lambda: block_configs),
|
||||
)
|
||||
|
||||
api = module.DefaultRagPipelineBlockConfigsApi()
|
||||
handler = unwrap_all(api.get)
|
||||
|
||||
response = handler(api, _pipeline())
|
||||
|
||||
assert response == block_configs
|
||||
|
||||
|
||||
def test_draft_rag_pipeline_second_step_parameters_serializes_variables(app, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
variables = [
|
||||
{
|
||||
"belong_to_node_id": "shared",
|
||||
"type": "number",
|
||||
"label": "Chunk size",
|
||||
"variable": "chunk_size",
|
||||
"default_value": 1024,
|
||||
"required": True,
|
||||
}
|
||||
]
|
||||
monkeypatch.setattr(
|
||||
module,
|
||||
"RagPipelineService",
|
||||
lambda: SimpleNamespace(get_second_step_parameters=lambda **_kwargs: variables),
|
||||
)
|
||||
|
||||
api = module.DraftRagPipelineSecondStepApi()
|
||||
handler = unwrap_all(api.get)
|
||||
|
||||
with app.test_request_context("/?node_id=node-1"):
|
||||
response = handler(api, _pipeline())
|
||||
|
||||
assert response["variables"] == variables
|
||||
|
||||
|
||||
def test_rag_pipeline_recommended_plugins_serializes_known_envelope(app, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
recommended_plugins = {
|
||||
"installed_recommended_plugins": [{"name": "Dify Extractor", "meta": {"version": "1.0.0"}}],
|
||||
"uninstalled_recommended_plugins": [{"plugin_id": "langgenius/notion_datasource"}],
|
||||
}
|
||||
monkeypatch.setattr(
|
||||
module,
|
||||
"RagPipelineService",
|
||||
lambda: SimpleNamespace(get_recommended_plugins=lambda *_args: recommended_plugins),
|
||||
)
|
||||
|
||||
api = module.RagPipelineRecommendedPluginApi()
|
||||
handler = unwrap_all(api.get)
|
||||
|
||||
with app.test_request_context("/?type=tool"):
|
||||
response = handler(api, "tenant-1", _account())
|
||||
|
||||
assert response == recommended_plugins
|
||||
|
||||
+41
-6
@@ -325,10 +325,12 @@ class TestPipelineRunApiEntity:
|
||||
def test_entity_missing_required_field(self):
|
||||
"""Test entity raises on missing required field."""
|
||||
with pytest.raises(ValueError):
|
||||
PipelineRunApiEntity(
|
||||
inputs={},
|
||||
datasource_type="online_document",
|
||||
# missing datasource_info_list, start_node_id, etc.
|
||||
PipelineRunApiEntity.model_validate(
|
||||
{
|
||||
"inputs": {},
|
||||
"datasource_type": "online_document",
|
||||
# missing datasource_info_list, start_node_id, etc.
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@@ -382,8 +384,19 @@ class TestDatasourcePluginsApiGet:
|
||||
mock_dataset = Mock()
|
||||
mock_db.session.scalar.return_value = mock_dataset
|
||||
|
||||
datasource_plugins = [
|
||||
{
|
||||
"node_id": "node-datasource-1",
|
||||
"plugin_id": "plugin-a",
|
||||
"provider_name": "provider-a",
|
||||
"datasource_type": "online_document",
|
||||
"title": "Online Docs",
|
||||
"user_input_variables": [{"variable": "url", "label": "URL", "type": "text-input", "required": True}],
|
||||
"credentials": [{"id": "cred-1", "name": "Default credential", "type": "oauth2", "is_default": True}],
|
||||
}
|
||||
]
|
||||
mock_svc_instance = Mock()
|
||||
mock_svc_instance.get_datasource_plugins.return_value = [{"name": "plugin_a"}]
|
||||
mock_svc_instance.get_datasource_plugins.return_value = datasource_plugins
|
||||
mock_svc_cls.return_value = mock_svc_instance
|
||||
|
||||
with app.test_request_context("/datasets/test/pipeline/datasource-plugins?is_published=true"):
|
||||
@@ -391,11 +404,33 @@ class TestDatasourcePluginsApiGet:
|
||||
response, status = api.get(tenant_id=tenant_id, dataset_id=dataset_id)
|
||||
|
||||
assert status == 200
|
||||
assert response == [{"name": "plugin_a"}]
|
||||
assert response == datasource_plugins
|
||||
mock_svc_instance.get_datasource_plugins.assert_called_once_with(
|
||||
tenant_id=tenant_id, dataset_id=dataset_id, is_published=True
|
||||
)
|
||||
|
||||
@patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.db")
|
||||
@patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.RagPipelineService")
|
||||
def test_get_plugins_parses_false_is_published_query(self, mock_svc_cls, mock_db, app: Flask):
|
||||
"""Test false query string is parsed as boolean False."""
|
||||
tenant_id = str(uuid.uuid4())
|
||||
dataset_id = str(uuid.uuid4())
|
||||
|
||||
mock_db.session.scalar.return_value = Mock()
|
||||
mock_svc_instance = Mock()
|
||||
mock_svc_instance.get_datasource_plugins.return_value = []
|
||||
mock_svc_cls.return_value = mock_svc_instance
|
||||
|
||||
with app.test_request_context("/datasets/test/pipeline/datasource-plugins?is_published=false"):
|
||||
api = DatasourcePluginsApi()
|
||||
response, status = api.get(tenant_id=tenant_id, dataset_id=dataset_id)
|
||||
|
||||
assert status == 200
|
||||
assert response == []
|
||||
mock_svc_instance.get_datasource_plugins.assert_called_once_with(
|
||||
tenant_id=tenant_id, dataset_id=dataset_id, is_published=False
|
||||
)
|
||||
|
||||
@patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.db")
|
||||
def test_get_plugins_not_found(self, mock_db, app: Flask):
|
||||
"""Test NotFound when dataset check fails."""
|
||||
|
||||
+5
-20
@@ -2,9 +2,10 @@
|
||||
Unit tests for Service API knowledge pipeline file-upload serialization.
|
||||
"""
|
||||
|
||||
import importlib.util
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
|
||||
from controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow import PipelineUploadFileResponse
|
||||
from libs.helper import dump_response
|
||||
|
||||
|
||||
class FakeUploadFile:
|
||||
@@ -17,21 +18,7 @@ class FakeUploadFile:
|
||||
created_at: datetime | None
|
||||
|
||||
|
||||
def _load_serialize_upload_file():
|
||||
api_dir = Path(__file__).resolve().parents[5]
|
||||
serializers_path = api_dir / "controllers" / "service_api" / "dataset" / "rag_pipeline" / "serializers.py"
|
||||
|
||||
spec = importlib.util.spec_from_file_location("rag_pipeline_serializers", serializers_path)
|
||||
assert spec
|
||||
assert spec.loader
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module) # type: ignore[attr-defined]
|
||||
return module.serialize_upload_file
|
||||
|
||||
|
||||
def test_file_upload_created_at_is_isoformat_string():
|
||||
serialize_upload_file = _load_serialize_upload_file()
|
||||
|
||||
created_at = datetime(2026, 2, 8, 12, 0, 0, tzinfo=UTC)
|
||||
upload_file = FakeUploadFile()
|
||||
upload_file.id = "file-1"
|
||||
@@ -42,13 +29,11 @@ def test_file_upload_created_at_is_isoformat_string():
|
||||
upload_file.created_by = "account-1"
|
||||
upload_file.created_at = created_at
|
||||
|
||||
result = serialize_upload_file(upload_file)
|
||||
result = dump_response(PipelineUploadFileResponse, upload_file)
|
||||
assert result["created_at"] == created_at.isoformat()
|
||||
|
||||
|
||||
def test_file_upload_created_at_none_serializes_to_null():
|
||||
serialize_upload_file = _load_serialize_upload_file()
|
||||
|
||||
upload_file = FakeUploadFile()
|
||||
upload_file.id = "file-1"
|
||||
upload_file.name = "test.pdf"
|
||||
@@ -58,5 +43,5 @@ def test_file_upload_created_at_none_serializes_to_null():
|
||||
upload_file.created_by = "account-1"
|
||||
upload_file.created_at = None
|
||||
|
||||
result = serialize_upload_file(upload_file)
|
||||
result = dump_response(PipelineUploadFileResponse, upload_file)
|
||||
assert result["created_at"] is None
|
||||
|
||||
Reference in New Issue
Block a user