feat: add avatar chat and knowledge workflow
This commit is contained in:
136
digital-avatar-app/backend/embeddings.py
Normal file
136
digital-avatar-app/backend/embeddings.py
Normal file
@@ -0,0 +1,136 @@
|
||||
"""
|
||||
向量化服务:对文档/查询文本生成向量。
|
||||
|
||||
优先级:
|
||||
1. 若配置了环境变量 EMBEDDING_API_URL,则调用第三方「OpenAI 兼容」的 /embeddings 接口
|
||||
(需配置 EMBEDDING_API_KEY、EMBEDDING_MODEL,默认 text-embedding-3-small)。
|
||||
2. 否则使用本地「哈希 TF 嵌入」兜底,使向量检索在无外部依赖时也能端到端跑通,
|
||||
且相似文本(共享词汇)会得到更高余弦相似度,便于演示召回效果。
|
||||
"""
|
||||
import os
|
||||
import re
|
||||
import math
|
||||
import json
|
||||
import hashlib
|
||||
import urllib.request
|
||||
|
||||
EMBED_DIM = 256
|
||||
MODEL = os.getenv("EMBEDDING_MODEL", "mock-hash-embed-v1")
|
||||
|
||||
|
||||
def _tokenize(text):
|
||||
text = (text or "").lower()
|
||||
# 英文/数字按词,CJK 逐字(中文无空格,需拆到字级才能命中子词)
|
||||
tokens = re.findall(r"[a-z0-9]+", text)
|
||||
tokens += re.findall(r"[一-鿿]", text)
|
||||
return tokens
|
||||
|
||||
|
||||
def _hash_embedding(texts, dim=EMBED_DIM):
|
||||
vecs = []
|
||||
for text in texts:
|
||||
vec = [0.0] * dim
|
||||
tokens = _tokenize(text)
|
||||
if not tokens:
|
||||
tokens = list(text or "")
|
||||
for tok in tokens:
|
||||
h = int(hashlib.md5(tok.encode("utf-8")).hexdigest(), 16)
|
||||
vec[h % dim] += 1.0
|
||||
norm = math.sqrt(sum(v * v for v in vec))
|
||||
if norm > 0:
|
||||
vec = [v / norm for v in vec]
|
||||
vecs.append(vec)
|
||||
return vecs
|
||||
|
||||
|
||||
def embed(texts):
|
||||
"""返回 list[list[float]],与输入顺序一致。"""
|
||||
if not texts:
|
||||
return []
|
||||
api_url = os.getenv("EMBEDDING_API_URL")
|
||||
if api_url:
|
||||
api_key = os.getenv("EMBEDDING_API_KEY", "")
|
||||
model = os.getenv("EMBEDDING_MODEL", "text-embedding-3-small")
|
||||
payload = json.dumps({"input": texts, "model": model}).encode("utf-8")
|
||||
req = urllib.request.Request(
|
||||
api_url,
|
||||
data=payload,
|
||||
headers={
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer {api_key}" if api_key else "",
|
||||
},
|
||||
method="POST",
|
||||
)
|
||||
with urllib.request.urlopen(req, timeout=30) as resp:
|
||||
data = json.loads(resp.read().decode("utf-8"))
|
||||
items = data["data"]
|
||||
if items and "index" in items[0]:
|
||||
items = sorted(items, key=lambda x: x["index"])
|
||||
return [item["embedding"] for item in items]
|
||||
return _hash_embedding(texts)
|
||||
|
||||
|
||||
def cosine(a, b):
|
||||
dot = sum(x * y for x, y in zip(a, b))
|
||||
na = math.sqrt(sum(x * x for x in a))
|
||||
nb = math.sqrt(sum(y * y for y in b))
|
||||
if na == 0 or nb == 0:
|
||||
return 0.0
|
||||
return dot / (na * nb)
|
||||
|
||||
|
||||
def chunk_text(text, size=400, overlap=50):
|
||||
text = (text or "").strip()
|
||||
if not text:
|
||||
return []
|
||||
if len(text) <= size:
|
||||
return [text]
|
||||
chunks = []
|
||||
start = 0
|
||||
while start < len(text):
|
||||
end = min(start + size, len(text))
|
||||
chunks.append(text[start:end])
|
||||
if end == len(text):
|
||||
break
|
||||
start = end - overlap
|
||||
return chunks
|
||||
|
||||
|
||||
def extract_text(path, ext):
|
||||
"""抽取文档纯文本;未知格式拒绝,已知格式解析失败时保留占位文本。"""
|
||||
if ext not in {".txt", ".md", ".docx", ".xlsx", ".pdf", ".doc"}:
|
||||
raise ValueError(f"unsupported file extension: {ext}")
|
||||
try:
|
||||
if ext in {".txt", ".md"}:
|
||||
with open(path, "r", encoding="utf-8", errors="replace") as f:
|
||||
return f.read()
|
||||
if ext == ".docx":
|
||||
from docx import Document
|
||||
|
||||
doc = Document(path)
|
||||
return "\n".join(p.text for p in doc.paragraphs)
|
||||
if ext == ".xlsx":
|
||||
import openpyxl
|
||||
|
||||
wb = openpyxl.load_workbook(path, data_only=True, read_only=True)
|
||||
rows = []
|
||||
for ws in wb.worksheets:
|
||||
for row in ws.iter_rows(values_only=True):
|
||||
cells = [str(c) for c in row if c is not None]
|
||||
if cells:
|
||||
rows.append(" ".join(cells))
|
||||
return "\n".join(rows)
|
||||
if ext == ".pdf":
|
||||
try:
|
||||
from pypdf import PdfReader
|
||||
except ImportError:
|
||||
from PyPDF2 import PdfReader
|
||||
reader = PdfReader(path)
|
||||
return "\n".join((p.extract_text() or "") for p in reader.pages)
|
||||
if ext == ".doc":
|
||||
with open(path, "rb") as f:
|
||||
raw = f.read().decode("utf-8", errors="ignore")
|
||||
return re.sub(r"[\x00-\x08\x0b\x0c\x0e-\x1f]+", " ", raw)
|
||||
except Exception as e: # 解析失败时回退
|
||||
print("extract_text failed:", e)
|
||||
return f"文档:{os.path.basename(path)} 类型 {ext}"
|
||||
105
digital-avatar-app/backend/main.py
Normal file
105
digital-avatar-app/backend/main.py
Normal file
@@ -0,0 +1,105 @@
|
||||
from fastapi import FastAPI
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
|
||||
import os
|
||||
|
||||
from database import init_db, SessionLocal
|
||||
from models import Avatar, Authorization, Organization, TokenAccount, TokenPlan
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
import routers.avatars
|
||||
import routers.tokens
|
||||
import routers.authorizations
|
||||
import routers.organizations
|
||||
import routers.knowledge
|
||||
import routers.huihui_auth
|
||||
import routers.chat
|
||||
from responses import ok
|
||||
|
||||
app = FastAPI(title="会会数字分身 API", version="1.0.0")
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
allow_credentials=False,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
app.include_router(routers.avatars.router, prefix="/api")
|
||||
app.include_router(routers.tokens.router, prefix="/api")
|
||||
app.include_router(routers.authorizations.router, prefix="/api")
|
||||
app.include_router(routers.organizations.router, prefix="/api")
|
||||
app.include_router(routers.knowledge.router, prefix="/api")
|
||||
app.include_router(routers.huihui_auth.router, prefix="/api")
|
||||
app.include_router(routers.chat.router, prefix="/api")
|
||||
|
||||
UPLOAD_DIR = routers.knowledge.UPLOAD_DIR
|
||||
os.makedirs(UPLOAD_DIR, exist_ok=True)
|
||||
app.mount("/api/files", StaticFiles(directory=UPLOAD_DIR), name="knowledge-files")
|
||||
|
||||
|
||||
@app.get("/api/health")
|
||||
def health():
|
||||
return ok({"status": "ok"})
|
||||
|
||||
|
||||
def seed():
|
||||
db = SessionLocal()
|
||||
try:
|
||||
if db.query(TokenAccount).first() is None:
|
||||
db.add(TokenAccount(balance=1250))
|
||||
|
||||
if db.query(TokenPlan).count() == 0:
|
||||
plans = [
|
||||
TokenPlan(id="1", name="新手体验", amount=1000, price=9.9, desc="新手体验"),
|
||||
TokenPlan(id="2", name="热门套餐", amount=5000, price=39.9, badge="热门"),
|
||||
TokenPlan(id="3", name="超值套餐", amount=12000, price=89.9, badge="超值"),
|
||||
TokenPlan(id="4", name="企业推荐", amount=30000, price=199, badge="企业推荐", desc="适合高频使用"),
|
||||
]
|
||||
db.add_all(plans)
|
||||
|
||||
if db.query(Avatar).count() == 0:
|
||||
avatar = Avatar(
|
||||
name="我的数字分身",
|
||||
display_name="会会助手",
|
||||
description="我是您的AI数字分身,可以帮您管理日程、回复消息、处理任务。",
|
||||
emoji="🤖",
|
||||
status="active",
|
||||
token_balance=0,
|
||||
config={
|
||||
"replyStyle": "professional",
|
||||
"creativity": 50,
|
||||
"rigor": 50,
|
||||
"humor": 30,
|
||||
"responseLength": "medium",
|
||||
"systemPrompt": "",
|
||||
},
|
||||
)
|
||||
db.add(avatar)
|
||||
db.commit()
|
||||
db.refresh(avatar)
|
||||
|
||||
if db.query(Authorization).count() == 0:
|
||||
auths = [
|
||||
Authorization(avatar_id=avatar.id, target_type="application", target_name="微信小程序", permissions=["read", "reply"], status="active"),
|
||||
Authorization(avatar_id=avatar.id, target_type="user", target_name="张三", permissions=["read"], status="active"),
|
||||
Authorization(avatar_id=avatar.id, target_type="organization", target_name="产品团队", permissions=["read", "edit"], status="inactive"),
|
||||
]
|
||||
db.add_all(auths)
|
||||
|
||||
if db.query(Organization).count() == 0:
|
||||
orgs = [
|
||||
Organization(name="会会增长团队", description="负责会会产品的增长与运营", emoji="🚀", org_type="team", member_count=12),
|
||||
Organization(name="AI 实验室", description="探索前沿 AI 能力", emoji="💡", org_type="company", member_count=8),
|
||||
]
|
||||
db.add_all(orgs)
|
||||
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
@app.on_event("startup")
|
||||
def on_startup():
|
||||
init_db()
|
||||
seed()
|
||||
205
digital-avatar-app/backend/routers/chat.py
Normal file
205
digital-avatar-app/backend/routers/chat.py
Normal file
@@ -0,0 +1,205 @@
|
||||
import difflib
|
||||
import os
|
||||
import re
|
||||
import string
|
||||
from typing import Any, Callable
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Body, Depends, Header, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
import embeddings
|
||||
from database import get_db
|
||||
from models import Avatar, KnowledgeChunk, KnowledgeDoc, QAPair, User
|
||||
from responses import ok, fail
|
||||
|
||||
router = APIRouter(tags=["数字分身聊天"])
|
||||
|
||||
CHAT_API_URL = os.getenv("CHAT_API_URL", "https://dashscope.aliyuncs.com/compatible-mode/v1")
|
||||
CHAT_API_KEY = os.getenv("CHAT_API_KEY", "")
|
||||
CHAT_MODEL = os.getenv("CHAT_MODEL", "qwen-plus")
|
||||
MAX_MESSAGE_LENGTH = 4000
|
||||
MAX_HISTORY_MESSAGES = 10
|
||||
QA_SIMILARITY_THRESHOLD = 0.86
|
||||
|
||||
|
||||
class ChatMessage(BaseModel):
|
||||
role: str = Field(pattern="^(user|assistant)$")
|
||||
content: str = Field(min_length=1, max_length=MAX_MESSAGE_LENGTH)
|
||||
|
||||
|
||||
class ChatIn(BaseModel):
|
||||
message: str = Field(min_length=1, max_length=MAX_MESSAGE_LENGTH)
|
||||
history: list[ChatMessage] = Field(default_factory=list, max_length=MAX_HISTORY_MESSAGES)
|
||||
|
||||
|
||||
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="分身不存在")
|
||||
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
|
||||
|
||||
|
||||
def _normalize_question(value: str) -> str:
|
||||
value = (value or "").strip().lower()
|
||||
value = re.sub(r"\s+", "", value)
|
||||
return value.translate(str.maketrans("", "", string.punctuation + ",。!?;:、()【】「」‘’“”《》"))
|
||||
|
||||
|
||||
def _match_standard_qa(question: str, qa_pairs: list[Any]):
|
||||
normalized = _normalize_question(question)
|
||||
if not normalized:
|
||||
return None
|
||||
enabled = [qa for qa in qa_pairs if getattr(qa, "enabled", True)]
|
||||
for qa in enabled:
|
||||
if _normalize_question(getattr(qa, "question", "")) == normalized:
|
||||
return qa
|
||||
best = None
|
||||
best_score = 0.0
|
||||
for qa in enabled:
|
||||
candidate = _normalize_question(getattr(qa, "question", ""))
|
||||
if not candidate:
|
||||
continue
|
||||
score = difflib.SequenceMatcher(None, normalized, candidate).ratio()
|
||||
if score > best_score:
|
||||
best, best_score = qa, score
|
||||
return best if best_score >= QA_SIMILARITY_THRESHOLD else None
|
||||
|
||||
|
||||
def _config(avatar: Avatar) -> dict:
|
||||
config = getattr(avatar, "config", None) or {}
|
||||
return {
|
||||
"replyStyle": config.get("replyStyle", "professional"),
|
||||
"creativity": max(0, min(100, int(config.get("creativity", 50)))),
|
||||
"rigor": max(0, min(100, int(config.get("rigor", 50)))),
|
||||
"humor": max(0, min(100, int(config.get("humor", 30)))),
|
||||
"responseLength": config.get("responseLength", "medium"),
|
||||
"systemPrompt": (config.get("systemPrompt", "") or "").strip(),
|
||||
}
|
||||
|
||||
|
||||
def _build_prompt(avatar: Avatar, history: list[Any], question: str, knowledge_hits: list[dict]) -> list[dict]:
|
||||
config = _config(avatar)
|
||||
knowledge = "\n".join(
|
||||
f"[{hit.get('filename', '知识库')}] {hit.get('snippet', '')}"
|
||||
for hit in knowledge_hits
|
||||
if hit.get("snippet")
|
||||
)
|
||||
system = (
|
||||
"你是用户的专属数字分身。请基于已提供的知识库回答,不要编造事实;"
|
||||
f"回复风格:{config['replyStyle']};严谨度:{config['rigor']}/100;"
|
||||
f"幽默感:{config['humor']}/100;回复长度:{config['responseLength']}。"
|
||||
)
|
||||
if config["systemPrompt"]:
|
||||
system += f"\n额外系统提示词:{config['systemPrompt']}"
|
||||
if knowledge:
|
||||
system += f"\n以下是可参考的知识库内容:\n{knowledge}"
|
||||
messages = [{"role": "system", "content": system}]
|
||||
for item in history[-MAX_HISTORY_MESSAGES:]:
|
||||
messages.append({"role": item.role, "content": item.content} if hasattr(item, "role") else item)
|
||||
messages.append({"role": "user", "content": question.strip()})
|
||||
return messages
|
||||
|
||||
|
||||
def _search_knowledge(db: Session, avatar_id: str, question: str, top_k: int = 5) -> list[dict]:
|
||||
chunks = db.query(KnowledgeChunk).filter(KnowledgeChunk.avatar_id == avatar_id).all()
|
||||
if not chunks:
|
||||
return []
|
||||
qvec = embeddings.embed([question])[0]
|
||||
scored = []
|
||||
for chunk in chunks:
|
||||
try:
|
||||
vector = __import__("json").loads(chunk.vector)
|
||||
except Exception:
|
||||
continue
|
||||
scored.append((embeddings.cosine(qvec, vector), chunk))
|
||||
scored.sort(key=lambda item: item[0], reverse=True)
|
||||
results = []
|
||||
for score, chunk in scored[: max(1, top_k)]:
|
||||
doc = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == chunk.doc_id).first()
|
||||
results.append({
|
||||
"docId": chunk.doc_id,
|
||||
"filename": doc.filename if doc else "",
|
||||
"fileType": doc.file_type if doc else "",
|
||||
"snippet": chunk.content[:120] + ("…" if len(chunk.content) > 120 else ""),
|
||||
"score": round(score, 4),
|
||||
})
|
||||
return results
|
||||
|
||||
|
||||
def _call_qwen(messages: list[dict], temperature: float) -> str:
|
||||
if not CHAT_API_KEY:
|
||||
raise RuntimeError("Qwen 模型服务未配置 CHAT_API_KEY")
|
||||
url = f"{CHAT_API_URL.rstrip('/')}/chat/completions"
|
||||
payload = {
|
||||
"model": CHAT_MODEL,
|
||||
"messages": messages,
|
||||
"temperature": temperature,
|
||||
}
|
||||
try:
|
||||
response = httpx.post(
|
||||
url,
|
||||
headers={"Authorization": f"Bearer {CHAT_API_KEY}"},
|
||||
json=payload,
|
||||
timeout=30,
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
answer = data.get("choices", [{}])[0].get("message", {}).get("content", "")
|
||||
except (httpx.HTTPError, ValueError, KeyError, IndexError) as exc:
|
||||
raise RuntimeError("Qwen 模型服务暂时不可用") from exc
|
||||
if not isinstance(answer, str) or not answer.strip():
|
||||
raise RuntimeError("Qwen 模型没有返回有效回答")
|
||||
return answer.strip()
|
||||
|
||||
|
||||
def _resolve_reply(
|
||||
db: Session,
|
||||
avatar: Avatar,
|
||||
question: str,
|
||||
history: list[Any],
|
||||
*,
|
||||
qa_pairs: list[Any] | None = None,
|
||||
search_fn: Callable[..., list[dict]] | None = None,
|
||||
model_client: Callable[..., str] | None = None,
|
||||
) -> dict:
|
||||
if qa_pairs is None:
|
||||
qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all()
|
||||
matched = _match_standard_qa(question, qa_pairs)
|
||||
if matched:
|
||||
return {"answer": matched.answer, "source": "qa", "references": []}
|
||||
|
||||
search_fn = search_fn or (lambda query, avatar_id: _search_knowledge(db, avatar_id, query))
|
||||
hits = search_fn(question, avatar.id)
|
||||
messages = _build_prompt(avatar, history, question, hits)
|
||||
config = _config(avatar)
|
||||
temperature = 0.2 + config["creativity"] / 100 * 0.6
|
||||
model_client = model_client or _call_qwen
|
||||
answer = model_client(messages=messages, temperature=temperature)
|
||||
return {
|
||||
"answer": answer,
|
||||
"source": "knowledge" if hits else "qwen",
|
||||
"references": hits,
|
||||
}
|
||||
|
||||
|
||||
@router.post("/avatar/{avatar_id}/chat")
|
||||
def chat(avatar_id: str, body: ChatIn = Body(...), authorization: str = Header(None), db: Session = Depends(get_db)):
|
||||
avatar = _require_owned_avatar(db, avatar_id, authorization)
|
||||
try:
|
||||
return ok(_resolve_reply(db, avatar, body.message, body.history))
|
||||
except RuntimeError as exc:
|
||||
return fail(str(exc), code=502)
|
||||
278
digital-avatar-app/backend/routers/knowledge.py
Normal file
278
digital-avatar-app/backend/routers/knowledge.py
Normal file
@@ -0,0 +1,278 @@
|
||||
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",
|
||||
})
|
||||
1
digital-avatar-app/backend/tests/__init__.py
Normal file
1
digital-avatar-app/backend/tests/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
|
||||
95
digital-avatar-app/backend/tests/test_chat_orchestration.py
Normal file
95
digital-avatar-app/backend/tests/test_chat_orchestration.py
Normal file
@@ -0,0 +1,95 @@
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import Mock
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from models import Avatar, User
|
||||
from routers.chat import _build_prompt, _match_standard_qa, _require_owned_avatar, _resolve_reply
|
||||
|
||||
|
||||
class ChatOrchestrationTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.avatar = SimpleNamespace(
|
||||
id="avatar-1",
|
||||
owner_id="huihui-user-1",
|
||||
config={
|
||||
"replyStyle": "professional",
|
||||
"creativity": 50,
|
||||
"rigor": 80,
|
||||
"humor": 20,
|
||||
"responseLength": "medium",
|
||||
"systemPrompt": "不要编造政策。",
|
||||
},
|
||||
)
|
||||
self.qa = SimpleNamespace(question="公司地址?", answer="标准地址", enabled=True)
|
||||
self.disabled_qa = SimpleNamespace(question="公司地址?", answer="错误答案", enabled=False)
|
||||
|
||||
def test_enabled_qa_wins_without_calling_model(self):
|
||||
fake_model = Mock()
|
||||
result = _resolve_reply(
|
||||
None,
|
||||
self.avatar,
|
||||
" 公司地址? ",
|
||||
[],
|
||||
qa_pairs=[self.disabled_qa, self.qa],
|
||||
search_fn=lambda *_args, **_kwargs: [],
|
||||
model_client=fake_model,
|
||||
)
|
||||
self.assertEqual(result["source"], "qa")
|
||||
self.assertEqual(result["answer"], "标准地址")
|
||||
fake_model.assert_not_called()
|
||||
|
||||
def test_knowledge_context_is_sent_to_qwen_after_qa_miss(self):
|
||||
fake_model = Mock(return_value="根据知识库内容回答")
|
||||
knowledge_hit = {
|
||||
"filename": "退款.md",
|
||||
"snippet": "知识库内容:七日内可申请退款。",
|
||||
"score": 0.92,
|
||||
}
|
||||
result = _resolve_reply(
|
||||
None,
|
||||
self.avatar,
|
||||
"退款规则",
|
||||
[],
|
||||
qa_pairs=[],
|
||||
search_fn=lambda *_args, **_kwargs: [knowledge_hit],
|
||||
model_client=fake_model,
|
||||
)
|
||||
self.assertEqual(result["source"], "knowledge")
|
||||
self.assertIn("知识库内容", fake_model.call_args.kwargs["messages"][0]["content"])
|
||||
|
||||
def test_prompt_contains_personality_configuration(self):
|
||||
messages = _build_prompt(self.avatar, [], "你好", [])
|
||||
self.assertIn("严谨度", messages[0]["content"])
|
||||
self.assertIn("不要编造政策", messages[0]["content"])
|
||||
|
||||
def test_chat_rejects_avatar_owned_by_another_user(self):
|
||||
class Query:
|
||||
def __init__(self, value):
|
||||
self.value = value
|
||||
|
||||
def filter(self, *_args, **_kwargs):
|
||||
return self
|
||||
|
||||
def first(self):
|
||||
return self.value
|
||||
|
||||
self_avatar = self.avatar
|
||||
|
||||
class DB:
|
||||
avatar = self_avatar
|
||||
|
||||
def query(self, model):
|
||||
return Query(
|
||||
self.avatar if model is Avatar else SimpleNamespace(huihui_user_id="huihui-user-2")
|
||||
)
|
||||
|
||||
db = DB()
|
||||
with self.assertRaises(HTTPException) as caught:
|
||||
_require_owned_avatar(db, self.avatar.id, "Bearer other-token")
|
||||
self.assertEqual(caught.exception.status_code, 403)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
32
digital-avatar-app/backend/tests/test_embeddings.py
Normal file
32
digital-avatar-app/backend/tests/test_embeddings.py
Normal file
@@ -0,0 +1,32 @@
|
||||
import os
|
||||
import tempfile
|
||||
import unittest
|
||||
|
||||
import embeddings
|
||||
|
||||
|
||||
class TextExtractionTests(unittest.TestCase):
|
||||
def write_text(self, suffix, content):
|
||||
handle = tempfile.NamedTemporaryFile(suffix=suffix, delete=False)
|
||||
handle.close()
|
||||
self.addCleanup(lambda: os.path.exists(handle.name) and os.unlink(handle.name))
|
||||
with open(handle.name, "w", encoding="utf-8") as stream:
|
||||
stream.write(content)
|
||||
return handle.name
|
||||
|
||||
def test_extracts_utf8_markdown(self):
|
||||
path = self.write_text(".md", "# 退款规则\n\n七日内可申请退款。")
|
||||
self.assertEqual(embeddings.extract_text(path, ".md"), "# 退款规则\n\n七日内可申请退款。")
|
||||
|
||||
def test_extracts_utf8_text(self):
|
||||
path = self.write_text(".txt", "客服热线:400-123-4567")
|
||||
self.assertEqual(embeddings.extract_text(path, ".txt"), "客服热线:400-123-4567")
|
||||
|
||||
def test_rejects_unsupported_extension(self):
|
||||
path = self.write_text(".csv", "not supported")
|
||||
with self.assertRaises(ValueError):
|
||||
embeddings.extract_text(path, ".csv")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user