diff --git a/openhands/app_server/git/git_router.py b/openhands/app_server/git/git_router.py index 83b266e..31fccbc 100644 --- a/openhands/app_server/git/git_router.py +++ b/openhands/app_server/git/git_router.py @@ -20,7 +20,6 @@ ) from openhands.app_server.integrations.provider import ProviderHandler from openhands.app_server.integrations.service_types import ( - Branch, ProviderType, Repository, SuggestedTask, @@ -46,6 +45,19 @@ user_context_dependency = depends_user_context() +def _provider_page_number(page_id: str | None) -> int: + """Decode a provider page token, rejecting values this API never emits.""" + if page_id is None: + return 1 + page = decode_page_id(page_id) + if page is None or page <= 1: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail='Invalid page_id.', + ) + return page + + @router.get('/installations/search') async def search_user_installations( provider: ProviderType, @@ -138,6 +150,11 @@ async def search_repositories( status_code=status.HTTP_403_FORBIDDEN, # 403 not 401 to avoid frontend logout detail='Git provider token required (such as GitHub).', ) + if provider not in provider_tokens: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail='Requested Git provider is not connected.', + ) user_id = await user_context.get_user_id() # Cast to the expected type since we validated provider_tokens exists @@ -147,10 +164,7 @@ async def search_repositories( external_auth_id=user_id, ) - page = 1 - decoded_page_id = decode_page_id(page_id) - if decoded_page_id is not None: - page = decoded_page_id + page = _provider_page_number(page_id) # If query is provided, use search; otherwise get user's repositories if query: @@ -164,11 +178,17 @@ async def search_repositories( repos: list[Repository] = await client.search_repositories( selected_provider=provider, query=query, - per_page=limit + 1, + per_page=limit, sort=search_sort, order=order, app_mode=get_global_config().app_mode, + page=page, ) + # A limit+1 look-ahead changes the provider's numbered page width and + # drops that boundary item when page+1 is requested. A full page uses + # an optimistic token; the harmless terminal page may therefore be empty. + has_next_page = len(repos) >= limit + repos = repos[:limit] else: if sort_order: # TODO: This is a temporary state until we refactor the underlying API. @@ -188,11 +208,11 @@ async def search_repositories( per_page=limit + 1, installation_id=installation_id, ) + has_next_page = len(repos) > limit + if has_next_page: + repos = repos[:-1] - next_page_id = None - if len(repos) > limit: - repos = repos[:-1] - next_page_id = encode_page_id(page + 1) + next_page_id = encode_page_id(page + 1) if has_next_page else None return RepositoryPage(items=repos, next_page_id=next_page_id) @@ -229,6 +249,11 @@ async def search_branches( status_code=status.HTTP_403_FORBIDDEN, # 403 not 401 to avoid frontend logout detail='Git provider token required (such as GitHub).', ) + if provider not in provider_tokens: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail='Requested Git provider is not connected.', + ) user_id = await user_context.get_user_id() # Cast to the expected type since we validated provider_tokens exists @@ -238,41 +263,32 @@ async def search_branches( external_auth_id=user_id, ) - page = 1 - decoded_page_id = decode_page_id(page_id) - if decoded_page_id is not None: - page = decoded_page_id + page = _provider_page_number(page_id) - if query: - if page != 1: - # TODO(#13883): Support pagination for branch search after refactoring. - # The search_branches method does not support paging in the same way as - # get_branches - those should be merged into a single paginated method - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail='Pagination not yet supported for branch search queries. Use empty query to list all branches with pagination.', - ) - # Get search results - we'll handle pagination ourselves - branches: list[Branch] = await client.search_branches( - selected_provider=provider, - repository=repository, - query=query, - per_page=limit + 1, - ) - else: + try: current_page = await client.get_branches( repository=repository, specified_provider=provider, page=page, - per_page=limit + 1, + per_page=limit, + raise_on_error=True, ) - branches = current_page.branches - - next_page_id = None - if len(branches) > limit: - branches = branches[:-1] - next_page_id = encode_page_id(page + 1) - + except Exception: + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail='Git branch search is temporarily unavailable.', + ) from None + branches = current_page.branches + if query: + normalized_query = query.casefold() + branches = [ + branch + for branch in branches + if normalized_query in branch.name.casefold() + or branch.commit_sha.casefold().startswith(normalized_query) + ] + + next_page_id = encode_page_id(page + 1) if current_page.has_next_page else None return BranchPage(items=branches, next_page_id=next_page_id) diff --git a/openhands/app_server/integrations/azure_devops/service/repos.py b/openhands/app_server/integrations/azure_devops/service/repos.py index eb7793d..662b1bb 100644 --- a/openhands/app_server/integrations/azure_devops/service/repos.py +++ b/openhands/app_server/integrations/azure_devops/service/repos.py @@ -18,6 +18,7 @@ async def search_repositories( order: str = 'desc', public: bool = False, app_mode: AppMode = AppMode.OPENHANDS, + page: int = 1, ) -> list[Repository]: """Search for repositories in Azure DevOps.""" # Get all repositories across all projects in the organization @@ -32,8 +33,8 @@ async def search_repositories( repo for repo in repos if query.lower() in repo.get('name', '').lower() ] - # Limit to per_page - repos = repos[:per_page] + start = (page - 1) * per_page + repos = repos[start : start + per_page] return [ Repository( diff --git a/openhands/app_server/integrations/bitbucket/service/repos.py b/openhands/app_server/integrations/bitbucket/service/repos.py index dee1a95..7f9dbfe 100644 --- a/openhands/app_server/integrations/bitbucket/service/repos.py +++ b/openhands/app_server/integrations/bitbucket/service/repos.py @@ -20,11 +20,14 @@ async def search_repositories( order: str, public: bool, app_mode: AppMode, + page: int = 1, ) -> list[Repository]: """Search for repositories.""" repositories = [] if public: + if page != 1: + return [] # Extract workspace and repo from URL using robust URL parsing # URL format: https://{domain}/{workspace}/{repo}/{additional_params} try: @@ -54,7 +57,7 @@ async def search_repositories( if '/' in query: workspace_slug, repo_query = query.split('/', 1) return await self.get_paginated_repos( - 1, per_page, sort, workspace_slug, repo_query + page, per_page, sort, workspace_slug, repo_query ) all_installations = await self.get_installations() @@ -67,7 +70,7 @@ async def search_repositories( # Get repositories where query matches workspace name try: repos = await self.get_paginated_repos( - 1, per_page, sort, workspace_slug + page, per_page, sort, workspace_slug ) repositories.extend(repos) except Exception: @@ -77,7 +80,7 @@ async def search_repositories( # Get repositories in all workspaces where query matches repo name try: repos = await self.get_paginated_repos( - 1, per_page, sort, workspace_slug, query + page, per_page, sort, workspace_slug, query ) repositories.extend(repos) except Exception: diff --git a/openhands/app_server/integrations/bitbucket_data_center/service/repos.py b/openhands/app_server/integrations/bitbucket_data_center/service/repos.py index 09574ff..6e322c5 100644 --- a/openhands/app_server/integrations/bitbucket_data_center/service/repos.py +++ b/openhands/app_server/integrations/bitbucket_data_center/service/repos.py @@ -21,11 +21,14 @@ async def search_repositories( order: str, public: bool, app_mode: AppMode, + page: int = 1, ) -> list[Repository]: """Search for repositories.""" repositories = [] if public: + if page != 1: + return [] try: parsed_url = urlparse(query) path_segments = [ @@ -80,6 +83,8 @@ async def search_repositories( if repo_query.lower() in r.get('slug', '').lower() or repo_query.lower() in r.get('name', '').lower() ] + start = (page - 1) * per_page + raw_repos = raw_repos[start : start + per_page] return [await self._parse_repository(repo) for repo in raw_repos] # No '/' in query, search across all projects @@ -87,7 +92,7 @@ async def search_repositories( for project_key in all_projects: try: repos = await self.get_paginated_repos( - 1, per_page, sort, project_key, query + page, per_page, sort, project_key, query ) repositories.extend(repos) except Exception: diff --git a/openhands/app_server/integrations/forgejo/service/repos.py b/openhands/app_server/integrations/forgejo/service/repos.py index 364fae4..3855b00 100644 --- a/openhands/app_server/integrations/forgejo/service/repos.py +++ b/openhands/app_server/integrations/forgejo/service/repos.py @@ -16,6 +16,7 @@ async def search_repositories( order: str, public: bool, app_mode: AppMode, + page: int = 1, ) -> list[Repository]: # type: ignore[override] url = f'{self.BASE_URL}/repos/search' params = { @@ -24,6 +25,7 @@ async def search_repositories( 'sort': sort, 'order': order, 'mode': 'source', + 'page': page, } response, _ = await self._make_request(url, params) diff --git a/openhands/app_server/integrations/github/service/repos.py b/openhands/app_server/integrations/github/service/repos.py index 88f120d..3977a53 100644 --- a/openhands/app_server/integrations/github/service/repos.py +++ b/openhands/app_server/integrations/github/service/repos.py @@ -215,12 +215,14 @@ async def search_repositories( order: str, public: bool, app_mode: AppMode, + page: int = 1, ) -> list[Repository]: url = f'{self.BASE_URL}/search/repositories' params = { 'per_page': per_page, 'sort': sort, 'order': order, + 'page': page, } if public: diff --git a/openhands/app_server/integrations/gitlab/service/repos.py b/openhands/app_server/integrations/gitlab/service/repos.py index 5fb31c4..2551f37 100644 --- a/openhands/app_server/integrations/gitlab/service/repos.py +++ b/openhands/app_server/integrations/gitlab/service/repos.py @@ -81,6 +81,7 @@ async def search_repositories( order: str = 'desc', public: bool = False, app_mode: AppMode = AppMode.OPENHANDS, + page: int = 1, ) -> list[Repository]: if public: # When public=True, query is a GitLab URL that we need to parse @@ -91,7 +92,7 @@ async def search_repositories( repository = await self.get_repository_details_from_repo_name(repo_path) return [repository] - return await self.get_paginated_repos(1, per_page, sort, None, query) + return await self.get_paginated_repos(page, per_page, sort, None, query) async def get_paginated_repos( self, diff --git a/openhands/app_server/integrations/provider.py b/openhands/app_server/integrations/provider.py index aad340d..4e3f107 100644 --- a/openhands/app_server/integrations/provider.py +++ b/openhands/app_server/integrations/provider.py @@ -368,12 +368,13 @@ async def search_repositories( sort: str, order: str, app_mode: AppMode, + page: int = 1, ) -> list[Repository]: if selected_provider: service = self.get_service(selected_provider) public = self._is_repository_url(query, selected_provider) user_repos = await service.search_repositories( - query, per_page, sort, order, public, app_mode + query, per_page, sort, order, public, app_mode, page ) return self._deduplicate_repositories(user_repos) @@ -383,7 +384,7 @@ async def search_repositories( service = self.get_service(provider) public = self._is_repository_url(query, provider) service_repos = await service.search_repositories( - query, per_page, sort, order, public, app_mode + query, per_page, sort, order, public, app_mode, page ) all_repos.extend(service_repos) except Exception as e: @@ -470,6 +471,7 @@ async def get_branches( specified_provider: ProviderType | None = None, page: int = 1, per_page: int = 30, + raise_on_error: bool = False, ) -> PaginatedBranchesResponse: """Get branches for a repository @@ -487,6 +489,9 @@ async def get_branches( service = self.get_service(specified_provider) return await service.get_paginated_branches(repository, page, per_page) except Exception as e: + if raise_on_error: + logger.warning(f'Error fetching branches from {specified_provider}') + raise logger.warning( f'Error fetching branches from {specified_provider}: {e}' ) diff --git a/openhands/app_server/integrations/service_types.py b/openhands/app_server/integrations/service_types.py index b775c65..97c19ca 100644 --- a/openhands/app_server/integrations/service_types.py +++ b/openhands/app_server/integrations/service_types.py @@ -265,6 +265,7 @@ async def search_repositories( order: str, public: bool, app_mode: AppMode, + page: int = 1, ) -> list[Repository]: """Search for public repositories""" ... diff --git a/openhands/app_server/settings/settings_router.py b/openhands/app_server/settings/settings_router.py index a830b6d..858f8a2 100644 --- a/openhands/app_server/settings/settings_router.py +++ b/openhands/app_server/settings/settings_router.py @@ -10,8 +10,13 @@ from fastapi import APIRouter, Body, Depends, HTTPException, Path, status from fastapi.responses import JSONResponse -from pydantic import BaseModel, Field +from pydantic import BaseModel, ConfigDict, Field +from openhands.agent_server.mcp_router import ( + MCPTestRequest, + MCPTestSuccess, + _probe_mcp_server, +) from openhands.analytics import get_analytics_service from openhands.app_server.integrations.provider import ( PROVIDER_TOKEN_TYPE, @@ -98,6 +103,26 @@ def _merge_marketplaces( dependencies=get_dependencies(), ) +_MCP_SETTINGS_KEY_PATTERN = r'^[A-Za-z0-9][A-Za-z0-9._-]{0,127}$' +MCPSettingsKey = Annotated[ + str, + Path(min_length=1, max_length=128, pattern=_MCP_SETTINGS_KEY_PATTERN), +] + + +class StoredMCPProbeRequest(BaseModel): + """Bounded options for testing one already-stored MCP server.""" + + model_config = ConfigDict(extra='forbid') + + timeout: float = Field(default=15.0, gt=0, le=30) + + +class StoredMCPProbeResponse(BaseModel): + """Sanitized connectivity verdict for a stored MCP server.""" + + ok: bool + def _post_merge_llm_fixups(settings: Settings) -> None: """Apply LLM-specific fixups after merging settings. @@ -354,6 +379,62 @@ async def load_conversation_settings_schema() -> dict[str, Any]: return ConversationSettings.export_schema().model_dump(mode='json') +@router.post( + '/mcp/{settings_key}/test', + response_model=StoredMCPProbeResponse, +) +async def test_stored_mcp_server( + settings_key: MCPSettingsKey, + request: StoredMCPProbeRequest, + settings: Settings | None = Depends(get_user_settings), +) -> StoredMCPProbeResponse: + """Test a user's stored MCP server without returning its config or errors. + + The server configuration is resolved from authenticated user settings and + stays inside the app server. The detailed MCP result is deliberately reduced + to a boolean so credentials, provider messages, tools, and stack traces never + cross this API boundary. + """ + if settings is None: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail='MCP server was not found', + ) + server = settings.agent_settings.mcp_config.get(settings_key) + if server is None: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail='MCP server was not found', + ) + + server_data = server.model_dump( + mode='json', + context={'expose_secrets': 'plaintext'}, + exclude_none=True, + exclude_defaults=True, + ) + probe_request = MCPTestRequest( + name=settings_key, + server=server_data, + timeout=request.timeout, + ) + loop = asyncio.get_running_loop() + try: + result = await loop.run_in_executor( + None, + _probe_mcp_server, + probe_request, + None, + ) + except Exception: + logger.warning('Stored MCP validation failed unexpectedly') + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail='MCP validation is temporarily unavailable', + ) from None + return StoredMCPProbeResponse(ok=isinstance(result, MCPTestSuccess)) + + async def invalidate_legacy_secrets_store( settings: Settings, settings_store: SettingsStore, secrets_store: SecretsStore ) -> Secrets | None: diff --git a/tests/unit/app_server/test_git_router.py b/tests/unit/app_server/test_git_router.py index f2c8468..c370b6a 100644 --- a/tests/unit/app_server/test_git_router.py +++ b/tests/unit/app_server/test_git_router.py @@ -7,7 +7,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest -from fastapi import FastAPI, status +from fastapi import FastAPI, HTTPException, status from fastapi.testclient import TestClient from openhands.app_server.git.git_models import SortOrder @@ -21,6 +21,7 @@ from openhands.app_server.integrations.provider import ProviderToken from openhands.app_server.integrations.service_types import ( Branch, + PaginatedBranchesResponse, ProviderType, Repository, SuggestedTask, @@ -330,8 +331,9 @@ async def test_search_repositories_with_query(self, mock_handler_cls): assert call_kwargs.get('sort') == 'stars' assert call_kwargs.get('order') == 'desc' - # Verify per_page is limit + 1 - assert call_kwargs.get('per_page') == 11 + # Provider page width must match the public page width so page 2 starts + # immediately after the last item returned on page 1. + assert call_kwargs.get('per_page') == 10 # Verify results are returned assert len(result.items) == 2 @@ -623,6 +625,122 @@ async def test_returns_paginated_search_results(self, mock_handler_cls): ] assert result.next_page_id == encode_page_id(2) + @pytest.mark.asyncio + @patch('openhands.app_server.git.git_router.ProviderHandler') + async def test_search_pagination_keeps_the_provider_boundary_item( + self, mock_handler_cls + ): + """The item used to infer a next page must appear on one returned page.""" + repositories = [ + Repository( + id=str(index), + full_name=f'user/repo{index}', + git_provider=ProviderType.GITHUB, + is_public=True, + ) + for index in range(1, 4) + ] + + async def provider_page(**kwargs): + page = kwargs['page'] + per_page = kwargs['per_page'] + start = (page - 1) * per_page + return repositories[start : start + per_page] + + mock_handler = MagicMock() + mock_handler.search_repositories = AsyncMock(side_effect=provider_page) + mock_handler_cls.return_value = mock_handler + mock_context = _make_mock_user_context( + provider_tokens={ + ProviderType.GITHUB: ProviderToken(user_id='user-123', token='token') + }, + user_id='user-123', + ) + + first = await search_repositories( + provider=ProviderType.GITHUB, + query='user/repo', + page_id=None, + limit=2, + sort_order=None, + user_context=mock_context, + ) + second = await search_repositories( + provider=ProviderType.GITHUB, + query='user/repo', + page_id=first.next_page_id, + limit=2, + sort_order=None, + user_context=mock_context, + ) + + assert [repo.id for repo in [*first.items, *second.items]] == ['1', '2', '3'] + + @pytest.mark.asyncio + @patch('openhands.app_server.git.git_router.ProviderHandler') + async def test_search_query_passes_the_decoded_page_to_provider( + self, mock_handler_cls + ): + """A query page token must not silently replay the first provider page.""" + mock_handler = MagicMock() + mock_handler.search_repositories = AsyncMock(return_value=[]) + mock_handler_cls.return_value = mock_handler + mock_context = _make_mock_user_context( + provider_tokens={ + ProviderType.GITHUB: ProviderToken(user_id='user-123', token='token') + }, + user_id='user-123', + ) + + await search_repositories( + provider=ProviderType.GITHUB, + query='user/repo', + page_id=encode_page_id(2), + limit=20, + sort_order=None, + user_context=mock_context, + ) + + assert mock_handler.search_repositories.call_args.kwargs['page'] == 2 + + @pytest.mark.parametrize( + 'page_id', + [ + '', + encode_page_id(-1), + encode_page_id(0), + encode_page_id(1), + 'not-a-page-token', + ], + ) + @pytest.mark.asyncio + @patch('openhands.app_server.git.git_router.ProviderHandler') + async def test_rejects_non_emitted_or_malformed_query_page( + self, mock_handler_cls, page_id + ): + mock_handler = MagicMock() + mock_handler.search_repositories = AsyncMock(return_value=[]) + mock_handler_cls.return_value = mock_handler + mock_context = _make_mock_user_context( + provider_tokens={ + ProviderType.GITHUB: ProviderToken(user_id='user-123', token='token') + }, + user_id='user-123', + ) + + with pytest.raises(HTTPException) as exc_info: + await search_repositories( + provider=ProviderType.GITHUB, + query='user/repo', + page_id=page_id, + limit=20, + sort_order=None, + user_context=mock_context, + ) + + assert exc_info.value.status_code == status.HTTP_400_BAD_REQUEST + mock_handler.search_repositories.assert_not_awaited() + @pytest.mark.asyncio @patch('openhands.app_server.git.git_router.ProviderHandler') async def test_parses_sort_order_correctly(self, mock_handler_cls): @@ -697,6 +815,28 @@ def test_returns_403_when_no_provider_tokens(self, test_client, monkeypatch): ) assert response.status_code == status.HTTP_403_FORBIDDEN + @pytest.mark.asyncio + async def test_returns_403_when_selected_provider_is_not_connected(self): + mock_context = _make_mock_user_context( + provider_tokens={ + ProviderType.GITHUB: ProviderToken(user_id='user-123', token='token') + }, + user_id='user-123', + ) + + with pytest.raises(HTTPException) as exc_info: + await search_repositories( + provider=ProviderType.GITLAB, + query='user/repo', + page_id=None, + limit=20, + sort_order=None, + user_context=mock_context, + ) + + assert exc_info.value.status_code == status.HTTP_403_FORBIDDEN + assert exc_info.value.detail == 'Requested Git provider is not connected.' + @pytest.mark.asyncio class TestSearchBranches: @@ -708,12 +848,16 @@ async def test_returns_paginated_branches(self, mock_handler_cls): """Test that search branches are returned with pagination.""" # Arrange mock_handler = MagicMock() - mock_handler.search_branches = AsyncMock( - return_value=[ - Branch(name='main', commit_sha='abc123', protected=False), - Branch(name='develop', commit_sha='def456', protected=False), - Branch(name='feature-branch', commit_sha='ghi789', protected=False), - ] + mock_handler.get_branches = AsyncMock( + return_value=PaginatedBranchesResponse( + branches=[ + Branch(name='main', commit_sha='abc123', protected=False), + Branch(name='develop', commit_sha='def456', protected=False), + ], + has_next_page=True, + current_page=1, + per_page=2, + ) ) mock_handler_cls.return_value = mock_handler @@ -735,9 +879,8 @@ async def test_returns_paginated_branches(self, mock_handler_cls): ) # Assert - assert len(result.items) == 2 + assert len(result.items) == 1 assert result.items[0].name == 'main' - assert result.items[1].name == 'develop' assert result.next_page_id == encode_page_id(2) @pytest.mark.asyncio @@ -746,7 +889,14 @@ async def test_passes_parameters_to_provider(self, mock_handler_cls): """Test that all parameters are passed through to the provider.""" # Arrange mock_handler = MagicMock() - mock_handler.search_branches = AsyncMock(return_value=[]) + mock_handler.get_branches = AsyncMock( + return_value=PaginatedBranchesResponse( + branches=[], + has_next_page=False, + current_page=1, + per_page=10, + ) + ) mock_handler_cls.return_value = mock_handler mock_context = _make_mock_user_context( @@ -767,12 +917,84 @@ async def test_passes_parameters_to_provider(self, mock_handler_cls): ) # Assert - mock_handler.search_branches.assert_called_once() - call_kwargs = mock_handler.search_branches.call_args.kwargs - assert call_kwargs.get('selected_provider') == ProviderType.GITHUB + mock_handler.get_branches.assert_called_once() + call_kwargs = mock_handler.get_branches.call_args.kwargs + assert call_kwargs.get('specified_provider') == ProviderType.GITHUB assert call_kwargs.get('repository') == 'user/repo' - assert call_kwargs.get('query') == 'feature' - assert call_kwargs.get('per_page') == 11 # limit + 1 + assert call_kwargs.get('page') == 1 + assert call_kwargs.get('per_page') == 10 + + @pytest.mark.asyncio + @patch('openhands.app_server.git.git_router.ProviderHandler') + async def test_query_can_continue_on_a_later_branch_page(self, mock_handler_cls): + """A branch search page token scans that provider page and stays paginated.""" + mock_handler = MagicMock() + mock_handler.get_branches = AsyncMock( + return_value=PaginatedBranchesResponse( + branches=[ + Branch(name='release/2.0', commit_sha='abc123', protected=True), + Branch(name='feature/x', commit_sha='def456', protected=False), + ], + has_next_page=True, + current_page=2, + per_page=20, + ) + ) + mock_handler_cls.return_value = mock_handler + mock_context = _make_mock_user_context( + provider_tokens={ + ProviderType.GITHUB: ProviderToken(user_id='user-123', token='token') + }, + user_id='user-123', + ) + + result = await search_branches( + provider=ProviderType.GITHUB, + repository='user/repo', + query='abc1', + page_id=encode_page_id(2), + limit=20, + user_context=mock_context, + ) + + assert [branch.name for branch in result.items] == ['release/2.0'] + assert result.next_page_id == encode_page_id(3) + mock_handler.get_branches.assert_awaited_once_with( + repository='user/repo', + specified_provider=ProviderType.GITHUB, + page=2, + per_page=20, + raise_on_error=True, + ) + + @pytest.mark.asyncio + @patch('openhands.app_server.git.git_router.ProviderHandler') + async def test_provider_failure_is_a_sanitized_503(self, mock_handler_cls): + """A provider outage must not masquerade as an empty branch result.""" + sentinel = 'provider-internal-secret-sentinel' + mock_handler = MagicMock() + mock_handler.get_branches = AsyncMock(side_effect=RuntimeError(sentinel)) + mock_handler_cls.return_value = mock_handler + mock_context = _make_mock_user_context( + provider_tokens={ + ProviderType.GITHUB: ProviderToken(user_id='user-123', token='token') + }, + user_id='user-123', + ) + + with pytest.raises(HTTPException) as exc_info: + await search_branches( + provider=ProviderType.GITHUB, + repository='user/repo', + query='main', + page_id=None, + limit=20, + user_context=mock_context, + ) + + assert exc_info.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE + assert exc_info.value.detail == 'Git branch search is temporarily unavailable.' + assert sentinel not in str(exc_info.value.detail) def test_returns_403_when_no_provider_tokens(self, test_client): """Test that 403 is returned when no provider tokens.""" @@ -790,6 +1012,28 @@ def test_returns_403_when_no_provider_tokens(self, test_client): ) assert response.status_code == status.HTTP_403_FORBIDDEN + @pytest.mark.asyncio + async def test_returns_403_when_selected_provider_is_not_connected(self): + mock_context = _make_mock_user_context( + provider_tokens={ + ProviderType.GITHUB: ProviderToken(user_id='user-123', token='token') + }, + user_id='user-123', + ) + + with pytest.raises(HTTPException) as exc_info: + await search_branches( + provider=ProviderType.GITLAB, + repository='user/repo', + query='main', + page_id=None, + limit=20, + user_context=mock_context, + ) + + assert exc_info.value.status_code == status.HTTP_403_FORBIDDEN + assert exc_info.value.detail == 'Requested Git provider is not connected.' + @pytest.mark.asyncio class TestSearchSuggestedTasks: diff --git a/tests/unit/app_server/test_settings_api.py b/tests/unit/app_server/test_settings_api.py index ac627af..945eef7 100644 --- a/tests/unit/app_server/test_settings_api.py +++ b/tests/unit/app_server/test_settings_api.py @@ -6,6 +6,7 @@ from fastapi.testclient import TestClient from pydantic import SecretStr +from openhands.agent_server.mcp_router import MCPTestFailure, MCPTestSuccess from openhands.app_server.app import app from openhands.app_server.file_store.memory import InMemoryFileStore from openhands.app_server.integrations.provider import ProviderToken, ProviderType @@ -145,6 +146,114 @@ def test_get_conversation_settings_schema_endpoint(test_client): assert 'security_analyzer' in field_keys +def _store_preflight_mcp_server( + test_client: TestClient, + *, + secret: str = 'mcp-secret-sentinel', +) -> None: + response = test_client.post( + '/api/v1/settings', + json={ + 'agent_settings_diff': { + 'mcp_config': { + 'github': { + 'transport': 'http', + 'url': 'https://mcp.example.test/github', + 'auth': {'strategy': 'api_key', 'value': secret}, + } + } + } + }, + ) + assert response.status_code == 200 + + +def test_stored_mcp_preflight_probes_server_without_exposing_details(test_client): + secret = 'mcp-secret-sentinel' + _store_preflight_mcp_server(test_client, secret=secret) + + def probe(request, cipher): + assert cipher is None + auth = request.resolved_server.auth + assert auth is not None + assert auth.value.get_secret_value() == secret + return MCPTestSuccess(tools=['provider-tool']) + + with patch( + 'openhands.app_server.settings.settings_router._probe_mcp_server', + side_effect=probe, + ) as probe_mock: + response = test_client.post( + '/api/v1/settings/mcp/github/test', + json={'timeout': 10}, + ) + + assert response.status_code == 200 + assert response.json() == {'ok': True} + assert secret not in response.text + assert 'provider-tool' not in response.text + probe_mock.assert_called_once() + + +def test_stored_mcp_preflight_sanitizes_connection_failure(test_client): + secret = 'mcp-secret-sentinel' + provider_error = 'provider-internal-error-sentinel' + _store_preflight_mcp_server(test_client, secret=secret) + + with patch( + 'openhands.app_server.settings.settings_router._probe_mcp_server', + return_value=MCPTestFailure( + error=provider_error, + error_kind='connection', + ), + ): + response = test_client.post('/api/v1/settings/mcp/github/test', json={}) + + assert response.status_code == 200 + assert response.json() == {'ok': False} + assert secret not in response.text + assert provider_error not in response.text + + +def test_stored_mcp_preflight_rejects_missing_server(test_client): + missing = test_client.post('/api/v1/settings/mcp/missing/test', json={}) + + assert missing.status_code == 404 + assert missing.json() == {'detail': 'MCP server was not found'} + + +def test_stored_mcp_preflight_returns_sanitized_503_on_internal_failure(test_client): + secret = 'mcp-secret-sentinel' + internal_error = 'internal-stack-sentinel' + _store_preflight_mcp_server(test_client, secret=secret) + + with patch( + 'openhands.app_server.settings.settings_router._probe_mcp_server', + side_effect=RuntimeError(internal_error), + ): + response = test_client.post('/api/v1/settings/mcp/github/test', json={}) + + assert response.status_code == 503 + assert response.json() == {'detail': 'MCP validation is temporarily unavailable'} + assert secret not in response.text + assert internal_error not in response.text + + +@pytest.mark.parametrize( + ('path', 'body'), + [ + ('bad%2Fname', {}), + ('github', {'timeout': 0}), + ('github', {'timeout': 31}), + ('github', {'timeout': 10, 'secret': 'must-not-be-accepted'}), + ], +) +def test_stored_mcp_preflight_rejects_unbounded_input(test_client, path, body): + response = test_client.post(f'/api/v1/settings/mcp/{path}/test', json=body) + + assert response.status_code in {404, 422} + + @pytest.mark.asyncio async def test_settings_api_endpoints(test_client): """Test that the settings API endpoints work with the new auth system."""