mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
refactor: replace manual model_validate with @model_validate in app workflow controllers (#40235)
This commit is contained in:
@@ -59,6 +59,7 @@ from controllers.console.app.workflow import AdvancedChatWorkflowRunPayload, Syn
|
||||
from controllers.console.app.workflow_app_log import WorkflowAppLogQuery
|
||||
from controllers.console.app.workflow_draft_variable import (
|
||||
EnvironmentVariableUpdatePayload,
|
||||
WorkflowDraftVariableListQuery,
|
||||
WorkflowDraftVariableUpdatePayload,
|
||||
)
|
||||
from controllers.console.app.workflow_statistic import WorkflowStatisticQuery
|
||||
@@ -421,7 +422,13 @@ class TestSiteEndpoints:
|
||||
site = self._add_site(db.session)
|
||||
|
||||
with database_app.test_request_context("/", json={"title": "My Site", "input_placeholder": "Ask me anything"}):
|
||||
result = method(api, db.session, _make_account(), app_model=_make_app())
|
||||
result = method(
|
||||
api,
|
||||
AppSiteUpdatePayload(title="My Site", input_placeholder="Ask me anything"),
|
||||
db.session,
|
||||
_make_account(),
|
||||
app_model=_make_app(),
|
||||
)
|
||||
|
||||
db.session.refresh(site)
|
||||
assert isinstance(result, dict)
|
||||
@@ -488,7 +495,7 @@ class TestWorkflowAppLogEndpoints:
|
||||
)
|
||||
|
||||
with database_app.test_request_context("/?page=1&limit=20"):
|
||||
result = method(api, app_model=_make_app("app-1"))
|
||||
result = method(api, WorkflowAppLogQuery(page=1, limit=20), app_model=_make_app("app-1"))
|
||||
|
||||
assert result == {"page": 1, "limit": 20, "total": 0, "has_more": False, "data": []}
|
||||
|
||||
@@ -518,7 +525,12 @@ class TestWorkflowDraftVariableEndpoints:
|
||||
monkeypatch.setattr(workflow_draft_variable_module, "WorkflowService", DummyWorkflowService)
|
||||
|
||||
with database_app.test_request_context("/?page=1&limit=20"):
|
||||
result = method(api, _make_account(), app_model=_make_app("app-1"))
|
||||
result = method(
|
||||
api,
|
||||
WorkflowDraftVariableListQuery(page=1, limit=20),
|
||||
_make_account(),
|
||||
app_model=_make_app("app-1"),
|
||||
)
|
||||
|
||||
assert result == {"items": [], "total": 0}
|
||||
|
||||
@@ -571,7 +583,16 @@ class TestWorkflowDraftVariableEndpoints:
|
||||
"deleted_environment_variable_ids": ["env-b"],
|
||||
},
|
||||
):
|
||||
result = method(api, _make_account(), app_model=_make_app())
|
||||
result = method(
|
||||
api,
|
||||
EnvironmentVariableUpdatePayload(
|
||||
environment_variables=[{"id": "env-a", "name": "a", "value_type": "string", "value": "new-a"}],
|
||||
patch=True,
|
||||
deleted_environment_variable_ids=["env-b"],
|
||||
),
|
||||
_make_account(),
|
||||
app_model=_make_app(),
|
||||
)
|
||||
|
||||
assert result == {"result": "success"}
|
||||
assert [(variable.id, variable.value) for variable in captured["environment_variables"]] == [("env-a", "new-a")]
|
||||
@@ -617,7 +638,7 @@ class TestWorkflowStatisticEndpoints:
|
||||
with database_app.test_request_context("/"):
|
||||
account = _make_account()
|
||||
account.timezone = "UTC"
|
||||
response = method(api, account, app_model=_make_app("app-1", tenant_id="t1"))
|
||||
response = method(api, WorkflowStatisticQuery(), account, app_model=_make_app("app-1", tenant_id="t1"))
|
||||
|
||||
assert response.get_json() == {"data": [{"date": "2024-01-01"}]}
|
||||
|
||||
@@ -647,7 +668,7 @@ class TestWorkflowStatisticEndpoints:
|
||||
with database_app.test_request_context("/"):
|
||||
account = _make_account()
|
||||
account.timezone = "UTC"
|
||||
response = method(api, account, app_model=_make_app("app-1", tenant_id="t1"))
|
||||
response = method(api, WorkflowStatisticQuery(), account, app_model=_make_app("app-1", tenant_id="t1"))
|
||||
|
||||
assert response.get_json() == {"data": [{"date": "2024-01-02"}]}
|
||||
|
||||
@@ -677,7 +698,7 @@ class TestWorkflowTriggerEndpoints:
|
||||
db.session.commit()
|
||||
|
||||
with database_app.test_request_context("/?node_id=node-1"):
|
||||
result = method(api, app_model=_make_app())
|
||||
result = method(api, Parser(node_id="node-1"), app_model=_make_app())
|
||||
|
||||
assert isinstance(result, dict)
|
||||
assert {"id", "webhook_id", "webhook_url", "webhook_debug_url", "node_id", "created_at"} <= set(result.keys())
|
||||
|
||||
@@ -159,7 +159,7 @@ class TestAppImportApi:
|
||||
)
|
||||
|
||||
with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}):
|
||||
response, status = method(api, _make_account())
|
||||
response, status = method(api, app_import_module.AppImportPayload(mode="yaml-content"), _make_account())
|
||||
|
||||
assert transaction_events.rollbacks == 1
|
||||
assert transaction_events.commits == 0
|
||||
@@ -185,7 +185,7 @@ class TestAppImportApi:
|
||||
)
|
||||
|
||||
with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}):
|
||||
response, status = method(api, _make_account())
|
||||
response, status = method(api, app_import_module.AppImportPayload(mode="yaml-content"), _make_account())
|
||||
|
||||
assert transaction_events.commits == 1
|
||||
assert transaction_events.rollbacks == 0
|
||||
@@ -213,7 +213,7 @@ class TestAppImportApi:
|
||||
monkeypatch.setattr(app_import_module.EnterpriseService.WebAppAuth, "update_app_access_mode", update_access)
|
||||
|
||||
with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}):
|
||||
response, status = method(api, _make_account())
|
||||
response, status = method(api, app_import_module.AppImportPayload(mode="yaml-content"), _make_account())
|
||||
|
||||
assert transaction_events.commits == 1
|
||||
assert transaction_events.rollbacks == 0
|
||||
@@ -251,7 +251,7 @@ class TestAppImportApi:
|
||||
)
|
||||
|
||||
with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}):
|
||||
response, status = method()
|
||||
response, status = method(app_import_module.AppImportPayload(mode="yaml-content"))
|
||||
|
||||
assert transaction_events.commits == 1
|
||||
_assert_app_persistence(sqlite_app_engine, app_id, persisted=True)
|
||||
@@ -291,7 +291,7 @@ class TestAppImportApi:
|
||||
method="POST",
|
||||
json={"mode": "yaml-content", "app_id": "existing-app"},
|
||||
):
|
||||
response, status = method()
|
||||
response, status = method(app_import_module.AppImportPayload(mode="yaml-content", app_id="existing-app"))
|
||||
|
||||
assert transaction_events.commits == 1
|
||||
_assert_app_persistence(sqlite_app_engine, app_id, persisted=True)
|
||||
|
||||
@@ -17,6 +17,8 @@ from controllers.console.app.audio import (
|
||||
ChatMessageAudioApi,
|
||||
ChatMessageTextApi,
|
||||
TextModesApi,
|
||||
TextToSpeechPayload,
|
||||
TextToSpeechVoiceQuery,
|
||||
)
|
||||
from controllers.console.app.error import (
|
||||
AppUnavailableError,
|
||||
@@ -290,7 +292,7 @@ def test_console_text_api_success(app: Flask, monkeypatch: pytest.MonkeyPatch) -
|
||||
method="POST",
|
||||
json={"text": "hello", "voice": "v"},
|
||||
):
|
||||
response = handler(api, app_model=app_model)
|
||||
response = handler(api, TextToSpeechPayload(text="hello"), app_model=app_model)
|
||||
|
||||
assert response == {"audio": "ok"}
|
||||
|
||||
@@ -315,7 +317,7 @@ def test_console_text_api_builds_message_ref(app: Flask, monkeypatch: pytest.Mon
|
||||
),
|
||||
patch("controllers.console.app.audio.current_user", SimpleNamespace(id="account-1")),
|
||||
):
|
||||
response = handler(api, app_model=app_model)
|
||||
response = handler(api, TextToSpeechPayload(text="hello", message_id="message-1"), app_model=app_model)
|
||||
|
||||
assert response == {"audio": "ok"}
|
||||
assert calls["message_ref"] == MessageRef(AppRef("tenant-1", "app-1"), "message-1", account_id="account-1")
|
||||
@@ -334,7 +336,7 @@ def test_console_text_api_error_mapping(app: Flask, monkeypatch: pytest.MonkeyPa
|
||||
json={"text": "hello"},
|
||||
):
|
||||
with pytest.raises(ProviderQuotaExceededError):
|
||||
handler(api, app_model=app_model)
|
||||
handler(api, TextToSpeechPayload(text="hello"), app_model=app_model)
|
||||
|
||||
|
||||
def test_console_text_modes_success(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
@@ -346,7 +348,7 @@ def test_console_text_modes_success(app: Flask, monkeypatch: pytest.MonkeyPatch)
|
||||
app_model = SimpleNamespace(tenant_id="t1")
|
||||
|
||||
with app.test_request_context("/console/api/apps/app/text-to-audio/voices?language=en", method="GET"):
|
||||
response = handler(api, app_model=app_model)
|
||||
response = handler(api, TextToSpeechVoiceQuery(language="en-US"), app_model=app_model)
|
||||
|
||||
assert response == expected_voices
|
||||
|
||||
@@ -364,7 +366,7 @@ def test_console_text_modes_language_error(app: Flask, monkeypatch: pytest.Monke
|
||||
|
||||
with app.test_request_context("/console/api/apps/app/text-to-audio/voices?language=en", method="GET"):
|
||||
with pytest.raises(AppUnavailableError):
|
||||
handler(api, app_model=app_model)
|
||||
handler(api, TextToSpeechVoiceQuery(language="en-US"), app_model=app_model)
|
||||
|
||||
|
||||
def test_audio_to_text_success(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
@@ -424,7 +426,7 @@ def test_text_to_audio_success(app: Flask, monkeypatch: pytest.MonkeyPatch) -> N
|
||||
method="POST",
|
||||
json={"text": "hello"},
|
||||
):
|
||||
response = method(api, app_model=app_model)
|
||||
response = method(api, TextToSpeechPayload(text="hello"), app_model=app_model)
|
||||
|
||||
assert response == {"audio": "ok"}
|
||||
|
||||
@@ -443,7 +445,7 @@ def test_text_to_audio_voices_success(app: Flask, monkeypatch: pytest.MonkeyPatc
|
||||
method="GET",
|
||||
query_string={"language": "en-US"},
|
||||
):
|
||||
response = method(api, app_model=app_model)
|
||||
response = method(api, TextToSpeechVoiceQuery(language="en-US"), app_model=app_model)
|
||||
|
||||
assert response == expected_voices
|
||||
|
||||
@@ -481,7 +483,7 @@ def test_text_to_audio_with_language_param(app: Flask, monkeypatch: pytest.Monke
|
||||
method="POST",
|
||||
json={"text": "hello", "language": "en-US"},
|
||||
):
|
||||
response = method(api, app_model=app_model)
|
||||
response = method(api, TextToSpeechPayload(text="hello"), app_model=app_model)
|
||||
assert response == {"audio": "test"}
|
||||
|
||||
|
||||
@@ -501,5 +503,5 @@ def test_text_to_audio_voices_with_language_filter(app: Flask, monkeypatch: pyte
|
||||
"/console/api/apps/app-1/text-to-audio/voices?language=en-US",
|
||||
method="GET",
|
||||
):
|
||||
response = method(api, app_model=app_model)
|
||||
response = method(api, TextToSpeechVoiceQuery(language="en-US"), app_model=app_model)
|
||||
assert isinstance(response, list)
|
||||
|
||||
@@ -51,7 +51,11 @@ def test_get_conversation_variables_returns_paginated_response(
|
||||
method="GET",
|
||||
query_string={"conversation_id": "conv-1"},
|
||||
):
|
||||
response = method(api, app_model=SimpleNamespace(id="app-1"))
|
||||
response = method(
|
||||
api,
|
||||
conversation_variables_module.ConversationVariablesQuery(conversation_id="conv-1"),
|
||||
app_model=SimpleNamespace(id="app-1"),
|
||||
)
|
||||
|
||||
assert response["page"] == 1
|
||||
assert response["limit"] == 100
|
||||
@@ -90,7 +94,11 @@ def test_get_conversation_variables_normalizes_value_type_and_value(
|
||||
method="GET",
|
||||
query_string={"conversation_id": "conv-1"},
|
||||
):
|
||||
response = method(api, app_model=SimpleNamespace(id="app-1"))
|
||||
response = method(
|
||||
api,
|
||||
conversation_variables_module.ConversationVariablesQuery(conversation_id="conv-1"),
|
||||
app_model=SimpleNamespace(id="app-1"),
|
||||
)
|
||||
|
||||
assert response["data"][0]["value_type"] == "number"
|
||||
assert response["data"][0]["value"] == "42"
|
||||
@@ -102,4 +110,4 @@ def test_get_conversation_variables_requires_conversation_id(app) -> None:
|
||||
|
||||
with app.test_request_context("/console/api/apps/app-1/conversation-variables", method="GET"):
|
||||
with pytest.raises(ValidationError):
|
||||
method(api, app_model=SimpleNamespace(id="app-1"))
|
||||
conversation_variables_module.ConversationVariablesQuery.model_validate({})
|
||||
|
||||
@@ -53,7 +53,12 @@ def test_daily_message_statistic_returns_rows(app: Flask, monkeypatch: pytest.Mo
|
||||
_install_db(monkeypatch, rows)
|
||||
|
||||
with app.test_request_context("/console/api/apps/app-1/statistics/daily-messages", method="GET"):
|
||||
response = method(api, SimpleNamespace(timezone="UTC"), app_model=SimpleNamespace(id="app-1"))
|
||||
response = method(
|
||||
api,
|
||||
SimpleNamespace(start=None, end=None),
|
||||
SimpleNamespace(timezone="UTC"),
|
||||
app_model=SimpleNamespace(id="app-1"),
|
||||
)
|
||||
|
||||
assert _json_payload(response) == {"data": [{"date": "2024-01-01", "message_count": 3}]}
|
||||
|
||||
@@ -67,7 +72,12 @@ def test_daily_conversation_statistic_returns_rows(app: Flask, monkeypatch: pyte
|
||||
_install_db(monkeypatch, rows)
|
||||
|
||||
with app.test_request_context("/console/api/apps/app-1/statistics/daily-conversations", method="GET"):
|
||||
response = method(api, SimpleNamespace(timezone="UTC"), app_model=SimpleNamespace(id="app-1"))
|
||||
response = method(
|
||||
api,
|
||||
SimpleNamespace(start=None, end=None),
|
||||
SimpleNamespace(timezone="UTC"),
|
||||
app_model=SimpleNamespace(id="app-1"),
|
||||
)
|
||||
|
||||
assert _json_payload(response) == {"data": [{"date": "2024-01-02", "conversation_count": 5}]}
|
||||
|
||||
@@ -81,7 +91,12 @@ def test_daily_token_cost_statistic_returns_rows(app: Flask, monkeypatch: pytest
|
||||
_install_db(monkeypatch, rows)
|
||||
|
||||
with app.test_request_context("/console/api/apps/app-1/statistics/token-costs", method="GET"):
|
||||
response = method(api, SimpleNamespace(timezone="UTC"), app_model=SimpleNamespace(id="app-1"))
|
||||
response = method(
|
||||
api,
|
||||
SimpleNamespace(start=None, end=None),
|
||||
SimpleNamespace(timezone="UTC"),
|
||||
app_model=SimpleNamespace(id="app-1"),
|
||||
)
|
||||
|
||||
data = _json_payload(response)
|
||||
assert len(data["data"]) == 1
|
||||
@@ -99,7 +114,12 @@ def test_daily_terminals_statistic_returns_rows(app: Flask, monkeypatch: pytest.
|
||||
_install_db(monkeypatch, rows)
|
||||
|
||||
with app.test_request_context("/console/api/apps/app-1/statistics/daily-end-users", method="GET"):
|
||||
response = method(api, SimpleNamespace(timezone="UTC"), app_model=SimpleNamespace(id="app-1"))
|
||||
response = method(
|
||||
api,
|
||||
SimpleNamespace(start=None, end=None),
|
||||
SimpleNamespace(timezone="UTC"),
|
||||
app_model=SimpleNamespace(id="app-1"),
|
||||
)
|
||||
|
||||
assert _json_payload(response) == {"data": [{"date": "2024-01-04", "terminal_count": 7}]}
|
||||
|
||||
@@ -126,7 +146,12 @@ def test_daily_message_statistic_with_invalid_time_range(app: Flask, monkeypatch
|
||||
|
||||
with app.test_request_context("/console/api/apps/app-1/statistics/daily-messages", method="GET"):
|
||||
with pytest.raises(BadRequest):
|
||||
method(api, SimpleNamespace(timezone="UTC"), app_model=SimpleNamespace(id="app-1"))
|
||||
method(
|
||||
api,
|
||||
SimpleNamespace(start=None, end=None),
|
||||
SimpleNamespace(timezone="UTC"),
|
||||
app_model=SimpleNamespace(id="app-1"),
|
||||
)
|
||||
|
||||
|
||||
def test_daily_message_statistic_multiple_rows(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
@@ -142,7 +167,12 @@ def test_daily_message_statistic_multiple_rows(app: Flask, monkeypatch: pytest.M
|
||||
_install_db(monkeypatch, rows)
|
||||
|
||||
with app.test_request_context("/console/api/apps/app-1/statistics/daily-messages", method="GET"):
|
||||
response = method(api, SimpleNamespace(timezone="UTC"), app_model=SimpleNamespace(id="app-1"))
|
||||
response = method(
|
||||
api,
|
||||
SimpleNamespace(start=None, end=None),
|
||||
SimpleNamespace(timezone="UTC"),
|
||||
app_model=SimpleNamespace(id="app-1"),
|
||||
)
|
||||
|
||||
data = _json_payload(response)
|
||||
assert len(data["data"]) == 3
|
||||
@@ -156,7 +186,12 @@ def test_daily_message_statistic_empty_result(app: Flask, monkeypatch: pytest.Mo
|
||||
_install_db(monkeypatch, [])
|
||||
|
||||
with app.test_request_context("/console/api/apps/app-1/statistics/daily-messages", method="GET"):
|
||||
response = method(api, SimpleNamespace(timezone="UTC"), app_model=SimpleNamespace(id="app-1"))
|
||||
response = method(
|
||||
api,
|
||||
SimpleNamespace(start=None, end=None),
|
||||
SimpleNamespace(timezone="UTC"),
|
||||
app_model=SimpleNamespace(id="app-1"),
|
||||
)
|
||||
|
||||
assert _json_payload(response) == {"data": []}
|
||||
|
||||
@@ -175,7 +210,12 @@ def test_daily_conversation_statistic_with_time_range(app: Flask, monkeypatch: p
|
||||
monkeypatch.setattr(statistic_module, "convert_datetime_to_date", lambda field: field)
|
||||
|
||||
with app.test_request_context("/console/api/apps/app-1/statistics/daily-conversations", method="GET"):
|
||||
response = method(api, SimpleNamespace(timezone="UTC"), app_model=SimpleNamespace(id="app-1"))
|
||||
response = method(
|
||||
api,
|
||||
SimpleNamespace(start=None, end=None),
|
||||
SimpleNamespace(timezone="UTC"),
|
||||
app_model=SimpleNamespace(id="app-1"),
|
||||
)
|
||||
|
||||
assert _json_payload(response) == {"data": [{"date": "2024-01-02", "conversation_count": 5}]}
|
||||
|
||||
@@ -192,7 +232,12 @@ def test_daily_token_cost_with_multiple_currencies(app: Flask, monkeypatch: pyte
|
||||
_install_db(monkeypatch, rows)
|
||||
|
||||
with app.test_request_context("/console/api/apps/app-1/statistics/token-costs", method="GET"):
|
||||
response = method(api, SimpleNamespace(timezone="UTC"), app_model=SimpleNamespace(id="app-1"))
|
||||
response = method(
|
||||
api,
|
||||
SimpleNamespace(start=None, end=None),
|
||||
SimpleNamespace(timezone="UTC"),
|
||||
app_model=SimpleNamespace(id="app-1"),
|
||||
)
|
||||
|
||||
data = _json_payload(response)
|
||||
assert len(data["data"]) == 2
|
||||
|
||||
@@ -98,7 +98,11 @@ def test_workflow_run_list_returns_frontend_history_contract(app: Flask, monkeyp
|
||||
handler = unwrap(api.get)
|
||||
|
||||
with app.test_request_context("/apps/app-1/workflow-runs?limit=10", method="GET"):
|
||||
payload = handler(api, app_model=SimpleNamespace(id="app-1", tenant_id="tenant-1"))
|
||||
payload = handler(
|
||||
api,
|
||||
workflow_run_module.WorkflowRunListQuery(limit=10),
|
||||
app_model=SimpleNamespace(id="app-1", tenant_id="tenant-1"),
|
||||
)
|
||||
|
||||
response = _serialize_200_response(api.get, payload)
|
||||
|
||||
@@ -139,7 +143,11 @@ def test_advanced_chat_workflow_run_list_keeps_message_fields(app: Flask, monkey
|
||||
handler = unwrap(api.get)
|
||||
|
||||
with app.test_request_context("/apps/app-1/advanced-chat/workflow-runs?limit=1", method="GET"):
|
||||
payload = handler(api, app_model=SimpleNamespace(id="app-1", tenant_id="tenant-1"))
|
||||
payload = handler(
|
||||
api,
|
||||
workflow_run_module.WorkflowRunListQuery(limit=1),
|
||||
app_model=SimpleNamespace(id="app-1", tenant_id="tenant-1"),
|
||||
)
|
||||
|
||||
response = _serialize_200_response(api.get, payload)
|
||||
|
||||
|
||||
@@ -142,7 +142,12 @@ def test_app_trigger_enable_uses_injected_tenant_id(app: Flask, database_session
|
||||
app.test_request_context("/", json=payload),
|
||||
patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload),
|
||||
):
|
||||
response = method(api, app_model.tenant_id, app_model)
|
||||
response = method(
|
||||
api,
|
||||
workflow_trigger_module.ParserEnable(trigger_id=trigger.id, enable_trigger=True),
|
||||
app_model.tenant_id,
|
||||
app_model,
|
||||
)
|
||||
|
||||
assert response["id"] == trigger.id
|
||||
assert response["status"] == "enabled"
|
||||
|
||||
Reference in New Issue
Block a user