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
36 changes: 22 additions & 14 deletions ssapy/compute.py
Original file line number Diff line number Diff line change
Expand Up @@ -140,23 +140,26 @@ def _countTime(time):


def _countR(r):
# orbit is one of:
# 1) scalar r
# 2) vector r
# 3) list of scalar Orbit
# convert to (2), set nOrbit, squeezeOrbit, and orbit.
# r is one of:
# 1) one position, shape (3,)
# 2) one trajectory, shape (nTime, 3)
# 3) multiple trajectories, shape (nOrbit, nTime, 3)
r = np.asarray(r, dtype=float)
squeezeR = False
if np.shape(r)[-1] == 3:
pass
else:
raise ValueError(f"Incorrect r dimensions. Expected shape (n, 3), but got {np.shape(r)}.")
# check 1) and 2)
if r.ndim < 3: # scalar r
if r.ndim == 1 and r.shape == (3,):
nR = 1
r = np.reshape(np.atleast_3d(r), (nR, np.shape(r)[0], np.shape(r)[1]))
r = r.reshape(1, 1, 3)
squeezeR = True
elif r.ndim == 2 and r.shape[-1] == 3:
nR = 1
r = r[None, ...]
squeezeR = True
elif r.ndim == 3 and r.shape[-1] == 3:
nR = r.shape[0]
else:
nR = np.shape(r)[0]
raise ValueError(
f"Incorrect r dimensions. Expected shape (3,), (n, 3), or (m, n, 3), but got {r.shape}."
)
return nR, squeezeR, r


Expand Down Expand Up @@ -392,7 +395,12 @@ def groundTrack(orbit, time, propagator=KeplerianPropagator(), format='geodetic'
raise ValueError("Format must be either 'cartesian' or 'geodetic'")

nTime, squeezeTime, time = _countTime(time)
if isinstance(orbit, Orbit):
orbit_sequence = (
isinstance(orbit, (list, tuple))
and len(orbit) > 0
and all(isinstance(item, Orbit) for item in orbit)
)
if isinstance(orbit, Orbit) or orbit_sequence:
nOrbit, squeezeOrbit, orbit = _countOrbit(orbit)
r, v = rv(orbit, time, propagator=propagator) # (n, m, 3)
else:
Expand Down
38 changes: 38 additions & 0 deletions tests/test_compute.py
Original file line number Diff line number Diff line change
Expand Up @@ -323,3 +323,41 @@ def test_import_erfa_present():
assert ssapy.compute.erfa is fake_erfa
"""
subprocess.run([sys.executable, "-c", code], check=True)


def test_groundTrack_accepts_list_of_scalar_orbits(monkeypatch):
_patch_identity_ground_track_frame(monkeypatch)
orbits = [
ssapy.Orbit(np.array([7.0e6, 0.0, 0.0]), np.array([0.0, 7.5e3, 0.0]), 0.0),
ssapy.Orbit(np.array([7.1e6, 0.0, 0.0]), np.array([0.0, 7.4e3, 0.0]), 0.0),
]
time = np.array([0.0, 1.0])

with pytest.warns(DeprecationWarning, match="list of Orbit syntax"):
x, y, z = groundTrack(orbits, time, format="cartesian")
expected = np.stack([ssapy.rv(orbit, time)[0] for orbit in orbits])

np.testing.assert_allclose(x, expected[..., 0])
np.testing.assert_allclose(y, expected[..., 1])
np.testing.assert_allclose(z, expected[..., 2])


def test_groundTrack_accepts_python_position_lists(monkeypatch):
_patch_identity_ground_track_frame(monkeypatch)
positions = [[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]]

x, y, z = groundTrack(positions, [0.0, 1.0], format="cartesian")

np.testing.assert_allclose(x, [1.0, 4.0])
np.testing.assert_allclose(y, [2.0, 5.0])
np.testing.assert_allclose(z, [3.0, 6.0])


def test_groundTrack_accepts_single_position_list(monkeypatch):
_patch_identity_ground_track_frame(monkeypatch)

x, y, z = groundTrack([1.0, 2.0, 3.0], 0.0, format="cartesian")

assert x == pytest.approx(1.0)
assert y == pytest.approx(2.0)
assert z == pytest.approx(3.0)
Loading