refactor: replace manual model_validate with @model_validate in app workflow controllers (#40235)

This commit is contained in:
Likalikali
2026-08-09 08:20:10 +00:00
committed by GitHub
parent 516fd12b8d
commit fff6f7cf2f
22 changed files with 354 additions and 218 deletions
@@ -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"