diff --git a/floodlight/io/skillcorner.py b/floodlight/io/skillcorner.py index 97a17b44..403b5ff6 100644 --- a/floodlight/io/skillcorner.py +++ b/floodlight/io/skillcorner.py @@ -236,7 +236,7 @@ def read_position_data_json( y_column = x_column + 1 xy_objects[half]["Away"][t, x_column] = tracked_object["x"] xy_objects[half]["Away"][t, y_column] = tracked_object["y"] - elif pID is ball_id: + elif pID == ball_id: xy_objects[half]["Ball"][t, 0] = tracked_object["x"] xy_objects[half]["Ball"][t, 1] = tracked_object["y"] # do not track referee diff --git a/tests/test_io/test_skillcorner.py b/tests/test_io/test_skillcorner.py new file mode 100644 index 00000000..bb64a3dd --- /dev/null +++ b/tests/test_io/test_skillcorner.py @@ -0,0 +1,49 @@ +import json + +import numpy as np + +from floodlight.io.skillcorner import read_position_data_json + + +def test_read_position_data_associates_ball_id_by_value(tmp_path): + ball_id = 98765432109876543210987654321 + home_id = 71000000000000000000000000001 + away_id = 72000000000000000000000000002 + match = { + "home_team": {"id": 101}, + "away_team": {"id": 202}, + "referees": [], + "ball": {"trackable_object": ball_id}, + "players": [ + {"team_id": 101, "trackable_object": home_id}, + {"team_id": 202, "trackable_object": away_id}, + ], + "pitch_length": 105.0, + "pitch_width": 68.0, + } + positions = [ + { + "period": 1, + "possession": {"group": "home team", "trackable_object": home_id}, + "data": [ + {"trackable_object": home_id, "x": 1.25, "y": 2.5}, + {"trackable_object": away_id, "x": -4.75, "y": 5.5}, + {"trackable_object": ball_id, "x": 12.5, "y": -3.25}, + ], + } + ] + match_path = tmp_path / "match.json" + positions_path = tmp_path / "positions.json" + # Separate files require identifier linkage across independent JSON decodes. + match_path.write_text(json.dumps(match), encoding="utf-8") + positions_path.write_text(json.dumps(positions), encoding="utf-8") + + xy, _, _, _, _ = read_position_data_json(str(positions_path), str(match_path)) + + ball = xy["firstHalf"]["Ball"] + np.testing.assert_array_equal(ball.xy, [[12.5, -3.25]]) + assert np.isfinite(ball.xy).all() + np.testing.assert_array_equal(xy["firstHalf"]["Home"].xy, [[1.25, 2.5]]) + np.testing.assert_array_equal(xy["firstHalf"]["Away"].xy, [[-4.75, 5.5]]) + assert ball.xy.shape == (1, 2) + assert ball.framerate == 10