334 lines
11 KiB
Python
334 lines
11 KiB
Python
#!/usr/bin/env python3
|
||
"""RAG 问答模块:向量检索 + LLM 生成。
|
||
|
||
流程:
|
||
1. 将用户问题转为向量
|
||
2. 在 embeddings 表中做余弦相似度检索,取 TOP-K 相关文章
|
||
3. 将检索到的文章内容作为上下文,调用千问生成回答
|
||
4. 返回回答 + 引用来源(可追溯)
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import json
|
||
import os
|
||
import sqlite3
|
||
import urllib.error
|
||
import urllib.request
|
||
from pathlib import Path
|
||
|
||
import numpy as np
|
||
|
||
from dotenv import load_dotenv
|
||
from structurer import get_api_key, get_embeddings, DASHSCOPE_CHAT_URL
|
||
|
||
BASE_DIR = Path(__file__).resolve().parent
|
||
load_dotenv(BASE_DIR / ".env")
|
||
|
||
|
||
def cosine_similarity(a: np.ndarray, b: np.ndarray) -> float:
|
||
"""计算两个向量的余弦相似度。"""
|
||
norm_a = np.linalg.norm(a)
|
||
norm_b = np.linalg.norm(b)
|
||
if norm_a == 0 or norm_b == 0:
|
||
return 0.0
|
||
return float(np.dot(a, b) / (norm_a * norm_b))
|
||
|
||
|
||
def vector_search(conn: sqlite3.Connection, query_embedding: list[float], top_k: int = 5) -> list[dict]:
|
||
"""向量检索:在 embeddings 表中找到与查询向量最相似的文章。"""
|
||
query_vec = np.array(query_embedding, dtype=np.float32)
|
||
rows = conn.execute(
|
||
"""
|
||
SELECT e.article_id, e.embedding, e.chunk_text,
|
||
a.title, a.source_name, a.source_type, a.published_at, a.summary,
|
||
x.result_json
|
||
FROM embeddings e
|
||
JOIN articles a ON a.id = e.article_id
|
||
LEFT JOIN ai_analyses x ON x.id = (
|
||
SELECT x2.id FROM ai_analyses x2
|
||
WHERE x2.article_id = e.article_id AND x2.analysis_type = 'editorial'
|
||
ORDER BY x2.updated_at DESC, x2.id DESC LIMIT 1
|
||
)
|
||
WHERE a.crawl_status = 'success'
|
||
"""
|
||
).fetchall()
|
||
|
||
scored = []
|
||
for row in rows:
|
||
emb = json.loads(row["embedding"])
|
||
emb_vec = np.array(emb, dtype=np.float32)
|
||
score = cosine_similarity(query_vec, emb_vec)
|
||
scored.append((score, row))
|
||
scored.sort(key=lambda x: x[0], reverse=True)
|
||
|
||
results = []
|
||
seen_articles = set()
|
||
for score, row in scored:
|
||
if row["article_id"] in seen_articles:
|
||
continue
|
||
seen_articles.add(row["article_id"])
|
||
analysis = json.loads(row["result_json"]) if row["result_json"] else {}
|
||
results.append({
|
||
"article_id": row["article_id"],
|
||
"title": row["title"],
|
||
"source": row["source_name"] or row["source_type"],
|
||
"date": (row["published_at"] or "")[:10],
|
||
"summary": analysis.get("summary") or row["summary"] or "",
|
||
"key_points": analysis.get("key_points", []),
|
||
"content": row["chunk_text"] or "",
|
||
"score": score,
|
||
})
|
||
if len(results) >= top_k:
|
||
break
|
||
return results
|
||
|
||
|
||
def generate_answer(question: str, contexts: list[dict]) -> dict:
|
||
"""基于检索到的上下文,调用千问生成回答。"""
|
||
api_key = get_api_key()
|
||
model = os.getenv("QWEN_MODEL", "qwen-plus")
|
||
|
||
# 构建上下文
|
||
context_text = ""
|
||
sources = []
|
||
for i, ctx in enumerate(contexts, start=1):
|
||
context_text += f"\n--- 来源{i} ---\n标题:{ctx['title']}\n摘要:{ctx['summary']}\n要点:{'; '.join(ctx['key_points'][:3])}\n原文片段:{ctx['content']}\n"
|
||
sources.append({
|
||
"title": ctx["title"],
|
||
"source": ctx["source"],
|
||
"date": ctx["date"],
|
||
})
|
||
|
||
prompt = f"""你是餐饮行业知识库的问答助手。根据以下知识库中的情报内容回答用户问题。
|
||
|
||
要求:
|
||
1. 回答必须基于提供的来源内容,不要编造信息
|
||
2. 如果来源内容不足以回答问题,明确说明"知识库中暂无直接相关信息"
|
||
3. 回答要结构化、简洁、有经营参考价值
|
||
4. 用中文回答
|
||
|
||
知识库来源:
|
||
{context_text}
|
||
|
||
用户问题:{question}"""
|
||
|
||
payload = json.dumps({
|
||
"model": model,
|
||
"messages": [
|
||
{"role": "system", "content": "你是餐饮行业知识库问答助手,基于库内情报给出有出处的专业回答。"},
|
||
{"role": "user", "content": prompt},
|
||
],
|
||
"temperature": 0.3,
|
||
}, ensure_ascii=False).encode("utf-8")
|
||
|
||
req = urllib.request.Request(
|
||
DASHSCOPE_CHAT_URL,
|
||
data=payload,
|
||
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
|
||
method="POST",
|
||
)
|
||
try:
|
||
with urllib.request.urlopen(req, timeout=90) as response:
|
||
body = json.loads(response.read().decode("utf-8"))
|
||
except urllib.error.HTTPError as exc:
|
||
detail = exc.read().decode("utf-8", errors="replace")[:500]
|
||
raise RuntimeError(f"千问接口返回 {exc.code}:{detail}") from exc
|
||
|
||
answer = body["choices"][0]["message"]["content"].strip()
|
||
return {
|
||
"answer": answer,
|
||
"sources": sources,
|
||
"contexts": contexts,
|
||
}
|
||
|
||
|
||
def generate_answer_stream(question: str, contexts: list[dict]):
|
||
"""流式生成回答,逐 token yield。"""
|
||
api_key = get_api_key()
|
||
model = os.getenv("QWEN_MODEL", "qwen-plus")
|
||
|
||
context_text = ""
|
||
sources = []
|
||
for i, ctx in enumerate(contexts, start=1):
|
||
context_text += f"\n--- 来源{i} ---\n标题:{ctx['title']}\n摘要:{ctx['summary']}\n要点:{'; '.join(ctx['key_points'][:3])}\n原文片段:{ctx['content']}\n"
|
||
sources.append({
|
||
"title": ctx["title"],
|
||
"source": ctx["source"],
|
||
"date": ctx["date"],
|
||
})
|
||
|
||
prompt = f"""你是餐饮行业知识库的问答助手。根据以下知识库中的情报内容回答用户问题。
|
||
|
||
要求:
|
||
1. 回答必须基于提供的来源内容,不要编造信息
|
||
2. 如果来源内容不足以回答问题,明确说明"知识库中暂无直接相关信息"
|
||
3. 回答要结构化、简洁、有经营参考价值
|
||
4. 用中文回答,使用 Markdown 格式
|
||
|
||
知识库来源:
|
||
{context_text}
|
||
|
||
用户问题:{question}"""
|
||
|
||
payload = json.dumps({
|
||
"model": model,
|
||
"messages": [
|
||
{"role": "system", "content": "你是餐饮行业知识库问答助手,基于库内情报给出有出处的专业回答。使用 Markdown 格式输出。"},
|
||
{"role": "user", "content": prompt},
|
||
],
|
||
"temperature": 0.3,
|
||
"stream": True,
|
||
"stream_options": {"include_usage": False},
|
||
}, ensure_ascii=False).encode("utf-8")
|
||
|
||
req = urllib.request.Request(
|
||
DASHSCOPE_CHAT_URL,
|
||
data=payload,
|
||
headers={
|
||
"Authorization": f"Bearer {api_key}",
|
||
"Content-Type": "application/json",
|
||
"Accept": "text/event-stream",
|
||
},
|
||
method="POST",
|
||
)
|
||
try:
|
||
resp = urllib.request.urlopen(req, timeout=120)
|
||
except urllib.error.HTTPError as exc:
|
||
detail = exc.read().decode("utf-8", errors="replace")[:500]
|
||
raise RuntimeError(f"千问接口返回 {exc.code}:{detail}") from exc
|
||
|
||
buffer = b""
|
||
while True:
|
||
chunk = resp.read(4096)
|
||
if not chunk:
|
||
break
|
||
buffer += chunk
|
||
while b"\n" in buffer:
|
||
line, buffer = buffer.split(b"\n", 1)
|
||
line = line.strip()
|
||
if not line or line == b"data: [DONE]":
|
||
continue
|
||
if line.startswith(b"data:"):
|
||
try:
|
||
data = json.loads(line[5:].strip())
|
||
delta = data.get("choices", [{}])[0].get("delta", {})
|
||
content = delta.get("content")
|
||
if content:
|
||
yield content
|
||
except (json.JSONDecodeError, IndexError, KeyError):
|
||
pass
|
||
resp.close()
|
||
|
||
yield json.dumps({"__sources__": sources}, ensure_ascii=False)
|
||
|
||
|
||
def rag_qa(question: str, top_k: int = 5) -> dict:
|
||
"""完整的 RAG 问答流程:向量检索 → LLM 生成。"""
|
||
db_path = Path(os.getenv("CIBANK_DB", str(BASE_DIR / "data" / "cibank.db")))
|
||
conn = sqlite3.connect(db_path)
|
||
conn.row_factory = sqlite3.Row
|
||
|
||
# 检查是否有向量数据
|
||
count = conn.execute("SELECT COUNT(*) FROM embeddings").fetchone()[0]
|
||
if count == 0:
|
||
# 无向量数据时使用关键词检索作为降级方案
|
||
return keyword_fallback(conn, question)
|
||
|
||
# 1. 将问题转为向量
|
||
query_embeddings = get_embeddings([question])
|
||
|
||
# 2. 向量检索
|
||
contexts = vector_search(conn, query_embeddings[0], top_k)
|
||
|
||
if not contexts:
|
||
conn.close()
|
||
return {
|
||
"answer": "知识库中暂无与您问题相关的内容。请先运行采集和分析流程入库更多情报。",
|
||
"sources": [],
|
||
"contexts": [],
|
||
}
|
||
|
||
# 3. LLM 生成回答
|
||
result = generate_answer(question, contexts)
|
||
conn.close()
|
||
return result
|
||
|
||
|
||
def rag_qa_stream(question: str, top_k: int = 5):
|
||
"""流式 RAG 问答:先检索,再流式生成。yield (type, data) 元组。"""
|
||
db_path = Path(os.getenv("CIBANK_DB", str(BASE_DIR / "data" / "cibank.db")))
|
||
conn = sqlite3.connect(db_path)
|
||
conn.row_factory = sqlite3.Row
|
||
|
||
count = conn.execute("SELECT COUNT(*) FROM embeddings").fetchone()[0]
|
||
if count == 0:
|
||
conn.close()
|
||
yield ("error", "知识库中暂无向量数据,请先运行采集和分析流程。")
|
||
return
|
||
|
||
query_embeddings = get_embeddings([question])
|
||
contexts = vector_search(conn, query_embeddings[0], top_k)
|
||
conn.close()
|
||
|
||
if not contexts:
|
||
yield ("error", "知识库中暂无与您问题相关的内容。")
|
||
return
|
||
|
||
sources = [
|
||
{"title": ctx["title"], "source": ctx["source"], "date": ctx["date"]}
|
||
for ctx in contexts
|
||
]
|
||
yield ("sources", sources)
|
||
|
||
for token in generate_answer_stream(question, contexts):
|
||
yield ("token", token)
|
||
|
||
|
||
def keyword_fallback(conn: sqlite3.Connection, question: str) -> dict:
|
||
"""无向量数据时的关键词检索降级方案。"""
|
||
rows = conn.execute(
|
||
"""
|
||
SELECT a.id, a.title, a.source_name, a.source_type, a.published_at, a.summary,
|
||
a.content, x.result_json
|
||
FROM articles a
|
||
LEFT JOIN ai_analyses x ON x.article_id = a.id AND x.analysis_type = 'editorial'
|
||
WHERE a.crawl_status = 'success'
|
||
AND (a.title LIKE ? OR a.summary LIKE ? OR a.content LIKE ?)
|
||
ORDER BY a.published_at DESC
|
||
LIMIT 5
|
||
""",
|
||
(f"%{question}%", f"%{question}%", f"%{question}%"),
|
||
).fetchall()
|
||
|
||
if not rows:
|
||
conn.close()
|
||
return {
|
||
"answer": "知识库中暂无与您问题相关的内容。请先运行采集和分析流程入库更多情报。",
|
||
"sources": [],
|
||
"contexts": [],
|
||
}
|
||
|
||
contexts = []
|
||
for row in rows:
|
||
analysis = json.loads(row["result_json"]) if row["result_json"] else {}
|
||
contexts.append({
|
||
"article_id": row["id"],
|
||
"title": row["title"],
|
||
"source": row["source_name"] or row["source_type"],
|
||
"date": (row["published_at"] or "")[:10],
|
||
"summary": analysis.get("summary") or row["summary"] or "",
|
||
"key_points": analysis.get("key_points", []),
|
||
"content": row["content"][:3000] if row["content"] else "",
|
||
"score": 0.5,
|
||
})
|
||
conn.close()
|
||
return generate_answer(question, contexts)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
import sys
|
||
question = " ".join(sys.argv[1:]) or "现在开茶饮店还有机会吗?"
|
||
result = rag_qa(question)
|
||
print(json.dumps(result, ensure_ascii=False, indent=2))
|