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
19 changes: 9 additions & 10 deletions kloppy/_providers/wyscout.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,23 +32,25 @@ 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,
coordinate_system=coordinates,
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(
Expand Down Expand Up @@ -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
Expand Down
7 changes: 3 additions & 4 deletions kloppy/infra/serializers/event/wyscout/deserializer_v2.py
Original file line number Diff line number Diff line change
@@ -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,
Expand Down Expand Up @@ -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]):
Expand All @@ -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"]
Expand Down
3 changes: 1 addition & 2 deletions kloppy/infra/serializers/event/wyscout/deserializer_v3.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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"]
Expand Down
224 changes: 224 additions & 0 deletions kloppy/tests/test_wyscout.py
Original file line number Diff line number Diff line change
@@ -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,
Expand All @@ -13,6 +17,7 @@
DuelQualifier,
DuelType,
EventDataset,
EventFactory,
EventType,
FormationType,
GoalkeeperActionType,
Expand All @@ -29,6 +34,11 @@
ShotResult,
Time,
)
from kloppy.infra.serializers.event.wyscout import (
WyscoutDeserializerV2,
WyscoutDeserializerV3,
WyscoutInputs,
)


@pytest.fixture(scope="session")
Expand All @@ -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
):
Expand Down