# -*- coding: utf-8 -*- """ 向量检索层:Embedding(16011) + Chroma(16010) 纯 REST 实现,无第三方客户端依赖。 - embedding:OpenAI 兼容 /v1/embeddings(bge-large-zh-v1.5,1024维) - chroma :REST /api/v2(add/query 用 UUID,get/delete 用名字) 两个集合: - stock_news_v1 财经新闻正文语义索引(metadata: code/title/date/category/sentiment/news_id) - stock_profiles_v1 公司概况语义索引(metadata: code/name/industry) """ import json import logging import threading import urllib.request import urllib.error import requests from config import (CHROMA_HOST, CHROMA_PORT, EMBEDDING_API_URL, EMBEDDING_MODEL, RERANK_API_URL, RERANK_MODEL, USE_RERANK) log = logging.getLogger("vector") _BASE = f"http://{CHROMA_HOST}:{CHROMA_PORT}/api/v2/tenants/default_tenant/databases/default_database/collections" _lock = threading.Lock() # ------------------------------------------------------------------ Embedding def embed_texts(texts): """返回 [[float...], ...] 向量列表""" if isinstance(texts, str): texts = [texts] resp = requests.post(EMBEDDING_API_URL, json={"model": EMBEDDING_MODEL, "input": list(texts)}, timeout=120) resp.raise_for_status() data = resp.json().get("data", []) return [d["embedding"] for d in sorted(data, key=lambda x: x["index"])] def rerank(query, docs, top_k=5): if not USE_RERANK or not docs: return docs try: payload = {"query": query, "documents": docs, "model": RERANK_MODEL, "top_k": top_k} resp = requests.post(RERANK_API_URL, json=payload, timeout=60) resp.raise_for_status() data = resp.json().get("data", []) return [{"id": d["document"]["id"], "text": d["document"]["text"], "score": d["score"]} for d in data] except Exception as e: log.warning("rerank failed: %s", e) return docs # ------------------------------------------------------------------ Chroma def _http(method, url, payload=None, timeout=30): req = urllib.request.Request(url, method=method) if payload is not None: req.add_header("Content-Type", "application/json") req.data = json.dumps(payload).encode("utf-8") try: with urllib.request.urlopen(req, timeout=timeout) as r: body = r.read().decode("utf-8") return r.status, (json.loads(body) if body else {}) except urllib.error.HTTPError as e: body = e.read().decode("utf-8", "ignore") raise RuntimeError(f"Chroma {method} {url} -> {e.code}: {body[:300]}") def _get_collection_id(name): """按名字查集合,返回 (id, 是否存在)""" try: _, data = _http("GET", f"{_BASE}/{name}") return data.get("id"), True except RuntimeError: return None, False def ensure_collection(name, space="cosine"): """获取或创建集合,返回 collection_id""" with _lock: cid, exists = _get_collection_id(name) if exists: return cid _, data = _http("POST", _BASE, { "name": name, "configuration": {"hnsw": {"space": space}}, "get_or_create": True, }) return data["id"] def collection_count(name): try: cid, exists = _get_collection_id(name) if not exists: return 0 _, data = _http("GET", f"{_BASE}/{cid}/count") return int(data) if isinstance(data, int) else int(data.get("count", 0)) except Exception: return 0 def add_documents(ids, documents, metadatas, name): """按文档批量写入(内部自动 embedding)""" if not ids: return cid = ensure_collection(name) vectors = embed_texts(documents) payload = {"ids": list(ids), "embeddings": vectors, "documents": list(documents), "metadatas": list(metadatas)} _http("POST", f"{_BASE}/{cid}/add", payload, timeout=180) def query_vectors(query_text, n_results=5, where=None, name=None): """语义检索:返回 [{id, document, distance, metadata}, ...](升序按相似度)""" try: cid, exists = _get_collection_id(name) if not exists: return [] except RuntimeError: return [] vec = embed_texts(query_text)[0] payload = {"query_embeddings": [vec], "n_results": n_results, "include": ["documents", "metadatas", "distances"]} if where: payload["where"] = where _, data = _http("POST", f"{_BASE}/{cid}/query", payload) out = [] for i, doc in enumerate(data.get("documents", [[]])[0]): out.append({ "id": data["ids"][0][i], "document": doc, "distance": data["distances"][0][i], "metadata": data["metadatas"][0][i] if data.get("metadatas") else {}, }) return out def delete_collection(name): """删除集合(用于重灌)""" try: _http("DELETE", f"{_BASE}/{name}") except RuntimeError: pass