From 7bc68b7fc579d07cf18c25620e5eab5bc8aba3eb Mon Sep 17 00:00:00 2001 From: Julio Rodriguez <144072916+Litju@users.noreply.github.com> Date: Mon, 3 Aug 2026 23:44:24 +0000 Subject: [PATCH] fix: parse Wyscout direct streams once --- kloppy/_providers/wyscout.py | 19 +- .../event/wyscout/deserializer_v2.py | 7 +- .../event/wyscout/deserializer_v3.py | 3 +- kloppy/tests/test_wyscout.py | 224 ++++++++++++++++++ 4 files changed, 237 insertions(+), 16 deletions(-) diff --git a/kloppy/_providers/wyscout.py b/kloppy/_providers/wyscout.py index af4b19021..a413cbcc3 100644 --- a/kloppy/_providers/wyscout.py +++ b/kloppy/_providers/wyscout.py @@ -32,12 +32,15 @@ def load( Returns: The parsed event data. """ + with open_as_file(event_data) as event_data_fp: + parsed_event_data = json.load(event_data_fp) + if data_version == "V2": deserializer_class = WyscoutDeserializerV2 elif data_version == "V3": deserializer_class = WyscoutDeserializerV3 else: - deserializer_class = identify_deserializer(event_data) + deserializer_class = identify_deserializer(parsed_event_data) deserializer = deserializer_class( event_types=event_types, @@ -45,10 +48,9 @@ def load( event_factory=event_factory or get_config("event_factory"), ) - with open_as_file(event_data) as event_data_fp: - return deserializer.deserialize( - inputs=WyscoutInputs(event_data=event_data_fp), - ) + return deserializer.deserialize( + inputs=WyscoutInputs(event_data=parsed_event_data), + ) def load_open_data( @@ -91,12 +93,9 @@ def load_open_data( def identify_deserializer( - event_data: FileLike, + event_data: dict, ) -> Union[type[WyscoutDeserializerV3], type[WyscoutDeserializerV2]]: - with open_as_file(event_data) as event_data_fp: - events_with_meta = json.load(event_data_fp) - - events = events_with_meta["events"] + events = event_data["events"] first_event = events[0] deserializer = None diff --git a/kloppy/infra/serializers/event/wyscout/deserializer_v2.py b/kloppy/infra/serializers/event/wyscout/deserializer_v2.py index 2542fa589..cbd972ab1 100644 --- a/kloppy/infra/serializers/event/wyscout/deserializer_v2.py +++ b/kloppy/infra/serializers/event/wyscout/deserializer_v2.py @@ -1,8 +1,7 @@ from dataclasses import replace from datetime import timedelta -import json import logging -from typing import IO, NamedTuple, Optional +from typing import NamedTuple, Optional from kloppy.domain import ( BodyPart, @@ -457,7 +456,7 @@ def _players_to_dict(players: list[Player]): class WyscoutInputs(NamedTuple): - event_data: IO[bytes] + event_data: dict class WyscoutDeserializerV2(EventDataDeserializer[WyscoutInputs]): @@ -469,7 +468,7 @@ def _deserialize(self, inputs: WyscoutInputs) -> EventDataset: transformer = self.get_transformer() with performance_logging("load data", logger=logger): - raw_events = json.load(inputs.event_data) + raw_events = inputs.event_data for event in raw_events["events"]: if "eventId" not in event: event["eventId"] = event["eventName"] diff --git a/kloppy/infra/serializers/event/wyscout/deserializer_v3.py b/kloppy/infra/serializers/event/wyscout/deserializer_v3.py index 546747278..3cbf38532 100644 --- a/kloppy/infra/serializers/event/wyscout/deserializer_v3.py +++ b/kloppy/infra/serializers/event/wyscout/deserializer_v3.py @@ -1,7 +1,6 @@ from dataclasses import replace from datetime import datetime, timedelta, timezone from enum import Enum -import json import logging from typing import Optional import warnings @@ -763,7 +762,7 @@ def _deserialize(self, inputs: WyscoutInputs) -> EventDataset: transformer = self.get_transformer() with performance_logging("load data", logger=logger): - raw_events = json.load(inputs.event_data) + raw_events = inputs.event_data for event in raw_events["events"]: if "id" not in event: event["id"] = event["type"]["primary"] diff --git a/kloppy/tests/test_wyscout.py b/kloppy/tests/test_wyscout.py index cacb230ec..3d266c389 100644 --- a/kloppy/tests/test_wyscout.py +++ b/kloppy/tests/test_wyscout.py @@ -1,9 +1,13 @@ +from contextlib import contextmanager from datetime import datetime, timedelta, timezone +from io import BytesIO, UnsupportedOperation +import json from pathlib import Path import pytest from kloppy import wyscout +import kloppy._providers.wyscout as wyscout_provider from kloppy.domain import ( BodyPart, BodyPartQualifier, @@ -13,6 +17,7 @@ DuelQualifier, DuelType, EventDataset, + EventFactory, EventType, FormationType, GoalkeeperActionType, @@ -29,6 +34,11 @@ ShotResult, Time, ) +from kloppy.infra.serializers.event.wyscout import ( + WyscoutDeserializerV2, + WyscoutDeserializerV3, + WyscoutInputs, +) @pytest.fixture(scope="session") @@ -41,6 +51,220 @@ def event_v3_data(base_dir: Path) -> Path: return base_dir / "files" / "wyscout_events_v3.json" +class NonSeekableStream(BytesIO): + def seekable(self) -> bool: + return False + + def seek(self, *args, **kwargs): + raise UnsupportedOperation("seek") + + def tell(self): + raise UnsupportedOperation("tell") + + +class CountingEventFactory(EventFactory): + def __init__(self): + self.pass_calls = 0 + + def build_pass(self, **kwargs): + self.pass_calls += 1 + return super().build_pass(**kwargs) + + +@pytest.mark.parametrize( + ("version", "fixture_name", "record_count"), + [ + ("V2", "event_v2_data", 1835), + ("V3", "event_v3_data", 1896), + ], +) +@pytest.mark.parametrize("source_kind", ["path", "seekable", "nonseekable"]) +@pytest.mark.parametrize("automatic", [True, False]) +def test_parse_once_public_matrix( + monkeypatch, + request, + version, + fixture_name, + record_count, + source_kind, + automatic, +): + path = request.getfixturevalue(fixture_name) + if source_kind == "path": + source = path + elif source_kind == "seekable": + source = BytesIO(path.read_bytes()) + else: + source = NonSeekableStream(path.read_bytes()) + + opened_inputs = [] + opened_streams = [] + parsed_streams = [] + original_open = wyscout_provider.open_as_file + original_json_load = wyscout_provider.json.load + + @contextmanager + def counted_open(input_, mode="rb"): + opened_inputs.append(input_) + with original_open(input_, mode=mode) as stream: + opened_streams.append(stream) + yield stream + + def counted_json_load(stream): + parsed_streams.append(stream) + return original_json_load(stream) + + monkeypatch.setattr(wyscout_provider, "open_as_file", counted_open) + monkeypatch.setattr(wyscout_provider.json, "load", counted_json_load) + + dataset = wyscout.load( + event_data=source, + data_version=None if automatic else version, + ) + + assert len(dataset.records) == record_count + assert opened_inputs == [source] + assert len(opened_streams) == 1 + assert parsed_streams == opened_streams + + if source_kind != "path": + assert opened_streams[0] is source + assert not source.closed + assert source.read() == b"" + if source_kind == "seekable": + assert source.tell() == path.stat().st_size + + +@pytest.mark.parametrize( + ("version", "fixture_name", "record_count", "coordinates"), + [ + ("V2", "event_v2_data", 1835, Point(29.0, 6.0)), + ("V3", "event_v3_data", 1896, Point(32.0, 56.0)), + ], +) +def test_automatic_and_explicit_semantics_match( + request, version, fixture_name, record_count, coordinates +): + path = request.getfixturevalue(fixture_name) + + automatic = wyscout.load(event_data=path, coordinates="wyscout") + explicit = wyscout.load( + event_data=path, + coordinates="wyscout", + data_version=version, + ) + + assert len(automatic.records) == len(explicit.records) == record_count + assert automatic.metadata == explicit.metadata + assert automatic.metadata.periods == explicit.metadata.periods + assert automatic.to_records() == explicit.to_records() + assert automatic.records[2].coordinates == coordinates + assert explicit.records[2].coordinates == coordinates + + automatic_factory = CountingEventFactory() + explicit_factory = CountingEventFactory() + automatic_passes = wyscout.load( + event_data=path, + event_types=["PASS"], + event_factory=automatic_factory, + ) + explicit_passes = wyscout.load( + event_data=path, + event_types=["PASS"], + event_factory=explicit_factory, + data_version=version, + ) + + assert automatic_passes.to_records() == explicit_passes.to_records() + assert automatic_factory.pass_calls == explicit_factory.pass_calls + assert automatic_factory.pass_calls == len(automatic_passes.records) + + +@pytest.mark.parametrize("data_version", [None, "V2", "V3"]) +def test_malformed_json_preserves_error(data_version): + with pytest.raises(json.JSONDecodeError): + wyscout.load(BytesIO(b'{"events": invalid}'), data_version=data_version) + + +@pytest.mark.parametrize( + ("data_version", "exception", "message"), + [ + (None, IndexError, "list index out of range"), + ("V2", ValueError, "not enough values to unpack"), + ("V3", ValueError, "not enough values to unpack"), + ], +) +def test_empty_events_preserve_error(data_version, exception, message): + with pytest.raises(exception, match=message): + wyscout.load( + BytesIO(b'{"events": [], "teams": {}}'), data_version=data_version + ) + + +@pytest.mark.parametrize( + ("data_version", "exception", "message"), + [ + ( + None, + ValueError, + "Wyscout data version could not be recognized, please specify", + ), + ("V2", KeyError, "eventName"), + ("V3", KeyError, "primary"), + ], +) +def test_unknown_schema_preserves_error(data_version, exception, message): + data = b'{"events": [{"type": {}}], "teams": {}}' + with pytest.raises(exception, match=message): + wyscout.load(BytesIO(data), data_version=data_version) + + +@pytest.mark.parametrize("data_version", ["v2", "V4", "", "unexpected"]) +def test_nonstandard_version_uses_automatic_fallback( + event_v2_data, data_version +): + dataset = wyscout.load(event_v2_data, data_version=data_version) + assert len(dataset.records) == 1835 + + +@pytest.mark.parametrize( + ("version", "fixture_name"), + [("V2", "event_v2_data"), ("V3", "event_v3_data")], +) +def test_reusing_consumed_stream_preserves_error( + request, version, fixture_name +): + stream = BytesIO(request.getfixturevalue(fixture_name).read_bytes()) + + wyscout.load(stream, data_version=version) + + with pytest.raises(json.JSONDecodeError): + wyscout.load(stream, data_version=version) + assert not stream.closed + + +@pytest.mark.parametrize( + ("deserializer_class", "fixture_name"), + [ + (WyscoutDeserializerV2, "event_v2_data"), + (WyscoutDeserializerV3, "event_v3_data"), + ], +) +def test_parsed_inputs_keep_base_metadata_merge( + request, deserializer_class, fixture_name +): + parsed_event_data = json.loads( + request.getfixturevalue(fixture_name).read_bytes() + ) + + dataset = deserializer_class().deserialize( + WyscoutInputs(event_data=parsed_event_data), + additional_metadata={"game_id": "override"}, + ) + + assert dataset.metadata.game_id == "override" + + def test_correct_auto_recognize_deserialization( event_v2_data: Path, event_v3_data: Path ):