diff --git a/tests/integration/_support/database_relay.py b/tests/integration/_support/database_relay.py index 986849c3403..65cf213340c 100644 --- a/tests/integration/_support/database_relay.py +++ b/tests/integration/_support/database_relay.py @@ -182,12 +182,15 @@ 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 try: while chunk := await reader.read(65536): - if inspect and self._armed.is_set() and self._trigger in chunk: + window: Final = tail + chunk + if inspect and self._armed.is_set() and self._trigger in window: 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 new file mode 100644 index 00000000000..a86e0fc6ca8 --- /dev/null +++ b/tests/unit/integration_support/test_database_relay.py @@ -0,0 +1,46 @@ +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 + +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.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_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""