131 lines
4.3 KiB
Python
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,
|
|
)
|