feat: introduce LLMModelService with DI, replacing legacy config-based model discovery (#14237)

Co-authored-by: openhands <openhands@all-hands.dev>
This commit is contained in:
Tim O'Farrell
2026-04-30 11:35:59 -06:00
committed by GitHub
co-authored by openhands
parent e94b504eb4
commit e4a6c4ae47
15 changed files with 732 additions and 589 deletions
-7
View File
@@ -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)
@@ -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)
+39
View File
@@ -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)
+20 -101
View File
@@ -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,
)
@@ -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,
)
@@ -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
@@ -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
@@ -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
+7 -59
View File
@@ -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/<name>"`` 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))
+14 -15
View File
@@ -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)
@@ -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()
@@ -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'
+13 -53
View File
@@ -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
@@ -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',
@@ -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