feat: 增加AI回复手动重试

This commit is contained in:
2026-08-03 17:55:02 +08:00
parent 8e6da99dcb
commit 1954f461af
10 changed files with 234 additions and 32 deletions

View File

@@ -268,7 +268,13 @@ async def _chat_stream(payload: ChatCompletionRequest, db: Session, current_user
waitingCount=queue_request.waiting_count,
reasoningVisible=reasoning_visible,
)
chunks = ChatStreamService.stream_answer_async(db, current_user, payload.sessionId, payload.message)
chunks = ChatStreamService.stream_answer_async(
db,
current_user,
payload.sessionId,
payload.message,
retry_failed_question=payload.retry,
)
async for segment in ReasoningPolicyService.iter_segments(chunks):
if segment.kind == "content":
yield _sse_event("content", content=segment.content)

View File

@@ -37,6 +37,7 @@ class UpdateSessionTitleRequest(BaseModel):
class ChatCompletionRequest(BaseModel):
sessionId: int
message: str = Field(min_length=1, max_length=4000)
retry: bool = False
class StopChatRequest(BaseModel):

View File

@@ -10,7 +10,7 @@ from fastapi import HTTPException, status
from sqlalchemy import select
from sqlalchemy.orm import Session
from app.models.chat import ChatMessage
from app.models.chat import ChatMessage, TopicSession
from app.models.knowledge import KnowledgeRetrievalLog
from app.models.user import User
from app.services.ai_request_log_service import AiRequestLogService
@@ -212,7 +212,14 @@ class ChatStreamService:
db.commit()
@staticmethod
async def stream_answer_async(db: Session, user: User, session_id: int, question: str) -> AsyncIterator[str]:
async def stream_answer_async(
db: Session,
user: User,
session_id: int,
question: str,
*,
retry_failed_question: bool = False,
) -> AsyncIterator[str]:
user = ChatService.prepare_daily_quota(db, user)
session = ChatService._get_user_session(db, user, session_id)
ChatService._ensure_quota(user)
@@ -225,25 +232,37 @@ class ChatStreamService:
now = _now()
normalized_question = question.strip()
topic = TopicSessionService.get_or_create_active(
db,
user=user,
session=session,
question=normalized_question,
deduct_quota=entitlement.deduct_quota,
user_message = (
_retryable_user_message(
db,
session_id=session.id,
user_id=user.id,
question=normalized_question,
)
if retry_failed_question
else None
)
user_message = ChatMessage(
session_id=session.id,
topic_session_id=topic.id,
user_id=user.id,
role="user",
content=normalized_question,
message_status="FINISHED",
created_at=now,
)
db.add(user_message)
TopicSessionService.attach_user_message(user_message, topic)
db.flush()
topic = db.get(TopicSession, user_message.topic_session_id) if user_message and user_message.topic_session_id else None
if user_message is None or topic is None:
topic = TopicSessionService.get_or_create_active(
db,
user=user,
session=session,
question=normalized_question,
deduct_quota=entitlement.deduct_quota,
)
user_message = ChatMessage(
session_id=session.id,
topic_session_id=topic.id,
user_id=user.id,
role="user",
content=normalized_question,
message_status="FINISHED",
created_at=now,
)
db.add(user_message)
TopicSessionService.attach_user_message(user_message, topic)
db.flush()
history = list(
db.scalars(
@@ -362,6 +381,27 @@ def _now() -> datetime:
return datetime.now(UTC).replace(tzinfo=None)
def _retryable_user_message(
db: Session,
*,
session_id: int,
user_id: int,
question: str,
) -> ChatMessage | None:
"""Reuse only the latest unanswered user turn to avoid duplicate retry messages."""
latest = db.scalar(
select(ChatMessage)
.where(ChatMessage.session_id == session_id, ChatMessage.user_id == user_id)
.order_by(ChatMessage.id.desc())
.limit(1)
)
if latest is None or latest.role != "user":
return None
if latest.content.strip() != question.strip():
return None
return latest
def _rough_token_count(text: str) -> int:
return max(1, len(text.strip()) // 2)

View 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

View File

@@ -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())]