import os import json import logging 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() logger = logging.getLogger(__name__) 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 = 10 * 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() ) # Older synchronous uploads could be interrupted after persisting "parsing". # New uploads are committed only after indexing finishes, so these rows are stale. stale_docs = [doc for doc in docs if doc.status == "parsing"] if stale_docs: for doc in stale_docs: doc.status = "failed" doc.vectorized = False doc.chunk_count = 0 db.commit() 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("文件不能超过 10MB", 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", ) # Complete extraction and embedding before the first database commit so a # process restart cannot leave a permanent "parsing" row behind. try: text = embeddings.extract_text(path, ext) chunks = embeddings.chunk_text(text) if not chunks: raise ValueError("文档没有可建立索引的文字内容") vectors = embeddings.embed(chunks) if len(vectors) != len(chunks): raise ValueError("向量服务返回数量与文档分段不一致") doc.vectorized = True doc.embedding_model = embeddings.MODEL doc.chunk_count = len(chunks) doc.vectorized_at = datetime.now(timezone.utc) doc.status = "ready" db.add(doc) for i, (chunk, vector) in enumerate(zip(chunks, vectors)): db.add( KnowledgeChunk( doc_id=doc.id, avatar_id=avatar_id, content=chunk, vector=json.dumps(vector), chunk_index=i, embedding_model=embeddings.MODEL, ) ) db.commit() db.refresh(doc) except Exception as exc: db.rollback() doc.status = "failed" doc.vectorized = False doc.embedding_model = "" doc.chunk_count = 0 doc.vectorized_at = None db.add(doc) db.commit() db.refresh(doc) logger.exception("knowledge vectorization failed for %s: %s", doc.id, exc) 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", })