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
113 changes: 101 additions & 12 deletions src/agents/sandbox/session/dependencies.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from __future__ import annotations

import asyncio
import inspect
from collections.abc import Awaitable, Callable, Mapping
from dataclasses import dataclass
Expand Down Expand Up @@ -73,7 +74,11 @@ class Dependencies:
def __init__(self) -> None:
self._bindings: dict[DependencyKey, _Binding] = {}
self._cache: dict[DependencyKey, object] = {}
self._pending: dict[DependencyKey, asyncio.Task[object]] = {}
self._active_tasks: set[asyncio.Task[object]] = set()
self._cleanup_tasks: set[asyncio.Task[None]] = set()
self._owned_results: list[object] = []
self._close_task: asyncio.Task[None] | None = None
self._closed = False

@classmethod
Expand Down Expand Up @@ -144,6 +149,9 @@ def _bind(
raise DependenciesBindingError(f"Dependency `{key}` is already bound")
self._bindings[key] = binding
self._cache.pop(key, None)
pending = self._pending.pop(key, None)
if pending is not None:
pending.cancel()

async def get(self, key: DependencyKey) -> object | None:
binding = self._bindings.get(key)
Expand Down Expand Up @@ -173,24 +181,96 @@ async def _resolve(self, key: DependencyKey, binding: _Binding) -> object:
return binding.value

assert isinstance(binding, _FactoryBinding)
if self._closed:
raise DependenciesError(f"Dependencies container is closed; cannot resolve `{key}`")
if binding.cache and key in self._cache:
return self._cache[key]

produced = binding.factory(self)
value = (
await cast(Awaitable[object], produced) if inspect.isawaitable(produced) else produced
)

if binding.cache:
self._cache[key] = value
if binding.owns_result:
self._owned_results.append(value)
return value
task = self._pending.get(key)
if task is None:
task = self._create_factory_task(key, binding)
self._pending[key] = task
return await asyncio.shield(task)

task = self._create_factory_task(key, binding)
return await task

def _create_factory_task(
self, key: DependencyKey, binding: _FactoryBinding
) -> asyncio.Task[object]:
task = asyncio.create_task(self._run_factory(key, binding))
self._active_tasks.add(task)
task.add_done_callback(_consume_task_exception)
return task

async def _run_factory(self, key: DependencyKey, binding: _FactoryBinding) -> object:
try:
produced = binding.factory(self)
value = (
await cast(Awaitable[object], produced)
if inspect.isawaitable(produced)
else produced
)

if self._closed:
if binding.owns_result:
await self._discard_factory_result(value)
raise DependenciesError(f"Dependencies container closed while resolving `{key}`")

if self._bindings.get(key) is not binding:
if binding.owns_result:
await self._discard_factory_result(value)
raise DependenciesBindingError(
f"Dependency `{key}` was rebound while its factory was resolving"
)

if binding.cache:
self._cache[key] = value
if binding.owns_result:
self._owned_results.append(value)
return value
except asyncio.CancelledError:
if self._closed:
raise DependenciesError(
f"Dependencies container closed while resolving `{key}`"
) from None
if self._bindings.get(key) is not binding:
raise DependenciesBindingError(
f"Dependency `{key}` was rebound while its factory was resolving"
) from None
raise
finally:
task = asyncio.current_task()
if task is not None:
self._active_tasks.discard(task)
if self._pending.get(key) is task:
self._pending.pop(key, None)

async def _discard_factory_result(self, value: object) -> None:
task = asyncio.create_task(_close_best_effort(value))
self._cleanup_tasks.add(task)
task.add_done_callback(self._cleanup_tasks.discard)
await asyncio.shield(task)

async def aclose(self) -> None:
if self._closed:
return
self._closed = True
task = self._close_task
if task is None:
self._closed = True
task = asyncio.create_task(self._close())
self._close_task = task
await asyncio.shield(task)

async def _close(self) -> None:
active_tasks = tuple(self._active_tasks)
for task in active_tasks:
task.cancel()
if active_tasks:
await asyncio.gather(*active_tasks, return_exceptions=True)

while self._cleanup_tasks:
cleanup_tasks = tuple(self._cleanup_tasks)
await asyncio.gather(*cleanup_tasks, return_exceptions=True)

seen_ids: set[int] = set()
for value in reversed(self._owned_results):
Expand All @@ -199,3 +279,12 @@ async def aclose(self) -> None:
continue
seen_ids.add(value_id)
await _close_best_effort(value)

self._pending.clear()
self._cache.clear()
self._owned_results.clear()


def _consume_task_exception(task: asyncio.Task[object]) -> None:
if not task.cancelled():
task.exception()
Loading