from collections.abc import Sequence from functools import wraps from json.decoder import JSONDecodeError from typing import Callable, Concatenate, Generic, Optional, ParamSpec, TypedDict, TypeGuard, TypeVar, Union from charts.api import app from marshmallow import Schema, fields from marshmallow.exceptions import ValidationError from charts.connectors import redis DataLoaderIdInput = TypeVar('DataLoaderIdInput') DataLoaderParamSpec = ParamSpec('DataLoaderParamSpec') DataLoaderResponseType = TypeVar('DataLoaderResponseType') class DataLoaderResponse(Generic[DataLoaderResponseType], TypedDict): data: Optional[DataLoaderResponseType] def dataloader( func: Callable[Concatenate[list[DataLoaderIdInput], DataLoaderParamSpec], list[DataLoaderResponse[DataLoaderResponseType]]] ) -> Callable[Concatenate[list[DataLoaderIdInput], DataLoaderParamSpec], list[DataLoaderResponse[DataLoaderResponseType]]]: """Decorate a function as a dataloader. This means that the function takes an array of ids (of any type) as its first positional parameter. The function may take any number of other positional or keyword arguments, but these must apply to all items. The number of returned values will be verified against the number of inputs. Args: func: The function to be decorated Return: The decorated function """ @wraps(func) def wrapped_func( ids: list[DataLoaderIdInput], /, *args: DataLoaderParamSpec.args, **kwargs: DataLoaderParamSpec.kwargs, ) -> list[DataLoaderResponse[DataLoaderResponseType]]: if not isinstance(ids, Sequence): raise Exception('Not a dataloader.') result = func(ids, *args, **kwargs) if len(result) != len(ids): raise Exception('Violated dataloader contract: Incorrect row count returned.') return result return wrapped_func def _default_should_write_to_redis( _key: DataLoaderIdInput, value: DataLoaderResponse[DataLoaderResponseType], /, *args: DataLoaderParamSpec.args, **kwargs: DataLoaderParamSpec.kwargs, ) -> bool: return value.get('data') is not None def _is_complete_result_list(arr: list[DataLoaderResponse[DataLoaderResponseType] | None]) -> TypeGuard[list[DataLoaderResponse[DataLoaderResponseType]]]: return all(x is not None for x in arr) class MarshmallowRedisDataLoader(Generic[DataLoaderIdInput, DataLoaderParamSpec, DataLoaderResponseType]): expiry: Callable[Concatenate[DataLoaderIdInput, DataLoaderResponse[DataLoaderResponseType], DataLoaderParamSpec], int] def __init__( self, *, dataloader: Callable[Concatenate[list[DataLoaderIdInput], DataLoaderParamSpec], list[DataLoaderResponse[DataLoaderResponseType]]], expiry: Union[int, Callable[Concatenate[DataLoaderIdInput, DataLoaderResponse[DataLoaderResponseType], DataLoaderParamSpec], int]], key_serializer: Callable[Concatenate[DataLoaderIdInput, DataLoaderParamSpec], str], schema: type[Schema], should_write_to_redis: Callable[Concatenate[DataLoaderIdInput, DataLoaderResponse[DataLoaderResponseType], DataLoaderParamSpec], bool] = _default_should_write_to_redis, ): self.dataloader = dataloader self.expiry = expiry if callable(expiry) else lambda _id, _key, /, *args, **kwargs: expiry self.key_serializer = key_serializer class DataLoaderItemSchema(Schema): data = fields.Nested(schema, required=True, allow_none=True) self.item_schema = DataLoaderItemSchema() self.should_write_to_redis = should_write_to_redis def _redis_read( self, ids: list[DataLoaderIdInput], /, *args: DataLoaderParamSpec.args, **kwargs: DataLoaderParamSpec.kwargs, ) -> list[Optional[DataLoaderResponse[DataLoaderResponseType]]]: redis_keys = [ self.key_serializer(id, *args, **kwargs) for id in ids ] try: # Redis sync types are incorrect. See https://github.com/redis/redis-py/issues/2897 redis_result: list[str | None] = redis.client.mget(redis_keys) # type: ignore[assignment] except redis.exceptions.ConnectionError as e: app.logger.warning(f'Redis connection error: {e}') return [None for _ in redis_keys] redis_values = [] for i, result in enumerate(redis_result): value = None if result: try: value = self.item_schema.loads(result) except JSONDecodeError as e: app.logger.warning(f'Loaded invalid JSON value from redis for key {redis_keys[i]}: {e}') except ValidationError as e: app.logger.warning(f'Loaded invalid schema value from redis for key {redis_keys[i]}: {e}') redis_values.append(value) return redis_values def load_many( self, ids: list[DataLoaderIdInput], /, *args: DataLoaderParamSpec.args, **kwargs: DataLoaderParamSpec.kwargs, ) -> list[DataLoaderResponse[DataLoaderResponseType]]: redis_values = self._redis_read(ids, *args, **kwargs) if _is_complete_result_list(redis_values): return redis_values non_found_ids = [id for i, id in enumerate(ids) if redis_values[i] is None] non_found_id_idxs = [i for i in range(len(redis_values)) if redis_values[i] is None] sub_results = self.load_many_no_cache_read(non_found_ids, *args, **kwargs) result = [ redis_value if redis_value else None for redis_value in redis_values ] for idx, sub_result in zip(non_found_id_idxs, sub_results): result[idx] = sub_result assert _is_complete_result_list(result) return result def load_many_no_cache_read( self, ids: list[DataLoaderIdInput], /, *args: DataLoaderParamSpec.args, **kwargs: DataLoaderParamSpec.kwargs, ) -> list[DataLoaderResponse[DataLoaderResponseType]]: sub_results = self.dataloader(ids, *args, **kwargs) try: for id, sub_result in zip(ids, sub_results): if self.should_write_to_redis(id, sub_result, *args, **kwargs): value_expiry = self.expiry(id, sub_result, *args, **kwargs) redis.client.setex( self.key_serializer(id, *args, **kwargs), value_expiry, self.item_schema.dumps(sub_result), ) except redis.exceptions.ConnectionError as e: app.logger.warning(f'Redis connection error: {e}') return sub_results