diff --git a/setup.cfg b/setup.cfg index 088736f..72beb40 100644 --- a/setup.cfg +++ b/setup.cfg @@ -56,6 +56,7 @@ enable_error_code = redundant-self, explicit_package_bases = true +mypy_path = src ignore_missing_imports = true strict = true warn_unreachable = true diff --git a/src/typesense/async_/analytics_rule_v1.py b/src/typesense/async_/analytics_rule_v1.py index d640853..5584623 100644 --- a/src/typesense/async_/analytics_rule_v1.py +++ b/src/typesense/async_/analytics_rule_v1.py @@ -74,11 +74,9 @@ async def retrieve( Union[RuleSchemaForQueries, RuleSchemaForCounters]: The schema containing the rule details. """ - response: typing.Union[ - RuleSchemaForQueries, RuleSchemaForCounters - ] = await self.api_call.get( + response = await self.api_call.get( self._endpoint_path, - entity_type=dict, + entity_type=typing.Dict[str, typing.Any], as_json=True, ) return typing.cast( @@ -101,7 +99,7 @@ async def delete(self) -> RuleDeleteSchema: return response @property - @warn_deprecation( # type: ignore[untyped-decorator] + @warn_deprecation( "AsyncAnalyticsRuleV1 is deprecated on v30+. Use client.analytics.rules[rule_id] instead.", flag_name="analytics_rules_v1_deprecation", ) diff --git a/src/typesense/async_/analytics_rules_v1.py b/src/typesense/async_/analytics_rules_v1.py index 1aac207..2e905e4 100644 --- a/src/typesense/async_/analytics_rules_v1.py +++ b/src/typesense/async_/analytics_rules_v1.py @@ -89,7 +89,7 @@ def __getitem__(self, rule_id: str) -> AsyncAnalyticsRuleV1: self.rules[rule_id] = AsyncAnalyticsRuleV1(self.api_call, rule_id) return self.rules[rule_id] - @warn_deprecation( # type: ignore[untyped-decorator] + @warn_deprecation( "AsyncAnalyticsRulesV1 is deprecated on v30+. Use client.analytics instead.", flag_name="analytics_rules_v1_deprecation", ) @@ -115,21 +115,19 @@ async def create( The created rule. Returns RuleSchemaForCounters for counter rules and RuleSchemaForQueries for query rules. """ - response: typing.Union[ - RuleSchemaForCounters, RuleSchemaForQueries - ] = await self.api_call.post( + response = await self.api_call.post( AsyncAnalyticsRulesV1.resource_path, body=rule, params=rule_parameters, as_json=True, - entity_type=dict, + entity_type=typing.Dict[str, typing.Any], ) return typing.cast( typing.Union[RuleSchemaForCounters, RuleSchemaForQueries], response, ) - @warn_deprecation( # type: ignore[untyped-decorator] + @warn_deprecation( "AsyncAnalyticsRulesV1 is deprecated on v30+. Use client.analytics instead.", flag_name="analytics_rules_v1_deprecation", ) @@ -148,19 +146,17 @@ async def upsert( Returns: Union[RuleSchemaForCounters, RuleCreateSchemaForQueries]: The upserted rule. """ - response: typing.Union[ - RuleSchemaForCounters, RuleCreateSchemaForQueries - ] = await self.api_call.put( + response = await self.api_call.put( "/".join([self.resource_path, rule_id]), body=rule, - entity_type=dict, + entity_type=typing.Dict[str, typing.Any], ) return typing.cast( typing.Union[RuleSchemaForCounters, RuleCreateSchemaForQueries], response, ) - @warn_deprecation( # type: ignore[untyped-decorator] + @warn_deprecation( "AsyncAnalyticsRulesV1 is deprecated on v30+. Use client.analytics instead.", flag_name="analytics_rules_v1_deprecation", ) diff --git a/src/typesense/async_/api_call.py b/src/typesense/async_/api_call.py index be1a83d..9bc98ec 100644 --- a/src/typesense/async_/api_call.py +++ b/src/typesense/async_/api_call.py @@ -37,6 +37,7 @@ import httpx +from typesense.concurrency_limit import AsyncConcurrencyLimit from typesense.configuration import Configuration, Node from typesense.exceptions import ( HTTPStatus0Error, @@ -59,7 +60,7 @@ import typing_extensions as typing TEntityDict = typing.TypeVar("TEntityDict") -TParams = typing.TypeVar("TParams", bound=typing.Dict[str, typing.Any]) +TParams = typing.TypeVar("TParams", bound=typing.Mapping[str, object]) TBody = typing.TypeVar( "TBody", bound=typing.Union[str, bytes, typing.Mapping[str, typing.Any]] ) @@ -94,7 +95,7 @@ class SessionFunctionKwargs(typing.Generic[TParams, TBody], typing.TypedDict): params: typing.NotRequired[typing.Union[TParams, None]] data: typing.NotRequired[typing.Union[TBody, None]] - content: typing.NotRequired[typing.Union[TBody, str, None]] + content: typing.NotRequired[typing.Union[str, bytes, None]] headers: typing.NotRequired[typing.Dict[str, str]] timeout: typing.NotRequired[float] @@ -135,6 +136,20 @@ class SessionFunctionKwargs(typing.Generic[TParams, TBody], typing.TypedDict): ServiceUnavailable, ) +_CLIENT_ERRORS: typing.Final[ + typing.Tuple[ + typing.Type[httpx.PoolTimeout], + typing.Type[httpx.LocalProtocolError], + typing.Type[httpx.DecodingError], + typing.Type[httpx.TooManyRedirects], + ] +] = ( + httpx.PoolTimeout, + httpx.LocalProtocolError, + httpx.DecodingError, + httpx.TooManyRedirects, +) + class AsyncApiCall: """ @@ -160,7 +175,17 @@ def __init__(self, config: Configuration): self.node_manager = NodeManager(config) self.request_handler = RequestHandler(config) self._client = httpx.AsyncClient( - timeout=config.connection_timeout_seconds, + timeout=httpx.Timeout( + config.connection_timeout_seconds, + pool=config.pool_timeout_seconds, + ), + limits=httpx.Limits( + max_connections=config.max_connections, + max_keepalive_connections=config.max_keepalive_connections, + ), + ) + self._concurrency_limit = AsyncConcurrencyLimit( + config.max_concurrent_requests, ) async def __aenter__(self) -> "AsyncApiCall": @@ -473,11 +498,14 @@ async def _execute_request( try: return await self._make_request_and_process_response( method, + node, url, entity_type, as_json, **request_kwargs, ) + except _CLIENT_ERRORS: + raise except _SERVER_ERRORS as server_error: self.node_manager.set_node_health(node, is_healthy=False) if num_retries < self.config.num_retries: @@ -495,24 +523,23 @@ async def _execute_request( async def _make_request_and_process_response( self, method: str, + node: Node, url: str, entity_type: typing.Type[TEntityDict], as_json: bool, **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], ) -> typing.Union[TEntityDict, str]: - """Make the async API request and process the response.""" - request_response = await self.request_handler.make_request( - method=method, - url=url, - as_json=as_json, - entity_type=entity_type, - client=self._client, - **kwargs, - ) - self.node_manager.set_node_health( - self.node_manager.get_node(), - is_healthy=True, - ) + """Make the async API request to `node` and process the response.""" + async with self._concurrency_limit: + request_response = await self.request_handler.make_request( + method=method, + url=url, + as_json=as_json, + entity_type=entity_type, + client=self._client, + **kwargs, + ) + self.node_manager.set_node_health(node, is_healthy=True) return ( typing.cast(TEntityDict, request_response) if as_json diff --git a/src/typesense/async_/client.py b/src/typesense/async_/client.py index 1ecb807..8175bcb 100644 --- a/src/typesense/async_/client.py +++ b/src/typesense/async_/client.py @@ -164,5 +164,5 @@ def typed_collection( """ if name is None: name = model.__name__.lower() - collection: AsyncCollection[TDoc] = self.collections[name] - return collection + # ``collections`` is typed for the default DocumentSchema; narrow it to the model. + return typing.cast(AsyncCollection[TDoc], self.collections[name]) diff --git a/src/typesense/async_/documents.py b/src/typesense/async_/documents.py index 8228762..399c82d 100644 --- a/src/typesense/async_/documents.py +++ b/src/typesense/async_/documents.py @@ -61,6 +61,17 @@ None, ] +# One line of an import response. ``ImportResponse`` is a union of lists, one per +# return mode, so the helpers below build a list of these and ``import_`` casts it +# to the list type its overloads promise. +_ImportResponseItem = typing.Union[ + ImportResponseWithDoc[TDoc], + ImportResponseWithId, + ImportResponseWithDocAndId[TDoc], + ImportResponseSuccess, + ImportResponseFail[TDoc], +] + class AsyncDocuments(typing.Generic[TDoc]): """ @@ -125,12 +136,14 @@ async def create( Returns: TDoc: The created document. """ - dirty_values_parameters = dirty_values_parameters or {} - dirty_values_parameters["action"] = "create" + write_parameters: typing.Dict[str, object] = { + **(dirty_values_parameters or {}), + "action": "create", + } response = await self.api_call.post( self._endpoint_path(), body=document, - params=dirty_values_parameters, + params=write_parameters, as_json=True, entity_type=typing.Dict[str, str], ) @@ -154,7 +167,14 @@ async def create_many( The list of import responses. """ logger.warn("`create_many` is deprecated: please use `import_`.") - return await self.import_(documents, dirty_values_parameters) + # Dirty values parameters are a subset of the write parameters. + return await self.import_( + documents, + typing.cast( + typing.Optional[DocumentWriteParameters], + dirty_values_parameters, + ), + ) async def upsert( self, @@ -172,12 +192,14 @@ async def upsert( Returns: TDoc: The upserted document. """ - dirty_values_parameters = dirty_values_parameters or {} - dirty_values_parameters["action"] = "upsert" + write_parameters: typing.Dict[str, object] = { + **(dirty_values_parameters or {}), + "action": "upsert", + } response = await self.api_call.post( self._endpoint_path(), body=document, - params=dirty_values_parameters, + params=write_parameters, as_json=True, entity_type=typing.Dict[str, str], ) @@ -199,12 +221,14 @@ async def update( Returns: UpdateByFilterResponse: The response containing information about the update. """ - dirty_values_parameters = dirty_values_parameters or {} - dirty_values_parameters["action"] = "update" + update_parameters: typing.Dict[str, object] = { + **(dirty_values_parameters or {}), + "action": "update", + } response: UpdateByFilterResponse = await self.api_call.patch( self._endpoint_path(), body=document, - params=dirty_values_parameters, + params=update_parameters, entity_type=UpdateByFilterResponse, ) return response @@ -301,9 +325,14 @@ async def import_( return await self._import_raw(documents, import_parameters) if batch_size: - return await self._batch_import(documents, import_parameters, batch_size) - - return await self._bulk_import(documents, import_parameters) + response_objs = await self._batch_import( + documents, + import_parameters, + batch_size, + ) + else: + response_objs = await self._bulk_import(documents, import_parameters) + return typing.cast(ImportResponse[TDoc], response_objs) async def export( self, @@ -410,9 +439,9 @@ async def _batch_import( documents: typing.List[TDoc], import_parameters: _ImportParameters, batch_size: int, - ) -> ImportResponse[TDoc]: + ) -> typing.List[_ImportResponseItem[TDoc]]: """Import documents in batches.""" - response_objs: ImportResponse[TDoc] = [] + response_objs: typing.List[_ImportResponseItem[TDoc]] = [] for batch_index in range(0, len(documents), batch_size): batch = documents[batch_index : batch_index + batch_size] api_response = await self._bulk_import(batch, import_parameters) @@ -423,7 +452,7 @@ async def _bulk_import( self, documents: typing.List[TDoc], import_parameters: _ImportParameters, - ) -> ImportResponse[TDoc]: + ) -> typing.List[_ImportResponseItem[TDoc]]: """Import a list of documents in bulk.""" document_strs = [json.dumps(doc) for doc in documents] if not document_strs: @@ -439,9 +468,12 @@ async def _bulk_import( ) return self._parse_import_response(res) - def _parse_import_response(self, response: str) -> ImportResponse[TDoc]: + def _parse_import_response( + self, + response: str, + ) -> typing.List[_ImportResponseItem[TDoc]]: """Parse the import response string into a list of response objects.""" - response_objs: typing.List[ImportResponse] = [] + response_objs: typing.List[_ImportResponseItem[TDoc]] = [] for res_obj_str in response.split("\n"): try: res_obj_json = json.loads(res_obj_str) diff --git a/src/typesense/async_/keys.py b/src/typesense/async_/keys.py index 0dd8d94..4639113 100644 --- a/src/typesense/async_/keys.py +++ b/src/typesense/async_/keys.py @@ -29,7 +29,6 @@ ApiKeyCreateResponseSchema, ApiKeyCreateSchema, ApiKeyRetrieveSchema, - ApiKeySchema, ) if sys.version_info >= (3, 11): @@ -103,11 +102,11 @@ async def create(self, schema: ApiKeyCreateSchema) -> ApiKeyCreateResponseSchema ... } ... ) """ - response: ApiKeySchema = await self.api_call.post( + response: ApiKeyCreateResponseSchema = await self.api_call.post( AsyncKeys.resource_path, as_json=True, body=schema, - entity_type=ApiKeySchema, + entity_type=ApiKeyCreateResponseSchema, ) return response diff --git a/src/typesense/async_/operations.py b/src/typesense/async_/operations.py index ca61a1f..4a36608 100644 --- a/src/typesense/async_/operations.py +++ b/src/typesense/async_/operations.py @@ -60,8 +60,10 @@ def __init__(self, api_call: AsyncApiCall): """ self.api_call = api_call + # The generic ``str`` overload below also matches "schema_changes"; overloads are + # tried in order, so this one wins. @typing.overload - async def perform( + async def perform( # type: ignore[overload-overlap] self, operation_name: typing.Literal["schema_changes"], query_params: None = None, @@ -132,36 +134,36 @@ async def perform( @typing.overload async def perform( self, - operation_name: str, - query_params: typing.Union[typing.Dict[str, str], None] = None, + operation_name: typing.Literal["snapshot"], + query_params: SnapshotParameters, ) -> OperationResponse: """ - Perform a generic operation. + Perform a snapshot operation. Args: - operation_name (str): The name of the operation. - query_params (Union[Dict[str, str], None], optional): - Query parameters for the operation. + operation_name (Literal["snapshot"]): The name of the operation. + query_params (SnapshotParameters): Query parameters for the snapshot operation. Returns: - OperationResponse: The response from the operation. + OperationResponse: The response from the snapshot operation. """ @typing.overload async def perform( self, - operation_name: typing.Literal["snapshot"], - query_params: SnapshotParameters, + operation_name: str, + query_params: typing.Union[typing.Dict[str, str], None] = None, ) -> OperationResponse: """ - Perform a snapshot operation. + Perform a generic operation. Args: - operation_name (Literal["snapshot"]): The name of the operation. - query_params (SnapshotParameters): Query parameters for the snapshot operation. + operation_name (str): The name of the operation. + query_params (Union[Dict[str, str], None], optional): + Query parameters for the operation. Returns: - OperationResponse: The response from the snapshot operation. + OperationResponse: The response from the operation. """ async def perform( @@ -181,7 +183,7 @@ async def perform( typing.Dict[str, str], None, ] = None, - ) -> OperationResponse: + ) -> typing.Union[OperationResponse, typing.List[SchemaChangesResponse]]: """ Perform an operation on the Typesense API. @@ -202,13 +204,16 @@ async def perform( >>> response = await operations.perform("vote") >>> health = await operations.is_healthy() """ - response: OperationResponse = await self.api_call.post( + response = await self.api_call.post( self._endpoint_path(operation_name), params=query_params, as_json=True, - entity_type=OperationResponse, + entity_type=object, + ) + return typing.cast( + typing.Union[OperationResponse, typing.List[SchemaChangesResponse]], + response, ) - return response async def is_healthy(self) -> bool: """ @@ -222,16 +227,14 @@ async def is_healthy(self) -> bool: >>> healthy = await operations.is_healthy() >>> print(healthy) """ - call_resp: HealthCheckResponse = await self.api_call.get( + call_resp: object = await self.api_call.get( AsyncOperations.health_path, as_json=True, entity_type=HealthCheckResponse, ) - if isinstance(call_resp, typing.Dict): - is_ok: bool = call_resp.get("ok", False) - else: - is_ok = False - return is_ok + if isinstance(call_resp, dict): + return bool(call_resp.get("ok", False)) + return False async def toggle_slow_request_log( self, diff --git a/src/typesense/async_/override.py b/src/typesense/async_/override.py index 58e5a26..3e3e0f8 100644 --- a/src/typesense/async_/override.py +++ b/src/typesense/async_/override.py @@ -87,7 +87,7 @@ async def delete(self) -> OverrideDeleteSchema: return response @property - @warn_deprecation( # type: ignore[untyped-decorator] + @warn_deprecation( "The override API (collections/{collection}/overrides/{override_id}) is deprecated is removed on v30+. " "Use curation sets (curation_sets) instead.", flag_name="overrides_deprecation", diff --git a/src/typesense/async_/overrides.py b/src/typesense/async_/overrides.py index b8e725b..99d1190 100644 --- a/src/typesense/async_/overrides.py +++ b/src/typesense/async_/overrides.py @@ -129,7 +129,7 @@ async def retrieve(self) -> OverrideRetrieveSchema: ) return response - @warn_deprecation( # type: ignore[untyped-decorator] + @warn_deprecation( "AsyncOverrides is deprecated on v30+. Use client.curation_sets instead.", flag_name="overrides_deprecation", ) diff --git a/src/typesense/async_/synonym.py b/src/typesense/async_/synonym.py index 3ad6bc2..73cd46c 100644 --- a/src/typesense/async_/synonym.py +++ b/src/typesense/async_/synonym.py @@ -79,7 +79,7 @@ async def delete(self) -> SynonymDeleteSchema: ) @property - @warn_deprecation( # type: ignore[untyped-decorator] + @warn_deprecation( "The synonym API (collections/{collection}/synonyms/{synonym_id}) is deprecated is removed on v30+. " "Use synonym sets (synonym_sets) instead.", flag_name="synonyms_deprecation", diff --git a/src/typesense/async_/synonyms.py b/src/typesense/async_/synonyms.py index 027172e..ea2b3aa 100644 --- a/src/typesense/async_/synonyms.py +++ b/src/typesense/async_/synonyms.py @@ -124,7 +124,7 @@ async def retrieve(self) -> SynonymsRetrieveSchema: ) return response - @warn_deprecation( # type: ignore[untyped-decorator] + @warn_deprecation( "The synonyms API (collections/{collection}/synonyms) is deprecated is removed on v30+. " "Use synonym sets (synonym_sets) instead.", flag_name="synonyms_deprecation", diff --git a/src/typesense/concurrency_limit.py b/src/typesense/concurrency_limit.py new file mode 100644 index 0000000..34ebdc5 --- /dev/null +++ b/src/typesense/concurrency_limit.py @@ -0,0 +1,89 @@ +""" +Optional caps on the number of requests a client sends at once. + +``AsyncConcurrencyLimit`` is used by the async client and ``ConcurrencyLimit`` by the +sync client (``utils/run-unasync.py`` maps one name to the other). Both are no-ops +when ``max_concurrent_requests`` is ``None``. + +Keeping the cap below the httpx pool's ``max_connections`` means requests queue here +instead of in the pool, so a burst of slow requests cannot exhaust the pool and +raise ``httpx.PoolTimeout``. +""" + +import asyncio +import sys +import threading +from types import TracebackType + +if sys.version_info >= (3, 11): + import typing +else: + import typing_extensions as typing + + +class AsyncConcurrencyLimit: + """Async context manager that holds a slot for the duration of a request.""" + + def __init__(self, max_concurrent_requests: typing.Optional[int]) -> None: + """ + Initialize the limit. + + Args: + max_concurrent_requests (Optional[int]): The maximum number of requests + in flight at once, or ``None`` for no limit. + """ + self._max_concurrent_requests = max_concurrent_requests + # Created on first use, inside the running event loop. On Python < 3.10 a + # semaphore binds to the loop that is current when it is constructed. + self._semaphore: typing.Optional[asyncio.Semaphore] = None + + async def __aenter__(self) -> None: + """Wait for a free slot.""" + if self._max_concurrent_requests is None: + return + if self._semaphore is None: + self._semaphore = asyncio.Semaphore(self._max_concurrent_requests) + await self._semaphore.acquire() + + async def __aexit__( + self, + exc_type: typing.Optional[typing.Type[BaseException]], + exc_val: typing.Optional[BaseException], + exc_tb: typing.Optional[TracebackType], + ) -> None: + """Release the slot.""" + if self._semaphore is not None: + self._semaphore.release() + + +class ConcurrencyLimit: + """Context manager that holds a slot for the duration of a request.""" + + def __init__(self, max_concurrent_requests: typing.Optional[int]) -> None: + """ + Initialize the limit. + + Args: + max_concurrent_requests (Optional[int]): The maximum number of requests + in flight at once, or ``None`` for no limit. + """ + self._semaphore: typing.Optional[threading.Semaphore] = ( + None + if max_concurrent_requests is None + else threading.Semaphore(max_concurrent_requests) + ) + + def __enter__(self) -> None: + """Wait for a free slot.""" + if self._semaphore is not None: + self._semaphore.acquire() + + def __exit__( + self, + exc_type: typing.Optional[typing.Type[BaseException]], + exc_val: typing.Optional[BaseException], + exc_tb: typing.Optional[TracebackType], + ) -> None: + """Release the slot.""" + if self._semaphore is not None: + self._semaphore.release() diff --git a/src/typesense/configuration.py b/src/typesense/configuration.py index aaa741e..4f9144a 100644 --- a/src/typesense/configuration.py +++ b/src/typesense/configuration.py @@ -82,6 +82,23 @@ class ConfigDict(typing.TypedDict): connection_timeout_seconds (float): The connection timeout in seconds. suppress_deprecation_warnings (bool): Whether to suppress deprecation warnings. + + pool_timeout_seconds (float): How long a request waits for a free connection + in the pool before raising ``httpx.PoolTimeout``. Defaults to + ``connection_timeout_seconds``. Setting it lower than + ``connection_timeout_seconds`` makes the httpcore connection leak + (encode/httpcore#1093) more likely under load. + + max_connections (int): The maximum number of connections in the pool. + Defaults to 100. + + max_keepalive_connections (int): The maximum number of idle connections + kept alive in the pool. Defaults to 20. + + max_concurrent_requests (int): The maximum number of requests in flight at + once; further requests wait for a slot. Keep it below + ``max_connections`` so a burst of slow requests cannot exhaust the pool. + Defaults to no limit. """ nodes: typing.List[typing.Union[str, NodeConfigDict]] @@ -100,6 +117,10 @@ class ConfigDict(typing.TypedDict): ] # deprecated connection_timeout_seconds: typing.NotRequired[float] suppress_deprecation_warnings: typing.NotRequired[bool] + pool_timeout_seconds: typing.NotRequired[float] + max_connections: typing.NotRequired[int] + max_keepalive_connections: typing.NotRequired[int] + max_concurrent_requests: typing.NotRequired[int] class Node: @@ -188,6 +209,10 @@ class Configuration: retry_interval_seconds (float): The interval in seconds between retries. healthcheck_interval_seconds (int): The interval in seconds between health checks. verify (bool): Whether to verify the SSL certificate. + pool_timeout_seconds (float): How long to wait for a free pooled connection. + max_connections (int): The maximum number of connections in the pool. + max_keepalive_connections (int): The maximum number of idle pooled connections. + max_concurrent_requests (int | None): The maximum number of requests in flight. """ def __init__( @@ -232,6 +257,18 @@ def __init__( self.suppress_deprecation_warnings = config_dict.get( "suppress_deprecation_warnings", False ) + self.pool_timeout_seconds = config_dict.get( + "pool_timeout_seconds", + self.connection_timeout_seconds, + ) + self.max_connections = config_dict.get("max_connections", 100) + self.max_keepalive_connections = config_dict.get( + "max_keepalive_connections", + 20, + ) + self.max_concurrent_requests: typing.Optional[int] = config_dict.get( + "max_concurrent_requests", + ) def _handle_nearest_node( self, @@ -295,6 +332,32 @@ def validate_config_dict(config_dict: ConfigDict) -> None: if nearest_node: ConfigurationValidations.validate_nearest_node(nearest_node) + ConfigurationValidations.validate_connection_pool(config_dict) + + @staticmethod + def validate_connection_pool(config_dict: ConfigDict) -> None: + """ + Validate the connection pool and concurrency settings. + + Args: + config_dict (ConfigDict): The configuration dictionary to validate. + + Raises: + ConfigError: If a pool or concurrency setting is out of range. + """ + positive_settings: typing.Dict[str, typing.Optional[float]] = { + "pool_timeout_seconds": config_dict.get("pool_timeout_seconds"), + "max_connections": config_dict.get("max_connections"), + "max_concurrent_requests": config_dict.get("max_concurrent_requests"), + } + for key, config_value in positive_settings.items(): + if config_value is not None and config_value <= 0: + raise ConfigError(f"`{key}` must be greater than 0.") + + max_keepalive_connections = config_dict.get("max_keepalive_connections") + if max_keepalive_connections is not None and max_keepalive_connections < 0: + raise ConfigError("`max_keepalive_connections` must not be negative.") + @staticmethod def validate_required_config_fields(config_dict: ConfigDict) -> None: """ diff --git a/src/typesense/preprocess.py b/src/typesense/preprocess.py index b45db0c..15e13d7 100644 --- a/src/typesense/preprocess.py +++ b/src/typesense/preprocess.py @@ -110,7 +110,9 @@ def process_param_list( return ",".join(stringified_list) -def stringify_search_params(parameter_dict: ParamSchema) -> StringifiedParamSchema: +def stringify_search_params( + parameter_dict: typing.Mapping[str, object], +) -> StringifiedParamSchema: """ Convert the search parameters to strings. @@ -118,7 +120,8 @@ def stringify_search_params(parameter_dict: ParamSchema) -> StringifiedParamSche to their string representations. List values are converted to comma-separated strings. Args: - parameter_dict (ParamSchema): The search parameters. + parameter_dict (Mapping[str, object]): The search parameters, e.g. a + ``SearchParameters`` TypedDict or a ``ParamSchema`` dictionary. Returns: StringifiedParamSchema: The search parameters as strings. diff --git a/src/typesense/request_handler.py b/src/typesense/request_handler.py index 38e6c24..1e8d81e 100644 --- a/src/typesense/request_handler.py +++ b/src/typesense/request_handler.py @@ -48,8 +48,23 @@ ) TEntityDict = typing.TypeVar("TEntityDict") -TParams = typing.TypeVar("TParams", bound=typing.Dict[str, typing.Any]) -TBody = typing.TypeVar("TBody", bound=typing.Union[str, bytes]) +TParams = typing.TypeVar("TParams", bound=typing.Mapping[str, object]) +TBody = typing.TypeVar( + "TBody", bound=typing.Union[str, bytes, typing.Mapping[str, typing.Any]] +) + +# The query parameter values httpx accepts, once booleans are normalized to strings. +_QueryParams = typing.Mapping[ + str, + typing.Union[ + str, + int, + float, + bool, + None, + typing.Sequence[typing.Union[str, int, float, bool, None]], + ], +] _ERROR_CODE_MAP: typing.Mapping[str, typing.Type[TypesenseClientError]] = ( MappingProxyType( @@ -99,7 +114,7 @@ class SessionFunctionKwargs(typing.Generic[TParams, TBody], typing.TypedDict): data: typing.NotRequired[ typing.Union[TBody, str, typing.Dict[str, typing.Any], None] ] - content: typing.NotRequired[typing.Union[TBody, str, None]] + content: typing.NotRequired[typing.Union[str, bytes, None]] headers: typing.NotRequired[typing.Dict[str, str]] timeout: typing.NotRequired[float] @@ -128,6 +143,30 @@ def __init__(self, config: Configuration): """ self.config = config + @typing.overload + def make_request( + self, + *, + method: str, + url: str, + entity_type: typing.Type[TEntityDict], + as_json: typing.Union[typing.Literal[True], typing.Literal[False]] = True, + client: httpx.AsyncClient, + **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], + ) -> typing.Coroutine[typing.Any, typing.Any, typing.Union[TEntityDict, str]]: ... + + @typing.overload + def make_request( + self, + *, + method: str, + url: str, + entity_type: typing.Type[TEntityDict], + as_json: typing.Union[typing.Literal[True], typing.Literal[False]] = True, + client: httpx.Client, + **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], + ) -> typing.Union[TEntityDict, str]: ... + def make_request( self, *, @@ -183,9 +222,9 @@ def make_request( request_kwargs["params"] = params if body := kwargs.get("data"): - if not isinstance(body, (str, bytes)): - body = json.dumps(body) - request_kwargs["content"] = typing.cast(TBody, body) + request_kwargs["content"] = ( + body if isinstance(body, (str, bytes)) else json.dumps(body) + ) if isinstance(client, httpx.AsyncClient): return self._make_async_request( @@ -207,13 +246,13 @@ def _make_sync_request( ) -> typing.Union[TEntityDict, str]: """Make a synchronous HTTP request using httpx.Client.""" params: typing.Union[TParams, None] = request_kwargs.get("params") - content: typing.Union[TBody, str, None] = request_kwargs.get("content") + content: typing.Union[str, bytes, None] = request_kwargs.get("content") headers: typing.Dict[str, str] = request_kwargs.get("headers", {}) response = client.request( method, url, - params=params, + params=typing.cast(typing.Optional[_QueryParams], params), content=content, headers=headers, ) @@ -242,13 +281,13 @@ async def _make_async_request( ) -> typing.Union[TEntityDict, str]: """Make an asynchronous HTTP request using httpx.AsyncClient.""" params: typing.Union[TParams, None] = request_kwargs.get("params") - content: typing.Union[TBody, str, None] = request_kwargs.get("content") + content: typing.Union[str, bytes, None] = request_kwargs.get("content") headers: typing.Dict[str, str] = request_kwargs.get("headers", {}) response = await client.request( method, url, - params=params, + params=typing.cast(typing.Optional[_QueryParams], params), content=content, headers=headers, ) @@ -267,17 +306,19 @@ async def _make_async_request( return response.text @staticmethod - def normalize_params(params: typing.Dict[str, typing.Any]) -> None: + def normalize_params(params: typing.Mapping[str, object]) -> None: """ - Normalize boolean parameters in the request. + Normalize boolean parameters in the request, in place. Args: - params (Dict[str, Any]): The parameters to normalize. + params (Mapping[str, object]): The parameters to normalize. They are + typed as read-only so TypedDict parameters are accepted, but must + be a ``dict`` at runtime. Raises: ValueError: If params is not a dictionary. """ - if not isinstance(params, typing.Dict): + if not isinstance(params, dict): raise ValueError("Params must be a dictionary.") for key, parameter_value in params.items(): if isinstance(parameter_value, bool): diff --git a/src/typesense/sync/analytics_rule_v1.py b/src/typesense/sync/analytics_rule_v1.py index 38e8f41..9eec662 100644 --- a/src/typesense/sync/analytics_rule_v1.py +++ b/src/typesense/sync/analytics_rule_v1.py @@ -74,11 +74,9 @@ def retrieve( Union[RuleSchemaForQueries, RuleSchemaForCounters]: The schema containing the rule details. """ - response: typing.Union[ - RuleSchemaForQueries, RuleSchemaForCounters - ] = self.api_call.get( + response = self.api_call.get( self._endpoint_path, - entity_type=dict, + entity_type=typing.Dict[str, typing.Any], as_json=True, ) return typing.cast( @@ -101,7 +99,7 @@ def delete(self) -> RuleDeleteSchema: return response @property - @warn_deprecation( # type: ignore[untyped-decorator] + @warn_deprecation( "SyncAnalyticsRuleV1 is deprecated on v30+. Use client.analytics.rules[rule_id] instead.", flag_name="analytics_rules_v1_deprecation", ) diff --git a/src/typesense/sync/analytics_rules_v1.py b/src/typesense/sync/analytics_rules_v1.py index e63f802..30edf75 100644 --- a/src/typesense/sync/analytics_rules_v1.py +++ b/src/typesense/sync/analytics_rules_v1.py @@ -89,7 +89,7 @@ def __getitem__(self, rule_id: str) -> AnalyticsRuleV1: self.rules[rule_id] = AnalyticsRuleV1(self.api_call, rule_id) return self.rules[rule_id] - @warn_deprecation( # type: ignore[untyped-decorator] + @warn_deprecation( "SyncAnalyticsRulesV1 is deprecated on v30+. Use client.analytics instead.", flag_name="analytics_rules_v1_deprecation", ) @@ -115,21 +115,19 @@ def create( The created rule. Returns RuleSchemaForCounters for counter rules and RuleSchemaForQueries for query rules. """ - response: typing.Union[ - RuleSchemaForCounters, RuleSchemaForQueries - ] = self.api_call.post( + response = self.api_call.post( AnalyticsRulesV1.resource_path, body=rule, params=rule_parameters, as_json=True, - entity_type=dict, + entity_type=typing.Dict[str, typing.Any], ) return typing.cast( typing.Union[RuleSchemaForCounters, RuleSchemaForQueries], response, ) - @warn_deprecation( # type: ignore[untyped-decorator] + @warn_deprecation( "SyncAnalyticsRulesV1 is deprecated on v30+. Use client.analytics instead.", flag_name="analytics_rules_v1_deprecation", ) @@ -148,19 +146,17 @@ def upsert( Returns: Union[RuleSchemaForCounters, RuleCreateSchemaForQueries]: The upserted rule. """ - response: typing.Union[ - RuleSchemaForCounters, RuleCreateSchemaForQueries - ] = self.api_call.put( + response = self.api_call.put( "/".join([self.resource_path, rule_id]), body=rule, - entity_type=dict, + entity_type=typing.Dict[str, typing.Any], ) return typing.cast( typing.Union[RuleSchemaForCounters, RuleCreateSchemaForQueries], response, ) - @warn_deprecation( # type: ignore[untyped-decorator] + @warn_deprecation( "SyncAnalyticsRulesV1 is deprecated on v30+. Use client.analytics instead.", flag_name="analytics_rules_v1_deprecation", ) diff --git a/src/typesense/sync/api_call.py b/src/typesense/sync/api_call.py index 402a0dc..09c6107 100644 --- a/src/typesense/sync/api_call.py +++ b/src/typesense/sync/api_call.py @@ -37,6 +37,7 @@ import httpx +from typesense.concurrency_limit import ConcurrencyLimit from typesense.configuration import Configuration, Node from typesense.exceptions import ( HTTPStatus0Error, @@ -59,7 +60,7 @@ import typing_extensions as typing TEntityDict = typing.TypeVar("TEntityDict") -TParams = typing.TypeVar("TParams", bound=typing.Dict[str, typing.Any]) +TParams = typing.TypeVar("TParams", bound=typing.Mapping[str, object]) TBody = typing.TypeVar( "TBody", bound=typing.Union[str, bytes, typing.Mapping[str, typing.Any]] ) @@ -94,7 +95,7 @@ class SessionFunctionKwargs(typing.Generic[TParams, TBody], typing.TypedDict): params: typing.NotRequired[typing.Union[TParams, None]] data: typing.NotRequired[typing.Union[TBody, None]] - content: typing.NotRequired[typing.Union[TBody, str, None]] + content: typing.NotRequired[typing.Union[str, bytes, None]] headers: typing.NotRequired[typing.Dict[str, str]] timeout: typing.NotRequired[float] @@ -135,6 +136,20 @@ class SessionFunctionKwargs(typing.Generic[TParams, TBody], typing.TypedDict): ServiceUnavailable, ) +_CLIENT_ERRORS: typing.Final[ + typing.Tuple[ + typing.Type[httpx.PoolTimeout], + typing.Type[httpx.LocalProtocolError], + typing.Type[httpx.DecodingError], + typing.Type[httpx.TooManyRedirects], + ] +] = ( + httpx.PoolTimeout, + httpx.LocalProtocolError, + httpx.DecodingError, + httpx.TooManyRedirects, +) + class ApiCall: """ @@ -160,7 +175,17 @@ def __init__(self, config: Configuration): self.node_manager = NodeManager(config) self.request_handler = RequestHandler(config) self._client = httpx.Client( - timeout=config.connection_timeout_seconds, + timeout=httpx.Timeout( + config.connection_timeout_seconds, + pool=config.pool_timeout_seconds, + ), + limits=httpx.Limits( + max_connections=config.max_connections, + max_keepalive_connections=config.max_keepalive_connections, + ), + ) + self._concurrency_limit = ConcurrencyLimit( + config.max_concurrent_requests, ) def __enter__(self) -> "ApiCall": @@ -473,11 +498,14 @@ def _execute_request( try: return self._make_request_and_process_response( method, + node, url, entity_type, as_json, **request_kwargs, ) + except _CLIENT_ERRORS: + raise except _SERVER_ERRORS as server_error: self.node_manager.set_node_health(node, is_healthy=False) if num_retries < self.config.num_retries: @@ -495,24 +523,23 @@ def _execute_request( def _make_request_and_process_response( self, method: str, + node: Node, url: str, entity_type: typing.Type[TEntityDict], as_json: bool, **kwargs: typing.Unpack[SessionFunctionKwargs[TParams, TBody]], ) -> typing.Union[TEntityDict, str]: - """Make the async API request and process the response.""" - request_response = self.request_handler.make_request( - method=method, - url=url, - as_json=as_json, - entity_type=entity_type, - client=self._client, - **kwargs, - ) - self.node_manager.set_node_health( - self.node_manager.get_node(), - is_healthy=True, - ) + """Make the async API request to `node` and process the response.""" + with self._concurrency_limit: + request_response = self.request_handler.make_request( + method=method, + url=url, + as_json=as_json, + entity_type=entity_type, + client=self._client, + **kwargs, + ) + self.node_manager.set_node_health(node, is_healthy=True) return ( typing.cast(TEntityDict, request_response) if as_json diff --git a/src/typesense/sync/client.py b/src/typesense/sync/client.py index ef1afb0..b11e542 100644 --- a/src/typesense/sync/client.py +++ b/src/typesense/sync/client.py @@ -164,5 +164,5 @@ def typed_collection( """ if name is None: name = model.__name__.lower() - collection: Collection[TDoc] = self.collections[name] - return collection + # ``collections`` is typed for the default DocumentSchema; narrow it to the model. + return typing.cast(Collection[TDoc], self.collections[name]) diff --git a/src/typesense/sync/documents.py b/src/typesense/sync/documents.py index b22ef69..0c7d7f7 100644 --- a/src/typesense/sync/documents.py +++ b/src/typesense/sync/documents.py @@ -61,6 +61,17 @@ None, ] +# One line of an import response. ``ImportResponse`` is a union of lists, one per +# return mode, so the helpers below build a list of these and ``import_`` casts it +# to the list type its overloads promise. +_ImportResponseItem = typing.Union[ + ImportResponseWithDoc[TDoc], + ImportResponseWithId, + ImportResponseWithDocAndId[TDoc], + ImportResponseSuccess, + ImportResponseFail[TDoc], +] + class Documents(typing.Generic[TDoc]): """ @@ -125,12 +136,14 @@ def create( Returns: TDoc: The created document. """ - dirty_values_parameters = dirty_values_parameters or {} - dirty_values_parameters["action"] = "create" + write_parameters: typing.Dict[str, object] = { + **(dirty_values_parameters or {}), + "action": "create", + } response = self.api_call.post( self._endpoint_path(), body=document, - params=dirty_values_parameters, + params=write_parameters, as_json=True, entity_type=typing.Dict[str, str], ) @@ -154,7 +167,14 @@ def create_many( The list of import responses. """ logger.warn("`create_many` is deprecated: please use `import_`.") - return self.import_(documents, dirty_values_parameters) + # Dirty values parameters are a subset of the write parameters. + return self.import_( + documents, + typing.cast( + typing.Optional[DocumentWriteParameters], + dirty_values_parameters, + ), + ) def upsert( self, @@ -172,12 +192,14 @@ def upsert( Returns: TDoc: The upserted document. """ - dirty_values_parameters = dirty_values_parameters or {} - dirty_values_parameters["action"] = "upsert" + write_parameters: typing.Dict[str, object] = { + **(dirty_values_parameters or {}), + "action": "upsert", + } response = self.api_call.post( self._endpoint_path(), body=document, - params=dirty_values_parameters, + params=write_parameters, as_json=True, entity_type=typing.Dict[str, str], ) @@ -199,12 +221,14 @@ def update( Returns: UpdateByFilterResponse: The response containing information about the update. """ - dirty_values_parameters = dirty_values_parameters or {} - dirty_values_parameters["action"] = "update" + update_parameters: typing.Dict[str, object] = { + **(dirty_values_parameters or {}), + "action": "update", + } response: UpdateByFilterResponse = self.api_call.patch( self._endpoint_path(), body=document, - params=dirty_values_parameters, + params=update_parameters, entity_type=UpdateByFilterResponse, ) return response @@ -301,9 +325,14 @@ def import_( return self._import_raw(documents, import_parameters) if batch_size: - return self._batch_import(documents, import_parameters, batch_size) - - return self._bulk_import(documents, import_parameters) + response_objs = self._batch_import( + documents, + import_parameters, + batch_size, + ) + else: + response_objs = self._bulk_import(documents, import_parameters) + return typing.cast(ImportResponse[TDoc], response_objs) def export( self, @@ -410,9 +439,9 @@ def _batch_import( documents: typing.List[TDoc], import_parameters: _ImportParameters, batch_size: int, - ) -> ImportResponse[TDoc]: + ) -> typing.List[_ImportResponseItem[TDoc]]: """Import documents in batches.""" - response_objs: ImportResponse[TDoc] = [] + response_objs: typing.List[_ImportResponseItem[TDoc]] = [] for batch_index in range(0, len(documents), batch_size): batch = documents[batch_index : batch_index + batch_size] api_response = self._bulk_import(batch, import_parameters) @@ -423,7 +452,7 @@ def _bulk_import( self, documents: typing.List[TDoc], import_parameters: _ImportParameters, - ) -> ImportResponse[TDoc]: + ) -> typing.List[_ImportResponseItem[TDoc]]: """Import a list of documents in bulk.""" document_strs = [json.dumps(doc) for doc in documents] if not document_strs: @@ -439,9 +468,12 @@ def _bulk_import( ) return self._parse_import_response(res) - def _parse_import_response(self, response: str) -> ImportResponse[TDoc]: + def _parse_import_response( + self, + response: str, + ) -> typing.List[_ImportResponseItem[TDoc]]: """Parse the import response string into a list of response objects.""" - response_objs: typing.List[ImportResponse] = [] + response_objs: typing.List[_ImportResponseItem[TDoc]] = [] for res_obj_str in response.split("\n"): try: res_obj_json = json.loads(res_obj_str) diff --git a/src/typesense/sync/keys.py b/src/typesense/sync/keys.py index b70ec5e..9bd494e 100644 --- a/src/typesense/sync/keys.py +++ b/src/typesense/sync/keys.py @@ -29,7 +29,6 @@ ApiKeyCreateResponseSchema, ApiKeyCreateSchema, ApiKeyRetrieveSchema, - ApiKeySchema, ) if sys.version_info >= (3, 11): @@ -103,11 +102,11 @@ def create(self, schema: ApiKeyCreateSchema) -> ApiKeyCreateResponseSchema: ... } ... ) """ - response: ApiKeySchema = self.api_call.post( + response: ApiKeyCreateResponseSchema = self.api_call.post( Keys.resource_path, as_json=True, body=schema, - entity_type=ApiKeySchema, + entity_type=ApiKeyCreateResponseSchema, ) return response diff --git a/src/typesense/sync/operations.py b/src/typesense/sync/operations.py index e560b76..450c345 100644 --- a/src/typesense/sync/operations.py +++ b/src/typesense/sync/operations.py @@ -60,8 +60,10 @@ def __init__(self, api_call: ApiCall): """ self.api_call = api_call + # The generic ``str`` overload below also matches "schema_changes"; overloads are + # tried in order, so this one wins. @typing.overload - def perform( + def perform( # type: ignore[overload-overlap] self, operation_name: typing.Literal["schema_changes"], query_params: None = None, @@ -132,36 +134,36 @@ def perform( @typing.overload def perform( self, - operation_name: str, - query_params: typing.Union[typing.Dict[str, str], None] = None, + operation_name: typing.Literal["snapshot"], + query_params: SnapshotParameters, ) -> OperationResponse: """ - Perform a generic operation. + Perform a snapshot operation. Args: - operation_name (str): The name of the operation. - query_params (Union[Dict[str, str], None], optional): - Query parameters for the operation. + operation_name (Literal["snapshot"]): The name of the operation. + query_params (SnapshotParameters): Query parameters for the snapshot operation. Returns: - OperationResponse: The response from the operation. + OperationResponse: The response from the snapshot operation. """ @typing.overload def perform( self, - operation_name: typing.Literal["snapshot"], - query_params: SnapshotParameters, + operation_name: str, + query_params: typing.Union[typing.Dict[str, str], None] = None, ) -> OperationResponse: """ - Perform a snapshot operation. + Perform a generic operation. Args: - operation_name (Literal["snapshot"]): The name of the operation. - query_params (SnapshotParameters): Query parameters for the snapshot operation. + operation_name (str): The name of the operation. + query_params (Union[Dict[str, str], None], optional): + Query parameters for the operation. Returns: - OperationResponse: The response from the snapshot operation. + OperationResponse: The response from the operation. """ def perform( @@ -181,7 +183,7 @@ def perform( typing.Dict[str, str], None, ] = None, - ) -> OperationResponse: + ) -> typing.Union[OperationResponse, typing.List[SchemaChangesResponse]]: """ Perform an operation on the Typesense API. @@ -202,13 +204,16 @@ def perform( >>> response = await operations.perform("vote") >>> health = await operations.is_healthy() """ - response: OperationResponse = self.api_call.post( + response = self.api_call.post( self._endpoint_path(operation_name), params=query_params, as_json=True, - entity_type=OperationResponse, + entity_type=object, + ) + return typing.cast( + typing.Union[OperationResponse, typing.List[SchemaChangesResponse]], + response, ) - return response def is_healthy(self) -> bool: """ @@ -222,16 +227,14 @@ def is_healthy(self) -> bool: >>> healthy = await operations.is_healthy() >>> print(healthy) """ - call_resp: HealthCheckResponse = self.api_call.get( + call_resp: object = self.api_call.get( Operations.health_path, as_json=True, entity_type=HealthCheckResponse, ) - if isinstance(call_resp, typing.Dict): - is_ok: bool = call_resp.get("ok", False) - else: - is_ok = False - return is_ok + if isinstance(call_resp, dict): + return bool(call_resp.get("ok", False)) + return False def toggle_slow_request_log( self, diff --git a/src/typesense/sync/override.py b/src/typesense/sync/override.py index 8a24e9e..78aff10 100644 --- a/src/typesense/sync/override.py +++ b/src/typesense/sync/override.py @@ -87,7 +87,7 @@ def delete(self) -> OverrideDeleteSchema: return response @property - @warn_deprecation( # type: ignore[untyped-decorator] + @warn_deprecation( "The override API (collections/{collection}/overrides/{override_id}) is deprecated is removed on v30+. " "Use curation sets (curation_sets) instead.", flag_name="overrides_deprecation", diff --git a/src/typesense/sync/overrides.py b/src/typesense/sync/overrides.py index 7682ff5..99c4667 100644 --- a/src/typesense/sync/overrides.py +++ b/src/typesense/sync/overrides.py @@ -129,7 +129,7 @@ def retrieve(self) -> OverrideRetrieveSchema: ) return response - @warn_deprecation( # type: ignore[untyped-decorator] + @warn_deprecation( "SyncOverrides is deprecated on v30+. Use client.curation_sets instead.", flag_name="overrides_deprecation", ) diff --git a/src/typesense/sync/synonym.py b/src/typesense/sync/synonym.py index d091fdd..27ff6cf 100644 --- a/src/typesense/sync/synonym.py +++ b/src/typesense/sync/synonym.py @@ -79,7 +79,7 @@ def delete(self) -> SynonymDeleteSchema: ) @property - @warn_deprecation( # type: ignore[untyped-decorator] + @warn_deprecation( "The synonym API (collections/{collection}/synonyms/{synonym_id}) is deprecated is removed on v30+. " "Use synonym sets (synonym_sets) instead.", flag_name="synonyms_deprecation", diff --git a/src/typesense/sync/synonyms.py b/src/typesense/sync/synonyms.py index d6e055b..f3c4e45 100644 --- a/src/typesense/sync/synonyms.py +++ b/src/typesense/sync/synonyms.py @@ -124,7 +124,7 @@ def retrieve(self) -> SynonymsRetrieveSchema: ) return response - @warn_deprecation( # type: ignore[untyped-decorator] + @warn_deprecation( "The synonyms API (collections/{collection}/synonyms) is deprecated is removed on v30+. " "Use synonym sets (synonym_sets) instead.", flag_name="synonyms_deprecation", diff --git a/tests/api_call_test.py b/tests/api_call_test.py index b7c4888..9c77d03 100644 --- a/tests/api_call_test.py +++ b/tests/api_call_test.py @@ -1,8 +1,11 @@ """Unit Tests for the ApiCall class.""" +import asyncio import logging import sys +import threading import time +from concurrent.futures import ThreadPoolExecutor from pytest_mock import MockFixture @@ -461,6 +464,25 @@ def test_selects_next_available_node_on_timeout( assert len(respx.calls) == 3 +def test_client_errors_do_not_mark_nodes_unhealthy( + fake_api_call: ApiCall, + mocker: MockerFixture, +) -> None: + """Pool exhaustion is local to the client and must not trigger failover.""" + node = fake_api_call.node_manager.get_node() + make_request = mocker.patch.object( + fake_api_call.request_handler, + "make_request", + side_effect=httpx.PoolTimeout("No connection available"), + ) + + with pytest.raises(httpx.PoolTimeout): + fake_api_call.get("/test", as_json=True, entity_type=typing.Dict[str, str]) + + assert node.healthy is True + make_request.assert_called_once() + + def test_get_node_no_healthy_nodes( fake_api_call: ApiCall, mocker: MockFixture, @@ -665,3 +687,259 @@ async def test_async_sleeps_retry_interval_between_retries( assert sleep_call == mocker.call( fake_async_api_call.config.retry_interval_seconds, ) + + +@pytest.mark.parametrize( + "client_side_error", + [ + httpx.PoolTimeout("Pool timeout"), + httpx.LocalProtocolError("Local protocol error"), + httpx.DecodingError("Decoding error"), + httpx.TooManyRedirects("Too many redirects"), + ], +) +def test_client_side_error_does_not_mark_node_unhealthy( + fake_api_call: ApiCall, + client_side_error: httpx.HTTPError, +) -> None: + """Test that client-side httpx errors propagate without failing over.""" + with respx.mock: + respx.get("http://nearest:8108/").mock(side_effect=client_side_error) + node0_route = respx.get("http://node0:8108/").mock( + return_value=httpx.Response(200, json={"key": "value"}), + ) + + with pytest.raises(type(client_side_error)): + fake_api_call.get("/", entity_type=typing.Dict[str, str]) + + assert len(respx.calls) == 1 + assert not node0_route.called + + assert fake_api_call.config.nearest_node.healthy is True + + +@pytest.mark.parametrize( + "client_side_error", + [ + httpx.PoolTimeout("Pool timeout"), + httpx.LocalProtocolError("Local protocol error"), + httpx.DecodingError("Decoding error"), + httpx.TooManyRedirects("Too many redirects"), + ], +) +async def test_async_client_side_error_does_not_mark_node_unhealthy( + fake_async_api_call: AsyncApiCall, + client_side_error: httpx.HTTPError, +) -> None: + """Test that client-side httpx errors propagate without failing over (async).""" + with respx.mock: + respx.get("http://nearest:8108/").mock(side_effect=client_side_error) + node0_route = respx.get("http://node0:8108/").mock( + return_value=httpx.Response(200, json={"key": "value"}), + ) + + with pytest.raises(type(client_side_error)): + await fake_async_api_call.get("/", entity_type=typing.Dict[str, str]) + + assert len(respx.calls) == 1 + assert not node0_route.called + + assert fake_async_api_call.config.nearest_node.healthy is True + + +def test_round_robin_visits_each_node_in_turn(fake_api_call: ApiCall) -> None: + """Test that successful requests advance the round-robin by one node each.""" + fake_api_call.config.nearest_node = None + + with respx.mock: + for host in ("node0", "node1", "node2"): + respx.get(f"http://{host}:8108/").mock( + return_value=httpx.Response(200, json={"key": "value"}), + ) + + for _ in range(6): + fake_api_call.get("/", entity_type=typing.Dict[str, str]) + + assert [str(call.request.url) for call in respx.calls] == [ + "http://node0:8108/", + "http://node1:8108/", + "http://node2:8108/", + "http://node0:8108/", + "http://node1:8108/", + "http://node2:8108/", + ] + + +async def test_async_round_robin_visits_each_node_in_turn( + fake_async_api_call: AsyncApiCall, +) -> None: + """Test that successful requests advance the round-robin by one node each (async).""" + fake_async_api_call.config.nearest_node = None + + with respx.mock: + for host in ("node0", "node1", "node2"): + respx.get(f"http://{host}:8108/").mock( + return_value=httpx.Response(200, json={"key": "value"}), + ) + + for _ in range(6): + await fake_async_api_call.get("/", entity_type=typing.Dict[str, str]) + + assert [str(call.request.url) for call in respx.calls] == [ + "http://node0:8108/", + "http://node1:8108/", + "http://node2:8108/", + "http://node0:8108/", + "http://node1:8108/", + "http://node2:8108/", + ] + + +def test_success_marks_only_the_answering_node_healthy( + fake_api_call: ApiCall, +) -> None: + """Test that a success refreshes the node that answered and no other.""" + fake_api_call.config.nearest_node = None + answering_node, unhealthy_node, _ = fake_api_call.node_manager.nodes + answering_node.last_access_ts = 0 + unhealthy_node.healthy = False + unhealthy_node.last_access_ts = int(time.time()) + + with respx.mock: + respx.get("http://node0:8108/").mock( + return_value=httpx.Response(200, json={"key": "value"}), + ) + + fake_api_call.get("/", entity_type=typing.Dict[str, str]) + + assert answering_node.healthy is True + assert answering_node.last_access_ts > 0 + assert unhealthy_node.healthy is False + + +def test_client_uses_connection_pool_settings( + fake_config: Configuration, + mocker: MockerFixture, +) -> None: + """Test that the httpx client is built from the connection pool settings.""" + client_mock = mocker.patch("typesense.sync.api_call.httpx.Client") + fake_config.connection_timeout_seconds = 3.0 + fake_config.pool_timeout_seconds = 1.5 + fake_config.max_connections = 200 + fake_config.max_keepalive_connections = 50 + + ApiCall(fake_config) + + client_mock.assert_called_once_with( + timeout=httpx.Timeout(3.0, pool=1.5), + limits=httpx.Limits(max_connections=200, max_keepalive_connections=50), + ) + + +def test_async_client_uses_connection_pool_settings( + fake_config: Configuration, + mocker: MockerFixture, +) -> None: + """Test that the httpx async client is built from the connection pool settings.""" + client_mock = mocker.patch("typesense.async_.api_call.httpx.AsyncClient") + fake_config.connection_timeout_seconds = 3.0 + fake_config.pool_timeout_seconds = 1.5 + fake_config.max_connections = 200 + fake_config.max_keepalive_connections = 50 + + AsyncApiCall(fake_config) + + client_mock.assert_called_once_with( + timeout=httpx.Timeout(3.0, pool=1.5), + limits=httpx.Limits(max_connections=200, max_keepalive_connections=50), + ) + + +def _count_requests_in_flight( + concurrent_requests: int, + max_concurrent_requests: typing.Optional[int], + fake_config: Configuration, +) -> int: + """Send requests from several threads and return the peak number in flight.""" + fake_config.max_concurrent_requests = max_concurrent_requests + api_call = ApiCall(fake_config) + lock = threading.Lock() + in_flight = 0 + peak = 0 + + def slow_response(request: httpx.Request) -> httpx.Response: + nonlocal in_flight, peak + with lock: + in_flight += 1 + peak = max(peak, in_flight) + time.sleep(0.05) + with lock: + in_flight -= 1 + return httpx.Response(200, json={"key": "value"}) + + with respx.mock: + respx.get("http://nearest:8108/").mock(side_effect=slow_response) + with ThreadPoolExecutor(max_workers=concurrent_requests) as executor: + for _ in range(concurrent_requests): + executor.submit(api_call.get, "/", entity_type=typing.Dict[str, str]) + + return peak + + +def test_max_concurrent_requests_caps_requests_in_flight( + fake_config: Configuration, +) -> None: + """Test that no more than ``max_concurrent_requests`` requests are in flight.""" + assert _count_requests_in_flight(6, 2, fake_config) == 2 + + +def test_requests_in_flight_are_unlimited_by_default( + fake_config: Configuration, +) -> None: + """Test that requests are not capped when ``max_concurrent_requests`` is unset.""" + assert _count_requests_in_flight(6, None, fake_config) == 6 + + +async def _async_count_requests_in_flight( + concurrent_requests: int, + max_concurrent_requests: typing.Optional[int], + fake_config: Configuration, +) -> int: + """Send concurrent async requests and return the peak number in flight.""" + fake_config.max_concurrent_requests = max_concurrent_requests + api_call = AsyncApiCall(fake_config) + in_flight = 0 + peak = 0 + + async def slow_response(request: httpx.Request) -> httpx.Response: + nonlocal in_flight, peak + in_flight += 1 + peak = max(peak, in_flight) + await asyncio.sleep(0.01) + in_flight -= 1 + return httpx.Response(200, json={"key": "value"}) + + with respx.mock: + respx.get("http://nearest:8108/").mock(side_effect=slow_response) + await asyncio.gather( + *( + api_call.get("/", entity_type=typing.Dict[str, str]) + for _ in range(concurrent_requests) + ), + ) + + return peak + + +async def test_async_max_concurrent_requests_caps_requests_in_flight( + fake_config: Configuration, +) -> None: + """Test that no more than ``max_concurrent_requests`` requests are in flight (async).""" + assert await _async_count_requests_in_flight(6, 2, fake_config) == 2 + + +async def test_async_requests_in_flight_are_unlimited_by_default( + fake_config: Configuration, +) -> None: + """Test that async requests are not capped when ``max_concurrent_requests`` is unset.""" + assert await _async_count_requests_in_flight(6, None, fake_config) == 6 diff --git a/tests/configuration_test.py b/tests/configuration_test.py index 626c477..092c93b 100644 --- a/tests/configuration_test.py +++ b/tests/configuration_test.py @@ -207,3 +207,46 @@ def test_configuration_invalid_nearest_node_url() -> None: match="Node URL does not contain the port.", ): Configuration(config) + + +def test_configuration_connection_pool_defaults() -> None: + """Test the connection pool defaults, with the pool timeout following the connection timeout.""" + configuration = Configuration( + { + "nodes": [DEFAULT_NODE], + "api_key": "xyz", + "connection_timeout_seconds": 7.0, + }, + ) + + expected = { + "pool_timeout_seconds": 7.0, + "max_connections": 100, + "max_keepalive_connections": 20, + "max_concurrent_requests": None, + } + + assert_to_contain_object(configuration, expected) + + +def test_configuration_connection_pool_explicit() -> None: + """Test the connection pool settings with explicit values.""" + configuration = Configuration( + { + "nodes": [DEFAULT_NODE], + "api_key": "xyz", + "pool_timeout_seconds": 1.5, + "max_connections": 200, + "max_keepalive_connections": 50, + "max_concurrent_requests": 150, + }, + ) + + expected = { + "pool_timeout_seconds": 1.5, + "max_connections": 200, + "max_keepalive_connections": 50, + "max_concurrent_requests": 150, + } + + assert_to_contain_object(configuration, expected) diff --git a/tests/configuration_validations_test.py b/tests/configuration_validations_test.py index d408e05..8cf8061 100644 --- a/tests/configuration_validations_test.py +++ b/tests/configuration_validations_test.py @@ -1,7 +1,13 @@ """Tests for the ConfigurationValidations class.""" +import sys import types +if sys.version_info >= (3, 11): + import typing +else: + import typing_extensions as typing + import pytest from typesense.configuration import ConfigDict, ConfigurationValidations @@ -199,3 +205,34 @@ def test_validate_config_dict_with_wrong_nearest_node() -> None: "api_key": "xyz", }, ) + + +@pytest.mark.parametrize( + ("key", "config_value", "message"), + [ + ("pool_timeout_seconds", 0, "`pool_timeout_seconds` must be greater than 0."), + ("max_connections", 0, "`max_connections` must be greater than 0."), + ( + "max_concurrent_requests", + -1, + "`max_concurrent_requests` must be greater than 0.", + ), + ( + "max_keepalive_connections", + -1, + "`max_keepalive_connections` must not be negative.", + ), + ], +) +def test_validate_config_dict_with_invalid_connection_pool( + key: str, + config_value: float, + message: str, +) -> None: + """Test validate_config_dict with out-of-range connection pool settings.""" + config_dict = {"nodes": [DEFAULT_NODE], "api_key": "xyz", key: config_value} + + with pytest.raises(ConfigError, match=message): + ConfigurationValidations.validate_config_dict( + typing.cast(ConfigDict, config_dict), + ) diff --git a/utils/run-unasync.py b/utils/run-unasync.py index aa4dcbd..7d836d4 100644 --- a/utils/run-unasync.py +++ b/utils/run-unasync.py @@ -27,6 +27,8 @@ def collect_class_replacements(source_dir: Path) -> dict[str, str]: # client (unasync strips ``await``); map the module token so the import and call # are rewritten too. replacements["asyncio"] = "time" + # Defined in the shared ``typesense.concurrency_limit`` module, outside async_. + replacements["AsyncConcurrencyLimit"] = "ConcurrencyLimit" return replacements