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
This commit is contained in:
mateo-berri 2026-10-05 14:35:07 -07:00
parent fb9fef7d48
commit 1300ba707e
2 changed files with 28 additions and 36 deletions

View file

@ -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):

View file

@ -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")