Refactor backend and frontend cleanup paths
This commit is contained in:
+41
-166
@@ -1,27 +1,7 @@
|
||||
"""
|
||||
AI provider configuration service for DocuVault.
|
||||
|
||||
Provides HKDF/Fernet encryption helpers, a provider config loader that reads
|
||||
from the system_settings DB table, and a startup seed function that populates
|
||||
the default provider row from env vars on first boot.
|
||||
|
||||
Security design (D-05, T-07-02):
|
||||
HKDF domain separation — the info bytes b"ai-provider-settings" differ from
|
||||
b"cloud-credentials" used by storage/cloud_utils.py. Both use the same master
|
||||
key (settings.cloud_creds_key) but produce DIFFERENT derived Fernet keys, so
|
||||
a leaked cloud credential cannot decrypt an AI API key and vice versa.
|
||||
|
||||
AlreadyFinalized warning (RESEARCH.md Pitfall 2 / .continue-here.md anti-pattern):
|
||||
The cryptography library raises AlreadyFinalized if .derive() is called twice
|
||||
on the same HKDF instance. _derive_ai_settings_key() creates a FRESH HKDF(...)
|
||||
object on every call — never cache or reuse the HKDF object between calls.
|
||||
|
||||
Pattern reference: storage/cloud_utils.py:_derive_fernet_key().
|
||||
"""
|
||||
"""AI provider configuration and API-key encryption helpers."""
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import logging
|
||||
from typing import Optional
|
||||
|
||||
import structlog
|
||||
@@ -36,167 +16,78 @@ from config import settings
|
||||
|
||||
logger = structlog.get_logger(__name__)
|
||||
|
||||
AI_SETTINGS_KEY_INFO = b"ai-provider-settings"
|
||||
|
||||
# ── HKDF key derivation ───────────────────────────────────────────────────────
|
||||
|
||||
def _derive_ai_settings_key(master_key: bytes, provider_id: str) -> Fernet:
|
||||
"""Derive a per-provider Fernet encryption key using HKDF-SHA256.
|
||||
|
||||
Security notes:
|
||||
- A FRESH HKDF instance is created on every call. The cryptography library
|
||||
raises AlreadyFinalized if .derive() is called twice on the same instance.
|
||||
Never cache or reuse the HKDF object (RESEARCH.md Pitfall 2).
|
||||
- salt = provider_id.encode("utf-8") provides per-provider derivation
|
||||
(deterministic: same provider → same key for encrypt/decrypt consistency).
|
||||
- info = b"ai-provider-settings" provides domain separation from
|
||||
b"cloud-credentials" — same master key, different derived keys.
|
||||
A leaked cloud credential cannot decrypt an AI API key (T-07-02 mitigated).
|
||||
|
||||
Args:
|
||||
master_key: The CLOUD_CREDS_KEY env var as bytes.
|
||||
provider_id: The provider slug, e.g. "openai", "anthropic" (used as HKDF salt).
|
||||
|
||||
Returns:
|
||||
A Fernet instance ready for encrypt/decrypt operations.
|
||||
"""
|
||||
# Create a FRESH HKDF instance — never cache (AlreadyFinalized guard)
|
||||
"""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=b"ai-provider-settings", # domain-separated from b"cloud-credentials"
|
||||
info=AI_SETTINGS_KEY_INFO,
|
||||
)
|
||||
raw_key: bytes = hkdf.derive(master_key)
|
||||
fernet_key = base64.urlsafe_b64encode(raw_key)
|
||||
return Fernet(fernet_key)
|
||||
|
||||
|
||||
# ── Encryption helpers ────────────────────────────────────────────────────────
|
||||
|
||||
def encrypt_api_key(master_key: bytes, provider_id: str, api_key: str) -> str:
|
||||
"""Encrypt a plaintext API key string to a Fernet token.
|
||||
|
||||
The returned string is safe to store in system_settings.api_key_enc.
|
||||
No JSON wrapping — the raw API key string is encrypted directly.
|
||||
|
||||
Args:
|
||||
master_key: The CLOUD_CREDS_KEY env var as bytes.
|
||||
provider_id: The provider slug (used as HKDF salt for key derivation).
|
||||
api_key: The plaintext API key, e.g. "sk-proj-...".
|
||||
|
||||
Returns:
|
||||
A URL-safe base64 Fernet token (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 Fernet token back to the original plaintext API key.
|
||||
|
||||
Args:
|
||||
master_key: The CLOUD_CREDS_KEY env var as bytes.
|
||||
provider_id: The provider slug (used as HKDF salt for key derivation).
|
||||
api_key_enc: The Fernet token string from the database.
|
||||
|
||||
Returns:
|
||||
The original plaintext API key string.
|
||||
"""
|
||||
"""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")
|
||||
|
||||
|
||||
# ── Provider config loader ────────────────────────────────────────────────────
|
||||
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]:
|
||||
"""Load the active AI provider config from the system_settings table.
|
||||
|
||||
Returns a ProviderConfig built from the row where is_active=True,
|
||||
decrypting api_key_enc when present. Returns None when no active row exists.
|
||||
|
||||
Args:
|
||||
session: An open AsyncSession.
|
||||
|
||||
Returns:
|
||||
A ProviderConfig if an active row exists; None if the table is empty or
|
||||
no row is marked active.
|
||||
"""
|
||||
from db.models import SystemSettings # local import to avoid circular deps
|
||||
|
||||
stmt = select(SystemSettings).where(SystemSettings.is_active.is_(True))
|
||||
result = await session.execute(stmt)
|
||||
row = result.scalar_one_or_none()
|
||||
row = await _load_settings_row(session, SystemSettings.is_active.is_(True))
|
||||
return _config_from_settings_row(row) if row else 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: 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,
|
||||
)
|
||||
|
||||
|
||||
# ── 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()
|
||||
row = await _load_settings_row(session, SystemSettings.provider_id == provider_id)
|
||||
return _config_from_settings_row(row) if row else 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,
|
||||
)
|
||||
|
||||
|
||||
# ── Provider ID validator (D-11 migration) ───────────────────────────────────
|
||||
|
||||
def validate_provider_id(v: str) -> str:
|
||||
"""Service-layer provider_id validator; raises ValueError per CLAUDE.md service-vs-API rule."""
|
||||
@@ -207,21 +98,8 @@ def validate_provider_id(v: str) -> str:
|
||||
return v
|
||||
|
||||
|
||||
# ── Startup seed ──────────────────────────────────────────────────────────────
|
||||
|
||||
async def seed_system_settings_from_env(session: AsyncSession) -> None:
|
||||
"""Populate system_settings with a default provider row on first boot.
|
||||
|
||||
Reads settings.default_ai_provider and settings.default_ai_model from config.
|
||||
If no row exists for that provider_id, inserts one with is_active=True.
|
||||
Never overwrites an existing row — idempotent across restarts (D-04).
|
||||
|
||||
This function is called from the FastAPI lifespan in main.py after the
|
||||
session factory is available. Caller is responsible for committing.
|
||||
|
||||
Args:
|
||||
session: An open AsyncSession.
|
||||
"""
|
||||
"""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
|
||||
@@ -232,13 +110,10 @@ async def seed_system_settings_from_env(session: AsyncSession) -> None:
|
||||
existing = result.scalar_one_or_none()
|
||||
|
||||
if existing is not None:
|
||||
# Row already exists — never overwrite (idempotent)
|
||||
return
|
||||
|
||||
# Use PROVIDER_DEFAULTS context_chars for the provider, fallback to 8000
|
||||
context_chars = PROVIDER_DEFAULTS.get(provider_id, {}).get("context_chars", 8000)
|
||||
|
||||
# Insert default row with no API key (local providers like Ollama don't need one)
|
||||
row = SystemSettings(
|
||||
provider_id=provider_id,
|
||||
model_name=model_name,
|
||||
|
||||
Reference in New Issue
Block a user