279 lines
9.3 KiB
Python
279 lines
9.3 KiB
Python
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",
|
|
})
|