""" Few-shot 经验样本的 Chroma 向量库存储与检索。 与 SchemaIndexer 一致:使用项目统一 ``get_embedder()``;默认 Chroma 为内存 (``EphemeralClient``)。离线灌库脚本可设 ``persist_to_disk=True`` 写入磁盘。 集合名默认 ``fewshot_samples``。 """ from __future__ import annotations import json import logging import os from pathlib import Path from typing import Any, Dict, List, Optional, Tuple import chromadb from chromadb.config import Settings as ChromaSettings logger = logging.getLogger(__name__) DEFAULT_FEWSHOT_CHROMA_DIR = "./data/embeddings/chroma_fewshot" COLLECTION_NAME = "fewshot_samples" # Chroma metadata 单值不宜过大,SQL/说明超长时截断 _MAX_SQL_META = 16000 _MAX_EXPLAIN_META = 6000 _MAX_QEN_META = 2000 class FewShotChromaStore: """Few-shot JSONL → Chroma 写入与按向量检索。""" def __init__( self, embedder: Any, persist_dir: Optional[str] = None, collection_name: Optional[str] = None, *, persist_to_disk: bool = False, ): """ Args: embedder: 编码器实例。 persist_dir: 磁盘模式下的持久化目录;内存模式下仍解析为配置引用路径。 collection_name: 集合名;可空,空则读环境变量或默认。 persist_to_disk: 为 True 时使用 ``PersistentClient`` 落盘(如灌库脚本)。 """ self.embedder = embedder # Chroma 的 collection 名;显式参数优先,否则读 FEWSHOT_CHROMA_COLLECTION,再回退默认 explicit = (collection_name or "").strip() from_env = (os.getenv("FEWSHOT_CHROMA_COLLECTION") or "").strip() self.collection_name = explicit or from_env or COLLECTION_NAME self.persist_dir = Path( (persist_dir or os.getenv("FEWSHOT_CHROMA_PATH") or DEFAULT_FEWSHOT_CHROMA_DIR).strip() ) self.persist_to_disk = persist_to_disk if persist_to_disk: self.persist_dir.mkdir(parents=True, exist_ok=True) self.client = chromadb.PersistentClient( path=str(self.persist_dir), settings=ChromaSettings(anonymized_telemetry=False), ) else: self.client = chromadb.EphemeralClient( settings=ChromaSettings(anonymized_telemetry=False), ) self.collection = self.client.get_or_create_collection( name=self.collection_name, metadata={"hnsw:space": "cosine"}, ) backend = f"磁盘 {self.persist_dir}" if persist_to_disk else f"内存(配置路径={self.persist_dir})" logger.info( f"[OK] FewShotChromaStore: {backend} collection={self.collection_name} " f"count={self.collection.count()}" ) def count(self) -> int: return int(self.collection.count()) def clear(self) -> None: """删除集合内全部文档(用于 force rebuild)。""" if self.collection.count() > 0: self.client.delete_collection(self.collection_name) self.collection = self.client.get_or_create_collection( name=self.collection_name, metadata={"hnsw:space": "cosine"}, ) logger.info("[OK] Few-shot Chroma 集合已清空并重建") def build_from_samples( self, samples: List["ExperienceSample"], *, force_rebuild: bool = False, batch_size: int = 32, ) -> int: """ 将已解析的样本列表写入 Chroma(向量由 question_zh 编码)。 Returns: 写入条数 """ if not samples: logger.warning("Few-shot Chroma: 无样本,跳过构建") return 0 if force_rebuild and self.count() > 0: self.clear() if not force_rebuild and self.count() > 0: logger.info("Few-shot Chroma 已有 %s 条,跳过构建(加 --force 可重建)", self.count()) return self.count() texts = [(s.question_zh or "").strip() or " " for s in samples] # Chroma 要求 ids 为唯一字符串 ids = [str(s.qid) for s in samples] metadatas: List[Dict[str, Any]] = [] for s in samples: metadatas.append( { "qid": str(s.qid), "question_en": ((s.question_en or "")[:_MAX_QEN_META]), "sql": (s.sql or "")[:_MAX_SQL_META], "explanation": (s.explanation or "")[:_MAX_EXPLAIN_META], "rating": int(s.rating) if s.rating is not None else -1, "difficulty": (s.difficulty or "medium")[:32], "tags_json": json.dumps(s.tags or [], ensure_ascii=False)[:4000], } ) logger.info("Few-shot Chroma: 计算 %s 条 embedding...", len(texts)) embeddings = self.embedder.encode( texts, batch_size=min(batch_size, len(texts)), normalize=True, show_progress=True, ) self.collection.add( embeddings=embeddings.tolist(), documents=texts, metadatas=metadatas, ids=ids, ) logger.info("[OK] Few-shot Chroma 索引完成: %s 条", len(ids)) return len(ids) def search_raw( self, query: str, top_k: int = 20, ) -> List[Tuple[float, Dict[str, Any], str]]: """ Returns: (score, metadata_dict, document) 列表,score 同 SchemaIndexer 为 1-distance """ n = self.count() if n == 0: return [] q_emb = self.embedder.encode( [(query or "").strip() or " "], batch_size=1, normalize=True, show_progress=False, ) k = min(max(top_k, 1), n) results = self.collection.query( query_embeddings=q_emb.tolist(), n_results=k, ) out: List[Tuple[float, Dict[str, Any], str]] = [] if not results["ids"] or not results["ids"][0]: return out for tid, dist, meta, doc in zip( results["ids"][0], results["distances"][0], results["metadatas"][0], results["documents"][0], ): score = 1.0 - float(dist) m = dict(meta or {}) m["qid"] = tid out.append((score, m, doc or "")) return out def get_all_as_samples(self) -> List[Any]: """导出集合中全部样本(用于统计/按标签列举;条数大时慎用)。""" if self.count() == 0: return [] data = self.collection.get(include=["metadatas", "documents"]) ids = data.get("ids") or [] metas = data.get("metadatas") or [] docs = data.get("documents") or [] out: List[Any] = [] for i, qid in enumerate(ids): meta = dict(metas[i] or {}) if i < len(metas) else {} meta["qid"] = qid doc = docs[i] if i < len(docs) else "" out.append(sample_from_chroma_metadata(meta, doc)) return out def sample_from_chroma_metadata(meta: Dict[str, Any], document: str) -> "ExperienceSample": from utils.fewshot_selector import ExperienceSample tags_raw = meta.get("tags_json") or "[]" try: tags = json.loads(tags_raw) if isinstance(tags_raw, str) else [] except json.JSONDecodeError: tags = [] if not isinstance(tags, list): tags = [] r = meta.get("rating") try: ri = int(r) if r is not None else -1 except (TypeError, ValueError): ri = -1 rating = ri if ri >= 0 else None qen = meta.get("question_en") return ExperienceSample( qid=str(meta.get("qid", "")), question_zh=document or str(meta.get("question_zh", "")), question_en=qen if qen else None, sql=str(meta.get("sql", "")), explanation=str(meta.get("explanation", "")), rating=rating, tags=tags, difficulty=str(meta.get("difficulty", "medium") or "medium"), )