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
14 changes: 10 additions & 4 deletions datastore_api/api/datastores/data/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,9 @@
from datastore_api.config import environment
from datastore_api.domain.data import (
DataReader,
generate_data_filter,
generate_fixed_filter,
generate_time_filter,
generate_time_period_filter,
)
from datastore_api.domain.data.models import (
ErrorMessage,
Expand All @@ -36,7 +38,9 @@
def stream_result_event(
input_query: InputTimePeriodQuery,
data_reader: Annotated[DataReader, Depends(get_data_reader)],
data_filter: Annotated[dataset.Expression, Depends(generate_data_filter)],
data_filter: Annotated[
dataset.Expression, Depends(generate_time_period_filter)
],
) -> PlainTextResponse:
"""
Create Result set of data with temporality type event,
Expand All @@ -57,7 +61,7 @@ def stream_result_event(
def stream_result_status(
input_query: InputTimeQuery,
data_reader: Annotated[DataReader, Depends(get_data_reader)],
data_filter: Annotated[dataset.Expression, Depends(generate_data_filter)],
data_filter: Annotated[dataset.Expression, Depends(generate_time_filter)],
) -> PlainTextResponse:
"""
Create result set of data with temporality type status,
Expand All @@ -80,7 +84,9 @@ def stream_result_status(
def stream_result_fixed(
input_query: InputFixedQuery,
data_reader: Annotated[DataReader, Depends(get_data_reader)],
data_filter: Annotated[dataset.Expression, Depends(generate_data_filter)],
data_filter: Annotated[
dataset.Expression | None, Depends(generate_fixed_filter)
],
) -> PlainTextResponse:
"""
Create result set of data with temporality type fixed,
Expand Down
49 changes: 27 additions & 22 deletions datastore_api/domain/data/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -194,26 +194,31 @@ def select_data_reader(
return UnencryptedDataReader(parquet_path=parquet_path, columns=columns)


def generate_data_filter(
input_query: InputTimePeriodQuery | InputTimeQuery | InputFixedQuery,
def generate_fixed_filter(
input_query: InputFixedQuery,
) -> dataset.Expression | None:
return filters.generate_fixed_filter(
population_filter=input_query.population,
value_filter=input_query.values,
)


def generate_time_filter(
input_query: InputTimeQuery,
) -> dataset.Expression:
if isinstance(input_query, InputTimePeriodQuery):
return filters.generate_time_period_filter(
start=input_query.startDate,
stop=input_query.stopDate,
population_filter=input_query.population,
value_filter=input_query.values,
)
elif isinstance(input_query, InputTimeQuery):
return filters.generate_time_filter(
date=input_query.date,
population_filter=input_query.population,
value_filter=input_query.values,
)
elif isinstance(input_query, InputFixedQuery):
return filters.generate_fixed_filter(
population_filter=input_query.population,
value_filter=input_query.values,
)
else:
raise ValueError("Unsupported query type")
return filters.generate_time_filter(
date=input_query.date,
population_filter=input_query.population,
value_filter=input_query.values,
)


def generate_time_period_filter(
input_query: InputTimePeriodQuery,
) -> dataset.Expression:
return filters.generate_time_period_filter(
start=input_query.startDate,
stop=input_query.stopDate,
population_filter=input_query.population,
value_filter=input_query.values,
)
95 changes: 86 additions & 9 deletions tests/unit/api/data/test_data_routes.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,18 +11,15 @@
from datastore_api.api.common.dependencies import (
get_data_reader,
)
from datastore_api.domain.data import (
generate_data_filter,
)
from datastore_api.main import app

FAKE_RESULT_FILE_NAME = "fake_result_file_name"
MOCK_RESULT = pq.read_table("tests/resources/results/mocked_result.parquet")


class FakeDataReader:
def read_data(self, data_filter, *, row_cap=None):
return MOCK_RESULT
@pytest.fixture
def mock_data_reader():
return Mock(read_data=Mock(return_value=MOCK_RESULT))


@pytest.fixture
Expand All @@ -42,11 +39,10 @@ def mock_auth_deps():


@pytest.fixture
def client(mock_db_client: Mock, mock_auth_deps: dict):
def client(mock_db_client: Mock, mock_auth_deps: dict, mock_data_reader: Mock):
app.dependency_overrides[db.get_database_client] = lambda: mock_db_client
app.dependency_overrides[authorize_user] = lambda: mock_auth_deps["user"]()
app.dependency_overrides[generate_data_filter] = lambda: None
app.dependency_overrides[get_data_reader] = FakeDataReader
app.dependency_overrides[get_data_reader] = lambda: mock_data_reader
yield TestClient(app)
app.dependency_overrides.clear()

Expand Down Expand Up @@ -94,3 +90,84 @@ def test_data_fixed_stream_result(client: TestClient, mock_auth_deps: dict):
reader = pa.BufferReader(response.content)
assert response.status_code == 200
assert pq.read_table(reader) == MOCK_RESULT


@pytest.mark.parametrize(
"temporality, expected_ids, filtered_ids",
[
("fixed", [1, 2, 3, 4, 5], [1, 3, 4]),
("status", [3, 5], [3]),
("event", [1, 2, 3, 5], [1, 3]),
],
)
@pytest.mark.parametrize("with_filters", [False, True])
def test_data_stream_uses_endpoint_filter(
client,
mock_data_reader,
temporality,
expected_ids,
filtered_ids,
with_filters,
):
table = pa.table(
{
"unit_id": [1, 2, 3, 4, 5],
"value": ["A", "A", "A", "A", "B"],
"start_epoch_days": [0, 5, 10, 20, 0],
"stop_epoch_days": [4, 9, None, None, None],
}
)
# Extra date fields must not change which filter the endpoint uses.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

💯

payload = {
"version": "1.0.0.0",
"dataStructureName": "FAKE_NAME",
"date": 10,
"startDate": 4,
"stopDate": 10,
}
if with_filters:
payload.update(population=[1, 3, 4, 5], values=["A"])

response = client.post(
f"/datastores/no.ssb.test/data/{temporality}/stream",
json=payload,
)

assert response.status_code == 200
mock_data_reader.read_data.assert_called_once()
data_filter = mock_data_reader.read_data.call_args.args[0]
if temporality == "fixed" and not with_filters:
assert data_filter is None
else:
table = table.filter(data_filter)
assert table["unit_id"].to_pylist() == (
filtered_ids if with_filters else expected_ids
)


@pytest.mark.parametrize(
"temporality, dates, missing_field",
[
("status", {}, "date"),
("event", {"stopDate": 10}, "startDate"),
("event", {"startDate": 4}, "stopDate"),
],
)
def test_data_stream_requires_dates(
client, mock_data_reader, temporality, dates, missing_field
):
response = client.post(
f"/datastores/no.ssb.test/data/{temporality}/stream",
json={
"version": "1.0.0.0",
"dataStructureName": "FAKE_NAME",
**dates,
},
)

assert response.status_code == 400
assert any(
error["loc"] == ["body", missing_field]
for error in response.json()["details"]
)
mock_data_reader.read_data.assert_not_called()
24 changes: 13 additions & 11 deletions tests/unit/domain/data/test_data.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,9 @@
EncryptedDataReader,
UnencryptedDataReader,
_get_parquet_path,
generate_data_filter,
generate_fixed_filter,
generate_time_filter,
generate_time_period_filter,
select_data_reader,
)
from datastore_api.domain.data.models import (
Expand Down Expand Up @@ -86,7 +88,7 @@ def fixed_dataset_parquet(scope="module"):

def test_valid_event_request():
payload = test_resources.VALID_EVENT_QUERY_PERSON_INCOME_ALL
data_filter = generate_data_filter(payload)
data_filter = generate_time_period_filter(payload)
columns = ALL_COLUMNS if payload.includeAttributes else ALL_COLUMNS[:2]
file_name = UnencryptedDataReader(
parquet_path=_get_parquet_path(
Expand All @@ -103,7 +105,7 @@ def test_valid_event_request():

def test_valid_event_request_partitioned():
payload = test_resources.VALID_EVENT_QUERY_TEST_STUDIEPOENG_ALL
data_filter = generate_data_filter(payload)
data_filter = generate_time_period_filter(payload)
columns = ALL_COLUMNS if payload.includeAttributes else ALL_COLUMNS[:2]
file_name = UnencryptedDataReader(
parquet_path=_get_parquet_path(
Expand All @@ -120,7 +122,7 @@ def test_valid_event_request_partitioned():

def test_event_request_causing_empty_result():
payload = test_resources.INVALID_EVENT_QUERY_INVALID_STOP_DATE
data_filter = generate_data_filter(payload)
data_filter = generate_time_period_filter(payload)
columns = ALL_COLUMNS if payload.includeAttributes else ALL_COLUMNS[:2]
result = UnencryptedDataReader(
parquet_path=_get_parquet_path(
Expand All @@ -137,7 +139,7 @@ def test_event_request_causing_empty_result():

def test_valid_status_request():
payload = test_resources.VALID_STATUS_QUERY_PERSON_INCOME_LAST_ROW
data_filter = generate_data_filter(payload)
data_filter = generate_time_filter(payload)
columns = ALL_COLUMNS if payload.includeAttributes else ALL_COLUMNS[:2]
file_name = UnencryptedDataReader(
parquet_path=_get_parquet_path(
Expand All @@ -154,7 +156,7 @@ def test_valid_status_request():

def test_invalid_status_request():
payload = test_resources.INVALID_STATUS_QUERY_NOT_FOUND
data_filter = generate_data_filter(payload)
data_filter = generate_time_filter(payload)
columns = ALL_COLUMNS if payload.includeAttributes else ALL_COLUMNS[:2]
with pytest.raises(NotFoundException) as e:
UnencryptedDataReader(
Expand All @@ -172,7 +174,7 @@ def test_invalid_status_request():

def test_valid_fixed_request():
payload = test_resources.VALID_FIXED_QUERY_PERSON_INCOME_ALL
data_filter = generate_data_filter(payload)
data_filter = generate_fixed_filter(payload)
columns = ALL_COLUMNS if payload.includeAttributes else ALL_COLUMNS[:2]
file_name = UnencryptedDataReader(
parquet_path=_get_parquet_path(
Expand All @@ -189,7 +191,7 @@ def test_valid_fixed_request():

def test_invalid_fixed_request():
payload = test_resources.INVALID_FIXED_QUERY_NOT_FOUND
data_filter = generate_data_filter(payload)
data_filter = generate_fixed_filter(payload)
columns = ALL_COLUMNS if payload.includeAttributes else ALL_COLUMNS[:2]
with pytest.raises(NotFoundException) as e:
UnencryptedDataReader(
Expand Down Expand Up @@ -420,7 +422,7 @@ def test_read_parquet_time_with_pop_filter():

def test_read_parquet_with_exact_string_value_filter(fixed_dataset_parquet):
expected_values = ["0012", "0100"]
data_filter = generate_data_filter(
data_filter = generate_fixed_filter(
InputFixedQuery(
values=expected_values,
dataStructureName="TEST_FIXED_DATASET",
Expand All @@ -446,7 +448,7 @@ def test_read_parquet_with_exact_string_value_filter(fixed_dataset_parquet):

def test_read_parquet_with_wildcard_value_filter(fixed_dataset_parquet):
expected_values = ["0020", "0025", "2100"]
data_filter = generate_data_filter(
data_filter = generate_fixed_filter(
InputFixedQuery(
values=["002*", "2*"],
dataStructureName="TEST_FIXED_DATASET",
Expand Down Expand Up @@ -481,7 +483,7 @@ def test_read_parquet_with_combined_population_and_value_filter(
includeAttributes=True,
values=["001*"],
)
data_filter = generate_data_filter(payload)
data_filter = generate_fixed_filter(payload)
columns = ALL_COLUMNS if payload.includeAttributes else ALL_COLUMNS[:2]
result_dict = (
UnencryptedDataReader(
Expand Down
8 changes: 4 additions & 4 deletions tests/unit/domain/data/test_data_big_parquet.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@
InputFixedQuery,
UnencryptedDataReader,
_get_parquet_path,
generate_data_filter,
generate_fixed_filter,
)

DATASTORE_DIR = Path("tests/resources/test_datastore")
Expand Down Expand Up @@ -86,11 +86,11 @@ def test_read_big_parquet_with_big_pop_and_value_filter(
payload = InputFixedQuery(
dataStructureName=DATASET_NAME,
version=Version.from_str("1.0.0.0"), # NOSONAR
population=[1, 3],
population=population_filter,

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🐛

includeAttributes=True,
values=["001*"],
values=value_filter,
)
data_filter = generate_data_filter(payload)
data_filter = generate_fixed_filter(payload)
columns = ALL_COLUMNS if payload.includeAttributes else ALL_COLUMNS[:2]
result = UnencryptedDataReader(
parquet_path=_get_parquet_path(
Expand Down