mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
a6272b2325
commit
fb9fef7d48
2 changed files with 50 additions and 1 deletions
|
|
@ -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):
|
||||
|
|
|
|||
46
tests/unit/integration_support/test_database_relay.py
Normal file
46
tests/unit/integration_support/test_database_relay.py
Normal 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""
|
||||
Loading…
Add table
Reference in a new issue