"""AI provider configuration and API-key encryption helpers.""" from __future__ import annotations import base64 from typing import Optional import structlog from cryptography.fernet import Fernet from cryptography.hazmat.primitives import hashes from cryptography.hazmat.primitives.kdf.hkdf import HKDF from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from ai.provider_config import ProviderConfig, PROVIDER_DEFAULTS from config import settings logger = structlog.get_logger(__name__) AI_SETTINGS_KEY_INFO = b"ai-provider-settings" def _derive_ai_settings_key(master_key: bytes, provider_id: str) -> Fernet: """Derive a per-provider Fernet key with domain-separated HKDF-SHA256.""" hkdf = HKDF( algorithm=hashes.SHA256(), length=32, salt=provider_id.encode("utf-8"), info=AI_SETTINGS_KEY_INFO, ) raw_key: bytes = hkdf.derive(master_key) fernet_key = base64.urlsafe_b64encode(raw_key) return Fernet(fernet_key) def encrypt_api_key(master_key: bytes, provider_id: str, api_key: str) -> str: """Encrypt a plaintext API key for storage in system_settings.api_key_enc.""" f = _derive_ai_settings_key(master_key, provider_id) return f.encrypt(api_key.encode("utf-8")).decode("utf-8") def decrypt_api_key(master_key: bytes, provider_id: str, api_key_enc: str) -> str: """Decrypt a stored API-key token.""" f = _derive_ai_settings_key(master_key, provider_id) return f.decrypt(api_key_enc.encode("utf-8")).decode("utf-8") def _config_from_settings_row(row) -> ProviderConfig: 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: 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, ) async def _load_settings_row(session: AsyncSession, *criteria): from db.models import SystemSettings # local import avoids circular deps stmt = select(SystemSettings) if criteria: stmt = stmt.where(*criteria) result = await session.execute(stmt) return result.scalar_one_or_none() async def load_provider_config(session: AsyncSession) -> Optional[ProviderConfig]: from db.models import SystemSettings # local import to avoid circular deps row = await _load_settings_row(session, SystemSettings.is_active.is_(True)) return _config_from_settings_row(row) if row else None async def load_provider_config_by_id(session: AsyncSession, provider_id: str) -> Optional[ProviderConfig]: from db.models import SystemSettings # local import to avoid circular deps row = await _load_settings_row(session, SystemSettings.provider_id == provider_id) return _config_from_settings_row(row) if row else None def validate_provider_id(v: str) -> str: """Service-layer provider_id validator; raises ValueError per CLAUDE.md service-vs-API rule.""" if v not in PROVIDER_DEFAULTS: raise ValueError( f"Unknown provider_id {v!r}. Must be one of: {list(PROVIDER_DEFAULTS.keys())}" ) return v async def seed_system_settings_from_env(session: AsyncSession) -> None: """Seed the default provider row from env settings if it does not exist.""" from db.models import SystemSettings # local import to avoid circular deps provider_id = settings.default_ai_provider model_name = settings.default_ai_model stmt = select(SystemSettings).where(SystemSettings.provider_id == provider_id) result = await session.execute(stmt) existing = result.scalar_one_or_none() if existing is not None: return context_chars = PROVIDER_DEFAULTS.get(provider_id, {}).get("context_chars", 8000) row = SystemSettings( provider_id=provider_id, model_name=model_name, context_chars=context_chars, is_active=True, api_key_enc=None, base_url=None, ) session.add(row) logger.info( "ai_config.seed_system_settings_from_env: seeded default provider", provider_id=provider_id, model_name=model_name, )