Skip to content

Commit 13b2727

Browse files
committed
Fix issue with string types in files.get_elems()
1 parent d8010ac commit 13b2727

2 files changed

Lines changed: 44 additions & 7 deletions

File tree

src/epr/_epr.pyx

Lines changed: 12 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -758,18 +758,23 @@ cdef class Field(EprObject):
758758
pyepr_null_ptr_error(msg)
759759
elif etype == e_tid_string:
760760
if shape[0] != 1:
761-
raise ValueError(f"unexpected number of elements: {shape[0]}")
762-
ndim = 0
763-
dtype = np.NPY_STRING
764-
buf = <char*>epr_get_field_elem_as_str(self._ptr)
765-
if buf is NULL:
766-
pyepr_null_ptr_error(msg)
761+
raise ValueError(
762+
f"unexpected number of elements: {shape[0]} "
763+
"for e_tid_string"
764+
)
765+
elem = self.get_elem()
766+
return np.asarray(elem)
767+
# ndim = 0
768+
# dtype = np.NPY_STRING
769+
# buf = <char*>epr_get_field_elem_as_str(self._ptr)
770+
# if buf is NULL:
771+
# pyepr_null_ptr_error(msg)
767772
# elif etype == e_tid_unknown:
768773
# pass
769774
# elif etype = e_tid_spare:
770775
# pass
771776
else:
772-
raise ValueError("invalid field type")
777+
raise ValueError(f"invalid field type: {etype}")
773778

774779
out = np.PyArray_SimpleNewFromData(ndim, shape, dtype, <void*>buf)
775780
# np.PyArray_CLEARFLAG(out, NPY_ARRAY_WRITEABLE) # new in numpy 1.7

tests/test_all.py

Lines changed: 32 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2134,6 +2134,38 @@ def test_tot_size(self):
21342134
self.assertEqual(self.field.tot_size, tot_size)
21352135

21362136

2137+
class TestFieldGetElems(unittest.TestCase):
2138+
def setUp(self):
2139+
self.product = epr.Product(PRODUCT_FILE)
2140+
self.record = self.product.get_sph()
2141+
2142+
def test_string_get_elem_int(self):
2143+
field = self.record.get_field("LINE_LENGTH")
2144+
value = field.get_elem()
2145+
self.assertIsInstance(value, int)
2146+
self.assertEqual(value, 1452)
2147+
2148+
def test_string_get_elem_str(self):
2149+
field = self.record.get_field("SPH_DESCRIPTOR")
2150+
value = field.get_elem()
2151+
self.assertIsInstance(value, bytes)
2152+
self.assertEqual(value, b"AP Mode Medium Res. Image")
2153+
2154+
def test_string_get_elems_int(self):
2155+
field = self.record.get_field("LINE_LENGTH")
2156+
value = field.get_elems()
2157+
self.assertIsInstance(value, np.ndarray)
2158+
self.assertTrue(np.isdtype(value.dtype, "integral"))
2159+
self.assertEqual(value.item(), 1452)
2160+
2161+
def test_string_get_elems_str(self):
2162+
field = self.record.get_field("SPH_DESCRIPTOR")
2163+
value = field.get_elems()
2164+
self.assertIsInstance(value, np.ndarray)
2165+
self.assertFalse(np.isdtype(value.dtype, "numeric"))
2166+
self.assertEqual(value.item(), b"AP Mode Medium Res. Image")
2167+
2168+
21372169
class TestFieldRW(TestField):
21382170
OPEN_MODE = "rb+"
21392171

0 commit comments

Comments
 (0)