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
3 changes: 3 additions & 0 deletions pymilvus/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@
TopHits,
)
from .client.search_result import Hit, Hits, SearchResult
from .client.telemetry import TelemetryConfig, new_client_request_id
from .client.types import (
BulkInsertState,
DataType,
Expand Down Expand Up @@ -142,6 +143,7 @@
"Shard",
"Status",
"StructFieldSchema",
"TelemetryConfig",
"TopHits",
"WeightedRanker",
"__version__",
Expand All @@ -167,6 +169,7 @@
"mkts_from_datetime",
"mkts_from_hybridts",
"mkts_from_unixtime",
"new_client_request_id",
"reset_password",
"transfer_node",
"transfer_replica",
Expand Down
164 changes: 163 additions & 1 deletion pymilvus/client/async_grpc_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,9 @@
import base64
import logging
import socket
import threading
import time
import weakref
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple, Union
from urllib import parse
Expand Down Expand Up @@ -46,6 +48,11 @@
from .embedding_list import EmbeddingList
from .prepare import Prepare
from .search_result import SearchResult
from .telemetry import (
AsyncClientTelemetryManager,
AsyncTelemetryUnaryUnaryInterceptor,
telemetry_operation,
)
from .types import (
AnalyzeResult,
CompactionState,
Expand Down Expand Up @@ -89,13 +96,36 @@ def __init__(
) -> None:
self._async_stub = None
self._async_channel = channel
self._client_telemetry_lock = threading.RLock()
self._client_telemetry_bindings = weakref.WeakKeyDictionary()
self._client_telemetry_generation = 0
self._client_owned_telemetry = bool(kwargs.pop("_client_owned_telemetry", False))

addr = kwargs.get("address")
self._address = addr if addr is not None else self.__get_address(uri, host, port)
self._log_level = None
self._user = kwargs.get("user")
self._connect_reserved = kwargs.get("option", {})
self._grpc_options = kwargs.get("grpc_options", {})
self._current_db_name = kwargs.get("db_name", "")
self._telemetry_stub = None
self._telemetry = None
self._telemetry_interceptor = None
if not self._client_owned_telemetry:
self._telemetry = AsyncClientTelemetryManager(
lambda: self._telemetry_stub,
kwargs.get("telemetry_config"),
user=self._user,
database_provider=lambda: self._current_db_name,
config_provider=lambda: {
"address": self._address,
"username": self._user or "",
"db_name": self._current_db_name,
"secure": self._secure,
},
runtime_client_id=kwargs.get("_telemetry_client_id", ""),
)
self._telemetry_interceptor = AsyncTelemetryUnaryUnaryInterceptor(self._telemetry)
self._set_authorization(**kwargs)
self._reconnect_lock = asyncio.Lock()
self._setup_grpc_channel(**kwargs)
Expand All @@ -104,6 +134,104 @@ def __init__(
self.callbacks = [] # Do nothing
self._server_info_cache = None

def _rebind_telemetry_stub(self, stub: Any) -> None:
Comment thread
xiaofan-luan marked this conversation as resolved.
client_lock = getattr(self, "_client_telemetry_lock", None)
if client_lock is None:
self._telemetry_stub = stub
telemetry = getattr(self, "_telemetry", None)
if telemetry is not None:
telemetry.rebind_stub(stub)
return
with client_lock:
self._telemetry_stub = stub
self._client_telemetry_generation += 1
generation = self._client_telemetry_generation
telemetry = getattr(self, "_telemetry", None)
bindings = list(self._client_telemetry_bindings.items())

if telemetry is not None:
self._stabilize_legacy_telemetry(telemetry, stub, generation)
for manager, (binding_token, _database) in bindings:
try:
self._stabilize_client_telemetry(manager, binding_token, stub, generation)
except BaseException:
logger.warning("failed to rebind logical async client telemetry", exc_info=True)

def _stabilize_legacy_telemetry(
self, telemetry: AsyncClientTelemetryManager, stub: Any, generation: int
) -> None:
while True:
telemetry.rebind_stub(stub)
with self._client_telemetry_lock:
if telemetry is not self._telemetry:
return
if generation == self._client_telemetry_generation:
return
stub = self._telemetry_stub
generation = self._client_telemetry_generation

def _stabilize_client_telemetry(
self,
manager: AsyncClientTelemetryManager,
binding_token: Any,
stub: Any,
generation: int,
) -> None:
while True:
manager.rebind_transport(stub, binding_token)
with self._client_telemetry_lock:
current = self._client_telemetry_bindings.get(manager)
if current is None or current[0] is not binding_token:
return
if generation == self._client_telemetry_generation:
return
stub = self._telemetry_stub
generation = self._client_telemetry_generation

def register_client_telemetry(
self, manager: AsyncClientTelemetryManager, database: str = ""
) -> None:
"""Attach one logical client's manager to this authenticated transport."""

with self._client_telemetry_lock:
existing = self._client_telemetry_bindings.get(manager)
if existing is not None and existing[1] == database:
return
binding_token = existing[0] if existing is not None else object()
self._client_telemetry_bindings[manager] = (binding_token, database)
stub = self._telemetry_stub
generation = self._client_telemetry_generation
try:
manager.bind_transport(stub, database, binding_token)
except BaseException:
with self._client_telemetry_lock:
current = self._client_telemetry_bindings.get(manager)
if current is not None and current[0] is binding_token:
self._client_telemetry_bindings.pop(manager, None)
manager.unbind_transport(binding_token)
raise

while True:
with self._client_telemetry_lock:
current = self._client_telemetry_bindings.get(manager)
detached = current is None or current[0] is not binding_token
stable = generation == self._client_telemetry_generation
if not stable:
stub = self._telemetry_stub
generation = self._client_telemetry_generation
if detached:
manager.unbind_transport(binding_token)
return
if stable:
return
manager.rebind_transport(stub, binding_token)

def unregister_client_telemetry(self, manager: AsyncClientTelemetryManager) -> None:
with self._client_telemetry_lock:
binding = self._client_telemetry_bindings.pop(manager, None)
if binding is not None:
manager.unbind_transport(binding[0])

def __get_address(self, uri: str, host: str, port: str) -> str:
if host != "" and port != "" and is_legal_host(host) and is_legal_port(port):
return f"{host}:{port}"
Expand Down Expand Up @@ -135,6 +263,9 @@ def __exit__(self: object, exc_type: object, exc_val: object, exc_tb: object):

async def close(self):
async with self._reconnect_lock:
if self._telemetry is not None:
await self._telemetry.stop_async()
self._rebind_telemetry_stub(None)
if self._async_channel:
await self._async_channel.close()
self._async_channel = None
Expand Down Expand Up @@ -184,6 +315,9 @@ async def reconnect(self, address: Optional[str] = None, timeout: float = 10):
self._final_channel = new_final_channel
self._async_stub = new_stub
self._async_identifier_interceptor = new_identifier_interceptor
self._rebind_telemetry_stub(
milvus_pb2_grpc.ClientTelemetryServiceStub(new_final_channel)
)
self._is_channel_ready = True

if old_channel:
Expand Down Expand Up @@ -285,6 +419,8 @@ def _build_stub(self, channel: Any, **kwargs: Any) -> Tuple[Any, Any]:
)
final_channel._unary_unary_interceptors.append(async_log_level_interceptor)
self._log_level = None
if self._telemetry_interceptor is not None:
final_channel._unary_unary_interceptors.append(self._telemetry_interceptor)
return final_channel, milvus_pb2_grpc.MilvusServiceStub(final_channel)

def _setup_grpc_channel(self, **kwargs):
Expand Down Expand Up @@ -316,7 +452,12 @@ async def ensure_channel_ready(self, timeout: Optional[float] = None):
timeout=wait_timeout,
)

self._rebind_telemetry_stub(
milvus_pb2_grpc.ClientTelemetryServiceStub(self._final_channel)
)
self._is_channel_ready = True
if self._telemetry is not None:
self._telemetry.start()
except (grpc.FutureTimeoutError, asyncio.TimeoutError, grpc.RpcError) as e:
raise MilvusException(
code=Status.CONNECT_FAILED,
Expand Down Expand Up @@ -344,6 +485,18 @@ async def _setup_identifier_interceptor_for_channel(
stub = milvus_pb2_grpc.MilvusServiceStub(final_channel)
return async_identifier_interceptor, final_channel, stub

@property
def telemetry(self) -> Optional[AsyncClientTelemetryManager]:
"""Return the legacy direct-handler telemetry manager, if enabled."""

return self._telemetry

@property
def telemetry_stub(self) -> Any:
"""Return the current authenticated heartbeat transport."""

return self._telemetry_stub

@retry_on_rpc_failure()
async def create_collection(
self,
Expand Down Expand Up @@ -733,6 +886,7 @@ async def release_collection(
)
check_status(response)

@telemetry_operation("Insert")
@retry_on_rpc_failure()
@retry_on_schema_mismatch()
async def insert_rows(
Expand Down Expand Up @@ -810,6 +964,7 @@ async def get_persistent_segment_infos(
check_status(response.status)
return response.infos

@telemetry_operation("Delete")
async def delete(
self,
collection_name: str,
Expand Down Expand Up @@ -880,6 +1035,7 @@ async def _prepare_batch_upsert_request(
)
)

@telemetry_operation("Upsert")
@retry_on_rpc_failure()
async def upsert(
self,
Expand Down Expand Up @@ -942,6 +1098,7 @@ async def _prepare_row_upsert_request(
field_ops=field_ops,
)

@telemetry_operation("Upsert")
@retry_on_rpc_failure()
@retry_on_schema_mismatch()
async def upsert_rows(
Expand Down Expand Up @@ -1004,6 +1161,7 @@ async def _execute_hybrid_search(
round_decimal = kwargs.get("round_decimal", -1)
return SearchResult(response.results, round_decimal, status=response.status)

@telemetry_operation("Search")
@retry_on_rpc_failure()
async def search(
self,
Expand Down Expand Up @@ -1084,6 +1242,7 @@ async def search(
request, timeout, context=context, round_decimal=round_decimal, **kwargs
)

@telemetry_operation("HybridSearch")
@retry_on_rpc_failure()
async def hybrid_search(
self,
Expand Down Expand Up @@ -1433,6 +1592,7 @@ async def list_partitions(
check_status(response.status)
return list(response.partition_names)

@telemetry_operation("Query")
@retry_on_rpc_failure()
async def query(
self,
Expand Down Expand Up @@ -1952,9 +2112,10 @@ async def list_aliases(
def reset_db_name(self, db_name: str):
"""Deprecated: db_name is now passed per-request via kwargs.

This method is kept for backward compatibility but does nothing.
This method is kept for backward compatibility and updates telemetry identity.
Use AsyncMilvusClient.use_database() instead.
"""
self._current_db_name = db_name

@retry_on_rpc_failure()
async def create_database(
Expand Down Expand Up @@ -2665,6 +2826,7 @@ async def get_compaction_state(
response.completedPlanNo,
)

@telemetry_operation("RunAnalyzer")
@retry_on_rpc_failure()
async def run_analyzer(
self,
Expand Down
Loading
Loading