From d91672c6d3fa4b1602426e6fa72db6b2f45a7e9d Mon Sep 17 00:00:00 2001 From: Kiveiru3 <168050332+Kiveiru3@users.noreply.github.com> Date: Mon, 7 Sep 2026 05:14:18 +0000 Subject: [PATCH] refactor(controllers): inject message request models (#41884) --- api/controllers/console/app/message.py | 75 +++++++++++++++---- .../console/agent/test_agent_controllers.py | 29 ++++++- .../console/app/test_message_api.py | 17 ++++- 3 files changed, 100 insertions(+), 21 deletions(-) diff --git a/api/controllers/console/app/message.py b/api/controllers/console/app/message.py index 5234ea9309e..c5bbba1f738 100644 --- a/api/controllers/console/app/message.py +++ b/api/controllers/console/app/message.py @@ -2,7 +2,6 @@ import logging from typing import Literal from uuid import UUID -from flask import request from flask_restx import Resource from pydantic import BaseModel, Field, field_validator from sqlalchemy import exists, func, select @@ -158,8 +157,14 @@ class ChatMessageListApi(Resource): @rbac_permission_required(RBACCheck(RBACPermission.APP_VIEW_LAYOUT, PlainApp())) @with_session(write=False) @get_app_model(mode=[AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT, AppMode.AGENT]) - def get(self, session: Session, current_user: Account, app_model: App): - return _list_chat_messages(session=session, app_model=app_model, current_user=current_user) + @model_validate(ChatMessagesQuery) + def get(self, req_data: ChatMessagesQuery, session: Session, current_user: Account, app_model: App): + return _list_chat_messages( + args=req_data, + session=session, + app_model=app_model, + current_user=current_user, + ) @console_ns.route("/agent//chat-messages") @@ -178,13 +183,26 @@ class AgentChatMessageListApi(Resource): @with_current_user @with_current_tenant_id @with_session(write=False) - def get(self, session: Session, current_tenant_id: str, current_user: Account, agent_id: UUID): + @model_validate(ChatMessagesQuery) + def get( + self, + req_data: ChatMessagesQuery, + session: Session, + current_tenant_id: str, + current_user: Account, + agent_id: UUID, + ): app_model = resolve_agent_runtime_app_model( session=session, tenant_id=current_tenant_id, agent_id=agent_id, ) - return _list_chat_messages(session=session, app_model=app_model, current_user=current_user) + return _list_chat_messages( + args=req_data, + session=session, + app_model=app_model, + current_user=current_user, + ) @console_ns.route("/apps//feedbacks") @@ -202,8 +220,14 @@ class MessageFeedbackApi(Resource): @with_current_user @with_session @get_app_model - def post(self, session: Session, current_user: Account, app_model: App): - return _update_message_feedback(session=session, current_user=current_user, app_model=app_model) + @model_validate(MessageFeedbackPayload) + def post(self, req_data: MessageFeedbackPayload, session: Session, current_user: Account, app_model: App): + return _update_message_feedback( + args=req_data, + session=session, + current_user=current_user, + app_model=app_model, + ) @console_ns.route("/agent//feedbacks") @@ -221,13 +245,26 @@ class AgentMessageFeedbackApi(Resource): @with_current_user @with_current_tenant_id @with_session - def post(self, session: Session, current_tenant_id: str, current_user: Account, agent_id: UUID): + @model_validate(MessageFeedbackPayload) + def post( + self, + req_data: MessageFeedbackPayload, + session: Session, + current_tenant_id: str, + current_user: Account, + agent_id: UUID, + ): app_model = resolve_agent_runtime_app_model( session=session, tenant_id=current_tenant_id, agent_id=agent_id, ) - return _update_message_feedback(session=session, current_user=current_user, app_model=app_model) + return _update_message_feedback( + args=req_data, + session=session, + current_user=current_user, + app_model=app_model, + ) @console_ns.route("/apps//annotations/count") @@ -389,9 +426,13 @@ class AgentMessageApi(Resource): return _get_message_detail(session=session, app_model=app_model, message_id=message_id) -def _list_chat_messages(*, session: Session, app_model: App, current_user: Account | None = None): - args = ChatMessagesQuery.model_validate(request.args.to_dict()) - +def _list_chat_messages( + *, + args: ChatMessagesQuery, + session: Session, + app_model: App, + current_user: Account | None = None, +): if AppMode.value_of(app_model.mode) == AppMode.AGENT and current_user is not None: try: conversation = ConversationService.get_conversation( @@ -468,9 +509,13 @@ def _list_chat_messages(*, session: Session, app_model: App, current_user: Accou ) -def _update_message_feedback(*, session: Session, current_user: Account, app_model: App): - args = MessageFeedbackPayload.model_validate(console_ns.payload) - +def _update_message_feedback( + *, + args: MessageFeedbackPayload, + session: Session, + current_user: Account, + app_model: App, +): message_id = args.message_id message = session.scalar(select(Message).where(Message.id == message_id, Message.app_id == app_model.id).limit(1)) diff --git a/api/tests/unit_tests/controllers/console/agent/test_agent_controllers.py b/api/tests/unit_tests/controllers/console/agent/test_agent_controllers.py index f39a8e8a3c9..17397b0076b 100644 --- a/api/tests/unit_tests/controllers/console/agent/test_agent_controllers.py +++ b/api/tests/unit_tests/controllers/console/agent/test_agent_controllers.py @@ -2141,14 +2141,24 @@ def test_agent_chat_message_routes_resolve_app_from_agent_id( monkeypatch.setattr(message_controller, "_get_message_suggested_questions", get_message_suggested_questions) monkeypatch.setattr(message_controller, "_get_message_detail", get_message_detail) assert unwrap(AgentChatMessageListApi.get)( - AgentChatMessageListApi(), unbound_session, "tenant-1", current_user, agent_id + AgentChatMessageListApi(), + message_controller.ChatMessagesQuery(conversation_id="00000000-0000-0000-0000-000000000010"), + unbound_session, + "tenant-1", + current_user, + agent_id, ) == {"data": []} list_call = cast(dict[str, object], captured["list"]) assert list_call["session"] is unbound_session assert list_call["app_model"] is app_model with app.test_request_context(json={"message_id": message_id, "rating": "like"}): assert unwrap(AgentMessageFeedbackApi.post)( - AgentMessageFeedbackApi(), unbound_session, "tenant-1", current_user, agent_id + AgentMessageFeedbackApi(), + message_controller.MessageFeedbackPayload(message_id=message_id, rating="like"), + unbound_session, + "tenant-1", + current_user, + agent_id, ) == {"result": "success"} feedback_call = cast(dict[str, object], captured["feedback"]) assert feedback_call["session"] is unbound_session @@ -2221,7 +2231,13 @@ def test_list_chat_messages_supports_first_id_pagination( f"/console/api/agent/agent-1/chat-messages?conversation_id={conversation_id}&first_id={first_message_id}&limit=1" ): result = message_controller._list_chat_messages( - session=sqlite_session, app_model=_app_detail_obj(id=app_id, mode=AppMode.CHAT) + args=message_controller.ChatMessagesQuery( + conversation_id=conversation_id, + first_id=first_message_id, + limit=1, + ), + session=sqlite_session, + app_model=_app_detail_obj(id=app_id, mode=AppMode.CHAT), ) assert result == {"data": [older_message_id], "limit": 1, "has_more": True} @@ -2263,7 +2279,10 @@ def test_list_agent_chat_messages_uses_current_user_conversation( monkeypatch.setattr(message_controller, "MessageInfiniteScrollPaginationResponse", FakeMessagePaginationResponse) with app.test_request_context(f"/console/api/agent/agent-1/chat-messages?conversation_id={conversation_id}"): result = message_controller._list_chat_messages( - session=sqlite_session, app_model=app_model, current_user=current_user + args=message_controller.ChatMessagesQuery(conversation_id=conversation_id), + session=sqlite_session, + app_model=app_model, + current_user=current_user, ) assert result == {"data": [message_id], "limit": 20, "has_more": False} assert captured.pop("session") is sqlite_session @@ -2282,6 +2301,7 @@ def test_list_agent_chat_messages_rejects_foreign_conversation( with app.test_request_context(f"/console/api/agent/agent-1/chat-messages?conversation_id={conversation_id}"): with pytest.raises(NotFound): message_controller._list_chat_messages( + args=message_controller.ChatMessagesQuery(conversation_id=conversation_id), session=unbound_session, app_model=_app_detail_obj(id="app-1", mode=AppMode.AGENT), current_user=_account(), @@ -2303,6 +2323,7 @@ def test_update_message_feedback_rejects_empty_rating_without_existing_feedback( with app.test_request_context(json={"message_id": message_id, "rating": None}): with pytest.raises(ValueError, match="rating cannot be None"): message_controller._update_message_feedback( + args=message_controller.MessageFeedbackPayload(message_id=message_id, rating=None), session=sqlite_session, current_user=_account(), app_model=_app_detail_obj(id=app_id), diff --git a/api/tests/unit_tests/controllers/console/app/test_message_api.py b/api/tests/unit_tests/controllers/console/app/test_message_api.py index ae6e994b6d5..bc797ef75e5 100644 --- a/api/tests/unit_tests/controllers/console/app/test_message_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_message_api.py @@ -81,10 +81,18 @@ def test_app_message_routes_pass_injected_session( monkeypatch.setattr(message_module, "_get_message_detail", get_message_detail) assert unwrap(message_module.ChatMessageListApi.get)( - message_module.ChatMessageListApi(), session, current_user, app_model + message_module.ChatMessageListApi(), + message_module.ChatMessagesQuery(conversation_id="550e8400-e29b-41d4-a716-446655440001"), + session, + current_user, + app_model, ) == {"data": []} assert unwrap(message_module.MessageFeedbackApi.post)( - message_module.MessageFeedbackApi(), session, current_user, app_model + message_module.MessageFeedbackApi(), + message_module.MessageFeedbackPayload(message_id=message_id, rating="like"), + session, + current_user, + app_model, ) == {"result": "success"} assert unwrap(message_module.MessageSuggestedQuestionApi.get)( message_module.MessageSuggestedQuestionApi(), session, current_user, app_model, message_id @@ -111,6 +119,11 @@ def test_update_message_feedback_commits_injected_session(app: Flask, sqlite_ses with app.test_request_context(json={"message_id": message_id, "rating": "like", "content": "helpful"}): result = message_module._update_message_feedback( + args=message_module.MessageFeedbackPayload( + message_id=message_id, + rating="like", + content="helpful", + ), session=session, current_user=SimpleNamespace(id="account-1"), app_model=SimpleNamespace(id="app-1"),