472 lines
17 KiB
Python
472 lines
17 KiB
Python
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",
|
|
})
|