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
39 changes: 29 additions & 10 deletions openhands/app_server/sandbox/remote_sandbox_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@
from fastapi import Request
from pydantic import Field
from sqlalchemy import String, func, select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
from sqlalchemy.orm import Mapped, mapped_column

from openhands.agent_server.models import (
Expand Down Expand Up @@ -122,6 +122,7 @@ class RemoteSandboxService(SandboxService):
user_context: UserContext
httpx_client: httpx.AsyncClient
db_session: AsyncSession
async_session_maker: async_sessionmaker[AsyncSession]

async def _send_runtime_api_request(
self, method: str, path: str, **kwargs: Any
Expand Down Expand Up @@ -235,13 +236,32 @@ async def _secure_select(self):
query = query.where(StoredRemoteSandbox.created_by_user_id == user_id)
return query

async def _isolated_read_one(self, stmt: Any) -> StoredRemoteSandbox | None:
"""Run a scalar read without touching the request-scoped transaction."""
async with self.async_session_maker() as read_session:
result = await read_session.execute(stmt)
return result.scalar_one_or_none()

async def _isolated_read_all(self, stmt: Any) -> list[StoredRemoteSandbox]:
"""Materialize scalar rows in a short-lived read-only session."""
async with self.async_session_maker() as read_session:
result = await read_session.execute(stmt)
return list(result.scalars().all())

async def _get_stored_sandbox(self, sandbox_id: str) -> StoredRemoteSandbox | None:
stmt = await self._secure_select()
stmt = stmt.where(StoredRemoteSandbox.id == sandbox_id)
result = await self.db_session.execute(stmt)
stored_sandbox = result.scalar_one_or_none()
return stored_sandbox

async def _get_stored_sandbox_for_read(
self, sandbox_id: str
) -> StoredRemoteSandbox | None:
stmt = await self._secure_select()
stmt = stmt.where(StoredRemoteSandbox.id == sandbox_id)
return await self._isolated_read_one(stmt)

async def _get_runtime(self, sandbox_id: str) -> dict[str, Any]:
response = await self._send_runtime_api_request(
'GET',
Expand Down Expand Up @@ -328,8 +348,7 @@ async def search_sandboxes(
# Apply limit and get one extra to check if there are more results
stmt = stmt.limit(limit + 1).order_by(StoredRemoteSandbox.created_at.desc())

result = await self.db_session.execute(stmt)
stored_sandboxes = result.scalars().all()
stored_sandboxes = await self._isolated_read_all(stmt)

# Check if there are more results
has_more = len(stored_sandboxes) > limit
Expand All @@ -355,7 +374,7 @@ async def search_sandboxes(

async def get_sandbox(self, sandbox_id: str) -> SandboxInfo | None:
"""Get a single sandbox by checking its corresponding runtime."""
stored_sandbox = await self._get_stored_sandbox(sandbox_id)
stored_sandbox = await self._get_stored_sandbox_for_read(sandbox_id)
if stored_sandbox is None:
return None

Expand All @@ -379,8 +398,7 @@ async def get_sandbox_by_session_api_key(
stmt = stmt.where(
StoredRemoteSandbox.session_api_key_hash == session_api_key_hash
)
result = await self.db_session.execute(stmt)
stored_sandbox = result.scalar_one_or_none()
stored_sandbox = await self._isolated_read_one(stmt)

if stored_sandbox is None:
return None
Expand Down Expand Up @@ -415,8 +433,7 @@ async def _get_user_running_sandboxes(self) -> list[StoredRemoteSandbox]:
query = query.filter(StoredRemoteSandbox.id.in_(running_session_ids)).order_by(
StoredRemoteSandbox.created_at.asc()
)
result = await self.db_session.execute(query)
return list(result.scalars().all())
return await self._isolated_read_all(query)

async def get_sandbox_record_by_session_api_key(
self, session_api_key: str
Expand Down Expand Up @@ -830,9 +847,9 @@ async def batch_get_sandboxes(
return []
query = await self._secure_select()
query = query.filter(StoredRemoteSandbox.id.in_(sandbox_ids))
stored_remote_sandboxes = await self.db_session.execute(query)
stored_remote_sandboxes = await self._isolated_read_all(query)
stored_remote_sandboxes_by_id = {
stored_remote_sandbox[0].id: stored_remote_sandbox[0]
stored_remote_sandbox.id: stored_remote_sandbox
for stored_remote_sandbox in stored_remote_sandboxes
}

Expand Down Expand Up @@ -1117,6 +1134,7 @@ async def inject(
# If no public facing web url is defined, poll for changes as callbacks will be unavailable.
# This is primarily used for local development rather than production
config = get_global_config()
async_session_maker = await config.db_session.get_async_session_maker()
web_url = config.web_url
if web_url is None or 'localhost' in web_url:
global polling_task
Expand Down Expand Up @@ -1146,4 +1164,5 @@ async def inject(
user_context=user_context,
httpx_client=httpx_client,
db_session=db_session,
async_session_maker=async_session_maker,
)
152 changes: 129 additions & 23 deletions tests/unit/app_server/test_remote_sandbox_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@

import httpx
import pytest
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession

from openhands.app_server.errors import SandboxDeleteRetryError, SandboxError
Expand All @@ -30,6 +31,7 @@
WEBHOOK_CALLBACK_VARIABLE,
RemoteSandboxService,
StoredRemoteSandbox,
_hash_session_api_key,
)
from openhands.app_server.sandbox.sandbox_models import (
AGENT_SERVER,
Expand Down Expand Up @@ -79,9 +81,22 @@ def mock_db_session():
return AsyncMock(spec=AsyncSession)


@pytest.fixture
def mock_async_session_maker(mock_db_session):
"""Create isolated-read contexts backed by the test's DB session mock."""
context = AsyncMock()
context.__aenter__.return_value = mock_db_session
context.__aexit__.return_value = False
return MagicMock(return_value=context)


@pytest.fixture
def remote_sandbox_service(
mock_sandbox_spec_service, mock_user_context, mock_httpx_client, mock_db_session
mock_sandbox_spec_service,
mock_user_context,
mock_httpx_client,
mock_db_session,
mock_async_session_maker,
):
"""Create RemoteSandboxService instance with mocked dependencies."""
return RemoteSandboxService(
Expand All @@ -96,6 +111,7 @@ def remote_sandbox_service(
user_context=mock_user_context,
httpx_client=mock_httpx_client,
db_session=mock_db_session,
async_session_maker=mock_async_session_maker,
)


Expand Down Expand Up @@ -1084,7 +1100,7 @@ async def test_get_sandbox_exists(self, remote_sandbox_service):
"""Test getting an existing sandbox."""
# Setup
stored_sandbox = create_stored_sandbox()
remote_sandbox_service._get_stored_sandbox = AsyncMock(
remote_sandbox_service._get_stored_sandbox_for_read = AsyncMock(
return_value=stored_sandbox
)
remote_sandbox_service._to_sandbox_info = MagicMock(
Expand All @@ -1104,15 +1120,17 @@ async def test_get_sandbox_exists(self, remote_sandbox_service):
# Verify
assert result is not None
assert result.id == 'test-sandbox-123'
remote_sandbox_service._get_stored_sandbox.assert_called_once_with(
remote_sandbox_service._get_stored_sandbox_for_read.assert_called_once_with(
'test-sandbox-123'
)

@pytest.mark.asyncio
async def test_get_sandbox_not_exists(self, remote_sandbox_service):
"""Test getting a non-existent sandbox."""
# Setup
remote_sandbox_service._get_stored_sandbox = AsyncMock(return_value=None)
remote_sandbox_service._get_stored_sandbox_for_read = AsyncMock(
return_value=None
)

# Execute
result = await remote_sandbox_service.get_sandbox('non-existent')
Expand Down Expand Up @@ -1857,12 +1875,9 @@ async def test_batch_get_sandboxes_success(self, remote_sandbox_service):
stored_sandbox_2 = create_stored_sandbox(sandbox_id='sandbox-2')
runtime_1 = create_runtime_data(session_id='sandbox-1', status='running')

# Mock DB query result
mock_result = MagicMock()
mock_result.__iter__ = MagicMock(
return_value=iter([(stored_sandbox_1,), (stored_sandbox_2,)])
remote_sandbox_service._isolated_read_all = AsyncMock(
return_value=[stored_sandbox_1, stored_sandbox_2]
)
remote_sandbox_service.db_session.execute = AsyncMock(return_value=mock_result)

# Mock successful runtime batch response
remote_sandbox_service._get_runtimes_batch = AsyncMock(
Expand Down Expand Up @@ -1906,12 +1921,9 @@ async def test_batch_get_sandboxes_graceful_fallback_on_timeout(
stored_sandbox_1 = create_stored_sandbox(sandbox_id='sandbox-1')
stored_sandbox_2 = create_stored_sandbox(sandbox_id='sandbox-2')

# Mock DB query result
mock_result = MagicMock()
mock_result.__iter__ = MagicMock(
return_value=iter([(stored_sandbox_1,), (stored_sandbox_2,)])
remote_sandbox_service._isolated_read_all = AsyncMock(
return_value=[stored_sandbox_1, stored_sandbox_2]
)
remote_sandbox_service.db_session.execute = AsyncMock(return_value=mock_result)

# Mock runtime API timeout
remote_sandbox_service._get_runtimes_batch = AsyncMock(
Expand Down Expand Up @@ -1944,10 +1956,9 @@ async def test_batch_get_sandboxes_graceful_fallback_on_http_error(
sandbox_ids = ['sandbox-1']
stored_sandbox_1 = create_stored_sandbox(sandbox_id='sandbox-1')

# Mock DB query result
mock_result = MagicMock()
mock_result.__iter__ = MagicMock(return_value=iter([(stored_sandbox_1,)]))
remote_sandbox_service.db_session.execute = AsyncMock(return_value=mock_result)
remote_sandbox_service._isolated_read_all = AsyncMock(
return_value=[stored_sandbox_1]
)

# Mock runtime API HTTP error
remote_sandbox_service._get_runtimes_batch = AsyncMock(
Expand Down Expand Up @@ -1977,10 +1988,9 @@ async def test_batch_get_sandboxes_graceful_fallback_on_raise_for_status(
sandbox_ids = ['sandbox-1']
stored_sandbox_1 = create_stored_sandbox(sandbox_id='sandbox-1')

# Mock DB query result
mock_result = MagicMock()
mock_result.__iter__ = MagicMock(return_value=iter([(stored_sandbox_1,)]))
remote_sandbox_service.db_session.execute = AsyncMock(return_value=mock_result)
remote_sandbox_service._isolated_read_all = AsyncMock(
return_value=[stored_sandbox_1]
)

# Mock HTTP status error from raise_for_status()
remote_sandbox_service._get_runtimes_batch = AsyncMock(
Expand Down Expand Up @@ -2669,6 +2679,93 @@ async def test_archive_sibling_conversations_distinct_keys(self, monkeypatch):
assert any('/shared-sandbox/convb/' in p for p in patch_blobs)


class TestIsolatedSandboxReads:
"""Pure reads must not commit sibling writes on the request session."""

@pytest.fixture
async def service_and_sessions(
self, tmp_path, mock_sandbox_spec_service, mock_user_context
):
from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine

from openhands.app_server.utils.sql_utils import Base

engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "reads.db"}')
async with engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
maker = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False)
async with maker() as shared_session:
service = RemoteSandboxService(
sandbox_spec_service=mock_sandbox_spec_service,
api_url='https://api.example.com',
api_key='test-api-key',
web_url='https://web.example.com',
resource_factor=1,
runtime_class='gvisor',
start_sandbox_timeout=120,
max_num_sandboxes=10,
user_context=mock_user_context,
httpx_client=AsyncMock(spec=httpx.AsyncClient),
db_session=shared_session,
async_session_maker=maker,
)
yield service, shared_session, maker
await engine.dispose()

@pytest.mark.asyncio
@pytest.mark.parametrize(
'read_path',
[
'get_sandbox',
'search_sandboxes',
'get_sandbox_by_session_api_key',
'batch_get_sandboxes',
'get_user_running_sandboxes',
],
)
async def test_read_path_does_not_commit_shared_session(
self, service_and_sessions, read_path
):
service, shared_session, maker = service_and_sessions
session_key = 'session-key'

committed = create_stored_sandbox(
sandbox_id='committed',
session_api_key_hash=_hash_session_api_key(session_key),
)
async with maker() as seed_session:
seed_session.add(committed)
await seed_session.commit()

pending = create_stored_sandbox(sandbox_id='pending-sibling-write')
shared_session.add(pending)

service._get_runtime = AsyncMock(return_value=None)
service._get_runtimes_batch = AsyncMock(return_value={})
list_response = MagicMock()
list_response.raise_for_status.return_value = None
list_response.json.return_value = {'runtimes': [{'session_id': committed.id}]}
service._send_runtime_api_request = AsyncMock(return_value=list_response)

if read_path == 'get_sandbox':
await service.get_sandbox(committed.id)
elif read_path == 'search_sandboxes':
await service.search_sandboxes()
elif read_path == 'get_sandbox_by_session_api_key':
await service.get_sandbox_by_session_api_key(session_key)
elif read_path == 'batch_get_sandboxes':
await service.batch_get_sandboxes([committed.id])
else:
await service._get_user_running_sandboxes()

assert pending in shared_session.new
async with maker() as verification_session:
result = await verification_session.execute(
select(StoredRemoteSandbox).where(StoredRemoteSandbox.id == pending.id)
)
assert result.scalar_one_or_none() is None


class TestDeleteSandboxKeyHandling:
"""The session_api_key_hash is invalidated UP FRONT on delete (a delete is
often a revoke of a leaked key). When a transient error keeps the row for
Expand Down Expand Up @@ -2706,8 +2803,14 @@ async def real_session(self, async_engine):

@pytest.fixture
def service_with_real_db(
self, mock_sandbox_spec_service, mock_user_context, real_session
self,
mock_sandbox_spec_service,
mock_user_context,
real_session,
async_engine,
):
from sqlalchemy.ext.asyncio import async_sessionmaker

return RemoteSandboxService(
sandbox_spec_service=mock_sandbox_spec_service,
api_url='https://api.example.com',
Expand All @@ -2720,6 +2823,9 @@ def service_with_real_db(
user_context=mock_user_context,
httpx_client=AsyncMock(spec=httpx.AsyncClient),
db_session=real_session,
async_session_maker=async_sessionmaker(
async_engine, class_=AsyncSession, expire_on_commit=False
),
)

@pytest.mark.asyncio
Expand Down