fix(proxy): count pgbouncer ready only once its own unix socket answers, not any listener on the port

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yassin 2026-09-04 02:23:18 +00:00
parent 4630956814
commit 2212c39c2a
2 changed files with 152 additions and 25 deletions

View file

@ -258,6 +258,20 @@ def _port_open(port: int) -> bool:
return False
def _unix_socket_open(path: Path) -> bool:
with socket.socket(socket.AF_UNIX) as probe:
probe.settimeout(0.5)
try:
probe.connect(str(path))
except OSError:
return False
return True
def unix_socket_path(runtime_dir: Path, port: int) -> Path:
return runtime_dir / f".s.PGSQL.{port}"
def _end(process: subprocess.Popen[bytes]) -> None:
if process.poll() is not None:
return
@ -278,17 +292,24 @@ class PgBouncerProcess:
listening again. A replacement that cannot be spawned, finds its port
taken, exits again or never starts listening is retried every
``restart_delay_seconds`` until ``stop`` is called.
A connect probe of ``port`` cannot tell the child from another process
that grabbed the port after the availability check, so readiness also
needs ``socket_path``: the unix socket PgBouncer creates in the private
runtime directory, which it only does once every TCP listener is bound.
"""
def __init__(
self,
argv: Sequence[str],
port: int,
socket_path: Path,
restart_delay_seconds: float = PGBOUNCER_RESTART_DELAY_SECONDS,
ready_timeout_seconds: float = PGBOUNCER_READY_TIMEOUT_SECONDS,
) -> None:
self.argv: Final = tuple(argv)
self.port: Final = port
self.socket_path: Final = socket_path
self.restart_delay_seconds: Final = restart_delay_seconds
self.ready_timeout_seconds: Final = ready_timeout_seconds
self._stopping: Final = threading.Event()
@ -323,16 +344,20 @@ class PgBouncerProcess:
while time.monotonic() < deadline:
if process.poll() is not None:
return PgBouncerError(f"pgbouncer exited with status {process.returncode} during startup")
if _port_open(self.port):
if _port_open(self.port) and _unix_socket_open(self.socket_path):
return None
time.sleep(0.1)
if _port_open(self.port):
return PgBouncerError(
f"{PGBOUNCER_LISTEN_ADDR}:{self.port} is served by another process, not the pgbouncer that was started"
)
return PgBouncerError(
f"pgbouncer did not start listening on {PGBOUNCER_LISTEN_ADDR}:{self.port} "
f"within {self.ready_timeout_seconds:.0f}s"
)
def start(self) -> PgBouncerError | None:
"""Spawn PgBouncer, wait until it accepts connections, then supervise it from a daemon thread."""
"""Spawn PgBouncer, wait until it listens on port and unix socket, then supervise it from a daemon thread."""
process: Final = self._spawn()
if process is None:
return PgBouncerError("pgbouncer was stopped before it started")
@ -411,7 +436,11 @@ def start_in_container_pgbouncer(
if isinstance(plan, PgBouncerError):
return plan
ini_path: Final = write_pgbouncer_files(plan, runtime_dir, run_as_user)
pooler: Final = PgBouncerProcess(argv=(settings.binary, str(ini_path)), port=settings.port)
pooler: Final = PgBouncerProcess(
argv=(settings.binary, str(ini_path)),
port=settings.port,
socket_path=unix_socket_path(runtime_dir, settings.port),
)
failed: Final = pooler.start()
if failed is not None:
return failed

View file

@ -9,8 +9,9 @@ import textwrap
import time
import urllib.parse
from collections.abc import Callable
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
from typing import Final
from typing import Final, cast
import pytest
@ -22,6 +23,7 @@ from litellm.proxy.db.pgbouncer import (
PgBouncerSettings,
plan_pgbouncer,
start_in_container_pgbouncer,
unix_socket_path,
write_pgbouncer_files,
)
@ -152,33 +154,58 @@ class TestWritePgBouncerFiles:
assert stat.S_IMODE(path.stat().st_mode) == 0o600
def _bound_port(sock: socket.socket) -> int:
return cast(tuple[str, int], sock.getsockname())[1]
def _free_port() -> int:
with socket.socket() as probe:
probe.bind(("127.0.0.1", 0))
return probe.getsockname()[1]
return _bound_port(probe)
def _fake_pooler(tmp_path: Path, port: int, exit_immediately: bool = False, port_file: Path | None = None) -> Path:
"""An executable that listens on ``port`` like PgBouncer would (or exits at once), ignoring its ini argument.
def _fake_pooler(
tmp_path: Path,
port: int,
exit_immediately: bool = False,
port_file: Path | None = None,
bind_delay_seconds: float = 0.0,
) -> Path:
"""An executable that listens like PgBouncer: on the TCP port first, then on ``.s.PGSQL.<port>`` in the socket dir.
With ``port_file`` each start reads the port to listen on from that file instead.
Port and socket dir come from the ini it is given, else from ``port`` and
``tmp_path``. With ``port_file`` each start reads the port from that file
instead. ``bind_delay_seconds`` holds the bind back, like a slow start.
"""
script: Final = tmp_path / "fake-pgbouncer"
script.write_text(
textwrap.dedent(
f"""\
#!{sys.executable}
import pathlib, socket, sys, time
import configparser, os, pathlib, select, socket, sys, time
if {exit_immediately!r}:
sys.exit(3)
port = {port} if {port_file is None!r} else int(pathlib.Path({str(port_file)!r}).read_text())
ini = configparser.ConfigParser()
ini.read(sys.argv[1:2])
port = ini.getint("pgbouncer", "listen_port", fallback={port})
if not {port_file is None!r}:
port = int(pathlib.Path({str(port_file)!r}).read_text())
socket_dir = ini.get("pgbouncer", "unix_socket_dir", fallback={str(tmp_path)!r})
socket_path = f"{{socket_dir}}/.s.PGSQL.{{port}}"
time.sleep({bind_delay_seconds!r})
listener = socket.socket()
listener.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
listener.bind(("127.0.0.1", port))
listener.listen()
if os.path.exists(socket_path):
os.unlink(socket_path)
unix_listener = socket.socket(socket.AF_UNIX)
unix_listener.bind(socket_path)
unix_listener.listen()
while True:
conn, _ = listener.accept()
conn.close()
for ready in select.select([listener, unix_listener], [], [])[0]:
conn, _ = ready.accept()
conn.close()
"""
)
)
@ -206,7 +233,9 @@ def _wait_until(condition: Callable[[], bool], timeout_seconds: float = 5.0) ->
class TestPgBouncerProcess:
def test_start_waits_for_the_listener_and_stop_ends_it(self, tmp_path: Path):
port: Final = _free_port()
pooler: Final = PgBouncerProcess(argv=(str(_fake_pooler(tmp_path, port)),), port=port)
pooler: Final = PgBouncerProcess(
argv=(str(_fake_pooler(tmp_path, port)),), port=port, socket_path=unix_socket_path(tmp_path, port)
)
assert pooler.start() is None
assert _listening(port)
pid: Final = pooler.pid
@ -219,7 +248,10 @@ class TestPgBouncerProcess:
def test_a_crashed_pooler_is_restarted_with_a_new_pid(self, tmp_path: Path):
port: Final = _free_port()
pooler: Final = PgBouncerProcess(
argv=(str(_fake_pooler(tmp_path, port)),), port=port, restart_delay_seconds=0.1
argv=(str(_fake_pooler(tmp_path, port)),),
port=port,
socket_path=unix_socket_path(tmp_path, port),
restart_delay_seconds=0.1,
)
assert pooler.start() is None
first_pid: Final = pooler.pid
@ -234,7 +266,9 @@ class TestPgBouncerProcess:
):
port: Final = _free_port()
script: Final = _fake_pooler(tmp_path, port)
pooler: Final = PgBouncerProcess(argv=(str(script),), port=port, restart_delay_seconds=0.1)
pooler: Final = PgBouncerProcess(
argv=(str(script),), port=port, socket_path=unix_socket_path(tmp_path, port), restart_delay_seconds=0.1
)
assert pooler.start() is None
first_pid: Final = pooler.pid
assert first_pid is not None
@ -254,7 +288,11 @@ class TestPgBouncerProcess:
port_file.write_text(str(port))
script: Final = _fake_pooler(tmp_path, port, port_file=port_file)
pooler: Final = PgBouncerProcess(
argv=(str(script),), port=port, restart_delay_seconds=0.1, ready_timeout_seconds=0.3
argv=(str(script),),
port=port,
socket_path=unix_socket_path(tmp_path, port),
restart_delay_seconds=0.1,
ready_timeout_seconds=0.3,
)
assert pooler.start() is None
first_pid: Final = pooler.pid
@ -274,7 +312,10 @@ class TestPgBouncerProcess:
def test_stopping_during_the_restart_delay_leaves_no_pooler_behind(self, tmp_path: Path):
port: Final = _free_port()
pooler: Final = PgBouncerProcess(
argv=(str(_fake_pooler(tmp_path, port)),), port=port, restart_delay_seconds=0.3
argv=(str(_fake_pooler(tmp_path, port)),),
port=port,
socket_path=unix_socket_path(tmp_path, port),
restart_delay_seconds=0.3,
)
assert pooler.start() is None
first_pid: Final = pooler.pid
@ -289,7 +330,10 @@ class TestPgBouncerProcess:
def test_a_stopped_pooler_is_not_restarted(self, tmp_path: Path, caplog: pytest.LogCaptureFixture):
port: Final = _free_port()
pooler: Final = PgBouncerProcess(
argv=(str(_fake_pooler(tmp_path, port)),), port=port, restart_delay_seconds=0.1
argv=(str(_fake_pooler(tmp_path, port)),),
port=port,
socket_path=unix_socket_path(tmp_path, port),
restart_delay_seconds=0.1,
)
assert pooler.start() is None
with caplog.at_level(logging.ERROR, logger=verbose_proxy_logger.name):
@ -300,13 +344,19 @@ class TestPgBouncerProcess:
def test_a_pooler_that_exits_during_startup_is_reported(self, tmp_path: Path):
port: Final = _free_port()
pooler: Final = PgBouncerProcess(argv=(str(_fake_pooler(tmp_path, port, exit_immediately=True)),), port=port)
pooler: Final = PgBouncerProcess(
argv=(str(_fake_pooler(tmp_path, port, exit_immediately=True)),),
port=port,
socket_path=unix_socket_path(tmp_path, port),
)
outcome: Final = pooler.start()
assert isinstance(outcome, PgBouncerError)
assert "status 3" in outcome.reason
def test_a_missing_binary_is_reported(self):
outcome: Final = PgBouncerProcess(argv=("/nonexistent/pgbouncer",), port=_free_port()).start()
def test_a_missing_binary_is_reported(self, tmp_path: Path):
outcome: Final = PgBouncerProcess(
argv=("/nonexistent/pgbouncer",), port=_free_port(), socket_path=tmp_path / "sock"
).start()
assert isinstance(outcome, PgBouncerError)
assert "/nonexistent/pgbouncer" in outcome.reason
@ -314,8 +364,10 @@ class TestPgBouncerProcess:
with socket.socket() as squatter:
squatter.bind(("127.0.0.1", 0))
squatter.listen()
port: Final = squatter.getsockname()[1]
pooler: Final = PgBouncerProcess(argv=(str(_fake_pooler(tmp_path, port)),), port=port)
port: Final = _bound_port(squatter)
pooler: Final = PgBouncerProcess(
argv=(str(_fake_pooler(tmp_path, port)),), port=port, socket_path=unix_socket_path(tmp_path, port)
)
outcome: Final = pooler.start()
assert isinstance(outcome, PgBouncerError)
assert f"127.0.0.1:{port} is already in use" in outcome.reason
@ -326,7 +378,10 @@ class TestPgBouncerProcess:
):
port: Final = _free_port()
pooler: Final = PgBouncerProcess(
argv=(str(_fake_pooler(tmp_path, port)),), port=port, restart_delay_seconds=0.5
argv=(str(_fake_pooler(tmp_path, port)),),
port=port,
socket_path=unix_socket_path(tmp_path, port),
restart_delay_seconds=0.5,
)
assert pooler.start() is None
first_pid: Final = pooler.pid
@ -343,10 +398,53 @@ class TestPgBouncerProcess:
pooler.stop()
assert _wait_until(lambda: not _listening(port))
def test_a_listener_that_grabs_the_port_after_the_spawn_is_not_taken_for_the_pooler(self, tmp_path: Path):
port: Final = _free_port()
pooler: Final = PgBouncerProcess(
argv=(str(_fake_pooler(tmp_path, port, bind_delay_seconds=0.5)),),
port=port,
socket_path=unix_socket_path(tmp_path, port),
ready_timeout_seconds=3.0,
)
with socket.socket() as squatter, ThreadPoolExecutor(max_workers=1) as starter:
squatter.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
starting: Final = starter.submit(pooler.start)
assert _wait_until(lambda: pooler.pid is not None)
squatter.bind(("127.0.0.1", port))
squatter.listen()
outcome: Final = starting.result()
assert isinstance(outcome, PgBouncerError)
assert "exited with status 1" in outcome.reason
def test_a_port_served_by_a_stranger_while_the_pooler_is_still_starting_is_reported(self, tmp_path: Path):
port: Final = _free_port()
pooler: Final = PgBouncerProcess(
argv=(str(_fake_pooler(tmp_path, port, bind_delay_seconds=30.0)),),
port=port,
socket_path=unix_socket_path(tmp_path, port),
ready_timeout_seconds=0.5,
)
with socket.socket() as squatter, ThreadPoolExecutor(max_workers=1) as starter:
squatter.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
starting: Final = starter.submit(pooler.start)
assert _wait_until(lambda: pooler.pid is not None)
squatter.bind(("127.0.0.1", port))
squatter.listen()
outcome: Final = starting.result()
assert isinstance(outcome, PgBouncerError)
assert f"127.0.0.1:{port} is served by another process" in outcome.reason
pid: Final = pooler.pid
assert pid is not None
with pytest.raises(ProcessLookupError):
os.kill(pid, 0)
def test_a_pooler_that_never_listens_times_out(self, tmp_path: Path):
port: Final = _free_port()
pooler: Final = PgBouncerProcess(
argv=(str(_fake_pooler(tmp_path, _free_port())),), port=port, ready_timeout_seconds=0.5
argv=(str(_fake_pooler(tmp_path, _free_port())),),
port=port,
socket_path=unix_socket_path(tmp_path, port),
ready_timeout_seconds=0.5,
)
outcome: Final = pooler.start()
assert isinstance(outcome, PgBouncerError)