304 lines
10 KiB
Python
304 lines
10 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import json
|
|
import logging
|
|
from collections.abc import AsyncIterator
|
|
|
|
from fastapi import APIRouter, Depends, HTTPException, Query
|
|
from fastapi.responses import StreamingResponse
|
|
from sqlalchemy.orm import Session
|
|
|
|
from app.core.database import get_db
|
|
from app.core.dependencies import get_current_user
|
|
from app.core.responses import api_success
|
|
from app.models.user import User
|
|
from app.schemas.chat import (
|
|
ChatCompletionRequest,
|
|
ChatMessageRead,
|
|
ChatSessionRead,
|
|
CreateSessionResponse,
|
|
StopChatRequest,
|
|
UpdateSessionTitleRequest,
|
|
)
|
|
from app.services.chat_service import ChatService
|
|
from app.services.chat_queue_runtime import (
|
|
cancel_chat_waiter,
|
|
release_chat_slot,
|
|
request_chat_slot,
|
|
wait_for_chat_slot,
|
|
)
|
|
from app.services.chat_queue_service import load_chat_queue_config
|
|
from app.services.chat_stream_service import ChatStreamService
|
|
from app.services.growth_profile_service import GrowthProfileService
|
|
from app.services.help_card_service import HelpCardService, help_card_dict
|
|
from app.services.reasoning_policy_service import ReasoningPolicyService
|
|
from app.services.share_draft_service import ShareDraftService, share_draft_dict
|
|
|
|
router = APIRouter()
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
@router.post("/session")
|
|
def create_session(db: Session = Depends(get_db), current_user: User = Depends(get_current_user)) -> dict:
|
|
session = ChatService.create_session(db, current_user)
|
|
return api_success(CreateSessionResponse(sessionId=session.id).model_dump())
|
|
|
|
|
|
@router.get("/session/list")
|
|
def list_sessions(db: Session = Depends(get_db), current_user: User = Depends(get_current_user)) -> dict:
|
|
sessions = ChatService.list_sessions(db, current_user)
|
|
return api_success([ChatSessionRead.model_validate(session).model_dump() for session in sessions])
|
|
|
|
|
|
@router.get("/history")
|
|
def history(
|
|
sessionId: int = Query(...),
|
|
db: Session = Depends(get_db),
|
|
current_user: User = Depends(get_current_user),
|
|
) -> dict:
|
|
messages = ChatService.get_history(db, current_user, sessionId)
|
|
reasoning_visible = ReasoningPolicyService.is_visible(db)
|
|
result = []
|
|
for message in messages:
|
|
item = ChatMessageRead.model_validate(message).model_dump(mode="json")
|
|
if message.role == "assistant" and not reasoning_visible:
|
|
item["content"] = ReasoningPolicyService.strip_reasoning(item["content"])
|
|
result.append(item)
|
|
return api_success(result)
|
|
|
|
|
|
@router.put("/session/title")
|
|
def update_title(
|
|
payload: UpdateSessionTitleRequest,
|
|
db: Session = Depends(get_db),
|
|
current_user: User = Depends(get_current_user),
|
|
) -> dict:
|
|
session = ChatService.update_title(db, current_user, payload.sessionId, payload.title)
|
|
return api_success(ChatSessionRead.model_validate(session).model_dump())
|
|
|
|
|
|
@router.delete("/session/{session_id}")
|
|
def delete_session(
|
|
session_id: int,
|
|
db: Session = Depends(get_db),
|
|
current_user: User = Depends(get_current_user),
|
|
) -> dict:
|
|
ChatService.delete_session(db, current_user, session_id)
|
|
return api_success()
|
|
|
|
|
|
@router.post("/session/{session_id}/topic/finish")
|
|
def finish_topic(
|
|
session_id: int,
|
|
db: Session = Depends(get_db),
|
|
current_user: User = Depends(get_current_user),
|
|
) -> dict:
|
|
session = ChatService._get_user_session(db, current_user, session_id)
|
|
return api_success(GrowthProfileService.finish_active_topic(db, user=current_user, session=session))
|
|
|
|
|
|
@router.get("/topic/settlement/{summary_id}")
|
|
def topic_settlement(
|
|
summary_id: int,
|
|
db: Session = Depends(get_db),
|
|
current_user: User = Depends(get_current_user),
|
|
) -> dict:
|
|
return api_success(GrowthProfileService.topic_settlement_result(db, user=current_user, summary_id=summary_id))
|
|
|
|
|
|
@router.post("/session/{session_id}/help-card")
|
|
def generate_help_card(
|
|
session_id: int,
|
|
db: Session = Depends(get_db),
|
|
current_user: User = Depends(get_current_user),
|
|
) -> dict:
|
|
session = ChatService._get_user_session(db, current_user, session_id)
|
|
card = HelpCardService.generate_for_session(db, user=current_user, session=session)
|
|
return api_success(help_card_dict(card))
|
|
|
|
|
|
@router.get("/help-card/list")
|
|
def list_help_cards(
|
|
limit: int = Query(default=20, ge=1, le=50),
|
|
db: Session = Depends(get_db),
|
|
current_user: User = Depends(get_current_user),
|
|
) -> dict:
|
|
return api_success([help_card_dict(card) for card in HelpCardService.list_user_cards(db, user=current_user, limit=limit)])
|
|
|
|
|
|
@router.post("/help-card/{card_id}/copied")
|
|
def mark_help_card_copied(
|
|
card_id: int,
|
|
db: Session = Depends(get_db),
|
|
current_user: User = Depends(get_current_user),
|
|
) -> dict:
|
|
return api_success(help_card_dict(HelpCardService.mark_copied(db, user=current_user, card_id=card_id)))
|
|
|
|
|
|
@router.delete("/help-card/{card_id}")
|
|
def delete_help_card(
|
|
card_id: int,
|
|
db: Session = Depends(get_db),
|
|
current_user: User = Depends(get_current_user),
|
|
) -> dict:
|
|
HelpCardService.delete(db, user=current_user, card_id=card_id)
|
|
return api_success()
|
|
|
|
|
|
@router.post("/session/{session_id}/share-draft")
|
|
def generate_share_draft(
|
|
session_id: int,
|
|
db: Session = Depends(get_db),
|
|
current_user: User = Depends(get_current_user),
|
|
) -> dict:
|
|
session = ChatService._get_user_session(db, current_user, session_id)
|
|
draft = ShareDraftService.generate_for_session(db, user=current_user, session=session)
|
|
return api_success(share_draft_dict(draft))
|
|
|
|
|
|
@router.get("/share-draft/list")
|
|
def list_share_drafts(
|
|
limit: int = Query(default=20, ge=1, le=50),
|
|
db: Session = Depends(get_db),
|
|
current_user: User = Depends(get_current_user),
|
|
) -> dict:
|
|
return api_success([share_draft_dict(draft) for draft in ShareDraftService.list_user_drafts(db, user=current_user, limit=limit)])
|
|
|
|
|
|
@router.post("/share-draft/{draft_id}/copied")
|
|
def mark_share_draft_copied(
|
|
draft_id: int,
|
|
db: Session = Depends(get_db),
|
|
current_user: User = Depends(get_current_user),
|
|
) -> dict:
|
|
return api_success(share_draft_dict(ShareDraftService.mark_copied(db, user=current_user, draft_id=draft_id)))
|
|
|
|
|
|
@router.delete("/share-draft/{draft_id}")
|
|
def delete_share_draft(
|
|
draft_id: int,
|
|
db: Session = Depends(get_db),
|
|
current_user: User = Depends(get_current_user),
|
|
) -> dict:
|
|
ShareDraftService.delete(db, user=current_user, draft_id=draft_id)
|
|
return api_success()
|
|
|
|
|
|
@router.post("/completions")
|
|
def completions(
|
|
payload: ChatCompletionRequest,
|
|
db: Session = Depends(get_db),
|
|
current_user: User = Depends(get_current_user),
|
|
) -> StreamingResponse:
|
|
return StreamingResponse(
|
|
_chat_stream(payload, db, current_user),
|
|
media_type="text/event-stream",
|
|
headers={
|
|
"Cache-Control": "no-cache, no-transform",
|
|
"Connection": "keep-alive",
|
|
"X-Accel-Buffering": "no",
|
|
},
|
|
)
|
|
|
|
|
|
@router.post("/stop")
|
|
def stop(
|
|
payload: StopChatRequest,
|
|
db: Session = Depends(get_db),
|
|
current_user: User = Depends(get_current_user),
|
|
) -> dict:
|
|
ChatService.stop_generation(db, current_user, payload.sessionId)
|
|
return api_success()
|
|
|
|
|
|
async def _chat_stream(payload: ChatCompletionRequest, db: Session, current_user: User) -> AsyncIterator[str]:
|
|
config = load_chat_queue_config(db)
|
|
queue_lease = await request_chat_slot(config)
|
|
queue_request = queue_lease.request
|
|
acquired = queue_request.status == "acquired"
|
|
should_cancel_waiter = queue_request.status == "queued"
|
|
|
|
if queue_request.status == "rejected":
|
|
yield _sse_event(
|
|
"error",
|
|
message="当前请求过多,请稍后再试。",
|
|
activeCount=queue_request.active_count,
|
|
waitingCount=queue_request.waiting_count,
|
|
)
|
|
yield _sse_done()
|
|
return
|
|
|
|
if queue_request.status == "queued" and queue_request.token is not None:
|
|
yield _sse_event(
|
|
"queued",
|
|
message=f"当前请求较多,正在排队中,前方约 {max(queue_request.position - 1, 0)} 个请求。",
|
|
position=queue_request.position,
|
|
activeCount=queue_request.active_count,
|
|
waitingCount=queue_request.waiting_count,
|
|
)
|
|
queue_lease = await wait_for_chat_slot(queue_lease, config)
|
|
queue_request = queue_lease.request
|
|
acquired = queue_request.status == "acquired"
|
|
should_cancel_waiter = queue_request.status == "queued"
|
|
if queue_request.status == "rejected":
|
|
yield _sse_event(
|
|
"error",
|
|
message="当前请求过多,请稍后再试。",
|
|
activeCount=queue_request.active_count,
|
|
waitingCount=queue_request.waiting_count,
|
|
)
|
|
yield _sse_done()
|
|
return
|
|
if queue_request.status == "timeout":
|
|
yield _sse_event("error", message="排队等待超时,请稍后再试。")
|
|
yield _sse_done()
|
|
return
|
|
if queue_request.status == "canceled":
|
|
yield _sse_event("error", message="排队请求已取消,请稍后再试。")
|
|
yield _sse_done()
|
|
return
|
|
|
|
try:
|
|
reasoning_visible = ReasoningPolicyService.is_visible(db)
|
|
yield _sse_event(
|
|
"generating",
|
|
message="思考中",
|
|
activeCount=queue_request.active_count,
|
|
waitingCount=queue_request.waiting_count,
|
|
reasoningVisible=reasoning_visible,
|
|
)
|
|
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)
|
|
elif reasoning_visible:
|
|
yield _sse_event("reasoning", content=segment.content)
|
|
except HTTPException as exc:
|
|
yield _sse_event("error", message=str(exc.detail))
|
|
except asyncio.CancelledError:
|
|
raise
|
|
except Exception:
|
|
logger.exception("Chat stream failed")
|
|
yield _sse_event("error", message="AI 回复生成失败,请稍后再试。")
|
|
finally:
|
|
if acquired:
|
|
await release_chat_slot(queue_lease)
|
|
elif should_cancel_waiter:
|
|
await cancel_chat_waiter(queue_lease)
|
|
|
|
yield _sse_done()
|
|
|
|
def _sse_event(event_type: str, **payload: object) -> str:
|
|
return f"data: {json.dumps({'type': event_type, **payload}, ensure_ascii=False)}\n\n"
|
|
|
|
|
|
def _sse_done() -> str:
|
|
return "data: [DONE]\n\n"
|