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 import select from sqlalchemy.orm import Session from app.core.database import get_db from app.core.auth_context import ChatAccessScope, UserAuthContext from app.core.dependencies import get_current_user_context from app.core.responses import api_success from app.models.chat import TopicSession from app.models.growth import TopicSummary from app.schemas.chat import ( ChatCompletionRequest, ChatMessageRead, ChatSessionRead, CreateSessionRequest, 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( payload: CreateSessionRequest | None = None, db: Session = Depends(get_db), current: UserAuthContext = Depends(get_current_user_context), ) -> dict: session = ChatService.create_session( db, current.user, current.chat_scope, current_session_id=payload.currentSessionId if payload else None, ) return api_success(CreateSessionResponse(sessionId=session.id).model_dump()) @router.get("/session/list") def list_sessions(db: Session = Depends(get_db), current: UserAuthContext = Depends(get_current_user_context)) -> dict: sessions = ChatService.list_sessions(db, current.user, current.chat_scope) 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: UserAuthContext = Depends(get_current_user_context), ) -> dict: messages = ChatService.get_history(db, current.user, sessionId, current.chat_scope) 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: UserAuthContext = Depends(get_current_user_context), ) -> dict: session = ChatService.update_title(db, current.user, payload.sessionId, payload.title, current.chat_scope) 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: UserAuthContext = Depends(get_current_user_context), ) -> dict: ChatService.delete_session(db, current.user, session_id, current.chat_scope) return api_success() @router.post("/session/{session_id}/topic/finish") def finish_topic( session_id: int, db: Session = Depends(get_db), current: UserAuthContext = Depends(get_current_user_context), ) -> dict: session = ChatService._get_user_session(db, current.user, session_id, current.chat_scope) 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: UserAuthContext = Depends(get_current_user_context), ) -> dict: topic_session_id = db.scalar( select(TopicSession.chat_session_id) .join(TopicSummary, TopicSummary.topic_session_id == TopicSession.id) .where(TopicSummary.id == summary_id, TopicSummary.user_id == current.user.id) ) if topic_session_id is None: raise HTTPException(status_code=404, detail="主题沉淀任务不存在") ChatService._get_user_session(db, current.user, topic_session_id, current.chat_scope) 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: UserAuthContext = Depends(get_current_user_context), ) -> dict: session = ChatService._get_user_session(db, current.user, session_id, current.chat_scope) 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: UserAuthContext = Depends(get_current_user_context), ) -> dict: cards = HelpCardService.list_user_cards( db, user=current.user, scope=current.chat_scope, limit=limit, ) return api_success([help_card_dict(card) for card in cards]) @router.post("/help-card/{card_id}/copied") def mark_help_card_copied( card_id: int, db: Session = Depends(get_db), current: UserAuthContext = Depends(get_current_user_context), ) -> dict: return api_success(help_card_dict(HelpCardService.mark_copied(db, user=current.user, scope=current.chat_scope, card_id=card_id))) @router.delete("/help-card/{card_id}") def delete_help_card( card_id: int, db: Session = Depends(get_db), current: UserAuthContext = Depends(get_current_user_context), ) -> dict: HelpCardService.delete(db, user=current.user, scope=current.chat_scope, 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: UserAuthContext = Depends(get_current_user_context), ) -> dict: session = ChatService._get_user_session(db, current.user, session_id, current.chat_scope) 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: UserAuthContext = Depends(get_current_user_context), ) -> dict: drafts = ShareDraftService.list_user_drafts( db, user=current.user, scope=current.chat_scope, limit=limit, ) return api_success([share_draft_dict(draft) for draft in drafts]) @router.post("/share-draft/{draft_id}/copied") def mark_share_draft_copied( draft_id: int, db: Session = Depends(get_db), current: UserAuthContext = Depends(get_current_user_context), ) -> dict: return api_success(share_draft_dict(ShareDraftService.mark_copied(db, user=current.user, scope=current.chat_scope, draft_id=draft_id))) @router.delete("/share-draft/{draft_id}") def delete_share_draft( draft_id: int, db: Session = Depends(get_db), current: UserAuthContext = Depends(get_current_user_context), ) -> dict: ShareDraftService.delete(db, user=current.user, scope=current.chat_scope, draft_id=draft_id) return api_success() @router.post("/completions") def completions( payload: ChatCompletionRequest, db: Session = Depends(get_db), current: UserAuthContext = Depends(get_current_user_context), ) -> StreamingResponse: return StreamingResponse( _chat_stream(payload, db, current), 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: UserAuthContext = Depends(get_current_user_context), ) -> dict: ChatService.stop_generation(db, current.user, payload.sessionId, current.chat_scope) return api_success() async def _chat_stream(payload: ChatCompletionRequest, db: Session, current: UserAuthContext) -> AsyncIterator[str]: # Keep the internal stream helper compatible with direct service/test callers; # HTTP requests always pass a fully validated UserAuthContext. current_user = getattr(current, "user", current) chat_scope = getattr(current, "chat_scope", ChatAccessScope.direct()) 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, scope=chat_scope, ) 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"