1044 lines
38 KiB
Python
1044 lines
38 KiB
Python
import aiohttp
|
||
import asyncio
|
||
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):
|
||
"""RAGFlow API错误异常"""
|
||
def __init__(self, code: int, message: str):
|
||
self.code = code
|
||
self.message = message
|
||
super().__init__(f"Error {code}: {message}")
|
||
|
||
|
||
class AsyncRAGFlowClient:
|
||
"""异步RAGFlow API客户端"""
|
||
|
||
def __init__(self, base_url: str, api_key: str, timeout: int = 30):
|
||
"""
|
||
初始化RAGFlow客户端
|
||
|
||
Args:
|
||
base_url: RAGFlow服务器地址
|
||
api_key: API密钥
|
||
timeout: 请求超时时间(秒)
|
||
"""
|
||
try:
|
||
self.base_url = base_url.rstrip('/')
|
||
except:
|
||
self.base_url = base_url
|
||
self.api_key = api_key
|
||
self.timeout = timeout
|
||
self.headers = {
|
||
'Authorization': f'Bearer {api_key}',
|
||
'Content-Type': 'application/json'
|
||
}
|
||
self._session = None
|
||
|
||
async def __aenter__(self):
|
||
"""异步上下文管理器入口"""
|
||
await self.create_session()
|
||
return self
|
||
|
||
async def __aexit__(self, exc_type, exc_val, exc_tb):
|
||
"""异步上下文管理器出口"""
|
||
await self.close_session()
|
||
|
||
async def create_session(self):
|
||
"""创建aiohttp会话"""
|
||
if self._session is None:
|
||
timeout = aiohttp.ClientTimeout(total=self.timeout)
|
||
self._session = aiohttp.ClientSession(timeout=timeout)
|
||
|
||
async def close_session(self):
|
||
"""关闭aiohttp会话"""
|
||
if self._session:
|
||
await self._session.close()
|
||
self._session = None
|
||
|
||
async def _request(self, method: str, endpoint: str, **kwargs) -> Dict[str, Any]:
|
||
"""发送HTTP请求"""
|
||
if not self._session:
|
||
await self.create_session()
|
||
|
||
url = f"{self.base_url}{endpoint}"
|
||
|
||
# 处理headers
|
||
headers = kwargs.pop('headers', self.headers.copy())
|
||
|
||
async with self._session.request(method, url, headers=headers, **kwargs) as response:
|
||
try:
|
||
result = await response.json()
|
||
except (aiohttp.ContentTypeError, json.JSONDecodeError):
|
||
if response.status == 200:
|
||
content = await response.read()
|
||
return {'code': 0, 'data': content}
|
||
else:
|
||
text = await response.text()
|
||
raise RAGFlowError(response.status, text)
|
||
|
||
if result.get('code', 0) != 0:
|
||
raise RAGFlowError(result.get('code'), result.get('message', 'Unknown error'))
|
||
|
||
return result
|
||
|
||
async def _stream_request(self, method: str, endpoint: str, **kwargs) -> AsyncGenerator[Dict[str, Any], None]:
|
||
"""发送流式HTTP请求"""
|
||
if not self._session:
|
||
await self.create_session()
|
||
|
||
url = f"{self.base_url}{endpoint}"
|
||
headers = kwargs.pop('headers', self.headers.copy())
|
||
|
||
async with self._session.request(method, url, headers=headers, **kwargs) as response:
|
||
async for line in response.content:
|
||
if line:
|
||
line_str = line.decode('utf-8').strip()
|
||
if line_str.startswith('data:'):
|
||
try:
|
||
data = json.loads(line_str[5:].strip())
|
||
yield data
|
||
except json.JSONDecodeError:
|
||
continue
|
||
|
||
# ====================
|
||
# OpenAI兼容API
|
||
# ====================
|
||
|
||
async def create_chat_completion(self, chat_id: str, model: str, messages: List[Dict[str, str]],
|
||
stream: bool = False) -> Union[Dict[str, Any], AsyncGenerator[Dict[str, Any], None]]:
|
||
"""
|
||
创建聊天完成
|
||
|
||
Args:
|
||
chat_id: 聊天ID
|
||
model: 模型名称
|
||
messages: 消息列表
|
||
stream: 是否流式返回
|
||
"""
|
||
endpoint = f"/api/v1/chats_openai/{chat_id}/chat/completions"
|
||
data = {
|
||
"model": model,
|
||
"messages": messages,
|
||
"stream": stream
|
||
}
|
||
|
||
if stream:
|
||
return self._stream_request('POST', endpoint, json=data)
|
||
else:
|
||
return await self._request('POST', endpoint, json=data)
|
||
|
||
async def create_agent_completion(self, agent_id: str, model: str, messages: List[Dict[str, str]],
|
||
stream: bool = False) -> Union[Dict[str, Any], AsyncGenerator[Dict[str, Any], None]]:
|
||
"""
|
||
创建代理完成
|
||
|
||
Args:
|
||
agent_id: 代理ID
|
||
model: 模型名称
|
||
messages: 消息列表
|
||
stream: 是否流式返回
|
||
"""
|
||
endpoint = f"/api/v1/agents_openai/{agent_id}/chat/completions"
|
||
data = {
|
||
"model": model,
|
||
"messages": messages,
|
||
"stream": stream
|
||
}
|
||
|
||
if stream:
|
||
return self._stream_request('POST', endpoint, json=data)
|
||
else:
|
||
return await self._request('POST', endpoint, json=data)
|
||
|
||
# ====================
|
||
# 数据集管理
|
||
# ====================
|
||
|
||
async def create_dataset(self, name: str, avatar: Optional[str] = None, description: Optional[str] = None,
|
||
embedding_model: Optional[str] = None, permission: str = "me",
|
||
chunk_method: str = "naive",
|
||
parser_config: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
|
||
"""
|
||
创建数据集
|
||
|
||
Args:
|
||
name: 数据集名称
|
||
avatar: Base64编码的头像
|
||
description: 描述
|
||
embedding_model: 嵌入模型
|
||
permission: 权限设置 ("me" 或 "team")
|
||
chunk_method: 分块方法
|
||
parser_config: 解析器配置
|
||
"""
|
||
endpoint = "/api/v1/datasets"
|
||
data = {
|
||
"name": name,
|
||
"permission": permission,
|
||
"chunk_method": chunk_method,
|
||
}
|
||
|
||
if avatar:
|
||
data["avatar"] = avatar
|
||
if description:
|
||
data["description"] = description
|
||
if embedding_model:
|
||
data["embedding_model"] = embedding_model
|
||
if parser_config:
|
||
data["parser_config"] = parser_config
|
||
|
||
return await self._request('POST', endpoint, json=data)
|
||
|
||
async def delete_datasets(self, ids: Optional[List[str]] = None) -> Dict[str, Any]:
|
||
"""
|
||
删除数据集
|
||
|
||
Args:
|
||
ids: 要删除的数据集ID列表,None表示删除所有
|
||
"""
|
||
endpoint = "/api/v1/datasets"
|
||
data = {"ids": ids}
|
||
return await self._request('DELETE', endpoint, json=data)
|
||
|
||
async def update_dataset(self, dataset_id: str, name: Optional[str] = None,
|
||
avatar: Optional[str] = None, description: Optional[str] = None,
|
||
embedding_model: Optional[str] = None, permission: Optional[str] = None,
|
||
chunk_method: Optional[str] = None, pagerank: Optional[int] = None,
|
||
parser_config: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
|
||
"""
|
||
更新数据集
|
||
"""
|
||
endpoint = f"/api/v1/datasets/{dataset_id}"
|
||
data = {}
|
||
|
||
if name is not None:
|
||
data["name"] = name
|
||
if avatar is not None:
|
||
data["avatar"] = avatar
|
||
if description is not None:
|
||
data["description"] = description
|
||
if embedding_model is not None:
|
||
data["embedding_model"] = embedding_model
|
||
if permission is not None:
|
||
data["permission"] = permission
|
||
if chunk_method is not None:
|
||
data["chunk_method"] = chunk_method
|
||
if pagerank is not None:
|
||
data["pagerank"] = pagerank
|
||
if parser_config is not None:
|
||
data["parser_config"] = parser_config
|
||
|
||
return await self._request('PUT', endpoint, json=data)
|
||
|
||
async def list_datasets(self, page: int = 1, page_size: int = 30, orderby: str = "create_time",
|
||
desc: str = "true", name: Optional[str] = None,
|
||
dataset_id: Optional[str] = None) -> Dict[str, Any]:
|
||
"""
|
||
列出数据集
|
||
"""
|
||
endpoint = "/api/v1/datasets"
|
||
params = {
|
||
"page": page,
|
||
"page_size": page_size,
|
||
"orderby": orderby,
|
||
"desc": desc
|
||
}
|
||
|
||
if name:
|
||
params["name"] = name
|
||
if dataset_id:
|
||
params["id"] = dataset_id
|
||
|
||
return await self._request('GET', endpoint, params=params)
|
||
|
||
# ====================
|
||
# 文档管理
|
||
# ====================
|
||
|
||
async def upload_documents(self, dataset_id: str, file_paths: List[str]) -> Dict[str, Any]:
|
||
"""
|
||
上传文档到数据集
|
||
|
||
Args:
|
||
dataset_id: 数据集ID
|
||
file_paths: 文件路径列表
|
||
"""
|
||
if not self._session:
|
||
await self.create_session()
|
||
|
||
endpoint = f"/api/v1/datasets/{dataset_id}/documents"
|
||
url = f"{self.base_url}{endpoint}"
|
||
|
||
# 准备multipart数据
|
||
data = aiohttp.FormData()
|
||
for file_path in file_paths:
|
||
if os.path.exists(file_path):
|
||
file_name = Path(file_path).name
|
||
with open(file_path, 'rb') as f:
|
||
data.add_field('file', f.read(), filename=file_name)
|
||
|
||
headers = {'Authorization': f'Bearer {self.api_key}'}
|
||
|
||
async with self._session.post(url, headers=headers, data=data) as response:
|
||
try:
|
||
result = await response.json()
|
||
if result.get('code', 0) != 0:
|
||
raise RAGFlowError(result.get('code'), result.get('message'))
|
||
return result
|
||
except (aiohttp.ContentTypeError, json.JSONDecodeError):
|
||
text = await response.text()
|
||
raise RAGFlowError(response.status, text)
|
||
|
||
async def upload_documents_bytes(self, dataset_id: str, file_bytes: List) -> Dict[str, Any]:
|
||
"""
|
||
上传文档到数据集
|
||
|
||
Args:
|
||
dataset_id: 数据集ID
|
||
file_name: 文件名
|
||
file_bytes: 文件二进制列表
|
||
"""
|
||
if not self._session:
|
||
await self.create_session()
|
||
|
||
endpoint = f"/api/v1/datasets/{dataset_id}/documents"
|
||
url = f"{self.base_url}{endpoint}"
|
||
|
||
# 准备multipart数据
|
||
data = aiohttp.FormData()
|
||
for file in file_bytes:
|
||
|
||
data.add_field('file', file.file.read(),filename=file.filename)
|
||
|
||
headers = {'Authorization': f'Bearer {self.api_key}'}
|
||
|
||
async with self._session.post(url, headers=headers, data=data) as response:
|
||
try:
|
||
result = await response.json()
|
||
if result.get('code', 0) != 0:
|
||
raise RAGFlowError(result.get('code'), result.get('message'))
|
||
return result
|
||
except (aiohttp.ContentTypeError, json.JSONDecodeError):
|
||
text = await response.text()
|
||
raise RAGFlowError(response.status, text)
|
||
|
||
async def update_document(self, dataset_id: str, document_id: str, name: Optional[str] = None,
|
||
meta_fields: Optional[Dict[str, Any]] = None,
|
||
chunk_method: Optional[str] = None,
|
||
parser_config: Optional[Dict[str, Any]] = None,
|
||
) -> Dict[str, Any]:
|
||
"""
|
||
更新文档配置
|
||
"""
|
||
endpoint = f"/api/v1/datasets/{dataset_id}/documents/{document_id}"
|
||
data = {}
|
||
|
||
if name is not None:
|
||
data["name"] = name
|
||
if meta_fields is not None:
|
||
data["meta_fields"] = meta_fields
|
||
if chunk_method is not None:
|
||
data["chunk_method"] = chunk_method
|
||
if parser_config is not None:
|
||
data["parser_config"] = parser_config
|
||
|
||
|
||
return await self._request('PUT', endpoint, json=data)
|
||
|
||
async def download_document(self, dataset_id: str, document_id: str, save_path: str) -> None:
|
||
"""
|
||
下载文档
|
||
"""
|
||
if not self._session:
|
||
await self.create_session()
|
||
|
||
endpoint = f"/api/v1/datasets/{dataset_id}/documents/{document_id}"
|
||
url = f"{self.base_url}{endpoint}"
|
||
headers = {'Authorization': f'Bearer {self.api_key}'}
|
||
|
||
async with self._session.get(url, headers=headers) as response:
|
||
if response.status == 200:
|
||
content = await response.read()
|
||
with open(save_path, 'wb') as f:
|
||
f.write(content)
|
||
else:
|
||
try:
|
||
error = await response.json()
|
||
raise RAGFlowError(error.get('code'), error.get('message'))
|
||
except (aiohttp.ContentTypeError, json.JSONDecodeError):
|
||
text = await response.text()
|
||
raise RAGFlowError(response.status, text)
|
||
|
||
async def list_documents(self, dataset_id: str, page: int = 1, page_size: int = 30,
|
||
orderby: str = "create_time", desc: str = "true",
|
||
keywords: Optional[str] = None, document_id: Optional[str] = None,
|
||
document_name: Optional[str] = None) -> Dict[str, Any]:
|
||
"""
|
||
列出数据集中的文档
|
||
"""
|
||
endpoint = f"/api/v1/datasets/{dataset_id}/documents"
|
||
params = {
|
||
"page": page,
|
||
"page_size": page_size,
|
||
"orderby": orderby,
|
||
"desc": desc
|
||
}
|
||
|
||
if keywords:
|
||
params["keywords"] = keywords
|
||
if document_id:
|
||
params["id"] = document_id
|
||
if document_name:
|
||
params["name"] = document_name
|
||
|
||
return await self._request('GET', endpoint, params=params)
|
||
|
||
async def delete_documents(self, dataset_id: str, ids: Optional[List[str]] = None) -> Dict[str, Any]:
|
||
"""
|
||
删除文档
|
||
"""
|
||
endpoint = f"/api/v1/datasets/{dataset_id}/documents"
|
||
data = {"ids": ids} if ids else {}
|
||
return await self._request('DELETE', endpoint, json=data)
|
||
|
||
async def parse_documents(self, dataset_id: str, document_ids: List[str]) -> Dict[str, Any]:
|
||
"""
|
||
解析文档
|
||
"""
|
||
endpoint = f"/api/v1/datasets/{dataset_id}/chunks"
|
||
data = {"document_ids": document_ids}
|
||
return await self._request('POST', endpoint, json=data)
|
||
|
||
async def stop_parsing_documents(self, dataset_id: str, document_ids: List[str]) -> Dict[str, Any]:
|
||
"""
|
||
停止解析文档
|
||
"""
|
||
endpoint = f"/api/v1/datasets/{dataset_id}/chunks"
|
||
data = {"document_ids": document_ids}
|
||
return await self._request('DELETE', endpoint, json=data)
|
||
|
||
# ====================
|
||
# 分块管理
|
||
# ====================
|
||
|
||
async def add_chunk(self, dataset_id: str, document_id: str, content: str,
|
||
important_keywords: Optional[List[str]] = None,
|
||
questions: Optional[List[str]] = None) -> Dict[str, Any]:
|
||
"""
|
||
添加分块
|
||
"""
|
||
endpoint = f"/api/v1/datasets/{dataset_id}/documents/{document_id}/chunks"
|
||
data = {"content": content}
|
||
|
||
if important_keywords:
|
||
data["important_keywords"] = important_keywords
|
||
if questions:
|
||
data["questions"] = questions
|
||
|
||
return await self._request('POST', endpoint, json=data)
|
||
|
||
async def list_chunks(self, dataset_id: str, document_id: str, keywords: Optional[str] = None,
|
||
page: int = 1, page_size: int = 1024,
|
||
chunk_id: Optional[str] = None) -> Dict[str, Any]:
|
||
"""
|
||
列出分块
|
||
"""
|
||
endpoint = f"/api/v1/datasets/{dataset_id}/documents/{document_id}/chunks"
|
||
params = {"page": page, "page_size": page_size}
|
||
|
||
if keywords:
|
||
params["keywords"] = keywords
|
||
if chunk_id:
|
||
params["id"] = chunk_id
|
||
|
||
return await self._request('GET', endpoint, params=params)
|
||
|
||
async def delete_chunks(self, dataset_id: str, document_id: str,
|
||
chunk_ids: Optional[List[str]] = None) -> Dict[str, Any]:
|
||
"""
|
||
删除分块
|
||
"""
|
||
endpoint = f"/api/v1/datasets/{dataset_id}/documents/{document_id}/chunks"
|
||
data = {"chunk_ids": chunk_ids} if chunk_ids else {}
|
||
return await self._request('DELETE', endpoint, json=data)
|
||
|
||
async def update_chunk(self, dataset_id: str, document_id: str, chunk_id: str,
|
||
content: Optional[str] = None,
|
||
important_keywords: Optional[List[str]] = None,
|
||
available: Optional[bool] = None) -> Dict[str, Any]:
|
||
"""
|
||
更新分块
|
||
"""
|
||
endpoint = f"/api/v1/datasets/{dataset_id}/documents/{document_id}/chunks/{chunk_id}"
|
||
data = {}
|
||
|
||
if content is not None:
|
||
data["content"] = content
|
||
if important_keywords is not None:
|
||
data["important_keywords"] = important_keywords
|
||
if available is not None:
|
||
data["available"] = available
|
||
|
||
return await self._request('PUT', endpoint, json=data)
|
||
|
||
async def retrieve_chunks(self, question: str, dataset_ids: Optional[List[str]] = None,
|
||
document_ids: Optional[List[str]] = None, page: int = 1,
|
||
page_size: int = 30, similarity_threshold: float = 0.2,
|
||
vector_similarity_weight: float = 0.3, top_k: int = 1024,
|
||
rerank_id: Optional[str] = None, keyword: bool = False,
|
||
highlight: bool = False) -> Dict[str, Any]:
|
||
"""
|
||
检索分块
|
||
"""
|
||
endpoint = "/api/v1/retrieval"
|
||
data = {
|
||
"question": question,
|
||
"page": page,
|
||
"page_size": page_size,
|
||
"similarity_threshold": similarity_threshold,
|
||
"vector_similarity_weight": vector_similarity_weight,
|
||
"top_k": top_k,
|
||
"keyword": keyword,
|
||
"highlight": highlight
|
||
}
|
||
|
||
if dataset_ids:
|
||
data["dataset_ids"] = dataset_ids
|
||
if document_ids:
|
||
data["document_ids"] = document_ids
|
||
if rerank_id:
|
||
data["rerank_id"] = rerank_id
|
||
|
||
return await self._request('POST', endpoint, json=data)
|
||
|
||
# ====================
|
||
# 聊天助手管理
|
||
# ====================
|
||
|
||
async def create_chat_assistant(self, name: str, avatar: Optional[str] = None,
|
||
dataset_ids: Optional[List[str]] = None,
|
||
llm: Optional[Dict[str, Any]] = None,
|
||
prompt: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
|
||
"""
|
||
创建聊天助手
|
||
"""
|
||
endpoint = "/api/v1/chats"
|
||
data = {"name": name}
|
||
|
||
if avatar:
|
||
data["avatar"] = avatar
|
||
if dataset_ids:
|
||
data["dataset_ids"] = dataset_ids
|
||
if llm:
|
||
data["llm"] = llm
|
||
if prompt:
|
||
data["prompt"] = prompt
|
||
|
||
return await self._request('POST', endpoint, json=data)
|
||
|
||
async def update_chat_assistant(self, chat_id: str, name: Optional[str] = None,
|
||
avatar: Optional[str] = None,
|
||
dataset_ids: Optional[List[str]] = None,
|
||
llm: Optional[Dict[str, Any]] = None,
|
||
prompt: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
|
||
"""
|
||
更新聊天助手
|
||
"""
|
||
endpoint = f"/api/v1/chats/{chat_id}"
|
||
data = {}
|
||
|
||
if name is not None:
|
||
data["name"] = name
|
||
if avatar is not None:
|
||
data["avatar"] = avatar
|
||
if dataset_ids is not None:
|
||
data["dataset_ids"] = dataset_ids
|
||
if llm is not None:
|
||
data["llm"] = llm
|
||
if prompt is not None:
|
||
data["prompt"] = prompt
|
||
|
||
return await self._request('PUT', endpoint, json=data)
|
||
|
||
async def delete_chat_assistants(self, ids: Optional[List[str]] = None) -> Dict[str, Any]:
|
||
"""
|
||
删除聊天助手
|
||
"""
|
||
endpoint = "/api/v1/chats"
|
||
data = {"ids": ids} if ids else {}
|
||
return await self._request('DELETE', endpoint, json=data)
|
||
|
||
async def list_chat_assistants(self, page: int = 1, page_size: int = 30,
|
||
orderby: str = "create_time", desc: str = "true",
|
||
name: Optional[str] = None,
|
||
chat_id: Optional[str] = None) -> Dict[str, Any]:
|
||
"""
|
||
列出聊天助手
|
||
"""
|
||
endpoint = "/api/v1/chats"
|
||
params = {
|
||
"page": page,
|
||
"page_size": page_size,
|
||
"orderby": orderby,
|
||
"desc": desc
|
||
}
|
||
|
||
if name:
|
||
params["name"] = name
|
||
if chat_id:
|
||
params["id"] = chat_id
|
||
|
||
return await self._request('GET', endpoint, params=params)
|
||
|
||
# ====================
|
||
# 会话管理
|
||
# ====================
|
||
|
||
async def create_session_with_chat(self, chat_id: str, name: str,
|
||
user_id: Optional[str] = None) -> Dict[str, Any]:
|
||
"""
|
||
创建与聊天助手的会话
|
||
"""
|
||
endpoint = f"/api/v1/chats/{chat_id}/sessions"
|
||
data = {"name": name}
|
||
|
||
if user_id:
|
||
data["user_id"] = user_id
|
||
|
||
return await self._request('POST', endpoint, json=data)
|
||
|
||
async def update_chat_session(self, chat_id: str, session_id: str, name: Optional[str] = None,
|
||
user_id: Optional[str] = None) -> Dict[str, Any]:
|
||
"""
|
||
更新聊天会话
|
||
"""
|
||
endpoint = f"/api/v1/chats/{chat_id}/sessions/{session_id}"
|
||
data = {}
|
||
|
||
if name is not None:
|
||
data["name"] = name
|
||
if user_id is not None:
|
||
data["user_id"] = user_id
|
||
|
||
return await self._request('PUT', endpoint, json=data)
|
||
|
||
async def list_chat_sessions(self, chat_id: str, page: int = 1, page_size: int = 30,
|
||
orderby: str = "create_time", desc: str = "true",
|
||
name: Optional[str] = None, session_id: Optional[str] = None,
|
||
user_id: Optional[str] = None) -> Dict[str, Any]:
|
||
"""
|
||
列出与指定聊天助手相关的聊天会话
|
||
"""
|
||
endpoint = f"/api/v1/chats/{chat_id}/sessions"
|
||
params = {
|
||
"page": page,
|
||
"page_size": page_size,
|
||
"orderby": orderby,
|
||
"desc": desc
|
||
}
|
||
|
||
if name:
|
||
params["name"] = name
|
||
if session_id:
|
||
params["id"] = session_id
|
||
if user_id:
|
||
params["user_id"] = user_id
|
||
|
||
return await self._request('GET', endpoint, params=params)
|
||
|
||
async def delete_chat_sessions(self, chat_id: str, ids: Optional[List[str]] = None) -> Dict[str, Any]:
|
||
"""
|
||
删除聊天会话
|
||
"""
|
||
endpoint = f"/api/v1/chats/{chat_id}/sessions"
|
||
data = {"ids": ids} if ids else {}
|
||
return await self._request('DELETE', endpoint, json=data)
|
||
|
||
async def converse_with_chat_assistant(self, chat_id: str, question: str, stream: bool = True,
|
||
session_id: Optional[str] = None,
|
||
user_id: Optional[str] = None) -> Union[Dict[str, Any], AsyncGenerator[Dict[str, Any], None]]:
|
||
"""
|
||
与聊天助手对话
|
||
"""
|
||
endpoint = f"/api/v1/chats/{chat_id}/completions"
|
||
data = {"question": question, "stream": stream}
|
||
|
||
if session_id:
|
||
data["session_id"] = session_id
|
||
|
||
print(f"开始对话: {question} {session_id}")
|
||
if user_id:
|
||
data["user_id"] = user_id
|
||
|
||
if stream:
|
||
return self._stream_request('POST', endpoint, json=data)
|
||
else:
|
||
return await self._request('POST', endpoint, json=data)
|
||
|
||
# ====================
|
||
# 代理管理
|
||
# ====================
|
||
|
||
async def create_session_with_agent(self, agent_id: str, user_id: Optional[str] = None,
|
||
file_data: Optional[Dict[str, Any]] = None,
|
||
**kwargs) -> Dict[str, Any]:
|
||
"""
|
||
创建与代理的会话
|
||
"""
|
||
if not self._session:
|
||
await self.create_session()
|
||
|
||
endpoint = f"/api/v1/agents/{agent_id}/sessions"
|
||
url = f"{self.base_url}{endpoint}"
|
||
params = {}
|
||
|
||
if user_id:
|
||
params["user_id"] = user_id
|
||
|
||
if file_data:
|
||
# 处理文件上传
|
||
data = aiohttp.FormData()
|
||
|
||
for key, file_path in file_data.items():
|
||
if os.path.exists(file_path):
|
||
file_name = Path(file_path).name
|
||
with open(file_path, 'rb') as f:
|
||
data.add_field(key, f.read(), filename=file_name)
|
||
|
||
headers = {'Authorization': f'Bearer {self.api_key}'}
|
||
|
||
async with self._session.post(url, headers=headers, data=data, params=params) as response:
|
||
try:
|
||
result = await response.json()
|
||
if result.get('code', 0) != 0:
|
||
raise RAGFlowError(result.get('code'), result.get('message'))
|
||
return result
|
||
except (aiohttp.ContentTypeError, json.JSONDecodeError):
|
||
text = await response.text()
|
||
raise RAGFlowError(response.status, text)
|
||
else:
|
||
# 普通JSON请求
|
||
data = kwargs
|
||
return await self._request('POST', endpoint, json=data, params=params)
|
||
|
||
async def converse_with_agent(self, agent_id: str, question: str, stream: bool = True,
|
||
session_id: Optional[str] = None, user_id: Optional[str] = None,
|
||
sync_dsl: bool = False, **kwargs) -> Union[Dict[str, Any], AsyncGenerator[Dict[str, Any], None]]:
|
||
"""
|
||
与代理对话
|
||
"""
|
||
endpoint = f"/api/v1/agents/{agent_id}/completions"
|
||
data = {"question": question, "stream": stream, "sync_dsl": sync_dsl}
|
||
|
||
if session_id:
|
||
data["session_id"] = session_id
|
||
if user_id:
|
||
data["user_id"] = user_id
|
||
|
||
# 添加其他Begin组件参数
|
||
data.update(kwargs)
|
||
|
||
if stream:
|
||
return self._stream_request('POST', endpoint, json=data)
|
||
else:
|
||
return await self._request('POST', endpoint, json=data)
|
||
|
||
async def list_agent_sessions(self, agent_id: str, page: int = 1, page_size: int = 30,
|
||
orderby: str = "create_time", desc: str = "true",
|
||
session_id: Optional[str] = None, user_id: Optional[str] = None,
|
||
dsl: bool = True) -> Dict[str, Any]:
|
||
"""
|
||
列出代理会话
|
||
"""
|
||
endpoint = f"/api/v1/agents/{agent_id}/sessions"
|
||
params = {
|
||
"page": page,
|
||
"page_size": page_size,
|
||
"orderby": orderby,
|
||
"desc": desc,
|
||
"dsl": dsl
|
||
}
|
||
|
||
if session_id:
|
||
params["id"] = session_id
|
||
if user_id:
|
||
params["user_id"] = user_id
|
||
|
||
return await self._request('GET', endpoint, params=params)
|
||
|
||
async def delete_agent_sessions(self, agent_id: str, ids: Optional[List[str]] = None) -> Dict[str, Any]:
|
||
"""
|
||
删除代理会话
|
||
"""
|
||
endpoint = f"/api/v1/agents/{agent_id}/sessions"
|
||
data = {"ids": ids} if ids else {}
|
||
return await self._request('DELETE', endpoint, json=data)
|
||
|
||
async def get_related_questions(self, question: str, login_token: str) -> Dict[str, Any]:
|
||
"""
|
||
生成相关问题
|
||
"""
|
||
endpoint = "/v1/sessions/related_questions"
|
||
headers = {
|
||
'Authorization': f'Bearer {login_token}',
|
||
'Content-Type': 'application/json'
|
||
}
|
||
data = {"question": question}
|
||
|
||
return await self._request('POST', endpoint, headers=headers, json=data)
|
||
|
||
async def list_agents(self, page: int = 1, page_size: int = 30, orderby: str = "create_time",
|
||
desc: str = "true", name: Optional[str] = None,
|
||
agent_id: Optional[str] = None) -> Dict[str, Any]:
|
||
"""
|
||
列出代理
|
||
"""
|
||
endpoint = "/api/v1/agents"
|
||
params = {
|
||
"page": page,
|
||
"page_size": page_size,
|
||
"orderby": orderby,
|
||
"desc": desc
|
||
}
|
||
|
||
if name:
|
||
params["name"] = name
|
||
if agent_id:
|
||
params["id"] = agent_id
|
||
|
||
return await self._request('GET', endpoint, params=params)
|
||
|
||
async def create_agent(self, title: str, description: Optional[str] = None,
|
||
dsl: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
|
||
"""
|
||
创建代理
|
||
"""
|
||
endpoint = "/api/v1/agents"
|
||
data = {"title": title}
|
||
|
||
if description is not None:
|
||
data["description"] = description
|
||
if dsl is not None:
|
||
data["dsl"] = dsl
|
||
|
||
return await self._request('POST', endpoint, json=data)
|
||
|
||
async def update_agent(self, agent_id: str, title: Optional[str] = None,
|
||
description: Optional[str] = None,
|
||
dsl: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
|
||
"""
|
||
更新代理
|
||
|
||
Args:
|
||
agent_id: 代理ID
|
||
title: 新标题
|
||
description: 新描述
|
||
dsl: 新DSL配置
|
||
"""
|
||
endpoint = f"/api/v1/agents/{agent_id}"
|
||
data = {}
|
||
|
||
if title is not None:
|
||
data["title"] = title
|
||
if description is not None:
|
||
data["description"] = description
|
||
if dsl is not None:
|
||
data["dsl"] = dsl
|
||
|
||
return await self._request('PUT', endpoint, json=data)
|
||
|
||
async def delete_agent(self, agent_id: str) -> Dict[str, Any]:
|
||
"""
|
||
删除代理
|
||
|
||
Args:
|
||
agent_id: 代理ID
|
||
"""
|
||
endpoint = f"/api/v1/agents/{agent_id}"
|
||
return await self._request('DELETE', endpoint)
|
||
|
||
|
||
# ====================
|
||
# 使用示例
|
||
# ====================
|
||
|
||
async def example_usage():
|
||
"""
|
||
异步RAGFlow SDK使用示例
|
||
"""
|
||
# 使用异步上下文管理器
|
||
async with AsyncRAGFlowClient(
|
||
base_url="http://10.0.0.202:82",
|
||
api_key="ragflow-hlMjRmNzE2ODNiNTExZjA4ZTNlMDI0Mm"
|
||
) as client:
|
||
|
||
try:
|
||
# 删除数据集
|
||
await client.delete_datasets(ids=["afe1387883bb11f0a0fd0242ac170006"])
|
||
print("删除数据集成功")
|
||
|
||
# 1. 创建数据集
|
||
dataset = await client.create_dataset(
|
||
name="我的数据集",
|
||
description="这是一个测试数据集",
|
||
chunk_method="naive"
|
||
)
|
||
dataset_id = dataset['data']['id']
|
||
print(f"创建数据集成功: {dataset_id}")
|
||
|
||
# 2. 上传文档
|
||
documents = await client.upload_documents(
|
||
dataset_id=dataset_id,
|
||
file_paths=[
|
||
"/home/admin-root/haotian/康达瑞贝斯机器人后台/ruoyi-fastapi-backend/requirements.txt",
|
||
"/home/admin-root/haotian/康达瑞贝斯机器人后台/ruoyi-fastapi-backend/requirements-pg.txt"
|
||
]
|
||
)
|
||
print("文档上传成功")
|
||
|
||
# 3. 解析文档
|
||
document_ids = [doc['id'] for doc in documents['data']]
|
||
await client.parse_documents(dataset_id, document_ids)
|
||
print("开始解析文档")
|
||
|
||
# 等待解析完成
|
||
await asyncio.sleep(5)
|
||
|
||
# 4. 创建聊天助手
|
||
chat_assistant = await client.create_chat_assistant(
|
||
name="我的AI助手",
|
||
dataset_ids=[dataset_id]
|
||
)
|
||
chat_id = chat_assistant['data']['id']
|
||
print(f"创建聊天助手成功: {chat_id}")
|
||
|
||
# 5. 创建会话
|
||
session = await client.create_session_with_chat(
|
||
chat_id=chat_id,
|
||
name="测试会话"
|
||
)
|
||
session_id = session['data']['id']
|
||
print(f"创建会话成功: {session_id}")
|
||
|
||
# 6. 开始对话(流式)
|
||
responses = client.converse_with_chat_assistant(
|
||
chat_id=chat_id,
|
||
question="你好,请介绍一下自己",
|
||
stream=True,
|
||
session_id=session_id
|
||
)
|
||
|
||
print("AI回复:")
|
||
async for response in responses:
|
||
if response.get('data') and isinstance(response['data'], dict):
|
||
answer = response['data'].get('answer', '')
|
||
if answer:
|
||
print(answer, end='', flush=True)
|
||
print()
|
||
|
||
# 7. 检索相关文档块
|
||
chunks = await client.retrieve_chunks(
|
||
question="RAGFlow的优势是什么?",
|
||
dataset_ids=[dataset_id],
|
||
top_k=5,
|
||
highlight=True
|
||
)
|
||
print(f"检索到 {chunks['data']['total']} 个相关文档块")
|
||
|
||
# 8. 列出数据集
|
||
datasets = await client.list_datasets(page=1, page_size=10)
|
||
print(f"当前有 {len(datasets['data'])} 个数据集")
|
||
|
||
except RAGFlowError as e:
|
||
print(f"RAGFlow API错误: {e}")
|
||
except Exception as e:
|
||
print(f"其他错误: {e}")
|
||
|
||
|
||
async def example_usage_1():
|
||
"""测试获取列表方法的异步版本"""
|
||
|
||
# 使用异步上下文管理器
|
||
async with AsyncRAGFlowClient(
|
||
base_url="http://10.0.0.202:82",
|
||
api_key="ragflow-hlMjRmNzE2ODNiNTExZjA4ZTNlMDI0Mm"
|
||
) as client:
|
||
|
||
try:
|
||
# 1. 获取数据集列表
|
||
results_dataset = await client.list_datasets()
|
||
print(f"获取数据集列表成功,共有 {len(results_dataset['data'])} 个数据集")
|
||
print("数据集ID:\n", [result["id"] for result in results_dataset['data']])
|
||
|
||
# 2. 获取数据集中文档列表
|
||
for result in results_dataset['data']:
|
||
print(f"数据集 {result['id']} 的文档列表为:")
|
||
results_doc = await client.list_documents(dataset_id=result['id'])
|
||
# 文档名称
|
||
print([t["name"] for t in results_doc["data"]["docs"]])
|
||
|
||
except RAGFlowError as e:
|
||
print(f"RAGFlow API错误: {e}")
|
||
except Exception as e:
|
||
print(f"其他错误: {e}")
|
||
|
||
|
||
async def batch_operations_example():
|
||
"""
|
||
批量操作示例 - 演示异步并发处理
|
||
"""
|
||
async with AsyncRAGFlowClient(
|
||
base_url="http://10.0.0.202:82",
|
||
api_key="ragflow-hlMjRmNzE2ODNiNTExZjA4ZTNlMDI0Mm"
|
||
) as client:
|
||
|
||
try:
|
||
# 并发获取数据集和代理列表
|
||
datasets_task = client.list_datasets()
|
||
agents_task = client.list_agents()
|
||
chats_task = client.list_chat_assistants()
|
||
|
||
# 等待所有任务完成
|
||
datasets, agents, chats = await asyncio.gather(
|
||
datasets_task, agents_task, chats_task
|
||
)
|
||
|
||
print(f"数据集数量: {len(datasets['data'])}")
|
||
print(f"代理数量: {len(agents['data'])}")
|
||
print(f"聊天助手数量: {len(chats['data'])}")
|
||
|
||
# 如果有数据集,并发获取每个数据集的文档
|
||
if datasets['data']:
|
||
doc_tasks = [
|
||
client.list_documents(dataset_id=dataset['id'])
|
||
for dataset in datasets['data'][:3] # 限制前3个
|
||
]
|
||
|
||
doc_results = await asyncio.gather(*doc_tasks, return_exceptions=True)
|
||
|
||
for i, result in enumerate(doc_results):
|
||
if isinstance(result, Exception):
|
||
print(f"获取数据集 {datasets['data'][i]['id']} 的文档时出错: {result}")
|
||
else:
|
||
print(f"数据集 {datasets['data'][i]['id']} 有 {len(result['data']['docs'])} 个文档")
|
||
|
||
except Exception as e:
|
||
print(f"批量操作出错: {e}")
|
||
|
||
|
||
# 运行示例的辅助函数
|
||
def run_example():
|
||
"""运行异步示例"""
|
||
# 可以选择运行不同的示例
|
||
asyncio.run(example_usage_1())
|
||
# asyncio.run(example_usage())
|
||
# asyncio.run(batch_operations_example())
|
||
|
||
|
||
if __name__ == "__main__":
|
||
# 运行示例
|
||
run_example() |