diff --git a/mkdocs/docs/reference/env.md b/mkdocs/docs/reference/env.md index 4a169fc22f..c721de395c 100644 --- a/mkdocs/docs/reference/env.md +++ b/mkdocs/docs/reference/env.md @@ -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`. diff --git a/src/dstack/_internal/proxy/gateway/models.py b/src/dstack/_internal/proxy/gateway/models.py index 3e0a76c032..80ce0b8abe 100644 --- a/src/dstack/_internal/proxy/gateway/models.py +++ b/src/dstack/_internal/proxy/gateway/models.py @@ -4,6 +4,7 @@ from pydantic import AnyHttpUrl +from dstack._internal.proxy.lib.const import DEFAULT_SERVICE_READ_TIMEOUT from dstack._internal.proxy.lib.models import ImmutableModel @@ -11,6 +12,7 @@ class ModelEntrypoint(ImmutableModel): project_name: str domain: str https: bool + read_timeout: int = DEFAULT_SERVICE_READ_TIMEOUT class ACMESettings(ImmutableModel): diff --git a/src/dstack/_internal/proxy/gateway/resources/nginx/entrypoint.jinja2 b/src/dstack/_internal/proxy/gateway/resources/nginx/entrypoint.jinja2 index bcb66e1dd4..ce0160f558 100644 --- a/src/dstack/_internal/proxy/gateway/resources/nginx/entrypoint.jinja2 +++ b/src/dstack/_internal/proxy/gateway/resources/nginx/entrypoint.jinja2 @@ -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 %} diff --git a/src/dstack/_internal/proxy/gateway/resources/nginx/service.jinja2 b/src/dstack/_internal/proxy/gateway/resources/nginx/service.jinja2 index 4f2adda389..e3eced5f6d 100644 --- a/src/dstack/_internal/proxy/gateway/resources/nginx/service.jinja2 +++ b/src/dstack/_internal/proxy/gateway/resources/nginx/service.jinja2 @@ -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 %} @@ -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 %} diff --git a/src/dstack/_internal/proxy/gateway/routers/registry.py b/src/dstack/_internal/proxy/gateway/routers/registry.py index c5fb7c0d6a..954ae2933a 100644 --- a/src/dstack/_internal/proxy/gateway/routers/registry.py +++ b/src/dstack/_internal/proxy/gateway/routers/registry.py @@ -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, @@ -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, ) diff --git a/src/dstack/_internal/proxy/gateway/schemas/registry.py b/src/dstack/_internal/proxy/gateway/schemas/registry.py index 9a800c2ceb..25f801252a 100644 --- a/src/dstack/_internal/proxy/gateway/schemas/registry.py +++ b/src/dstack/_internal/proxy/gateway/schemas/registry.py @@ -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 @@ -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): @@ -68,3 +70,4 @@ class RegisterReplicaRequest(BaseModel): class RegisterEntrypointRequest(BaseModel): domain: str https: bool + read_timeout: int = DEFAULT_SERVICE_READ_TIMEOUT diff --git a/src/dstack/_internal/proxy/gateway/services/nginx.py b/src/dstack/_internal/proxy/gateway/services/nginx.py index 1486ea0c22..df71594f5d 100644 --- a/src/dstack/_internal/proxy/gateway/services/nginx.py +++ b/src/dstack/_internal/proxy/gateway/services/nginx.py @@ -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: diff --git a/src/dstack/_internal/proxy/gateway/services/registry.py b/src/dstack/_internal/proxy/gateway/services/registry.py index c69b57136a..2ce9ea9ffd 100644 --- a/src/dstack/_internal/proxy/gateway/services/registry.py +++ b/src/dstack/_internal/proxy/gateway/services/registry.py @@ -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, @@ -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: @@ -252,6 +254,7 @@ async def register_model_entrypoint( project_name: str, domain: str, https: bool, + read_timeout: int, repo: GatewayProxyRepo, nginx: Nginx, ) -> None: @@ -259,6 +262,7 @@ async def register_model_entrypoint( 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) @@ -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, ) @@ -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) diff --git a/src/dstack/_internal/proxy/lib/const.py b/src/dstack/_internal/proxy/lib/const.py index 43ede03ac8..0824a5edb7 100644 --- a/src/dstack/_internal/proxy/lib/const.py +++ b/src/dstack/_internal/proxy/lib/const.py @@ -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. diff --git a/src/dstack/_internal/proxy/lib/models.py b/src/dstack/_internal/proxy/lib/models.py index 511f27d70c..44d5878bbd 100644 --- a/src/dstack/_internal/proxy/lib/models.py +++ b/src/dstack/_internal/proxy/lib/models.py @@ -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 @@ -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 diff --git a/src/dstack/_internal/proxy/lib/services/service_connection.py b/src/dstack/_internal/proxy/lib/services/service_connection.py index 37bdc5083a..c8229ad53a 100644 --- a/src/dstack/_internal/proxy/lib/services/service_connection.py +++ b/src/dstack/_internal/proxy/lib/services/service_connection.py @@ -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): @@ -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() @@ -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) diff --git a/src/dstack/_internal/server/background/pipeline_tasks/gateway_replicas.py b/src/dstack/_internal/server/background/pipeline_tasks/gateway_replicas.py index 44fd687cda..7e51f90eff 100644 --- a/src/dstack/_internal/server/background/pipeline_tasks/gateway_replicas.py +++ b/src/dstack/_internal/server/background/pipeline_tasks/gateway_replicas.py @@ -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, diff --git a/src/dstack/_internal/server/services/gateways/client.py b/src/dstack/_internal/server/services/gateways/client.py index ed6c8d78b9..0c9defecca 100644 --- a/src/dstack/_internal/server/services/gateways/client.py +++ b/src/dstack/_internal/server/services/gateways/client.py @@ -45,6 +45,7 @@ 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, @@ -52,7 +53,7 @@ async def register_service( ): 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, @@ -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, @@ -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: diff --git a/src/dstack/_internal/server/services/proxy/repo.py b/src/dstack/_internal/server/services/proxy/repo.py index d0afd41def..7ea7df43cd 100644 --- a/src/dstack/_internal/server/services/proxy/repo.py +++ b/src/dstack/_internal/server/services/proxy/repo.py @@ -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) @@ -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]: diff --git a/src/dstack/_internal/server/settings.py b/src/dstack/_internal/server/settings.py index efe6e320fc..c077feb4ea 100644 --- a/src/dstack/_internal/server/settings.py +++ b/src/dstack/_internal/server/settings.py @@ -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 = ( diff --git a/src/tests/_internal/proxy/gateway/routers/test_registry.py b/src/tests/_internal/proxy/gateway/routers/test_registry.py index 48c7a41d5f..393b901acc 100644 --- a/src/tests/_internal/proxy/gateway/routers/test_registry.py +++ b/src/tests/_internal/proxy/gateway/routers/test_registry.py @@ -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, @@ -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: @@ -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( @@ -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( diff --git a/src/tests/_internal/server/background/pipeline_tasks/test_gateway_replicas.py b/src/tests/_internal/server/background/pipeline_tasks/test_gateway_replicas.py index 07f0a00b5f..aaddd0b7eb 100644 --- a/src/tests/_internal/server/background/pipeline_tasks/test_gateway_replicas.py +++ b/src/tests/_internal/server/background/pipeline_tasks/test_gateway_replicas.py @@ -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, @@ -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) @@ -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, @@ -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,