diff --git a/app/controller/add_knowledge_route.py b/app/controller/add_knowledge_route.py index 6ba49a4..a0e9862 100644 --- a/app/controller/add_knowledge_route.py +++ b/app/controller/add_knowledge_route.py @@ -7,16 +7,14 @@ from starlette.requests import Request from app.utils.db_utils import AsyncSessionDep from app.utils.knowledge_utils import knowledge_service -from app.utils.milvus_utils import milvus_service +from app.utils.milvus_utils import milvus_service, KnowledgeQueryParam 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/search") + async def knowledge_search(param: KnowledgeQueryParam): + return await milvus_service.async_search(param) + # @app.post("/knowledge/embed_text") # async def knowledge_embed_text(text_list: List[str]): # id_list = await next_id(len(text_list)) diff --git a/app/utils/milvus_utils.py b/app/utils/milvus_utils.py index efca688..4217def 100644 --- a/app/utils/milvus_utils.py +++ b/app/utils/milvus_utils.py @@ -4,12 +4,19 @@ from typing import List, Optional from llama_index.core import Document, VectorStoreIndex from llama_index.core.vector_stores import MetadataFilters, MetadataFilter, FilterOperator from llama_index.vector_stores.milvus import MilvusVectorStore +from pydantic import BaseModel, Field from app.config.env import env -from app.utils.llm_utils import create_embeddings, create_llama_index_llm +from app.utils.llm_utils import create_embeddings from app.utils.nltk_utils import load_nltk +class KnowledgeQueryParam(BaseModel): + question: str = Field(..., description="问题文本") + kb_code: str = Field(..., description="知识库编码") + top_k: int = Field(default=5, description="返回结果数量") + + class MilvusService: def __init__(self): milvus_vector_store: Optional[MilvusVectorStore] = None @@ -65,7 +72,7 @@ class MilvusService: return vector_index # 检索milvus文档 - async def async_search(self, query: str, top_k: int = 5) -> List[dict]: + async def async_search(self, param: KnowledgeQueryParam) -> List[dict]: """异步向量搜索""" # 创建查询引擎,设置top_k参数 vector_index = VectorStoreIndex.from_vector_store( @@ -73,15 +80,11 @@ class MilvusService: embed_model=self.embeddings, ) - filters = MetadataFilters(filters=[MetadataFilter(key="kb_id", value="abc", operator=FilterOperator.EQ)]) + filters = MetadataFilters(filters=[MetadataFilter(key="parent_code", value=param.kb_code, operator=FilterOperator.EQ)]) - # query_engine = vector_index.as_query_engine( - # streaming=False, - # similarity_top_k=top_k, # 这里正确使用top_k参数 - # llm=create_llama_index_llm() - # ) - retriever = vector_index.as_retriever(filters=filters, similarity_top_k=5) - result_nodes = await retriever.aretrieve(query) + retriever = vector_index.as_retriever(filters=filters, similarity_top_k=param.top_k) + + result_nodes = await retriever.aretrieve(param.question) # 异步执行查询 # response = await query_engine.aquery(query)