diff --git a/openhands-agent-server/openhands/agent_server/agent_profiles_router.py b/openhands-agent-server/openhands/agent_server/agent_profiles_router.py index 9a448c05e7..30e86cda60 100644 --- a/openhands-agent-server/openhands/agent_server/agent_profiles_router.py +++ b/openhands-agent-server/openhands/agent_server/agent_profiles_router.py @@ -54,7 +54,7 @@ agent_profiles_router = APIRouter(prefix="/agent-profiles", tags=["Agent Profiles"]) -MAX_AGENT_PROFILES = 50 +MAX_AGENT_PROFILES = 500 ProfileName = Annotated[ str, diff --git a/openhands-agent-server/openhands/agent_server/api.py b/openhands-agent-server/openhands/agent_server/api.py index 541edf0829..e0fc27f887 100644 --- a/openhands-agent-server/openhands/agent_server/api.py +++ b/openhands-agent-server/openhands/agent_server/api.py @@ -48,6 +48,7 @@ init_router, require_initialized, ) +from openhands.agent_server.llm_providers import providers_router from openhands.agent_server.llm_router import llm_router from openhands.agent_server.mcp_router import mcp_router from openhands.agent_server.middleware import CORSDispatcher @@ -440,6 +441,7 @@ def _add_api_routes(app: FastAPI) -> None: api_router.include_router(plugins_router) api_router.include_router(hooks_router) api_router.include_router(llm_router) + api_router.include_router(providers_router) api_router.include_router(mcp_router) api_router.include_router(settings_router) api_router.include_router(workspaces_router) @@ -694,6 +696,7 @@ def create_app(config: Config | None = None) -> FastAPI: _add_api_routes(app) _setup_static_files(app, config) + app.add_middleware( CORSDispatcher, allow_origins=config.allow_cors_origins, diff --git a/openhands-agent-server/openhands/agent_server/llm_providers.py b/openhands-agent-server/openhands/agent_server/llm_providers.py new file mode 100644 index 0000000000..5694591058 --- /dev/null +++ b/openhands-agent-server/openhands/agent_server/llm_providers.py @@ -0,0 +1,581 @@ +"""Model provider endpoints: connect a provider once, manage its models under it. + +A "model provider" is the persisted record for the provider-centric flow in +OpenHands/OpenHands#15492. One API key is held on the provider and shared by +every model nested under it; the user manages those models (add / edit / remove) +directly on the provider. The key is stored as a *named secret* in the +SecretsStore — the provider record keeps only ``secret_name`` and the raw key is +never returned (responses carry ``api_key_set`` instead). + +Endpoints (mounted under ``/api/llm``). The list of *available provider kinds* +for the "add provider" preset picker stays at ``GET /api/llm/providers`` +(see ``llm_router``); these configured-provider records live under +``/api/llm/model-providers`` to avoid shadowing it: + + - GET /model-providers list providers (masked, no keys) + - POST /model-providers create {display_name, kind, + base_url, wire_api, key, + custom_headers, models?} + - GET /model-providers/{id} a single provider (masked) + - PATCH /model-providers/{id} update fields / rotate key + - DELETE /model-providers/{id} remove provider + its named secret + - POST /model-providers/{id}/models add a nested model {name, wire_api?} + - PATCH /model-providers/{id}/models/{name} edit a nested model + - DELETE /model-providers/{id}/models/{name} remove a nested model + - POST /model-providers/{id}/test optional key probe; NEVER mutates + the curated model list +""" + +from __future__ import annotations + +import time +import uuid + +from fastapi import APIRouter, HTTPException, Request, status +from pydantic import BaseModel, Field, SecretStr + +from openhands.agent_server._secrets_exposure import get_config +from openhands.agent_server.persistence import ( + ModelProvider, + PersistedProviders, + ProviderModel, + get_providers_store, + get_secrets_store, +) +from openhands.agent_server.persistence.models import WireApi +from openhands.sdk.llm.utils.unverified_models import ( + _extract_model_and_provider, + _get_litellm_provider_names, + get_supported_llm_models, +) +from openhands.sdk.llm.utils.verified_models import VERIFIED_MODELS +from openhands.sdk.logger import get_logger + + +logger = get_logger(__name__) + +providers_router = APIRouter(prefix="/llm/model-providers", tags=["Model Providers"]) + +# Cap on saved providers. Generous headroom while bounding the config document. +MAX_PROVIDERS = 64 +# Cap on models nested under a single provider. +MAX_MODELS_PER_PROVIDER = 256 + +_SECRET_NAME_PREFIX = "llm_provider_" + + +def _secret_name(provider_id: str) -> str: + return f"{_SECRET_NAME_PREFIX}{provider_id}" + + +def _now() -> int: + return int(time.time()) + + +# ── Request / Response models ──────────────────────────────────────────── + + +class ProviderModelPayload(BaseModel): + name: str = Field(..., min_length=1, max_length=256) + wire_api: WireApi | None = None + + +class ProviderCreateRequest(BaseModel): + """Create a provider. ``key`` is written to the SecretsStore; never echoed.""" + + display_name: str = Field(..., min_length=1, max_length=128) + kind: str = Field(default="custom", max_length=128) + key: SecretStr = Field(..., min_length=1) + base_url: str | None = Field(default=None, max_length=2048) + wire_api: WireApi = "auto" + custom_headers: dict[str, str] = Field(default_factory=dict) + models: list[ProviderModelPayload] = Field(default_factory=list) + + +class ProviderUpdateRequest(BaseModel): + """Partial update. ``key`` rotates the named secret. At least one field.""" + + display_name: str | None = Field(default=None, min_length=1, max_length=128) + kind: str | None = Field(default=None, max_length=128) + key: SecretStr | None = None + base_url: str | None = Field(default=None, max_length=2048) + wire_api: WireApi | None = None + custom_headers: dict[str, str] | None = None + + +class ModelResponse(BaseModel): + name: str + wire_api: WireApi | None = None + + +class ProviderResponse(BaseModel): + """Safe provider view — never includes the raw key or the secret name.""" + + id: str + display_name: str + kind: str + base_url: str | None = None + wire_api: WireApi = "auto" + custom_headers: dict[str, str] = Field(default_factory=dict) + models: list[ModelResponse] = Field(default_factory=list) + created_at: int + updated_at: int + api_key_set: bool = False + + +class TestResponse(BaseModel): + """Result of probing a provider's key. Never mutates the model list. + + ``verified`` is True only when a live network probe confirmed the provider + accepted the key. ``suggested_models`` is the provider's advertised catalog, + offered purely as a convenience for populating model rows. + """ + + id: str + ok: bool + verified: bool = False + suggested_models: list[str] = Field(default_factory=list) + error: str | None = None + + +# ── Helpers ────────────────────────────────────────────────────────────── + + +def _to_response(p: ModelProvider, *, api_key_set: bool) -> ProviderResponse: + return ProviderResponse( + id=p.id, + display_name=p.display_name, + kind=p.kind, + base_url=p.base_url, + wire_api=p.wire_api, + custom_headers=dict(p.custom_headers), + models=[ModelResponse(name=m.name, wire_api=m.wire_api) for m in p.models], + created_at=p.created_at, + updated_at=p.updated_at, + api_key_set=api_key_set, + ) + + +def _api_key_set(secret_name: str) -> bool: + value = get_secrets_store().get_secret(secret_name) + return bool(value and value.strip()) + + +def _get_provider_or_404( + persisted: PersistedProviders | None, provider_id: str +) -> ModelProvider: + for p in persisted.providers if persisted else []: + if p.id == provider_id: + return p + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=f"Provider '{provider_id}' not found", + ) + + +def _provider_catalog(kind: str) -> list[str]: + """Provider's advertised model catalog (no network call), for suggestions.""" + all_models = get_supported_llm_models() + verified = set(VERIFIED_MODELS.get(kind, ())) + out: list[str] = [] + for model in all_models: + model_provider, _, _ = _extract_model_and_provider(model) + if model_provider == kind or model in verified: + prefix = f"{kind}/" + out.append(model[len(prefix) :] if model.startswith(prefix) else model) + out.extend(verified) + return sorted(set(out)) + + +def _live_probe( + kind: str, + key: str, + *, + base_url: str | None, +) -> tuple[bool, str | None]: + """Cheaply check a key against a provider over the network. + + Uses LiteLLM's provider-endpoint listing, which validates the key without + spending tokens. Auth/permission rejections map to ``(False, cause)``; + connectivity failures are surfaced as an error but do not assert invalidity. + """ + import litellm + from litellm.exceptions import AuthenticationError, PermissionDeniedError + + try: + litellm.get_valid_models( + check_provider_endpoint=True, + custom_llm_provider=kind, + api_key=key, + api_base=base_url, + ) + return True, None + except (AuthenticationError, PermissionDeniedError) as e: + return False, f"Provider rejected the key: {str(e)[:200]}" + except Exception as e: # noqa: BLE001 - connectivity/other; don't assert invalid + logger.warning(f"Live probe failed for {kind}: {e}") + return False, f"Could not reach {kind} to verify the key: {str(e)[:200]}" + + +# ── Provider endpoints ─────────────────────────────────────────────────── + + +@providers_router.get("", response_model=list[ProviderResponse]) +async def list_providers(request: Request) -> list[ProviderResponse]: + """List all saved model providers (keys never returned).""" + store = get_providers_store(get_config(request)) + persisted = store.load() + providers = persisted.providers if persisted else [] + return [_to_response(p, api_key_set=_api_key_set(p.secret_name)) for p in providers] + + +@providers_router.post( + "", response_model=ProviderResponse, status_code=status.HTTP_201_CREATED +) +async def create_provider( + request: Request, body: ProviderCreateRequest +) -> ProviderResponse: + """Create a provider: store the key as a named secret, then the record.""" + config = get_config(request) + store = get_providers_store(config) + secrets_store = get_secrets_store(config) + + provider_id = uuid.uuid4().hex + secret_name = _secret_name(provider_id) + now = _now() + + def add(persisted: PersistedProviders) -> PersistedProviders: + if len(persisted.providers) >= MAX_PROVIDERS: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail=( + f"Provider limit reached ({MAX_PROVIDERS}). " + "Remove one before adding a new provider." + ), + ) + if len(body.models) > MAX_MODELS_PER_PROVIDER: + raise HTTPException( + status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, + detail=f"Too many models (max {MAX_MODELS_PER_PROVIDER})", + ) + persisted.providers.append( + ModelProvider( + id=provider_id, + display_name=body.display_name, + kind=body.kind, + base_url=body.base_url, + wire_api=body.wire_api, + custom_headers=dict(body.custom_headers), + secret_name=secret_name, + models=[ + ProviderModel(name=m.name, wire_api=m.wire_api) for m in body.models + ], + created_at=now, + updated_at=now, + ) + ) + return persisted + + try: + secrets_store.set_secret( + name=secret_name, + value=body.key.get_secret_value(), + description=f"LLM provider key for {body.display_name}", + ) + except RuntimeError as e: + logger.error(f"Provider create blocked (secrets): {e}") + raise HTTPException( + status_code=500, + detail="Secrets file is corrupted or encrypted with a different key", + ) + + try: + persisted = store.update(add) + except HTTPException: + # Roll back the secret we just wrote so we don't leak an orphaned key. + try: + secrets_store.delete_secret(secret_name) + except Exception: # noqa: BLE001 - best-effort cleanup + logger.warning(f"Failed to roll back secret {secret_name}") + raise + + provider = next(p for p in persisted.providers if p.id == provider_id) + logger.info("Created model provider", extra={"provider_id": provider_id}) + return _to_response(provider, api_key_set=True) + + +@providers_router.get("/{provider_id}", response_model=ProviderResponse) +async def get_provider(request: Request, provider_id: str) -> ProviderResponse: + """Get a single provider (key never returned).""" + store = get_providers_store(get_config(request)) + provider = _get_provider_or_404(store.load(), provider_id) + return _to_response(provider, api_key_set=_api_key_set(provider.secret_name)) + + +@providers_router.patch("/{provider_id}", response_model=ProviderResponse) +async def update_provider( + request: Request, provider_id: str, body: ProviderUpdateRequest +) -> ProviderResponse: + """Update provider fields or rotate its key. Models are managed separately.""" + if all( + v is None + for v in ( + body.display_name, + body.kind, + body.key, + body.base_url, + body.wire_api, + body.custom_headers, + ) + ): + raise HTTPException( + status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, + detail=( + "Provide at least one of: display_name, kind, key, base_url, " + "wire_api, custom_headers" + ), + ) + + config = get_config(request) + store = get_providers_store(config) + secrets_store = get_secrets_store(config) + + def patch(persisted: PersistedProviders) -> PersistedProviders: + # Runs under the providers lock; rotate the secret only after confirming + # the provider still exists so a concurrent delete can't orphan a key. + for p in persisted.providers: + if p.id == provider_id: + if body.key is not None: + try: + secrets_store.set_secret( + name=p.secret_name, + value=body.key.get_secret_value(), + description=f"LLM provider key for {p.display_name}", + ) + except RuntimeError as e: + logger.error(f"Provider rotate blocked (secrets): {e}") + raise HTTPException( + status_code=500, + detail=( + "Secrets file is corrupted or encrypted with a " + "different key" + ), + ) + updates: dict = {"updated_at": _now()} + if body.display_name is not None: + updates["display_name"] = body.display_name + if body.kind is not None: + updates["kind"] = body.kind + if body.base_url is not None: + updates["base_url"] = body.base_url + if body.wire_api is not None: + updates["wire_api"] = body.wire_api + if body.custom_headers is not None: + updates["custom_headers"] = dict(body.custom_headers) + updated = p.model_copy(update=updates) + persisted.providers = [ + updated if x.id == provider_id else x for x in persisted.providers + ] + return persisted + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=f"Provider '{provider_id}' not found", + ) + + persisted = store.update(patch) + provider = next(p for p in persisted.providers if p.id == provider_id) + return _to_response(provider, api_key_set=_api_key_set(provider.secret_name)) + + +@providers_router.delete( + "/{provider_id}", status_code=status.HTTP_200_OK, response_model=ProviderResponse +) +async def delete_provider(request: Request, provider_id: str) -> ProviderResponse: + """Remove a provider and its named secret.""" + config = get_config(request) + store = get_providers_store(config) + secrets_store = get_secrets_store(config) + + removed: dict[str, ModelProvider] = {} + + def remove(persisted: PersistedProviders) -> PersistedProviders: + for i, p in enumerate(persisted.providers): + if p.id == provider_id: + removed["p"] = persisted.providers.pop(i) + return persisted + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=f"Provider '{provider_id}' not found", + ) + + store.update(remove) + provider = removed["p"] + try: + secrets_store.delete_secret(provider.secret_name) + except Exception: # noqa: BLE001 - record already gone; best-effort + logger.warning(f"Failed to delete secret {provider.secret_name}") + logger.info("Deleted model provider", extra={"provider_id": provider_id}) + return _to_response(provider, api_key_set=False) + + +# ── Nested model endpoints ─────────────────────────────────────────────── + + +@providers_router.post( + "/{provider_id}/models", + response_model=ProviderResponse, + status_code=status.HTTP_201_CREATED, +) +async def add_model( + request: Request, provider_id: str, body: ProviderModelPayload +) -> ProviderResponse: + """Add a model under the provider (shares the provider's key/endpoint).""" + store = get_providers_store(get_config(request)) + + def mutate(persisted: PersistedProviders) -> PersistedProviders: + for p in persisted.providers: + if p.id == provider_id: + if any(m.name == body.name for m in p.models): + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail=f"Model '{body.name}' already exists", + ) + if len(p.models) >= MAX_MODELS_PER_PROVIDER: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail=f"Model limit reached ({MAX_MODELS_PER_PROVIDER})", + ) + models = [ + *p.models, + ProviderModel(name=body.name, wire_api=body.wire_api), + ] + updated = p.model_copy(update={"models": models, "updated_at": _now()}) + persisted.providers = [ + updated if x.id == provider_id else x for x in persisted.providers + ] + return persisted + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=f"Provider '{provider_id}' not found", + ) + + persisted = store.update(mutate) + provider = next(p for p in persisted.providers if p.id == provider_id) + return _to_response(provider, api_key_set=_api_key_set(provider.secret_name)) + + +@providers_router.patch( + "/{provider_id}/models/{model_name}", response_model=ProviderResponse +) +async def update_model( + request: Request, + provider_id: str, + model_name: str, + body: ProviderModelPayload, +) -> ProviderResponse: + """Rename a model and/or change its per-model wire-API override.""" + store = get_providers_store(get_config(request)) + + def mutate(persisted: PersistedProviders) -> PersistedProviders: + for p in persisted.providers: + if p.id == provider_id: + if not any(m.name == model_name for m in p.models): + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=f"Model '{model_name}' not found", + ) + if body.name != model_name and any( + m.name == body.name for m in p.models + ): + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail=f"Model '{body.name}' already exists", + ) + models = [ + ProviderModel(name=body.name, wire_api=body.wire_api) + if m.name == model_name + else m + for m in p.models + ] + updated = p.model_copy(update={"models": models, "updated_at": _now()}) + persisted.providers = [ + updated if x.id == provider_id else x for x in persisted.providers + ] + return persisted + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=f"Provider '{provider_id}' not found", + ) + + persisted = store.update(mutate) + provider = next(p for p in persisted.providers if p.id == provider_id) + return _to_response(provider, api_key_set=_api_key_set(provider.secret_name)) + + +@providers_router.delete( + "/{provider_id}/models/{model_name}", response_model=ProviderResponse +) +async def remove_model( + request: Request, provider_id: str, model_name: str +) -> ProviderResponse: + """Remove a model from the provider.""" + store = get_providers_store(get_config(request)) + + def mutate(persisted: PersistedProviders) -> PersistedProviders: + for p in persisted.providers: + if p.id == provider_id: + if not any(m.name == model_name for m in p.models): + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=f"Model '{model_name}' not found", + ) + models = [m for m in p.models if m.name != model_name] + updated = p.model_copy(update={"models": models, "updated_at": _now()}) + persisted.providers = [ + updated if x.id == provider_id else x for x in persisted.providers + ] + return persisted + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=f"Provider '{provider_id}' not found", + ) + + persisted = store.update(mutate) + provider = next(p for p in persisted.providers if p.id == provider_id) + return _to_response(provider, api_key_set=_api_key_set(provider.secret_name)) + + +# ── Optional test probe ────────────────────────────────────────────────── + + +@providers_router.post("/{provider_id}/test", response_model=TestResponse) +async def test_provider(request: Request, provider_id: str) -> TestResponse: + """Probe the provider's stored key and suggest catalog models. + + This never mutates the provider's curated model list — ``suggested_models`` + is offered only as a convenience for the "add model" affordance. + """ + config = get_config(request) + store = get_providers_store(config) + secrets_store = get_secrets_store(config) + + provider = _get_provider_or_404(store.load(), provider_id) + key = secrets_store.get_secret(provider.secret_name) or "" + if not key.strip(): + return TestResponse(id=provider_id, ok=False, error="No API key stored") + + suggested = _provider_catalog(provider.kind) + if provider.kind not in _get_litellm_provider_names(): + # Unknown/custom endpoint: can't probe, only offer the catalog. + return TestResponse( + id=provider_id, ok=True, verified=False, suggested_models=suggested + ) + + ok, error = _live_probe(provider.kind, key, base_url=provider.base_url) + return TestResponse( + id=provider_id, + ok=ok, + verified=ok, + suggested_models=suggested if ok else [], + error=error, + ) diff --git a/openhands-agent-server/openhands/agent_server/persistence/__init__.py b/openhands-agent-server/openhands/agent_server/persistence/__init__.py index b41360a259..81ff9efd52 100644 --- a/openhands-agent-server/openhands/agent_server/persistence/__init__.py +++ b/openhands-agent-server/openhands/agent_server/persistence/__init__.py @@ -8,25 +8,32 @@ from openhands.agent_server.persistence.models import ( PERSISTED_SETTINGS_SCHEMA_VERSION, + PROVIDERS_SCHEMA_VERSION, SECRET_NAME_PATTERN, WORKSPACES_SCHEMA_VERSION, CustomSecret, + ModelProvider, + PersistedProviders, PersistedSettings, PersistedWorkspaces, + ProviderModel, Secrets, SettingsUpdatePayload, WorkspaceItem, WorkspaceParentItem, ) from openhands.agent_server.persistence.store import ( + FileProvidersStore, FileSecretsStore, FileSettingsStore, FileWorkspacesStore, + ProvidersStore, SecretsStore, SettingsStore, WorkspacesStore, get_agent_profile_store, get_llm_profile_store, + get_providers_store, get_secrets_store, get_settings_store, get_workspaces_store, @@ -37,25 +44,32 @@ __all__ = [ # Constants "PERSISTED_SETTINGS_SCHEMA_VERSION", + "PROVIDERS_SCHEMA_VERSION", "SECRET_NAME_PATTERN", "WORKSPACES_SCHEMA_VERSION", # Models "CustomSecret", + "ModelProvider", + "PersistedProviders", "PersistedSettings", "PersistedWorkspaces", + "ProviderModel", "Secrets", "SettingsUpdatePayload", "WorkspaceItem", "WorkspaceParentItem", # Stores + "FileProvidersStore", "FileSecretsStore", "FileSettingsStore", "FileWorkspacesStore", + "ProvidersStore", "SecretsStore", "SettingsStore", "WorkspacesStore", "get_agent_profile_store", "get_llm_profile_store", + "get_providers_store", "get_secrets_store", "get_settings_store", "get_workspaces_store", diff --git a/openhands-agent-server/openhands/agent_server/persistence/models.py b/openhands-agent-server/openhands/agent_server/persistence/models.py index 2c4547a1f7..0096f6a1f6 100644 --- a/openhands-agent-server/openhands/agent_server/persistence/models.py +++ b/openhands-agent-server/openhands/agent_server/persistence/models.py @@ -9,7 +9,7 @@ import re from collections.abc import Mapping -from typing import Any, TypedDict +from typing import Any, Literal, TypedDict from pydantic import ( BaseModel, @@ -537,6 +537,83 @@ def from_persisted(cls, data: Any) -> PersistedWorkspaces: return cls.model_validate(payload) +# ── Model Providers ────────────────────────────────────────────────────── +# +# A "Provider" is the persisted record for "connect a provider once, then manage +# its models under it" (OpenHands/OpenHands#15492). One key is held on the +# provider and shared by every model nested under it. The key lives in the +# SecretsStore; the provider record stores only ``secret_name`` (never the value), +# so rotating the key is a single SecretsStore write. Models are a nested list +# the user manages (add / edit / remove) — not a fan-out of standalone records. + +PROVIDERS_SCHEMA_VERSION = 1 + +WireApi = Literal["auto", "chat", "responses"] + + +class ProviderModel(BaseModel): + """A model offered by a provider. Inherits the provider's key and endpoint. + + ``wire_api`` optionally overrides the provider default for this one model. + """ + + name: str = Field(..., min_length=1, max_length=256) + wire_api: WireApi | None = Field(default=None) + + model_config = ConfigDict(populate_by_name=True) + + +class ModelProvider(BaseModel): + """A saved model provider (one key, many nested models). + + ``secret_name`` references the key stored in the SecretsStore; the value is + never held here. Responses to clients must mask the key (``api_key_set``) + and never return ``secret_name``. + """ + + id: str = Field(..., min_length=1, max_length=128) + display_name: str = Field(..., min_length=1, max_length=128) + kind: str = Field( + default="custom", + max_length=128, + description="Preset id or litellm provider key (e.g. 'openai', 'custom').", + ) + base_url: str | None = Field(default=None, max_length=2048) + wire_api: WireApi = Field(default="auto") + custom_headers: dict[str, str] = Field(default_factory=dict) + secret_name: str = Field(..., min_length=1, max_length=128) + models: list[ProviderModel] = Field(default_factory=list) + created_at: int = Field(..., description="Unix epoch seconds.") + updated_at: int = Field(..., description="Unix epoch seconds.") + + model_config = ConfigDict(populate_by_name=True) + + +class PersistedProviders(BaseModel): + """Container for all model providers (single JSON document).""" + + schema_version: int = Field(default=PROVIDERS_SCHEMA_VERSION) + providers: list[ModelProvider] = Field(default_factory=list) + + model_config = ConfigDict(populate_by_name=True) + + @classmethod + def from_persisted(cls, data: Any) -> PersistedProviders: + if not isinstance(data, dict): + return cls.model_validate(data) + payload = dict(data) + version = payload.get("schema_version", PROVIDERS_SCHEMA_VERSION) + if not isinstance(version, int): + raise ValueError("PersistedProviders schema_version must be an integer") + if version > PROVIDERS_SCHEMA_VERSION: + raise ValueError( + f"PersistedProviders schema_version {version} is newer than " + f"supported {PROVIDERS_SCHEMA_VERSION}" + ) + payload["schema_version"] = PROVIDERS_SCHEMA_VERSION + return cls.model_validate(payload) + + # ── Helper Functions ───────────────────────────────────────────────────── # # Note: API request/response models have been moved to the SDK to enable diff --git a/openhands-agent-server/openhands/agent_server/persistence/store.py b/openhands-agent-server/openhands/agent_server/persistence/store.py index d661a0cd87..4f1f6a2d7b 100644 --- a/openhands-agent-server/openhands/agent_server/persistence/store.py +++ b/openhands-agent-server/openhands/agent_server/persistence/store.py @@ -24,6 +24,7 @@ from openhands.agent_server.persistence.models import ( CustomSecret, + PersistedProviders, PersistedSettings, PersistedWorkspaces, Secrets, @@ -786,11 +787,94 @@ def update( return updated +class ProvidersStore(ABC): + """Abstract base class for model-provider storage.""" + + @abstractmethod + def load(self) -> PersistedProviders | None: + """Load providers from storage.""" + + @abstractmethod + def save(self, providers: PersistedProviders) -> None: + """Save providers to storage.""" + + @abstractmethod + def update( + self, + update_fn: Callable[[PersistedProviders], PersistedProviders], + ) -> PersistedProviders: + """Atomically update providers with file locking.""" + + +class FileProvidersStore(ProvidersStore): + """File-based storage for model providers. + + Persists a single JSON document at ``/providers.json`` + using the same atomic-write + file-lock primitives as ``FileWorkspacesStore``. + Provider records hold a ``secret_name`` reference to the key (stored in the + SecretsStore), never the key value itself, so no cipher is needed here. + """ + + def __init__( + self, + persistence_dir: Path | str, + filename: str = "providers.json", + ): + _validate_filename(filename) + self.persistence_dir = Path(persistence_dir) + self.filename = filename + self._path = self.persistence_dir / filename + self._lock_path = self.persistence_dir / ".providers.lock" + + def load(self) -> PersistedProviders | None: + if not self._path.exists(): + return None + + try: + with self._path.open("r", encoding="utf-8") as f: + data = json.load(f) + return PersistedProviders.from_persisted(data) + except (PermissionError, OSError) as e: + logger.error(f"Cannot access providers file: {e}") + raise + except json.JSONDecodeError as e: + logger.error(f"Providers file is corrupted: {e}") + return None + except Exception: + logger.error("Failed to load providers", exc_info=True) + return None + + def save(self, providers: PersistedProviders) -> None: + _ensure_secure_directory(self.persistence_dir) + data = providers.model_dump(mode="json", exclude_none=True) + _atomic_write_json(self._path, data) + logger.debug(f"Providers saved to {self._path}") + + def update( + self, + update_fn: Callable[[PersistedProviders], PersistedProviders], + ) -> PersistedProviders: + with _file_lock(self._lock_path): + providers = self.load() + if providers is None: + if self._path.exists(): + raise RuntimeError( + f"Cannot load providers from {self._path}. " + "File may be corrupted. " + "Refusing to overwrite with defaults to prevent data loss." + ) + providers = PersistedProviders() + updated = update_fn(providers) + self.save(updated) + return updated + + # ── Global Store Access ────────────────────────────────────────────────── _settings_store: FileSettingsStore | None = None _secrets_store: FileSecretsStore | None = None _workspaces_store: FileWorkspacesStore | None = None +_providers_store: FileProvidersStore | None = None _llm_profile_store: LLMProfileStore | None = None _agent_profile_store: AgentProfileStore | None = None _store_lock = threading.Lock() @@ -914,6 +998,28 @@ def get_workspaces_store(config: Config | None = None) -> FileWorkspacesStore: return _workspaces_store +def get_providers_store(config: Config | None = None) -> FileProvidersStore: # noqa: ARG001 + """Get the global model-providers store instance (thread-safe). + + Provider records hold only a ``secret_name`` reference to the key (the key + lives in the SecretsStore), so no cipher is used here. Stored in the profile + persistence dir (same as secrets/profiles) so credentials stay in the user's + config directory, never workspace-relative. ``config`` is accepted for parity + with the other store factories; the providers dir is resolved from + ``OH_PERSISTENCE_DIR`` / ``~/.openhands`` (see ``_get_profile_persistence_dir``). + """ + global _providers_store + if _providers_store is not None: + return _providers_store + + with _store_lock: + if _providers_store is None: + _providers_store = FileProvidersStore( + persistence_dir=_get_profile_persistence_dir(), + ) + return _providers_store + + def get_llm_profile_store() -> LLMProfileStore: """Get the global ``LLMProfileStore`` instance (thread-safe). @@ -957,10 +1063,12 @@ def get_agent_profile_store() -> AgentProfileStore: def reset_stores() -> None: """Reset global store instances (for testing).""" global _settings_store, _secrets_store, _workspaces_store + global _providers_store global _llm_profile_store, _agent_profile_store with _store_lock: _settings_store = None _secrets_store = None _workspaces_store = None + _providers_store = None _llm_profile_store = None _agent_profile_store = None diff --git a/openhands-agent-server/openhands/agent_server/profiles_router.py b/openhands-agent-server/openhands/agent_server/profiles_router.py index cea019baee..65d7a62aec 100644 --- a/openhands-agent-server/openhands/agent_server/profiles_router.py +++ b/openhands-agent-server/openhands/agent_server/profiles_router.py @@ -37,7 +37,7 @@ profiles_router = APIRouter(prefix="/profiles", tags=["Profiles"]) -MAX_PROFILES = 50 +MAX_PROFILES = 500 ProfileName = Annotated[ str, diff --git a/openhands-sdk/openhands/sdk/llm/llm.py b/openhands-sdk/openhands/sdk/llm/llm.py index aa4ee34641..1c7c43104b 100644 --- a/openhands-sdk/openhands/sdk/llm/llm.py +++ b/openhands-sdk/openhands/sdk/llm/llm.py @@ -123,6 +123,8 @@ logger = get_logger(__name__) + + _serialized_is_subscription = ContextVar( "serialized_is_subscription", default=False, diff --git a/tests/agent_server/test_agent_profiles_router.py b/tests/agent_server/test_agent_profiles_router.py index 735be21376..c2ab1c3825 100644 --- a/tests/agent_server/test_agent_profiles_router.py +++ b/tests/agent_server/test_agent_profiles_router.py @@ -19,7 +19,6 @@ from openhands.agent_server.api import create_app from openhands.agent_server.config import Config from openhands.agent_server.persistence import reset_stores -from openhands.agent_server.profiles_router import MAX_PROFILES from openhands.sdk.llm import LLM from openhands.sdk.llm.llm_profile_store import LLMProfileStore from openhands.sdk.profiles import ( @@ -219,7 +218,9 @@ def test_seed_does_not_clobber_differently_cased_default_llm_profile( assert reloaded.model == "existing/model" -def test_seed_llm_profile_limit_reached_does_not_500(client, default_llm_profile_store): +def test_seed_llm_profile_limit_reached_does_not_500( + client, default_llm_profile_store, monkeypatch +): """Hitting the LLM profile cap during backfill warns and continues instead of 500ing. @@ -229,7 +230,9 @@ def test_seed_llm_profile_limit_reached_does_not_500(client, default_llm_profile (``openhands.sdk.profiles``) — catching the wrong class let the real one propagate as an unhandled 500. """ - for i in range(MAX_PROFILES): + monkeypatch.setattr(router_module, "MAX_PROFILES", 2) + + for i in range(2): default_llm_profile_store.save(f"other-{i}", LLM(model="x")) response = client.get("/api/agent-profiles") diff --git a/tests/agent_server/test_llm_providers.py b/tests/agent_server/test_llm_providers.py new file mode 100644 index 0000000000..48701c93f1 --- /dev/null +++ b/tests/agent_server/test_llm_providers.py @@ -0,0 +1,219 @@ +"""Tests for the Model Provider endpoints (OpenHands/OpenHands#15492).""" + +from __future__ import annotations + +import tempfile +from pathlib import Path +from unittest.mock import patch + +import pytest +from fastapi.testclient import TestClient + +from openhands.agent_server import llm_providers as prov_module +from openhands.agent_server.api import create_app +from openhands.agent_server.config import Config +from openhands.agent_server.persistence import ( + FileProvidersStore, + get_secrets_store, + reset_stores, +) + + +@pytest.fixture +def temp_dirs(): + with tempfile.TemporaryDirectory() as tmpdir: + base = Path(tmpdir) + (base / "profiles").mkdir(parents=True, exist_ok=True) + yield base + + +@pytest.fixture +def client(temp_dirs, monkeypatch): + reset_stores() + monkeypatch.setenv("OH_PERSISTENCE_DIR", str(temp_dirs)) + config = Config(static_files_path=None, session_api_keys=[], secret_key=None) + with patch( + "openhands.agent_server.llm_providers.get_providers_store", + lambda *_a, **_kw: FileProvidersStore(persistence_dir=temp_dirs), + ): + app = create_app(config) + yield TestClient(app) + reset_stores() + + +def _create(client, **overrides): + body = { + "display_name": "OpenAI", + "kind": "openai", + "key": "sk-test", + "base_url": "https://api.openai.com/v1", + "wire_api": "chat", + "custom_headers": {"X-Org": "eng"}, + "models": [{"name": "gpt-5.6-luna"}], + } + body.update(overrides) + return client.post("/api/llm/model-providers", json=body) + + +def test_list_empty(client): + r = client.get("/api/llm/model-providers") + assert r.status_code == 200 + assert r.json() == [] + + +def test_create_then_list(client): + r = _create(client) + assert r.status_code == 201 + body = r.json() + assert body["display_name"] == "OpenAI" + assert body["kind"] == "openai" + assert body["base_url"] == "https://api.openai.com/v1" + assert body["wire_api"] == "chat" + assert body["custom_headers"] == {"X-Org": "eng"} + assert body["models"] == [{"name": "gpt-5.6-luna", "wire_api": None}] + assert body["api_key_set"] is True + # Key/secret never echoed. + assert "key" not in body + assert "secret_name" not in body + + r2 = client.get("/api/llm/model-providers") + assert r2.status_code == 200 + listed = r2.json() + assert len(listed) == 1 + assert listed[0]["id"] == body["id"] + assert listed[0]["api_key_set"] is True + + +def test_key_stored_as_named_secret(client, temp_dirs, monkeypatch): + r = _create(client) + pid = r.json()["id"] + monkeypatch.setenv("OH_PERSISTENCE_DIR", str(temp_dirs)) + reset_stores() + store = get_secrets_store() + assert store.get_secret(f"llm_provider_{pid}") == "sk-test" + + +def test_get_and_404(client): + r = _create(client) + pid = r.json()["id"] + assert client.get(f"/api/llm/model-providers/{pid}").status_code == 200 + assert client.get("/api/llm/model-providers/nope").status_code == 404 + + +def test_update_fields_and_rotate_key(client, temp_dirs, monkeypatch): + pid = _create(client).json()["id"] + r = client.patch( + f"/api/llm/model-providers/{pid}", + json={ + "display_name": "OpenAI Prod", + "key": "sk-rotated", + "wire_api": "responses", + }, + ) + assert r.status_code == 200 + body = r.json() + assert body["display_name"] == "OpenAI Prod" + assert body["wire_api"] == "responses" + + monkeypatch.setenv("OH_PERSISTENCE_DIR", str(temp_dirs)) + reset_stores() + assert get_secrets_store().get_secret(f"llm_provider_{pid}") == "sk-rotated" + + +def test_update_requires_a_field(client): + pid = _create(client).json()["id"] + r = client.patch(f"/api/llm/model-providers/{pid}", json={}) + assert r.status_code == 422 + + +def test_delete_removes_provider_and_secret(client, temp_dirs, monkeypatch): + pid = _create(client).json()["id"] + r = client.delete(f"/api/llm/model-providers/{pid}") + assert r.status_code == 200 + assert r.json()["api_key_set"] is False + assert client.get(f"/api/llm/model-providers/{pid}").status_code == 404 + + monkeypatch.setenv("OH_PERSISTENCE_DIR", str(temp_dirs)) + reset_stores() + assert get_secrets_store().get_secret(f"llm_provider_{pid}") is None + + +def test_add_edit_remove_model(client): + pid = _create(client).json()["id"] + + # Add + r = client.post( + f"/api/llm/model-providers/{pid}/models", + json={"name": "gpt-5.6-sol", "wire_api": "responses"}, + ) + assert r.status_code == 201 + names = [m["name"] for m in r.json()["models"]] + assert names == ["gpt-5.6-luna", "gpt-5.6-sol"] + + # Duplicate add -> 409 + dup = client.post( + f"/api/llm/model-providers/{pid}/models", json={"name": "gpt-5.6-sol"} + ) + assert dup.status_code == 409 + + # Edit (rename + change wire api) + r = client.patch( + f"/api/llm/model-providers/{pid}/models/gpt-5.6-sol", + json={"name": "gpt-5.6-terra", "wire_api": "chat"}, + ) + assert r.status_code == 200 + models = {m["name"]: m["wire_api"] for m in r.json()["models"]} + assert models == {"gpt-5.6-luna": None, "gpt-5.6-terra": "chat"} + + # Remove + r = client.delete(f"/api/llm/model-providers/{pid}/models/gpt-5.6-luna") + assert r.status_code == 200 + assert [m["name"] for m in r.json()["models"]] == ["gpt-5.6-terra"] + + # Remove missing -> 404 + assert ( + client.delete(f"/api/llm/model-providers/{pid}/models/nope").status_code == 404 + ) + + +def test_test_probe_never_mutates_models(client, monkeypatch): + pid = _create(client).json()["id"] + + monkeypatch.setattr(prov_module, "_live_probe", lambda *a, **k: (True, None)) + r = client.post(f"/api/llm/model-providers/{pid}/test") + assert r.status_code == 200 + body = r.json() + assert body["ok"] is True + assert body["verified"] is True + assert isinstance(body["suggested_models"], list) + + # The provider's curated model list is unchanged. + after = client.get(f"/api/llm/model-providers/{pid}").json() + assert [m["name"] for m in after["models"]] == ["gpt-5.6-luna"] + + +def test_test_probe_reports_bad_key(client, monkeypatch): + pid = _create(client).json()["id"] + monkeypatch.setattr( + prov_module, "_live_probe", lambda *a, **k: (False, "401 invalid key") + ) + r = client.post(f"/api/llm/model-providers/{pid}/test") + assert r.status_code == 200 + body = r.json() + assert body["ok"] is False + assert body["verified"] is False + assert body["suggested_models"] == [] + assert "401" in body["error"] + + +def test_custom_endpoint_test_offers_catalog_without_probe(client): + # A kind litellm doesn't recognize (a custom OpenAI-compatible endpoint): + # ``test`` can't probe it, so it returns ok=True but verified=False. + pid = _create(client, kind="my-vllm", base_url="http://localhost:1234/v1").json()[ + "id" + ] + r = client.post(f"/api/llm/model-providers/{pid}/test") + assert r.status_code == 200 + body = r.json() + assert body["ok"] is True + assert body["verified"] is False