init
This commit is contained in:
@@ -20,6 +20,15 @@ class RagflowClient:
|
||||
def _build_url(self) -> str:
|
||||
return self._base_url.rstrip("/") + "/" + self._retrieval_path.lstrip("/")
|
||||
|
||||
@staticmethod
|
||||
def _normalize_dataset_ids(dataset_id: Optional[str]) -> list[str]:
|
||||
"""将配置值规范化为 RAGFlow 需要的 list[string]"""
|
||||
if not dataset_id:
|
||||
return []
|
||||
# 兼容逗号分隔配置
|
||||
parts = [p.strip() for p in str(dataset_id).split(",") if p.strip()]
|
||||
return parts
|
||||
|
||||
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:
|
||||
@@ -29,10 +38,12 @@ class RagflowClient:
|
||||
url = self._build_url()
|
||||
headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else {}
|
||||
payload = {
|
||||
"dataset_ids": dataset_id or "",
|
||||
"query": query,
|
||||
"dataset_ids": self._normalize_dataset_ids(dataset_id),
|
||||
"question": query,
|
||||
"top_k": top_k,
|
||||
}
|
||||
# 兼容部分版本字段
|
||||
payload["query"] = query
|
||||
if document_ids:
|
||||
payload["document_ids"] = document_ids
|
||||
|
||||
@@ -57,6 +68,17 @@ def extract_table_name(record: Dict[str, Any]) -> Optional[str]:
|
||||
return record.get(key)
|
||||
|
||||
content = record.get("content") or record.get("text") or ""
|
||||
|
||||
# 兼容 content 为 JSON 字符串:{"table":"xxx", ...}
|
||||
try:
|
||||
parsed = json.loads(str(content))
|
||||
if isinstance(parsed, dict):
|
||||
for key in ("table", "table_name"):
|
||||
if parsed.get(key):
|
||||
return str(parsed.get(key))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
for line in str(content).splitlines():
|
||||
if line.lower().startswith("table:"):
|
||||
return line.split(":", 1)[1].strip()
|
||||
|
||||
Reference in New Issue
Block a user