From e4a6c4ae477c4e78face6cf20bbcf0469a790022 Mon Sep 17 00:00:00 2001 From: Tim O'Farrell Date: Thu, 30 Apr 2026 11:35:59 -0600 Subject: [PATCH] feat: introduce LLMModelService with DI, replacing legacy config-based model discovery (#14237) Co-authored-by: openhands --- enterprise/saas_server.py | 7 - .../verified_models/verified_model_router.py | 61 +++- openhands/app_server/config.py | 39 +++ .../app_server/config_api/config_router.py | 121 ++------ .../config_api/default_llm_model_service.py | 241 +++++++++++++++ .../config_api/llm_model_service.py | 61 ++++ .../app_server/settings/settings_models.py | 91 ------ .../app_server/user_auth/default_user_auth.py | 5 - openhands/app_server/utils/llm.py | 66 +--- tests/unit/app_server/test_config_router.py | 29 +- .../unit/app_server/test_llm_model_service.py | 289 ++++++++++++++++++ .../core/config/test_mcp_settings_merge.py | 184 ----------- tests/unit/mcp/test_mcp_integration.py | 66 +--- .../unit/storage/data_models/test_settings.py | 53 ---- .../settings/test_file_settings_store.py | 8 +- 15 files changed, 732 insertions(+), 589 deletions(-) create mode 100644 openhands/app_server/config_api/default_llm_model_service.py create mode 100644 openhands/app_server/config_api/llm_model_service.py create mode 100644 tests/unit/app_server/test_llm_model_service.py delete mode 100644 tests/unit/core/config/test_mcp_settings_merge.py diff --git a/enterprise/saas_server.py b/enterprise/saas_server.py index 048fad957d..77080bf422 100644 --- a/enterprise/saas_server.py +++ b/enterprise/saas_server.py @@ -56,9 +56,6 @@ from server.sharing.shared_event_router import ( # noqa: E402 from server.verified_models.verified_model_router import ( # noqa: E402 api_router as verified_models_router, ) -from server.verified_models.verified_model_router import ( # noqa: E402 - override_llm_models_dependency, -) from openhands.server.app import app as base_app # noqa: E402 from openhands.server.middleware import ( # noqa: E402 @@ -135,10 +132,6 @@ base_app.include_router( verified_models_router ) # Add routes for verified models management -# Override the default LLM models implementation with SaaS version -# This must happen after all routers are included -override_llm_models_dependency(base_app) - # Override the /api/v1/users/me endpoint to include organization info # This replaces the OSS endpoint with a SAAS version that adds org_id, org_name, role, permissions override_users_me_endpoint(base_app) diff --git a/enterprise/server/verified_models/verified_model_router.py b/enterprise/server/verified_models/verified_model_router.py index 3c7cb625e3..33a89b1e6d 100644 --- a/enterprise/server/verified_models/verified_model_router.py +++ b/enterprise/server/verified_models/verified_model_router.py @@ -1,6 +1,7 @@ """API routes for managing verified LLM models (admin only).""" -from typing import Annotated +import logging +from typing import Annotated, AsyncGenerator from fastapi import APIRouter, Depends, HTTPException, Query, Request, status from server.email_validation import get_admin_user_id @@ -16,9 +17,18 @@ from server.verified_models.verified_model_service import ( ) from openhands.app_server.config import get_db_session -from openhands.app_server.config_api.config_router import get_llm_models_dependency +from openhands.app_server.config_api.default_llm_model_service import ( + DefaultLLMModelService, +) +from openhands.app_server.config_api.llm_model_service import ( + LLMModelService, + LLMModelServiceInjector, +) +from openhands.app_server.services.injector import InjectorState from openhands.app_server.utils.llm import ModelsResponse, get_supported_llm_models +_logger = logging.getLogger(__name__) + api_router = APIRouter(prefix='/api/admin/verified-models', tags=['Verified Models']) @@ -117,25 +127,46 @@ async def delete_verified_model( ) -async def get_saas_llm_models_dependency(request: Request) -> ModelsResponse: - """SaaS implementation for the LLM models endpoint.""" - async with get_db_session(request.state, request) as db_session: - # Prevent circular import - from openhands.server.shared import config +class SaaSLLMModelService(DefaultLLMModelService): + """SaaS implementation that reads verified models from the database. - verified_model_service = VerifiedModelService(db_session) + Inherits filtering, pagination, and provider logic from + ``DefaultLLMModelService`` — only the verified-model list is different. + """ + + def __init__(self, db_session) -> None: + super().__init__() + self._db_session = db_session + + async def _get_models_response( + self, + verified_models: list[str] | None = None, + ) -> ModelsResponse: + if self._cached_response is not None: + return self._cached_response + + verified_model_service = VerifiedModelService(self._db_session) page = await verified_model_service.search_verified_models(enabled_only=True) if page.next_page_id: raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail='Too many models defined in database', ) - verified_models = [f'{m.provider}/{m.model_name}' for m in page.items] - return get_supported_llm_models(config, verified_models) + db_verified = [f'{m.provider}/{m.model_name}' for m in page.items] + self._cached_response = get_supported_llm_models(db_verified) + return self._cached_response -# Override the default implementation with SaaS implementation -# This must be called after the app is created in saas_server.py -def override_llm_models_dependency(app): - """Override the default LLM models implementation with SaaS version.""" - app.dependency_overrides[get_llm_models_dependency] = get_saas_llm_models_dependency +class SaaSLLMModelServiceInjector(LLMModelServiceInjector): + """Injector that provides the SaaS LLM model service. + + Activate via the environment variable:: + + OH_LLM_MODEL_KIND=server.verified_models.verified_model_router.SaaSLLMModelServiceInjector + """ + + async def inject( + self, state: InjectorState, request: Request | None = None + ) -> AsyncGenerator[LLMModelService, None]: + async with get_db_session(state, request) as db_session: + yield SaaSLLMModelService(db_session) diff --git a/openhands/app_server/config.py b/openhands/app_server/config.py index df4517308c..b1f4cce517 100644 --- a/openhands/app_server/config.py +++ b/openhands/app_server/config.py @@ -28,6 +28,10 @@ from openhands.app_server.app_lifespan.app_lifespan_service import AppLifespanSe from openhands.app_server.app_lifespan.oss_app_lifespan_service import ( OssAppLifespanService, ) +from openhands.app_server.config_api.llm_model_service import ( + LLMModelService, + LLMModelServiceInjector, +) from openhands.app_server.event.event_service import EventService, EventServiceInjector from openhands.app_server.event_callback.event_callback_service import ( EventCallbackService, @@ -183,6 +187,7 @@ class AppServerConfig(OpenHandsModel): description='Base URL for the OpenHands provider', ) # Dependency Injection Injectors + llm_model: LLMModelServiceInjector | None = None event: EventServiceInjector | None = None event_callback: EventCallbackServiceInjector | None = None sandbox: SandboxServiceInjector | None = None @@ -254,6 +259,26 @@ def config_from_env() -> AppServerConfig: config: AppServerConfig = from_env(AppServerConfig, 'OH') # type: ignore + if config.llm_model is None: + from openhands.app_server.config_api.default_llm_model_service import ( + DefaultLLMModelServiceInjector, + ) + + llm_model_kwargs: dict = {} + aws_region = os.getenv('AWS_REGION_NAME') + aws_key = os.getenv('AWS_ACCESS_KEY_ID') + aws_secret = os.getenv('AWS_SECRET_ACCESS_KEY') + if aws_region and aws_key and aws_secret: + llm_model_kwargs['aws_region_name'] = aws_region + llm_model_kwargs['aws_access_key_id'] = SecretStr(aws_key) + llm_model_kwargs['aws_secret_access_key'] = SecretStr(aws_secret) + + ollama_url = os.getenv('OLLAMA_BASE_URL') + if ollama_url: + llm_model_kwargs['ollama_base_url'] = ollama_url + + config.llm_model = DefaultLLMModelServiceInjector(**llm_model_kwargs) + if config.event is None: provider = get_storage_provider() @@ -552,3 +577,17 @@ def depends_jwt_service(): def depends_db_session(): return Depends(get_global_config().db_session.depends) + + +def get_llm_model_service( + state: InjectorState, request: Request | None = None +) -> AsyncContextManager[LLMModelService]: + injector = get_global_config().llm_model + assert injector is not None + return injector.context(state, request) + + +def depends_llm_model_service(): + injector = get_global_config().llm_model + assert injector is not None + return Depends(injector.depends) diff --git a/openhands/app_server/config_api/config_router.py b/openhands/app_server/config_api/config_router.py index 4e3e87ec3d..28d9cac76b 100644 --- a/openhands/app_server/config_api/config_router.py +++ b/openhands/app_server/config_api/config_router.py @@ -6,31 +6,14 @@ provider search with pagination support. from typing import Annotated -from fastapi import APIRouter, Depends, Query, Request +from fastapi import APIRouter, Query -from openhands.app_server.config_api.config_models import ( - LLMModel, - LLMModelPage, - Provider, - ProviderPage, -) +from openhands.app_server.config import depends_llm_model_service +from openhands.app_server.config_api.config_models import LLMModelPage, ProviderPage +from openhands.app_server.config_api.llm_model_service import LLMModelService from openhands.app_server.utils.dependencies import get_dependencies -from openhands.app_server.utils.llm import ModelsResponse, get_supported_llm_models -from openhands.app_server.utils.paging_utils import ( - paginate_results, -) -from openhands.sdk.llm.utils.verified_models import VERIFIED_MODELS -from openhands.server.shared import config - - -async def get_llm_models_dependency(request: Request) -> ModelsResponse: - """Returns a callable that provides the LLM models implementation. - - Returns a factory that produces the actual implementation function. - Override this in enterprise/saas mode via app.dependency_overrides. - """ - return get_supported_llm_models(config) +llm_model_service_dependency = depends_llm_model_service() # We use the get_dependencies method here to signal to the OpenAPI docs that this endpoint # is protected. The actual protection is provided by SetAuthCookieMiddleware @@ -63,29 +46,20 @@ async def search_models( str | None, Query(title='Filter by provider name (exact match)'), ] = None, - models: ModelsResponse = Depends(get_llm_models_dependency), + llm_model_service: LLMModelService = llm_model_service_dependency, ) -> LLMModelPage: """Search for LLM models with pagination and filtering. Returns a paginated list of models that can be filtered by name (contains), verified status, and provider. """ - filtered_models = _get_all_models_with_verified(models) - - if query is not None: - query_lower = query.lower() - filtered_models = [m for m in filtered_models if query_lower in m.name.lower()] - - if verified__eq is not None: - filtered_models = [m for m in filtered_models if m.verified == verified__eq] - - if provider__eq is not None: - filtered_models = [m for m in filtered_models if m.provider == provider__eq] - - # Apply pagination - items, next_page_id = paginate_results(filtered_models, page_id, limit) - - return LLMModelPage(items=items, next_page_id=next_page_id) + return await llm_model_service.search_llm_models( + query=query, + verified_eq=verified__eq, + provider_eq=provider__eq, + page_id=page_id, + limit=limit, + ) @router.get('/providers/search') @@ -106,71 +80,16 @@ async def search_providers( bool | None, Query(title='Filter by verified status (true/false, omit for all)'), ] = None, - models: ModelsResponse = Depends(get_llm_models_dependency), + llm_model_service: LLMModelService = llm_model_service_dependency, ) -> ProviderPage: """Search for LLM providers with pagination and filtering. Returns a paginated list of providers extracted from the available models. Each provider indicates whether it is verified by OpenHands. """ - providers = _get_all_providers(models) - - if query is not None: - query_lower = query.lower() - providers = [p for p in providers if query_lower in p.name.lower()] - - if verified__eq is not None: - providers = [p for p in providers if p.verified == verified__eq] - - items, next_page_id = paginate_results(providers, page_id, limit) - - return ProviderPage(items=items, next_page_id=next_page_id) - - -def _get_verified_models() -> set[str]: - verified_models = set() - for provider, models in VERIFIED_MODELS.items(): - for name in models: - verified_models.add(f'{provider}/{name}') - return verified_models - - -def _get_all_models_with_verified(models: ModelsResponse) -> list[LLMModel]: - verified_models = _get_verified_models() - results = [] - for model_name in models.models: - verified = model_name in verified_models - parts = model_name.split('/', 1) - if len(parts) == 2: - provider, name = parts - else: - provider = None - name = parts[0] - result = LLMModel( - provider=provider, - name=name, - verified=verified, - ) - results.append(result) - return results - - -def _get_all_providers(models: ModelsResponse) -> list[Provider]: - """Extract unique providers from the models list, sorted verified-first.""" - verified_set = set(models.verified_providers) - seen: set[str] = set() - providers: list[Provider] = [] - - for model_name in models.models: - parts = model_name.split('/', 1) - if len(parts) == 2: - name = parts[0] - else: - continue # skip bare model names without a provider - if name not in seen: - seen.add(name) - providers.append(Provider(name=name, verified=name in verified_set)) - - # Sort: verified providers first, then alphabetically within each group - providers.sort(key=lambda p: (not p.verified, p.name)) - return providers + return await llm_model_service.search_providers( + query=query, + verified_eq=verified__eq, + page_id=page_id, + limit=limit, + ) diff --git a/openhands/app_server/config_api/default_llm_model_service.py b/openhands/app_server/config_api/default_llm_model_service.py new file mode 100644 index 0000000000..1bf22bc4b6 --- /dev/null +++ b/openhands/app_server/config_api/default_llm_model_service.py @@ -0,0 +1,241 @@ +"""Default LLM model discovery service. + +Discovers models from litellm's built-in catalogue, optional AWS Bedrock, +and optional Ollama instances. Filtering and pagination are applied +in-memory so that the router stays thin. +""" + +import logging +from typing import Any, AsyncGenerator + +import httpx +from fastapi import Request +from pydantic import Field, SecretStr + +from openhands.app_server.config_api.config_models import ( + LLMModel, + LLMModelPage, + Provider, + ProviderPage, +) +from openhands.app_server.config_api.llm_model_service import ( + LLMModelService, + LLMModelServiceInjector, +) +from openhands.app_server.services.injector import InjectorState +from openhands.app_server.utils.async_utils import call_sync_from_async +from openhands.app_server.utils.llm import ( + ModelsResponse, + get_supported_llm_models, +) +from openhands.app_server.utils.paging_utils import paginate_results +from openhands.sdk.llm.utils.verified_models import VERIFIED_MODELS + +_logger = logging.getLogger(__name__) + +_VERIFIED_MODEL_SET: set[str] = { + f'{provider}/{name}' + for provider, models in VERIFIED_MODELS.items() + for name in models +} + + +def _to_llm_models(models_response: ModelsResponse) -> list[LLMModel]: + """Convert raw model strings into ``LLMModel`` objects with verified flags.""" + results: list[LLMModel] = [] + for model_name in models_response.models: + parts = model_name.split('/', 1) + if len(parts) == 2: + provider, name = parts + else: + provider = None + name = parts[0] + results.append( + LLMModel( + provider=provider, + name=name, + verified=model_name in _VERIFIED_MODEL_SET, + ) + ) + return results + + +def _to_providers(models_response: ModelsResponse) -> list[Provider]: + """Extract unique providers, sorted verified-first then alphabetically.""" + verified_set = set(models_response.verified_providers) + seen: set[str] = set() + providers: list[Provider] = [] + for model_name in models_response.models: + parts = model_name.split('/', 1) + if len(parts) != 2: + continue + name = parts[0] + if name not in seen: + seen.add(name) + providers.append(Provider(name=name, verified=name in verified_set)) + providers.sort(key=lambda p: (not p.verified, p.name)) + return providers + + +class DefaultLLMModelService(LLMModelService): + """Model discovery via litellm catalogue, optional Bedrock, and optional Ollama.""" + + def __init__( + self, + *, + bedrock_client: Any | None = None, + ollama_base_url: str | None = None, + ) -> None: + self._bedrock_client = bedrock_client + self._ollama_base_url = ollama_base_url + self._cached_response: ModelsResponse | None = None + + def _list_foundation_models(self) -> list[str]: + """Query AWS Bedrock for available foundation models. + + This is a synchronous boto3 call; callers should run it via + ``call_sync_from_async`` to avoid blocking the event loop. + """ + if self._bedrock_client is None: + return [] + try: + response = self._bedrock_client.list_foundation_models( + byOutputModality='TEXT', byInferenceType='ON_DEMAND' + ) + return ['bedrock/' + m['modelId'] for m in response['modelSummaries']] + except Exception as e: + _logger.warning( + '%s. Please config AWS_REGION_NAME AWS_ACCESS_KEY_ID' + ' AWS_SECRET_ACCESS_KEY if you want use bedrock model.', + e, + ) + return [] + + async def _get_models_response( + self, + verified_models: list[str] | None = None, + ) -> ModelsResponse: + """Fetch the raw ``ModelsResponse`` from all configured sources. + + The result is cached on the service instance so that multiple + calls (e.g. ``search_llm_models`` + ``search_providers``) within + the same request do not repeat expensive discovery work. + """ + if self._cached_response is not None: + return self._cached_response + + extra_models: list[str] = [] + + if self._bedrock_client is not None: + bedrock_models: list[str] = await call_sync_from_async( + self._list_foundation_models + ) + extra_models.extend(bedrock_models) + + if self._ollama_base_url: + ollama_url = self._ollama_base_url.strip('/') + '/api/tags' + try: + async with httpx.AsyncClient() as client: + resp = await client.get(ollama_url, timeout=3) + ollama_models_list = resp.json()['models'] + extra_models.extend('ollama/' + m['name'] for m in ollama_models_list) + except httpx.HTTPError as e: + _logger.error(f'Error getting OLLAMA models: {e}') + + self._cached_response = get_supported_llm_models( + verified_models=verified_models, + extra_models=extra_models or None, + ) + return self._cached_response + + # ------------------------------------------------------------------ + # LLMModelService interface + # ------------------------------------------------------------------ + + async def search_llm_models( + self, + *, + query: str | None = None, + verified_eq: bool | None = None, + provider_eq: str | None = None, + page_id: str | None = None, + limit: int = 50, + ) -> LLMModelPage: + raw = await self._get_models_response() + models = _to_llm_models(raw) + + if query is not None: + query_lower = query.lower() + models = [m for m in models if query_lower in m.name.lower()] + if verified_eq is not None: + models = [m for m in models if m.verified == verified_eq] + if provider_eq is not None: + models = [m for m in models if m.provider == provider_eq] + + items, next_page_id = paginate_results(models, page_id, limit) + return LLMModelPage(items=items, next_page_id=next_page_id) + + async def search_providers( + self, + *, + query: str | None = None, + verified_eq: bool | None = None, + page_id: str | None = None, + limit: int = 50, + ) -> ProviderPage: + raw = await self._get_models_response() + providers = _to_providers(raw) + + if query is not None: + query_lower = query.lower() + providers = [p for p in providers if query_lower in p.name.lower()] + if verified_eq is not None: + providers = [p for p in providers if p.verified == verified_eq] + + items, next_page_id = paginate_results(providers, page_id, limit) + return ProviderPage(items=items, next_page_id=next_page_id) + + +class DefaultLLMModelServiceInjector(LLMModelServiceInjector): + """Injector that reads AWS / Ollama credentials from its own fields. + + When AWS credentials are provided, a ``boto3`` Bedrock client is created + once and passed to every service instance, avoiding repeated credential + negotiation. + """ + + aws_region_name: str | None = None + aws_access_key_id: SecretStr | None = None + aws_secret_access_key: SecretStr | None = None + ollama_base_url: str | None = Field( + default=None, + description='Base URL for a local Ollama instance (e.g. http://localhost:11434)', + ) + + _bedrock_client: Any | None = None + + def _get_bedrock_client(self) -> Any | None: + if self._bedrock_client is not None: + return self._bedrock_client + if ( + self.aws_region_name + and self.aws_access_key_id + and self.aws_secret_access_key + ): + import boto3 + + self._bedrock_client = boto3.client( + service_name='bedrock', + region_name=self.aws_region_name, + aws_access_key_id=self.aws_access_key_id.get_secret_value(), + aws_secret_access_key=self.aws_secret_access_key.get_secret_value(), + ) + return self._bedrock_client + + async def inject( + self, state: InjectorState, request: Request | None = None + ) -> AsyncGenerator[LLMModelService, None]: + yield DefaultLLMModelService( + bedrock_client=self._get_bedrock_client(), + ollama_base_url=self.ollama_base_url, + ) diff --git a/openhands/app_server/config_api/llm_model_service.py b/openhands/app_server/config_api/llm_model_service.py new file mode 100644 index 0000000000..b78ff4c0be --- /dev/null +++ b/openhands/app_server/config_api/llm_model_service.py @@ -0,0 +1,61 @@ +"""LLM model discovery service. + +Provides an abstract interface for discovering available LLM models. +Concrete implementations handle different model sources (litellm, AWS Bedrock, +database-backed verified models for SaaS, etc.). +""" + +from abc import ABC, abstractmethod + +from openhands.app_server.config_api.config_models import LLMModelPage, ProviderPage +from openhands.app_server.services.injector import Injector +from openhands.sdk.utils.models import DiscriminatedUnionMixin + + +class LLMModelService(ABC): + """Service for discovering available LLM models.""" + + @abstractmethod + async def search_llm_models( + self, + *, + query: str | None = None, + verified_eq: bool | None = None, + provider_eq: str | None = None, + page_id: str | None = None, + limit: int = 50, + ) -> LLMModelPage: + """Search models with optional filtering and pagination. + + Args: + query: Case-insensitive substring match on the model name. + verified_eq: If provided, only return models whose verified + flag matches this value. + provider_eq: If provided, only return models from this + provider (exact match). + page_id: Opaque pagination token from a previous response. + limit: Maximum number of results per page. + """ + + @abstractmethod + async def search_providers( + self, + *, + query: str | None = None, + verified_eq: bool | None = None, + page_id: str | None = None, + limit: int = 50, + ) -> ProviderPage: + """Search providers with optional filtering and pagination. + + Args: + query: Case-insensitive substring match on the provider name. + verified_eq: If provided, only return providers whose verified + flag matches this value. + page_id: Opaque pagination token from a previous response. + limit: Maximum number of results per page. + """ + + +class LLMModelServiceInjector(DiscriminatedUnionMixin, Injector[LLMModelService], ABC): + pass diff --git a/openhands/app_server/settings/settings_models.py b/openhands/app_server/settings/settings_models.py index 88c7815214..c3ee294d85 100644 --- a/openhands/app_server/settings/settings_models.py +++ b/openhands/app_server/settings/settings_models.py @@ -35,9 +35,7 @@ from openhands.app_server.utils.sdk_settings_compat import ( default_agent_settings, validate_agent_settings, ) -from openhands.core.config.llm_config import LLMConfig from openhands.core.config.mcp_config import MCPConfig -from openhands.core.config.utils import load_openhands_config from openhands.sdk.settings import ConversationSettings @@ -61,38 +59,6 @@ def _coerce_dict_secrets(d: dict[str, Any]) -> dict[str, Any]: return out -def _merge_sdk_mcp_configs( - base_config: SDKMCPConfig | None, extra_config: SDKMCPConfig | None -) -> SDKMCPConfig | None: - if base_config is None: - return extra_config - if extra_config is None: - return base_config - - merged_servers: dict[str, Any] = {} - - def _add_server(server_name: str, server_config: dict[str, Any]) -> None: - candidate = server_name or 'server' - if candidate not in merged_servers: - merged_servers[candidate] = server_config - return - - suffix = 1 - while f'{candidate}_{suffix}' in merged_servers: - suffix += 1 - merged_servers[f'{candidate}_{suffix}'] = server_config - - for config in (base_config, extra_config): - raw_config = config.model_dump(exclude_none=True) - for server_name, server_config in raw_config.get('mcpServers', {}).items(): - _add_server(server_name, server_config) - - if not merged_servers: - return None - - return SDKMCPConfig.model_validate({'mcpServers': merged_servers}) - - class SandboxGroupingStrategy(str, Enum): """Strategy for grouping conversations within sandboxes.""" @@ -366,63 +332,6 @@ class Settings(BaseModel): def secrets_store_serializer(self, secrets: Any, info: SerializationInfo): return {'provider_tokens': {}} - # ── Factory methods ───────────────────────────────────────────── - - @staticmethod - def from_config() -> Settings | None: - app_config = load_openhands_config() - llm_config: LLMConfig = app_config.get_llm_config() - if llm_config.api_key is None: - return None - - agent_settings_dict: dict[str, Any] = { - 'agent': app_config.default_agent, - 'llm': { - 'model': llm_config.model, - 'api_key': ( - llm_config.api_key.get_secret_value() - if isinstance(llm_config.api_key, SecretStr) - else llm_config.api_key - ), - 'base_url': llm_config.base_url, - }, - } - if hasattr(app_config, 'mcp') and app_config.mcp: - agent_settings_dict['mcp_config'] = _coerce_value(app_config.mcp) - - return Settings( - language='en', - search_api_key=app_config.search_api_key, - max_budget_per_task=app_config.max_budget_per_task, - # Always LLM for config-file-sourced settings - agent_settings=LLMAgentSettings(**agent_settings_dict), - conversation_settings=ConversationSettings.model_validate( - { - 'confirmation_mode': False, - 'security_analyzer': None, - 'max_iterations': app_config.max_iterations, - } - ), - ) - - def merge_with_config_settings(self) -> 'Settings': - """Merge config.toml MCP settings with stored SDK agent_settings.""" - if not isinstance(self.agent_settings, LLMAgentSettings): - return self - config_settings = Settings.from_config() - if not config_settings: - return self - - merged_mcp = _merge_sdk_mcp_configs( - config_settings.agent_settings.mcp_config, - self.agent_settings.mcp_config, - ) - if merged_mcp is None: - return self - - self.agent_settings.mcp_config = merged_mcp - return self - def to_agent_settings(self) -> LLMAgentSettings | ACPAgentSettings: return self.agent_settings diff --git a/openhands/app_server/user_auth/default_user_auth.py b/openhands/app_server/user_auth/default_user_auth.py index 5937fe2785..e191cf8017 100644 --- a/openhands/app_server/user_auth/default_user_auth.py +++ b/openhands/app_server/user_auth/default_user_auth.py @@ -60,11 +60,6 @@ class DefaultUserAuth(UserAuth): return settings settings_store = await self.get_user_settings_store() settings = await settings_store.load() - - # Merge config.toml settings with stored settings - if settings: - settings = settings.merge_with_config_settings() - self._settings = settings return settings diff --git a/openhands/app_server/utils/llm.py b/openhands/app_server/utils/llm.py index 4067b439e5..2813f4d0d6 100644 --- a/openhands/app_server/utils/llm.py +++ b/openhands/app_server/utils/llm.py @@ -1,7 +1,5 @@ import warnings -import boto3 -import httpx from pydantic import BaseModel with warnings.catch_warnings(): @@ -10,7 +8,6 @@ with warnings.catch_warnings(): from litellm import LlmProviders, ProviderConfigManager, get_llm_provider from openhands.app_server.utils.logger import openhands_logger as logger -from openhands.core.config import LLMConfig, OpenHandsConfig # --------------------------------------------------------------------------- # The ``openhands-sdk`` package is the **single source of truth** for which @@ -264,8 +261,8 @@ def _derive_verified_models(openhands_models: list[str]) -> list[str]: def get_supported_llm_models( - config: OpenHandsConfig, verified_models: list[str] | None = None, + extra_models: list[str] | None = None, ) -> ModelsResponse: """Collect every model available to this server and return structured data. @@ -278,41 +275,17 @@ def get_supported_llm_models( * the recommended default model. Args: - config: The OpenHands configuration. verified_models: Optional list of ``"openhands/"`` strings from the database (SaaS mode). When provided these replace the hardcoded ``OPENHANDS_MODELS``. + extra_models: Optional list of additional model names to include + (e.g. from Bedrock or Ollama discovery). """ litellm_model_list = litellm.model_list + list(litellm.model_cost.keys()) - litellm_model_list_without_bedrock = remove_error_modelId(litellm_model_list) - # TODO: for bedrock, this is using the default config - llm_config: LLMConfig = config.get_llm_config() - bedrock_model_list: list[str] = [] - if ( - llm_config.aws_region_name - and llm_config.aws_access_key_id - and llm_config.aws_secret_access_key - ): - bedrock_model_list = list_foundation_models( - llm_config.aws_region_name, - llm_config.aws_access_key_id.get_secret_value(), - llm_config.aws_secret_access_key.get_secret_value(), - ) - model_list = litellm_model_list_without_bedrock + bedrock_model_list - for llm_config in config.llms.values(): - ollama_base_url = llm_config.ollama_base_url - if llm_config.model.startswith('ollama'): - if not ollama_base_url: - ollama_base_url = llm_config.base_url - if ollama_base_url: - ollama_url = ollama_base_url.strip('/') + '/api/tags' - try: - ollama_models_list = httpx.get(ollama_url, timeout=3).json()['models'] # noqa: ASYNC100 - for model in ollama_models_list: - model_list.append('ollama/' + model['name']) - break - except httpx.HTTPError as e: - logger.error(f'Error getting OLLAMA models: {e}') + model_list = remove_error_modelId(litellm_model_list) + + if extra_models: + model_list = model_list + extra_models openhands_models = get_openhands_models(verified_models) @@ -330,30 +303,5 @@ def get_supported_llm_models( ) -def list_foundation_models( - aws_region_name: str, aws_access_key_id: str, aws_secret_access_key: str -) -> list[str]: - try: - # The AWS bedrock model id is not queried, if no AWS parameters are configured. - client = boto3.client( - service_name='bedrock', - region_name=aws_region_name, - aws_access_key_id=aws_access_key_id, - aws_secret_access_key=aws_secret_access_key, - ) - foundation_models_list = client.list_foundation_models( - byOutputModality='TEXT', byInferenceType='ON_DEMAND' - ) - model_summaries = foundation_models_list['modelSummaries'] - return ['bedrock/' + model['modelId'] for model in model_summaries] - except Exception as err: - logger.warning( - '%s. Please config AWS_REGION_NAME AWS_ACCESS_KEY_ID AWS_SECRET_ACCESS_KEY' - ' if you want use bedrock model.', - err, - ) - return [] - - def remove_error_modelId(model_list: list[str]) -> list[str]: return list(filter(lambda m: not m.startswith('bedrock'), model_list)) diff --git a/tests/unit/app_server/test_config_router.py b/tests/unit/app_server/test_config_router.py index b142f57f4a..5ce4e48101 100644 --- a/tests/unit/app_server/test_config_router.py +++ b/tests/unit/app_server/test_config_router.py @@ -9,15 +9,14 @@ from fastapi import FastAPI, status from fastapi.testclient import TestClient from openhands.app_server.config_api.config_models import LLMModel, Provider -from openhands.app_server.config_api.config_router import ( - _get_all_models_with_verified, - _get_all_providers, - router, +from openhands.app_server.config_api.config_router import router +from openhands.app_server.config_api.default_llm_model_service import ( + _to_llm_models, + _to_providers, ) from openhands.app_server.utils.dependencies import check_session_api_key from openhands.app_server.utils.llm import get_supported_llm_models from openhands.app_server.utils.paging_utils import encode_page_id, paginate_results -from openhands.server.shared import config class TestLLMModel: @@ -85,39 +84,39 @@ class TestPagination: assert next_page_id is None -class TestGetAllModelsWithVerified: - """Test suite for _get_all_models_with_verified function.""" +class TestToLLMModels: + """Test suite for _to_llm_models conversion function.""" def test_returns_list_of_llm_models(self): - models = _get_all_models_with_verified(get_supported_llm_models(config)) + models = _to_llm_models(get_supported_llm_models()) assert isinstance(models, list) assert all(isinstance(m, LLMModel) for m in models) def test_models_verified_mix(self): - models = _get_all_models_with_verified(get_supported_llm_models(config)) + models = _to_llm_models(get_supported_llm_models()) assert any(m.verified is True for m in models) assert any(m.verified is False for m in models) -class TestGetAllProviders: - """Test suite for _get_all_providers function.""" +class TestToProviders: + """Test suite for _to_providers conversion function.""" def test_returns_list_of_providers(self): - providers = _get_all_providers(get_supported_llm_models(config)) + providers = _to_providers(get_supported_llm_models()) assert isinstance(providers, list) assert all(isinstance(p, Provider) for p in providers) def test_providers_are_unique(self): - providers = _get_all_providers(get_supported_llm_models(config)) + providers = _to_providers(get_supported_llm_models()) names = [p.name for p in providers] assert len(names) == len(set(names)) def test_verified_providers_sorted_first(self): - providers = _get_all_providers(get_supported_llm_models(config)) + providers = _to_providers(get_supported_llm_models()) # Find the boundary between verified and unverified found_unverified = False for p in providers: @@ -127,7 +126,7 @@ class TestGetAllProviders: pytest.fail('Verified provider found after unverified provider') def test_contains_verified_and_unverified(self): - providers = _get_all_providers(get_supported_llm_models(config)) + providers = _to_providers(get_supported_llm_models()) assert any(p.verified for p in providers) assert any(not p.verified for p in providers) diff --git a/tests/unit/app_server/test_llm_model_service.py b/tests/unit/app_server/test_llm_model_service.py new file mode 100644 index 0000000000..bb48b1063a --- /dev/null +++ b/tests/unit/app_server/test_llm_model_service.py @@ -0,0 +1,289 @@ +"""Unit tests for the LLM model service.""" + +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from openhands.app_server.config_api.config_models import ( + LLMModelPage, + ProviderPage, +) +from openhands.app_server.config_api.default_llm_model_service import ( + DefaultLLMModelService, + DefaultLLMModelServiceInjector, +) +from openhands.app_server.config_api.llm_model_service import LLMModelService + + +class TestDefaultLLMModelServiceSearchModels: + """Test suite for DefaultLLMModelService.search_llm_models.""" + + @pytest.mark.asyncio + async def test_returns_model_page(self): + service = DefaultLLMModelService() + result = await service.search_llm_models() + + assert isinstance(result, LLMModelPage) + assert len(result.items) > 0 + + @pytest.mark.asyncio + async def test_includes_openhands_models(self): + service = DefaultLLMModelService() + result = await service.search_llm_models(limit=10000) + + providers = {m.provider for m in result.items} + assert 'openhands' in providers + + @pytest.mark.asyncio + async def test_includes_clarifai_models(self): + service = DefaultLLMModelService() + result = await service.search_llm_models(limit=10000) + + providers = {m.provider for m in result.items} + assert 'clarifai' in providers + + @pytest.mark.asyncio + async def test_filters_by_query(self): + service = DefaultLLMModelService() + result = await service.search_llm_models(query='gpt', limit=10000) + + assert len(result.items) > 0 + for m in result.items: + assert 'gpt' in m.name.lower() + + @pytest.mark.asyncio + async def test_filters_by_verified_eq(self): + service = DefaultLLMModelService() + + verified = await service.search_llm_models(verified_eq=True, limit=10000) + assert all(m.verified for m in verified.items) + + unverified = await service.search_llm_models(verified_eq=False, limit=10000) + assert all(not m.verified for m in unverified.items) + + @pytest.mark.asyncio + async def test_filters_by_provider_eq(self): + service = DefaultLLMModelService() + result = await service.search_llm_models(provider_eq='openai', limit=10000) + + assert len(result.items) > 0 + for m in result.items: + assert m.provider == 'openai' + + @pytest.mark.asyncio + async def test_pagination(self): + service = DefaultLLMModelService() + + page1 = await service.search_llm_models(limit=2) + assert len(page1.items) == 2 + assert page1.next_page_id is not None + + page2 = await service.search_llm_models(limit=2, page_id=page1.next_page_id) + assert len(page2.items) == 2 + # Pages should not overlap + names1 = {m.name for m in page1.items} + names2 = {m.name for m in page2.items} + assert names1.isdisjoint(names2) + + @pytest.mark.asyncio + async def test_no_bedrock_without_client(self): + """Without a bedrock client, _list_foundation_models returns empty.""" + service = DefaultLLMModelService() + assert service._list_foundation_models() == [] + + @pytest.mark.asyncio + async def test_bedrock_models_with_client(self): + mock_client = MagicMock() + mock_client.list_foundation_models.return_value = { + 'modelSummaries': [ + {'modelId': 'anthropic.claude-v2'}, + {'modelId': 'amazon.titan-text'}, + ] + } + + service = DefaultLLMModelService(bedrock_client=mock_client) + result = await service.search_llm_models(provider_eq='bedrock', limit=10000) + + model_names = [m.name for m in result.items] + assert 'anthropic.claude-v2' in model_names + assert 'amazon.titan-text' in model_names + + @pytest.mark.asyncio + async def test_bedrock_error_handled_gracefully(self): + mock_client = MagicMock() + mock_client.list_foundation_models.side_effect = Exception('AWS error') + + service = DefaultLLMModelService(bedrock_client=mock_client) + result = await service.search_llm_models() + + assert isinstance(result, LLMModelPage) + assert len(result.items) > 0 + + @pytest.mark.asyncio + async def test_ollama_models_with_url(self): + mock_response = MagicMock() + mock_response.json.return_value = { + 'models': [{'name': 'llama3'}, {'name': 'codellama'}] + } + + mock_client = AsyncMock() + mock_client.get.return_value = mock_response + mock_client.__aenter__ = AsyncMock(return_value=mock_client) + mock_client.__aexit__ = AsyncMock(return_value=None) + + with patch('httpx.AsyncClient', return_value=mock_client): + service = DefaultLLMModelService( + ollama_base_url='http://localhost:11434', + ) + result = await service.search_llm_models(provider_eq='ollama', limit=10000) + + model_names = [m.name for m in result.items] + assert 'llama3' in model_names + assert 'codellama' in model_names + + @pytest.mark.asyncio + async def test_ollama_error_handled_gracefully(self): + import httpx + + mock_client = AsyncMock() + mock_client.get.side_effect = httpx.ConnectError('Connection refused') + mock_client.__aenter__ = AsyncMock(return_value=mock_client) + mock_client.__aexit__ = AsyncMock(return_value=None) + + with patch('httpx.AsyncClient', return_value=mock_client): + service = DefaultLLMModelService( + ollama_base_url='http://localhost:11434', + ) + result = await service.search_llm_models() + + assert isinstance(result, LLMModelPage) + assert len(result.items) > 0 + + @pytest.mark.asyncio + async def test_response_is_cached(self): + service = DefaultLLMModelService() + + result1 = await service.search_llm_models() + await service.search_providers() + + # Both calls should have populated the same cached response + assert service._cached_response is not None + assert result1.items[0].name in [ + m.split('/', 1)[-1] for m in service._cached_response.models + ] + + +class TestDefaultLLMModelServiceSearchProviders: + """Test suite for DefaultLLMModelService.search_providers.""" + + @pytest.mark.asyncio + async def test_returns_provider_page(self): + service = DefaultLLMModelService() + result = await service.search_providers() + + assert isinstance(result, ProviderPage) + assert len(result.items) > 0 + + @pytest.mark.asyncio + async def test_filters_by_query(self): + service = DefaultLLMModelService() + result = await service.search_providers(query='openai', limit=10000) + + assert len(result.items) > 0 + for p in result.items: + assert 'openai' in p.name.lower() + + @pytest.mark.asyncio + async def test_filters_by_verified_eq(self): + service = DefaultLLMModelService() + verified = await service.search_providers(verified_eq=True, limit=10000) + assert all(p.verified for p in verified.items) + + @pytest.mark.asyncio + async def test_pagination(self): + service = DefaultLLMModelService() + + page1 = await service.search_providers(limit=2) + assert len(page1.items) == 2 + assert page1.next_page_id is not None + + page2 = await service.search_providers(limit=2, page_id=page1.next_page_id) + names1 = {p.name for p in page1.items} + names2 = {p.name for p in page2.items} + assert names1.isdisjoint(names2) + + +class TestDefaultLLMModelServiceInjector: + """Test suite for the injector.""" + + @pytest.mark.asyncio + async def test_inject_creates_service(self): + injector = DefaultLLMModelServiceInjector() + + from starlette.datastructures import State + + state = State() + async for service in injector.inject(state): + assert isinstance(service, DefaultLLMModelService) + assert isinstance(service, LLMModelService) + + @pytest.mark.asyncio + async def test_inject_passes_ollama_url(self): + injector = DefaultLLMModelServiceInjector( + ollama_base_url='http://ollama:11434', + ) + + from starlette.datastructures import State + + state = State() + async for service in injector.inject(state): + assert service._ollama_base_url == 'http://ollama:11434' + assert service._bedrock_client is None + + @pytest.mark.asyncio + async def test_inject_creates_bedrock_client(self): + from pydantic import SecretStr + + injector = DefaultLLMModelServiceInjector( + aws_region_name='us-west-2', + aws_access_key_id=SecretStr('AKIATEST'), + aws_secret_access_key=SecretStr('secret123'), + ) + + mock_client = MagicMock() + with patch('boto3.client', return_value=mock_client) as mock_boto3: + from starlette.datastructures import State + + state = State() + async for service in injector.inject(state): + assert service._bedrock_client is mock_client + + mock_boto3.assert_called_once_with( + service_name='bedrock', + region_name='us-west-2', + aws_access_key_id='AKIATEST', + aws_secret_access_key='secret123', + ) + + @pytest.mark.asyncio + async def test_inject_reuses_bedrock_client(self): + from pydantic import SecretStr + + injector = DefaultLLMModelServiceInjector( + aws_region_name='us-west-2', + aws_access_key_id=SecretStr('AKIATEST'), + aws_secret_access_key=SecretStr('secret123'), + ) + + mock_client = MagicMock() + with patch('boto3.client', return_value=mock_client) as mock_boto3: + from starlette.datastructures import State + + state = State() + # Inject twice — boto3.client should only be called once + async for service in injector.inject(state): + assert service._bedrock_client is mock_client + async for service in injector.inject(state): + assert service._bedrock_client is mock_client + + mock_boto3.assert_called_once() diff --git a/tests/unit/core/config/test_mcp_settings_merge.py b/tests/unit/core/config/test_mcp_settings_merge.py deleted file mode 100644 index 842c6e03e9..0000000000 --- a/tests/unit/core/config/test_mcp_settings_merge.py +++ /dev/null @@ -1,184 +0,0 @@ -"""Test MCP settings merging functionality.""" - -import os -from unittest.mock import patch - -import pytest - -from openhands.app_server.settings.settings_models import Settings -from openhands.core.config.mcp_config import ( - MCPConfig, - RemoteMCPServer, - StdioMCPServer, -) -from openhands.sdk.llm import LLM -from openhands.sdk.settings import AgentSettings - - -@pytest.fixture(autouse=True) -def allow_short_context_windows(): - with patch.dict(os.environ, {'ALLOW_SHORT_CONTEXT_WINDOWS': 'true'}, clear=False): - yield - - -def _mcp_config(settings: Settings) -> MCPConfig | None: - mcp = settings.agent_settings.mcp_config - return mcp if mcp and mcp.mcpServers else None - - -_DEFAULT_LLM = LLM(model='test-model') - - -def _settings_with_mcp(mcp_config, llm=None): - """Helper: create Settings with mcp_config set via agent_settings.""" - s = Settings(agent_settings=AgentSettings(llm=llm or _DEFAULT_LLM)) - s.agent_settings.mcp_config = mcp_config - return s - - -@pytest.mark.asyncio -async def test_mcp_settings_merge_config_only(): - """Test merging when only config.toml has MCP settings.""" - mock_config_settings = _settings_with_mcp( - MCPConfig( - mcpServers={ - 'config': RemoteMCPServer( - url='http://config-server.com', transport='sse' - ) - } - ) - ) - - frontend_settings = Settings(agent_settings=AgentSettings(llm=LLM(model='gpt-4'))) - - with patch( - 'openhands.app_server.settings.settings_models.Settings.from_config', - return_value=mock_config_settings, - ): - merged_settings = frontend_settings.merge_with_config_settings() - - merged_mcp_config = _mcp_config(merged_settings) - assert merged_mcp_config is not None - assert len(merged_mcp_config.mcpServers) == 1 - assert 'config' in merged_mcp_config.mcpServers - assert merged_settings.agent_settings.llm.model == 'gpt-4' - - -@pytest.mark.asyncio -async def test_mcp_settings_merge_frontend_only(): - """Test merging when only frontend has MCP settings.""" - mock_config_settings = Settings( - agent_settings=AgentSettings(llm=LLM(model='claude-3')) - ) - - frontend_settings = _settings_with_mcp( - MCPConfig( - mcpServers={ - 'frontend': RemoteMCPServer( - url='http://frontend-server.com', transport='sse' - ) - } - ), - llm=LLM(model='gpt-4'), - ) - - with patch( - 'openhands.app_server.settings.settings_models.Settings.from_config', - return_value=mock_config_settings, - ): - merged_settings = frontend_settings.merge_with_config_settings() - - merged_mcp_config = _mcp_config(merged_settings) - assert merged_mcp_config is not None - assert len(merged_mcp_config.mcpServers) == 1 - assert 'frontend' in merged_mcp_config.mcpServers - assert merged_settings.agent_settings.llm.model == 'gpt-4' - - -@pytest.mark.asyncio -async def test_mcp_settings_merge_both_present(): - """Test merging when both config.toml and frontend have MCP settings.""" - mock_config_settings = _settings_with_mcp( - MCPConfig( - mcpServers={ - 'config-sse': RemoteMCPServer( - url='http://config-server.com', transport='sse' - ), - 'config-stdio': StdioMCPServer(command='config-cmd', args=['arg1']), - } - ) - ) - - frontend_settings = _settings_with_mcp( - MCPConfig( - mcpServers={ - 'frontend-sse': RemoteMCPServer( - url='http://frontend-server.com', transport='sse' - ), - 'frontend-stdio': StdioMCPServer(command='frontend-cmd', args=['arg2']), - } - ), - llm=LLM(model='gpt-4'), - ) - - with patch( - 'openhands.app_server.settings.settings_models.Settings.from_config', - return_value=mock_config_settings, - ): - merged_settings = frontend_settings.merge_with_config_settings() - - merged_mcp_config = _mcp_config(merged_settings) - assert merged_mcp_config is not None - assert len(merged_mcp_config.mcpServers) == 4 - assert 'config-sse' in merged_mcp_config.mcpServers - assert 'frontend-sse' in merged_mcp_config.mcpServers - assert 'config-stdio' in merged_mcp_config.mcpServers - assert 'frontend-stdio' in merged_mcp_config.mcpServers - assert merged_settings.agent_settings.llm.model == 'gpt-4' - - -@pytest.mark.asyncio -async def test_mcp_settings_merge_no_config(): - """Test merging when config.toml has no MCP settings.""" - mock_config_settings = None - - frontend_settings = _settings_with_mcp( - MCPConfig( - mcpServers={ - 'frontend': RemoteMCPServer( - url='http://frontend-server.com', transport='sse' - ) - } - ), - llm=LLM(model='gpt-4'), - ) - - with patch( - 'openhands.app_server.settings.settings_models.Settings.from_config', - return_value=mock_config_settings, - ): - merged_settings = frontend_settings.merge_with_config_settings() - - merged_mcp_config = _mcp_config(merged_settings) - assert merged_mcp_config is not None - assert len(merged_mcp_config.mcpServers) == 1 - assert merged_settings.agent_settings.llm.model == 'gpt-4' - - -@pytest.mark.asyncio -async def test_mcp_settings_merge_neither_present(): - """Test merging when neither config.toml nor frontend have MCP settings.""" - mock_config_settings = Settings( - agent_settings=AgentSettings(llm=LLM(model='claude-3')) - ) - - frontend_settings = Settings(agent_settings=AgentSettings(llm=LLM(model='gpt-4'))) - - with patch( - 'openhands.app_server.settings.settings_models.Settings.from_config', - return_value=mock_config_settings, - ): - merged_settings = frontend_settings.merge_with_config_settings() - - assert _mcp_config(merged_settings) is None - assert merged_settings.agent_settings.llm.model == 'gpt-4' diff --git a/tests/unit/mcp/test_mcp_integration.py b/tests/unit/mcp/test_mcp_integration.py index 20f3c4515b..b11f43ed02 100644 --- a/tests/unit/mcp/test_mcp_integration.py +++ b/tests/unit/mcp/test_mcp_integration.py @@ -1,4 +1,4 @@ -"""Integration test for MCP settings merging in the full flow.""" +"""Integration test for user auth settings flow.""" from unittest.mock import AsyncMock, patch @@ -12,26 +12,9 @@ from openhands.sdk.llm import LLM from openhands.sdk.settings import AgentSettings -def _sdk_mcp_config(settings: Settings) -> MCPConfig | None: - return settings.agent_settings.mcp_config - - @pytest.mark.asyncio -async def test_user_auth_mcp_merging_integration(): - """Test that MCP merging works in the user auth flow.""" - config_settings = Settings( - agent_settings=AgentSettings( - llm=LLM(model='config-model'), - mcp_config=MCPConfig( - mcpServers={ - 'config': RemoteMCPServer( - url='http://config-server.com', transport='sse' - ) - } - ), - ), - ) - +async def test_user_auth_returns_stored_settings(): + """Test that user auth returns stored settings.""" stored_settings = Settings( agent_settings=AgentSettings( llm=LLM(model='anthropic/claude-sonnet-4-5-20250929'), @@ -53,37 +36,19 @@ async def test_user_auth_mcp_merging_integration(): with patch.object( user_auth, 'get_user_settings_store', return_value=mock_settings_store ): - with patch.object(Settings, 'from_config', return_value=config_settings): - merged_settings = await user_auth.get_user_settings() + settings = await user_auth.get_user_settings() - assert merged_settings is not None - merged_mcp = _sdk_mcp_config(merged_settings) - assert ( - merged_settings.agent_settings.llm.model - == 'anthropic/claude-sonnet-4-5-20250929' - ) - assert merged_mcp is not None - assert len(merged_mcp.mcpServers) == 2 - assert 'config' in merged_mcp.mcpServers - assert 'frontend' in merged_mcp.mcpServers + assert settings is not None + assert settings.agent_settings.llm.model == 'anthropic/claude-sonnet-4-5-20250929' + mcp = settings.agent_settings.mcp_config + assert mcp is not None + assert len(mcp.mcpServers) == 1 + assert 'frontend' in mcp.mcpServers @pytest.mark.asyncio async def test_user_auth_caching_behavior(): - """Test that user auth caches the merged settings correctly.""" - config_settings = Settings( - agent_settings=AgentSettings( - llm=LLM(model='config-model'), - mcp_config=MCPConfig( - mcpServers={ - 'config': RemoteMCPServer( - url='http://config-server.com', transport='sse' - ) - } - ), - ), - ) - + """Test that user auth caches settings correctly.""" stored_settings = Settings( agent_settings=AgentSettings( llm=LLM(model='anthropic/claude-sonnet-4-5-20250929'), @@ -105,16 +70,11 @@ async def test_user_auth_caching_behavior(): with patch.object( user_auth, 'get_user_settings_store', return_value=mock_settings_store ): - with patch.object( - Settings, 'from_config', return_value=config_settings - ) as mock_from_config: - settings1 = await user_auth.get_user_settings() - settings2 = await user_auth.get_user_settings() + settings1 = await user_auth.get_user_settings() + settings2 = await user_auth.get_user_settings() assert settings1 is settings2 - assert len(_sdk_mcp_config(settings1).mcpServers) == 2 mock_settings_store.load.assert_called_once() - mock_from_config.assert_called_once() @pytest.mark.asyncio diff --git a/tests/unit/storage/data_models/test_settings.py b/tests/unit/storage/data_models/test_settings.py index eaab7a4e47..7004afa324 100644 --- a/tests/unit/storage/data_models/test_settings.py +++ b/tests/unit/storage/data_models/test_settings.py @@ -1,6 +1,5 @@ import importlib import warnings -from unittest.mock import patch import pytest from fastmcp.mcp_config import MCPConfig @@ -10,8 +9,6 @@ import openhands.app_server.settings.settings_models as settings_module from openhands.app_server.settings.llm_profiles import ProfileNotFoundError from openhands.app_server.settings.settings_models import Settings from openhands.app_server.settings.settings_router import LITE_LLM_API_URL -from openhands.core.config.llm_config import LLMConfig -from openhands.core.config.openhands_config import OpenHandsConfig from openhands.sdk.llm import LLM from openhands.sdk.settings import ( AGENT_SETTINGS_SCHEMA_VERSION, @@ -21,56 +18,6 @@ from openhands.sdk.settings import ( from openhands.sdk.settings.model import CondenserSettings, VerificationSettings -def test_settings_from_config(): - mock_app_config = OpenHandsConfig( - default_agent='test-agent', - max_iterations=100, - llms={ - 'llm': LLMConfig( - model='test-model', - api_key=SecretStr('test-key'), - base_url='https://test.example.com', - ) - }, - ) - - with patch( - 'openhands.app_server.settings.settings_models.load_openhands_config', - return_value=mock_app_config, - ): - settings = Settings.from_config() - - assert settings is not None - assert settings.language == 'en' - assert settings.agent_settings.agent == 'test-agent' - assert settings.conversation_settings.max_iterations == 100 - assert settings.conversation_settings.security_analyzer is None - assert settings.conversation_settings.confirmation_mode is False - assert settings.agent_settings.llm.model == 'test-model' - assert settings.agent_settings.llm.api_key.get_secret_value() == 'test-key' - assert settings.agent_settings.llm.base_url == 'https://test.example.com' - assert not settings.secrets_store.provider_tokens - - -def test_settings_from_config_no_api_key(): - mock_app_config = OpenHandsConfig( - default_agent='test-agent', - max_iterations=100, - llms={ - 'llm': LLMConfig( - model='test-model', api_key=None, base_url='https://test.example.com' - ) - }, - ) - - with patch( - 'openhands.app_server.settings.settings_models.load_openhands_config', - return_value=mock_app_config, - ): - settings = Settings.from_config() - assert settings is None - - def test_settings_handles_sensitive_data(): settings = Settings( language='en', diff --git a/tests/unit/storage/settings/test_file_settings_store.py b/tests/unit/storage/settings/test_file_settings_store.py index 0120d86ce4..a5a387b459 100644 --- a/tests/unit/storage/settings/test_file_settings_store.py +++ b/tests/unit/storage/settings/test_file_settings_store.py @@ -30,12 +30,8 @@ def file_settings_store(mock_file_store): @pytest.mark.asyncio async def test_load_nonexistent_data(file_settings_store): - with patch( - 'openhands.app_server.settings.settings_models.load_openhands_config', - MagicMock(return_value=OpenHandsConfig()), - ): - file_settings_store.file_store.read.side_effect = FileNotFoundError() - assert await file_settings_store.load() is None + file_settings_store.file_store.read.side_effect = FileNotFoundError() + assert await file_settings_store.load() is None @pytest.mark.asyncio