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
119 changes: 76 additions & 43 deletions src/openenv/core/env_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -359,6 +359,8 @@ def __init__(
self._websocket_ping_interval_s = websocket_ping_interval_s
self._websocket_ping_timeout_s = websocket_ping_timeout_s
self._provider = provider
self._provider_stopped = False
self._provider_cleanup_pending = False
self._start_provider_on_connect = base_url is None
self._child_clients: list[EnvClient[Any, Any, Any]] = []
self._ws: Optional[ClientConnection] = None
Expand Down Expand Up @@ -391,6 +393,16 @@ def _set_base_url(self, base_url: str) -> None:
self._ws_url = f"{ws_url}/ws"

def _start_provider_if_needed(self) -> None:
# A missing URL does not mean the previous resource was released.
# Retry its cleanup before start can overwrite the provider's handle.
if self._provider_cleanup_pending:
if not self._start_provider_on_connect:
raise RuntimeError(
"Provider cleanup is pending for this client with an existing "
"base URL. Retry close() to finish cleanup, then create a new "
"client instead of reconnecting to the old URL."
)
self._stop_provider()
Comment thread
cursor[bot] marked this conversation as resolved.
if self._ws_url is not None:
return
Comment thread
cursor[bot] marked this conversation as resolved.
if self._provider is None:
Expand All @@ -405,41 +417,50 @@ def _start_provider_if_needed(self) -> None:
f"{required}. Start the provider manually and pass base_url, "
"or configure a provider with a constructor-owned image/source."
)
self._provider_stopped = False
base_url = self._provider.start_container()
self._provider.wait_for_ready(base_url)
elif hasattr(self._provider, "start"):
self._provider_stopped = False
base_url = self._provider.start()
self._provider.wait_for_ready()
else:
raise TypeError("provider must define start_container() or start().")
self._set_base_url(base_url)

def _create_session_client(self) -> "EnvClient[Any, Any, Any]":
self._start_provider_if_needed()
if self._base_url is None:
raise RuntimeError("EnvClient has no base URL.")

signature = inspect.signature(type(self))
accepts_kwargs = any(
parameter.kind == inspect.Parameter.VAR_KEYWORD
for parameter in signature.parameters.values()
)
candidate_kwargs = {
"base_url": self._base_url,
"connect_timeout_s": self._connect_timeout,
"message_timeout_s": self._message_timeout,
"max_message_size_mb": self._max_message_size / (1024 * 1024),
"websocket_ping_interval_s": self._websocket_ping_interval_s,
"websocket_ping_timeout_s": self._websocket_ping_timeout_s,
"mode": self._mode,
}
constructor_kwargs = {}
for name, value in candidate_kwargs.items():
if accepts_kwargs or name in signature.parameters:
constructor_kwargs[name] = value
# Match _start_provider_if_needed's startup condition. A provider with
# an existing URL may already be serving the parent or other children.
starts_provider_here = self._provider is not None and self._ws_url is None
try:
self._start_provider_if_needed()
if self._base_url is None:
raise RuntimeError("EnvClient has no base URL.")

client = type(self)(**constructor_kwargs)
return client
signature = inspect.signature(type(self))
accepts_kwargs = any(
parameter.kind == inspect.Parameter.VAR_KEYWORD
for parameter in signature.parameters.values()
)
candidate_kwargs = {
"base_url": self._base_url,
"connect_timeout_s": self._connect_timeout,
"message_timeout_s": self._message_timeout,
"max_message_size_mb": self._max_message_size / (1024 * 1024),
"websocket_ping_interval_s": self._websocket_ping_interval_s,
"websocket_ping_timeout_s": self._websocket_ping_timeout_s,
"mode": self._mode,
}
constructor_kwargs = {}
for name, value in candidate_kwargs.items():
if accepts_kwargs or name in signature.parameters:
constructor_kwargs[name] = value

return type(self)(**constructor_kwargs)
except Exception:
if starts_provider_here and not self._provider_cleanup_pending:
self._stop_provider_best_effort()
raise
Comment thread
cursor[bot] marked this conversation as resolved.

async def new_session(self) -> "EnvClient[Any, Any, Any]":
"""
Expand Down Expand Up @@ -557,7 +578,11 @@ async def _connect_async(self) -> "EnvClient":
try:
self._start_provider_if_needed()
except Exception:
await self.close()
# A failed cleanup retry must propagate without attempting stop
# again. For a new startup failure, preserve its original error.
if not self._provider_cleanup_pending:
with suppress(Exception):
await self.close()
Comment thread
cursor[bot] marked this conversation as resolved.
raise

assert self._ws_url is not None
Expand Down Expand Up @@ -1059,21 +1084,35 @@ async def _close_async(self) -> None:
if deferred_cancellation is None:
deferred_cancellation = exc
finally:
try:
if self._provider is not None:
# Handle both ContainerProvider and RuntimeProvider
if hasattr(self._provider, "stop_container"):
self._provider.stop_container()
elif hasattr(self._provider, "stop"):
self._provider.stop()
finally:
if self._start_provider_on_connect:
self._base_url = None
self._ws_url = None
self._stop_provider()

if deferred_cancellation is not None:
raise deferred_cancellation

def _stop_provider(self) -> None:
"""Stop the provider once, retaining it for retries and later startup."""
try:
provider = self._provider
if provider is None or self._provider_stopped:
return

if hasattr(provider, "stop_container"):
stop = provider.stop_container
elif hasattr(provider, "stop"):
stop = provider.stop
else:
return

self._provider_cleanup_pending = True
stop()
# Only a successful stop discharges our cleanup responsibility.
self._provider_cleanup_pending = False
self._provider_stopped = True
finally:
if self._start_provider_on_connect:
self._base_url = None
self._ws_url = None
Comment thread
cursor[bot] marked this conversation as resolved.
Comment thread
cursor[bot] marked this conversation as resolved.

def _stop_provider_best_effort(self) -> None:
"""Stop the underlying provider directly, ignoring any errors.

Expand All @@ -1082,14 +1121,8 @@ def _stop_provider_best_effort(self) -> None:
after the provider started but before the connection is established, so
routing cleanup through the (possibly broken) sync loop is not an option.
"""
provider = self._provider
if provider is None:
return
with suppress(Exception):
if hasattr(provider, "stop_container"):
provider.stop_container()
elif hasattr(provider, "stop"):
provider.stop()
self._stop_provider()

async def __aenter__(self) -> "EnvClient":
"""Enter async context manager, ensuring connection is established."""
Expand Down
Loading
Loading