test: use sqlite3 session in test_based_generate_task_pipeline (#38738)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
Asuka Minato 2026-07-21 16:37:20 +09:00 committed by GitHub
parent 1e672530c4
commit 2622b1ce0a
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -1,15 +1,52 @@
from decimal import Decimal
from types import SimpleNamespace
from unittest.mock import Mock
import pytest
from sqlalchemy.orm import Session
from clients.agent_backend.errors import AgentBackendRunFailedError
from core.app.apps.base_app_generate_response_converter import AppGenerateResponseConverter
from core.app.entities.app_invoke_entities import InvokeFrom
from core.app.entities.queue_entities import QueueErrorEvent
from core.app.task_pipeline.based_generate_task_pipeline import BasedGenerateTaskPipeline
from core.errors.error import QuotaExceededError
from graphon.model_runtime.errors.invoke import InvokeAuthorizationError, InvokeError, InvokeRateLimitError
from models.enums import MessageStatus
from models.enums import ConversationFromSource, MessageStatus
from models.model import AppMode, Message
def _persist_message(session: Session, *, message_id: str) -> Message:
message = Message(
id=message_id,
app_id="app-1",
model_provider=None,
model_id=None,
override_model_configs=None,
conversation_id="conversation-1",
inputs={},
query="query",
message={},
message_unit_price=Decimal(0),
answer="",
answer_unit_price=Decimal(0),
parent_message_id=None,
total_price=None,
currency="USD",
status=MessageStatus.NORMAL,
error=None,
message_metadata=None,
invoke_from=InvokeFrom.WEB_APP,
from_source=ConversationFromSource.CONSOLE,
from_end_user_id=None,
from_account_id="account-1",
workflow_run_id=None,
app_mode=AppMode.COMPLETION,
)
session.add(message)
session.commit()
session.expunge_all()
return message
class TestBasedGenerateTaskPipeline:
@ -58,26 +95,33 @@ class TestBasedGenerateTaskPipeline:
assert "Knowledge retrieval failed" in str(err)
assert "agent_run_id=run-1" in str(err)
def test_handle_error_updates_message_when_found(self, pipeline):
@pytest.mark.parametrize("sqlite_session", [(Message,)], indirect=True)
def test_handle_error_updates_message_when_found(self, pipeline, sqlite_session: Session):
event = QueueErrorEvent(error=ValueError("oops"))
message = SimpleNamespace(status=MessageStatus.NORMAL, error=None)
session = Mock()
session.scalar.return_value = message
_persist_message(sqlite_session, message_id="msg-1")
err = pipeline.handle_error(event=event, session=session, message_id="msg-1")
err = pipeline.handle_error(event=event, session=sqlite_session, message_id="msg-1")
assert err is event.error
assert message.status == MessageStatus.ERROR
assert message.error == "oops"
sqlite_session.flush()
sqlite_session.expire_all()
updated_message = sqlite_session.get(Message, "msg-1")
assert updated_message is not None
assert updated_message.status == MessageStatus.ERROR
assert updated_message.error == "oops"
def test_handle_error_returns_err_when_message_missing(self, pipeline):
@pytest.mark.parametrize("sqlite_session", [(Message,)], indirect=True)
def test_handle_error_returns_err_when_message_missing(self, pipeline, sqlite_session: Session):
event = QueueErrorEvent(error=ValueError("oops"))
session = Mock()
session.scalar.return_value = None
_persist_message(sqlite_session, message_id="other-message")
err = pipeline.handle_error(event=event, session=session, message_id="msg-1")
err = pipeline.handle_error(event=event, session=sqlite_session, message_id="msg-1")
assert err is event.error
untouched_message = sqlite_session.get(Message, "other-message")
assert untouched_message is not None
assert untouched_message.status == MessageStatus.NORMAL
assert untouched_message.error is None
def test_error_to_stream_response_and_ping(self, pipeline):
error_response = pipeline.error_to_stream_response(ValueError("boom"))