197 lines
6.3 KiB
Python
197 lines
6.3 KiB
Python
from sqlalchemy.ext.asyncio import AsyncSession
|
||
from module_admin.entity.vo.ragflow_vo import RagflowListQueryModel, ListDocumentsQueryModel, UpdateFileModel, DeleteFileModel, CreateDatasetModel, DocumentIdsModel, UpdateChatAssistantModel \
|
||
,CreateSessionWithChatModel, ConverseWithChatAssistantModel
|
||
from utils.ragflow_client_manager import get_ragflow_client
|
||
|
||
class RAGFlowService:
|
||
"""
|
||
RAGFlow服务
|
||
"""
|
||
|
||
# 获取数据集列表
|
||
@classmethod
|
||
async def get_ragflow_dataset_list_services(cls, query_db: AsyncSession, rage_flow_query: RagflowListQueryModel):
|
||
"""
|
||
获取数据集列表
|
||
"""
|
||
|
||
client = await get_ragflow_client()
|
||
result = await client.list_datasets(**(rage_flow_query.model_dump()))
|
||
|
||
# 获取分页数据
|
||
return result
|
||
# 创建数据集
|
||
@classmethod
|
||
async def create_dataset_services(cls, create_dataset_params: CreateDatasetModel):
|
||
"""创建数据集
|
||
|
||
Args:
|
||
create_dataset_params (CreateDatasetModel): 创建参数
|
||
|
||
Returns:
|
||
_type_: _description_
|
||
"""
|
||
client = await get_ragflow_client()
|
||
result = await client.create_dataset(
|
||
**(create_dataset_params.model_dump())
|
||
)
|
||
|
||
return result
|
||
|
||
# 更新数据集
|
||
@classmethod
|
||
async def update_dataset_services(
|
||
cls,
|
||
dataset_id: str,
|
||
update_dataset_params: CreateDatasetModel,
|
||
):
|
||
"""更新数据集信息
|
||
|
||
Args:
|
||
dataset_id (str): 数据集id
|
||
update_dataset_params (CreateDatasetModel): 更新参数
|
||
|
||
Returns:
|
||
_type_: _description_
|
||
"""
|
||
client = await get_ragflow_client()
|
||
result = await client.update_dataset(
|
||
dataset_id=dataset_id, **(update_dataset_params.model_dump())
|
||
)
|
||
|
||
return result
|
||
|
||
|
||
# 获取数据集中文档列表
|
||
@classmethod
|
||
async def list_documents_services(
|
||
cls,
|
||
query_db: AsyncSession,
|
||
dataset_id: str,
|
||
list_documents_query: ListDocumentsQueryModel,
|
||
):
|
||
client = await get_ragflow_client()
|
||
result = await client.list_documents(dataset_id=dataset_id, **(list_documents_query.model_dump()))
|
||
|
||
return result
|
||
|
||
# 上传文档到数据集
|
||
@classmethod
|
||
async def upload_file_dataset_services(
|
||
cls,
|
||
dataset_id: str,
|
||
files,
|
||
):
|
||
client = await get_ragflow_client()
|
||
result = await client.upload_documents_bytes(dataset_id=dataset_id, file_bytes=files)
|
||
|
||
return result
|
||
# 开始解析文档
|
||
@classmethod
|
||
async def parse_documents_services(
|
||
cls,
|
||
dataset_id: str,
|
||
parse_params: DocumentIdsModel,
|
||
):
|
||
client = await get_ragflow_client()
|
||
result = await client.parse_documents(dataset_id=dataset_id, document_ids=parse_params.documnet_ids)
|
||
|
||
return result
|
||
|
||
# 停止解析文档
|
||
@classmethod
|
||
async def stop_parse_documents_services(
|
||
cls,
|
||
dataset_id: str,
|
||
parse_params: DocumentIdsModel,
|
||
):
|
||
client = await get_ragflow_client()
|
||
result = await client.stop_parsing_documents(dataset_id=dataset_id, document_ids=parse_params.documnet_ids)
|
||
|
||
return result
|
||
|
||
|
||
# 更新文档内容
|
||
@classmethod
|
||
async def update_file_dataset_services(
|
||
cls,
|
||
dataset_id: str,
|
||
document_id: str,
|
||
update_params: UpdateFileModel,
|
||
):
|
||
client = await get_ragflow_client()
|
||
result = await client.update_document(dataset_id=dataset_id, document_id=document_id, **(update_params.model_dump()))
|
||
|
||
return result
|
||
|
||
# 删除文档
|
||
@classmethod
|
||
async def delete_file_services(
|
||
cls,
|
||
dataset_id: str,
|
||
delete_params: DeleteFileModel,
|
||
):
|
||
client = await get_ragflow_client()
|
||
result = await client.delete_documents(dataset_id=dataset_id, **(delete_params.model_dump()))
|
||
return result
|
||
|
||
# 删除数据集
|
||
@classmethod
|
||
async def delete_datasets_services(
|
||
cls,
|
||
delete_params: DeleteFileModel,
|
||
):
|
||
client = await get_ragflow_client()
|
||
result = await client.delete_datasets(**(delete_params.model_dump()))
|
||
return result
|
||
|
||
|
||
# 查看聊天助手列表
|
||
@classmethod
|
||
async def get_chat_assistant_list_services(
|
||
cls,
|
||
query_params: RagflowListQueryModel,
|
||
):
|
||
client = await get_ragflow_client()
|
||
result = await client.list_chat_assistants(**(query_params.model_dump()))
|
||
return result
|
||
|
||
# 修改聊天助手
|
||
@classmethod
|
||
async def update_chat_assistant_services(cls, update_params: UpdateChatAssistantModel):
|
||
client = await get_ragflow_client()
|
||
result = await client.update_chat_assistant(**(update_params.model_dump()))
|
||
return result
|
||
|
||
# 创建助手会话
|
||
@classmethod
|
||
async def create_session_with_chat_services(cls, create_params: CreateSessionWithChatModel):
|
||
client = await get_ragflow_client()
|
||
result = await client.create_session_with_chat(**(create_params.model_dump()))
|
||
return result
|
||
|
||
|
||
# 与助手聊天
|
||
@classmethod
|
||
async def converse_with_chat_assistant_services(cls, converse_params: ConverseWithChatAssistantModel):
|
||
client = await get_ragflow_client()
|
||
# 修复:直接返回AsyncGenerator,不使用await消费流式数据
|
||
return client.converse_with_chat_assistant(**(converse_params.model_dump()))
|
||
|
||
# 与助手聊天 (OpenAI Compatible)
|
||
@classmethod
|
||
async def converse_with_chat_assistant_services_openai(cls, converse_params: ConverseWithChatAssistantModel):
|
||
client = await get_ragflow_client()
|
||
# Construct messages list for OpenAI format
|
||
messages = [{"role": "user", "content": converse_params.question}]
|
||
# Uses defaults for model name as per user indication "server will parse this automatically"
|
||
return await client.create_chat_completion(
|
||
chat_id=converse_params.chat_id,
|
||
model="ragflow",
|
||
messages=messages,
|
||
stream=converse_params.stream
|
||
)
|
||
|
||
|
||
|