From af80c4c5f7200166bdb004a064b1e4e956e7e769 Mon Sep 17 00:00:00 2001 From: DanielElisenberg Date: Fri, 11 Sep 2026 16:01:37 +0200 Subject: [PATCH] fix_data_endpoint_filters: make filters less clever --- datastore_api/api/datastores/data/__init__.py | 14 ++- datastore_api/domain/data/__init__.py | 49 +++++----- tests/unit/api/data/test_data_routes.py | 95 +++++++++++++++++-- tests/unit/domain/data/test_data.py | 24 ++--- .../unit/domain/data/test_data_big_parquet.py | 8 +- 5 files changed, 140 insertions(+), 50 deletions(-) diff --git a/datastore_api/api/datastores/data/__init__.py b/datastore_api/api/datastores/data/__init__.py index 14031a7..ee02ac0 100644 --- a/datastore_api/api/datastores/data/__init__.py +++ b/datastore_api/api/datastores/data/__init__.py @@ -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, @@ -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, @@ -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, @@ -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, diff --git a/datastore_api/domain/data/__init__.py b/datastore_api/domain/data/__init__.py index a0de142..89875a5 100644 --- a/datastore_api/domain/data/__init__.py +++ b/datastore_api/domain/data/__init__.py @@ -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, + ) diff --git a/tests/unit/api/data/test_data_routes.py b/tests/unit/api/data/test_data_routes.py index 18028a6..33477d5 100644 --- a/tests/unit/api/data/test_data_routes.py +++ b/tests/unit/api/data/test_data_routes.py @@ -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 @@ -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() @@ -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. + 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() diff --git a/tests/unit/domain/data/test_data.py b/tests/unit/domain/data/test_data.py index cbe8df4..8044010 100644 --- a/tests/unit/domain/data/test_data.py +++ b/tests/unit/domain/data/test_data.py @@ -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 ( @@ -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( @@ -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( @@ -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( @@ -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( @@ -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( @@ -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( @@ -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( @@ -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", @@ -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", @@ -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( diff --git a/tests/unit/domain/data/test_data_big_parquet.py b/tests/unit/domain/data/test_data_big_parquet.py index f29d74c..f85eb94 100644 --- a/tests/unit/domain/data/test_data_big_parquet.py +++ b/tests/unit/domain/data/test_data_big_parquet.py @@ -12,7 +12,7 @@ InputFixedQuery, UnencryptedDataReader, _get_parquet_path, - generate_data_filter, + generate_fixed_filter, ) DATASTORE_DIR = Path("tests/resources/test_datastore") @@ -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, 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(