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
1 change: 1 addition & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -123,6 +123,7 @@ uv run plugrl-run-env-client <ENVIRONMENT_TYPE> [OPTIONS]
| Type | Extra required | Description |
| :--- | :--- | :--- |
| **dummy-v1** | — | Dummy environment for protocol and connectivity tests |
| **probe-v1** | — | The probe environment of SPEC.md section 8.1, for `plugrl-conformance --probe` |
| **mujoco-v1** | `mujoco` | Gymnasium MuJoCo control, default `HalfCheetah-v5` |
| **classic-v1** | `classic` | Classic control environments (e.g. CartPole) |
| **atari-v1** | `atari` | Atari games via ALE |
Expand Down
162 changes: 99 additions & 63 deletions src/plugrl_env_client/agent/websocket_env_client_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,26 @@ class ServerStopped(RuntimeError):
"""Raised when the server explicitly requests env clients to stop."""


class ConnectionReplaced(RuntimeError):
"""The connection the caller's actions came from is gone.

The server keeps each environment's half of a transition - the previous
observation, the policy step state - on the connection that answered the
infer, so a new connection starts with none of it (SPEC section 7.6). Any
action chunk the caller is still executing was answered on the old one:
its feedback can never be completed, and must not be sent on the new
connection. The caller drops every chunk in flight and asks again for
all of its environments.

It used not to be told. The agent reconnected inside its next call and
carried on, so environments in the middle of a chunk kept executing it
and then sent its feedback on the new connection, and a feedback whose
send failed was followed by a reconnect inside the next `feedback` call,
whose first message was then a stale feedback. `plugrl-conformance
--probe --scenario resync` caught both.
"""


class WebSocketEnvClientAgent(_base_agent.BaseAgent):
def __init__(
self,
Expand All @@ -46,6 +66,11 @@ def __init__(
self._api_key = api_key
self._reconnect_on_server_stop = reconnect_on_server_stop
self._ws, self._server_metadata = self._wait_for_server()
# Which connection this is, and which one answered the caller's last
# infer. When they differ, the caller is holding actions from a
# connection that is gone; see ConnectionReplaced.
self._connection_id = 1
self._actions_from: int | None = None

def _close_connection(self) -> None:
if self._ws is None:
Expand Down Expand Up @@ -130,6 +155,12 @@ def _ensure_connection(self) -> None:
return
logger.warning("Connection closed. Attempting to re-establish connection.")
self._ws, self._server_metadata = self._wait_for_server()
self._connection_id += 1

def _replaced(self, why: str) -> ConnectionReplaced:
"""Tell the caller once; its next infer then goes out as normal."""
self._actions_from = None
return ConnectionReplaced(why)

def infer(
self,
Expand All @@ -143,6 +174,15 @@ def infer(
ws = self._ws
if ws is None:
raise RuntimeError("WebSocket connection is not available.")
if (
self._actions_from is not None
and self._actions_from != self._connection_id
):
# Nothing is sent: the caller has to re-plan every env first.
raise self._replaced(
"The connection was replaced; the action chunks in flight "
"came from the old one and are void."
)

try:
packed_data = self._packer.pack(
Expand All @@ -159,6 +199,7 @@ def infer(
if isinstance(response, str):
raise RuntimeError(f"Error in inference server:\n{response}")

self._actions_from = self._connection_id
return msgpack_numpy.unpackb(response)["data"]
except ConnectionClosedOK as exc:
close_code, close_reason = _get_close_details(exc)
Expand Down Expand Up @@ -220,71 +261,66 @@ def feedback(
makes the server build a transition out of an empty observation and
store it, with nothing downstream able to tell. SPEC section 7.6.

So a closed connection here costs exactly one transition, and that is
the cheap outcome. The next infer reconnects and resyncs.
So a closed connection here costs the transitions in flight, and that
is the cheap outcome. This method never reconnects: a feedback is only
ever sent on the connection that answered its infer. When that
connection is gone it raises ConnectionReplaced, so the caller drops
every chunk it is still executing, and the next infer reconnects.
"""
while True:
self._ensure_connection()
ws = self._ws
if ws is None:
raise RuntimeError("WebSocket connection is not available.")

try:
packed_data = self._packer.pack(
{
"message_type": str(MessageType.FEEDBACK),
"env_indices": env_indices,
"step_ids": step_ids,
"data": {
"obs": obs,
"rewards": rewards,
"terminated": terminated,
"truncated": truncated,
"info": info,
},
}
)
ws.send(packed_data)
return
except ConnectionClosedOK as exc:
close_code, close_reason = _get_close_details(exc)
self._close_connection()
ws = self._ws
if ws is None or self._actions_from != self._connection_id:
raise self._replaced(
"The connection that answered this feedback's infer is gone; "
"dropping the feedback."
)

if close_reason == SERVER_STOP_REASON:
if self._reconnect_on_server_stop:
logger.info(
"Server requested env client shutdown during FEEDBACK send, "
"but reconnect_on_server_stop is enabled. "
"Waiting for server to come back and retrying..."
)
continue
raise ServerStopped(
"Server requested env client shutdown after algorithm stop."
) from exc

if close_reason == SERVER_RESYNC_REASON:
logger.info(
"Server requested session resync during FEEDBACK send. "
"Dropping stale feedback and resuming from the next infer request."
)
return

logger.warning(
"Connection closed during FEEDBACK send. Dropping this "
"transition and resuming from the next infer request. "
f"code={close_code}, reason={close_reason or '<empty>'}"
)
return
except ConnectionClosedError as exc:
logger.warning(
"Connection closed during FEEDBACK send. Dropping this "
f"transition and resuming from the next infer request. {exc}"
)
self._close_connection()
return
except Exception:
self._close_connection()
raise
try:
packed_data = self._packer.pack(
{
"message_type": str(MessageType.FEEDBACK),
"env_indices": env_indices,
"step_ids": step_ids,
"data": {
"obs": obs,
"rewards": rewards,
"terminated": terminated,
"truncated": truncated,
"info": info,
},
}
)
ws.send(packed_data)
except ConnectionClosedOK as exc:
close_code, close_reason = _get_close_details(exc)
self._close_connection()

if (
close_reason == SERVER_STOP_REASON
and not self._reconnect_on_server_stop
):
raise ServerStopped(
"Server requested env client shutdown after algorithm stop."
) from exc
# A resync, a stop the client waits out, or any other close: the
# feedback is dropped either way. Before, a stop with
# reconnect_on_server_stop retried, which resent it on the new
# connection - the one path that broke SPEC section 7.6.
logger.info(
"Connection closed during FEEDBACK send. Dropping the "
"transitions in flight and resuming from the next infer request. "
f"code={close_code}, reason={close_reason or '<empty>'}"
)
raise self._replaced("dropped feedback after a close") from exc
except ConnectionClosedError as exc:
logger.warning(
"Connection closed during FEEDBACK send. Dropping the "
f"transitions in flight and resuming from the next infer request. {exc}"
)
self._close_connection()
raise self._replaced("dropped feedback after a close") from exc
except Exception:
self._close_connection()
raise

def reset(self) -> None:
return
106 changes: 106 additions & 0 deletions src/plugrl_env_client/envs/probe_env.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,106 @@
"""The probe environment of SPEC.md section 8.1, for `plugrl-conformance --probe`.

Env i counts the steps of its episode in `states["t"]`, keeps the action it
applied last in `states["a"]`, pays a reward of 1 per step, and ends its
episode, terminated, when t reaches 3 + 2 * (i % 3). It never truncates.

The conformance checker sends actions whose values encode where in the chunk
they are, so from each feedback it can tell what this client did with them:
whether it summed the chunk's reward, sent the terminal observation, and
applied the actions time-major and in order. It needs no extra:

plugrl-conformance --probe --scenario all --action-dim 3 --client \
"plugrl-run-env-client probe-v1 --server-host 127.0.0.1 --server-port 8000 \
--num-envs 3 --num-episodes 1000000"

`--env.action-dim` must match the checker's `--action-dim`.
"""

from __future__ import annotations

import dataclasses

import gymnasium as gym
import numpy as np

from plugrl_env_client.utils.registration import register_env, register_env_config

from .base_env import (
Action,
BaseEnv,
BaseEnvConfig,
BoolArray,
Observation,
RewardArray,
)

UID = "Probe-v1"


@register_env_config(UID)
@dataclasses.dataclass
class ProbeEnvConfig(BaseEnvConfig):
action_dim: int = 3


# No time limit: the probe env never truncates, and its episodes are 3 to 7
# steps long.
@register_env(UID, max_episode_steps=None)
class ProbeEnv(BaseEnv):
def __init__(
self,
config: ProbeEnvConfig,
num_envs: int = 1,
process_id: int | None = None,
total_processes: int | None = None,
):
super().__init__(
config=config,
num_envs=num_envs,
process_id=process_id,
total_processes=total_processes,
)
self.action_dim = config.action_dim
# float64, so the checker's values arrive exactly whatever they are.
self.single_action_space = gym.spaces.Box(
-np.inf, np.inf, (self.action_dim,), np.float64
)
self.action_space = self.single_action_space
# Env indices on the wire are the positions in this vector env.
self.lengths = 3 + 2 * (np.arange(num_envs) % 3)
self.t = np.zeros(num_envs, dtype=np.int64)
self.a = np.zeros((num_envs, self.action_dim), dtype=np.float64)

def _observation(self) -> Observation:
return Observation(
images={},
states={
"t": self.t[:, None].astype(np.float64),
"a": self.a.copy(),
},
text="probe",
)

def reset(self, *, seed: int | None = None, options: dict | None = None) -> tuple:
self.seed_rngs(seed)
indices = np.arange(self.num_envs)
if options is not None and options.get("reset_indices") is not None:
indices = np.asarray(options["reset_indices"], dtype=np.int64)
self.t[indices] = 0
self.a[indices] = 0.0
return self._observation(), {}

def step(
self, actions: Action
) -> tuple[Observation, RewardArray, BoolArray, BoolArray, dict]:
actions = np.asarray(actions, dtype=np.float64).reshape(
self.num_envs, self.action_dim
)
self.t += 1
self.a = actions.copy()
reward = np.ones(self.num_envs, dtype=np.float32)
terminated = self.t == self.lengths
truncated = np.zeros(self.num_envs, dtype=np.bool_)
# No reset here: the client sends this terminal observation, then
# resets the env itself (SPEC.md section 5.4).
return self._observation(), reward, terminated, truncated, {}
Loading
Loading