mirror of
https://github.com/OpenHands/OpenHands.git
synced 2026-10-07 13:18:19 +08:00
perf: skip runtime API call in webhook auth via SandboxRecord (#14732)
Co-authored-by: openhands <openhands@all-hands.dev>
This commit is contained in:
co-authored by
openhands
parent
3dc83a055d
commit
2a3f06a75d
@@ -37,7 +37,7 @@ from openhands.app_server.event_callback.set_title_callback_processor import (
|
||||
SetTitleCallbackProcessor,
|
||||
)
|
||||
from openhands.app_server.integrations.provider import ProviderType
|
||||
from openhands.app_server.sandbox.sandbox_models import SandboxInfo
|
||||
from openhands.app_server.sandbox.sandbox_models import SandboxRecord
|
||||
from openhands.app_server.services.injector import InjectorState
|
||||
from openhands.app_server.services.jwt_service import JwtService
|
||||
from openhands.app_server.user.auth_user_context import AuthUserContext
|
||||
@@ -232,8 +232,8 @@ async def valid_sandbox(
|
||||
session_api_key: str = Depends(
|
||||
APIKeyHeader(name='X-Session-API-Key', auto_error=False)
|
||||
),
|
||||
) -> SandboxInfo:
|
||||
"""Use a session api key for validation, and get a sandbox. Subsequent actions
|
||||
) -> SandboxRecord:
|
||||
"""Use a session api key for validation, and get a sandbox record. Subsequent actions
|
||||
are executed in the context of the owner of the sandbox"""
|
||||
if not session_api_key:
|
||||
raise HTTPException(
|
||||
@@ -246,36 +246,36 @@ async def valid_sandbox(
|
||||
# Since we need access to all sandboxes, this is executed in the context of the admin.
|
||||
setattr(state, USER_CONTEXT_ATTR, ADMIN)
|
||||
async with get_sandbox_service(state) as sandbox_service:
|
||||
sandbox_info = await sandbox_service.get_sandbox_by_session_api_key(
|
||||
sandbox_record = await sandbox_service.get_sandbox_record_by_session_api_key(
|
||||
session_api_key
|
||||
)
|
||||
if sandbox_info is None:
|
||||
if sandbox_record is None:
|
||||
raise HTTPException(
|
||||
status.HTTP_401_UNAUTHORIZED, detail='Invalid session API key'
|
||||
)
|
||||
|
||||
# In SAAS Mode there is always a user, so we set the owner of the sandbox
|
||||
# as the current user (Validated by the session_api_key they provided)
|
||||
if sandbox_info.created_by_user_id:
|
||||
if sandbox_record.created_by_user_id:
|
||||
setattr(
|
||||
request.state,
|
||||
USER_CONTEXT_ATTR,
|
||||
SpecifyUserContext(sandbox_info.created_by_user_id),
|
||||
SpecifyUserContext(sandbox_record.created_by_user_id),
|
||||
)
|
||||
elif app_mode == AppMode.SAAS:
|
||||
_logger.error(
|
||||
'Sandbox had no user specified', extra={'sandbox_id': sandbox_info.id}
|
||||
'Sandbox had no user specified', extra={'sandbox_id': sandbox_record.id}
|
||||
)
|
||||
raise HTTPException(
|
||||
status.HTTP_401_UNAUTHORIZED, detail='Sandbox had no user specified'
|
||||
)
|
||||
|
||||
return sandbox_info
|
||||
return sandbox_record
|
||||
|
||||
|
||||
async def valid_conversation(
|
||||
conversation_id: UUID,
|
||||
sandbox_info: SandboxInfo = Depends(valid_sandbox),
|
||||
sandbox_record: SandboxRecord = Depends(valid_sandbox),
|
||||
app_conversation_info_service: AppConversationInfoService = app_conversation_info_service_dependency,
|
||||
) -> AppConversationInfo:
|
||||
app_conversation_info = (
|
||||
@@ -285,12 +285,12 @@ async def valid_conversation(
|
||||
# Conversation does not yet exist - create a stub
|
||||
return AppConversationInfo(
|
||||
id=conversation_id,
|
||||
sandbox_id=sandbox_info.id,
|
||||
created_by_user_id=sandbox_info.created_by_user_id,
|
||||
sandbox_id=sandbox_record.id,
|
||||
created_by_user_id=sandbox_record.created_by_user_id,
|
||||
)
|
||||
|
||||
# Sanity check - Make sure that the conversation and sandbox were created by the same user
|
||||
if app_conversation_info.created_by_user_id != sandbox_info.created_by_user_id:
|
||||
if app_conversation_info.created_by_user_id != sandbox_record.created_by_user_id:
|
||||
raise AuthError()
|
||||
|
||||
return app_conversation_info
|
||||
@@ -299,7 +299,7 @@ async def valid_conversation(
|
||||
@router.post('/conversations')
|
||||
async def on_conversation_update(
|
||||
conversation_info: ConversationInfo,
|
||||
sandbox_info: SandboxInfo = Depends(valid_sandbox),
|
||||
sandbox_record: SandboxRecord = Depends(valid_sandbox),
|
||||
app_conversation_info_service: AppConversationInfoService = app_conversation_info_service_dependency,
|
||||
) -> Success:
|
||||
"""Webhook callback for when a conversation starts, pauses, resumes, or deletes.
|
||||
@@ -309,7 +309,7 @@ async def on_conversation_update(
|
||||
accepted on this single endpoint.
|
||||
"""
|
||||
existing = await valid_conversation(
|
||||
conversation_info.id, sandbox_info, app_conversation_info_service
|
||||
conversation_info.id, sandbox_record, app_conversation_info_service
|
||||
)
|
||||
|
||||
# If the conversation is being deleted, no action is required...
|
||||
@@ -329,7 +329,7 @@ async def on_conversation_update(
|
||||
existing.trigger,
|
||||
merged_tags,
|
||||
conversation_id=str(conversation_info.id),
|
||||
sandbox_id=sandbox_info.id,
|
||||
sandbox_id=sandbox_record.id,
|
||||
)
|
||||
|
||||
# Trust the discriminated-union payload over any stored ``agent_kind``
|
||||
@@ -350,8 +350,8 @@ async def on_conversation_update(
|
||||
app_conversation_info = AppConversationInfo(
|
||||
id=conversation_info.id,
|
||||
title=existing.title or f'Conversation {conversation_info.id.hex}',
|
||||
sandbox_id=sandbox_info.id,
|
||||
created_by_user_id=sandbox_info.created_by_user_id,
|
||||
sandbox_id=sandbox_record.id,
|
||||
created_by_user_id=sandbox_record.created_by_user_id,
|
||||
llm_model=llm_model,
|
||||
agent_kind=agent_kind,
|
||||
# Git parameters
|
||||
@@ -378,7 +378,7 @@ async def on_conversation_update(
|
||||
setattr(
|
||||
state,
|
||||
USER_CONTEXT_ATTR,
|
||||
SpecifyUserContext(sandbox_info.created_by_user_id),
|
||||
SpecifyUserContext(sandbox_record.created_by_user_id),
|
||||
)
|
||||
async with get_event_callback_service(state) as event_callback_service:
|
||||
await event_callback_service.save_event_callback(
|
||||
@@ -391,8 +391,8 @@ async def on_conversation_update(
|
||||
|
||||
# Analytics: conversation created
|
||||
analytics = get_analytics_service()
|
||||
if analytics and sandbox_info.created_by_user_id:
|
||||
ctx = await resolve_analytics_context(sandbox_info.created_by_user_id)
|
||||
if analytics and sandbox_record.created_by_user_id:
|
||||
ctx = await resolve_analytics_context(sandbox_record.created_by_user_id)
|
||||
analytics.track_conversation_created(
|
||||
ctx=ctx,
|
||||
conversation_id=str(conversation_info.id),
|
||||
|
||||
@@ -24,6 +24,7 @@ from openhands.app_server.sandbox.sandbox_models import (
|
||||
ExposedUrl,
|
||||
SandboxInfo,
|
||||
SandboxPage,
|
||||
SandboxRecord,
|
||||
SandboxStatus,
|
||||
)
|
||||
from openhands.app_server.sandbox.sandbox_service import (
|
||||
@@ -357,6 +358,26 @@ class DockerSandboxService(SandboxService):
|
||||
except (NotFound, APIError):
|
||||
return None
|
||||
|
||||
async def get_sandbox_record_by_session_api_key(
|
||||
self, session_api_key: str
|
||||
) -> SandboxRecord | None:
|
||||
"""Get persisted sandbox identity by session API key."""
|
||||
try:
|
||||
all_containers = self.docker_client.containers.list(all=True)
|
||||
for container in all_containers:
|
||||
if container.name and container.name.startswith(
|
||||
self.container_name_prefix
|
||||
):
|
||||
env_vars = self._get_container_env_vars(container)
|
||||
if env_vars.get(SESSION_API_KEY_VARIABLE) == session_api_key:
|
||||
return SandboxRecord(
|
||||
id=container.name,
|
||||
created_by_user_id=None,
|
||||
)
|
||||
return None
|
||||
except (NotFound, APIError):
|
||||
return None
|
||||
|
||||
async def start_sandbox(
|
||||
self, sandbox_spec_id: str | None = None, sandbox_id: str | None = None
|
||||
) -> SandboxInfo:
|
||||
|
||||
@@ -29,6 +29,7 @@ from openhands.app_server.sandbox.sandbox_models import (
|
||||
ExposedUrl,
|
||||
SandboxInfo,
|
||||
SandboxPage,
|
||||
SandboxRecord,
|
||||
SandboxStatus,
|
||||
)
|
||||
from openhands.app_server.sandbox.sandbox_service import (
|
||||
@@ -287,6 +288,18 @@ class ProcessSandboxService(SandboxService):
|
||||
|
||||
return None
|
||||
|
||||
async def get_sandbox_record_by_session_api_key(
|
||||
self, session_api_key: str
|
||||
) -> SandboxRecord | None:
|
||||
"""Get persisted sandbox identity by session API key."""
|
||||
for sandbox_id, process_info in _processes.items():
|
||||
if process_info.session_api_key == session_api_key:
|
||||
return SandboxRecord(
|
||||
id=sandbox_id,
|
||||
created_by_user_id=process_info.user_id,
|
||||
)
|
||||
return None
|
||||
|
||||
async def start_sandbox(
|
||||
self, sandbox_spec_id: str | None = None, sandbox_id: str | None = None
|
||||
) -> SandboxInfo:
|
||||
|
||||
@@ -40,6 +40,7 @@ from openhands.app_server.sandbox.sandbox_models import (
|
||||
ExposedUrl,
|
||||
SandboxInfo,
|
||||
SandboxPage,
|
||||
SandboxRecord,
|
||||
SandboxStatus,
|
||||
)
|
||||
from openhands.app_server.sandbox.sandbox_service import (
|
||||
@@ -378,6 +379,27 @@ class RemoteSandboxService(SandboxService):
|
||||
)
|
||||
return self._to_sandbox_info(stored_sandbox, None)
|
||||
|
||||
async def get_sandbox_record_by_session_api_key(
|
||||
self, session_api_key: str
|
||||
) -> SandboxRecord | None:
|
||||
"""Get persisted sandbox identity by session API key — DB lookup only, no runtime call."""
|
||||
session_api_key_hash = _hash_session_api_key(session_api_key)
|
||||
|
||||
stmt = await self._secure_select()
|
||||
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()
|
||||
|
||||
if stored_sandbox is None:
|
||||
return None
|
||||
|
||||
return SandboxRecord(
|
||||
id=stored_sandbox.id,
|
||||
created_by_user_id=stored_sandbox.created_by_user_id,
|
||||
)
|
||||
|
||||
async def start_sandbox(
|
||||
self, sandbox_spec_id: str | None = None, sandbox_id: str | None = None
|
||||
) -> SandboxInfo:
|
||||
|
||||
@@ -30,6 +30,21 @@ WORKER_1 = 'WORKER_1'
|
||||
WORKER_2 = 'WORKER_2'
|
||||
|
||||
|
||||
class SandboxRecord(BaseModel):
|
||||
"""Persisted identity fields for a sandbox — no live runtime data.
|
||||
|
||||
Contains only what is stored in the app server's own database (id and
|
||||
owner). Use this when you need to authenticate a session key or verify
|
||||
ownership without paying for a runtime API round-trip.
|
||||
|
||||
Use ``SandboxInfo`` when you additionally need live status, exposed URLs,
|
||||
or the plain-text session key returned by the runtime.
|
||||
"""
|
||||
|
||||
id: str
|
||||
created_by_user_id: str | None
|
||||
|
||||
|
||||
class SandboxInfo(BaseModel):
|
||||
"""Information about a sandbox."""
|
||||
|
||||
|
||||
@@ -10,6 +10,7 @@ from openhands.app_server.sandbox.sandbox_models import (
|
||||
AGENT_SERVER,
|
||||
SandboxInfo,
|
||||
SandboxPage,
|
||||
SandboxRecord,
|
||||
SandboxStatus,
|
||||
)
|
||||
from openhands.app_server.services.injector import Injector
|
||||
@@ -47,6 +48,20 @@ class SandboxService(ABC):
|
||||
) -> SandboxInfo | None:
|
||||
"""Get a single sandbox by session API key. Return None if the sandbox was not found."""
|
||||
|
||||
@abstractmethod
|
||||
async def get_sandbox_record_by_session_api_key(
|
||||
self, session_api_key: str
|
||||
) -> SandboxRecord | None:
|
||||
"""Get persisted sandbox identity by session API key without querying the runtime.
|
||||
|
||||
Returns only the fields stored in the app server's own database (id and
|
||||
owner). Use this for authentication paths that do not need live status,
|
||||
exposed URLs, or the plain-text session key — callers avoid a runtime
|
||||
API round-trip.
|
||||
|
||||
Return None if no sandbox matches the key.
|
||||
"""
|
||||
|
||||
async def batch_get_sandboxes(
|
||||
self, sandbox_ids: list[str]
|
||||
) -> list[SandboxInfo | None]:
|
||||
|
||||
@@ -16,6 +16,7 @@ import pytest
|
||||
from openhands.app_server.sandbox.sandbox_models import (
|
||||
SandboxInfo,
|
||||
SandboxPage,
|
||||
SandboxRecord,
|
||||
SandboxStatus,
|
||||
)
|
||||
from openhands.app_server.sandbox.sandbox_service import SandboxService
|
||||
@@ -28,6 +29,7 @@ class MockSandboxService(SandboxService):
|
||||
self.search_sandboxes_mock = AsyncMock()
|
||||
self.get_sandbox_mock = AsyncMock()
|
||||
self.get_sandbox_by_session_api_key_mock = AsyncMock()
|
||||
self.get_sandbox_record_by_session_api_key_mock = AsyncMock()
|
||||
self.start_sandbox_mock = AsyncMock()
|
||||
self.resume_sandbox_mock = AsyncMock()
|
||||
self.pause_sandbox_mock = AsyncMock()
|
||||
@@ -46,6 +48,11 @@ class MockSandboxService(SandboxService):
|
||||
) -> SandboxInfo | None:
|
||||
return await self.get_sandbox_by_session_api_key_mock(session_api_key)
|
||||
|
||||
async def get_sandbox_record_by_session_api_key(
|
||||
self, session_api_key: str
|
||||
) -> SandboxRecord | None:
|
||||
return await self.get_sandbox_record_by_session_api_key_mock(session_api_key)
|
||||
|
||||
async def start_sandbox(
|
||||
self, sandbox_spec_id: str | None = None, sandbox_id: str | None = None
|
||||
) -> SandboxInfo:
|
||||
|
||||
@@ -65,7 +65,7 @@ def service(async_session) -> SQLAppConversationInfoService:
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sandbox_info():
|
||||
def sandbox_record():
|
||||
sandbox = MagicMock()
|
||||
sandbox.id = 'sandbox_acp_test'
|
||||
sandbox.created_by_user_id = 'user_123'
|
||||
@@ -114,7 +114,9 @@ def _make_acp_conversation_info(acp_command: list[str]) -> ConversationInfo:
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_llm_conversation_stores_llm_model(async_session, service, sandbox_info):
|
||||
async def test_llm_conversation_stores_llm_model(
|
||||
async_session, service, sandbox_record
|
||||
):
|
||||
"""LLM path stores the real model in llm_model and sets agent_kind='openhands'."""
|
||||
llm_info = _make_llm_conversation_info()
|
||||
conversation_id = llm_info.id
|
||||
@@ -122,8 +124,8 @@ async def test_llm_conversation_stores_llm_model(async_session, service, sandbox
|
||||
existing = AppConversationInfo(
|
||||
id=conversation_id,
|
||||
title='Test',
|
||||
sandbox_id=sandbox_info.id,
|
||||
created_by_user_id=sandbox_info.created_by_user_id,
|
||||
sandbox_id=sandbox_record.id,
|
||||
created_by_user_id=sandbox_record.created_by_user_id,
|
||||
)
|
||||
|
||||
with patch(
|
||||
@@ -132,7 +134,7 @@ async def test_llm_conversation_stores_llm_model(async_session, service, sandbox
|
||||
):
|
||||
result = await on_conversation_update(
|
||||
conversation_info=llm_info,
|
||||
sandbox_info=sandbox_info,
|
||||
sandbox_record=sandbox_record,
|
||||
app_conversation_info_service=service,
|
||||
)
|
||||
|
||||
@@ -145,7 +147,7 @@ async def test_llm_conversation_stores_llm_model(async_session, service, sandbox
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acp_conversation_sets_agent_kind(async_session, service, sandbox_info):
|
||||
async def test_acp_conversation_sets_agent_kind(async_session, service, sandbox_record):
|
||||
"""ACP path sets agent_kind='acp' and leaves llm_model null."""
|
||||
acp_info = _make_acp_conversation_info(
|
||||
acp_command=['npx', '-y', '@agentclientprotocol/claude-agent-acp']
|
||||
@@ -155,8 +157,8 @@ async def test_acp_conversation_sets_agent_kind(async_session, service, sandbox_
|
||||
existing = AppConversationInfo(
|
||||
id=conversation_id,
|
||||
title='Test',
|
||||
sandbox_id=sandbox_info.id,
|
||||
created_by_user_id=sandbox_info.created_by_user_id,
|
||||
sandbox_id=sandbox_record.id,
|
||||
created_by_user_id=sandbox_record.created_by_user_id,
|
||||
)
|
||||
|
||||
with patch(
|
||||
@@ -165,7 +167,7 @@ async def test_acp_conversation_sets_agent_kind(async_session, service, sandbox_
|
||||
):
|
||||
result = await on_conversation_update(
|
||||
conversation_info=acp_info,
|
||||
sandbox_info=sandbox_info,
|
||||
sandbox_record=sandbox_record,
|
||||
app_conversation_info_service=service,
|
||||
)
|
||||
|
||||
@@ -179,7 +181,7 @@ async def test_acp_conversation_sets_agent_kind(async_session, service, sandbox_
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acp_server_tag_preserved_on_webhook_update(
|
||||
async_session, service, sandbox_info
|
||||
async_session, service, sandbox_record
|
||||
):
|
||||
"""``tags['acp_server']`` set during creation must survive a webhook update.
|
||||
|
||||
@@ -194,8 +196,8 @@ async def test_acp_server_tag_preserved_on_webhook_update(
|
||||
existing = AppConversationInfo(
|
||||
id=conversation_id,
|
||||
title='Test',
|
||||
sandbox_id=sandbox_info.id,
|
||||
created_by_user_id=sandbox_info.created_by_user_id,
|
||||
sandbox_id=sandbox_record.id,
|
||||
created_by_user_id=sandbox_record.created_by_user_id,
|
||||
tags={'acp_server': 'claude-code'},
|
||||
)
|
||||
|
||||
@@ -205,7 +207,7 @@ async def test_acp_server_tag_preserved_on_webhook_update(
|
||||
):
|
||||
await on_conversation_update(
|
||||
conversation_info=acp_info,
|
||||
sandbox_info=sandbox_info,
|
||||
sandbox_record=sandbox_record,
|
||||
app_conversation_info_service=service,
|
||||
)
|
||||
|
||||
@@ -221,7 +223,7 @@ async def test_acp_server_tag_preserved_on_webhook_update(
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acp_conversation_analytics_llm_model_is_null(
|
||||
async_session, service, sandbox_info
|
||||
async_session, service, sandbox_record
|
||||
):
|
||||
"""``track_conversation_created`` must receive ``llm_model=None`` for ACP.
|
||||
|
||||
@@ -234,8 +236,8 @@ async def test_acp_conversation_analytics_llm_model_is_null(
|
||||
existing = AppConversationInfo(
|
||||
id=acp_info.id,
|
||||
title='Test',
|
||||
sandbox_id=sandbox_info.id,
|
||||
created_by_user_id=sandbox_info.created_by_user_id,
|
||||
sandbox_id=sandbox_record.id,
|
||||
created_by_user_id=sandbox_record.created_by_user_id,
|
||||
)
|
||||
analytics = MagicMock()
|
||||
|
||||
@@ -255,7 +257,7 @@ async def test_acp_conversation_analytics_llm_model_is_null(
|
||||
):
|
||||
await on_conversation_update(
|
||||
conversation_info=acp_info,
|
||||
sandbox_info=sandbox_info,
|
||||
sandbox_record=sandbox_record,
|
||||
app_conversation_info_service=service,
|
||||
)
|
||||
|
||||
|
||||
@@ -19,7 +19,7 @@ from openhands.app_server.event_callback.webhook_router import (
|
||||
valid_conversation,
|
||||
valid_sandbox,
|
||||
)
|
||||
from openhands.app_server.sandbox.sandbox_models import SandboxInfo, SandboxStatus
|
||||
from openhands.app_server.sandbox.sandbox_models import SandboxRecord
|
||||
from openhands.app_server.user.specifiy_user_context import (
|
||||
USER_CONTEXT_ATTR,
|
||||
SpecifyUserContext,
|
||||
@@ -73,16 +73,13 @@ class TestValidSandbox:
|
||||
# Arrange
|
||||
session_api_key = 'valid-api-key-123'
|
||||
user_id = 'user-123'
|
||||
expected_sandbox = SandboxInfo(
|
||||
expected_sandbox = SandboxRecord(
|
||||
id='sandbox-123',
|
||||
status=SandboxStatus.RUNNING,
|
||||
session_api_key=session_api_key,
|
||||
created_by_user_id=user_id,
|
||||
sandbox_spec_id='spec-123',
|
||||
)
|
||||
|
||||
mock_sandbox_service = AsyncMock()
|
||||
mock_sandbox_service.get_sandbox_by_session_api_key = AsyncMock(
|
||||
mock_sandbox_service.get_sandbox_record_by_session_api_key = AsyncMock(
|
||||
return_value=expected_sandbox
|
||||
)
|
||||
|
||||
@@ -100,7 +97,7 @@ class TestValidSandbox:
|
||||
|
||||
# Assert
|
||||
assert result == expected_sandbox
|
||||
mock_sandbox_service.get_sandbox_by_session_api_key.assert_called_once_with(
|
||||
mock_sandbox_service.get_sandbox_record_by_session_api_key.assert_called_once_with(
|
||||
session_api_key
|
||||
)
|
||||
|
||||
@@ -116,16 +113,13 @@ class TestValidSandbox:
|
||||
# Arrange
|
||||
session_api_key = 'valid-api-key'
|
||||
sandbox_owner_id = 'sandbox-owner-user-id'
|
||||
expected_sandbox = SandboxInfo(
|
||||
expected_sandbox = SandboxRecord(
|
||||
id='sandbox-456',
|
||||
status=SandboxStatus.RUNNING,
|
||||
session_api_key=session_api_key,
|
||||
created_by_user_id=sandbox_owner_id,
|
||||
sandbox_spec_id='spec-456',
|
||||
)
|
||||
|
||||
mock_sandbox_service = AsyncMock()
|
||||
mock_sandbox_service.get_sandbox_by_session_api_key = AsyncMock(
|
||||
mock_sandbox_service.get_sandbox_record_by_session_api_key = AsyncMock(
|
||||
return_value=expected_sandbox
|
||||
)
|
||||
|
||||
@@ -152,16 +146,13 @@ class TestValidSandbox:
|
||||
"""Test that user_context is not set when sandbox has no created_by_user_id."""
|
||||
# Arrange
|
||||
session_api_key = 'valid-api-key'
|
||||
expected_sandbox = SandboxInfo(
|
||||
expected_sandbox = SandboxRecord(
|
||||
id='sandbox-789',
|
||||
status=SandboxStatus.RUNNING,
|
||||
session_api_key=session_api_key,
|
||||
created_by_user_id=None, # No user ID
|
||||
sandbox_spec_id='spec-789',
|
||||
created_by_user_id=None,
|
||||
)
|
||||
|
||||
mock_sandbox_service = AsyncMock()
|
||||
mock_sandbox_service.get_sandbox_by_session_api_key = AsyncMock(
|
||||
mock_sandbox_service.get_sandbox_record_by_session_api_key = AsyncMock(
|
||||
return_value=expected_sandbox
|
||||
)
|
||||
|
||||
@@ -190,16 +181,13 @@ class TestValidSandbox:
|
||||
"""Test that user_context is not set when sandbox has no created_by_user_id."""
|
||||
# Arrange
|
||||
session_api_key = 'valid-api-key'
|
||||
expected_sandbox = SandboxInfo(
|
||||
expected_sandbox = SandboxRecord(
|
||||
id='sandbox-789',
|
||||
status=SandboxStatus.RUNNING,
|
||||
session_api_key=session_api_key,
|
||||
created_by_user_id=None, # No user ID
|
||||
sandbox_spec_id='spec-789',
|
||||
created_by_user_id=None,
|
||||
)
|
||||
|
||||
mock_sandbox_service = AsyncMock()
|
||||
mock_sandbox_service.get_sandbox_by_session_api_key = AsyncMock(
|
||||
mock_sandbox_service.get_sandbox_record_by_session_api_key = AsyncMock(
|
||||
return_value=expected_sandbox
|
||||
)
|
||||
|
||||
@@ -245,7 +233,7 @@ class TestValidSandbox:
|
||||
# Arrange
|
||||
session_api_key = 'invalid-api-key'
|
||||
mock_sandbox_service = AsyncMock()
|
||||
mock_sandbox_service.get_sandbox_by_session_api_key = AsyncMock(
|
||||
mock_sandbox_service.get_sandbox_record_by_session_api_key = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
|
||||
@@ -283,7 +271,7 @@ class TestValidSandbox:
|
||||
assert exc_info.value.status_code == status.HTTP_401_UNAUTHORIZED
|
||||
assert 'X-Session-API-Key header is required' in exc_info.value.detail
|
||||
# Verify the sandbox service was NOT called (rejected before lookup)
|
||||
mock_sandbox_service.get_sandbox_by_session_api_key.assert_not_called()
|
||||
mock_sandbox_service.get_sandbox_record_by_session_api_key.assert_not_called()
|
||||
|
||||
|
||||
class TestValidConversation:
|
||||
@@ -294,12 +282,9 @@ class TestValidConversation:
|
||||
"""Test that existing conversation returns info."""
|
||||
# Arrange
|
||||
conversation_id = uuid4()
|
||||
sandbox_info = SandboxInfo(
|
||||
sandbox_record = SandboxRecord(
|
||||
id='sandbox-123',
|
||||
status=SandboxStatus.RUNNING,
|
||||
session_api_key='api-key',
|
||||
created_by_user_id='user-123',
|
||||
sandbox_spec_id='spec-123',
|
||||
)
|
||||
|
||||
expected_info = MagicMock()
|
||||
@@ -311,7 +296,7 @@ class TestValidConversation:
|
||||
# Act
|
||||
result = await valid_conversation(
|
||||
conversation_id=conversation_id,
|
||||
sandbox_info=sandbox_info,
|
||||
sandbox_record=sandbox_record,
|
||||
app_conversation_info_service=mock_service,
|
||||
)
|
||||
|
||||
@@ -323,12 +308,9 @@ class TestValidConversation:
|
||||
"""Test that non-existing conversation creates a stub."""
|
||||
# Arrange
|
||||
conversation_id = uuid4()
|
||||
sandbox_info = SandboxInfo(
|
||||
sandbox_record = SandboxRecord(
|
||||
id='sandbox-123',
|
||||
status=SandboxStatus.RUNNING,
|
||||
session_api_key='api-key',
|
||||
created_by_user_id='user-123',
|
||||
sandbox_spec_id='spec-123',
|
||||
)
|
||||
|
||||
mock_service = AsyncMock()
|
||||
@@ -337,26 +319,23 @@ class TestValidConversation:
|
||||
# Act
|
||||
result = await valid_conversation(
|
||||
conversation_id=conversation_id,
|
||||
sandbox_info=sandbox_info,
|
||||
sandbox_record=sandbox_record,
|
||||
app_conversation_info_service=mock_service,
|
||||
)
|
||||
|
||||
# Assert
|
||||
assert result.id == conversation_id
|
||||
assert result.sandbox_id == sandbox_info.id
|
||||
assert result.created_by_user_id == sandbox_info.created_by_user_id
|
||||
assert result.sandbox_id == sandbox_record.id
|
||||
assert result.created_by_user_id == sandbox_record.created_by_user_id
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_valid_conversation_different_user_raises_auth_error(self):
|
||||
"""Test that conversation from different user raises AuthError."""
|
||||
# Arrange
|
||||
conversation_id = uuid4()
|
||||
sandbox_info = SandboxInfo(
|
||||
sandbox_record = SandboxRecord(
|
||||
id='sandbox-123',
|
||||
status=SandboxStatus.RUNNING,
|
||||
session_api_key='api-key',
|
||||
created_by_user_id='user-123',
|
||||
sandbox_spec_id='spec-123',
|
||||
)
|
||||
|
||||
# Conversation created by different user
|
||||
@@ -374,7 +353,7 @@ class TestValidConversation:
|
||||
with pytest.raises(AuthError):
|
||||
await valid_conversation(
|
||||
conversation_id=conversation_id,
|
||||
sandbox_info=sandbox_info,
|
||||
sandbox_record=sandbox_record,
|
||||
app_conversation_info_service=mock_service,
|
||||
)
|
||||
|
||||
@@ -384,12 +363,9 @@ class TestValidConversation:
|
||||
# Arrange
|
||||
conversation_id = uuid4()
|
||||
user_id = 'user-123'
|
||||
sandbox_info = SandboxInfo(
|
||||
sandbox_record = SandboxRecord(
|
||||
id='sandbox-123',
|
||||
status=SandboxStatus.RUNNING,
|
||||
session_api_key='api-key',
|
||||
created_by_user_id=user_id,
|
||||
sandbox_spec_id='spec-123',
|
||||
)
|
||||
|
||||
# Conversation created by same user
|
||||
@@ -402,7 +378,7 @@ class TestValidConversation:
|
||||
# Act
|
||||
result = await valid_conversation(
|
||||
conversation_id=conversation_id,
|
||||
sandbox_info=sandbox_info,
|
||||
sandbox_record=sandbox_record,
|
||||
app_conversation_info_service=mock_service,
|
||||
)
|
||||
|
||||
@@ -418,17 +394,14 @@ class TestWebhookAuthenticationIntegration:
|
||||
"""Test complete auth flow with valid API key."""
|
||||
# Arrange
|
||||
session_api_key = 'valid-api-key'
|
||||
sandbox_info = SandboxInfo(
|
||||
sandbox_record = SandboxRecord(
|
||||
id='sandbox-123',
|
||||
status=SandboxStatus.RUNNING,
|
||||
session_api_key=session_api_key,
|
||||
created_by_user_id='user-123',
|
||||
sandbox_spec_id='spec-123',
|
||||
)
|
||||
|
||||
mock_sandbox_service = AsyncMock()
|
||||
mock_sandbox_service.get_sandbox_by_session_api_key = AsyncMock(
|
||||
return_value=sandbox_info
|
||||
mock_sandbox_service.get_sandbox_record_by_session_api_key = AsyncMock(
|
||||
return_value=sandbox_record
|
||||
)
|
||||
|
||||
conversation_info = MagicMock()
|
||||
@@ -454,7 +427,7 @@ class TestWebhookAuthenticationIntegration:
|
||||
# Then call valid_conversation
|
||||
conversation_result = await valid_conversation(
|
||||
conversation_id=uuid4(),
|
||||
sandbox_info=sandbox_result,
|
||||
sandbox_record=sandbox_result,
|
||||
app_conversation_info_service=mock_conversation_service,
|
||||
)
|
||||
|
||||
@@ -468,7 +441,7 @@ class TestWebhookAuthenticationIntegration:
|
||||
# Arrange
|
||||
session_api_key = 'invalid-api-key'
|
||||
mock_sandbox_service = AsyncMock()
|
||||
mock_sandbox_service.get_sandbox_by_session_api_key = AsyncMock(
|
||||
mock_sandbox_service.get_sandbox_record_by_session_api_key = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
|
||||
@@ -492,17 +465,14 @@ class TestWebhookAuthenticationIntegration:
|
||||
"""Test complete auth flow with valid key but wrong user fails."""
|
||||
# Arrange
|
||||
session_api_key = 'valid-api-key'
|
||||
sandbox_info = SandboxInfo(
|
||||
sandbox_record = SandboxRecord(
|
||||
id='sandbox-123',
|
||||
status=SandboxStatus.RUNNING,
|
||||
session_api_key=session_api_key,
|
||||
created_by_user_id='user-123',
|
||||
sandbox_spec_id='spec-123',
|
||||
)
|
||||
|
||||
mock_sandbox_service = AsyncMock()
|
||||
mock_sandbox_service.get_sandbox_by_session_api_key = AsyncMock(
|
||||
return_value=sandbox_info
|
||||
mock_sandbox_service.get_sandbox_record_by_session_api_key = AsyncMock(
|
||||
return_value=sandbox_record
|
||||
)
|
||||
|
||||
# Conversation created by different user
|
||||
@@ -532,7 +502,7 @@ class TestWebhookAuthenticationIntegration:
|
||||
with pytest.raises(AuthError):
|
||||
await valid_conversation(
|
||||
conversation_id=uuid4(),
|
||||
sandbox_info=sandbox_result,
|
||||
sandbox_record=sandbox_result,
|
||||
app_conversation_info_service=mock_conversation_service,
|
||||
)
|
||||
|
||||
|
||||
@@ -26,7 +26,7 @@ from openhands.app_server.event_callback.set_title_callback_processor import (
|
||||
SetTitleCallbackProcessor,
|
||||
)
|
||||
from openhands.app_server.event_callback.webhook_router import on_conversation_update
|
||||
from openhands.app_server.sandbox.sandbox_models import SandboxInfo, SandboxStatus
|
||||
from openhands.app_server.sandbox.sandbox_models import SandboxRecord
|
||||
from openhands.app_server.user.specifiy_user_context import SpecifyUserContext
|
||||
from openhands.app_server.utils.sql_utils import Base
|
||||
from openhands.sdk.conversation import ConversationExecutionStatus
|
||||
@@ -73,14 +73,11 @@ def app_conversation_info_service(
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sandbox_info() -> SandboxInfo:
|
||||
def sandbox_record() -> SandboxRecord:
|
||||
"""Create a test sandbox info."""
|
||||
return SandboxInfo(
|
||||
return SandboxRecord(
|
||||
id='sandbox_123',
|
||||
status=SandboxStatus.RUNNING,
|
||||
session_api_key='test_session_key',
|
||||
created_by_user_id='user_123',
|
||||
sandbox_spec_id='spec_123',
|
||||
)
|
||||
|
||||
|
||||
@@ -114,7 +111,7 @@ class TestOnConversationUpdateAutoTitle:
|
||||
self,
|
||||
async_session,
|
||||
app_conversation_info_service,
|
||||
sandbox_info,
|
||||
sandbox_record,
|
||||
mock_conversation_info,
|
||||
):
|
||||
"""Test that SetTitleCallbackProcessor is registered for new conversations.
|
||||
@@ -163,7 +160,7 @@ class TestOnConversationUpdateAutoTitle:
|
||||
):
|
||||
result = await on_conversation_update(
|
||||
conversation_info=mock_conversation_info,
|
||||
sandbox_info=sandbox_info,
|
||||
sandbox_record=sandbox_record,
|
||||
app_conversation_info_service=app_conversation_info_service,
|
||||
)
|
||||
|
||||
@@ -182,7 +179,7 @@ class TestOnConversationUpdateAutoTitle:
|
||||
self,
|
||||
async_session,
|
||||
app_conversation_info_service,
|
||||
sandbox_info,
|
||||
sandbox_record,
|
||||
mock_conversation_info,
|
||||
):
|
||||
"""Test that SetTitleCallbackProcessor is NOT registered for existing conversations.
|
||||
@@ -231,7 +228,7 @@ class TestOnConversationUpdateAutoTitle:
|
||||
):
|
||||
result = await on_conversation_update(
|
||||
conversation_info=mock_conversation_info,
|
||||
sandbox_info=sandbox_info,
|
||||
sandbox_record=sandbox_record,
|
||||
app_conversation_info_service=app_conversation_info_service,
|
||||
)
|
||||
|
||||
@@ -246,7 +243,7 @@ class TestOnConversationUpdateAutoTitle:
|
||||
self,
|
||||
async_session,
|
||||
app_conversation_info_service,
|
||||
sandbox_info,
|
||||
sandbox_record,
|
||||
mock_conversation_info,
|
||||
):
|
||||
"""Test that SetTitleCallbackProcessor is NOT registered for deleting conversations.
|
||||
@@ -299,7 +296,7 @@ class TestOnConversationUpdateAutoTitle:
|
||||
):
|
||||
result = await on_conversation_update(
|
||||
conversation_info=mock_conversation_info,
|
||||
sandbox_info=sandbox_info,
|
||||
sandbox_record=sandbox_record,
|
||||
app_conversation_info_service=app_conversation_info_service,
|
||||
)
|
||||
|
||||
@@ -314,18 +311,18 @@ class TestOnConversationUpdateAutoTitle:
|
||||
self,
|
||||
async_session,
|
||||
app_conversation_info_service,
|
||||
sandbox_info,
|
||||
sandbox_record,
|
||||
mock_conversation_info,
|
||||
):
|
||||
"""Test that the callback registration uses the user_id from sandbox_info.
|
||||
"""Test that the callback registration uses the user_id from sandbox_record.
|
||||
|
||||
Arrange:
|
||||
- Create a stub conversation (title=None)
|
||||
- sandbox_info has specific user_id
|
||||
- sandbox_record has specific user_id
|
||||
Act:
|
||||
- Call on_conversation_update webhook
|
||||
Assert:
|
||||
- InjectorState is created with sandbox_info.created_by_user_id
|
||||
- InjectorState is created with sandbox_record.created_by_user_id
|
||||
"""
|
||||
# Arrange
|
||||
conversation_id = mock_conversation_info.id
|
||||
@@ -362,7 +359,7 @@ class TestOnConversationUpdateAutoTitle:
|
||||
):
|
||||
result = await on_conversation_update(
|
||||
conversation_info=mock_conversation_info,
|
||||
sandbox_info=sandbox_info,
|
||||
sandbox_record=sandbox_record,
|
||||
app_conversation_info_service=app_conversation_info_service,
|
||||
)
|
||||
|
||||
@@ -376,14 +373,14 @@ class TestOnConversationUpdateAutoTitle:
|
||||
user_context = getattr(captured_state, USER_CONTEXT_ATTR)
|
||||
# get_user_id() is async, so we need to await it
|
||||
user_id = await user_context.get_user_id()
|
||||
assert user_id == sandbox_info.created_by_user_id
|
||||
assert user_id == sandbox_record.created_by_user_id
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_conversation_saved_before_callback_registration(
|
||||
self,
|
||||
async_session,
|
||||
app_conversation_info_service,
|
||||
sandbox_info,
|
||||
sandbox_record,
|
||||
mock_conversation_info,
|
||||
):
|
||||
"""Test that conversation is saved before SetTitleCallbackProcessor is registered.
|
||||
@@ -447,7 +444,7 @@ class TestOnConversationUpdateAutoTitle:
|
||||
):
|
||||
result = await on_conversation_update(
|
||||
conversation_info=mock_conversation_info,
|
||||
sandbox_info=sandbox_info,
|
||||
sandbox_record=sandbox_record,
|
||||
app_conversation_info_service=app_conversation_info_service,
|
||||
)
|
||||
|
||||
|
||||
@@ -23,7 +23,7 @@ from openhands.app_server.app_conversation.sql_app_conversation_info_service imp
|
||||
)
|
||||
from openhands.app_server.event_callback.webhook_router import on_conversation_update
|
||||
from openhands.app_server.integrations.provider import ProviderType
|
||||
from openhands.app_server.sandbox.sandbox_models import SandboxInfo, SandboxStatus
|
||||
from openhands.app_server.sandbox.sandbox_models import SandboxRecord
|
||||
from openhands.app_server.user.specifiy_user_context import SpecifyUserContext
|
||||
from openhands.app_server.utils.sql_utils import Base
|
||||
from openhands.sdk.conversation import ConversationExecutionStatus
|
||||
@@ -78,14 +78,11 @@ def app_conversation_info_service(
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sandbox_info() -> SandboxInfo:
|
||||
def sandbox_record() -> SandboxRecord:
|
||||
"""Create a test sandbox info."""
|
||||
return SandboxInfo(
|
||||
return SandboxRecord(
|
||||
id='sandbox_123',
|
||||
status=SandboxStatus.RUNNING,
|
||||
session_api_key='test_session_key',
|
||||
created_by_user_id='user_123',
|
||||
sandbox_spec_id='spec_123',
|
||||
)
|
||||
|
||||
|
||||
@@ -119,7 +116,7 @@ class TestOnConversationUpdateParentConversationId:
|
||||
self,
|
||||
async_session,
|
||||
app_conversation_info_service,
|
||||
sandbox_info,
|
||||
sandbox_record,
|
||||
mock_conversation_info,
|
||||
):
|
||||
"""Test that parent_conversation_id is preserved when it exists in existing conversation.
|
||||
@@ -154,7 +151,7 @@ class TestOnConversationUpdateParentConversationId:
|
||||
):
|
||||
result = await on_conversation_update(
|
||||
conversation_info=mock_conversation_info,
|
||||
sandbox_info=sandbox_info,
|
||||
sandbox_record=sandbox_record,
|
||||
app_conversation_info_service=app_conversation_info_service,
|
||||
)
|
||||
|
||||
@@ -172,7 +169,7 @@ class TestOnConversationUpdateParentConversationId:
|
||||
self,
|
||||
async_session,
|
||||
app_conversation_info_service,
|
||||
sandbox_info,
|
||||
sandbox_record,
|
||||
mock_conversation_info,
|
||||
):
|
||||
"""Test that parent_conversation_id remains None when it doesn't exist.
|
||||
@@ -204,7 +201,7 @@ class TestOnConversationUpdateParentConversationId:
|
||||
):
|
||||
result = await on_conversation_update(
|
||||
conversation_info=mock_conversation_info,
|
||||
sandbox_info=sandbox_info,
|
||||
sandbox_record=sandbox_record,
|
||||
app_conversation_info_service=app_conversation_info_service,
|
||||
)
|
||||
|
||||
@@ -221,7 +218,7 @@ class TestOnConversationUpdateParentConversationId:
|
||||
async def test_parent_conversation_id_none_for_new_conversation(
|
||||
self,
|
||||
app_conversation_info_service,
|
||||
sandbox_info,
|
||||
sandbox_record,
|
||||
mock_conversation_info,
|
||||
):
|
||||
"""Test that new conversations (stubs) have parent_conversation_id as None.
|
||||
@@ -240,8 +237,8 @@ class TestOnConversationUpdateParentConversationId:
|
||||
# Create stub conversation (simulating valid_conversation for new conversation)
|
||||
stub_conv = AppConversationInfo(
|
||||
id=conversation_id,
|
||||
sandbox_id=sandbox_info.id,
|
||||
created_by_user_id=sandbox_info.created_by_user_id,
|
||||
sandbox_id=sandbox_record.id,
|
||||
created_by_user_id=sandbox_record.created_by_user_id,
|
||||
)
|
||||
|
||||
# Act - call on_conversation_update directly with mocked valid_conversation
|
||||
@@ -258,7 +255,7 @@ class TestOnConversationUpdateParentConversationId:
|
||||
):
|
||||
result = await on_conversation_update(
|
||||
conversation_info=mock_conversation_info,
|
||||
sandbox_info=sandbox_info,
|
||||
sandbox_record=sandbox_record,
|
||||
app_conversation_info_service=app_conversation_info_service,
|
||||
)
|
||||
|
||||
@@ -276,7 +273,7 @@ class TestOnConversationUpdateParentConversationId:
|
||||
self,
|
||||
async_session,
|
||||
app_conversation_info_service,
|
||||
sandbox_info,
|
||||
sandbox_record,
|
||||
mock_conversation_info,
|
||||
):
|
||||
"""Test that parent_conversation_id is preserved alongside other metadata.
|
||||
@@ -314,7 +311,7 @@ class TestOnConversationUpdateParentConversationId:
|
||||
):
|
||||
result = await on_conversation_update(
|
||||
conversation_info=mock_conversation_info,
|
||||
sandbox_info=sandbox_info,
|
||||
sandbox_record=sandbox_record,
|
||||
app_conversation_info_service=app_conversation_info_service,
|
||||
)
|
||||
|
||||
@@ -341,7 +338,7 @@ class TestOnConversationUpdateParentConversationId:
|
||||
self,
|
||||
async_session,
|
||||
app_conversation_info_service,
|
||||
sandbox_info,
|
||||
sandbox_record,
|
||||
mock_conversation_info,
|
||||
):
|
||||
"""Test that parent_conversation_id remains stable across multiple updates.
|
||||
@@ -386,7 +383,7 @@ class TestOnConversationUpdateParentConversationId:
|
||||
):
|
||||
result = await on_conversation_update(
|
||||
conversation_info=mock_conversation_info,
|
||||
sandbox_info=sandbox_info,
|
||||
sandbox_record=sandbox_record,
|
||||
app_conversation_info_service=app_conversation_info_service,
|
||||
)
|
||||
assert isinstance(result, Success)
|
||||
@@ -403,7 +400,7 @@ class TestOnConversationUpdateParentConversationId:
|
||||
self,
|
||||
async_session,
|
||||
app_conversation_info_service,
|
||||
sandbox_info,
|
||||
sandbox_record,
|
||||
mock_conversation_info,
|
||||
):
|
||||
"""Test that deleting conversations skips all updates including parent_conversation_id.
|
||||
@@ -444,7 +441,7 @@ class TestOnConversationUpdateParentConversationId:
|
||||
):
|
||||
result = await on_conversation_update(
|
||||
conversation_info=mock_conversation_info,
|
||||
sandbox_info=sandbox_info,
|
||||
sandbox_record=sandbox_record,
|
||||
app_conversation_info_service=app_conversation_info_service,
|
||||
)
|
||||
|
||||
@@ -464,7 +461,7 @@ class TestOnConversationUpdateParentConversationId:
|
||||
self,
|
||||
async_session,
|
||||
app_conversation_info_service,
|
||||
sandbox_info,
|
||||
sandbox_record,
|
||||
mock_conversation_info,
|
||||
):
|
||||
"""Test that parent_conversation_id is preserved when title changes.
|
||||
@@ -504,7 +501,7 @@ class TestOnConversationUpdateParentConversationId:
|
||||
):
|
||||
result = await on_conversation_update(
|
||||
conversation_info=mock_conversation_info,
|
||||
sandbox_info=sandbox_info,
|
||||
sandbox_record=sandbox_record,
|
||||
app_conversation_info_service=app_conversation_info_service,
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user