feat: 文档检索接口
This commit is contained in:
@@ -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))
|
||||
|
||||
+13
-10
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user