Skip to content
Open
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
120 changes: 118 additions & 2 deletions python/ray/serve/_private/long_poll.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,9 @@
from ray.serve._private.constants import (
DEFAULT_LATENCY_BUCKET_MS,
RAY_SERVE_COMPACT_LONG_POLL_METRIC_TAGS,
SERVE_CONTROLLER_NAME,
SERVE_LOGGER_NAME,
SERVE_NAMESPACE,
)
from ray.serve.generated.serve_pb2 import (
DeploymentTargetInfo,
Expand All @@ -53,6 +55,25 @@
float(os.environ.get("LISTEN_FOR_CHANGE_REQUEST_TIMEOUT_S_UPPER_BOUND", "60")),
)

# How long a client keeps trying to re-resolve a dead host before retiring
# itself. Bounded so an intentional `serve.shutdown()` still retires clients.
LONG_POLL_RECONNECT_TIMEOUT_S = float(
os.environ.get("RAY_SERVE_LONG_POLL_RECONNECT_TIMEOUT_S", "60")
)
LONG_POLL_RECONNECT_BACKOFF_S = (0.5, 8.0)


def _actor_id_str(host_actor: Any) -> str:
try:
return host_actor._actor_id.hex()
except AttributeError:
return "<unknown>"


def _resolve_serve_controller() -> Any:
"""Look up the current controller, which may be a replacement actor."""
return ray.get_actor(SERVE_CONTROLLER_NAME, namespace=SERVE_NAMESPACE)


class LongPollNamespace(Enum):
def __repr__(self):
Expand Down Expand Up @@ -119,6 +140,9 @@ class LongPollClient:
to post the callback into.
client_id: identifier reported back to the host if this client
disables itself.
host_actor_resolver: called to look the host up again after it dies.
Defaults to resolving the Serve controller, which every production
consumer polls; pass None to retire the client instead.
"""

def __init__(
Expand All @@ -127,6 +151,7 @@ def __init__(
key_listeners: Dict[KeyType, UpdateStateCallable],
call_in_event_loop: AbstractEventLoop,
client_id: str,
host_actor_resolver: Optional[Callable[[], Any]] = _resolve_serve_controller,
) -> None:
# We used to allow this to be optional, but due to Ray Client issue
# we now enforce all long poll client to post callback to event loop
Expand All @@ -144,6 +169,8 @@ def __init__(
self.key_listeners.keys(), -1
)
self.is_running = True
self._host_actor_resolver = host_actor_resolver
self._reconnect_task: Optional[asyncio.Task] = None
Comment thread
johntaylor-cell marked this conversation as resolved.

# Metric to track end-to-end latency from controller to client
self.long_poll_latency_histogram = metrics.Histogram(
Expand All @@ -166,6 +193,11 @@ def __init__(
def stop(self) -> None:
"""Stop the long poll client after the next RPC returns."""
self.is_running = False
# Otherwise a reconnect already in flight lingers for a backoff
# interval. Cancel on the loop thread, since callers may be elsewhere.
task = self._reconnect_task
if task is not None and not task.done() and self.event_loop.is_running():
self.event_loop.call_soon_threadsafe(task.cancel)

def add_key_listeners(
self, key_listeners: Dict[KeyType, UpdateStateCallable]
Expand Down Expand Up @@ -248,15 +280,99 @@ def _schedule_to_event_loop(self, callback):
f"{self.client_id!r} disabled itself."
)

def _process_update(self, updates: Dict[str, UpdatedObject]):
if isinstance(updates, (ray.exceptions.RayActorError)):
def _start_reconnect(self, error: ray.exceptions.RayActorError) -> None:
"""Begin re-resolving the host. Runs on the event loop."""
if not self.is_running:
return

if self._host_actor_resolver is None:
# This can happen during shutdown where the controller is
# intentionally killed, the client should just gracefully
# exit.
logger.debug("LongPollClient failed to connect to host. Shutting down.")
self.is_running = False
return

if self._reconnect_task is None or self._reconnect_task.done():
self._reconnect_task = self.event_loop.create_task(self._reconnect(error))

def _resolve_host_actor(self) -> Optional[Any]:
"""Resolve the host by name, or None while no replacement exists."""
resolver = self._host_actor_resolver
if resolver is None or not ray.is_initialized():
# Resolving auto-inits Ray, so never reach for it once it is gone.
return None

try:
return resolver()
except Exception:
return None

async def _reconnect(self, error: ray.exceptions.RayActorError) -> None:
"""Swap in a replacement host, or retire the client if none appears.

A replaced controller keeps its registered name but gets a new actor
ID, so the handle captured at construction is stale forever.
"""
logger.warning(
f"LongPollClient {self.client_id!r} lost its host "
f"{_actor_id_str(self.host_actor)}: {type(error).__name__}. "
"Trying to re-resolve it."
)
deadline = time.monotonic() + LONG_POLL_RECONNECT_TIMEOUT_S
backoff = LONG_POLL_RECONNECT_BACKOFF_S[0]
while self.is_running and time.monotonic() < deadline:
if not ray.is_initialized():
# This process shut Ray down, so there is nothing left to
# reconnect to and no one to deliver updates to.
self.is_running = False
return

# This blocks on the GCS, so it must not run on the event loop nor
# on the Ray callback thread that delivers the reply it waits for.
host_actor = await self.event_loop.run_in_executor(
None, self._resolve_host_actor
)
# The name can still point at the dead actor for a short window
# after it exits; rebinding to it would spin instead of recovering.
if (
self.is_running
and host_actor is not None
and _actor_id_str(host_actor) != _actor_id_str(self.host_actor)
):
self._rebind_host_actor(host_actor)
return

await asyncio.sleep(backoff)
backoff = min(backoff * 2, LONG_POLL_RECONNECT_BACKOFF_S[1])

if self.is_running:
logger.warning(
f"LongPollClient {self.client_id!r} could not re-resolve its host "
f"within {LONG_POLL_RECONNECT_TIMEOUT_S}s and has been disabled; "
"it will no longer receive updates."
)
self.is_running = False

def _rebind_host_actor(self, host_actor: Any) -> None:
"""Swap in the replacement host and ask it for a full state refresh.

Snapshot IDs are per-host counters seeded randomly, so IDs held for the
dead host mean nothing to its replacement; -1 forces a full resend.
"""
logger.warning(
f"LongPollClient {self.client_id!r} reconnected to host "
f"{_actor_id_str(host_actor)}, replacing {_actor_id_str(self.host_actor)}."
)
self.host_actor = host_actor
self.snapshot_ids = dict.fromkeys(self.key_listeners.keys(), -1)
self._poll_next()

def _process_update(self, updates: Dict[str, UpdatedObject]):
if isinstance(updates, (ray.exceptions.RayActorError)):
self._schedule_to_event_loop(lambda: self._start_reconnect(updates))
return

if isinstance(updates, ConnectionError):
logger.warning("LongPollClient connection failed, shutting down.")
self.is_running = False
Expand Down
173 changes: 173 additions & 0 deletions python/ray/serve/tests/test_long_poll.py
Original file line number Diff line number Diff line change
Expand Up @@ -659,5 +659,178 @@ def test_long_poll_client_disable_propagates_to_host_log():
assert "not running" in output.lower(), output


@pytest.fixture
def ray_initialized(monkeypatch):
"""These tests mock the host, so Ray itself is not running."""
monkeypatch.setattr(ray, "is_initialized", lambda: True)


def _reconnecting_client(resolver, host_actor=None, client_id="test_reconnect"):
return LongPollClient(
host_actor if host_actor is not None else MagicMock(),
{"key_1": lambda _: None},
call_in_event_loop=get_or_create_event_loop(),
client_id=client_id,
host_actor_resolver=resolver,
)


@pytest.mark.asyncio
async def test_client_reconnects_to_replacement_host(serve_instance):
"""A host replaced under the same name must not wedge the client forever.

See https://github.com/ray-project/ray/issues/63784.
"""

@ray.remote
class NamedHost:
def __init__(self, value):
self.host = LongPollHost()
self.host.notify_changed({"key_1": value})

async def listen_for_change(self, keys_to_snapshot_ids):
return await self.host.listen_for_change(keys_to_snapshot_ids)

name = "test_replacement_host"
host = NamedHost.options(name=name).remote(100)

received = {}
client = LongPollClient(
host,
{"key_1": lambda result: received.__setitem__("key_1", result)},
call_in_event_loop=get_or_create_event_loop(),
client_id="test_client_reconnects",
host_actor_resolver=lambda: ray.get_actor(name),
)
try:
await async_wait_for_condition(
lambda: received.get("key_1") == 100, timeout=20, retry_interval_ms=100
)

ray.kill(host, no_restart=True)

def name_released():
try:
ray.get_actor(name)
return False
except ValueError:
return True

await async_wait_for_condition(name_released, timeout=20, retry_interval_ms=100)
# Hold the handle: a named actor dies with its last reference.
replacement = NamedHost.options(name=name).remote(999)
assert replacement._actor_id != host._actor_id

await async_wait_for_condition(
lambda: received.get("key_1") == 999, timeout=60, retry_interval_ms=200
)
assert client.is_running
finally:
client.stop()


def test_rebind_host_actor_resets_snapshot_ids():
"""IDs from the dead host are meaningless to its replacement."""
client = _reconnecting_client(MagicMock())
client.snapshot_ids["key_1"] = 7

replacement = MagicMock()
client._rebind_host_actor(replacement)

assert client.host_actor is replacement
assert client.snapshot_ids == {"key_1": -1}


@pytest.mark.asyncio
async def test_reconnect_declines_stale_name_resolution(monkeypatch, ray_initialized):
"""A name still pointing at the dead actor is not a replacement."""
monkeypatch.setattr(long_poll_module, "LONG_POLL_RECONNECT_TIMEOUT_S", 0.2)
dead = MagicMock()
client = _reconnecting_client(lambda: dead, host_actor=dead)

client._process_update(ray.exceptions.ActorDiedError())

await async_wait_for_condition(
lambda: client.is_running is False, timeout=20, retry_interval_ms=100
)
assert client.host_actor is dead


@pytest.mark.asyncio
async def test_client_disables_itself_when_host_never_resolves(
monkeypatch, ray_initialized
):
"""Reconnection is bounded so an intentional shutdown still retires it."""
monkeypatch.setattr(long_poll_module, "LONG_POLL_RECONNECT_TIMEOUT_S", 0.2)

def resolver():
raise ValueError("no such actor")

client = _reconnecting_client(resolver)
client._process_update(ray.exceptions.ActorDiedError())

await async_wait_for_condition(
lambda: client.is_running is False, timeout=20, retry_interval_ms=100
)


@pytest.mark.asyncio
async def test_client_without_resolver_retires_on_host_death():
"""Opting out preserves the pre-existing shutdown behavior."""
client = _reconnecting_client(None)

client._process_update(ray.exceptions.ActorDiedError())

await async_wait_for_condition(
lambda: client.is_running is False, timeout=20, retry_interval_ms=100
)


@pytest.mark.asyncio
async def test_stopped_client_does_not_reconnect():
"""A client stopped on purpose must not chase a replacement host."""
resolver = MagicMock()
client = _reconnecting_client(resolver)
client.stop()

client._process_update(ray.exceptions.ActorDiedError())

await asyncio.sleep(1)
resolver.assert_not_called()


@pytest.mark.asyncio
async def test_stop_cancels_in_flight_reconnect(ray_initialized):
"""A reconnect already in flight must not outlive stop()."""
dead = MagicMock()
client = _reconnecting_client(lambda: dead, host_actor=dead)
client._process_update(ray.exceptions.ActorDiedError())
await async_wait_for_condition(
lambda: client._reconnect_task is not None, timeout=20, retry_interval_ms=50
)

client.stop()

await async_wait_for_condition(
lambda: client._reconnect_task.done(), timeout=20, retry_interval_ms=50
)
assert client._reconnect_task.cancelled()


@pytest.mark.asyncio
async def test_reconnect_retires_when_ray_is_shut_down(monkeypatch):
"""Resolving auto-inits Ray, which would revive a process that shut it down."""
monkeypatch.setattr(ray, "is_initialized", lambda: False)
resolver = MagicMock()
client = _reconnecting_client(resolver)

client._process_update(ray.exceptions.ActorDiedError())

await async_wait_for_condition(
lambda: client.is_running is False, timeout=20, retry_interval_ms=50
)
resolver.assert_not_called()


if __name__ == "__main__":
sys.exit(pytest.main(["-v", "-s", __file__]))
Loading