feat: milvus基本的增删改查
This commit is contained in:
@@ -10,6 +10,11 @@ PG_DB_USERNAME=xxx # postgres数据库连接用户名
|
|||||||
PG_DB_PASSWORD=xxx # postgres数据库连接密码
|
PG_DB_PASSWORD=xxx # postgres数据库连接密码
|
||||||
PG_DB_DATABASE=xxx # postgres数据库连接的数据库名
|
PG_DB_DATABASE=xxx # postgres数据库连接的数据库名
|
||||||
|
|
||||||
|
MILVUS_URI=http://xxxx:7323 # milvus数据库的连接uri
|
||||||
|
MILVUS_COLLECTION=xxxx # milvus数据库集合名称
|
||||||
|
MILVUS_DATABASE=xxxx # milvus数据库数据库名称
|
||||||
|
MILVUS_DIMENSION=2560 # milvus数据库向量维度
|
||||||
|
|
||||||
LLM_KEY_LOCAL=123
|
LLM_KEY_LOCAL=123
|
||||||
LLM_KEY_HUOSHAN=a0311f2a-ba85-4428-b158-xxxxxxxxxxxx # 火山引擎模型服务平台key
|
LLM_KEY_HUOSHAN=a0311f2a-ba85-4428-b158-xxxxxxxxxxxx # 火山引擎模型服务平台key
|
||||||
LLM_KEY_BAILIAN=sk-51d13ba8ea044d128c66dxxxxxxxxxxxx # 阿里云百炼模型服务平台key
|
LLM_KEY_BAILIAN=sk-51d13ba8ea044d128c66dxxxxxxxxxxxx # 阿里云百炼模型服务平台key
|
||||||
|
|||||||
@@ -0,0 +1,41 @@
|
|||||||
|
from typing import List
|
||||||
|
|
||||||
|
from fastapi.params import Param
|
||||||
|
from llama_index.core import Document
|
||||||
|
|
||||||
|
from app.utils.milvus_utils import milvus_service
|
||||||
|
from app.utils.next_id import next_id
|
||||||
|
|
||||||
|
|
||||||
|
def add_knowledge_route(app):
|
||||||
|
@app.get("/knowledge/search")
|
||||||
|
async def knowledge_search(text: str = Param()):
|
||||||
|
print("text", text)
|
||||||
|
query_result = await milvus_service.async_search(text or "hello")
|
||||||
|
return query_result
|
||||||
|
|
||||||
|
@app.post("/knowledge/embed_text")
|
||||||
|
async def knowledge_embed_text(text_list: List[str]):
|
||||||
|
id_list = await next_id(len(text_list))
|
||||||
|
document_list = [
|
||||||
|
Document(
|
||||||
|
text=text,
|
||||||
|
doc_id=id_list[index],
|
||||||
|
metadata={"kb_id": "abc"}
|
||||||
|
)
|
||||||
|
for index, text in enumerate(text_list)
|
||||||
|
]
|
||||||
|
await milvus_service.async_create_index_from_documents(document_list)
|
||||||
|
search_result = await milvus_service.async_search("hello")
|
||||||
|
return {
|
||||||
|
"result": "嵌入成功:",
|
||||||
|
"origin_documents": [document.to_dict() for document in document_list],
|
||||||
|
"search_documents": search_result,
|
||||||
|
}
|
||||||
|
|
||||||
|
@app.post("/knowledge/delete")
|
||||||
|
async def knowledge_embed_text(body: dict):
|
||||||
|
await milvus_service.async_delete(body.get('id'))
|
||||||
|
return {
|
||||||
|
"result": "删除成功",
|
||||||
|
}
|
||||||
+8
-2
@@ -1,3 +1,4 @@
|
|||||||
|
import asyncio
|
||||||
from contextlib import asynccontextmanager
|
from contextlib import asynccontextmanager
|
||||||
|
|
||||||
from fastapi import FastAPI
|
from fastapi import FastAPI
|
||||||
@@ -8,6 +9,7 @@ from starlette.staticfiles import StaticFiles
|
|||||||
|
|
||||||
from app.middlewares.app_middlewares import add_app_middlewares
|
from app.middlewares.app_middlewares import add_app_middlewares
|
||||||
from app.utils.db_utils import check_database_connection
|
from app.utils.db_utils import check_database_connection
|
||||||
|
from app.utils.milvus_utils import milvus_service
|
||||||
from app.utils.postgres_checkpointer import check_postgres_connection, close_postgres_connection
|
from app.utils.postgres_checkpointer import check_postgres_connection, close_postgres_connection
|
||||||
|
|
||||||
|
|
||||||
@@ -15,8 +17,12 @@ def create_app():
|
|||||||
@asynccontextmanager
|
@asynccontextmanager
|
||||||
async def lifespan(app: FastAPI):
|
async def lifespan(app: FastAPI):
|
||||||
print("lifespan:启动阶段")
|
print("lifespan:启动阶段")
|
||||||
async_engine = await check_database_connection()
|
async_results = await asyncio.gather(
|
||||||
await check_postgres_connection()
|
asyncio.create_task(check_database_connection()),
|
||||||
|
asyncio.create_task(check_postgres_connection()),
|
||||||
|
asyncio.create_task(milvus_service.check_milvus_connection()),
|
||||||
|
)
|
||||||
|
async_engine = async_results[0]
|
||||||
yield
|
yield
|
||||||
print("lifespan:销毁阶段")
|
print("lifespan:销毁阶段")
|
||||||
await async_engine.dispose()
|
await async_engine.dispose()
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ from app.config.env import env
|
|||||||
from app.controller.add_approve_route import add_approve_route
|
from app.controller.add_approve_route import add_approve_route
|
||||||
from app.controller.add_file_route import add_file_route
|
from app.controller.add_file_route import add_file_route
|
||||||
from app.controller.add_hotel_route import add_hotel_route
|
from app.controller.add_hotel_route import add_hotel_route
|
||||||
|
from app.controller.add_knowledge_route import add_knowledge_route
|
||||||
from app.controller.add_langgraph_approve_route import add_langgraph_approve_route
|
from app.controller.add_langgraph_approve_route import add_langgraph_approve_route
|
||||||
from app.controller.add_langgraph_chat_route import add_langgraph_chat_route
|
from app.controller.add_langgraph_chat_route import add_langgraph_chat_route
|
||||||
from app.controller.add_langgraph_route import add_langgraph_route
|
from app.controller.add_langgraph_route import add_langgraph_route
|
||||||
@@ -64,6 +65,7 @@ add_approve_route(app)
|
|||||||
add_reimburse_route(app)
|
add_reimburse_route(app)
|
||||||
add_hotel_route(app)
|
add_hotel_route(app)
|
||||||
add_file_route(app)
|
add_file_route(app)
|
||||||
|
add_knowledge_route(app)
|
||||||
|
|
||||||
|
|
||||||
@app.get("/get_env")
|
@app.get("/get_env")
|
||||||
|
|||||||
@@ -0,0 +1,91 @@
|
|||||||
|
from typing import List, Optional
|
||||||
|
|
||||||
|
from llama_index.core import Document, VectorStoreIndex
|
||||||
|
from llama_index.vector_stores.milvus import MilvusVectorStore
|
||||||
|
|
||||||
|
from app.config.env import env
|
||||||
|
from app.utils.llm_utils import create_embeddings, create_llama_index_llm
|
||||||
|
|
||||||
|
|
||||||
|
class MilvusService:
|
||||||
|
def __init__(self):
|
||||||
|
milvus_vector_store: Optional[MilvusVectorStore] = None
|
||||||
|
self._milvus_vector_store = milvus_vector_store
|
||||||
|
|
||||||
|
self.embeddings = create_embeddings()
|
||||||
|
|
||||||
|
# 获取一个MilvusVectorStore实例,如果不存在就异步创建
|
||||||
|
async def get_vector_store(self):
|
||||||
|
if not self._milvus_vector_store:
|
||||||
|
print("create new_vector_store...")
|
||||||
|
self._milvus_vector_store = MilvusVectorStore(
|
||||||
|
uri=env.milvus_uri,
|
||||||
|
dim=env.milvus_dimension,
|
||||||
|
collection_name=env.milvus_collection,
|
||||||
|
db_name=env.milvus_database,
|
||||||
|
embedding_field="embedding",
|
||||||
|
id_field="id",
|
||||||
|
similarity_metric="COSINE",
|
||||||
|
consistency_level="Strong",
|
||||||
|
overwrite=False,
|
||||||
|
)
|
||||||
|
return self._milvus_vector_store
|
||||||
|
|
||||||
|
# 检查Milvus连接是否正常
|
||||||
|
async def check_milvus_connection(self):
|
||||||
|
retrieve_result = await self.async_search("hello")
|
||||||
|
print("✅ Milvus connection successful:", f"{env.milvus_uri}:{env.milvus_database}/{env.milvus_collection}", f"[{retrieve_result}]")
|
||||||
|
|
||||||
|
async def async_create_index_from_documents(self, documents: List[Document]) -> VectorStoreIndex:
|
||||||
|
vector_store = await self.get_vector_store()
|
||||||
|
"""异步创建向量索引并插入文档"""
|
||||||
|
vector_index: VectorStoreIndex = VectorStoreIndex.from_vector_store(
|
||||||
|
vector_store=vector_store,
|
||||||
|
embed_model=self.embeddings,
|
||||||
|
)
|
||||||
|
from llama_index.core.node_parser import SimpleNodeParser
|
||||||
|
parser = SimpleNodeParser()
|
||||||
|
nodes = parser.get_nodes_from_documents(documents)
|
||||||
|
await vector_index.ainsert_nodes(nodes)
|
||||||
|
return vector_index
|
||||||
|
|
||||||
|
# 检索milvus文档
|
||||||
|
async def async_search(self, query: str, top_k: int = 5) -> List[dict]:
|
||||||
|
"""异步向量搜索"""
|
||||||
|
# 创建查询引擎,设置top_k参数
|
||||||
|
vector_index = VectorStoreIndex.from_vector_store(
|
||||||
|
vector_store=await self.get_vector_store(),
|
||||||
|
embed_model=self.embeddings,
|
||||||
|
)
|
||||||
|
query_engine = vector_index.as_query_engine(
|
||||||
|
streaming=False,
|
||||||
|
similarity_top_k=top_k, # 这里正确使用top_k参数
|
||||||
|
llm=create_llama_index_llm()
|
||||||
|
)
|
||||||
|
|
||||||
|
# 异步执行查询
|
||||||
|
response = await query_engine.aquery(query)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"answer": str(response),
|
||||||
|
"sources": [
|
||||||
|
{
|
||||||
|
"text": node.node.text,
|
||||||
|
"metadata": node.node.metadata,
|
||||||
|
"score": node.score
|
||||||
|
}
|
||||||
|
for node in response.source_nodes
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
async def async_delete(self, doc_id: str):
|
||||||
|
vector_index = VectorStoreIndex.from_vector_store(
|
||||||
|
vector_store=await self.get_vector_store(),
|
||||||
|
embed_model=self.embeddings,
|
||||||
|
)
|
||||||
|
await vector_index.adelete_ref_doc(doc_id)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
# 全局实例
|
||||||
|
milvus_service = MilvusService()
|
||||||
Generated
+2622
-5
File diff suppressed because it is too large
Load Diff
@@ -25,6 +25,10 @@ pyjwt = "^2.10.1"
|
|||||||
python-multipart = "^0.0.20"
|
python-multipart = "^0.0.20"
|
||||||
langgraph = "^0.6.3"
|
langgraph = "^0.6.3"
|
||||||
langgraph-checkpoint-postgres = "^2.0.23"
|
langgraph-checkpoint-postgres = "^2.0.23"
|
||||||
|
llama-index = "0.12.42"
|
||||||
|
llama-index-vector-stores-milvus = "0.8.7"
|
||||||
|
llama-index-llms-openai-like = "0.4.0"
|
||||||
|
pymilvus = "2.6.1"
|
||||||
|
|
||||||
|
|
||||||
[tool.poetry.group.dev.dependencies]
|
[tool.poetry.group.dev.dependencies]
|
||||||
|
|||||||
Reference in New Issue
Block a user