Skip to content
Merged
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
1 change: 1 addition & 0 deletions setup.cfg
Original file line number Diff line number Diff line change
Expand Up @@ -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
8 changes: 3 additions & 5 deletions src/typesense/async_/analytics_rule_v1.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand All @@ -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",
)
Expand Down
18 changes: 7 additions & 11 deletions src/typesense/async_/analytics_rules_v1.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
)
Expand All @@ -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",
)
Expand All @@ -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",
)
Expand Down
59 changes: 43 additions & 16 deletions src/typesense/async_/api_call.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,7 @@

import httpx

from typesense.concurrency_limit import AsyncConcurrencyLimit
from typesense.configuration import Configuration, Node
from typesense.exceptions import (
HTTPStatus0Error,
Expand All @@ -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]]
)
Expand Down Expand Up @@ -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]

Expand Down Expand Up @@ -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:
"""
Expand All @@ -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":
Expand Down Expand Up @@ -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:
Expand All @@ -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
Expand Down
4 changes: 2 additions & 2 deletions src/typesense/async_/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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])
68 changes: 50 additions & 18 deletions src/typesense/async_/documents.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]):
"""
Expand Down Expand Up @@ -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],
)
Expand All @@ -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,
Expand All @@ -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],
)
Expand All @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand All @@ -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:
Expand All @@ -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)
Expand Down
5 changes: 2 additions & 3 deletions src/typesense/async_/keys.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,6 @@
ApiKeyCreateResponseSchema,
ApiKeyCreateSchema,
ApiKeyRetrieveSchema,
ApiKeySchema,
)

if sys.version_info >= (3, 11):
Expand Down Expand Up @@ -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

Expand Down
Loading
Loading