diff --git a/src/dstack/_internal/proxy/gateway/resources/nginx/service.jinja2 b/src/dstack/_internal/proxy/gateway/resources/nginx/service.jinja2 index 7c8a5fcc9..4f2adda38 100644 --- a/src/dstack/_internal/proxy/gateway/resources/nginx/service.jinja2 +++ b/src/dstack/_internal/proxy/gateway/resources/nginx/service.jinja2 @@ -89,6 +89,9 @@ server { proxy_set_header X-Real-IP $remote_addr; proxy_set_header Host $host; proxy_read_timeout 300s; + {% if not proxy_buffering %} + proxy_buffering off; + {% endif %} {% else %} return 503; {% endif %} diff --git a/src/dstack/_internal/proxy/gateway/services/nginx.py b/src/dstack/_internal/proxy/gateway/services/nginx.py index 74581c647..1486ea0c2 100644 --- a/src/dstack/_internal/proxy/gateway/services/nginx.py +++ b/src/dstack/_internal/proxy/gateway/services/nginx.py @@ -68,6 +68,7 @@ class ServiceConfig(SiteConfig): replicas: list[ReplicaConfig] has_router_replica: bool = False cors_enabled: bool = False + proxy_buffering: bool = True class ModelEntrypointConfig(SiteConfig): diff --git a/src/dstack/_internal/proxy/gateway/services/registry.py b/src/dstack/_internal/proxy/gateway/services/registry.py index 592a96492..c69b57136 100644 --- a/src/dstack/_internal/proxy/gateway/services/registry.py +++ b/src/dstack/_internal/proxy/gateway/services/registry.py @@ -62,6 +62,7 @@ async def register_service( replicas=(), has_router_replica=has_router_replica, cors_enabled=cors_enabled, + proxy_buffering=model is None, ) async with lock: @@ -407,6 +408,7 @@ async def get_nginx_service_config( replicas=sorted(replicas, key=lambda r: r.id), # sort for reproducible configs has_router_replica=service.has_router_replica, cors_enabled=service.cors_enabled, + proxy_buffering=service.proxy_buffering, ) diff --git a/src/dstack/_internal/proxy/lib/models.py b/src/dstack/_internal/proxy/lib/models.py index 53eb13e74..511f27d70 100644 --- a/src/dstack/_internal/proxy/lib/models.py +++ b/src/dstack/_internal/proxy/lib/models.py @@ -63,6 +63,9 @@ class Service(ImmutableModel): replicas: tuple[Replica, ...] has_router_replica: bool = False cors_enabled: bool = False # only used on gateways; enabled for openai-format models + 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.""" @model_validator(mode="before") @classmethod diff --git a/src/tests/_internal/proxy/gateway/routers/test_registry.py b/src/tests/_internal/proxy/gateway/routers/test_registry.py index 20a0e148c..48c7a41d5 100644 --- a/src/tests/_internal/proxy/gateway/routers/test_registry.py +++ b/src/tests/_internal/proxy/gateway/routers/test_registry.py @@ -226,6 +226,13 @@ async def test_register_with_model(self, tmp_path: Path, system_mocks: Mocks) -> format_spec=OpenAIChatModelFormat(prefix="/v1"), ) ] + 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 "proxy_buffering off;" in conf async def test_register_with_rate_limits(self, tmp_path: Path, system_mocks: Mocks) -> None: client = make_client(tmp_path) @@ -325,6 +332,7 @@ async def test_register(self, tmp_path: Path, system_mocks: Mocks) -> None: assert (m1 := re.search(r"server unix:/(.+)/replica.sock; # replica xxx-xxx", conf)) assert (m2 := re.search(r"server unix:/(.+)/replica.sock; # replica yyy-yyy", conf)) assert m1.group(1) != m2.group(1) + assert "proxy_buffering" not in conf assert system_mocks.reload_nginx.call_count == 3 assert system_mocks.open_conn.call_count == 2