Files
AIRegulation-DocAnalysis/backend/app/api/routes/rag.py
T
2026-07-02 22:03:39 +08:00

192 lines
8.0 KiB
Python

"""Define API routes for rag."""
from __future__ import annotations
import json
import os
import re
import tempfile
from typing import AsyncGenerator, Optional
from fastapi import APIRouter, Depends, File, UploadFile
from fastapi.responses import StreamingResponse
from loguru import logger
from app.api.dependencies.auth import get_current_user
from app.config.settings import settings
from app.domain.auth.models import UserClaims
from app.schemas.rag import RagChatRequest, QuickQuestionsResponse, QuickQuestion
from app.shared.async_utils import iter_in_thread
from app.shared.bootstrap import get_agent_conversation_service
# Maximum characters of document text injected as LLM context (≈ 6 000 tokens).
_MAX_CONTEXT_CHARS = 8_000
router = APIRouter(prefix="/rag", tags=["RAG问答"])
_DEFAULT_QUICK_QUESTIONS = [
{"id": "1", "question": "请总结最新入库法规对电池安全的核心要求", "category": "法规解读"},
{"id": "2", "question": "我上传的制度文档与新能源法规有哪些潜在冲突?", "category": "差距分析"},
{"id": "3", "question": "请给出法规依据,并按条款列出整改建议", "category": "整改建议"},
{"id": "4", "question": "请解释 UN-ECE 与 GB 标准在网络安全方面的差异", "category": "标准对比"},
{"id": "5", "question": "IATF 16949 对供应商质量管理有哪些强制要求?", "category": "法规解读"},
{"id": "6", "question": "ISO 45001 与 AQ 标准在职业健康安全方面的主要差异是什么?", "category": "标准对比"},
]
def _extract_text_from_bytes(content: bytes, filename: str) -> str:
"""Extract plain text from an uploaded file using the document parser.
Tries the configured parser first; falls back to raw UTF-8 decode for
plain-text formats (.txt, .md). Returns at most _MAX_CONTEXT_CHARS characters
so the text fits comfortably inside the LLM context window.
"""
suffix = os.path.splitext(filename or "doc.pdf")[1] or ".pdf"
# Fast path: plain-text files don't need a parser
if suffix.lower() in {".txt", ".md", ".csv"}:
try:
return content.decode("utf-8", errors="replace")[:_MAX_CONTEXT_CHARS]
except Exception:
pass
tmp_path = ""
try:
with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as tmp:
tmp.write(content)
tmp_path = tmp.name
from app.shared.bootstrap import get_document_command_service
svc = get_document_command_service()
parsed = svc.parser.parse(file_path=tmp_path, doc_id="ctx_extract", doc_name=filename)
if parsed.raw_text:
return parsed.raw_text[:_MAX_CONTEXT_CHARS]
# Fallback: join semantic blocks
return "\n".join(
b.get("text", "") for b in parsed.semantic_blocks if b.get("text")
)[:_MAX_CONTEXT_CHARS]
except Exception as exc:
logger.warning("Context text extraction failed for {}: {}", filename, exc)
return ""
finally:
if tmp_path:
try:
os.unlink(tmp_path)
except OSError:
pass
@router.post("/upload-context")
async def upload_context(
file: UploadFile = File(...),
current_user: UserClaims = Depends(get_current_user),
):
"""Extract text from an uploaded document and return it as conversation context.
The client stores the returned text and includes it in subsequent /rag/chat
requests via the context_text field — the LLM receives the document content
directly without requiring vector-store indexing.
"""
content = await file.read()
filename = file.filename or "document"
text = await __import__("asyncio").to_thread(_extract_text_from_bytes, content, filename)
if not text.strip():
from fastapi import HTTPException
raise HTTPException(status_code=422, detail="Could not extract text from the uploaded file.")
return {
"filename": filename,
"text": text,
"char_count": len(text),
"truncated": len(text) >= _MAX_CONTEXT_CHARS,
}
@router.post("/chat")
async def rag_chat(
request: RagChatRequest,
current_user: UserClaims = Depends(get_current_user),
):
"""Stream RAG Q&A using the real agent service.
When request.context_text is provided the document text is passed directly
to the answer generator as a dedicated document context section — RAG
retrieval still runs on the user's original question (not the document text)
so embedding quality is preserved for regulation chunk matching.
"""
session_id, event_stream = get_agent_conversation_service().stream_chat(
query=request.query,
session_id=request.session_id,
filters=request.filters,
top_k=request.top_k or settings.rag_top_k,
context_text=request.context_text,
context_filename=request.context_filename,
)
async def generate() -> AsyncGenerator[str, None]:
"""Translate agent SSE events to rag format."""
yield (
"event: message\n"
f"data: {json.dumps({'type': 'session', 'session_id': session_id}, ensure_ascii=False)}\n\n"
)
async for event in iter_in_thread(event_stream):
event_type = event.get("event", "")
data = event.get("data", "")
if event_type == "sources":
docs = [
{
"id": str(s.get("chunk_id") or s.get("doc_id") or idx + 1),
"score": s.get("score", 0),
"preview": s.get("text", s.get("content", ""))[:200],
"doc_name": s.get("doc_title", s.get("doc_name", "")),
"clause": s.get("section_title", "法规片段"),
"doc_id": s.get("doc_id"),
"download_url": (
f"/api/v1/documents/download/{s['doc_id']}" if s.get("doc_id") else None
),
}
for idx, s in enumerate(data if isinstance(data, list) else [])
]
yield (
"event: message\n"
f"data: {json.dumps({'type': 'retrieved', 'docs': docs}, ensure_ascii=False)}\n\n"
)
elif event_type == "content":
if data:
yield (
"event: message\n"
f"data: {json.dumps({'type': 'chunk', 'text': data}, ensure_ascii=False)}\n\n"
)
elif event_type == "done":
yield (
"event: message\n"
f"data: {json.dumps({'type': 'done', 'session_id': session_id}, ensure_ascii=False)}\n\n"
)
elif event_type == "status":
yield (
"event: message\n"
f"data: {json.dumps({'type': 'status', 'text': data}, ensure_ascii=False)}\n\n"
)
elif event_type == "error":
yield (
"event: message\n"
f"data: {json.dumps({'type': 'error', 'text': str(data)}, ensure_ascii=False)}\n\n"
)
return StreamingResponse(
generate(),
media_type="text/event-stream",
headers={"Cache-Control": "no-cache", "Connection": "keep-alive", "X-Accel-Buffering": "no"},
)
@router.get("/quick-questions", response_model=QuickQuestionsResponse)
async def get_quick_questions():
"""Return configurable quick questions from settings or defaults."""
raw = getattr(settings, "rag_quick_questions", None)
if raw and isinstance(raw, list):
questions = [
QuickQuestion(id=str(i + 1), question=q if isinstance(q, str) else q.get("question", ""), category=q.get("category", "法规问答") if isinstance(q, dict) else "法规问答")
for i, q in enumerate(raw)
]
else:
questions = [QuickQuestion(**q) for q in _DEFAULT_QUICK_QUESTIONS]
return QuickQuestionsResponse(questions=questions)