feat(image-gen): add /api/image-generation config endpoints with masked keys
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user