155 lines
5.3 KiB
Python
155 lines
5.3 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 _embedding_endpoint(api_url):
|
|
"""Accept either an OpenAI-compatible base URL or its full endpoint."""
|
|
api_url = (api_url or "").strip().rstrip("/")
|
|
if not api_url or api_url.endswith("/embeddings"):
|
|
return api_url
|
|
return f"{api_url}/embeddings"
|
|
|
|
|
|
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 = _embedding_endpoint(os.getenv("EMBEDDING_API_URL"))
|
|
if api_url:
|
|
api_key = os.getenv("EMBEDDING_API_KEY", "")
|
|
model = os.getenv("EMBEDDING_MODEL", "text-embedding-3-small")
|
|
try:
|
|
batch_size = max(1, int(os.getenv("EMBEDDING_BATCH_SIZE", "10")))
|
|
except ValueError:
|
|
batch_size = 10
|
|
embeddings = []
|
|
for start in range(0, len(texts), batch_size):
|
|
batch = texts[start:start + batch_size]
|
|
payload = json.dumps({"input": batch, "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"])
|
|
if len(items) != len(batch):
|
|
raise ValueError("embedding response count does not match request")
|
|
embeddings.extend(item["embedding"] for item in items)
|
|
return embeddings
|
|
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}"
|