From e678930b8d7377ad75cd0d7736c4b724b18da6ca Mon Sep 17 00:00:00 2001 From: curo1305 Date: Thu, 4 Jun 2026 23:19:40 +0200 Subject: [PATCH] feat(07-05): admin AI-config endpoints + load_provider_config_by_id + atomic is_active flip MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Add GET/PUT /api/admin/ai-config + GET /api/admin/ai-config/test-connection - _ai_config_to_dict whitelist excludes api_key_enc (T-07-01) - PUT: HKDF-encrypt api_key, atomic UPDATE SET is_active = (provider_id = target) (T-07-03) - PUT: upsert with PROVIDER_DEFAULTS for omitted fields; audit log records fields_changed only (T-07-14) - All 3 endpoints require get_current_admin (T-07-15) - Add load_provider_config_by_id to services/ai_config.py (ignores is_active) - Promote 3 xfail tests in test_admin_ai_config.py — all 3 pass --- backend/api/admin.py | 243 +++++++++++++++++++++++++- backend/services/ai_config.py | 45 +++++ backend/tests/test_admin_ai_config.py | 137 +++++++++++++-- 3 files changed, 405 insertions(+), 20 deletions(-) diff --git a/backend/api/admin.py b/backend/api/admin.py index 137e598..63a1274 100644 --- a/backend/api/admin.py +++ b/backend/api/admin.py @@ -28,14 +28,17 @@ from datetime import datetime from typing import Optional from fastapi import APIRouter, Depends, HTTPException, Request, status -from pydantic import BaseModel, EmailStr, Field, field_validator -from sqlalchemy import func, select +from pydantic import BaseModel, ConfigDict, EmailStr, Field, field_validator +from sqlalchemy import func, select, update from sqlalchemy.ext.asyncio import AsyncSession -from db.models import CloudConnection, Document, Quota, RefreshToken, Topic, User +from ai import get_provider +from ai.provider_config import ProviderConfig, PROVIDER_DEFAULTS +from db.models import CloudConnection, Document, Quota, RefreshToken, SystemSettings, Topic, User from deps.auth import get_current_admin from deps.db import get_db from deps.utils import get_client_ip +from services.ai_config import encrypt_api_key, load_provider_config_by_id from services.audit import write_audit_log from services.auth import hash_password, revoke_all_refresh_tokens, validate_password_strength, verify_password from storage import get_storage_backend, get_storage_backend_for_document @@ -48,7 +51,24 @@ _DEFAULT_QUOTA_BYTES = 104857600 # 100 MB free-tier default (D-06) -# ── Safe response helper ────────────────────────────────────────────────────── +# ── Safe response helpers ───────────────────────────────────────────────────── + +def _ai_config_to_dict(row: SystemSettings) -> dict: + """Return a safe subset of SystemSettings fields — explicitly excludes api_key_enc. + + has_api_key is the ONLY indicator that a key is stored (T-07-01 mitigated). + The raw encrypted value and any decrypted plaintext are NEVER returned. + """ + return { + "provider_id": row.provider_id, + "base_url": row.base_url, + "model_name": row.model_name, + "context_chars": row.context_chars, + "is_active": row.is_active, + "has_api_key": row.api_key_enc is not None, + "updated_at": row.updated_at.isoformat() if row.updated_at else None, + } + def _user_to_dict(user: User) -> dict: """Return a safe subset of User fields — never includes password_hash, @@ -99,11 +119,39 @@ class QuotaUpdate(BaseModel): return v -class AiConfigUpdate(BaseModel): +class UserAiConfigUpdate(BaseModel): ai_provider: Optional[str] = None ai_model: Optional[str] = None +class SystemAiConfigUpdate(BaseModel): + """Request model for PUT /api/admin/ai-config (system-level provider configuration). + + Security: extra="forbid" prevents mass-assignment of unexpected fields (T-07-13). + provider_id is validated against PROVIDER_DEFAULTS keys (T-07-13). + api_key is write-only: when None the existing api_key_enc is left untouched, + when "" the api_key_enc is cleared, when a non-empty string it is encrypted. + """ + + model_config = ConfigDict(extra="forbid") + + provider_id: str + api_key: Optional[str] = None + base_url: Optional[str] = None + model_name: Optional[str] = None + context_chars: Optional[int] = None + is_active: Optional[bool] = None + + @field_validator("provider_id") + @classmethod + def provider_must_be_known(cls, v: str) -> str: + if v not in PROVIDER_DEFAULTS: + raise ValueError( + f"Unknown provider_id {v!r}. Must be one of: {list(PROVIDER_DEFAULTS.keys())}" + ) + return v + + class SystemTopicCreate(BaseModel): """Request model for admin system topic creation (D-09).""" @@ -413,7 +461,7 @@ async def update_user_quota( @router.patch("/users/{user_id}/ai-config") async def update_ai_config( user_id: uuid.UUID, - body: AiConfigUpdate, + body: UserAiConfigUpdate, request: Request, session: AsyncSession = Depends(get_db), _admin: User = Depends(get_current_admin), @@ -576,3 +624,186 @@ async def create_system_topic( session, body.name, body.description, body.color, user_id=None ) return topic + + +# ── System AI Provider Configuration (D-08, D-15) ──────────────────────────── + +@router.get("/ai-config/test-connection") +async def test_ai_connection( + provider_id: str, + session: AsyncSession = Depends(get_db), + _admin: User = Depends(get_current_admin), +) -> dict: + """Test connectivity for an AI provider configuration (D-08). + + Loads the system_settings row for the given provider_id (ignoring is_active + so admins can test inactive configurations), builds a provider instance, and + calls health_check(). + + Returns {"ok": true/false, "provider_id": str} — never raises 5xx for + provider-side failures (surfaces as ok=False so the UI shows a clear status). + + Security: requires get_current_admin; provider_id sourced from query param, + provider instance built only from DB row (never from request body). + """ + config = await load_provider_config_by_id(session, provider_id) + if config is None: + return {"ok": False, "provider_id": provider_id, "reason": "not_configured"} + + try: + provider = get_provider(config) + await provider.health_check() + return {"ok": True, "provider_id": provider_id} + except Exception: + return {"ok": False, "provider_id": provider_id, "reason": "health_check_failed"} + + +@router.get("/ai-config") +async def get_ai_config( + session: AsyncSession = Depends(get_db), + _admin: User = Depends(get_current_admin), +) -> dict: + """Return all AI provider configurations for the admin panel (D-08). + + Includes DB rows for providers that have been saved AND synthesised stubs + for providers that only exist in PROVIDER_DEFAULTS (so the admin UI always + shows all 10 providers even before any have been configured). + + Security invariant: api_key_enc is NEVER returned (T-07-01). + Use has_api_key (bool) as the only indicator that a key is stored. + """ + result = await session.execute(select(SystemSettings)) + db_rows = result.scalars().all() + + # Build a lookup for DB rows + db_by_provider: dict[str, SystemSettings] = {r.provider_id: r for r in db_rows} + + providers_out = [] + for pid in PROVIDER_DEFAULTS: + if pid in db_by_provider: + providers_out.append(_ai_config_to_dict(db_by_provider[pid])) + else: + # Synthesise a stub entry for providers with no DB row yet + defaults = PROVIDER_DEFAULTS[pid] + providers_out.append({ + "provider_id": pid, + "base_url": defaults.get("base_url"), + "model_name": defaults.get("model", ""), + "context_chars": defaults.get("context_chars", 8000), + "is_active": False, + "has_api_key": False, + "updated_at": None, + }) + + return {"providers": providers_out} + + +@router.put("/ai-config") +async def update_system_ai_config( + body: SystemAiConfigUpdate, + request: Request, + session: AsyncSession = Depends(get_db), + _admin: User = Depends(get_current_admin), +) -> dict: + """Create or update a system-level AI provider configuration (D-08, D-15). + + Upsert semantics: if no row exists for body.provider_id, one is created using + PROVIDER_DEFAULTS for any omitted fields. + + API key handling (T-07-01 mitigated): + - body.api_key is None → leave existing api_key_enc untouched + - body.api_key == "" → clear api_key_enc (set to NULL) + - body.api_key is a non-empty string → HKDF-encrypt and store + + is_active=True handling (T-07-03 mitigated): + When body.is_active is True, a single atomic UPDATE flips all rows: + SET is_active = (provider_id = :target_id) + This guarantees COUNT(WHERE is_active) == 1 with no read-then-write race. + + Audit log (T-07-14 mitigated): + metadata_ contains only provider_id + fields_changed list — never the + api_key value itself. + """ + from config import settings as _settings # noqa: PLC0415 + + # Load existing row or create a new one from PROVIDER_DEFAULTS + stmt = select(SystemSettings).where(SystemSettings.provider_id == body.provider_id) + result = await session.execute(stmt) + row = result.scalar_one_or_none() + + is_new = row is None + if is_new: + defaults = PROVIDER_DEFAULTS[body.provider_id] + row = SystemSettings( + provider_id=body.provider_id, + model_name=defaults.get("model", ""), + context_chars=defaults.get("context_chars", 8000), + base_url=defaults.get("base_url"), + is_active=False, + api_key_enc=None, + ) + + # Track which fields the caller explicitly set (for audit log — never api_key value) + fields_changed: list[str] = [] + + # Apply provided fields + if body.api_key is not None: + fields_changed.append("api_key") + if body.api_key == "": + row.api_key_enc = None + else: + master_key_str = _settings.cloud_creds_key + master_key_bytes = ( + master_key_str.encode("utf-8") + if isinstance(master_key_str, str) + else master_key_str + ) + row.api_key_enc = encrypt_api_key(master_key_bytes, body.provider_id, body.api_key) + + if body.base_url is not None: + row.base_url = body.base_url + fields_changed.append("base_url") + + if body.model_name is not None: + row.model_name = body.model_name + fields_changed.append("model_name") + + if body.context_chars is not None: + row.context_chars = body.context_chars + fields_changed.append("context_chars") + + if body.is_active is not None: + fields_changed.append("is_active") + + if is_new: + session.add(row) + await session.flush() # ensure row has an id before UPDATE + + # Atomic is_active flip: SET is_active = (provider_id = :target) on ALL rows. + # Single UPDATE statement prevents dual-active race condition (T-07-03). + if body.is_active is True: + await session.execute( + update(SystemSettings).values( + is_active=(SystemSettings.provider_id == body.provider_id) + ) + ) + # Reflect the flip on the in-memory row + row.is_active = True + + _ip_addr = get_client_ip(request) + await write_audit_log( + session, + event_type="admin.ai_config_changed", + user_id=None, + actor_id=_admin.id, + resource_id=None, + ip_address=_ip_addr, + metadata_={"provider_id": body.provider_id, "fields_changed": fields_changed}, + ) + + await session.commit() + + # Reload to pick up DB-generated updated_at after commit + await session.refresh(row) + + return _ai_config_to_dict(row) diff --git a/backend/services/ai_config.py b/backend/services/ai_config.py index 10cfa3b..819e97f 100644 --- a/backend/services/ai_config.py +++ b/backend/services/ai_config.py @@ -151,6 +151,51 @@ async def load_provider_config(session: AsyncSession) -> Optional[ProviderConfig ) +# ── Provider config loader by ID ───────────────────────────────────────────── + +async def load_provider_config_by_id(session: AsyncSession, provider_id: str) -> Optional[ProviderConfig]: + """Load an AI provider config from system_settings by provider_id. + + Unlike load_provider_config(), this function does NOT require is_active=True. + Used by the admin test-connection endpoint so admins can test inactive rows. + + Args: + session: An open AsyncSession. + provider_id: The provider slug to load (e.g. "openai", "lmstudio"). + + Returns: + A ProviderConfig if a row with the given provider_id exists; None otherwise. + """ + from db.models import SystemSettings # local import to avoid circular deps + + stmt = select(SystemSettings).where(SystemSettings.provider_id == provider_id) + result = await session.execute(stmt) + row = result.scalar_one_or_none() + + if row is None: + return None + + # Decrypt API key if present + api_key = "" + if row.api_key_enc: + master_key = settings.cloud_creds_key.encode("utf-8") + try: + api_key = decrypt_api_key(master_key, row.provider_id, row.api_key_enc) + except Exception: + logger.warning( + "ai_config.load_provider_config_by_id: failed to decrypt api_key_enc", + provider_id=row.provider_id, + ) + + return ProviderConfig( + provider_id=row.provider_id, + api_key=api_key, + base_url=row.base_url, + model=row.model_name, + context_chars=row.context_chars, + ) + + # ── Startup seed ────────────────────────────────────────────────────────────── async def seed_system_settings_from_env(session: AsyncSession) -> None: diff --git a/backend/tests/test_admin_ai_config.py b/backend/tests/test_admin_ai_config.py index 2a99647..962f50d 100644 --- a/backend/tests/test_admin_ai_config.py +++ b/backend/tests/test_admin_ai_config.py @@ -1,25 +1,134 @@ """ -Wave 0 xfail stubs for Phase 7 admin AI config endpoint tests. +Integration tests for Phase 7 admin AI config endpoints. + +Promoted from Wave 0 xfail stubs in Plan 07-05. Covers: - T-07-01: GET /api/admin/ai-config never returns api_key_enc (D-05/D-08) - - D-08: PUT /api/admin/ai-config writes active provider (admin AI Providers panel) + - D-08: PUT /api/admin/ai-config writes active provider (admin AI Providers panel) + - T-07-15: PUT /api/admin/ai-config is admin-only (non-admin gets 403) -Each function is a placeholder to be promoted in Plan 07-05 once the admin -AI config endpoints exist. - -Stub policy (STATE.md decision: xfail(strict=False) for Wave 0): - - Body is a single pytest.xfail() call — no assertion code. - - strict=False so unexpected passes (xpass) never break CI. +All tests run against the in-memory SQLite fixture; no live services required. """ import pytest +import pytest_asyncio +from httpx import ASGITransport, AsyncClient +from sqlalchemy.ext.asyncio import AsyncSession -@pytest.mark.xfail(strict=False, reason="Wave 0 stub — promoted in Plan 07-05") -async def test_get_never_returns_key(): - pytest.xfail("not implemented yet — Plan 07-05") +# ── Helpers ─────────────────────────────────────────────────────────────────── + +async def _get_admin_client(db_session: AsyncSession) -> AsyncClient: + """Return an AsyncClient with the DB dependency overridden to db_session.""" + from deps.db import get_db + from main import app + + app.dependency_overrides[get_db] = lambda: db_session + return AsyncClient(transport=ASGITransport(app=app), base_url="http://test") -@pytest.mark.xfail(strict=False, reason="Wave 0 stub — promoted in Plan 07-05") -async def test_put_writes_active_provider(): - pytest.xfail("not implemented yet — Plan 07-05") +# ── Tests ───────────────────────────────────────────────────────────────────── + +@pytest.mark.asyncio +async def test_get_never_returns_key( + async_client: AsyncClient, + admin_user: dict, + db_session: AsyncSession, +): + """GET /api/admin/ai-config must never include api_key_enc or a plaintext key. + + Mitigates T-07-01: admin panel response whitelist (D-08). + """ + # First PUT a provider config with an API key so the DB row exists + put_resp = await async_client.put( + "/api/admin/ai-config", + json={"provider_id": "openai", "api_key": "sk-test-secret"}, + headers=admin_user["headers"], + ) + assert put_resp.status_code == 200 + + # GET must not expose the key in any form + get_resp = await async_client.get( + "/api/admin/ai-config", + headers=admin_user["headers"], + ) + assert get_resp.status_code == 200 + body_text = get_resp.text + + # Security assertions (T-07-01) + assert "api_key_enc" not in body_text, "api_key_enc must never appear in GET response" + assert "sk-test-secret" not in body_text, "plaintext API key must never appear in GET response" + + providers = get_resp.json()["providers"] + assert len(providers) > 0, "providers list must not be empty" + + openai_row = next((p for p in providers if p["provider_id"] == "openai"), None) + assert openai_row is not None, "openai provider must appear in response" + assert "api_key_enc" not in openai_row, "api_key_enc key must not be present in provider dict" + assert openai_row["has_api_key"] is True, "has_api_key must be True after saving a key" + + +@pytest.mark.asyncio +async def test_put_writes_active_provider( + async_client: AsyncClient, + admin_user: dict, + db_session: AsyncSession, +): + """PUT /api/admin/ai-config with is_active=True atomically flips active provider. + + Mitigates T-07-03: no dual-active race (D-08). + After two sequential PUTs with is_active=True on different providers, + exactly one row must have is_active=True. + """ + from sqlalchemy import func, select + from db.models import SystemSettings + + # PUT openai as active + r1 = await async_client.put( + "/api/admin/ai-config", + json={"provider_id": "openai", "api_key": "sk-test", "is_active": True}, + headers=admin_user["headers"], + ) + assert r1.status_code == 200 + assert r1.json()["has_api_key"] is True + assert r1.json()["is_active"] is True + + # PUT anthropic as active (should atomically flip openai to inactive) + r2 = await async_client.put( + "/api/admin/ai-config", + json={"provider_id": "anthropic", "is_active": True}, + headers=admin_user["headers"], + ) + assert r2.status_code == 200 + assert r2.json()["is_active"] is True + + # DB assertion: exactly one row must be active + count_result = await db_session.execute( + select(func.count(SystemSettings.id)).where( + SystemSettings.is_active.is_(True) + ) + ) + active_count = count_result.scalar_one() + assert active_count == 1, ( + f"Expected exactly 1 active provider after two PUTs with is_active=True, " + f"got {active_count}" + ) + + +@pytest.mark.asyncio +async def test_put_admin_only( + async_client: AsyncClient, + auth_user: dict, +): + """PUT /api/admin/ai-config by a regular (non-admin) user must return 403. + + Mitigates T-07-15: admin-only enforcement (D-08). + """ + resp = await async_client.put( + "/api/admin/ai-config", + json={"provider_id": "openai", "api_key": "sk-test"}, + headers=auth_user["headers"], + ) + assert resp.status_code == 403, ( + f"Expected 403 Forbidden for non-admin user, got {resp.status_code}" + )