diff --git a/src/debugpy/adapter/__main__.py b/src/debugpy/adapter/__main__.py index 916a42a1..52409a18 100644 --- a/src/debugpy/adapter/__main__.py +++ b/src/debugpy/adapter/__main__.py @@ -82,11 +82,9 @@ def main(): ) try: - ipv6 = localhost.count(":") > 1 - sock = sockets.create_client(ipv6) + sock = sockets.connect((localhost, args.for_server), 5, 0.1) try: sock.settimeout(None) - sock.connect((localhost, args.for_server)) sock_io = sock.makefile("wb", 0) try: sock_io.write(json.dumps(endpoints).encode("utf-8")) diff --git a/src/debugpy/common/sockets.py b/src/debugpy/common/sockets.py index aecb6d83..9e0ba9b3 100644 --- a/src/debugpy/common/sockets.py +++ b/src/debugpy/common/sockets.py @@ -5,6 +5,7 @@ import socket import sys import threading +import time from typing import Any, Callable, Union from debugpy.common import log @@ -99,6 +100,28 @@ def create_client(ipv6=False): return _new_sock(ipv6) +def connect( + address: tuple[str, int], attempts: int = 1, retry_interval: float = 0 +) -> socket.socket: + """Return a client socket connected to the given address.""" + assert attempts > 0 + ipv6 = address[0].count(":") > 1 + while True: + sock = create_client(ipv6) + try: + sock.connect(address) + return sock + except ConnectionRefusedError: + close_socket(sock) + attempts -= 1 + if attempts == 0: + raise + time.sleep(retry_interval) + except Exception: + close_socket(sock) + raise + + def _new_sock(ipv6=False): address_family = socket.AF_INET6 if ipv6 else socket.AF_INET sock = socket.socket(address_family, socket.SOCK_STREAM, socket.IPPROTO_TCP) diff --git a/tests/debugpy/common/test_socket.py b/tests/debugpy/common/test_socket.py index b59db21d..ad2b00ee 100644 --- a/tests/debugpy/common/test_socket.py +++ b/tests/debugpy/common/test_socket.py @@ -8,6 +8,28 @@ from debugpy.common import sockets +def test_connect_retries_connection_refused(monkeypatch): + class Client: + def __init__(self, refused=False): + self.refused = refused + self.closed = False + + def connect(self, address): + if self.refused: + raise ConnectionRefusedError + + def close(self): + self.closed = True + + refused_client = Client(refused=True) + connected_client = Client() + clients = [refused_client, connected_client] + monkeypatch.setattr(sockets, "create_client", lambda ipv6: clients.pop(0)) + connected = sockets.connect(("127.0.0.1", 5678), attempts=2) + assert connected is connected_client + assert refused_client.closed is True + + class TestSocketServerReuse(object): HOST1 = "127.0.0.1" # NOTE: Windows allows loopback range 127/8. Some flavors of Linux support