diff --git a/pymilvus/bulk_writer/bulk_import.py b/pymilvus/bulk_writer/bulk_import.py index cd2b3fc93..5a31eb423 100644 --- a/pymilvus/bulk_writer/bulk_import.py +++ b/pymilvus/bulk_writer/bulk_import.py @@ -21,7 +21,7 @@ logger = logging.getLogger(__name__) -def _http_headers(api_key: str, db_name: str = ""): +def _http_headers(api_key: str, db_name: str = "", idempotency_key: str = ""): headers = { "User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_7_0) AppleWebKit/535.11 (KHTML, like Gecko) " "Chrome/17.0.963.56 Safari/535.11", @@ -32,6 +32,8 @@ def _http_headers(api_key: str, db_name: str = ""): } if db_name: headers["DB-Name"] = db_name + if idempotency_key: + headers["Idempotency-Key"] = idempotency_key return headers @@ -72,10 +74,11 @@ def _post_request( requests.Response: Response object. """ db_name = kwargs.pop("db_name", "") + idempotency_key = kwargs.pop("idempotency_key", "") try: resp = requests.post( url=url, - headers=_http_headers(api_key, db_name), + headers=_http_headers(api_key, db_name, idempotency_key), json=params, timeout=timeout, verify=verify, @@ -124,6 +127,7 @@ def bulk_import( data_paths: [List[List[str]]] = None, verify: Optional[Union[bool, str]] = True, cert: Optional[Union[str, tuple]] = None, + idempotency_key: str = "", **kwargs, ) -> requests.Response: """call bulkinsert restful interface to import files @@ -153,6 +157,9 @@ def bulk_import( or a string, which must be server's certificate path. Defaults to `True`. cert (str, tuple, optional): if String, path to ssl client cert file. if Tuple, ('cert', 'key') pair. + idempotency_key (str, optional): Sent as the ``Idempotency-Key`` header. A retry + carrying the same key resolves to the original import job instead of + creating a second one. Keep one key per logical request. Returns: response of the restful interface @@ -253,6 +260,7 @@ def bulk_import( verify=verify, cert=cert, db_name=db_name, + idempotency_key=idempotency_key, **kwargs, ) _handle_response(request_url, resp.json()) diff --git a/pymilvus/client/call_context.py b/pymilvus/client/call_context.py index b6e1574c3..fea2427ef 100644 --- a/pymilvus/client/call_context.py +++ b/pymilvus/client/call_context.py @@ -10,16 +10,22 @@ def _api_level_md(context: Optional["CallContext"]) -> Optional[list]: class CallContext: - def __init__(self, db_name: str = "", client_request_id: str = ""): + def __init__(self, db_name: str = "", client_request_id: str = "", idempotency_key: str = ""): self._db_name = db_name self._client_request_id = client_request_id + self._idempotency_key = idempotency_key def to_grpc_metadata(self): - return [ + metadata = [ ("dbname", self._db_name), ("client-request-id", self._client_request_id), ("client-request-unixmsec", current_time_ms()), ] + # Only sent when set: the server treats an empty metadata value as a + # present-but-empty key, which is not the same as no key. + if self._idempotency_key: + metadata.append(("idempotency-key", self._idempotency_key)) + return metadata def get_db_name(self): return self._db_name diff --git a/pymilvus/milvus_client/async_milvus_client.py b/pymilvus/milvus_client/async_milvus_client.py index b748d0c62..cd1fb4816 100644 --- a/pymilvus/milvus_client/async_milvus_client.py +++ b/pymilvus/milvus_client/async_milvus_client.py @@ -487,6 +487,8 @@ async def upsert( partition_name (str, optional): Name of the partition to upsert into. **kwargs (dict): Extra keyword arguments. + * *idempotency_key* (str, optional): Sent as the ``idempotency-key`` gRPC + metadata. Keep one key per logical request. * *partial_update* (bool, optional): Whether this is a partial update operation. If True, only the specified fields will be updated while others remain unchanged Default is False. diff --git a/pymilvus/milvus_client/base.py b/pymilvus/milvus_client/base.py index cff67cf0e..d110620d8 100644 --- a/pymilvus/milvus_client/base.py +++ b/pymilvus/milvus_client/base.py @@ -20,7 +20,11 @@ class BaseMilvusClient: def _generate_call_context(self, **kwargs) -> CallContext: client_request_id = kwargs.get("client_request_id") or kwargs.get("client-request-id", "") - return CallContext(db_name=self._config.db_name, client_request_id=client_request_id) + return CallContext( + db_name=self._config.db_name, + client_request_id=client_request_id, + idempotency_key=kwargs.get("idempotency_key", ""), + ) def _with_cluster_id(self, kwargs: Dict) -> Dict: cluster_id = getattr(self, "_cluster_id", "") diff --git a/pymilvus/milvus_client/milvus_client.py b/pymilvus/milvus_client/milvus_client.py index 0a3d43baa..a31d5260b 100644 --- a/pymilvus/milvus_client/milvus_client.py +++ b/pymilvus/milvus_client/milvus_client.py @@ -244,6 +244,11 @@ def insert( cast to list. timeout (float, optional): The timeout to use, will override init timeout. Defaults to None. + **kwargs (dict): Extra keyword arguments. + + * *idempotency_key* (str, optional): Sent as the ``idempotency-key`` gRPC + metadata. A retry carrying the same key is applied at most once and + returns the original result. Keep one key per logical request. Raises: DataNotMatchException: If the data has missing fields an exception will be thrown. @@ -302,6 +307,8 @@ def upsert( partition_name (str, optional): Name of the partition to upsert into. **kwargs (dict): Extra keyword arguments. + * *idempotency_key* (str, optional): Sent as the ``idempotency-key`` gRPC + metadata. Keep one key per logical request. * *partial_update* (bool, optional): Whether this is a partial update operation. If True, only the specified fields will be updated while others remain unchanged Default is False. @@ -836,6 +843,10 @@ def delete( filter(str, optional): A filter to use for the deletion. Defaults to none. timeout (int, optional): Timeout to use, overides the client level assigned at init. Defaults to None. + **kwargs (dict): Extra keyword arguments. + + * *idempotency_key* (str, optional): Sent as the ``idempotency-key`` gRPC + metadata. Keep one key per logical request. Note: You need to passin either ids or filter, and they cannot be used at the same time. diff --git a/pymilvus/orm/collection.py b/pymilvus/orm/collection.py index f06ec7b8e..17c07c9c6 100644 --- a/pymilvus/orm/collection.py +++ b/pymilvus/orm/collection.py @@ -183,7 +183,8 @@ def _get_connection(self, **kwargs): context creation in some code paths (e.g., num_entities property). Args: - **kwargs: Optional kwargs for context generation (e.g., client_request_id). + **kwargs: Optional kwargs for context generation (e.g., client_request_id, + idempotency_key). Returns: tuple: (handler, context) tuple where handler is GrpcHandler/AsyncGrpcHandler diff --git a/pymilvus/orm/connections.py b/pymilvus/orm/connections.py index c0ec815bf..12e08c310 100644 --- a/pymilvus/orm/connections.py +++ b/pymilvus/orm/connections.py @@ -589,7 +589,11 @@ def _generate_call_context(self, alias: str, **kwargs) -> CallContext: config = self._alias_config.get(alias, {}) db_name = config.get("db_name", "") req_id = kwargs.get("client_request_id") or kwargs.get("client-request-id", "") - return CallContext(db_name=db_name, client_request_id=req_id) + return CallContext( + db_name=db_name, + client_request_id=req_id, + idempotency_key=kwargs.get("idempotency_key", ""), + ) # Singleton Mode in Python diff --git a/tests/unit/test_bulk_import.py b/tests/unit/test_bulk_import.py index 65f25850e..af188e5c4 100644 --- a/tests/unit/test_bulk_import.py +++ b/tests/unit/test_bulk_import.py @@ -32,6 +32,14 @@ def test_with_db_name(self): assert headers["DB-Name"] == "my_db" assert headers["Authorization"] == "Bearer my-key" + def test_with_idempotency_key(self): + headers = _http_headers(api_key="my-key", idempotency_key="run-1-batch-1") + assert headers["Idempotency-Key"] == "run-1-batch-1" + + def test_with_empty_idempotency_key(self): + headers = _http_headers(api_key="my-key", idempotency_key="") + assert "Idempotency-Key" not in headers + class TestPostRequest: @patch.object(bulk_import_mod.requests, "post") @@ -68,6 +76,23 @@ def test_without_db_name_has_no_db_header(self, mock_post): _, kwargs = mock_post.call_args assert "DB-Name" not in kwargs["headers"] + @patch.object(bulk_import_mod.requests, "post") + def test_pops_idempotency_key_and_adds_header(self, mock_post): + mock_resp = MagicMock() + mock_resp.status_code = 200 + mock_post.return_value = mock_resp + + _post_request( + url="http://example.com/api", + api_key="my-key", + params={"foo": "bar"}, + idempotency_key="run-1-batch-1", + ) + + _, kwargs = mock_post.call_args + assert kwargs["headers"]["Idempotency-Key"] == "run-1-batch-1" + assert "idempotency_key" not in kwargs + class TestGetImportProgress: @patch.object(bulk_import_mod.requests, "post") @@ -172,6 +197,43 @@ def test_without_db_name_has_no_db_header(self, mock_post): assert "DB-Name" not in kwargs["headers"] assert kwargs["json"]["dbName"] == "" + @patch.object(bulk_import_mod.requests, "post") + def test_sends_idempotency_key_header(self, mock_post): + mock_resp = MagicMock() + mock_resp.status_code = 200 + mock_resp.json.return_value = {"code": 0, "data": {}} + mock_post.return_value = mock_resp + + bulk_import( + url="http://example.com", + collection_name="my_collection", + api_key="my-key", + files=[["file1.parquet"]], + idempotency_key="run-1-batch-1", + ) + + _, kwargs = mock_post.call_args + assert kwargs["headers"]["Idempotency-Key"] == "run-1-batch-1" + assert "idempotencyKey" not in kwargs["json"] + assert "idempotency_key" not in kwargs + + @patch.object(bulk_import_mod.requests, "post") + def test_without_idempotency_key_has_no_header(self, mock_post): + mock_resp = MagicMock() + mock_resp.status_code = 200 + mock_resp.json.return_value = {"code": 0, "data": {}} + mock_post.return_value = mock_resp + + bulk_import( + url="http://example.com", + collection_name="my_collection", + api_key="my-key", + files=[["file1.parquet"]], + ) + + _, kwargs = mock_post.call_args + assert "Idempotency-Key" not in kwargs["headers"] + class TestListImportJobs: @patch.object(bulk_import_mod.requests, "post") diff --git a/tests/unit/test_idempotency_key.py b/tests/unit/test_idempotency_key.py new file mode 100644 index 000000000..45c7cbfc0 --- /dev/null +++ b/tests/unit/test_idempotency_key.py @@ -0,0 +1,143 @@ +"""Idempotency key propagation: CallContext -> gRPC metadata, and the client entry points.""" + +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +import pytest_asyncio +from pymilvus import AsyncMilvusClient, DataType, MilvusClient, connections +from pymilvus.client.call_context import CallContext +from pymilvus.client.connection_manager import AsyncConnectionManager, ConnectionManager + +IDEMPOTENCY_KEY_HEADER = "idempotency-key" +_SCHEMA = {"fields": [{"name": "id", "is_primary": True, "type": DataType.INT64}]} + + +def _metadata_values(context: CallContext, key: str): + return [v for k, v in context.to_grpc_metadata() if k == key] + + +def _context_of(mock_call): + _, kwargs = mock_call.call_args + return kwargs["context"] + + +@pytest.fixture(autouse=True) +def _reset_connection_managers(): + ConnectionManager._reset_instance() + AsyncConnectionManager._reset_instance() + yield + ConnectionManager._reset_instance() + AsyncConnectionManager._reset_instance() + + +class TestCallContext: + def test_metadata_carries_idempotency_key(self): + ctx = CallContext(db_name="db", idempotency_key="order-4711") + assert _metadata_values(ctx, IDEMPOTENCY_KEY_HEADER) == ["order-4711"] + + def test_metadata_omits_absent_key(self): + ctx = CallContext(db_name="db") + assert _metadata_values(ctx, IDEMPOTENCY_KEY_HEADER) == [] + + def test_metadata_omits_empty_key(self): + ctx = CallContext(db_name="db", idempotency_key="") + assert _metadata_values(ctx, IDEMPOTENCY_KEY_HEADER) == [] + + +def _sync_handler(): + handler = MagicMock() + handler.get_server_type.return_value = "milvus" + handler._wait_for_channel_ready = MagicMock() + handler._get_schema.return_value = (_SCHEMA, 100) + return handler + + +class TestMilvusClient: + def test_insert_forwards_idempotency_key(self): + handler = _sync_handler() + handler.insert_rows.return_value = MagicMock(insert_count=1, primary_keys=[1], cost=0) + with patch("pymilvus.client.grpc_handler.GrpcHandler", return_value=handler): + MilvusClient().insert("col", {"id": 1}, idempotency_key="order-4711") + + ctx = _context_of(handler.insert_rows) + assert _metadata_values(ctx, IDEMPOTENCY_KEY_HEADER) == ["order-4711"] + + def test_insert_without_key_sends_no_header(self): + handler = _sync_handler() + handler.insert_rows.return_value = MagicMock(insert_count=1, primary_keys=[1], cost=0) + with patch("pymilvus.client.grpc_handler.GrpcHandler", return_value=handler): + MilvusClient().insert("col", {"id": 1}) + + ctx = _context_of(handler.insert_rows) + assert _metadata_values(ctx, IDEMPOTENCY_KEY_HEADER) == [] + + def test_upsert_forwards_idempotency_key(self): + handler = _sync_handler() + handler.upsert_rows.return_value = MagicMock(upsert_count=1, primary_keys=[1], cost=0) + with patch("pymilvus.client.grpc_handler.GrpcHandler", return_value=handler): + MilvusClient().upsert("col", {"id": 1}, idempotency_key="order-4711") + + ctx = _context_of(handler.upsert_rows) + assert _metadata_values(ctx, IDEMPOTENCY_KEY_HEADER) == ["order-4711"] + + def test_delete_forwards_idempotency_key(self): + handler = _sync_handler() + handler.delete.return_value = MagicMock(delete_count=1, primary_keys=[], cost=0) + with patch("pymilvus.client.grpc_handler.GrpcHandler", return_value=handler): + MilvusClient().delete("col", ids=[1], idempotency_key="order-4711") + + ctx = _context_of(handler.delete) + assert _metadata_values(ctx, IDEMPOTENCY_KEY_HEADER) == ["order-4711"] + + +@pytest_asyncio.fixture +async def async_client_and_handler(): + handler = MagicMock() + handler.ensure_channel_ready = AsyncMock() + handler._get_schema = AsyncMock(return_value=(_SCHEMA, 100)) + with patch("pymilvus.client.async_grpc_handler.AsyncGrpcHandler", return_value=handler): + client = AsyncMilvusClient() + await client._connect() + yield client, handler + + +class TestAsyncMilvusClient: + @pytest.mark.asyncio + async def test_insert_forwards_idempotency_key(self, async_client_and_handler): + client, handler = async_client_and_handler + handler.insert_rows = AsyncMock( + return_value=MagicMock(insert_count=1, primary_keys=[1], cost=0) + ) + + await client.insert("col", {"id": 1}, idempotency_key="order-4711") + + ctx = _context_of(handler.insert_rows) + assert _metadata_values(ctx, IDEMPOTENCY_KEY_HEADER) == ["order-4711"] + + @pytest.mark.asyncio + async def test_upsert_forwards_idempotency_key(self, async_client_and_handler): + client, handler = async_client_and_handler + handler.upsert_rows = AsyncMock( + return_value=MagicMock(upsert_count=1, primary_keys=[1], cost=0) + ) + + await client.upsert("col", {"id": 1}, idempotency_key="order-4711") + + ctx = _context_of(handler.upsert_rows) + assert _metadata_values(ctx, IDEMPOTENCY_KEY_HEADER) == ["order-4711"] + + @pytest.mark.asyncio + async def test_delete_forwards_idempotency_key(self, async_client_and_handler): + client, handler = async_client_and_handler + handler.delete = AsyncMock(return_value=MagicMock(delete_count=1, primary_keys=[], cost=0)) + + await client.delete("col", ids=[1], idempotency_key="order-4711") + + ctx = _context_of(handler.delete) + assert _metadata_values(ctx, IDEMPOTENCY_KEY_HEADER) == ["order-4711"] + + +class TestOrmConnections: + def test_generate_call_context_carries_idempotency_key(self): + ctx = connections._generate_call_context("unknown_alias", idempotency_key="order-4711") + assert _metadata_values(ctx, IDEMPOTENCY_KEY_HEADER) == ["order-4711"]