mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-08-13 04:13:35 +08:00
test: extend /dify/retrieval unit test coverage (#17957)
## Summary - Extends unit test coverage for pi/apps/restful_apis/dify_retrieval_api.py (the Dify external knowledge base endpoint). Co-authored-by: zjm11902 <zjm11902@users.noreply.github.com> Co-authored-by: Jin Hai <haijin.chn@gmail.com>
This commit is contained in:
@@ -65,12 +65,17 @@ def _stub(monkeypatch, name, **attrs):
|
||||
|
||||
|
||||
class _FakeRetriever:
|
||||
def __init__(self, chunks=None):
|
||||
def __init__(self, chunks=None, raise_exc=None):
|
||||
self._chunks = chunks if chunks is not None else []
|
||||
self._raise_exc = raise_exc
|
||||
self.retrieval_calls = []
|
||||
self.last_kwargs = None
|
||||
|
||||
async def retrieval(self, question, embd_mdl, tenant_id, kb_ids, **kwargs):
|
||||
if self._raise_exc is not None:
|
||||
raise self._raise_exc
|
||||
self.retrieval_calls.append({"question": question, "tenant_id": tenant_id, "kb_ids": list(kb_ids)})
|
||||
self.last_kwargs = kwargs
|
||||
return {"chunks": list(self._chunks)}
|
||||
|
||||
def retrieval_by_children(self, chunks, _tenant_ids):
|
||||
@@ -78,11 +83,33 @@ class _FakeRetriever:
|
||||
|
||||
|
||||
class _FakeKGRetriever:
|
||||
def __init__(self, content=""):
|
||||
self._content = content
|
||||
|
||||
async def retrieval(self, *_a, **_k):
|
||||
return {"content_with_weight": ""}
|
||||
if not self._content:
|
||||
return {"content_with_weight": ""}
|
||||
return {
|
||||
"content_with_weight": self._content,
|
||||
"doc_id": "d1",
|
||||
"similarity": 0.95,
|
||||
"docnm_kwd": "graph.txt",
|
||||
}
|
||||
|
||||
|
||||
def _load_dify_retrieval(monkeypatch, *, kb, accessible, request_body, tenant_id, chunks=None):
|
||||
def _load_dify_retrieval(
|
||||
monkeypatch,
|
||||
*,
|
||||
kb,
|
||||
accessible,
|
||||
request_body,
|
||||
tenant_id,
|
||||
chunks=None,
|
||||
method="POST",
|
||||
request_args=None,
|
||||
kg_content="",
|
||||
meta_filter_ids=None,
|
||||
):
|
||||
"""Load dify_retrieval_api.py with minimum stubs to exercise the retrieval handler."""
|
||||
|
||||
def _add_tenant_id_to_kwargs(func):
|
||||
@@ -141,7 +168,7 @@ def _load_dify_retrieval(monkeypatch, *, kb, accessible, request_body, tenant_id
|
||||
_stub(
|
||||
monkeypatch,
|
||||
"common.metadata_utils",
|
||||
meta_filter=lambda *_a, **_k: [],
|
||||
meta_filter=lambda *_a, **_k: list(meta_filter_ids) if meta_filter_ids is not None else [],
|
||||
convert_conditions=lambda c: c,
|
||||
)
|
||||
|
||||
@@ -152,11 +179,11 @@ def _load_dify_retrieval(monkeypatch, *, kb, accessible, request_body, tenant_id
|
||||
monkeypatch,
|
||||
"common.settings",
|
||||
retriever=fake_retriever,
|
||||
kg_retriever=_FakeKGRetriever(),
|
||||
kg_retriever=_FakeKGRetriever(kg_content),
|
||||
)
|
||||
|
||||
quart_stub = ModuleType("quart")
|
||||
quart_stub.request = SimpleNamespace(method="POST", args={})
|
||||
quart_stub.request = SimpleNamespace(method=method, args=request_args or {})
|
||||
quart_stub.jsonify = lambda payload: payload
|
||||
monkeypatch.setitem(sys.modules, "quart", quart_stub)
|
||||
|
||||
@@ -265,3 +292,362 @@ class TestDifyRetrievalTenantCheck:
|
||||
|
||||
assert result["code"] == 404
|
||||
assert "not found" in result["message"].lower()
|
||||
|
||||
|
||||
@pytest.mark.p1
|
||||
class TestDifyRetrievalArgumentValidation:
|
||||
"""Argument validation paths of /dify/retrieval (POST and GET)."""
|
||||
|
||||
@pytest.mark.p1
|
||||
def test_empty_knowledge_id_returns_argument_error(self, monkeypatch):
|
||||
owner_kb = SimpleNamespace(id="kb-owner", tenant_id="tenant-owner", tenant_embd_id="", embd_id="bge")
|
||||
module = _load_dify_retrieval(
|
||||
monkeypatch,
|
||||
kb=(True, owner_kb),
|
||||
accessible=lambda _id, _u: True,
|
||||
request_body={"knowledge_id": "", "query": "hello"},
|
||||
tenant_id="tenant-owner",
|
||||
)
|
||||
|
||||
result = asyncio.run(module.retrieval())
|
||||
|
||||
assert result["code"] == 101
|
||||
assert "knowledge_id" in result["message"]
|
||||
assert module._fake_retriever.retrieval_calls == [], "retriever must not run on validation failure"
|
||||
|
||||
@pytest.mark.p1
|
||||
def test_missing_query_returns_argument_error(self, monkeypatch):
|
||||
owner_kb = SimpleNamespace(id="kb-owner", tenant_id="tenant-owner", tenant_embd_id="", embd_id="bge")
|
||||
module = _load_dify_retrieval(
|
||||
monkeypatch,
|
||||
kb=(True, owner_kb),
|
||||
accessible=lambda _id, _u: True,
|
||||
request_body={"knowledge_id": "kb-owner"},
|
||||
tenant_id="tenant-owner",
|
||||
)
|
||||
|
||||
result = asyncio.run(module.retrieval())
|
||||
|
||||
assert result["code"] == 101
|
||||
assert "query" in result["message"]
|
||||
|
||||
@pytest.mark.p1
|
||||
def test_missing_knowledge_id_field_returns_argument_error(self, monkeypatch):
|
||||
owner_kb = SimpleNamespace(id="kb-owner", tenant_id="tenant-owner", tenant_embd_id="", embd_id="bge")
|
||||
module = _load_dify_retrieval(
|
||||
monkeypatch,
|
||||
kb=(True, owner_kb),
|
||||
accessible=lambda _id, _u: True,
|
||||
request_body={"query": "hello"},
|
||||
tenant_id="tenant-owner",
|
||||
)
|
||||
|
||||
result = asyncio.run(module.retrieval())
|
||||
|
||||
assert result["code"] == 101
|
||||
assert "knowledge_id" in result["message"]
|
||||
assert module._fake_retriever.retrieval_calls == []
|
||||
|
||||
@pytest.mark.p1
|
||||
def test_empty_query_returns_argument_error(self, monkeypatch):
|
||||
owner_kb = SimpleNamespace(id="kb-owner", tenant_id="tenant-owner", tenant_embd_id="", embd_id="bge")
|
||||
module = _load_dify_retrieval(
|
||||
monkeypatch,
|
||||
kb=(True, owner_kb),
|
||||
accessible=lambda _id, _u: True,
|
||||
request_body={"knowledge_id": "kb-owner", "query": ""},
|
||||
tenant_id="tenant-owner",
|
||||
)
|
||||
|
||||
result = asyncio.run(module.retrieval())
|
||||
|
||||
assert result["code"] == 101
|
||||
assert "query" in result["message"]
|
||||
|
||||
@pytest.mark.p1
|
||||
def test_invalid_top_k_returns_argument_error(self, monkeypatch):
|
||||
owner_kb = SimpleNamespace(id="kb-owner", tenant_id="tenant-owner", tenant_embd_id="", embd_id="bge")
|
||||
module = _load_dify_retrieval(
|
||||
monkeypatch,
|
||||
kb=(True, owner_kb),
|
||||
accessible=lambda _id, _u: True,
|
||||
request_body={
|
||||
"knowledge_id": "kb-owner",
|
||||
"query": "hello",
|
||||
"retrieval_setting": {"top_k": "not-an-int"},
|
||||
},
|
||||
tenant_id="tenant-owner",
|
||||
)
|
||||
|
||||
result = asyncio.run(module.retrieval())
|
||||
|
||||
assert result["code"] == 101
|
||||
assert module._fake_retriever.retrieval_calls == []
|
||||
|
||||
@pytest.mark.p1
|
||||
def test_invalid_score_threshold_returns_argument_error(self, monkeypatch):
|
||||
owner_kb = SimpleNamespace(id="kb-owner", tenant_id="tenant-owner", tenant_embd_id="", embd_id="bge")
|
||||
module = _load_dify_retrieval(
|
||||
monkeypatch,
|
||||
kb=(True, owner_kb),
|
||||
accessible=lambda _id, _u: True,
|
||||
request_body={
|
||||
"knowledge_id": "kb-owner",
|
||||
"query": "hello",
|
||||
"retrieval_setting": {"score_threshold": "oops"},
|
||||
},
|
||||
tenant_id="tenant-owner",
|
||||
)
|
||||
|
||||
result = asyncio.run(module.retrieval())
|
||||
|
||||
assert result["code"] == 101
|
||||
|
||||
@pytest.mark.p1
|
||||
def test_retrieval_setting_values_are_forwarded(self, monkeypatch):
|
||||
owner_kb = SimpleNamespace(id="kb-owner", tenant_id="tenant-owner", tenant_embd_id="", embd_id="bge")
|
||||
module = _load_dify_retrieval(
|
||||
monkeypatch,
|
||||
kb=(True, owner_kb),
|
||||
accessible=lambda _id, _u: True,
|
||||
request_body={
|
||||
"knowledge_id": "kb-owner",
|
||||
"query": "hello",
|
||||
"retrieval_setting": {"top_k": 7, "score_threshold": 0.3},
|
||||
},
|
||||
tenant_id="tenant-owner",
|
||||
)
|
||||
|
||||
asyncio.run(module.retrieval())
|
||||
|
||||
kwargs = module._fake_retriever.last_kwargs
|
||||
assert kwargs["page_size"] == 7
|
||||
assert kwargs["top"] == 7
|
||||
assert kwargs["similarity_threshold"] == 0.3
|
||||
|
||||
@pytest.mark.p1
|
||||
def test_get_request_with_query_args(self, monkeypatch):
|
||||
owner_kb = SimpleNamespace(id="kb-owner", tenant_id="tenant-owner", tenant_embd_id="", embd_id="bge")
|
||||
module = _load_dify_retrieval(
|
||||
monkeypatch,
|
||||
kb=(True, owner_kb),
|
||||
accessible=lambda _id, _u: True,
|
||||
request_body={},
|
||||
request_args={"knowledge_id": "kb-owner", "query": "hello", "top_k": "5", "score_threshold": "0.2"},
|
||||
method="GET",
|
||||
tenant_id="tenant-owner",
|
||||
chunks=[{"doc_id": "d1", "content_with_weight": "hello world", "similarity": 0.8, "docnm_kwd": "doc.txt"}],
|
||||
)
|
||||
|
||||
result = asyncio.run(module.retrieval())
|
||||
|
||||
assert len(result["records"]) == 1
|
||||
kwargs = module._fake_retriever.last_kwargs
|
||||
assert kwargs["page_size"] == 5
|
||||
assert kwargs["similarity_threshold"] == 0.2
|
||||
|
||||
@pytest.mark.p1
|
||||
def test_get_request_invalid_top_k_returns_argument_error(self, monkeypatch):
|
||||
owner_kb = SimpleNamespace(id="kb-owner", tenant_id="tenant-owner", tenant_embd_id="", embd_id="bge")
|
||||
module = _load_dify_retrieval(
|
||||
monkeypatch,
|
||||
kb=(True, owner_kb),
|
||||
accessible=lambda _id, _u: True,
|
||||
request_body={},
|
||||
request_args={"knowledge_id": "kb-owner", "query": "hello", "top_k": "abc"},
|
||||
method="GET",
|
||||
tenant_id="tenant-owner",
|
||||
)
|
||||
|
||||
result = asyncio.run(module.retrieval())
|
||||
|
||||
assert result["code"] == 101
|
||||
|
||||
@pytest.mark.p1
|
||||
def test_get_request_invalid_score_threshold_returns_argument_error(self, monkeypatch):
|
||||
owner_kb = SimpleNamespace(id="kb-owner", tenant_id="tenant-owner", tenant_embd_id="", embd_id="bge")
|
||||
module = _load_dify_retrieval(
|
||||
monkeypatch,
|
||||
kb=(True, owner_kb),
|
||||
accessible=lambda _id, _u: True,
|
||||
request_body={},
|
||||
request_args={"knowledge_id": "kb-owner", "query": "hello", "score_threshold": "abc"},
|
||||
method="GET",
|
||||
tenant_id="tenant-owner",
|
||||
)
|
||||
|
||||
result = asyncio.run(module.retrieval())
|
||||
|
||||
assert result["code"] == 101
|
||||
|
||||
|
||||
@pytest.mark.p1
|
||||
class TestDifyRetrievalRetrievalBehavior:
|
||||
"""Retrieval behavior branches: metadata filters, knowledge graph, records mapping, errors."""
|
||||
|
||||
@pytest.mark.p1
|
||||
def test_metadata_condition_empty_result_uses_sentinel(self, monkeypatch):
|
||||
owner_kb = SimpleNamespace(id="kb-owner", tenant_id="tenant-owner", tenant_embd_id="", embd_id="bge")
|
||||
module = _load_dify_retrieval(
|
||||
monkeypatch,
|
||||
kb=(True, owner_kb),
|
||||
accessible=lambda _id, _u: True,
|
||||
request_body={
|
||||
"knowledge_id": "kb-owner",
|
||||
"query": "hello",
|
||||
"metadata_condition": {"conditions": [{"name": "k", "comparison_operator": "eq", "value": "v"}]},
|
||||
},
|
||||
meta_filter_ids=[],
|
||||
tenant_id="tenant-owner",
|
||||
)
|
||||
|
||||
asyncio.run(module.retrieval())
|
||||
|
||||
assert module._fake_retriever.last_kwargs["doc_ids"] == ["-999"]
|
||||
|
||||
@pytest.mark.p1
|
||||
def test_metadata_condition_matching_docs_are_forwarded(self, monkeypatch):
|
||||
owner_kb = SimpleNamespace(id="kb-owner", tenant_id="tenant-owner", tenant_embd_id="", embd_id="bge")
|
||||
module = _load_dify_retrieval(
|
||||
monkeypatch,
|
||||
kb=(True, owner_kb),
|
||||
accessible=lambda _id, _u: True,
|
||||
request_body={
|
||||
"knowledge_id": "kb-owner",
|
||||
"query": "hello",
|
||||
"metadata_condition": {"conditions": [{"name": "k", "comparison_operator": "eq", "value": "v"}]},
|
||||
},
|
||||
meta_filter_ids=["d1", "d2"],
|
||||
tenant_id="tenant-owner",
|
||||
)
|
||||
|
||||
asyncio.run(module.retrieval())
|
||||
|
||||
assert module._fake_retriever.last_kwargs["doc_ids"] == ["d1", "d2"]
|
||||
|
||||
@pytest.mark.p1
|
||||
def test_use_kg_inserts_graph_chunk_at_front(self, monkeypatch):
|
||||
owner_kb = SimpleNamespace(id="kb-owner", tenant_id="tenant-owner", tenant_embd_id="", embd_id="bge")
|
||||
module = _load_dify_retrieval(
|
||||
monkeypatch,
|
||||
kb=(True, owner_kb),
|
||||
accessible=lambda _id, _u: True,
|
||||
request_body={"knowledge_id": "kb-owner", "query": "hello", "use_kg": True},
|
||||
kg_content="graph result text",
|
||||
tenant_id="tenant-owner",
|
||||
chunks=[{"doc_id": "d1", "content_with_weight": "base chunk", "similarity": 0.5, "docnm_kwd": "doc.txt"}],
|
||||
)
|
||||
|
||||
result = asyncio.run(module.retrieval())
|
||||
|
||||
assert len(result["records"]) == 2
|
||||
assert result["records"][0]["content"] == "graph result text"
|
||||
assert result["records"][0]["title"] == "graph.txt"
|
||||
|
||||
@pytest.mark.p1
|
||||
def test_use_kg_empty_result_is_not_inserted(self, monkeypatch):
|
||||
owner_kb = SimpleNamespace(id="kb-owner", tenant_id="tenant-owner", tenant_embd_id="", embd_id="bge")
|
||||
module = _load_dify_retrieval(
|
||||
monkeypatch,
|
||||
kb=(True, owner_kb),
|
||||
accessible=lambda _id, _u: True,
|
||||
request_body={"knowledge_id": "kb-owner", "query": "hello", "use_kg": True},
|
||||
kg_content="",
|
||||
tenant_id="tenant-owner",
|
||||
chunks=[{"doc_id": "d1", "content_with_weight": "base chunk", "similarity": 0.5, "docnm_kwd": "doc.txt"}],
|
||||
)
|
||||
|
||||
result = asyncio.run(module.retrieval())
|
||||
|
||||
assert len(result["records"]) == 1
|
||||
assert result["records"][0]["content"] == "base chunk"
|
||||
|
||||
@pytest.mark.p1
|
||||
def test_records_carry_full_metadata(self, monkeypatch):
|
||||
owner_kb = SimpleNamespace(id="kb-owner", tenant_id="tenant-owner", tenant_embd_id="", embd_id="bge")
|
||||
module = _load_dify_retrieval(
|
||||
monkeypatch,
|
||||
kb=(True, owner_kb),
|
||||
accessible=lambda _id, _u: True,
|
||||
request_body={"knowledge_id": "kb-owner", "query": "hello"},
|
||||
tenant_id="tenant-owner",
|
||||
chunks=[{"doc_id": "d1", "content_with_weight": "hello world", "similarity": 0.8, "docnm_kwd": "doc.txt"}],
|
||||
)
|
||||
|
||||
result = asyncio.run(module.retrieval())
|
||||
|
||||
record = result["records"][0]
|
||||
assert record["content"] == "hello world"
|
||||
assert record["score"] == 0.8
|
||||
assert record["title"] == "doc.txt"
|
||||
assert record["metadata"]["doc_id"] == "d1"
|
||||
assert record["metadata"]["document_id"] == "d1"
|
||||
|
||||
@pytest.mark.p1
|
||||
def test_no_chunks_returns_empty_records(self, monkeypatch):
|
||||
owner_kb = SimpleNamespace(id="kb-owner", tenant_id="tenant-owner", tenant_embd_id="", embd_id="bge")
|
||||
module = _load_dify_retrieval(
|
||||
monkeypatch,
|
||||
kb=(True, owner_kb),
|
||||
accessible=lambda _id, _u: True,
|
||||
request_body={"knowledge_id": "kb-owner", "query": "hello"},
|
||||
tenant_id="tenant-owner",
|
||||
)
|
||||
|
||||
result = asyncio.run(module.retrieval())
|
||||
|
||||
assert result == {"records": []}
|
||||
|
||||
@pytest.mark.p1
|
||||
def test_not_found_exception_returns_404(self, monkeypatch):
|
||||
owner_kb = SimpleNamespace(id="kb-owner", tenant_id="tenant-owner", tenant_embd_id="", embd_id="bge")
|
||||
module = _load_dify_retrieval(
|
||||
monkeypatch,
|
||||
kb=(True, owner_kb),
|
||||
accessible=lambda _id, _u: True,
|
||||
request_body={"knowledge_id": "kb-owner", "query": "hello"},
|
||||
tenant_id="tenant-owner",
|
||||
chunks=None,
|
||||
)
|
||||
module._fake_retriever._raise_exc = Exception("index_not_found_exception")
|
||||
|
||||
result = asyncio.run(module.retrieval())
|
||||
|
||||
assert result["code"] == 404
|
||||
assert "no chunk found" in result["message"].lower()
|
||||
|
||||
@pytest.mark.p1
|
||||
def test_other_exception_returns_server_error(self, monkeypatch):
|
||||
owner_kb = SimpleNamespace(id="kb-owner", tenant_id="tenant-owner", tenant_embd_id="", embd_id="bge")
|
||||
module = _load_dify_retrieval(
|
||||
monkeypatch,
|
||||
kb=(True, owner_kb),
|
||||
accessible=lambda _id, _u: True,
|
||||
request_body={"knowledge_id": "kb-owner", "query": "hello"},
|
||||
tenant_id="tenant-owner",
|
||||
)
|
||||
module._fake_retriever._raise_exc = RuntimeError("boom")
|
||||
|
||||
result = asyncio.run(module.retrieval())
|
||||
|
||||
assert result["code"] == 500
|
||||
|
||||
|
||||
@pytest.mark.p1
|
||||
class TestDifyRetrievalHealth:
|
||||
"""Health endpoint for Dify external knowledge base connectivity checks."""
|
||||
|
||||
@pytest.mark.p1
|
||||
def test_health_endpoint(self, monkeypatch):
|
||||
module = _load_dify_retrieval(
|
||||
monkeypatch,
|
||||
kb=(True, SimpleNamespace()),
|
||||
accessible=lambda _id, _u: True,
|
||||
request_body={},
|
||||
tenant_id="tenant-owner",
|
||||
)
|
||||
|
||||
result = asyncio.run(module.retrieval_health_check())
|
||||
|
||||
assert result["code"] == 0
|
||||
assert result["data"] is True
|
||||
|
||||
Reference in New Issue
Block a user