diff --git a/pooltool/evolution/event_based/detect/ball_ball.py b/pooltool/evolution/event_based/detect/ball_ball.py index df782764..d7650359 100644 --- a/pooltool/evolution/event_based/detect/ball_ball.py +++ b/pooltool/evolution/event_based/detect/ball_ball.py @@ -54,6 +54,9 @@ def ball_ball_collision_time( if C[4] == 0.0: # C[3] must also be 0.0, and this is a quadratic assert C[3] == 0.0 + # Exact contact only collides at t=0 if the balls are closing. + if C[0] == 0.0 and C[1] >= 0.0: + return np.inf return get_real_positive_smallest_root(quadratic.solve(C[2], C[1], C[0])) return get_real_positive_smallest_root(quartic.solve(C[4], C[3], C[2], C[1], C[0])) @@ -129,6 +132,15 @@ def ball_ball_collision_time_2d( d = 2 * Bx * Cx + 2 * By * Cy e = Cx * Cx + Cy * Cy - 4 * R * R + if a == 0.0: + # If relative acceleration is zero, the cubic term must also be zero and the + # distance equation reduces to a quadratic. + assert b == 0.0 + # Exact contact only collides at t=0 if the balls are closing. + if e == 0.0 and d >= 0.0: + return np.inf + return get_real_positive_smallest_root(quadratic.solve(c, d, e)) + return get_real_positive_smallest_root(quartic.solve(a, b, c, d, e)) diff --git a/tests/evolution/event_based/test_ball_ball.py b/tests/evolution/event_based/test_ball_ball.py index f3b0986a..4c7ff7ac 100644 --- a/tests/evolution/event_based/test_ball_ball.py +++ b/tests/evolution/event_based/test_ball_ball.py @@ -7,6 +7,7 @@ from pooltool.events import EventType from pooltool.evolution.event_based.cache import CollisionCache from pooltool.evolution.event_based.detect.ball_ball import ( + ball_ball_collision_time_2d, get_next_ball_ball_event, ) from pooltool.physics.evolve import evolve_ball_motion @@ -14,6 +15,16 @@ from pooltool.system.datatypes import Ball, Cue, System, Table +def _make_rolling_ball(ball_id: str, xy: tuple[float, float], velocity: float) -> Ball: + ball = Ball.create(ball_id, xy=xy) + v = np.array([0.0, velocity, 0.0]) + w = ptmath.cross(np.array([0.0, 0.0, 1.0]), v) / ball.params.R + ball.state.rvw[1] = v + ball.state.rvw[2] = w + ball.state.s = const.rolling + return ball + + @pytest.mark.parametrize("is_3d", [True, False]) def test_sliding_ball_collision_time(is_3d: bool): table = Table.default() @@ -53,6 +64,134 @@ def test_sliding_ball_collision_time(is_3d: bool): assert np.isclose(actual, expected), f"actual={actual}, expected={expected}" +@pytest.mark.parametrize("is_3d", [True, False]) +def test_parallel_rolling_balls_do_not_collide(is_3d: bool): + """Parallel rolling balls at fixed separation never collide.""" + + table = Table.default() + cue = Cue.default() + + cue_ball = _make_rolling_ball("cue", (table.w / 2, table.l / 4), velocity=1.0) + one_ball = _make_rolling_ball("1", (table.w / 2, 3 * table.l / 4), velocity=1.0) + + system = System( + cue=cue, + table=table, + balls={ + "cue": cue_ball, + "1": one_ball, + }, + ) + + event = get_next_ball_ball_event(system, CollisionCache(), is_3d=is_3d) + assert event.time == np.inf + + if not is_3d: + actual = ball_ball_collision_time_2d( + rvw1=cue_ball.state.rvw, + rvw2=one_ball.state.rvw, + s1=cue_ball.state.s, + s2=one_ball.state.s, + mu1=cue_ball.params.u_r, + mu2=one_ball.params.u_r, + m1=cue_ball.params.m, + m2=one_ball.params.m, + g1=cue_ball.params.g, + g2=one_ball.params.g, + R=cue_ball.params.R, + ) + assert actual == np.inf + + +def test_parallel_rolling_balls_collide_from_quadratic_root_2d(): + """The 2D detector handles finite roots when the quartic term is zero.""" + + table = Table.default() + cue = Cue.default() + + cue_ball = _make_rolling_ball("cue", (table.w / 2, table.l / 4), velocity=2.0) + R = cue_ball.params.R + center_gap = 4 * R + one_ball = _make_rolling_ball( + "1", + (table.w / 2, table.l / 4 + center_gap), + velocity=1.0, + ) + + system = System( + cue=cue, + table=table, + balls={ + "cue": cue_ball, + "1": one_ball, + }, + ) + + expected = (center_gap - 2 * R) / ( + cue_ball.state.rvw[1, 1] - one_ball.state.rvw[1, 1] + ) + + event = get_next_ball_ball_event(system, CollisionCache(), is_3d=False) + assert np.isclose(event.time, expected) + + +@pytest.mark.parametrize( + ("cue_velocity", "one_velocity", "expected"), + [ + (2.0, 1.0, 0.0), + (1.0, 1.0, np.inf), + (1.0, 2.0, np.inf), + ], +) +@pytest.mark.parametrize("is_3d", [True, False]) +def test_tangent_parallel_rolling_balls_only_collide_when_closing( + cue_velocity: float, + one_velocity: float, + expected: float, + is_3d: bool, +): + table = Table.default() + cue = Cue.default() + + cue_ball = _make_rolling_ball("cue", (0.0, 0.0), velocity=cue_velocity) + one_ball = _make_rolling_ball( + "1", + (0.0, 2 * cue_ball.params.R), + velocity=one_velocity, + ) + + system = System( + cue=cue, + table=table, + balls={ + "cue": cue_ball, + "1": one_ball, + }, + ) + + direct = ball_ball_collision_time_2d( + rvw1=cue_ball.state.rvw, + rvw2=one_ball.state.rvw, + s1=cue_ball.state.s, + s2=one_ball.state.s, + mu1=cue_ball.params.u_r, + mu2=one_ball.params.u_r, + m1=cue_ball.params.m, + m2=one_ball.params.m, + g1=cue_ball.params.g, + g2=one_ball.params.g, + R=cue_ball.params.R, + ) + event = get_next_ball_ball_event(system, CollisionCache(), is_3d=is_3d) + + if expected == np.inf: + assert direct == np.inf + assert event.time == np.inf + else: + assert direct == expected + assert event.time == expected + + def test_airborne_balls_colliding(): """Tests two airborne balls colliding.