Skip to content
Draft
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
128 changes: 128 additions & 0 deletions areal/api/openenv_api.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,128 @@
# SPDX-License-Identifier: Apache-2.0
"""Config dataclasses and Protocols for the OpenEnv adapter.

OpenEnv (https://github.com/huggingface/OpenEnv) exposes agentic environments
via a uniform reset/step/state HTTP+WebSocket surface. This module defines the
configuration surface for AReaL's ``OpenEnvWorkflow`` and the pluggable
adapter Protocols (action parser + observation formatter) that let a workflow
target arbitrary OpenEnv environments without new code.
"""

from __future__ import annotations

from dataclasses import dataclass, field
from typing import Any, Literal, Protocol


class ObservationFormatter(Protocol):
"""Convert an OpenEnv observation into a chat message dict.

Implementations must be stateless with respect to a single episode; per-step
state should live on ``OpenEnvWorkflow`` or in the observation itself.
"""

def __call__(self, observation: Any, step: int) -> dict[str, str]:
"""Return a ``{"role": ..., "content": ...}`` message for the LLM."""
...


class ActionParser(Protocol):
"""Convert an LLM completion string into an OpenEnv action object.

Returning ``None`` marks the completion as unparsable and the workflow
treats the step as a failed no-op with zero reward.
"""

def __call__(self, completion: str, observation: Any) -> Any:
"""Return the parsed action, or ``None`` on parse failure."""
...


@dataclass
class OpenEnvConfig:
"""Configuration for :class:`areal.workflow.openenv.OpenEnvWorkflow`.

Attributes:
env_client_class: Fully-qualified import path to an
:class:`openenv.core.env_client.EnvClient` subclass or a factory
callable that returns one. Example: ``"echo_env.EchoEnv"``.
base_url: Base URL of a running environment server (``http://`` or
``ws://``). Mutually exclusive with ``provider``. Set exactly one.
provider: Provider spec for launching the environment locally. One of
``"uv"`` (uses ``UVProvider``, no Docker) or ``"docker"``
(``LocalDockerProvider``). Leave unset to require ``base_url``.
project_path: Project source for ``provider="uv"``. Either a local path
or a ``git+<url>`` spec, forwarded to ``UVProvider``.
docker_image: Image tag for ``provider="docker"``. Ignored otherwise.
action_class: Optional fully-qualified import path to the Action
dataclass expected by ``env.step`` (e.g. ``"echo_env.CallToolAction"``).
When set, ``ActionParser`` output is passed as ``action_class(**parsed)``
if the parser returns a dict; otherwise passed through as-is.
action_parser: Import path to an :class:`ActionParser` implementation,
or a shorthand: ``"json"`` (default), ``"tag"``
(``<action>...</action>``), or ``"passthrough"`` (raw string).
obs_formatter: Import path to an :class:`ObservationFormatter`, or the
shorthand ``"auto"`` (dataclass/dict → JSON, str → identity).
system_prompt: Optional system message prepended to every episode.
max_turns: Hard cap on env steps per episode.
step_discount: Per-step reward discount (``reward *= step_discount``
before accumulation). ``1.0`` disables discounting.
terminal_reward_only: When ``True``, only the last step's reward is
kept; intermediate rewards are discarded (episode-level scoring).
reset_kwargs: Keyword args forwarded to ``env.reset`` (e.g. ``{"seed": 0}``).
connect_timeout_s: WebSocket connect timeout.
message_timeout_s: WebSocket message timeout.
"""

env_client_class: str
base_url: str | None = None
provider: Literal["uv", "docker"] | None = None
project_path: str | None = None
docker_image: str | None = None
action_class: str | None = None
action_parser: str = "json"
obs_formatter: str = "auto"
system_prompt: str = ""
max_turns: int = 8
step_discount: float = 1.0
terminal_reward_only: bool = False
reset_kwargs: dict[str, Any] = field(default_factory=dict)
connect_timeout_s: float = 30.0
message_timeout_s: float = 60.0

def __post_init__(self) -> None:
if self.max_turns <= 0:
raise ValueError(f"max_turns must be positive, got {self.max_turns}")
if not 0.0 < self.step_discount <= 1.0:
raise ValueError(
f"step_discount must be in (0, 1], got {self.step_discount}"
)
if self.base_url is None and self.provider is None:
raise ValueError(
"OpenEnvConfig requires either base_url or provider to be set."
)
if self.base_url is not None and self.provider is not None:
raise ValueError(
"OpenEnvConfig.base_url and OpenEnvConfig.provider are "
"mutually exclusive; set exactly one."
)
if self.provider is not None and self.provider not in ("uv", "docker"):
raise ValueError(
f"OpenEnvConfig.provider must be 'uv' or 'docker', "
f"got {self.provider!r}"
)
if self.provider == "uv" and self.project_path is None:
raise ValueError(
"OpenEnvConfig.provider='uv' requires project_path to be set."
)
if self.provider == "docker" and self.docker_image is None:
raise ValueError(
"OpenEnvConfig.provider='docker' requires docker_image to be set."
)


__all__ = [
"ActionParser",
"ObservationFormatter",
"OpenEnvConfig",
]
1 change: 1 addition & 0 deletions areal/utils/logging.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,7 @@
"VisionRLVRWorkflow": "light_purple",
"MultiTurnWorkflow": "light_purple",
"MultiTurnV2Workflow": "light_purple",
"OpenEnvWorkflow": "light_purple",
# Controllers - white
"TrainController": "white",
"RolloutController": "white",
Expand Down
2 changes: 2 additions & 0 deletions areal/workflow/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,12 +3,14 @@
__all__ = [
"RLVRWorkflow",
"MultiTurnWorkflow",
"OpenEnvWorkflow",
"VisionRLVRWorkflow",
]

_LAZY_IMPORTS = {
"RLVRWorkflow": "areal.workflow.rlvr",
"MultiTurnWorkflow": "areal.workflow.multi_turn",
"OpenEnvWorkflow": "areal.workflow.openenv",
"VisionRLVRWorkflow": "areal.workflow.vision_rlvr",
}

Expand Down
247 changes: 247 additions & 0 deletions areal/workflow/openenv.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,247 @@
# SPDX-License-Identifier: Apache-2.0
"""Rollout workflow that drives any OpenEnv-compatible environment.

OpenEnv (https://github.com/huggingface/OpenEnv) is HuggingFace's uniform
adapter layer over agentic environments -- BrowserGym, OpenSpiel, Coding,
BlackJack, Terminal-Bench, etc. Each environment exposes
``reset() / step(action) / state()`` over WebSocket.

This workflow lets AReaL train against any such environment by pointing at an
EnvClient subclass and (optionally) an Action dataclass via config. No new
Python file is needed to plug in a new environment; only a YAML change.
"""

from __future__ import annotations

from copy import deepcopy
from typing import TYPE_CHECKING, Any

from areal import workflow_context
from areal.api import RolloutWorkflow
from areal.api.cli_args import GenerationHyperparameters
from areal.api.openenv_api import OpenEnvConfig
from areal.experimental.openai import ArealOpenAI
from areal.utils import logging, stats_tracker
from areal.workflow.openenv_utils import (
_import_from_string,
build_action,
resolve_action_parser,
resolve_obs_formatter,
)

if TYPE_CHECKING:
from transformers import PreTrainedTokenizerFast

from areal.api import InferenceEngine

logger = logging.getLogger("OpenEnvWorkflow")


def _instantiate_provider(cfg: OpenEnvConfig) -> Any:
"""Lazily import the OpenEnv provider requested by ``cfg``."""
if cfg.provider == "uv":
from openenv.core.containers.runtime.uv_provider import UVProvider

return UVProvider(project_path=cfg.project_path)
if cfg.provider == "docker":
from openenv.core.containers.runtime import LocalDockerProvider

return LocalDockerProvider(image=cfg.docker_image)
raise ValueError(f"Unknown OpenEnv provider: {cfg.provider!r}")


def _instantiate_env_client(cfg: OpenEnvConfig) -> Any:
"""Build the concrete ``EnvClient`` from ``cfg``.

Deferred import keeps ``openenv`` an optional dependency: users who never
touch this workflow do not need it installed.
"""
env_client_cls = _import_from_string(cfg.env_client_class)
kwargs: dict[str, Any] = {
"connect_timeout_s": cfg.connect_timeout_s,
"message_timeout_s": cfg.message_timeout_s,
}
if cfg.base_url is not None:
kwargs["base_url"] = cfg.base_url
else:
kwargs["provider"] = _instantiate_provider(cfg)
return env_client_cls(**kwargs)


class OpenEnvWorkflow(RolloutWorkflow):
"""Drive one episode against an OpenEnv environment.

The loop is: ``reset() -> [chat.completion -> parse action -> env.step()] *
N -> aggregate reward``. Each LLM turn is cached by :class:`ArealOpenAI`
with the environment reward for that step, so the exported trajectory
carries per-step supervision suitable for GRPO / PPO / RLOO.

Parameters
----------
config
:class:`OpenEnvConfig` describing which environment to launch and how
to convert observations/actions.
gconfig
Standard AReaL generation hyperparameters.
tokenizer
Tokenizer used by the underlying inference engine.
initial_user_prompt
Optional prompt string prepended before the first observation. When
``None``, only the observation-derived user message is sent.
reward_shaping_fn
Optional callable ``(step_result, step, episode_data) -> float`` that
overrides the raw environment reward. Return the value that should be
recorded for this step. Useful for cost/length penalties.
"""

def __init__(
self,
config: OpenEnvConfig,
gconfig: GenerationHyperparameters,
tokenizer: PreTrainedTokenizerFast | str,
initial_user_prompt: str | None = None,
reward_shaping_fn: Any = None,
) -> None:
self.config = config
if isinstance(tokenizer, str):
from areal.utils.hf_utils import load_hf_tokenizer

tokenizer = load_hf_tokenizer(tokenizer)
# Grouped rollout is external; each workflow instance produces one trajectory.
self.gconfig = gconfig.new_with_stop_and_pad_token_ids(tokenizer).new(
n_samples=1
)
self.tokenizer = tokenizer
self.initial_user_prompt = initial_user_prompt
self.reward_shaping_fn = reward_shaping_fn

self._obs_formatter = resolve_obs_formatter(config.obs_formatter)
self._action_parser = resolve_action_parser(config.action_parser)

def _initial_messages(self, observation: Any) -> list[dict[str, str]]:
messages: list[dict[str, str]] = []
if self.config.system_prompt:
messages.append({"role": "system", "content": self.config.system_prompt})
if self.initial_user_prompt:
messages.append({"role": "user", "content": self.initial_user_prompt})
messages.append(self._obs_formatter(observation, step=0))
return messages

async def _generate_step(
self, client: ArealOpenAI, messages: list[dict[str, str]]
) -> Any:
return await client.chat.completions.create( # type: ignore[arg-type]
messages=messages,
frequency_penalty=self.gconfig.frequency_penalty,
max_completion_tokens=self.gconfig.max_new_tokens,
stop=self.gconfig.stop,
store=True,
temperature=self.gconfig.temperature,
top_p=self.gconfig.top_p,
)

async def arun_episode(
self, engine: InferenceEngine, data: dict[str, Any]
) -> dict[str, Any]:
"""Run one full episode against the configured OpenEnv environment.

``data`` is forwarded verbatim to the tokenizer; recognized keys:

* ``seed``: passed to ``env.reset(seed=...)``. Overrides any
``seed`` inside ``config.reset_kwargs``.
* ``system_prompt``: overrides ``config.system_prompt`` for this
episode. Empty strings are ignored (workflow default is kept).
"""
openai_client = ArealOpenAI(engine=engine, tokenizer=self.tokenizer)
env_client = _instantiate_env_client(self.config)

# Per-episode data.seed is authoritative over the workflow-level
# default; otherwise the entire batch would share the same env init.
reset_kwargs = dict(self.config.reset_kwargs)
if "seed" in data:
reset_kwargs["seed"] = data["seed"]

step_rewards: list[float] = []
completion_ids: list[str] = []
terminal_done = False
parse_failed = False

async with env_client as env:
result = await env.reset(**reset_kwargs)
messages = self._initial_messages(result.observation)
if data.get("system_prompt"):
# Replace any existing system message with the per-episode override.
# Empty strings fall through so the workflow default is preserved.
messages = [m for m in messages if m.get("role") != "system"]
messages.insert(0, {"role": "system", "content": data["system_prompt"]})

for step in range(self.config.max_turns):
completion = await self._generate_step(openai_client, messages)
assistant_text = completion.choices[0].message.content or ""

parsed = self._action_parser(assistant_text, result.observation)
if parsed is None:
logger.debug(
f"Action parser returned None on step {step}; "
"recording the completion with zero reward and ending "
"the episode."
)
parse_failed = True
# Record the completion with zero reward but keep it OUT of
# step_rewards / completion_ids so the trajectory bookkeeping
# below sees only real environment interactions.
openai_client.set_reward(completion.id, 0.0)
break

action = build_action(parsed, self.config.action_class)
result = await env.step(action)

if self.reward_shaping_fn is not None:
step_reward = float(self.reward_shaping_fn(result, step, data))
else:
step_reward = float(result.reward or 0.0)
completion_ids.append(completion.id)
step_rewards.append(step_reward)
openai_client.set_reward(completion.id, step_reward)

# Append assistant + next observation for the following turn.
messages = deepcopy(messages)
messages.append({"role": "assistant", "content": assistant_text})
if not result.done and step + 1 < self.config.max_turns:
messages.append(
self._obs_formatter(result.observation, step=step + 1)
)

if result.done:
terminal_done = True
break

# Trajectory-level bookkeeping.
if self.config.terminal_reward_only and completion_ids:
# Zero out intermediate rewards; keep the last one as episode reward.
# Skipping the discount below is intentional: back-propagating from
# a lone terminal reward would smear non-zero credit onto the
# 'discarded' intermediate turns, contradicting the doc-promise.
terminal_reward = step_rewards[-1] if step_rewards else 0.0
for cid in completion_ids[:-1]:
openai_client.set_reward(cid, 0.0)
openai_client.set_reward(completion_ids[-1], terminal_reward)
elif self.config.step_discount < 1.0:
# Apply per-step discount by backward propagation. Gated off when
# terminal_reward_only is set, per the semantic conflict above.
openai_client.apply_reward_discount(self.config.step_discount)

# Log the raw environment reward; the trainer sees whatever the cache
# holds after the terminal/discount rewrites above, but for
# observability the un-shaped sum is more actionable.
episode_reward = float(sum(step_rewards))
stats_tracker.get(workflow_context.stat_scope()).scalar(
reward=episode_reward,
num_turns=len(step_rewards),
terminated=float(terminal_done),
parse_failed=float(parse_failed),
)
return openai_client.export_interactions("individual")


__all__ = ["OpenEnvWorkflow"]
Loading
Loading