import os import json import uuid from datetime import datetime, timezone 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 router = APIRouter() BASE_DIR = os.path.dirname(os.path.abspath(__file__)) 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 = 10 * 1024 * 1024 class QAIn(BaseModel): question: str = "" answer: str = "" enabled: bool = True class EnabledIn(BaseModel): enabled: bool = True 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([d.to_dict() 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("文件不能超过 10MB", code=400) with open(path, "wb") as f: f.write(content) doc = KnowledgeDoc( 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", ) db.add(doc) db.commit() db.refresh(doc) # 向量化:抽取文本 -> 分块 -> 调第三方/本地嵌入 -> 存切片 try: text = embeddings.extract_text(path, ext) chunks = embeddings.chunk_text(text) if chunks: vectors = embeddings.embed(chunks) for i, (c, v) in enumerate(zip(chunks, vectors)): db.add( KnowledgeChunk( doc_id=doc.id, avatar_id=avatar_id, content=c, vector=json.dumps(v), chunk_index=i, embedding_model=embeddings.MODEL, ) ) doc.vectorized = True doc.embedding_model = embeddings.MODEL doc.chunk_count = len(chunks) doc.vectorized_at = datetime.now(timezone.utc) doc.status = "ready" db.commit() db.refresh(doc) except Exception as e: print("vectorize failed:", e) doc.status = "ready" # 上传成功但向量化失败,仍可展示 db.commit() db.refresh(doc) return ok(doc.to_dict()) @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", })