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
1 change: 1 addition & 0 deletions changelog/651.bugfix.rst
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Fix sorting of UnifiedResponse table
16 changes: 12 additions & 4 deletions dkist/net/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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
Expand Down
22 changes: 22 additions & 0 deletions dkist/net/tests/test_client.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import copy
import json

import hypothesis.strategies as st
Expand Down Expand Up @@ -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())
Expand Down Expand Up @@ -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.
Expand Down
Loading