This commit is contained in:
2026-02-26 18:06:17 +08:00
parent 68200cdfe6
commit ad0082d811
33 changed files with 356 additions and 249 deletions
+8 -3
View File
@@ -14,22 +14,27 @@ class RagflowClient:
self._base_url = cfg.get("url", "")
self._api_key = cfg.get("api_key", "")
self._retrieval_path = cfg.get("retrieval", "/api/v1/retrieval")
self._dataset_ids = cfg.get("dataset_ids", "")
self._table_retrieval_dataset_id = cfg.get("table_retrieval_dataset_id", "")
self._sql_gen_dataset_id = cfg.get("sql_gen_dataset_id", "")
def _build_url(self) -> str:
return self._base_url.rstrip("/") + "/" + self._retrieval_path.lstrip("/")
def retrieve(self, query: str, top_k: int = 3) -> Dict[str, Any]:
def retrieve(self, query: str, top_k: int = 3, dataset_id: Optional[str] = None, document_ids: Optional[str] = None) -> Dict[str, Any]:
"""检索匹配文档"""
if not self._base_url or not self._retrieval_path:
raise RuntimeError("未配置 ragflow.url 或 ragflow.retrieval")
if not dataset_id:
raise ValueError("未提供 ragflow.dataset_id,无法进行检索")
url = self._build_url()
headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else {}
payload = {
"dataset_ids": self._dataset_ids,
"dataset_ids": dataset_id or "",
"query": query,
"top_k": top_k,
}
if document_ids:
payload["document_ids"] = document_ids
with httpx.Client(timeout=30) as client:
response = client.post(url, json=payload, headers=headers)