diff --git a/pygeoapi/util.py b/pygeoapi/util.py index b60e187a0..67f17bea2 100644 --- a/pygeoapi/util.py +++ b/pygeoapi/util.py @@ -426,10 +426,12 @@ def json_serial(obj: Any) -> str: return base64.b64encode(obj) elif isinstance(obj, Decimal): return float(obj) - elif type(obj).__name__ in ['int32', 'int64']: + elif _is_numpy_scalar(obj, 'integer'): return int(obj) - elif type(obj).__name__ in ['float32', 'float64']: + elif _is_numpy_scalar(obj, 'floating'): return float(obj) + elif _is_numpy_scalar(obj, 'bool', 'bool_'): + return bool(obj) elif isinstance(obj, l10n.Locale): return l10n.locale2str(obj) elif isinstance(obj, (pathlib.PurePath, Path)): @@ -442,6 +444,22 @@ def json_serial(obj: Any) -> str: raise TypeError(msg) +def _is_numpy_scalar(obj: Any, *type_names: str) -> bool: + """ + helper function to check whether an object is a NumPy scalar of + a given abstract type (e.g. `integer` covers `int8` to `uint64`), + without requiring NumPy to be installed + + :param obj: `object` to be evaluated + :param type_names: NumPy type names to match + + :returns: `bool` of whether the object is a matching NumPy scalar + """ + + return any(cls.__module__ == 'numpy' and cls.__name__ in type_names + for cls in type(obj).__mro__) + + def is_url(urlstring: str) -> bool: """ Validation function that determines whether a candidate URL should be diff --git a/tests/other/test_util.py b/tests/other/test_util.py index 7cb321019..9c8a9edab 100644 --- a/tests/other/test_util.py +++ b/tests/other/test_util.py @@ -31,10 +31,12 @@ from decimal import Decimal from copy import deepcopy from io import StringIO +import json from unittest import mock import uuid from xml.sax.saxutils import unescape +import numpy as np import pytest from pygeoapi import util @@ -367,3 +369,24 @@ def test_format_datetime(value, format_, result): ]) def test_format_duration(start, end, result): assert util.format_duration(start, end) == result + + +@pytest.mark.parametrize('value,expected', [ + (np.int8(-3), -3), + (np.int16(7), 7), + (np.int32(7), 7), + (np.int64(7), 7), + (np.uint8(255), 255), + (np.uint64(2**64 - 1), 2**64 - 1), + (np.float16(0.5), 0.5), + (np.float32(0.5), 0.5), + (np.float64(0.5), 0.5), + (np.bool_(True), True), +]) +def test_json_serial_numpy_scalars(value, expected): + result = util.json_serial(value) + assert result == expected + assert type(result) is type(expected) + assert json.loads(json.dumps({'value': value}, + default=util.json_serial)) == { + 'value': expected}