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.
37 lines
1.8 KiB
37 lines
1.8 KiB
from src.config_loader import VoicePromptSettings
|
|
from src.models import BoundingBox, TrackEvent, TrackedPerson, VoicePromptTrackState
|
|
from src.voice_rules import eligibility_reason, trigger_matches
|
|
|
|
|
|
def cfg(strategy: str = "trigger_line_crossing", **kw: object) -> VoicePromptSettings:
|
|
values = dict(enabled=True, mode="prompt", trigger_strategy=strategy, require_confirmed_track=True,
|
|
require_inside_region=True, skip_if_already_turned=True, one_prompt_per_track=True,
|
|
global_cooldown_sec=3., minimum_track_age_sec=0., response_window_sec=3.,
|
|
fixed_position_axis="x", fixed_position_value=50., fixed_position_direction="increasing")
|
|
values.update(kw)
|
|
return VoicePromptSettings(**values)
|
|
|
|
|
|
def person(x: float = 60, confirmed: bool = True, inside: bool = True) -> TrackedPerson:
|
|
return TrackedPerson(1, BoundingBox(40, 10, 80, 100), .9, 0, "person", 60, 55, x, 100,
|
|
inside, confirmed, 10, 0)
|
|
|
|
|
|
def event(kind: str) -> TrackEvent:
|
|
return TrackEvent(kind, 1, 1, "t", 1., 60, 100)
|
|
|
|
|
|
def test_event_and_position_triggers() -> None:
|
|
assert trigger_matches(cfg(), person(), [event("trigger_crossed")])
|
|
assert trigger_matches(cfg("region_entry"), person(), [event("region_entered")])
|
|
assert trigger_matches(cfg("fixed_x_position"), person(), [])
|
|
|
|
|
|
def test_eligibility_guards() -> None:
|
|
state = VoicePromptTrackState(1, "prompt")
|
|
assert eligibility_reason(cfg(), person(confirmed=False), state, None, 0) == "unconfirmed_track"
|
|
assert eligibility_reason(cfg(), person(inside=False), state, None, 0) == "outside_region"
|
|
assert eligibility_reason(cfg(), person(), state, None, 1) == "global_cooldown"
|
|
state.prompt_attempted = True
|
|
assert eligibility_reason(cfg(), person(), state, None, 0) == "already_prompted"
|
|
|
|
|