import asyncio from datetime import datetime from inspect import getfullargspec from typing import Any, Callable, Dict, Iterable, List, Tuple from aiojobs.aiohttp import spawn import config import context from server.core.cache.keys import get_full_kwargs, get_key_parts, make_custom_key from server.core.cache.results import BaseResult from server.core.cache.storage import get_items, save_item, save_items from server.core.constants import DataSource from server.core.utils import count_time, set_debug_data_source, set_debug_request_count def cache_data(result_handler: BaseResult, arg_index: int = 1, ignore_params: Tuple[str] = None) -> Callable: """Cache multiple results by custom keys. Gets values from cache if exist and loads rest if needed. Args: result_handler: Additional logic to process function results, get keys, get final results. arg_index: Args list index of iterable keys. ignore_params: Use all params except these as key parts. Returns: Wrapped function. """ def wrapper(f: Callable) -> Callable: args_spec = getfullargspec(f) async def wrapped(*args, **kwargs) -> List or Dict: keys_arg = args[arg_index] if not isinstance(keys_arg, list): keys_arg = [keys_arg] # get a sum of args and kwargs as kwargs full_kwargs = get_full_kwargs(args, kwargs, args_spec) # get key parts key_parts = get_key_parts(full_kwargs, ignore_params=ignore_params or ("self", args_spec.args[arg_index])) # get list of keys keys = result_handler.get_keys(keys_arg, key_parts) # try to get data from cache cached_data, min_created_at, max_created_at = await get_items( result_handler.collection_name, keys, context.DATA_TIMEOUT.get(), by_record_id=False ) # calc missing keys missing_keys = list(set(keys_arg) - set(cached_data.keys())) # return if everything is present in cache if not missing_keys: set_debug_data_source(DataSource.CACHE, min_created_at, max_created_at) return result_handler.process_result(cached_data) start_date = datetime.utcnow() # call func to get missing result = await count_time(context.REQUEST_TIME)(f)( *result_handler.replace_args(args, arg_index, missing_keys), **kwargs ) # get mapping for caching and combined result caching_map, combined_result = result_handler.get_results(result, cached_data, key_parts) created_at = datetime.utcnow() # set data to cache if context.BACKGROUND_SAVE.get(): for collection_name, records in caching_map.items(): await spawn(context.REQUEST.get(), save_items(collection_name, records, created_at)) else: tasks = [save_items(cn, dm, created_at) for cn, dm in caching_map.items()] await asyncio.gather(*tasks) if config.DEBUG_DATA_SOURCE: new_data_loaded = len(combined_result) > len(cached_data) set_debug_data_source( DataSource.COMBINED if cached_data and new_data_loaded else DataSource.API, min_created_at if cached_data else created_at, created_at if new_data_loaded else max_created_at, ) set_debug_request_count(1, start_date, created_at, force_replace=False) return combined_result return wrapped return wrapper def cache_request( collection_name: str, key_prefix: str = None, ignore_params: Tuple[str] = None, expires: int = None ) -> Callable: """Cache value by custom keys. Gets value from cache if exists and loads rest if needed. Args: collection_name: MongoDB collection name. key_prefix: Key prefix, should contain name of client, wrapped function. ignore_params: Use all params except these as key parts. expires: Data actuality in seconds. Returns: Wrapped function. """ def wrapper(f: Callable) -> Callable: args_spec = getfullargspec(f) async def wrapped(*args, **kwargs): full_kwargs = get_full_kwargs(args, kwargs, args_spec) key_parts = get_key_parts(full_kwargs, ignore_params=ignore_params or ("self",)) key_name = make_custom_key(key_prefix=key_prefix, key_args_params=key_parts) timeout = context.DATA_TIMEOUT.get() if expires: if timeout > 0: timeout = min(timeout, expires) elif timeout < 0: timeout = expires cached_value, created_at, _ = await get_items(collection_name, [key_name], timeout, by_record_id=True) if cached_value: set_debug_data_source(DataSource.CACHE, created_at, is_single=True) return list(cached_value.values())[0] start_date = datetime.utcnow() new_value = await count_time(context.REQUEST_TIME)(f)(*args, **kwargs) created_at = datetime.utcnow() if context.BACKGROUND_SAVE.get(): await spawn(context.REQUEST.get(), save_item(collection_name, key_name, new_value, created_at)) else: await save_item(collection_name, key_name, new_value, created_at) set_debug_data_source(DataSource.API, created_at, is_single=True) set_debug_request_count(1, start_date, created_at, force_replace=False) return new_value return wrapped return wrapper async def get_cached( id_list: Iterable[str], collection_name: str, key_kwargs: dict = None ) -> Tuple[Dict[str, Any], List[str]]: """Cache multiple results by custom keys. Gets values from cache if exist and loads rest if needed. Args: id_list: iterable of IDs to get cached values for collection_name: name of cache-collection for keys prefix key_kwargs: extra kwargs to be included in keys. Returns: Tuple of found cached collection and missing ids """ # get key parts key_parts = get_key_parts(key_kwargs or {}) # get list of keys keys = [make_custom_key(key_value=key_value, key_args_params=key_parts) for key_value in id_list] # try to get data from cache cached_data, min_created_at, max_created_at = await get_items( collection_name, keys, context.DATA_TIMEOUT.get(), by_record_id=False ) # calc missing keys missing_ids = list(set(id_list) - set(cached_data.keys())) return { "data": cached_data, "min_created_at": min_created_at, "max_created_at": max_created_at, } if cached_data else {}, missing_ids async def set_cached( result: Dict or List, result_handler: BaseResult, extra_result_data: dict = None, key_kwargs: dict = None ) -> Dict or List: """Cache new data and merge it with passed ones. Args: result: new data to be cached, data types is defined by result_handler, result_handler: handler to process cached data of the specific type. extra_result_data: some known data to be merged with the new data in combined result, extra_result_data is not being cached. key_kwargs: kwargs to add for cache keys. """ # get key parts key_parts = get_key_parts(key_kwargs or {}) # get mapping for caching and combined result caching_map, combined_result = result_handler.get_results(result, extra_result_data or {}, key_parts) created_at = datetime.utcnow() # set data to cache if context.BACKGROUND_SAVE.get(): for collection_name, records in caching_map.items(): await spawn(context.REQUEST.get(), save_items(collection_name, records, created_at)) else: tasks = [save_items(cn, dm, created_at) for cn, dm in caching_map.items()] await asyncio.gather(*tasks) if config.DEBUG_DATA_SOURCE: set_debug_data_source( DataSource.COMBINED if extra_result_data else DataSource.API, created_at, created_at, ) set_debug_request_count(1, None, created_at, force_replace=False) return combined_result