test(integration): match a dropped-connection trigger split across two reads

The dropped-connection relay checked each TCP read on its own, so a SQL
marker that straddled two reads never tripped it and the outage cells
would run without the outage they meant to exercise. Carry the tail of
the previous read into the next check, as the held-statement relay
already does, and pin that with a unit test that splits the trigger
across two writes.
This commit is contained in:
mateo-berri 2026-10-05 14:24:01 -07:00
parent a6272b2325
commit fb9fef7d48
2 changed files with 50 additions and 1 deletions

View file

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

View file

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