import os import time import psycopg2 import requests from pathlib import Path from typing import List from tenacity import retry, stop_after_attempt, wait_exponential, retry_if_exception_type from llama_index.core import Settings, SimpleDirectoryReader, VectorStoreIndex, StorageContext, Document from llama_index.core.embeddings import BaseEmbedding from llama_index.core.node_parser import SentenceSplitter from llama_index.llms.openai_like import OpenAILike from llama_index.vector_stores.postgres import PGVectorStore from llama_index.core.retrievers import VectorIndexRetriever from llama_index.core.query_engine import RetrieverQueryEngine from llama_index.core.vector_stores import MetadataFilter, MetadataFilters, FilterOperator
class LocalLlamaServerEmbedding(BaseEmbedding): api_base: str embed_model_name: str api_key: str = "dummy"
def _get_embedding(self, text: str) -> List[float]: max_text_len = 2048 text = text[:max_text_len] url = f"{self.api_base}/embeddings" payload = { "input": text, "model": self.embed_model_name } headers = {"Authorization": f"Bearer {self.api_key}"} try: resp = requests.post(url, json=payload, headers=headers, timeout=30) resp.raise_for_status() except requests.exceptions.RequestException as e: raise RuntimeError(f"Embedding服务调用失败: {e}") return resp.json()["data"][0]["embedding"]
def _get_text_embedding(self, text: str) -> List[float]: return self._get_embedding(text)
def _get_query_embedding(self, query: str) -> List[float]: return self._get_embedding(query)
async def _aget_query_embedding(self, query: str) -> List[float]: return self._get_embedding(query)
async def _aget_text_embedding(self, text: str) -> List[float]: return self._get_embedding(text)
class RAGService: def __init__( self, db_config: dict, embed_api_base: str, llm_base_url: str, llm_model: str, embed_model_name: str, chunk_size: int = 512, chunk_overlap: int = 50, batch_size: int = 10 ): self.db_config = db_config self.embed_api_base = embed_api_base self.llm_base_url = llm_base_url self.llm_model = llm_model self.embed_model_name = embed_model_name self.chunk_size = chunk_size self.chunk_overlap = chunk_overlap self.batch_size = batch_size self._vector_store = None self.splitter = SentenceSplitter(chunk_size=self.chunk_size, chunk_overlap=self.chunk_overlap) self._init_settings()
def _init_settings(self): """初始化LLM与Embedding全局配置""" os.environ["OPENAI_API_KEY"] = "dummy" os.environ["OPENAI_BASE_URL"] = self.llm_base_url llm = OpenAILike( model=self.llm_model, api_base=os.environ["OPENAI_BASE_URL"], api_key=os.environ["OPENAI_API_KEY"], is_chat_model=True, context_window=4096, temperature=0.1 ) Settings.llm = llm Settings.embed_model = LocalLlamaServerEmbedding( api_base=self.embed_api_base, embed_model_name=self.embed_model_name ) print("[+] LLM与Embedding模型初始化完成")
def _get_vector_store(self) -> PGVectorStore: """单例获取PGVectorStore""" if self._vector_store is None: print("[+] 初始化PGVectorStore连接") self._vector_store = PGVectorStore.from_params( database=self.db_config["database"], host=self.db_config["host"], password=self.db_config["password"], port=self.db_config["port"], user=self.db_config["user"], table_name=self.db_config["table_name"], embed_dim=self.db_config["embed_dim"], hnsw_kwargs={ "hnsw_m": 16, "hnsw_ef_construction": 64, "hnsw_ef_search": 40, "hnsw_dist_method": "vector_cosine_ops", }, ) return self._vector_store
def load_index_from_pg(self) -> VectorStoreIndex: """从PG加载已有索引""" vector_store = self._get_vector_store() storage_context = StorageContext.from_defaults(vector_store=vector_store) index = VectorStoreIndex.from_vector_store( vector_store, storage_context=storage_context ) return index
def add_or_update_knowledge(self, docs: List[Document]) -> VectorStoreIndex: """增量新增/更新文档:存在则删除旧chunk,再写入新文档""" vector_store = self._get_vector_store() doc_ids = [doc.metadata["doc_id"] for doc in docs] print(f"待处理文档doc_ids: {doc_ids}") filters = MetadataFilters( filters=[ MetadataFilter( key="doc_id", value=doc_ids, operator=FilterOperator.IN ) ] ) exist_nodes = vector_store.get_nodes(filters=filters) exist_doc_ids = {n.metadata["doc_id"] for n in exist_nodes} print(f"数据库中已存在的doc_ids: {exist_doc_ids}") new_docs = [] update_doc_ids = [] for d in docs: if d.metadata["doc_id"] in exist_doc_ids: update_doc_ids.append(d.metadata["doc_id"]) else: new_docs.append(d) if update_doc_ids: print(f"删除旧文档向量,doc_ids={update_doc_ids}") del_filters = MetadataFilters( filters=[ MetadataFilter( key="doc_id", value=update_doc_ids, operator=FilterOperator.IN ) ] ) vector_store.delete_nodes(filters=del_filters) if len(docs) > 0: storage_context = StorageContext.from_defaults(vector_store=vector_store) index = VectorStoreIndex.from_documents( docs, storage_context=storage_context, transformations=[self.splitter], show_progress=True ) print("[+] 知识库写入完成") return index else: print("[-] 没有待处理文档") return self.load_index_from_pg()
def delete_knowledge(self, doc_id: str): """根据doc_id删除文档全部向量片段""" vector_store = self._get_vector_store() del_filters = MetadataFilters( filters=[ MetadataFilter(key="doc_id", value=doc_id, operator=FilterOperator.EQ) ] ) vector_store.delete_nodes(filters=del_filters) print(f"[+] 已删除 doc_id={doc_id} 的所有向量片段")
def clear_all_vector(self) -> None: """清空整张向量表,如果表不存在则直接跳过""" vector_store = self._get_vector_store() table_name = vector_store.table_name print(f"[-] 准备清空向量表 [{table_name}] 全部数据") try: conn = psycopg2.connect( database=self.db_config["database"], host=self.db_config["host"], password=self.db_config["password"], port=self.db_config["port"], user=self.db_config["user"] ) cur = conn.cursor() cur.execute(""" SELECT EXISTS ( SELECT FROM information_schema.tables WHERE table_name = %s ); """, (table_name,)) exists = cur.fetchone()[0] if exists: cur.execute(f"TRUNCATE TABLE {table_name};") conn.commit() print(f"[+] 向量表 {table_name} 已全部清空") else: print(f"[*] 表 {table_name} 不存在,无需清空") cur.close() conn.close() except Exception as e: print(f"清空向量表失败: {str(e)}") raise
@retry( stop=stop_after_attempt(3), wait=wait_exponential(multiplier=1, min=1, max=5), retry=retry_if_exception_type((psycopg2.OperationalError, requests.exceptions.RequestException, RuntimeError)) ) def rag_query(self, query_str: str, filter_meta: dict = None, top_k: int = 3): """RAG问答查询,支持元数据过滤,带重试""" start_time = time.time() index = self.load_index_from_pg() filters = None if filter_meta: filter_list = [] for k, v in filter_meta.items(): if isinstance(v, list): op = FilterOperator.IN else: op = FilterOperator.EQ filter_list.append(MetadataFilter(key=k, value=v, operator=op)) filters = MetadataFilters(filters=filter_list) retriever = VectorIndexRetriever( index=index, similarity_top_k=top_k, filters=filters ) query_engine = RetrieverQueryEngine.from_args(retriever) response = query_engine.query(query_str) cost = time.time() - start_time print(f"Query: {query_str}, cost={cost:.2f}s, hit_chunk_count={len(response.source_nodes)}") return response
DB_CONFIG = { "database": "storage_db", "host": "8.122.231.178", "password": "1233", "port": "5432", "user": "storage_user", "table_name": "llama_rag_vector", "embed_dim": 1024 }
EMBEDDING_API_BASE = "http://127.0.0.1:11434/v1" LLM_BASE_URL = "http://127.0.0.1:11433/v1"
LLM_MODEL = "qwen2.5-1.5b-instruct-q4_k_m.gguf" EMBED_MODEL_NAME = "Qwen3-Embedding-0.6B-Q8_0.gguf"
CHUNK_SIZE = 512 CHUNK_OVERLAP = 50 BATCH_SIZE = 10
if __name__ == "__main__": rag_service = RAGService( db_config=DB_CONFIG, embed_api_base=EMBEDDING_API_BASE, llm_base_url=LLM_BASE_URL, llm_model=LLM_MODEL, embed_model_name=EMBED_MODEL_NAME, chunk_size=CHUNK_SIZE, chunk_overlap=CHUNK_OVERLAP, batch_size=BATCH_SIZE ) rag_service.clear_all_vector() docs = SimpleDirectoryReader( "./data/", required_exts=[".pdf", ".docx", ".txt"] ).load_data()
file_to_docid = {} for doc in docs: fname = Path(doc.metadata["file_path"]).name if fname not in file_to_docid: file_to_docid[fname] = f"file_{len(file_to_docid)}" doc.metadata["doc_id"] = file_to_docid[fname] doc.metadata["source"] = "./data/" doc.metadata["upload_time"] = time.strftime("%Y-%m-%d %H:%M:%S")
index = rag_service.add_or_update_knowledge(docs) print(f"[+] 文档 {len(docs)} 条已成功入库") print(index)
resp = rag_service.rag_query("概括文档内容,并返回中文。", filter_meta={"source": "./data/"}, top_k=1) print("---- LLM回答 ----") print(resp.response)
rag_service.delete_knowledge("file_0")
print("---- 检索到的源片段 ----") for node in resp.source_nodes: print(f"相似度分数:{node.score:.4f}") print(f"元数据:{node.metadata}")
|