feat: 增加AI回复手动重试
This commit is contained in:
59
ai_knowledge_base_v2/apps/backend/tests/test_chat_retry.py
Normal file
59
ai_knowledge_base_v2/apps/backend/tests/test_chat_retry.py
Normal file
@@ -0,0 +1,59 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import Session
|
||||
from sqlalchemy.pool import StaticPool
|
||||
|
||||
from app.models import Base
|
||||
from app.models.chat import ChatMessage, ChatSession
|
||||
from app.services.chat_stream_service import _retryable_user_message
|
||||
|
||||
|
||||
def _db() -> Session:
|
||||
engine = create_engine(
|
||||
"sqlite:///:memory:",
|
||||
connect_args={"check_same_thread": False},
|
||||
poolclass=StaticPool,
|
||||
)
|
||||
Base.metadata.create_all(engine)
|
||||
return Session(engine)
|
||||
|
||||
|
||||
def test_retry_reuses_only_latest_unanswered_matching_user_message():
|
||||
with _db() as db:
|
||||
session = ChatSession(id=1, user_id=7, title="重试测试", message_count=0, is_deleted=0)
|
||||
user_message = ChatMessage(
|
||||
id=1,
|
||||
session_id=1,
|
||||
user_id=7,
|
||||
role="user",
|
||||
content="原来的问题",
|
||||
message_status="FINISHED",
|
||||
)
|
||||
db.add_all([session, user_message])
|
||||
db.commit()
|
||||
|
||||
retried = _retryable_user_message(
|
||||
db,
|
||||
session_id=1,
|
||||
user_id=7,
|
||||
question=" 原来的问题 ",
|
||||
)
|
||||
assert retried is not None
|
||||
assert retried.id == user_message.id
|
||||
|
||||
assert _retryable_user_message(db, session_id=1, user_id=7, question="另一个问题") is None
|
||||
|
||||
db.add(
|
||||
ChatMessage(
|
||||
id=2,
|
||||
session_id=1,
|
||||
user_id=7,
|
||||
role="assistant",
|
||||
content="已经回答",
|
||||
message_status="FINISHED",
|
||||
)
|
||||
)
|
||||
db.commit()
|
||||
|
||||
assert _retryable_user_message(db, session_id=1, user_id=7, question="原来的问题") is None
|
||||
@@ -127,7 +127,7 @@ def test_queued_chat_reports_position_and_completes(monkeypatch):
|
||||
monkeypatch.setattr(chat.ChatStreamService, "stream_answer_async", stream_answer)
|
||||
monkeypatch.setattr(chat.ReasoningPolicyService, "is_visible", lambda _db: False)
|
||||
|
||||
payload = type("Payload", (), {"sessionId": 1, "message": "问题"})()
|
||||
payload = type("Payload", (), {"sessionId": 1, "message": "问题", "retry": False})()
|
||||
async def collect_events():
|
||||
return [item async for item in chat._chat_stream(payload, object(), object())]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user