Files
2026-07-20 19:49:27 +08:00

334 lines
11 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/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))