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
4 changes: 2 additions & 2 deletions spatialmath/base/argcheck.py
Original file line number Diff line number Diff line change
Expand Up @@ -510,8 +510,8 @@ def isvector(v: Any, dim: Optional[int] = None) -> bool:
if dim is None:
return (
(len(s) == 1 and s[0] > 0)
or (s[0] == 1 and s[1] > 0)
or (s[0] > 0 and s[1] == 1)
or (len(s) == 2 and s[0] == 1 and s[1] > 0)
or (len(s) == 2 and s[0] > 0 and s[1] == 1)
)
else:
return s == (dim,) or s == (1, dim) or s == (dim, 1)
Expand Down
23 changes: 22 additions & 1 deletion tests/base/test_argcheck.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
import numpy.testing as nt

from spatialmath.base.argcheck import *
from spatialmath.base.transformsNd import e2h, h2e


class Test_check(unittest.TestCase):
Expand Down Expand Up @@ -212,10 +213,30 @@ def test_isvector(self):
self.assertFalse(isvector(np.array([[1, 2, 3]]), 4))
self.assertFalse(isvector(np.array([[1], [2], [3]]), 4))

def test_isvector(self):
def test_assertvector(self):
l = [1, 2, 3]
nt.assert_raises(ValueError, assertvector, l, 4)

def test_isvector_unsupported_array_shapes(self):
for shape in [(), (2, 3), (1, 3, 2), (3, 1, 2), (1, 1, 1), (1, 3, 0)]:
with self.subTest(shape=shape):
value = np.ones(shape)
self.assertFalse(isvector(value))
with self.assertRaises(ValueError):
assertvector(value)

def test_isvector_empty_arrays(self):
for shape in [(0,), (1, 0), (0, 1), (0, 0)]:
with self.subTest(shape=shape):
self.assertFalse(isvector(np.empty(shape)))

def test_homogeneous_conversion_unsupported_array_shapes(self):
for convert in [e2h, h2e]:
for shape in [(), (1, 3, 2), (3, 1, 2), (1, 1, 1)]:
with self.subTest(convert=convert.__name__, shape=shape):
with self.assertRaises(ValueError):
convert(np.ones(shape))

def test_getvector(self):
l = [1, 2, 3]
t = (1, 2, 3)
Expand Down
Loading