From 1300ba707e099b7bd180c9d7de73f87784ccf99e Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 5 Oct 2026 14:35:07 -0700 Subject: [PATCH] test(integration): scan relay triggers through an in-process helper The dropped-connection relay now matches its SQL trigger through a TriggerScanner that carries the previous read's tail, and the unit test exercises that scanner directly instead of opening loopback sockets, which tests/unit forbids. The relay's end to end behavior stays covered by the integration cells --- tests/integration/_support/database_relay.py | 18 ++++++-- .../test_database_relay.py | 46 ++++++------------- 2 files changed, 28 insertions(+), 36 deletions(-) diff --git a/tests/integration/_support/database_relay.py b/tests/integration/_support/database_relay.py index 65cf213340c..20f23d0e010 100644 --- a/tests/integration/_support/database_relay.py +++ b/tests/integration/_support/database_relay.py @@ -146,6 +146,17 @@ class HeldStatementRelay: ) +class TriggerScanner: + def __init__(self, trigger: bytes) -> None: + self._trigger: Final = trigger + self._tail: bytes = b"" + + def feed(self, chunk: bytes) -> bool: + window: Final = self._tail + chunk + self._tail = window[-(len(self._trigger) - 1) :] + return self._trigger in window + + class DroppedConnectionRelay: def __init__(self, upstream_host: str, upstream_port: int, trigger: bytes) -> None: self.port: Final = _free_port() @@ -182,15 +193,14 @@ class DroppedConnectionRelay: server_reader, server_writer = await asyncio.open_connection(self._upstream_host, self._upstream_port) async def forward(reader: asyncio.StreamReader, writer: asyncio.StreamWriter, inspect: bool) -> None: - tail = b"" # rebind-ok: carries the previous read's end so a trigger split across reads still matches + scanner: Final = TriggerScanner(self._trigger) try: while chunk := await reader.read(65536): - window: Final = tail + chunk - if inspect and self._armed.is_set() and self._trigger in window: + matched: Final = scanner.feed(chunk) + if inspect and self._armed.is_set() and matched: self.dropped.set() client_writer.close() return - tail = window[-(len(self._trigger) - 1) :] writer.write(chunk) await writer.drain() except (ConnectionError, asyncio.IncompleteReadError): diff --git a/tests/unit/integration_support/test_database_relay.py b/tests/unit/integration_support/test_database_relay.py index a86e0fc6ca8..98501f84058 100644 --- a/tests/unit/integration_support/test_database_relay.py +++ b/tests/unit/integration_support/test_database_relay.py @@ -1,46 +1,28 @@ from __future__ import annotations -import socket -import socketserver -import threading -from collections.abc import Iterator from typing import Final -from urllib.parse import urlsplit import pytest -from tests.integration._support.database_relay import dropped_connection_relay +from tests.integration._support.database_relay import TriggerScanner TRIGGER: Final = b'SELECT "startTime" FROM "LiteLLM_SpendLogs"' -SPLIT_AT: Final = len(TRIGGER) // 2 -class _EchoHandler(socketserver.BaseRequestHandler): - def handle(self) -> None: - while chunk := self.request.recv(65536): - self.request.sendall(chunk) +@pytest.mark.parametrize("split_at", range(1, len(TRIGGER))) +def test_trigger_scanner_matches_a_trigger_split_across_two_reads(split_at: int) -> None: + scanner: Final = TriggerScanner(TRIGGER) + assert not scanner.feed(TRIGGER[:split_at]) + assert scanner.feed(TRIGGER[split_at:]) -@pytest.fixture -def echo_upstream_url() -> Iterator[str]: - server: Final = socketserver.ThreadingTCPServer(("127.0.0.1", 0), _EchoHandler) - server.daemon_threads = True - threading.Thread(target=server.serve_forever, daemon=True).start() - try: - yield f"postgresql://relay:relay@127.0.0.1:{server.server_address[1]}/relay" - finally: - server.shutdown() - server.server_close() +def test_trigger_scanner_matches_a_trigger_arriving_one_byte_at_a_time() -> None: + scanner: Final = TriggerScanner(TRIGGER) + hits: Final = tuple(scanner.feed(TRIGGER[i : i + 1]) for i in range(len(TRIGGER))) + assert hits == (False,) * (len(TRIGGER) - 1) + (True,) -def test_dropped_connection_relay_trips_on_a_trigger_split_across_two_reads(echo_upstream_url: str) -> None: - with dropped_connection_relay(echo_upstream_url, TRIGGER) as (relay, relayed_url): - relay.arm() - port: Final = urlsplit(relayed_url).port - assert port is not None - with socket.create_connection(("127.0.0.1", port), timeout=5) as client, client.makefile("rb") as echoed: - client.sendall(TRIGGER[:SPLIT_AT]) - assert echoed.read(SPLIT_AT) == TRIGGER[:SPLIT_AT] - client.sendall(TRIGGER[SPLIT_AT:]) - assert relay.dropped.wait(5), "the relay never saw the trigger that arrived in two reads" - assert echoed.read() == b"" +def test_trigger_scanner_reports_a_match_once() -> None: + scanner: Final = TriggerScanner(TRIGGER) + assert scanner.feed(b"x" + TRIGGER + b"y") + assert not scanner.feed(b"z")