Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 29 additions & 0 deletions providers/qdrant/docs/operators/qdrant.rst
Original file line number Diff line number Diff line change
Expand Up @@ -38,3 +38,32 @@ An example using the operator in this way:
:dedent: 4
:start-after: [START howto_operator_qdrant_ingest]
:end-before: [END howto_operator_qdrant_ingest]


.. _howto/operator:QdrantSearchOperator:

QdrantSearchOperator
======================

Use the :class:`~airflow.providers.qdrant.operators.qdrant.QdrantSearchOperator` to
run a vector similarity search against a Qdrant collection and pull the top-k
matches back into the DAG for downstream tasks (for example, a retrieval step
in a RAG pipeline).

Using the Operator
^^^^^^^^^^^^^^^^^^

Pass the target ``collection_name`` and a ``query`` (typically a dense embedding
vector, but any query form accepted by
:meth:`~qdrant_client.QdrantClient.query_points` also works). The operator pushes
the results to XCom as a list of dictionaries -- one per matched point, with
``id``, ``score``, ``payload`` and (optionally) ``vector`` keys -- so downstream
tasks can consume them without any custom serialization.

An example using the operator downstream of an ingest task:

.. exampleinclude:: /../../qdrant/tests/system/qdrant/example_dag_qdrant.py
:language: python
:dedent: 4
:start-after: [START howto_operator_qdrant_search]
:end-before: [END howto_operator_qdrant_search]
53 changes: 52 additions & 1 deletion providers/qdrant/src/airflow/providers/qdrant/hooks/qdrant.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,14 +18,17 @@
from __future__ import annotations

from functools import cached_property
from typing import Any
from typing import TYPE_CHECKING, Any

from grpc import RpcError
from qdrant_client import QdrantClient
from qdrant_client.http.exceptions import UnexpectedResponse

from airflow.providers.common.compat.sdk import BaseHook

if TYPE_CHECKING:
from qdrant_client import models


class QdrantHook(BaseHook):
"""
Expand Down Expand Up @@ -130,3 +133,51 @@ def verify_connection(self) -> tuple[bool, str]:
def test_connection(self) -> tuple[bool, str]:
"""Test the connection to the Qdrant instance."""
return self.verify_connection()

def search(
self,
collection_name: str,
query: Any,
*,
query_filter: models.Filter | None = None,
search_params: models.SearchParams | None = None,
limit: int = 10,
offset: int | None = None,
with_payload: bool | list[str] = True,
with_vectors: bool | list[str] = False,
score_threshold: float | None = None,
**kwargs: Any,
) -> list[dict[str, Any]]:
"""
Run a similarity search against a Qdrant collection and return the matches.

Wraps ``QdrantClient.query_points`` and returns plain, XCom-serializable
dictionaries (via ``ScoredPoint.model_dump``) instead of pydantic objects.

:param collection_name: Name of the collection to search.
:param query: The query. Commonly a dense vector (``list[float]``); it may
also be a point id, a named/sparse vector, or a ``qdrant_client.models``
query object. See the Qdrant ``query_points`` docs for all supported forms.
:param query_filter: Optional filter to restrict which points are considered.
:param search_params: Optional search-tuning parameters (e.g. ``hnsw_ef``).
:param limit: Maximum number of results to return (top-k). Defaults to 10.
:param offset: Number of results to skip, for pagination. Optional.
:param with_payload: Whether (or which payload fields) to include. Defaults to True.
:param with_vectors: Whether (or which vectors) to include. Defaults to False.
:param score_threshold: Minimal similarity score for a result to be returned.
:param kwargs: Additional keyword arguments forwarded to ``query_points``.
:return: A list of scored points as dictionaries, ordered by descending score.
"""
response = self.conn.query_points(
collection_name=collection_name,
query=query,
query_filter=query_filter,
search_params=search_params,
limit=limit,
offset=offset,
with_payload=with_payload,
with_vectors=with_vectors,
score_threshold=score_threshold,
**kwargs,
)
return [point.model_dump() for point in response.points]
84 changes: 84 additions & 0 deletions providers/qdrant/src/airflow/providers/qdrant/operators/qdrant.py
Original file line number Diff line number Diff line change
Expand Up @@ -108,3 +108,87 @@ def execute(self, context: Context) -> None:
max_retries=self.max_retries,
wait=self.wait,
)


class QdrantSearchOperator(BaseOperator):
"""
Run a similarity search against a Qdrant collection and return the matches.

The operator wraps :meth:`~airflow.providers.qdrant.hooks.qdrant.QdrantHook.search`,
which calls Qdrant's ``query_points`` API and converts each returned point to a
plain, XCom-serializable ``dict`` (id, score, payload, vector). Results land in
XCom as a ``list[dict]`` ready to be consumed by downstream tasks (e.g. an LLM
task that builds a prompt from the retrieved payloads in a RAG pipeline).

.. seealso::
For more information on how to use this operator, take a look at the guide:
:ref:`howto/operator:QdrantSearchOperator`

:param conn_id: The connection id to connect to a Qdrant instance.
:param collection_name: The name of the collection to search.
:param query: The query. Commonly a dense vector (``list[float]``); it may also
be a point id, a named/sparse vector, or any query form supported by
``qdrant_client.QdrantClient.query_points``.
:param query_filter: Optional filter to restrict which points are considered.
:param search_params: Optional Qdrant ``SearchParams`` (e.g. ``hnsw_ef``).
:param limit: Maximum number of results to return (top-k). Defaults to 10.
:param offset: Number of results to skip, for pagination. Optional.
:param with_payload: Whether (or which payload fields) to include. Defaults to True.
:param with_vectors: Whether (or which vectors) to include. Defaults to False.
:param score_threshold: Minimum similarity score for a result to be returned.
:param kwargs: Additional keyword arguments passed to the ``BaseOperator``
constructor.
"""

template_fields: Sequence[str] = (
"collection_name",
"query",
"query_filter",
"limit",
)

def __init__(
self,
*,
conn_id: str = QdrantHook.default_conn_name,
collection_name: str,
query: Any,
query_filter: Any = None,
search_params: Any = None,
limit: int = 10,
offset: int | None = None,
with_payload: bool | list[str] = True,
with_vectors: bool | list[str] = False,
score_threshold: float | None = None,
**kwargs: Any,
) -> None:
super().__init__(**kwargs)
self.conn_id = conn_id
self.collection_name = collection_name
self.query = query
self.query_filter = query_filter
self.search_params = search_params
self.limit = limit
self.offset = offset
self.with_payload = with_payload
self.with_vectors = with_vectors
self.score_threshold = score_threshold

@cached_property
def hook(self) -> QdrantHook:
"""Return an instance of QdrantHook."""
return QdrantHook(conn_id=self.conn_id)

def execute(self, context: Context) -> list[dict[str, Any]]:
"""Search the Qdrant collection and return the matching points as dicts."""
return self.hook.search(
collection_name=self.collection_name,
query=self.query,
query_filter=self.query_filter,
search_params=self.search_params,
limit=self.limit,
offset=self.offset,
with_payload=self.with_payload,
with_vectors=self.with_vectors,
score_threshold=self.score_threshold,
)
16 changes: 14 additions & 2 deletions providers/qdrant/tests/system/qdrant/example_dag_qdrant.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@
from datetime import datetime

from airflow import DAG
from airflow.providers.qdrant.operators.qdrant import QdrantIngestOperator
from airflow.providers.qdrant.operators.qdrant import QdrantIngestOperator, QdrantSearchOperator

with DAG(
"example_qdrant_ingest",
Expand All @@ -32,7 +32,7 @@
ids: list[str | int] = [32, 21, "b626f6a9-b14d-4af9-b7c3-43d8deb719a6"]
payload = [{"meta": "data"}, {"meta": "data_2"}, {"meta": "data_3", "extra": "data"}]

QdrantIngestOperator(
ingest = QdrantIngestOperator(
task_id="qdrant_ingest",
collection_name="test_collection",
vectors=vectors,
Expand All @@ -42,6 +42,18 @@
)
# [END howto_operator_qdrant_ingest]

# [START howto_operator_qdrant_search]
search = QdrantSearchOperator(
task_id="qdrant_search",
collection_name="test_collection",
query=[0.5, 0.5, 0.5, 0.5],
limit=3,
with_payload=True,
)
# [END howto_operator_qdrant_search]

ingest >> search


from tests_common.test_utils.system_tests import get_test_run

Expand Down
80 changes: 80 additions & 0 deletions providers/qdrant/tests/unit/qdrant/hooks/test_qdrant.py
Original file line number Diff line number Diff line change
Expand Up @@ -131,3 +131,83 @@ def test_delete_collection(self, conn):
self.qdrant_hook.conn.delete_collection(collection_name=self.collection_name)

conn.delete_collection.assert_called_once_with(collection_name=self.collection_name)

@patch("airflow.providers.qdrant.hooks.qdrant.QdrantHook.conn")
def test_search_returns_list_of_dicts(self, conn):
"""``search`` returns plain dicts by calling ``model_dump`` on each point.

Raw ``ScoredPoint`` pydantic objects are not XCom-serializable; the hook's
job is to convert them at the boundary so callers get JSON-safe results.
"""
point_a = Mock(spec=["model_dump"])
point_a.model_dump.return_value = {"id": "a", "score": 0.9, "payload": {"text": "hi"}}
point_b = Mock(spec=["model_dump"])
point_b.model_dump.return_value = {"id": "b", "score": 0.7, "payload": {"text": "yo"}}
conn.query_points.return_value = Mock(points=[point_a, point_b])

results = self.qdrant_hook.search(
collection_name=self.collection_name,
query=[0.1, 0.2, 0.3],
limit=5,
)

assert results == [
{"id": "a", "score": 0.9, "payload": {"text": "hi"}},
{"id": "b", "score": 0.7, "payload": {"text": "yo"}},
]
point_a.model_dump.assert_called_once_with()
point_b.model_dump.assert_called_once_with()

@patch("airflow.providers.qdrant.hooks.qdrant.QdrantHook.conn")
def test_search_uses_query_points_and_forwards_arguments(self, conn):
"""``search`` calls ``query_points`` (modern API) with all arguments forwarded.

Guards against a regression to the deprecated ``search()`` API and against
silently dropping optional parameters.
"""
conn.query_points.return_value = Mock(points=[])
query_filter = Mock(name="filter")
search_params = Mock(name="params")

self.qdrant_hook.search(
collection_name=self.collection_name,
query=[1.0, 2.0, 3.0],
query_filter=query_filter,
search_params=search_params,
limit=7,
offset=2,
with_payload=["title"],
with_vectors=False,
score_threshold=0.5,
)

conn.search.assert_not_called()
conn.query_points.assert_called_once_with(
collection_name=self.collection_name,
query=[1.0, 2.0, 3.0],
query_filter=query_filter,
search_params=search_params,
limit=7,
offset=2,
with_payload=["title"],
with_vectors=False,
score_threshold=0.5,
)

@patch("airflow.providers.qdrant.hooks.qdrant.QdrantHook.conn")
def test_search_forwards_extra_kwargs(self, conn):
"""Extra ``**kwargs`` (e.g. ``using`` for named vectors) reach ``query_points``.

Keeps the hook forward-compatible with future ``query_points`` parameters
(hybrid search, named vectors) without needing to enumerate them here.
"""
conn.query_points.return_value = Mock(points=[])

self.qdrant_hook.search(
collection_name=self.collection_name,
query=[0.1, 0.2, 0.3],
using="text-embedding",
)

_, kwargs = conn.query_points.call_args
assert kwargs["using"] == "text-embedding"
Loading