Skip to content

Commit 53ebca0

Browse files
authored
Refactor/shared replica tunnel (#3978)
* Refractor shared replica tunnel * Minor Update --------- Co-authored-by: Bihan Rana
1 parent 67d49cb commit 53ebca0

7 files changed

Lines changed: 237 additions & 134 deletions

File tree

src/dstack/_internal/server/background/scheduled_tasks/probes.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -12,9 +12,9 @@
1212
from dstack._internal.server.models import InstanceModel, JobModel, ProbeModel
1313
from dstack._internal.server.services.jobs import get_job_spec
1414
from dstack._internal.server.services.jobs.job_replica_http_client import (
15-
SSH_CONNECT_TIMEOUT,
1615
get_service_replica_client,
1716
)
17+
from dstack._internal.server.services.jobs.job_replica_tunnel import SSH_CONNECT_TIMEOUT
1818
from dstack._internal.server.services.locking import get_locker
1919
from dstack._internal.server.services.logging import fmt
2020
from dstack._internal.utils.common import get_current_datetime

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

Lines changed: 16 additions & 36 deletions
Original file line numberDiff line numberDiff line change
@@ -2,25 +2,14 @@
22

33
from collections.abc import AsyncGenerator
44
from contextlib import asynccontextmanager
5-
from datetime import timedelta
65
from pathlib import Path
7-
from tempfile import TemporaryDirectory
86
from typing import Any
97

108
import grpc
119

12-
from dstack._internal.core.services.ssh.tunnel import (
13-
SSH_DEFAULT_OPTIONS,
14-
IPSocket,
15-
SocketPair,
16-
UnixSocket,
17-
)
1810
from dstack._internal.server.models import JobModel
19-
from dstack._internal.server.services.jobs import get_job_spec
20-
from dstack._internal.server.services.ssh import container_ssh_tunnel
21-
from dstack._internal.utils.common import get_or_error
11+
from dstack._internal.server.services.jobs.job_replica_tunnel import get_service_replica_tunnel
2212

23-
SSH_CONNECT_TIMEOUT = timedelta(seconds=10)
2413
# Match router_worker_sync HTTP server_info cap (_MAX_SERVER_INFO_RESPONSE_BYTES).
2514
_MAX_GRPC_MESSAGE_BYTES = 256 * 1024
2615
_GRPC_CHANNEL_OPTIONS = (
@@ -29,29 +18,20 @@
2918
)
3019

3120

21+
@asynccontextmanager
22+
async def get_service_replica_grpc_channel_over_uds(
23+
uds_path: Path,
24+
) -> AsyncGenerator[Any, None]:
25+
target = f"unix://{uds_path}"
26+
channel = grpc.aio.insecure_channel(target, options=_GRPC_CHANNEL_OPTIONS)
27+
try:
28+
yield channel
29+
finally:
30+
await channel.close()
31+
32+
3233
@asynccontextmanager
3334
async def get_service_replica_grpc_client(job: JobModel) -> AsyncGenerator[Any, None]:
34-
options = {
35-
**SSH_DEFAULT_OPTIONS,
36-
"ConnectTimeout": str(int(SSH_CONNECT_TIMEOUT.total_seconds())),
37-
}
38-
job_spec = get_job_spec(job)
39-
with TemporaryDirectory() as temp_dir:
40-
# Keep the same socket file name as the HTTP helper for consistency.
41-
app_socket_path = (Path(temp_dir) / "replica.sock").absolute()
42-
async with container_ssh_tunnel(
43-
job=job,
44-
forwarded_sockets=[
45-
SocketPair(
46-
remote=IPSocket("localhost", get_or_error(job_spec.service_port)),
47-
local=UnixSocket(app_socket_path),
48-
),
49-
],
50-
options=options,
51-
):
52-
target = f"unix://{app_socket_path}"
53-
channel = grpc.aio.insecure_channel(target, options=_GRPC_CHANNEL_OPTIONS)
54-
try:
55-
yield channel
56-
finally:
57-
await channel.close()
35+
async with get_service_replica_tunnel(job) as uds_path:
36+
async with get_service_replica_grpc_channel_over_uds(uds_path) as channel:
37+
yield channel

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

Lines changed: 11 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -2,48 +2,26 @@
22

33
from collections.abc import AsyncGenerator
44
from contextlib import asynccontextmanager
5-
from datetime import timedelta
65
from pathlib import Path
7-
from tempfile import TemporaryDirectory
86

97
from httpx import AsyncClient, AsyncHTTPTransport
108

11-
from dstack._internal.core.services.ssh.tunnel import (
12-
SSH_DEFAULT_OPTIONS,
13-
IPSocket,
14-
SocketPair,
15-
UnixSocket,
16-
)
179
from dstack._internal.server.models import JobModel
18-
from dstack._internal.server.services.jobs import get_job_spec
19-
from dstack._internal.server.services.ssh import container_ssh_tunnel
20-
from dstack._internal.utils.common import get_or_error
10+
from dstack._internal.server.services.jobs.job_replica_tunnel import get_service_replica_tunnel
2111

22-
SSH_CONNECT_TIMEOUT = timedelta(seconds=10)
12+
13+
@asynccontextmanager
14+
async def get_service_replica_http_client_over_uds(
15+
uds_path: Path,
16+
) -> AsyncGenerator[AsyncClient, None]:
17+
async with AsyncClient(transport=AsyncHTTPTransport(uds=str(uds_path))) as client:
18+
yield client
2319

2420

2521
@asynccontextmanager
2622
async def get_service_replica_client(
2723
job: JobModel,
2824
) -> AsyncGenerator[AsyncClient, None]:
29-
options = {
30-
**SSH_DEFAULT_OPTIONS,
31-
"ConnectTimeout": str(int(SSH_CONNECT_TIMEOUT.total_seconds())),
32-
}
33-
job_spec = get_job_spec(job)
34-
with TemporaryDirectory() as temp_dir:
35-
app_socket_path = (Path(temp_dir) / "replica.sock").absolute()
36-
async with container_ssh_tunnel(
37-
job=job,
38-
forwarded_sockets=[
39-
SocketPair(
40-
remote=IPSocket("localhost", get_or_error(job_spec.service_port)),
41-
local=UnixSocket(app_socket_path),
42-
),
43-
],
44-
options=options,
45-
):
46-
async with AsyncClient(
47-
transport=AsyncHTTPTransport(uds=str(app_socket_path))
48-
) as client:
49-
yield client
25+
async with get_service_replica_tunnel(job) as uds_path:
26+
async with get_service_replica_http_client_over_uds(uds_path) as client:
27+
yield client
Lines changed: 43 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,43 @@
1+
"""SSH tunnel to a job replica's service port, exposed as a local Unix domain socket."""
2+
3+
from collections.abc import AsyncGenerator
4+
from contextlib import asynccontextmanager
5+
from datetime import timedelta
6+
from pathlib import Path
7+
from tempfile import TemporaryDirectory
8+
9+
from dstack._internal.core.services.ssh.tunnel import (
10+
SSH_DEFAULT_OPTIONS,
11+
IPSocket,
12+
SocketPair,
13+
UnixSocket,
14+
)
15+
from dstack._internal.server.models import JobModel
16+
from dstack._internal.server.services.jobs import get_job_spec
17+
from dstack._internal.server.services.ssh import container_ssh_tunnel
18+
from dstack._internal.utils.common import get_or_error
19+
20+
SSH_CONNECT_TIMEOUT = timedelta(seconds=10)
21+
_REPLICA_SOCKET_NAME = "replica.sock"
22+
23+
24+
@asynccontextmanager
25+
async def get_service_replica_tunnel(job: JobModel) -> AsyncGenerator[Path, None]:
26+
options = {
27+
**SSH_DEFAULT_OPTIONS,
28+
"ConnectTimeout": str(int(SSH_CONNECT_TIMEOUT.total_seconds())),
29+
}
30+
job_spec = get_job_spec(job)
31+
with TemporaryDirectory() as temp_dir:
32+
app_socket_path = (Path(temp_dir) / _REPLICA_SOCKET_NAME).absolute()
33+
async with container_ssh_tunnel(
34+
job=job,
35+
forwarded_sockets=[
36+
SocketPair(
37+
remote=IPSocket("localhost", get_or_error(job_spec.service_port)),
38+
local=UnixSocket(app_socket_path),
39+
),
40+
],
41+
options=options,
42+
):
43+
yield app_socket_path

0 commit comments

Comments
 (0)