Skip to content
Closed
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
Original file line number Diff line number Diff line change
Expand Up @@ -76,6 +76,9 @@ server {
location @ {
set $dstack_replica_hit 1;
{% if replicas %}
{% if not proxy_buffering %}
proxy_buffering off;
{% endif %}
{% if cors_enabled %}
proxy_hide_header 'Access-Control-Allow-Origin';
proxy_hide_header 'Access-Control-Allow-Methods';
Expand Down
1 change: 1 addition & 0 deletions src/dstack/_internal/proxy/gateway/services/nginx.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
16 changes: 16 additions & 0 deletions src/dstack/_internal/proxy/gateway/services/registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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,
)


Expand Down Expand Up @@ -446,10 +448,24 @@ async def _migrate_cors_enabled(repo: GatewayProxyRepo) -> None:
await repo.set_service(updated)


async def _migrate_proxy_buffering(repo: GatewayProxyRepo) -> None:
"""Disable buffering for model services saved by older gateways."""
services = await repo.list_services()
model_runs = {
(project_name, model.run_name)
for project_name in {service.project_name for service in services}
for model in await repo.list_models(project_name)
}
for service in services:
if service.proxy_buffering and (service.project_name, service.run_name) in model_runs:
await repo.set_service(service.model_copy(update={"proxy_buffering": False}))


async def apply_all(
repo: GatewayProxyRepo, nginx: Nginx, service_conn_pool: ServiceConnectionPool
) -> None:
await _migrate_cors_enabled(repo)
await _migrate_proxy_buffering(repo)
service_tasks = [
apply_service(
service=service,
Expand Down
1 change: 1 addition & 0 deletions src/dstack/_internal/proxy/lib/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,7 @@ 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 model services

@model_validator(mode="before")
@classmethod
Expand Down
28 changes: 28 additions & 0 deletions src/tests/_internal/proxy/gateway/routers/test_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -88,6 +88,34 @@ def sample_model_options(name: str = "test-model") -> dict:

@pytest.mark.asyncio
class TestRegisterService:
@pytest.mark.parametrize("model_format", [None, "openai", "tgi"])
async def test_proxy_buffering(
self, tmp_path: Path, system_mocks: Mocks, model_format: Optional[str]
) -> None:
repo = GatewayProxyRepo()
client = make_client(tmp_path, repo=repo)
options = None
if model_format is not None:
model = {"type": "chat", "name": "test-model", "format": model_format}
if model_format == "openai":
model["prefix"] = "/v1"
else:
model.update(chat_template="{{ messages }}", eos_token="</s>")
options = {"openai": {"model": model}}
response = await client.post(
"/api/registry/test-proj/services/register",
json=register_service_payload(options=options),
)
assert response.status_code == 200
response = await client.post(
"/api/registry/test-proj/services/test-run/replicas/register",
json=register_replica_payload(),
)
assert response.status_code == 200
conf = (tmp_path / "sites-enabled" / "443-test-run.gtw.test.conf").read_text()
service_location = conf.split("location @ {", 1)[1].split("}", 1)[0]
assert ("proxy_buffering off;" in service_location) == (model_format is not None)

async def test_register(self, tmp_path: Path, system_mocks: Mocks) -> None:
client = make_client(tmp_path)
resp = await client.post(
Expand Down
55 changes: 55 additions & 0 deletions src/tests/_internal/proxy/gateway/test_app.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
import json
from datetime import datetime
from pathlib import Path

import pytest
Expand All @@ -7,6 +9,11 @@
from dstack._internal.proxy.gateway.repo.repo import GatewayProxyRepo
from dstack._internal.proxy.gateway.services.nginx import Nginx
from dstack._internal.proxy.gateway.testing.common import Mocks
from dstack._internal.proxy.lib.models import (
ChatModel,
OpenAIChatModelFormat,
TGIChatModelFormat,
)
from dstack._internal.proxy.lib.testing.common import make_project, make_service


Expand All @@ -32,3 +39,51 @@ async def test_lifespan(tmp_path: Path, system_mocks: Mocks) -> None:
assert system_mocks.open_conn.call_count == 1
assert system_mocks.close_conn.call_count == 0
assert system_mocks.close_conn.call_count == 1


@pytest.mark.asyncio
@pytest.mark.parametrize("model_format", ["openai", "tgi"])
async def test_lifespan_migrates_proxy_buffering(
tmp_path: Path, system_mocks: Mocks, model_format: str
) -> None:
state_file = tmp_path / "state-v2.json"
repo = GatewayProxyRepo(file=state_file)
for project in ("model-proj", "web-proj"):
await repo.set_project(make_project(project))
await repo.set_service(
make_service(project, "same-run", domain=f"{project}.gtw.test", https=False)
)
await repo.set_model(
ChatModel(
project_name="model-proj",
name="test-model",
created_at=datetime(2026, 1, 1),
run_name="same-run",
format_spec=(
OpenAIChatModelFormat(prefix="/v1")
if model_format == "openai"
else TGIChatModelFormat(chat_template="{{ messages }}", eos_token="</s>")
),
)
)
# Old gateways saved services without a proxy_buffering field.
state = json.loads(state_file.read_text())
for services in state["services"].values():
for service in services.values():
service.pop("proxy_buffering", None)
state_file.write_text(json.dumps(state))
nginx_dir = tmp_path / "nginx"
conf_dir = nginx_dir / "sites-enabled"
conf_dir.mkdir(parents=True)
# Restart twice to cover both migration and subsequent persisted-state recovery.
for _ in range(2):
repo = GatewayProxyRepo.load(state_file)
app = make_app(repo=repo, nginx=Nginx(nginx_dir=nginx_dir))
async with lifespan(app):
model_conf = (conf_dir / "443-model-proj.gtw.test.conf").read_text()
web_conf = (conf_dir / "443-web-proj.gtw.test.conf").read_text()
assert "proxy_buffering off;" in model_conf
assert "proxy_buffering off;" not in web_conf
persisted = json.loads(state_file.read_text())
assert persisted["services"]["model-proj"]["same-run"]["proxy_buffering"] is False
assert persisted["services"]["web-proj"]["same-run"]["proxy_buffering"] is True