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", })