Skip to content
Open
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
22 changes: 20 additions & 2 deletions pygeoapi/util.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)):
Expand All @@ -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
Expand Down
23 changes: 23 additions & 0 deletions tests/other/test_util.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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}
Loading