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

472 lines
17 KiB
Python

import os
import json
import shutil
import time
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 = 50 * 1024 * 1024
UPLOAD_CHUNK_BYTES = 1024 * 1024
MULTIPART_CHUNK_BYTES = 5 * 1024 * 1024
MULTIPART_ROOT = ".multipart"
MULTIPART_TTL_SECONDS = 24 * 60 * 60
class QAIn(BaseModel):
question: str = ""
answer: str = ""
enabled: bool = True
class EnabledIn(BaseModel):
enabled: bool = True
class MultipartUploadIn(BaseModel):
filename: str
fileSize: int
totalChunks: int
def _validate_document(filename: str, file_size: int):
ext = os.path.splitext(filename or "")[1].lower()
if ext not in ALLOWED_EXT:
return None, f"不支持的文件类型:{ext or '空'},仅支持 md/txt/pdf/doc/docx/xlsx"
if file_size <= 0:
return None, "文件内容不能为空"
if file_size > MAX_UPLOAD_BYTES:
return None, "文件不能超过 50MB"
return ext, ""
def _multipart_dir(avatar_id: str, upload_id: str) -> str:
safe_avatar_id = os.path.basename(avatar_id)
safe_upload_id = os.path.basename(upload_id)
if (
safe_avatar_id != avatar_id
or safe_upload_id != upload_id
or len(upload_id) != 32
or any(character not in "0123456789abcdef" for character in upload_id)
):
raise HTTPException(status_code=400, detail="上传标识无效")
return os.path.join(UPLOAD_DIR, MULTIPART_ROOT, safe_avatar_id, safe_upload_id)
def _purge_stale_multipart_uploads(avatar_id: str):
avatar_upload_root = os.path.join(UPLOAD_DIR, MULTIPART_ROOT, os.path.basename(avatar_id))
if not os.path.isdir(avatar_upload_root):
return
cutoff = time.time() - MULTIPART_TTL_SECONDS
for entry in os.scandir(avatar_upload_root):
if entry.is_dir(follow_symlinks=False) and entry.stat(follow_symlinks=False).st_mtime < cutoff:
shutil.rmtree(entry.path, ignore_errors=True)
def _read_multipart_metadata(avatar_id: str, upload_id: str) -> tuple[str, dict]:
upload_dir = _multipart_dir(avatar_id, upload_id)
metadata_path = os.path.join(upload_dir, "metadata.json")
if not os.path.isfile(metadata_path):
raise HTTPException(status_code=404, detail="上传任务不存在或已过期")
with open(metadata_path, "r", encoding="utf-8") as stream:
return upload_dir, json.load(stream)
def _create_knowledge_doc(db: Session, avatar_id: str, filename: str, ext: str, file_size: int, stored: str):
doc = KnowledgeDoc(
id=uuid.uuid4().hex,
avatar_id=avatar_id,
filename=filename,
file_type=ext.lstrip("."),
file_size=file_size,
file_url=f"/api/files/{avatar_id}/{stored}",
status="parsing",
index_stage="queued",
index_progress=0,
)
db.add(doc)
db.commit()
db.refresh(doc)
knowledge_vectorizer.enqueue(doc.id)
return doc
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, validation_error = _validate_document(file.filename or "", 1)
if validation_error:
return fail(validation_error, 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)
file_size = 0
try:
# Stream large files to disk so a 100MB upload does not occupy 100MB RAM.
with open(path, "wb") as f:
while chunk := await file.read(UPLOAD_CHUNK_BYTES):
file_size += len(chunk)
if file_size > MAX_UPLOAD_BYTES:
raise ValueError("文件不能超过 50MB")
f.write(chunk)
except ValueError as exc:
if os.path.exists(path):
os.remove(path)
return fail(str(exc), code=400)
if file_size == 0:
if os.path.exists(path):
os.remove(path)
return fail("文件内容不能为空", code=400)
doc = _create_knowledge_doc(db, avatar_id, file.filename or stored, ext, file_size, stored)
return ok(_doc_payload(doc))
@router.post("/avatar/{avatar_id}/knowledge/uploads")
def create_multipart_upload(
avatar_id: str,
body: MultipartUploadIn,
authorization: str = Header(None),
db: Session = Depends(get_db),
):
_require_owned_avatar(db, avatar_id, authorization)
ext, validation_error = _validate_document(body.filename, body.fileSize)
if validation_error:
return fail(validation_error, code=400)
expected_chunks = (body.fileSize + MULTIPART_CHUNK_BYTES - 1) // MULTIPART_CHUNK_BYTES
if body.totalChunks != expected_chunks:
return fail("文件分片数量不正确", code=400)
_purge_stale_multipart_uploads(avatar_id)
upload_id = uuid.uuid4().hex
upload_dir = _multipart_dir(avatar_id, upload_id)
os.makedirs(upload_dir, exist_ok=False)
metadata = {
"filename": body.filename,
"fileSize": body.fileSize,
"totalChunks": body.totalChunks,
"extension": ext,
}
with open(os.path.join(upload_dir, "metadata.json"), "w", encoding="utf-8") as stream:
json.dump(metadata, stream, ensure_ascii=False)
return ok({"uploadId": upload_id, "chunkSize": MULTIPART_CHUNK_BYTES})
@router.post("/avatar/{avatar_id}/knowledge/uploads/{upload_id}/chunks/{chunk_index}")
async def upload_multipart_chunk(
avatar_id: str,
upload_id: str,
chunk_index: int,
file: UploadFile = File(...),
authorization: str = Header(None),
db: Session = Depends(get_db),
):
_require_owned_avatar(db, avatar_id, authorization)
upload_dir, metadata = _read_multipart_metadata(avatar_id, upload_id)
total_chunks = int(metadata["totalChunks"])
if chunk_index < 0 or chunk_index >= total_chunks:
return fail("文件分片序号不正确", code=400)
expected_size = min(
MULTIPART_CHUNK_BYTES,
int(metadata["fileSize"]) - chunk_index * MULTIPART_CHUNK_BYTES,
)
part_path = os.path.join(upload_dir, f"{chunk_index}.part")
temporary_path = f"{part_path}.uploading"
received = 0
try:
with open(temporary_path, "wb") as stream:
while chunk := await file.read(UPLOAD_CHUNK_BYTES):
received += len(chunk)
if received > expected_size:
raise ValueError("文件分片大小不正确")
stream.write(chunk)
if received != expected_size:
raise ValueError("文件分片大小不正确")
os.replace(temporary_path, part_path)
except ValueError as exc:
if os.path.exists(temporary_path):
os.remove(temporary_path)
return fail(str(exc), code=400)
return ok({"chunkIndex": chunk_index, "uploadedBytes": received})
@router.post("/avatar/{avatar_id}/knowledge/uploads/{upload_id}/complete")
def complete_multipart_upload(
avatar_id: str,
upload_id: str,
authorization: str = Header(None),
db: Session = Depends(get_db),
):
_require_owned_avatar(db, avatar_id, authorization)
upload_dir, metadata = _read_multipart_metadata(avatar_id, upload_id)
total_chunks = int(metadata["totalChunks"])
part_paths = [os.path.join(upload_dir, f"{index}.part") for index in range(total_chunks)]
if not all(os.path.isfile(path) for path in part_paths):
return fail("文件分片尚未上传完整", code=400)
if sum(os.path.getsize(path) for path in part_paths) != int(metadata["fileSize"]):
return fail("文件分片总大小不正确", code=400)
avatar_dir = os.path.join(UPLOAD_DIR, avatar_id)
os.makedirs(avatar_dir, exist_ok=True)
stored = f"{uuid.uuid4().hex}{metadata['extension']}"
final_path = os.path.join(avatar_dir, stored)
temporary_path = f"{final_path}.assembling"
try:
with open(temporary_path, "wb") as output:
for part_path in part_paths:
with open(part_path, "rb") as source:
shutil.copyfileobj(source, output, UPLOAD_CHUNK_BYTES)
os.replace(temporary_path, final_path)
doc = _create_knowledge_doc(
db,
avatar_id,
metadata["filename"],
metadata["extension"],
int(metadata["fileSize"]),
stored,
)
except Exception:
if os.path.exists(temporary_path):
os.remove(temporary_path)
raise
shutil.rmtree(upload_dir, ignore_errors=True)
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 = ""
doc.index_stage = "queued"
doc.index_progress = 0
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",
})