Files

131 lines
4.3 KiB
Python

"""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,
)