diff --git a/changelog/651.bugfix.rst b/changelog/651.bugfix.rst new file mode 100644 index 000000000..a74129c6a --- /dev/null +++ b/changelog/651.bugfix.rst @@ -0,0 +1 @@ +Fix sorting of UnifiedResponse table diff --git a/dkist/net/client.py b/dkist/net/client.py index e22d3df30..7db1616c8 100644 --- a/dkist/net/client.py +++ b/dkist/net/client.py @@ -32,6 +32,17 @@ __all__ = ["DKISTClient", "DKISTQueryResponseTable"] +def process_nones(results, key, replacement=np.nan): + # We need to replace Nones with nans here for sorting purposes + # We also need to recreate the whole row so that it can have a numerical dtype + # Without this is doesn't sort properly and any nans up in strange places + if key in results.colnames: + old_r = results[key] + results[key] = [replacement] * len(results) + notnone = results[key] != None + results[key][notnone] = old_r[notnone] + + class DKISTQueryResponseTable(QueryResponseTable): """ Results of a DKIST Dataset search. @@ -98,10 +109,7 @@ def _process_table(results: "DKISTQueryResponseTable") -> "DKISTQueryResponseTab results[colname][none_values] = np.nan results[colname] = u.Quantity(results[colname], unit=unit) - if "Average Fried Parameter" in results.colnames: - r_none_values = np.array(results["Average Fried Parameter"] == None) - if r_none_values.any(): - results["Average Fried Parameter"][r_none_values] = np.nan + process_nones(results, "Average Fried Parameter") if results and "Wavelength" not in results.colnames: results["Wavelength"] = u.Quantity([results["Wavelength Min"], results["Wavelength Max"]]).T diff --git a/dkist/net/tests/test_client.py b/dkist/net/tests/test_client.py index 55266767f..879766c0d 100644 --- a/dkist/net/tests/test_client.py +++ b/dkist/net/tests/test_client.py @@ -1,3 +1,4 @@ +import copy import json import hypothesis.strategies as st @@ -104,6 +105,20 @@ def example_api_response(): } +@pytest.fixture +def example_api_response_multiple_r0(example_api_response): + """ + A larger dummy API response with varied Fried parameter values to test sorting + """ + for _ in range(4): + example_api_response["searchResults"].append(copy.copy(example_api_response["searchResults"][0])) + example_api_response["searchResults"][1]["qualityAverageFriedParameter"] = 5 + example_api_response["searchResults"][3]["qualityAverageFriedParameter"] = 1 + example_api_response["searchResults"][4]["qualityAverageFriedParameter"] = 3 + + return example_api_response + + @pytest.fixture def expected_table_keys(): translated_keys = set(INVENTORY_KEY_MAP.values()) @@ -151,6 +166,13 @@ def test_query_response_from_results(empty_query_response, example_api_response, assert np.isnan(qr["Average Fried Parameter"][0]) +def test_sort_fried_parameter(example_api_response_multiple_r0): + qr = DKISTQueryResponseTable.from_results([example_api_response_multiple_r0], client=DKISTClient()) + qr.sort("Average Fried Parameter") + assert all(qr["Average Fried Parameter"][:3] == [1.0, 3.0, 5.0]) + assert all(np.isnan(qr["Average Fried Parameter"][3:])) + + def test_query_response_from_results_unknown_field(empty_query_response, example_api_response, expected_table_keys): """ This test asserts that if the API starts returning new fields we don't error, they get passed though verbatim.