diff --git a/litellm/proxy/db/pgbouncer.py b/litellm/proxy/db/pgbouncer.py index 313ab5967b3..be99c0eb60b 100644 --- a/litellm/proxy/db/pgbouncer.py +++ b/litellm/proxy/db/pgbouncer.py @@ -275,9 +275,9 @@ class PgBouncerProcess: Prisma reconnects by itself after a failed query, so a PgBouncer crash costs the requests in flight plus one failed query per idle pooled connection the crash severed, and nothing else once the replacement is - listening again. A replacement that cannot be spawned, exits again or - never starts listening is retried every ``restart_delay_seconds`` until - ``stop`` is called. + 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. """ def __init__( @@ -300,12 +300,21 @@ class PgBouncerProcess: with self._lock: return None if self._process is None else self._process.pid - def _spawn(self) -> subprocess.Popen[bytes] | None: - """Start a child, or None once ``stop`` ran; both take the lock so no child can slip in after a stop.""" + def _spawn(self) -> subprocess.Popen[bytes] | PgBouncerError | None: + """Start a child, or None once ``stop`` ran; both take the lock so no child can slip in after a stop. + + The port has to be free first: a listener that is already there would + pass the readiness check while the child fails to bind. + """ with self._lock: if self._stopping.is_set(): return None - process: Final = subprocess.Popen(self.argv) + if _port_open(self.port): + return PgBouncerError(f"{PGBOUNCER_LISTEN_ADDR}:{self.port} is already in use by another process") + try: + process: Final = subprocess.Popen(self.argv) + except OSError as spawn_error: + return PgBouncerError(f"could not start {self.argv[0]!r}: {spawn_error}") self._process = process return process @@ -324,12 +333,11 @@ class PgBouncerProcess: def start(self) -> PgBouncerError | None: """Spawn PgBouncer, wait until it accepts connections, then supervise it from a daemon thread.""" - try: - process: Final = self._spawn() - except OSError as spawn_error: - return PgBouncerError(f"could not start {self.argv[0]!r}: {spawn_error}") + process: Final = self._spawn() if process is None: return PgBouncerError("pgbouncer was stopped before it started") + if isinstance(process, PgBouncerError): + return process not_ready: Final = self._wait_ready(process) if not_ready is not None: self.stop() @@ -356,13 +364,12 @@ class PgBouncerProcess: def _restart_after_delay(self) -> None: time.sleep(self.restart_delay_seconds) - try: - process: Final = self._spawn() - except OSError as spawn_error: - self._retry_restart(str(spawn_error)) - return + process: Final = self._spawn() if process is None: return + if isinstance(process, PgBouncerError): + self._retry_restart(process.reason) + return not_ready: Final = self._wait_ready(process) if not_ready is None: self._watch(process) diff --git a/tests/test_litellm/proxy/db/test_pgbouncer.py b/tests/test_litellm/proxy/db/test_pgbouncer.py index 97d19c93634..21038ec6089 100644 --- a/tests/test_litellm/proxy/db/test_pgbouncer.py +++ b/tests/test_litellm/proxy/db/test_pgbouncer.py @@ -310,6 +310,39 @@ class TestPgBouncerProcess: assert isinstance(outcome, PgBouncerError) assert "/nonexistent/pgbouncer" in outcome.reason + def test_a_port_owned_by_someone_else_is_refused_before_spawning(self, tmp_path: Path): + 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) + outcome: Final = pooler.start() + assert isinstance(outcome, PgBouncerError) + assert f"127.0.0.1:{port} is already in use" in outcome.reason + assert pooler.pid is None + + def test_a_replacement_waits_until_a_squatter_leaves_the_port( + 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.5 + ) + assert pooler.start() is None + first_pid: Final = pooler.pid + assert first_pid is not None + os.kill(first_pid, signal.SIGKILL) + assert _wait_until(lambda: not _listening(port)) + with socket.socket() as squatter, caplog.at_level(logging.ERROR, logger=verbose_proxy_logger.name): + squatter.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) + squatter.bind(("127.0.0.1", port)) + squatter.listen() + assert _wait_until(lambda: any("already in use" in record.message for record in caplog.records)) + assert pooler.pid == first_pid + assert _wait_until(lambda: pooler.pid not in (None, first_pid) and _listening(port)) + pooler.stop() + assert _wait_until(lambda: not _listening(port)) + def test_a_pooler_that_never_listens_times_out(self, tmp_path: Path): port: Final = _free_port() pooler: Final = PgBouncerProcess(