mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
fb9fef7d48
commit
1300ba707e
2 changed files with 28 additions and 36 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue