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 = 50 * 1024 * 1024 UPLOAD_CHUNK_BYTES = 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) 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) doc = KnowledgeDoc( id=uuid.uuid4().hex, avatar_id=avatar_id, filename=file.filename, file_type=ext.lstrip("."), file_size=file_size, 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", })