You can not select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
94 lines
4.6 KiB
94 lines
4.6 KiB
from src.config_loader import CrossingSettings, TrackManagementSettings
|
|
from src.models import BoundingBox, FrameInfo, PixelRegion, PixelTriggerLine, TrackedPerson
|
|
from src.track_manager import TrackManager
|
|
|
|
|
|
def management(**overrides: object) -> TrackManagementSettings:
|
|
values = dict(trajectory_max_points=20, direction_window_points=5, minimum_direction_displacement_pixels=5,
|
|
region_entry_confirm_frames=2, region_exit_confirm_frames=3,
|
|
minimum_visible_frames_for_count=3, minimum_duration_sec_for_count=0.1,
|
|
require_region_entry_for_count=True, require_trigger_crossing_for_count=True)
|
|
values.update(overrides)
|
|
return TrackManagementSettings(**values) # type: ignore[arg-type]
|
|
|
|
|
|
def crossing() -> CrossingSettings:
|
|
return CrossingSettings(True, True, 2, 5, 3, True)
|
|
|
|
|
|
def info(frame: int) -> FrameInfo:
|
|
return FrameInfo(frame, f"2026-01-01T00:00:{frame:02d}", frame * 0.1, 100, 100, 20.0)
|
|
|
|
|
|
def person(track_id: int, frame: int, x: int, y: int, confirmed: bool = True) -> TrackedPerson:
|
|
bbox = BoundingBox(max(0, x - 5), max(0, y - 20), min(99, x + 5), min(99, y))
|
|
return TrackedPerson(track_id, bbox, 0.9, 0, "person", x, y - 10, x, y, False, confirmed, frame, 0)
|
|
|
|
|
|
def manager(**overrides: object) -> TrackManager:
|
|
return TrackManager(management(**overrides), PixelRegion(20, 20, 80, 80), PixelTriggerLine("vertical", 50), crossing(), 2)
|
|
|
|
|
|
def test_create_initial_entry_trajectory_and_unique_count() -> None:
|
|
value = manager()
|
|
events = value.update([person(1, 1, 30, 60)], info(1))
|
|
assert [event.event_type for event in events] == ["track_created", "region_entered"]
|
|
value.update([person(1, 2, 40, 60)], info(2))
|
|
assert len(value.get_state(1).trajectory) == 2 # type: ignore[union-attr]
|
|
assert value.total_unique_tracks == value.region_visitors == 1
|
|
|
|
|
|
def test_entry_exit_hysteresis_and_no_duplicates() -> None:
|
|
value = manager()
|
|
value.update([person(1, 1, 10, 60)], info(1))
|
|
assert not value.update([person(1, 2, 30, 60)], info(2))
|
|
assert [e.event_type for e in value.update([person(1, 3, 35, 60)], info(3))] == ["region_entered"]
|
|
assert "region_exited" not in [e.event_type for e in value.update([person(1, 4, 90, 60)], info(4))]
|
|
assert "region_exited" not in [e.event_type for e in value.update([person(1, 5, 90, 60)], info(5))]
|
|
assert "region_exited" in [e.event_type for e in value.update([person(1, 6, 90, 60)], info(6))]
|
|
assert "region_exited" not in [e.event_type for e in value.update([person(1, 7, 90, 60)], info(7))]
|
|
|
|
|
|
def test_crossing_direction_finalization_and_valid_count() -> None:
|
|
value = manager()
|
|
all_events = []
|
|
for frame, x in enumerate((30, 45, 60), 1):
|
|
all_events.extend(value.update([person(1, frame, x, 60)], info(frame)))
|
|
assert sum(event.event_type == "trigger_crossed" for event in all_events) == 1
|
|
state = value.get_state(1)
|
|
assert state and state.movement_direction == "left_to_right"
|
|
events = value.finalize_all(info(4))
|
|
assert events[0].event_type == "track_finalized"
|
|
assert value.get_finalized_states()[0].valid_for_counting
|
|
assert value.completed_passers == 1
|
|
|
|
|
|
def test_stale_track_and_short_track_invalid() -> None:
|
|
value = manager(minimum_visible_frames_for_count=5)
|
|
value.update([person(1, 1, 30, 60)], info(1))
|
|
value.update([], info(2))
|
|
value.update([], info(3))
|
|
value.update([], info(4))
|
|
events = value.finalize_stale_tracks(info(4))
|
|
assert [event.event_type for event in events] == ["track_lost", "track_finalized"]
|
|
assert value.get_state(1).invalid_reason == "insufficient_visible_frames" # type: ignore[union-attr]
|
|
|
|
|
|
def test_crossing_direction_wins_when_track_stops_before_finalization() -> None:
|
|
"""Reproduce a right-to-left crossing followed by stationary tail points."""
|
|
value = manager(direction_window_points=3, minimum_direction_displacement_pixels=5)
|
|
all_events = []
|
|
positions = (70, 55, 40, 40, 40, 40)
|
|
for frame, x in enumerate(positions, 1):
|
|
all_events.extend(value.update([person(1, frame, x, 60)], info(frame)))
|
|
|
|
crossing_event = next(event for event in all_events if event.event_type == "trigger_crossed")
|
|
assert crossing_event.direction == "right_to_left"
|
|
state = value.get_state(1)
|
|
assert state and state.trigger_crossing_direction == "right_to_left"
|
|
assert state.movement_direction == "right_to_left"
|
|
|
|
finalized_event = value.finalize_all(info(7))[0]
|
|
finalized_state = value.get_finalized_states()[0]
|
|
assert finalized_event.direction == "right_to_left"
|
|
assert finalized_state.movement_direction == "right_to_left"
|
|
|