Skip to content

Commit 42c8f6e

Browse files
Andrey Cheptsovclaude
andcommitted
Adopt an existing job server socket instead of overwriting it
With multiple server replicas, each replica used to overwrite the job's /run/dstack/server.sock on every open, so ownership churned between replicas (they repeatedly stole the socket from each other) and every overwrite briefly broke access. Probe the socket first and (re)create the reverse forward only when it is missing or unreachable; otherwise adopt the existing one. This keeps a single stable owner. If the owner replica dies, another takes over on its next health check. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
1 parent e428da5 commit 42c8f6e

2 files changed

Lines changed: 74 additions & 30 deletions

File tree

src/dstack/_internal/server/services/jobs/server_connection.py

Lines changed: 26 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -91,17 +91,6 @@ async def open(self) -> None:
9191
try:
9292
remote_dir = shlex.quote(str(_REMOTE_SOCKET_PATH.parent))
9393
remote_socket = shlex.quote(str(_REMOTE_SOCKET_PATH))
94-
# A new server owner replaces the stable path, making an orphaned forward unreachable.
95-
await self._tunnel.aexec(
96-
f"mkdir -p {remote_dir} && chmod 755 {remote_dir} && rm -f {remote_socket}"
97-
)
98-
server_socket = _get_server_socket()
99-
self._tunnel.reverse_forwarded_sockets = [
100-
SocketPair(
101-
local=server_socket,
102-
remote=UnixSocket(path=_REMOTE_SOCKET_PATH),
103-
)
104-
]
10594
# Probe through the job socket itself: a socket path can remain after its listener
10695
# becomes unreachable.
10796
self._tunnel.forwarded_sockets = [
@@ -111,14 +100,34 @@ async def open(self) -> None:
111100
)
112101
]
113102
await self._tunnel.aopen()
114-
# The socket carries no credentials. World access inside the isolated job container
115-
# lets configurations using a non-root `user` reach it as well.
116-
await self._tunnel.aexec(f"chmod 666 {remote_socket}")
103+
# With multiple server replicas, do not overwrite a socket that another replica already
104+
# serves: create the reverse forward only when the current one is missing or
105+
# unreachable. This keeps a single stable owner and avoids ownership churn (replicas
106+
# repeatedly stealing the socket) and the brief unavailability of overwriting a live one.
117107
if not await self._server_is_reachable():
118-
raise SSHError(
119-
"dstack server is not reachable from the job"
120-
f" (forward target {server_socket.render()})"
108+
# A new server owner replaces the stable path, making an orphaned forward
109+
# unreachable.
110+
await self._tunnel.aexec(
111+
f"mkdir -p {remote_dir} && chmod 755 {remote_dir} && rm -f {remote_socket}"
121112
)
113+
server_socket = _get_server_socket()
114+
self._tunnel.reverse_forwarded_sockets = [
115+
SocketPair(
116+
local=server_socket,
117+
remote=UnixSocket(path=_REMOTE_SOCKET_PATH),
118+
)
119+
]
120+
# The probe forward is already established; do not request it again.
121+
self._tunnel.forwarded_sockets = []
122+
await self._tunnel.aopen()
123+
# The socket carries no credentials. World access inside the isolated job container
124+
# lets configurations using a non-root `user` reach it as well.
125+
await self._tunnel.aexec(f"chmod 666 {remote_socket}")
126+
if not await self._server_is_reachable():
127+
raise SSHError(
128+
"dstack server is not reachable from the job"
129+
f" (forward target {server_socket.render()})"
130+
)
122131
except Exception:
123132
await self._tunnel.aclose()
124133
raise

src/tests/_internal/server/services/jobs/test_server_connection.py

Lines changed: 48 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,8 @@
2020
def tunnel_mock(tmp_path, monkeypatch: pytest.MonkeyPatch):
2121
monkeypatch.setattr(server_connection, "CONNECTIONS_DIR", tmp_path)
2222
tunnel = MagicMock()
23+
tunnel.forwarded_sockets = []
24+
tunnel.reverse_forwarded_sockets = []
2325
tunnel.acheck = AsyncMock(return_value=False)
2426
tunnel.aopen = AsyncMock()
2527
tunnel.aexec = AsyncMock(return_value="")
@@ -64,33 +66,66 @@ def test_get_server_socket_follows_server_bind_host(
6466

6567
class TestJobServerConnection:
6668
@pytest.mark.asyncio
67-
async def test_opens_private_reverse_socket(self, tunnel_mock):
69+
async def test_becomes_owner_when_socket_unreachable(self, tunnel_mock):
6870
tunnel, tunnel_class = tunnel_mock
6971
job = Mock(id=uuid.uuid4())
7072
connection = JobServerConnection(job, job_runtime_data=None)
71-
connection._server_is_reachable = AsyncMock(return_value=True)
73+
# The probe after the local forward is unreachable (no healthy owner) -> become owner;
74+
# the probe after the reverse forward confirms reachability.
75+
connection._server_is_reachable = AsyncMock(side_effect=[False, True])
76+
77+
snapshots = []
78+
79+
async def record_aopen():
80+
snapshots.append(
81+
(list(tunnel.forwarded_sockets), list(tunnel.reverse_forwarded_sockets))
82+
)
83+
84+
tunnel.aopen.side_effect = record_aopen
7285

7386
await connection.open()
7487

7588
tunnel_class.assert_called_once()
76-
assert tunnel.aopen.await_count == 2
77-
assert tunnel.reverse_forwarded_sockets == [
78-
SocketPair(
79-
local=IPSocket(host="127.0.0.1", port=server_connection.settings.SERVER_PORT),
80-
remote=UnixSocket(path=server_connection._REMOTE_SOCKET_PATH),
81-
)
89+
probe_pair = SocketPair(
90+
local=UnixSocket(path=connection._probe_socket_path),
91+
remote=UnixSocket(path=server_connection._REMOTE_SOCKET_PATH),
92+
)
93+
reverse_pair = SocketPair(
94+
local=IPSocket(host="127.0.0.1", port=server_connection.settings.SERVER_PORT),
95+
remote=UnixSocket(path=server_connection._REMOTE_SOCKET_PATH),
96+
)
97+
# master (no forwards) -> add probe forward -> add reverse forward
98+
assert snapshots == [
99+
([], []),
100+
([probe_pair], []),
101+
([], [reverse_pair]),
102+
]
103+
commands = [call.args[0] for call in tunnel.aexec.await_args_list]
104+
assert commands == [
105+
"mkdir -p /run/dstack && chmod 755 /run/dstack && rm -f /run/dstack/server.sock",
106+
"chmod 666 /run/dstack/server.sock",
82107
]
108+
109+
@pytest.mark.asyncio
110+
async def test_adopts_existing_healthy_socket(self, tunnel_mock):
111+
tunnel, _ = tunnel_mock
112+
job = Mock(id=uuid.uuid4())
113+
connection = JobServerConnection(job, job_runtime_data=None)
114+
# A healthy owner (another replica) already serves the socket.
115+
connection._server_is_reachable = AsyncMock(return_value=True)
116+
117+
await connection.open()
118+
119+
# master + probe forward only; the reverse forward is not (re)created
120+
assert tunnel.aopen.await_count == 2
121+
assert tunnel.reverse_forwarded_sockets == []
122+
tunnel.aexec.assert_not_awaited()
83123
assert tunnel.forwarded_sockets == [
84124
SocketPair(
85125
local=UnixSocket(path=connection._probe_socket_path),
86126
remote=UnixSocket(path=server_connection._REMOTE_SOCKET_PATH),
87127
)
88128
]
89-
commands = [call.args[0] for call in tunnel.aexec.await_args_list]
90-
assert commands == [
91-
"mkdir -p /run/dstack && chmod 755 /run/dstack && rm -f /run/dstack/server.sock",
92-
"chmod 666 /run/dstack/server.sock",
93-
]
94129

95130
@pytest.mark.asyncio
96131
async def test_reuses_live_tunnel_with_existing_socket(self, tunnel_mock):

0 commit comments

Comments
 (0)