实现将文档上传到指定数据集接口

This commit is contained in:
haotian 2025-09-04 14:42:30 +08:00
parent 19fed3551e
commit 39ed13daf5
3 changed files with 35 additions and 6 deletions

View File

@ -1,5 +1,6 @@
from datetime import datetime
from fastapi import APIRouter, Depends, Request
from typing import List
from fastapi import APIRouter, Depends, Request, UploadFile, File, Form
from pydantic_validation_decorator import ValidateFields
from sqlalchemy.ext.asyncio import AsyncSession
from config.enums import BusinessType
@ -52,6 +53,21 @@ async def list_documents_by_dataset_id(
列出数据集中文档列表
"""
print(list_documents_query)
result = await RAGFlowService.list_documents(None, dataset_id, list_documents_query)
result = await RAGFlowService.list_documents_services(None, dataset_id, list_documents_query)
return ResponseUtil.success(data = result)
# 上传文件到数据集
@ragflowController.post("/upload_file/{dataset_id}")
async def upload_file_dataset(
dataset_id: str,
files: List[UploadFile] = File(...),
# query_db: AsyncSession = Depends(get_db),
):
"""
上传文件到数据集
"""
# print(file)
result = await RAGFlowService.upload_file_dataset_services(None, dataset_id ,files)
return ResponseUtil.success(data = result)

View File

@ -23,7 +23,7 @@ class RAGFlowService:
# 获取数据集中文档列表
@classmethod
async def list_documents(
async def list_documents_services(
cls,
query_db: AsyncSession,
dataset_id: str,
@ -33,3 +33,14 @@ class RAGFlowService:
result = await client.list_documents(dataset_id=dataset_id, **(list_documents_query.model_dump()))
return result.get('data', None)
# 上传文档到数据集
async def upload_file_dataset_services(
cls,
dataset_id: str,
files,
):
async with AsyncRAGFlowClient(RAGFlowConfig.RAGFLOW_BASE_URL, RAGFlowConfig.RAGFLOW_API_KEY) as client:
result = await client.upload_documents_bytes(dataset_id=dataset_id, file_bytes=files)
return result.get('data', None)

View File

@ -4,6 +4,7 @@ import json
from typing import Optional, List, Dict, Any, Union, AsyncGenerator
import os
from pathlib import Path
from urllib.parse import unquote
class RAGFlowError(Exception):
@ -292,13 +293,14 @@ class AsyncRAGFlowClient:
text = await response.text()
raise RAGFlowError(response.status, text)
async def upload_documents_bytes(self, dataset_id: str, file_name, file_bytes: List) -> Dict[str, Any]:
async def upload_documents_bytes(self, dataset_id: str, file_bytes: List) -> Dict[str, Any]:
"""
上传文档到数据集
Args:
dataset_id: 数据集ID
file_paths: 文件路径列表
file_name: 文件名
file_bytes: 文件二进制列表
"""
if not self._session:
await self.create_session()
@ -310,7 +312,7 @@ class AsyncRAGFlowClient:
data = aiohttp.FormData()
for file in file_bytes:
data.add_field('file', file, filename=file_name)
data.add_field('file', file.file.read(),filename=file.filename)
headers = {'Authorization': f'Bearer {self.api_key}'}