diff --git a/EvoScientist/image_gen/config.py b/EvoScientist/image_gen/config.py index 805356b..279e2e4 100644 --- a/EvoScientist/image_gen/config.py +++ b/EvoScientist/image_gen/config.py @@ -132,3 +132,28 @@ def load_image_generation_settings( raise ImageGenError( f"invalid {IMAGE_GENERATION_SECTION} section: {details}" ) from None + + +def save_image_generation_settings( + settings: ImageGenerationSettings, *, config_path: Path | None = None +) -> None: + """Replace the image_generation section, preserving other sections.""" + if config_path is None: + from EvoScientist.config.settings import get_config_path + + config_path = get_config_path() + data: dict[str, Any] = {} + if config_path.exists(): + try: + loaded = yaml.safe_load(config_path.read_text(encoding="utf-8")) + except yaml.YAMLError: + loaded = None + if isinstance(loaded, dict): + data = loaded + data[IMAGE_GENERATION_SECTION] = settings.model_dump(mode="json") + config_path.parent.mkdir(parents=True, exist_ok=True) + config_path.write_text( + yaml.safe_dump(data, allow_unicode=True, sort_keys=False), + encoding="utf-8", + ) + config_path.chmod(0o600) diff --git a/EvoScientist/model_registry/http_api.py b/EvoScientist/model_registry/http_api.py index 0ac501b..f953c77 100644 --- a/EvoScientist/model_registry/http_api.py +++ b/EvoScientist/model_registry/http_api.py @@ -45,6 +45,13 @@ from starlette.requests import Request from starlette.responses import JSONResponse, Response from starlette.routing import Route +from EvoScientist.image_gen.adapters.base import ImageGenError +from EvoScientist.image_gen.config import ( + ImageGenerationSettings, + load_image_generation_settings, + save_image_generation_settings, +) + from .adapters import find_adapter_spec, resolve_parameters from .auth import ActorContext, BffAuthenticator from .endpoint_policy import EndpointPolicy @@ -181,6 +188,29 @@ class SnapshotBindRequest(BaseModel): langgraph_run_id: NonEmptyString +class ImageModelWrite(BaseModel): + """One image model entry in a PUT; blank api_key keeps the stored key.""" + + id: str + name: str = "" + provider: Literal["openai", "gemini"] = "openai" + api_key: str = "" + base_url: str = "" + supports_generation: bool = True + supports_edit: bool = True + default_size: str = "1024x1024" + default_quality: str = "auto" + params: dict[str, Any] = Field(default_factory=dict) + + +class PutImageGenerationRequest(BaseModel): + """The image_generation section save request.""" + + default_model: str = "" + timeout_seconds: float = 120.0 + models: list[ImageModelWrite] = Field(default_factory=list) + + class RuntimeOptionsPublic(BaseModel): """The public runtime subset of a frozen model config (section 8.2).""" @@ -474,6 +504,16 @@ class ModelRegistryHttpApi: methods=["POST"], ), Route("/api/models", self.get_selectable_models, methods=["GET"]), + Route( + "/api/image-generation", + self.get_image_generation, + methods=["GET"], + ), + Route( + "/api/image-generation", + self.put_image_generation, + methods=["PUT"], + ), Route("/api/runtime-snapshots", self.create_snapshot, methods=["POST"]), Route( "/api/runtime-snapshots/{snapshot_id}/bind", @@ -607,6 +647,37 @@ class ModelRegistryHttpApi: return _error_response(exc, request_id) return JSONResponse(response.model_dump(mode="json")) + # --- image-generation config API ----------------------------------------- + + async def get_image_generation(self, request: Request) -> Response: + request_id = uuid.uuid4().hex + try: + await self._authenticate( + request, required_scope="model_config:read", require_thread_id=False + ) + response = await asyncio.to_thread(self._image_generation_response) + except ModelRegistryError as exc: + return _error_response(exc, request_id) + return JSONResponse(response) + + async def put_image_generation(self, request: Request) -> Response: + request_id = uuid.uuid4().hex + try: + await self._authenticate( + request, required_scope="model_config:write", require_thread_id=False + ) + body = await self._body(request) + response = await asyncio.to_thread(self._save_image_generation, body) + except ModelRegistryError as exc: + return _error_response(exc, request_id) + except ValidationError as exc: + return _validation_error_response(exc, request_id) + except ImageGenError as exc: + return _error_response( + ModelRegistryError(VALIDATION_FAILED, str(exc)), request_id + ) + return JSONResponse(response) + # --- snapshot API ---------------------------------------------------------- async def create_snapshot(self, request: Request) -> Response: @@ -784,6 +855,50 @@ class ModelRegistryHttpApi: ) return GetSelectableModelsResponse(models=models, defaults=registry.defaults) + @staticmethod + def _image_generation_response() -> dict[str, Any]: + settings = load_image_generation_settings() + models = [] + for entry in settings.models: + payload = entry.model_dump(mode="json") + key = payload.pop("api_key") + configured = bool(key) + hint = f"...{key[-4:]}" if len(key) >= 4 else None + models.append( + { + **payload, + "api_key": "", + "api_key_configured": configured, + "api_key_hint": hint, + } + ) + return { + "default_model": settings.default_model, + "timeout_seconds": settings.timeout_seconds, + "models": models, + } + + @staticmethod + def _save_image_generation(body: Any) -> dict[str, Any]: + request = PutImageGenerationRequest.model_validate(body) + existing = load_image_generation_settings() + existing_keys = {entry.id: entry.api_key for entry in existing.models} + entries = [] + for model in request.models: + payload = model.model_dump(mode="json") + if not payload["api_key"]: + payload["api_key"] = existing_keys.get(model.id, "") + entries.append(payload) + settings = ImageGenerationSettings.model_validate( + { + "default_model": request.default_model, + "timeout_seconds": request.timeout_seconds, + "models": entries, + } + ) + save_image_generation_settings(settings) + return ModelRegistryHttpApi._image_generation_response() + @staticmethod def _create_snapshot( services: ApiServices, actor: ActorContext, body: Any diff --git a/tests/test_image_generation_http.py b/tests/test_image_generation_http.py new file mode 100644 index 0000000..c2b71f5 --- /dev/null +++ b/tests/test_image_generation_http.py @@ -0,0 +1,228 @@ +"""Contract tests for the /api/image-generation endpoints.""" + +from __future__ import annotations + +import json +import time +import uuid +from pathlib import Path + +import jwt +import pytest +import yaml +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import ec +from starlette.applications import Starlette +from starlette.testclient import TestClient + +from EvoScientist.image_gen import config as image_config +from EvoScientist.image_gen.config import ( + ImageGenerationSettings, + ImageModelEntry, + save_image_generation_settings, +) +from EvoScientist.model_registry import http_api +from EvoScientist.model_registry.auth import BffAuthenticator +from EvoScientist.model_registry.endpoint_policy import EndpointPolicy +from EvoScientist.model_registry.hashing import configuration_hash # noqa: F401 (fixture parity) +from EvoScientist.model_registry.http_api import ApiServices, model_registry_routes +from EvoScientist.model_registry.platform import DelegationPublicKey +from EvoScientist.model_registry.resolver import ModelRegistryResolver +from EvoScientist.model_registry.snapshots import SnapshotService +from EvoScientist.model_registry.store import ModelRuntimeStore + +SERVICE_TOKEN = "bff-service-token" + +_PRIVATE_KEY = ec.generate_private_key(ec.SECP256R1()) +_PRIVATE_PEM = _PRIVATE_KEY.private_bytes( + serialization.Encoding.PEM, + serialization.PrivateFormat.PKCS8, + serialization.NoEncryption(), +) +_PUBLIC_PEM = _PRIVATE_KEY.public_key().public_bytes( + serialization.Encoding.PEM, + serialization.PublicFormat.SubjectPublicKeyInfo, +) + + +@pytest.fixture +def config_path(tmp_path): + path = tmp_path / "config.yaml" + path.write_text("other_section: {keep: true}\n", encoding="utf-8") + return path + + +@pytest.fixture +def services(tmp_path, config_path, monkeypatch): + store = ModelRuntimeStore(config_dir=tmp_path / "runtime") + monkeypatch.setattr( + http_api, + "load_image_generation_settings", + lambda: image_config.load_image_generation_settings(config_path=config_path), + ) + monkeypatch.setattr( + http_api, + "save_image_generation_settings", + lambda settings: image_config.save_image_generation_settings( + settings, config_path=config_path + ), + ) + resolver = ModelRegistryResolver(store) + return ApiServices( + store=store, + resolver=resolver, + snapshot_service=SnapshotService(store, resolver), + endpoint_policy=EndpointPolicy([]), + authenticator=BffAuthenticator( + service_token=SERVICE_TOKEN, + service_token_hash=None, + delegation_keys=( + DelegationPublicKey(deployment_id="webui-1", public_key=_PUBLIC_PEM), + ), + jti_store=store, + ), + ) + + +@pytest.fixture +def client(services): + app = Starlette(routes=model_registry_routes(lambda: services)) + return TestClient(app) + + +def _headers(scopes): + now = int(time.time()) + claims = { + "iss": "WebUI", + "aud": "EvoScientist", + "sub": "user-1", + "scopes": list(scopes), + "deployment_id": "webui-1", + "iat": now, + "exp": now + 30, + "jti": uuid.uuid4().hex, + } + return { + "Authorization": f"Bearer {SERVICE_TOKEN}", + "X-Evo-Actor": jwt.encode(claims, _PRIVATE_PEM, algorithm="ES256"), + } + + +def _admin_headers(): + return _headers(["model_config:read", "model_config:write"]) + + +def _settings_payload(**overrides): + payload = { + "default_model": "gpt-image-2", + "timeout_seconds": 120.0, + "models": [ + { + "id": "gpt-image-2", + "name": "GPT Image 2", + "provider": "openai", + "api_key": "sk-live-9876abcd", + "base_url": "", + "supports_generation": True, + "supports_edit": True, + "default_size": "1024x1024", + "default_quality": "auto", + "params": {}, + } + ], + } + payload.update(overrides) + return payload + + +def test_put_then_get_masks_api_key(client, config_path): + put = client.put( + "/api/image-generation", headers=_admin_headers(), json=_settings_payload() + ) + assert put.status_code == 200 + body = put.json() + model = body["models"][0] + assert model["api_key"] == "" + assert model["api_key_configured"] is True + assert model["api_key_hint"] == "...abcd" + assert "sk-live-9876abcd" not in json.dumps(body) + + get = client.get("/api/image-generation", headers=_admin_headers()) + assert get.status_code == 200 + assert get.json() == body + + # The plaintext key really landed in config.yaml, other sections intact. + raw = yaml.safe_load(config_path.read_text(encoding="utf-8")) + assert raw["other_section"] == {"keep": True} + assert raw["image_generation"]["models"][0]["api_key"] == "sk-live-9876abcd" + + +def test_put_blank_api_key_keeps_existing(client): + client.put( + "/api/image-generation", headers=_admin_headers(), json=_settings_payload() + ) + update = _settings_payload() + update["models"][0]["api_key"] = "" + update["models"][0]["name"] = "Renamed" + response = client.put( + "/api/image-generation", headers=_admin_headers(), json=update + ) + assert response.status_code == 200 + model = response.json()["models"][0] + assert model["name"] == "Renamed" + assert model["api_key_configured"] is True + assert model["api_key_hint"] == "...abcd" + + +def test_put_new_model_without_key_is_unconfigured(client): + payload = _settings_payload(models=[_settings_payload()["models"][0] | {"api_key": ""}]) + response = client.put( + "/api/image-generation", headers=_admin_headers(), json=payload + ) + assert response.status_code == 200 + model = response.json()["models"][0] + assert model["api_key_configured"] is False + assert model["api_key_hint"] is None + + +def test_short_key_never_leaks_through_hint(client): + payload = _settings_payload() + payload["models"][0]["api_key"] = "abc" + response = client.put( + "/api/image-generation", headers=_admin_headers(), json=payload + ) + assert response.status_code == 200 + model = response.json()["models"][0] + assert model["api_key_configured"] is True + assert model["api_key_hint"] is None + + +def test_get_empty_when_unconfigured(client): + response = client.get("/api/image-generation", headers=_admin_headers()) + assert response.status_code == 200 + assert response.json() == { + "default_model": "", + "timeout_seconds": 120.0, + "models": [], + } + + +def test_put_rejects_invalid_payload(client): + payload = _settings_payload() + payload["models"][0]["provider"] = "midjourney" + response = client.put( + "/api/image-generation", headers=_admin_headers(), json=payload + ) + assert response.status_code == 422 + body = response.json() + assert body["code"] == "VALIDATION_FAILED" + assert "midjourney" not in json.dumps(body) + + +def test_read_scope_cannot_write(client): + response = client.put( + "/api/image-generation", + headers=_headers(["model_config:read"]), + json=_settings_payload(), + ) + assert response.status_code == 403