mirror of
https://github.com/OpenHands/OpenHands.git
synced 2026-10-07 13:38:55 +08:00
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:
co-authored by
openhands
parent
e94b504eb4
commit
e4a6c4ae47
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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'
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user