Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion mkdocs/docs/reference/env.md
Original file line number Diff line number Diff line change
Expand Up @@ -132,7 +132,7 @@ For more details on the options below, refer to the [server deployment](../guide
- `DSTACK_OTEL_METRICS_ENABLED`{ #DSTACK_OTEL_METRICS_ENABLED } – Enables OpenTelemetry metrics if set to any value. Requires the `otel` extra.
- `DSTACK_OTEL_METRICS_EXPORTERS`{ #DSTACK_OTEL_METRICS_EXPORTERS } – A comma-separated list of OpenTelemetry metrics exporters: `otlp` (push via OTLP) and/or `prometheus` (expose on the `/metrics` endpoint). Defaults to `otlp`.
- `DSTACK_DEFAULT_SERVICE_CLIENT_MAX_BODY_SIZE`{ #DSTACK_DEFAULT_SERVICE_CLIENT_MAX_BODY_SIZE } – Request body size limit for services running with a gateway, in bytes. Defaults to 64 MiB.
- `DSTACK_SERVICE_CLIENT_TIMEOUT`{ #DSTACK_SERVICE_CLIENT_TIMEOUT } – Timeout in seconds for HTTP requests sent from the in-server proxy and gateways to service replicas. Defaults to 60.
- `DSTACK_SERVICE_CLIENT_TIMEOUT`{ #DSTACK_SERVICE_CLIENT_TIMEOUT } – How long (in seconds) the in-server proxy and gateways wait for a service replica to send data. Defaults to 300.
- `DSTACK_FORBID_SERVICES_WITHOUT_GATEWAY`{ #DSTACK_FORBID_SERVICES_WITHOUT_GATEWAY } – Forbids registering new services without a gateway if set to any value.
- `DSTACK_FORBID_DSTACK_IN_RUNS`{ #DSTACK_FORBID_DSTACK_IN_RUNS } – Forbids submitting runs with `dstack: true` (dstack server access inside runs) if set to any value.
- `DSTACK_SERVER_CODE_UPLOAD_LIMIT`{ #DSTACK_SERVER_CODE_UPLOAD_LIMIT } - The repo size limit when uploading diffs or local repos, in bytes. Set to `0` to disable size limits. Defaults to `2MiB`.
Expand Down
2 changes: 2 additions & 0 deletions src/dstack/_internal/proxy/gateway/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,13 +4,15 @@

from pydantic import AnyHttpUrl

from dstack._internal.proxy.lib.const import DEFAULT_SERVICE_READ_TIMEOUT
from dstack._internal.proxy.lib.models import ImmutableModel


class ModelEntrypoint(ImmutableModel):
project_name: str
domain: str
https: bool
read_timeout: int = DEFAULT_SERVICE_READ_TIMEOUT


class ACMESettings(ImmutableModel):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@ server {
proxy_pass http://localhost:{{ proxy_port }}/api/models/{{ project_name }}/;
proxy_set_header X-Real-IP $remote_addr;
proxy_set_header Host $host;
proxy_read_timeout 300s;
proxy_read_timeout {{ read_timeout }}s;
}
listen 80;
{% if https %}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -68,7 +68,7 @@ server {
proxy_http_version 1.1;
proxy_set_header Upgrade $http_upgrade;
proxy_set_header Connection "Upgrade";
proxy_read_timeout 300s;
proxy_read_timeout {{ read_timeout }}s;
{% else %}
return 503;
{% endif %}
Expand All @@ -88,7 +88,7 @@ server {
proxy_pass http://{{ domain }}.upstream;
proxy_set_header X-Real-IP $remote_addr;
proxy_set_header Host $host;
proxy_read_timeout 300s;
proxy_read_timeout {{ read_timeout }}s;
{% if not proxy_buffering %}
proxy_buffering off;
{% endif %}
Expand Down
2 changes: 2 additions & 0 deletions src/dstack/_internal/proxy/gateway/routers/registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@ async def register_service(
rate_limits=body.rate_limits,
auth=body.auth,
client_max_body_size=body.client_max_body_size,
read_timeout=body.read_timeout,
model=body.options.openai.model if body.options.openai is not None else None,
ssh_private_key=body.ssh_private_key,
repo=repo,
Expand Down Expand Up @@ -139,6 +140,7 @@ async def register_entrypoint(
project_name=project_name.lower(),
domain=body.domain.lower(),
https=body.https,
read_timeout=body.read_timeout,
repo=repo,
nginx=nginx,
)
Expand Down
3 changes: 3 additions & 0 deletions src/dstack/_internal/proxy/gateway/schemas/registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
from pydantic import BaseModel, Field

from dstack._internal.core.models.instances import SSHConnectionParams
from dstack._internal.proxy.lib.const import DEFAULT_SERVICE_READ_TIMEOUT
from dstack._internal.proxy.lib.models import RateLimit


Expand Down Expand Up @@ -47,6 +48,7 @@ class RegisterServiceRequest(BaseModel):
ssh_private_key: str
rate_limits: tuple[RateLimit, ...] = ()
has_router_replica: bool = False
read_timeout: int = DEFAULT_SERVICE_READ_TIMEOUT


class SetServiceIdRequest(BaseModel):
Expand All @@ -68,3 +70,4 @@ class RegisterReplicaRequest(BaseModel):
class RegisterEntrypointRequest(BaseModel):
domain: str
https: bool
read_timeout: int = DEFAULT_SERVICE_READ_TIMEOUT
2 changes: 2 additions & 0 deletions src/dstack/_internal/proxy/gateway/services/nginx.py
Original file line number Diff line number Diff line change
Expand Up @@ -69,11 +69,13 @@ class ServiceConfig(SiteConfig):
has_router_replica: bool = False
cors_enabled: bool = False
proxy_buffering: bool = True
read_timeout: int


class ModelEntrypointConfig(SiteConfig):
type: Literal["entrypoint"] = "entrypoint"
project_name: str
read_timeout: int


class Nginx:
Expand Down
6 changes: 6 additions & 0 deletions src/dstack/_internal/proxy/gateway/services/registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,7 @@ async def register_service(
rate_limits: tuple[models.RateLimit, ...],
auth: bool,
client_max_body_size: int,
read_timeout: int,
model: Optional[schemas.AnyModel],
ssh_private_key: str,
repo: GatewayProxyRepo,
Expand All @@ -63,6 +64,7 @@ async def register_service(
has_router_replica=has_router_replica,
cors_enabled=cors_enabled,
proxy_buffering=model is None,
read_timeout=read_timeout,
)

async with lock:
Expand Down Expand Up @@ -252,13 +254,15 @@ async def register_model_entrypoint(
project_name: str,
domain: str,
https: bool,
read_timeout: int,
repo: GatewayProxyRepo,
nginx: Nginx,
) -> None:
entrypoint = gateway_models.ModelEntrypoint(
project_name=project_name,
domain=domain,
https=https,
read_timeout=read_timeout,
)
logger.debug("Registering entrypoint %s in project %s", domain, project_name)
await apply_entrypoint(entrypoint, repo, nginx)
Expand Down Expand Up @@ -409,6 +413,7 @@ async def get_nginx_service_config(
has_router_replica=service.has_router_replica,
cors_enabled=service.cors_enabled,
proxy_buffering=service.proxy_buffering,
read_timeout=service.read_timeout,
)


Expand All @@ -419,6 +424,7 @@ async def apply_entrypoint(
domain=entrypoint.domain,
https=entrypoint.https,
project_name=entrypoint.project_name,
read_timeout=entrypoint.read_timeout,
)
acme = (await repo.get_config()).acme_settings
await nginx.register(config, acme)
Expand Down
2 changes: 2 additions & 0 deletions src/dstack/_internal/proxy/lib/const.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@
Shared constants for proxy components (gateway + in-server proxy).
"""

DEFAULT_SERVICE_READ_TIMEOUT = 300

# Inference endpoints exposed by the in-replica HTTP router. Applies to both
# SGLang's router and Dynamo's `dynamo.frontend` — they share the
# OpenAI-compatible endpoint surface.
Expand Down
3 changes: 3 additions & 0 deletions src/dstack/_internal/proxy/lib/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
from typing_extensions import Annotated

from dstack._internal.core.models.instances import SSHConnectionParams
from dstack._internal.proxy.lib.const import DEFAULT_SERVICE_READ_TIMEOUT
from dstack._internal.proxy.lib.errors import UnexpectedProxyError


Expand Down Expand Up @@ -66,6 +67,8 @@ class Service(ImmutableModel):
proxy_buffering: bool = True
"""Only used on gateways. Disabled for services with a model so that streamed responses
reach the client as the replica emits them."""
read_timeout: int = DEFAULT_SERVICE_READ_TIMEOUT
"""Seconds to wait for data from a replica between two successive reads."""

@model_validator(mode="before")
@classmethod
Expand Down
7 changes: 2 additions & 5 deletions src/dstack/_internal/proxy/lib/services/service_connection.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,14 +19,11 @@
from dstack._internal.proxy.lib.models import Project, Replica, Service
from dstack._internal.proxy.lib.repo import BaseProxyRepo
from dstack._internal.utils.common import get_or_error
from dstack._internal.utils.env import environ
from dstack._internal.utils.logging import get_logger
from dstack._internal.utils.path import FileContent

logger = get_logger(__name__)
OPEN_TUNNEL_TIMEOUT = 10
HTTP_TIMEOUT = environ.get_int("DSTACK_SERVICE_CLIENT_TIMEOUT", default=60)
# Same as default Nginx proxy timeout; override via DSTACK_SERVICE_CLIENT_TIMEOUT


class ServiceClient(httpx.AsyncClient):
Expand Down Expand Up @@ -75,7 +72,7 @@ def __init__(self, project: Project, service: Service, replica: Replica) -> None
# The hostname in base_url is there for troubleshooting, as it may appear in
# logs and in the Host header. The actual destination is the Unix socket.
base_url=f"http://{replica.id}-{service.run_name}/",
timeout=HTTP_TIMEOUT,
timeout=service.read_timeout,
)
self._is_open = asyncio.locks.Event()

Expand Down Expand Up @@ -148,7 +145,7 @@ async def get_service_replica_client(
return httpx.AsyncClient(
base_url="http://127.0.0.1",
headers={"Host": service.domain},
timeout=HTTP_TIMEOUT,
timeout=service.read_timeout,
)
# Nginx not available, forward directly to the tunnel
replica = random.choice(service.replicas)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1288,6 +1288,7 @@ async def _register_service(
),
auth=run_spec.configuration.auth,
client_max_body_size=settings.DEFAULT_SERVICE_CLIENT_MAX_BODY_SIZE,
read_timeout=settings.SERVICE_CLIENT_TIMEOUT,
options=service_spec.options,
rate_limits=run_spec.configuration.rate_limits,
ssh_private_key=run_model.project.ssh_private_key,
Expand Down
9 changes: 7 additions & 2 deletions src/dstack/_internal/server/services/gateways/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,14 +45,15 @@ async def register_service(
gateway_https: bool,
auth: bool,
client_max_body_size: int,
read_timeout: int,
options: dict,
rate_limits: list[RateLimit],
ssh_private_key: str,
has_router_replica: bool = False,
):
if "openai" in options:
entrypoint = f"gateway.{domain.split('.', maxsplit=1)[1]}"
await self.register_openai_entrypoint(project, entrypoint, gateway_https)
await self.register_openai_entrypoint(project, entrypoint, gateway_https, read_timeout)

payload = {
"id": run_id.hex,
Expand All @@ -61,6 +62,7 @@ async def register_service(
"https": service_https,
"auth": auth,
"client_max_body_size": client_max_body_size,
"read_timeout": read_timeout,
"options": options,
"rate_limits": [limit.model_dump() for limit in rate_limits],
"ssh_private_key": ssh_private_key,
Expand Down Expand Up @@ -160,12 +162,15 @@ async def set_service_id(self, project: str, run_name: str, run_id: uuid.UUID) -
resp.raise_for_status()
self.is_server_ready = True

async def register_openai_entrypoint(self, project: str, domain: str, https: bool):
async def register_openai_entrypoint(
self, project: str, domain: str, https: bool, read_timeout: int
):
resp = await self._client.post(
self._url(f"/api/registry/{project}/entrypoints/register"),
json={
"domain": domain,
"https": https,
"read_timeout": read_timeout,
},
)
if resp.status_code == 400:
Expand Down
6 changes: 5 additions & 1 deletion src/dstack/_internal/server/services/proxy/repo.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,10 @@
from dstack._internal.server.services.instances import get_instance_remote_connection_info
from dstack._internal.server.services.jobs import get_job_spec
from dstack._internal.server.services.runs import get_run_spec
from dstack._internal.server.settings import DEFAULT_SERVICE_CLIENT_MAX_BODY_SIZE
from dstack._internal.server.settings import (
DEFAULT_SERVICE_CLIENT_MAX_BODY_SIZE,
SERVICE_CLIENT_TIMEOUT,
)
from dstack._internal.utils.common import get_or_error

_ANY_MODEL_ADAPTER = pydantic.TypeAdapter(AnyModel)
Expand Down Expand Up @@ -136,6 +139,7 @@ async def get_service(self, project_name: str, run_name: str) -> Optional[Servic
strip_prefix=run_spec.configuration.strip_prefix,
replicas=tuple(replicas),
has_router_replica=has_router_replica,
read_timeout=SERVICE_CLIENT_TIMEOUT,
)

async def list_models(self, project_name: str) -> List[ChatModel]:
Expand Down
1 change: 1 addition & 0 deletions src/dstack/_internal/server/settings.py
Original file line number Diff line number Diff line change
Expand Up @@ -180,6 +180,7 @@ def get_database_url() -> str:
DEFAULT_SERVICE_CLIENT_MAX_BODY_SIZE = int(
os.getenv("DSTACK_DEFAULT_SERVICE_CLIENT_MAX_BODY_SIZE", 64 * 1024 * 1024)
)
SERVICE_CLIENT_TIMEOUT = environ.get_int("DSTACK_SERVICE_CLIENT_TIMEOUT", default=300)

SERVER_DEFAULT_DOCKER_REGISTRY = os.getenv("DSTACK_SERVER_DEFAULT_DOCKER_REGISTRY") or None
SERVER_DEFAULT_DOCKER_REGISTRY_USERNAME = (
Expand Down
39 changes: 38 additions & 1 deletion src/tests/_internal/proxy/gateway/routers/test_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,8 +31,9 @@ def register_service_payload(
client_max_body_size: int = 1024,
options: Optional[dict] = None,
rate_limits: Optional[list[dict]] = None,
read_timeout: Optional[int] = None,
) -> dict:
return {
payload = {
"id": uuid.uuid4().hex,
"run_name": run_name,
"domain": domain,
Expand All @@ -43,6 +44,9 @@ def register_service_payload(
"rate_limits": rate_limits or [],
"ssh_private_key": "private-key",
}
if read_timeout is not None:
payload["read_timeout"] = read_timeout
return payload


def register_replica_payload(job_id: str = "xxx-xxx") -> dict:
Expand Down Expand Up @@ -234,6 +238,26 @@ async def test_register_with_model(self, tmp_path: Path, system_mocks: Mocks) ->
conf = (tmp_path / "sites-enabled" / "443-test-run.gtw.test.conf").read_text()
assert "proxy_buffering off;" in conf

async def test_register_with_read_timeout(self, tmp_path: Path, system_mocks: Mocks) -> None:
repo = GatewayProxyRepo()
client = make_client(tmp_path, repo=repo)
resp = await client.post(
"/api/registry/test-proj/services/register",
json=register_service_payload(run_name="test-run", read_timeout=900),
)
assert resp.status_code == 200
service = await repo.get_service("test-proj", "test-run")
assert service is not None
assert service.read_timeout == 900
resp = await client.post(
"/api/registry/test-proj/services/test-run/replicas/register",
json=register_replica_payload(),
)
assert resp.status_code == 200
conf = (tmp_path / "sites-enabled" / "443-test-run.gtw.test.conf").read_text()
assert conf.count("proxy_read_timeout 900s;") == 2
assert "proxy_read_timeout 300s;" not in conf

async def test_register_with_rate_limits(self, tmp_path: Path, system_mocks: Mocks) -> None:
client = make_client(tmp_path)
resp = await client.post(
Expand Down Expand Up @@ -599,6 +623,19 @@ async def test_register(self, tmp_path: Path, system_mocks: Mocks) -> None:
assert system_mocks.reload_nginx.call_count == 1
assert system_mocks.run_certbot.call_count == 0

async def test_register_with_read_timeout(self, tmp_path: Path, system_mocks: Mocks) -> None:
repo = GatewayProxyRepo()
client = make_client(tmp_path, repo=repo)
resp = await client.post(
"/api/registry/test-proj/entrypoints/register",
json={"domain": "gateway.gtw.test", "https": False, "read_timeout": 900},
)
assert resp.status_code == 200
conf = (tmp_path / "sites-enabled" / "443-gateway.gtw.test.conf").read_text()
assert "proxy_read_timeout 900s;" in conf
[entrypoint] = await repo.list_entrypoints()
assert entrypoint.read_timeout == 900

async def test_register_with_https(self, tmp_path: Path, system_mocks: Mocks) -> None:
client = make_client(tmp_path)
resp = await client.post(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
from dstack._internal.core.models.instances import InstanceStatus
from dstack._internal.core.models.runs import JobStatus, RunStatus, ServiceSpec
from dstack._internal.proxy.gateway.schemas.services import ServiceListItem, ServiceListReplicaItem
from dstack._internal.server import settings
from dstack._internal.server.background.pipeline_tasks.gateway_replicas import (
GatewayReplicaFetcher,
GatewayReplicaPipeline,
Expand Down Expand Up @@ -817,7 +818,9 @@ async def test_registers_new_service_and_replica(
session: AsyncSession,
worker: GatewayReplicaWorker,
mock_gateway_connection: AsyncMock,
monkeypatch: pytest.MonkeyPatch,
):
monkeypatch.setattr(settings, "SERVICE_CLIENT_TIMEOUT", 900)
project = await create_project(session=session)
user = await create_user(session=session)
repo = await create_repo(session=session, project_id=project.id)
Expand Down Expand Up @@ -858,6 +861,7 @@ async def test_registers_new_service_and_replica(
gateway_https=ANY,
auth=ANY,
client_max_body_size=ANY,
read_timeout=900,
options={},
rate_limits=[],
ssh_private_key=project.ssh_private_key,
Expand Down Expand Up @@ -1432,6 +1436,7 @@ async def test_unregisters_and_reregisters_legacy_service_without_id_and_replica
gateway_https=ANY,
auth=ANY,
client_max_body_size=ANY,
read_timeout=ANY,
options={},
rate_limits=[],
ssh_private_key=project.ssh_private_key,
Expand Down
Loading