Files
huihuiSquare/digital-avatar-app/backend/routers/knowledge.py
T

286 lines
9.8 KiB
Python

import os
import uuid
from fastapi import APIRouter, UploadFile, File, Depends, Header, HTTPException
from pydantic import BaseModel
from sqlalchemy.orm import Session
from database import get_db
from models import KnowledgeDoc, QAPair, KnowledgeChunk, Avatar, User
from responses import ok, fail
import embeddings
from services.knowledge_vectorizer import knowledge_vectorizer
router = APIRouter()
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
UPLOAD_DIR = os.path.abspath(os.getenv("UPLOAD_DIR", os.path.join(BASE_DIR, "uploads")))
os.makedirs(UPLOAD_DIR, exist_ok=True)
ALLOWED_EXT = {".md", ".txt", ".pdf", ".doc", ".docx", ".xlsx"}
MAX_UPLOAD_BYTES = 20 * 1024 * 1024
class QAIn(BaseModel):
question: str = ""
answer: str = ""
enabled: bool = True
class EnabledIn(BaseModel):
enabled: bool = True
def _doc_payload(doc: KnowledgeDoc) -> dict:
payload = doc.to_dict()
stored_name = os.path.basename(doc.file_url or "")
stored_path = os.path.join(UPLOAD_DIR, doc.avatar_id, stored_name)
payload["filePresent"] = bool(stored_name and os.path.isfile(stored_path))
return payload
def _resolve_user(authorization: str | None, db: Session):
if not authorization:
return None
token = authorization.replace("Bearer ", "", 1).replace("bearer ", "", 1).strip()
return db.query(User).filter(User.app_token == token).first()
def _require_owned_avatar(db: Session, avatar_id: str, authorization: str | None):
avatar = db.query(Avatar).filter(Avatar.id == avatar_id).first()
if not avatar:
raise HTTPException(status_code=404, detail="avatar not found")
user = _resolve_user(authorization, db)
if not user:
raise HTTPException(status_code=401, detail="未登录")
if avatar.owner_id and avatar.owner_id != user.huihui_user_id:
raise HTTPException(status_code=403, detail="无权访问该分身")
return avatar
# ---------------- Documents ----------------
@router.get("/avatar/{avatar_id}/knowledge/docs")
def list_docs(avatar_id: str, authorization: str = Header(None), db: Session = Depends(get_db)):
_require_owned_avatar(db, avatar_id, authorization)
docs = (
db.query(KnowledgeDoc)
.filter(KnowledgeDoc.avatar_id == avatar_id)
.order_by(KnowledgeDoc.created_at.desc())
.all()
)
return ok([_doc_payload(d) for d in docs])
@router.post("/avatar/{avatar_id}/knowledge/docs")
async def upload_doc(avatar_id: str, file: UploadFile = File(...), authorization: str = Header(None), db: Session = Depends(get_db)):
_require_owned_avatar(db, avatar_id, authorization)
ext = os.path.splitext(file.filename or "")[1].lower()
if ext not in ALLOWED_EXT:
return fail(f"不支持的文件类型:{ext or '空'},仅支持 md/txt/pdf/doc/docx/xlsx", code=400)
avatar_dir = os.path.join(UPLOAD_DIR, avatar_id)
os.makedirs(avatar_dir, exist_ok=True)
stored = f"{uuid.uuid4().hex}{ext}"
path = os.path.join(avatar_dir, stored)
content = await file.read()
if len(content) > MAX_UPLOAD_BYTES:
return fail("文件不能超过 20MB", code=400)
with open(path, "wb") as f:
f.write(content)
doc = KnowledgeDoc(
id=uuid.uuid4().hex,
avatar_id=avatar_id,
filename=file.filename,
file_type=ext.lstrip("."),
file_size=len(content),
file_url=f"/api/files/{avatar_id}/{stored}",
status="parsing",
)
# Persist and acknowledge the upload first. Extraction and embeddings may take
# minutes for a PDF and must never consume the browser request timeout.
db.add(doc)
db.commit()
db.refresh(doc)
knowledge_vectorizer.enqueue(doc.id)
return ok(_doc_payload(doc))
@router.post("/avatar/{avatar_id}/knowledge/docs/{doc_id}/retry")
def retry_doc(avatar_id: str, doc_id: str, authorization: str = Header(None), db: Session = Depends(get_db)):
_require_owned_avatar(db, avatar_id, authorization)
doc = db.query(KnowledgeDoc).filter(
KnowledgeDoc.id == doc_id, KnowledgeDoc.avatar_id == avatar_id
).first()
if not doc:
return fail("文档不存在", code=404)
if doc.vectorized and doc.status == "ready":
return ok(_doc_payload(doc))
stored_name = os.path.basename(doc.file_url or "")
if not stored_name or not os.path.isfile(os.path.join(UPLOAD_DIR, avatar_id, stored_name)):
return fail("原文件不可用,请重新上传", code=400)
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == doc.id).delete()
doc.status = "parsing"
doc.vectorized = False
doc.embedding_model = ""
doc.chunk_count = 0
doc.vectorized_at = None
doc.error_message = ""
db.commit()
db.refresh(doc)
knowledge_vectorizer.enqueue(doc.id)
return ok(_doc_payload(doc))
@router.delete("/avatar/{avatar_id}/knowledge/docs/{doc_id}")
def delete_doc(avatar_id: str, doc_id: str, authorization: str = Header(None), db: Session = Depends(get_db)):
_require_owned_avatar(db, avatar_id, authorization)
doc = (
db.query(KnowledgeDoc)
.filter(KnowledgeDoc.id == doc_id, KnowledgeDoc.avatar_id == avatar_id)
.first()
)
if not doc:
return fail("文档不存在", code=404)
# 级联删除切片
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == doc_id).delete()
try:
fp = os.path.join(UPLOAD_DIR, avatar_id, os.path.basename(doc.file_url))
if os.path.exists(fp):
os.remove(fp)
except Exception:
pass
db.delete(doc)
db.commit()
return ok({"id": doc_id})
# ---------------- 向量检索 ----------------
@router.get("/avatar/{avatar_id}/knowledge/search")
def search_knowledge(avatar_id: str, q: str = "", top_k: int = 5, authorization: str = Header(None), db: Session = Depends(get_db)):
_require_owned_avatar(db, avatar_id, authorization)
q = (q or "").strip()
if not q:
return ok([])
chunks = (
db.query(KnowledgeChunk)
.filter(KnowledgeChunk.avatar_id == avatar_id)
.all()
)
if not chunks:
return ok([])
qvec = embeddings.embed([q])[0]
scored = []
for c in chunks:
try:
vec = json.loads(c.vector)
except Exception:
continue
scored.append((embeddings.cosine(qvec, vec), c))
scored.sort(key=lambda x: x[0], reverse=True)
results = []
for score, c in scored[: max(1, top_k)]:
doc = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == c.doc_id).first()
snippet = c.content[:120] + ("…" if len(c.content) > 120 else "")
results.append(
{
"docId": c.doc_id,
"filename": doc.filename if doc else "",
"fileType": doc.file_type if doc else "",
"snippet": snippet,
"score": round(score, 4),
}
)
return ok(results)
# ---------------- Standard Q&A pairs ----------------
@router.get("/avatar/{avatar_id}/knowledge/qa")
def list_qa(avatar_id: str, authorization: str = Header(None), db: Session = Depends(get_db)):
_require_owned_avatar(db, avatar_id, authorization)
items = (
db.query(QAPair)
.filter(QAPair.avatar_id == avatar_id)
.order_by(QAPair.created_at.desc())
.all()
)
return ok([q.to_dict() for q in items])
@router.post("/avatar/{avatar_id}/knowledge/qa")
def create_qa(avatar_id: str, body: QAIn, authorization: str = Header(None), db: Session = Depends(get_db)):
_require_owned_avatar(db, avatar_id, authorization)
q = QAPair(
avatar_id=avatar_id,
question=body.question,
answer=body.answer,
enabled=body.enabled,
)
db.add(q)
db.commit()
db.refresh(q)
return ok(q.to_dict())
@router.put("/avatar/{avatar_id}/knowledge/qa/{qa_id}")
def update_qa(avatar_id: str, qa_id: str, body: QAIn, authorization: str = Header(None), db: Session = Depends(get_db)):
_require_owned_avatar(db, avatar_id, authorization)
q = (
db.query(QAPair)
.filter(QAPair.id == qa_id, QAPair.avatar_id == avatar_id)
.first()
)
if not q:
return fail("问答对不存在", code=404)
q.question = body.question
q.answer = body.answer
q.enabled = body.enabled
db.commit()
db.refresh(q)
return ok(q.to_dict())
@router.put("/avatar/{avatar_id}/knowledge/qa/{qa_id}/enabled")
def set_qa_enabled(avatar_id: str, qa_id: str, body: EnabledIn, authorization: str = Header(None), db: Session = Depends(get_db)):
_require_owned_avatar(db, avatar_id, authorization)
q = (
db.query(QAPair)
.filter(QAPair.id == qa_id, QAPair.avatar_id == avatar_id)
.first()
)
if not q:
return fail("问答对不存在", code=404)
q.enabled = bool(body.enabled)
db.commit()
db.refresh(q)
return ok(q.to_dict())
@router.delete("/avatar/{avatar_id}/knowledge/qa/{qa_id}")
def delete_qa(avatar_id: str, qa_id: str, authorization: str = Header(None), db: Session = Depends(get_db)):
_require_owned_avatar(db, avatar_id, authorization)
q = (
db.query(QAPair)
.filter(QAPair.id == qa_id, QAPair.avatar_id == avatar_id)
.first()
)
if not q:
return fail("问答对不存在", code=404)
db.delete(q)
db.commit()
return ok({"id": qa_id})
# ---------------- HuiHui user profile (mock; plug real interface via HUIHUI_USER_API) ----------------
@router.get("/user/profile")
def user_profile():
# 接入真实会会接口:设置环境变量 HUIHUI_USER_API 后在此请求并映射字段
api = os.getenv("HUIHUI_USER_API")
if api:
# TODO: 调用会会用户接口,返回 { userId, nickname, avatarUrl }
pass
return ok({
"userId": "hh_10001",
"nickname": "会会用户",
"avatarUrl": "https://api.dicebear.com/7.x/initials/svg?seed=HuiHui&backgroundColor=F97316",
})