import pytest import flask import types from typing import TypedDict from unittest.mock import call, patch, MagicMock from marshmallow import Schema, fields from charts.api import app from charts.connectors import redis from charts.utils import dataloader def test_dataloader_not_a_dataloader_error() -> None: """Test dataloader decorator parameter type assertion.""" @dataloader.dataloader def test_func(input: list[None]) -> list[dataloader.DataLoaderResponse[None]]: return [] with pytest.raises(Exception, match='Not a dataloader.'): # test is to assert a runtime type check test_func(1) # type: ignore[arg-type] def test_dataloader_incorrect_row_count_error() -> None: """Test dataloader decorator return length assertion.""" @dataloader.dataloader def test_func(input: list[int]) -> list[dataloader.DataLoaderResponse[str]]: return [{'data': 'a'}] with pytest.raises(Exception, match='Violated dataloader contract: Incorrect row count returned.'): test_func([]) with pytest.raises(Exception, match='Violated dataloader contract: Incorrect row count returned.'): test_func([1, 2]) def test_dataloader_passthrough() -> None: """Test dataloader decorator return passthrough on success.""" @dataloader.dataloader def test_func(input: list[int]) -> list[dataloader.DataLoaderResponse[str]]: return [{'data': 'a'}] response = test_func([1]) assert response == [{'data': 'a'}] class TestSchema(Schema): field = fields.String(required=True) def test_marshmallow_redis_dataloader_redis_read() -> None: """Tests MarshmallowRedisDataloader redis_read function.""" class ReturnType(TypedDict): field: str @dataloader.dataloader def mock_dl(input: list[str], arg1: str, /, *, kwarg1: str) -> list[dataloader.DataLoaderResponse[ReturnType]]: raise Exception('Expected not to be called.') dl = dataloader.MarshmallowRedisDataLoader( dataloader=mock_dl, expiry=1000, key_serializer=lambda key, arg1, /, *, kwarg1: f'{arg1}-{kwarg1}-{key}', schema=TestSchema ) with patch.object(redis.client, 'mget') as mock_redis_mget, flask.Flask(__name__).test_request_context(): app.logger.warning = types.SimpleNamespace() app.logger.warning = MagicMock() mock_redis_mget.return_value = [None, 'not-json', '{"not_valid": ""}', '{"data": {"field": 1}}', '{"data": {"field": "value"}}'] result = dl._redis_read(['id1', 'id2', 'id3', 'id4', 'id5'], 'arg1', kwarg1='kwarg1') assert result == [None, None, None, None, {'data': {'field': 'value'}}] mock_redis_mget.assert_called_once_with([ 'arg1-kwarg1-id1', 'arg1-kwarg1-id2', 'arg1-kwarg1-id3', 'arg1-kwarg1-id4', 'arg1-kwarg1-id5', ]) app.logger.warning.assert_has_calls(( call('Loaded invalid JSON value from redis for key arg1-kwarg1-id2: Expecting value: line 1 column 1 (char 0)', ), call("Loaded invalid schema value from redis for key arg1-kwarg1-id3: {'data': ['Missing data for required field.'], 'not_valid': ['Unknown field.']}", ), call("Loaded invalid schema value from redis for key arg1-kwarg1-id4: {'data': {'field': ['Not a valid string.']}}", ), )) def test_marshmallow_redis_dataloader_load_many_no_cache_read() -> None: """Tests MarshmallowRedisDataloader load_many_no_cache_read function.""" class ReturnType(TypedDict): field: str @dataloader.dataloader def mock_dl(input: list[str], arg1: str, /, *, kwarg1: str) -> list[dataloader.DataLoaderResponse[ReturnType]]: return [ { 'data': None }, { 'data': { 'field': 'value', }, }, ] dl = dataloader.MarshmallowRedisDataLoader( dataloader=mock_dl, expiry=1000, key_serializer=lambda key, arg1, /, *, kwarg1: f'{arg1}-{kwarg1}-{key}', schema=TestSchema ) with patch.object(redis.client, 'setex') as mock_redis_setex: result = dl.load_many_no_cache_read(['id1', 'id2'], 'arg1', kwarg1='kwarg1') assert result == [ { 'data': None }, { 'data': { 'field': 'value', }, }, ] mock_redis_setex.assert_has_calls([ call('arg1-kwarg1-id2', 1000, '{"data": {"field": "value"}}'), ]) def test_marshmallow_redis_dataloader_load_many() -> None: """Tests MarshmallowRedisDataloader load_many function.""" class ReturnType(TypedDict): field: str @dataloader.dataloader def mock_dl(input: list[str], arg1: str, /, *, kwarg1: str) -> list[dataloader.DataLoaderResponse[ReturnType]]: raise Exception('Expected not to be called.') dl = dataloader.MarshmallowRedisDataLoader( dataloader=mock_dl, expiry=1000, key_serializer=lambda key, arg1, /, *, kwarg1: f'{arg1}-{kwarg1}-{key}', schema=TestSchema ) with patch.object(dl, '_redis_read') as mock_dl_redis_read,\ patch.object(dl, 'load_many_no_cache_read') as mock_dl_load_many_no_cache_read: mock_dl_redis_read.return_value = [ None, { 'data': { 'field': 'value2', }, }, None, { 'data': { 'field': 'value4', }, }, None, ] mock_dl_load_many_no_cache_read.return_value = [ { 'error': 'An error occurred', }, { 'data': None }, { 'data': { 'field': 'value5', }, }, ] result = dl.load_many(['id1', 'id2', 'id3', 'id4', 'id5'], 'arg1', kwarg1='kwarg1') assert result == [ { 'error': 'An error occurred', }, { 'data': { 'field': 'value2', }, }, { 'data': None }, { 'data': { 'field': 'value4', }, }, { 'data': { 'field': 'value5', }, }, ] mock_dl_redis_read.assert_called_once_with(['id1', 'id2', 'id3', 'id4', 'id5'], 'arg1', kwarg1='kwarg1') mock_dl_load_many_no_cache_read.assert_called_once_with(['id1', 'id3', 'id5'], 'arg1', kwarg1='kwarg1')