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
56 changes: 55 additions & 1 deletion envs/textarena_env/server/environment.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,12 +8,19 @@

from __future__ import annotations

import random
import sys
import threading
from typing import Any, Dict, Iterable, List, Optional
from uuid import uuid4

from openenv.core.env_server.interfaces import Environment

try:
import numpy as _np
except Exception: # pragma: no cover - numpy is optional
_np = None

try:
# When running as installed package
from textarena_env.models import (
Expand All @@ -38,6 +45,11 @@
_TEXTARENA_IMPORT_ERROR: Exception | None = None
_NLTK_DOWNLOADED: bool = False

# TextArena selects episodes/words via the process-global ``random`` (and
# sometimes ``numpy``) RNGs. Seeding+selection must be atomic across concurrent
# sessions, so guard it with a process-wide lock.
_SEED_LOCK = threading.Lock()


def _ensure_nltk_data() -> None:
"""Download NLTK data once per process."""
Expand Down Expand Up @@ -147,7 +159,7 @@ def reset(
if hasattr(env, "full_observations"):
env.full_observations = {}

self._ta_env.reset(num_players=self.num_players)
self._seeded_reset(seed)

for provider in self._reward_providers:
provider.reset()
Expand Down Expand Up @@ -203,6 +215,48 @@ def step(self, action: TextArenaAction) -> TextArenaObservation: # type: ignore
def state(self) -> TextArenaState:
return self._state

# ------------------------------------------------------------------
# Seeding
# ------------------------------------------------------------------
def _seeded_reset(self, seed: Optional[int]) -> None:
"""Reset the underlying TextArena env, honoring ``seed`` for reproducibility.

TextArena chooses episodes/words with the process-global ``random`` (and
sometimes ``numpy``) RNGs and does *not* apply ``reset(seed=...)`` to that
selection. So we seed the global RNGs ourselves immediately before the
underlying reset. This makes episodes reproducible, which is required, for
example, for GRPO-style training where every rollout in a group must share
the same episode (same Wordle word) for the group baseline to be valid.

A process-wide lock keeps seed+selection atomic across concurrent sessions,
and we restore the prior RNG state afterwards so that unseeded sessions
sharing the process are left undisturbed.
"""
if seed is None:

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

seed=None bypasses _SEED_LOCK, so an unseeded reset can run after another thread calls random.seed(seed) but before it selects/restores. That consumes the seeded stream, perturbs the unseeded session, and can change the seeded episode. Hold the same lock around every underlying reset; for seeded calls, save/seed/reset/restore while holding it. Also, TextArena 0.7.4 already forwards seed through SinglePlayerState.__init__ to State.__init__, which calls random.seed(seed), so update the docstring’s contrary claim.

self._ta_reset(seed)
return

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Unseeded reset skips seed lock

Medium Severity

_seeded_reset only acquires _SEED_LOCK when seed is set, so an unseeded reset (and __init__'s direct _ta_env.reset) can run while another session has temporarily reseeded the process-global RNGs. With SUPPORTS_CONCURRENT_SESSIONS and a thread-pool server, that interleaving can steal draws from a seeded episode or break same-seed reproducibility.

Additional Locations (1)
Fix in Cursor Fix in Web

Reviewed by Cursor Bugbot for commit 58723a9. Configure here.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Unseeded reset skips seed lock

Medium Severity

Unseeded _seeded_reset calls _ta_reset without _SEED_LOCK, so another session can consume the process-global random stream while a seeded reset holds the hijacked RNG. That breaks the atomic seed-plus-selection guarantee and can assign the wrong episode under concurrent sessions.

Fix in Cursor Fix in Web

Reviewed by Cursor Bugbot for commit 2131c0f. Configure here.


with _SEED_LOCK:
py_state = random.getstate()
np_state = _np.random.get_state() if _np is not None else None
try:
random.seed(seed)
if _np is not None:
_np.random.seed(seed)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Large seeds crash with numpy

Low Severity

When numpy is importable, _seeded_reset always calls numpy.random.seed, which rejects integers outside 0..2**32-1. A valid OpenEnv seed that random.seed accepts then raises ValueError and aborts reset, even for games that only use the Python random module.

Fix in Cursor Fix in Web

Reviewed by Cursor Bugbot for commit 2131c0f. Configure here.

self._ta_reset(seed)
finally:
random.setstate(py_state)
if np_state is not None:
_np.random.set_state(np_state)

def _ta_reset(self, seed: Optional[int]) -> None:
"""Call the underlying TextArena reset, forwarding ``seed`` when supported."""
try:
self._ta_env.reset(num_players=self.num_players, seed=seed)
except TypeError:
# Older/other TextArena games whose reset does not accept a seed.
self._ta_env.reset(num_players=self.num_players)

# ------------------------------------------------------------------
# Helpers
# ------------------------------------------------------------------
Expand Down
79 changes: 79 additions & 0 deletions tests/envs/test_textarena_seed.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,79 @@
# SPDX-License-Identifier: BSD-3-Clause

"""Reproducibility tests for TextArena environment seed handling.

TextArena chooses its episode/word with the global ``random`` RNG and does not
apply ``reset(seed=...)`` to that selection, so the OpenEnv wrapper seeds the
global RNGs itself. These tests pin that behaviour down: same seed -> same word,
seed drives selection, and seeding a reset does not disturb the global RNG stream.
"""

import random

import pytest

# Skip the whole module if the optional textarena dependency is not installed.
pytest.importorskip("textarena", reason="textarena is not installed")

from envs.textarena_env.server.environment import TextArenaEnvironment


def _secret_word(env: TextArenaEnvironment):
"""Best-effort extraction of the hidden Wordle word from the underlying env."""
inner = env._ta_env
while inner is not None:
game_state = getattr(getattr(inner, "state", None), "game_state", None)
if isinstance(game_state, dict) and "secret_word" in game_state:
return game_state["secret_word"]
inner = getattr(inner, "env", None)
return None


@pytest.fixture(scope="module")
def env():
return TextArenaEnvironment(env_id="Wordle-v0", num_players=1)


def test_same_seed_same_word(env):
env.reset(seed=123)
first = _secret_word(env)
env.reset(seed=123)
second = _secret_word(env)
assert first is not None
assert first == second


def test_same_seed_across_instances():
a = TextArenaEnvironment(env_id="Wordle-v0", num_players=1)
b = TextArenaEnvironment(env_id="Wordle-v0", num_players=1)
a.reset(seed=7)
b.reset(seed=7)
assert _secret_word(a) is not None
assert _secret_word(a) == _secret_word(b)


def test_seed_drives_word_selection(env):
words = set()
for seed in range(12):
env.reset(seed=seed)
words.add(_secret_word(env))
# Combined with the determinism tests above, more than one distinct word
# confirms the seed actually drives selection rather than being ignored.
assert len(words) > 1


def test_unseeded_reset_still_works(env):
observation = env.reset()
assert observation.done is False
assert _secret_word(env) is not None


def test_seed_restores_global_rng_state(env):
random.seed(999)
baseline = random.random()

random.seed(999)
env.reset(seed=1) # a seeded reset must not disturb the global RNG stream
after = random.random()

assert baseline == after
Loading