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
12 changes: 10 additions & 2 deletions pymilvus/bulk_writer/bulk_import.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand All @@ -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


Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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())
Expand Down
10 changes: 8 additions & 2 deletions pymilvus/client/call_context.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
2 changes: 2 additions & 0 deletions pymilvus/milvus_client/async_milvus_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
6 changes: 5 additions & 1 deletion pymilvus/milvus_client/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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", "")
Expand Down
11 changes: 11 additions & 0 deletions pymilvus/milvus_client/milvus_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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.

Expand Down
3 changes: 2 additions & 1 deletion pymilvus/orm/collection.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
6 changes: 5 additions & 1 deletion pymilvus/orm/connections.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
62 changes: 62 additions & 0 deletions tests/unit/test_bulk_import.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down Expand Up @@ -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")
Expand Down Expand Up @@ -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")
Expand Down
143 changes: 143 additions & 0 deletions tests/unit/test_idempotency_key.py
Original file line number Diff line number Diff line change
@@ -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"]
Loading