feat: 增加AI回复手动重试
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user