137 lines
4.5 KiB
Python
137 lines
4.5 KiB
Python
"""
|
||
向量化服务:对文档/查询文本生成向量。
|
||
|
||
优先级:
|
||
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}"
|