phase 3: secure credential vault

This commit is contained in:
shawnkim1997
2026-04-21 21:23:14 +01:00
parent 14de213acf
commit feea679f6e
13 changed files with 617 additions and 1 deletions
+135
View File
@@ -0,0 +1,135 @@
"""Envelope encryption utilities for locally stored API credentials."""
from __future__ import annotations
import base64
import hashlib
import os
import struct
from dataclasses import dataclass
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
MAGIC = b"ATLASV1"
NONCE_SIZE = 12
DEK_SIZE = 32
class SecretsError(Exception):
"""Base class for credential encryption failures."""
class MissingMasterKey(SecretsError):
"""Raised when ATLAS_MASTER_KEY is required but not configured."""
class InvalidSecretBlob(SecretsError):
"""Raised when an encrypted credential blob is malformed or cannot decrypt."""
def _decode_master_key(value: str) -> bytes:
raw = value.strip()
if not raw:
raise MissingMasterKey("ATLAS_MASTER_KEY is empty")
for candidate in (
raw,
raw + "=" * (-len(raw) % 4),
):
try:
decoded = base64.urlsafe_b64decode(candidate.encode("utf-8"))
if len(decoded) == DEK_SIZE:
return decoded
except Exception:
pass
try:
decoded_hex = bytes.fromhex(raw)
if len(decoded_hex) == DEK_SIZE:
return decoded_hex
except ValueError:
pass
# Allow token_urlsafe strings of any supported length while still handing
# AES-GCM a fixed-width key. The raw env var remains the root secret.
return hashlib.sha256(raw.encode("utf-8")).digest()
def load_master_key_from_env(env_var: str = "ATLAS_MASTER_KEY") -> bytes:
value = os.getenv(env_var, "")
if not value.strip():
raise MissingMasterKey(f"{env_var} is not configured")
return _decode_master_key(value)
@dataclass(frozen=True)
class SecretsVault:
"""Envelope encryption vault.
Format:
MAGIC || key_nonce || data_nonce || encrypted_dek_len:uint16 ||
encrypted_dek || ciphertext
"""
master_key: bytes
@classmethod
def from_env(cls) -> "SecretsVault":
return cls(load_master_key_from_env())
def __post_init__(self) -> None:
if len(self.master_key) not in {16, 24, 32}:
raise ValueError("master_key must be 16, 24, or 32 bytes for AES-GCM")
def _aad(self, user_id: str, provider: str) -> bytes:
return f"{user_id.strip()}:{provider.strip().lower()}".encode("utf-8")
def encrypt(self, plaintext: str, user_id: str, provider: str = "kis") -> bytes:
if not plaintext:
raise ValueError("plaintext must not be empty")
dek = os.urandom(DEK_SIZE)
key_nonce = os.urandom(NONCE_SIZE)
data_nonce = os.urandom(NONCE_SIZE)
aad = self._aad(user_id, provider)
encrypted_dek = AESGCM(self.master_key).encrypt(key_nonce, dek, aad)
ciphertext = AESGCM(dek).encrypt(data_nonce, plaintext.encode("utf-8"), aad)
return b"".join(
[
MAGIC,
key_nonce,
data_nonce,
struct.pack("!H", len(encrypted_dek)),
encrypted_dek,
ciphertext,
]
)
def decrypt(self, encrypted_blob: bytes, user_id: str, provider: str = "kis") -> str:
try:
if not encrypted_blob.startswith(MAGIC):
raise InvalidSecretBlob("invalid secret blob header")
offset = len(MAGIC)
key_nonce = encrypted_blob[offset : offset + NONCE_SIZE]
offset += NONCE_SIZE
data_nonce = encrypted_blob[offset : offset + NONCE_SIZE]
offset += NONCE_SIZE
encrypted_dek_len = struct.unpack("!H", encrypted_blob[offset : offset + 2])[0]
offset += 2
encrypted_dek = encrypted_blob[offset : offset + encrypted_dek_len]
offset += encrypted_dek_len
ciphertext = encrypted_blob[offset:]
if len(key_nonce) != NONCE_SIZE or len(data_nonce) != NONCE_SIZE or not encrypted_dek or not ciphertext:
raise InvalidSecretBlob("truncated secret blob")
aad = self._aad(user_id, provider)
dek = AESGCM(self.master_key).decrypt(key_nonce, encrypted_dek, aad)
plaintext = AESGCM(dek).decrypt(data_nonce, ciphertext, aad)
return plaintext.decode("utf-8")
except InvalidSecretBlob:
raise
except Exception as exc:
raise InvalidSecretBlob("secret blob could not be decrypted") from exc
@@ -0,0 +1,208 @@
"""Credential persistence and audit logging."""
from __future__ import annotations
import os
from datetime import datetime, timezone
from typing import Any
from server.db.database import get_db
def _use_postgres() -> bool:
return bool(os.getenv("DATABASE_URL", ""))
def _now_iso() -> str:
return datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ")
def _row_to_dict(row: Any) -> dict[str, Any]:
return dict(row) if row is not None else {}
async def upsert_credential(user_id: str, provider: str, encrypted_blob: bytes) -> None:
provider_key = provider.lower()
now = _now_iso()
if _use_postgres():
from server.db.pg_database import get_pg_pool
pool = await get_pg_pool()
if not pool:
raise RuntimeError("PostgreSQL pool is unavailable")
async with pool.acquire() as conn:
await conn.execute(
"""
INSERT INTO user_credentials (user_id, provider, encrypted_blob, created_at)
VALUES ($1, $2, $3, $4)
ON CONFLICT (user_id, provider)
DO UPDATE SET encrypted_blob = EXCLUDED.encrypted_blob
""",
user_id,
provider_key,
encrypted_blob,
now,
)
return
db = await get_db()
await db.execute(
"""
INSERT INTO user_credentials (user_id, provider, encrypted_blob, created_at)
VALUES (?, ?, ?, ?)
ON CONFLICT(user_id, provider)
DO UPDATE SET encrypted_blob = excluded.encrypted_blob
""",
(user_id, provider_key, encrypted_blob, now),
)
await db.commit()
async def get_credential_blob(user_id: str, provider: str) -> bytes | None:
provider_key = provider.lower()
if _use_postgres():
from server.db.pg_database import get_pg_pool
pool = await get_pg_pool()
if not pool:
return None
async with pool.acquire() as conn:
row = await conn.fetchrow(
"SELECT encrypted_blob FROM user_credentials WHERE user_id = $1 AND provider = $2",
user_id,
provider_key,
)
return bytes(row["encrypted_blob"]) if row else None
db = await get_db()
cursor = await db.execute(
"SELECT encrypted_blob FROM user_credentials WHERE user_id = ? AND provider = ?",
(user_id, provider_key),
)
row = await cursor.fetchone()
return bytes(row["encrypted_blob"]) if row else None
async def get_credential_status(user_id: str, provider: str) -> dict[str, Any] | None:
provider_key = provider.lower()
if _use_postgres():
from server.db.pg_database import get_pg_pool
pool = await get_pg_pool()
if not pool:
return None
async with pool.acquire() as conn:
row = await conn.fetchrow(
"""
SELECT user_id, provider, created_at, last_used_at
FROM user_credentials
WHERE user_id = $1 AND provider = $2
""",
user_id,
provider_key,
)
return dict(row) if row else None
db = await get_db()
cursor = await db.execute(
"""
SELECT user_id, provider, created_at, last_used_at
FROM user_credentials
WHERE user_id = ? AND provider = ?
""",
(user_id, provider_key),
)
row = await cursor.fetchone()
return _row_to_dict(row) if row else None
async def mark_credential_used(user_id: str, provider: str) -> None:
provider_key = provider.lower()
now = _now_iso()
if _use_postgres():
from server.db.pg_database import get_pg_pool
pool = await get_pg_pool()
if not pool:
return
async with pool.acquire() as conn:
await conn.execute(
"UPDATE user_credentials SET last_used_at = $3 WHERE user_id = $1 AND provider = $2",
user_id,
provider_key,
now,
)
return
db = await get_db()
await db.execute(
"UPDATE user_credentials SET last_used_at = ? WHERE user_id = ? AND provider = ?",
(now, user_id, provider_key),
)
await db.commit()
async def delete_credential(user_id: str, provider: str) -> bool:
provider_key = provider.lower()
if _use_postgres():
from server.db.pg_database import get_pg_pool
pool = await get_pg_pool()
if not pool:
return False
async with pool.acquire() as conn:
result = await conn.execute(
"DELETE FROM user_credentials WHERE user_id = $1 AND provider = $2",
user_id,
provider_key,
)
return result.endswith("1")
db = await get_db()
cursor = await db.execute(
"DELETE FROM user_credentials WHERE user_id = ? AND provider = ?",
(user_id, provider_key),
)
await db.commit()
return cursor.rowcount > 0
async def log_credential_access(
user_id: str,
provider: str,
action: str,
ip: str | None = None,
ua: str | None = None,
) -> None:
provider_key = provider.lower()
now = _now_iso()
if _use_postgres():
from server.db.pg_database import get_pg_pool
pool = await get_pg_pool()
if not pool:
return
async with pool.acquire() as conn:
await conn.execute(
"""
INSERT INTO credential_access_log (user_id, provider, action, ip, ua, at)
VALUES ($1, $2, $3, $4, $5, $6)
""",
user_id,
provider_key,
action,
ip,
ua,
now,
)
return
db = await get_db()
await db.execute(
"""
INSERT INTO credential_access_log (user_id, provider, action, ip, ua, at)
VALUES (?, ?, ?, ?, ?, ?)
""",
(user_id, provider_key, action, ip, ua, now),
)
await db.commit()
+26
View File
@@ -107,6 +107,32 @@ async def init_db() -> None:
"""
)
await db.execute(
"""
CREATE TABLE IF NOT EXISTS user_credentials (
user_id TEXT NOT NULL,
provider TEXT NOT NULL,
encrypted_blob BLOB NOT NULL,
created_at TEXT NOT NULL DEFAULT (datetime('now')),
last_used_at TEXT,
PRIMARY KEY (user_id, provider)
)
"""
)
await db.execute(
"""
CREATE TABLE IF NOT EXISTS credential_access_log (
user_id TEXT,
provider TEXT,
action TEXT,
ip TEXT,
ua TEXT,
at TEXT NOT NULL DEFAULT (datetime('now'))
)
"""
)
await db.commit()
+23
View File
@@ -99,10 +99,33 @@ async def init_pg_tables() -> None:
expires_at TIMESTAMPTZ
);
""")
await conn.execute("""
CREATE TABLE IF NOT EXISTS user_credentials (
user_id TEXT NOT NULL,
provider TEXT NOT NULL,
encrypted_blob BYTEA NOT NULL,
created_at TIMESTAMPTZ DEFAULT NOW(),
last_used_at TIMESTAMPTZ,
PRIMARY KEY (user_id, provider)
);
""")
await conn.execute("""
CREATE TABLE IF NOT EXISTS credential_access_log (
user_id TEXT,
provider TEXT,
action TEXT,
ip TEXT,
ua TEXT,
at TIMESTAMPTZ DEFAULT NOW()
);
""")
# Index for cache expiry cleanup
await conn.execute("""
CREATE INDEX IF NOT EXISTS idx_cache_expires ON cache(expires_at);
""")
await conn.execute("""
CREATE INDEX IF NOT EXISTS idx_credential_access_log_at ON credential_access_log(at);
""")
logger.info("PostgreSQL tables initialized.")
+2 -1
View File
@@ -81,7 +81,7 @@ app.add_middleware(
)
# --- Mount routers ---
from server.routers import edgar, analysis, valuation, market_data, news, crypto, fx, portfolio, technical, financials, estimates, earnings, insider, screener, markets, fmp, macro, dart, edinet, research, chat, copilot # noqa: E402
from server.routers import edgar, analysis, valuation, market_data, news, crypto, fx, portfolio, technical, financials, estimates, earnings, insider, screener, markets, fmp, macro, dart, edinet, research, chat, copilot, credentials # noqa: E402
app.include_router(edgar.router, prefix="/api/edgar", tags=["SEC EDGAR"])
app.include_router(analysis.router, prefix="/api/analysis", tags=["AI Analysis"])
@@ -105,6 +105,7 @@ app.include_router(edinet.router, prefix="/api/edinet", tags=["EDINET"])
app.include_router(research.router, prefix="/api/research", tags=["Research"])
app.include_router(chat.router, prefix="/api/chat", tags=["Chat"])
app.include_router(copilot.router, prefix="/api/copilot", tags=["Copilot"])
app.include_router(credentials.router, prefix="/api/credentials", tags=["Credentials"])
@app.get("/health")
@@ -0,0 +1,59 @@
"""Secure credential storage routes."""
from __future__ import annotations
from fastapi import APIRouter, HTTPException, Request
from pydantic import BaseModel, Field
from server.core.secrets import MissingMasterKey, SecretsError
from server.db.credentials_repo import delete_credential, get_credential_status, log_credential_access
from server.services.credentials_service import store_credential
router = APIRouter()
class CredentialUpsertRequest(BaseModel):
secret: str = Field(min_length=1, description="Plaintext API secret. It is envelope-encrypted before storage.")
user_id: str = Field(default="local", min_length=1)
def _request_meta(request: Request) -> tuple[str | None, str | None]:
ip = request.client.host if request.client else None
ua = request.headers.get("user-agent")
return ip, ua
@router.put("/{provider}")
async def upsert_provider_credential(provider: str, body: CredentialUpsertRequest, request: Request) -> dict[str, object]:
ip, ua = _request_meta(request)
try:
await store_credential(body.user_id, provider, body.secret, ip, ua)
except MissingMasterKey as exc:
raise HTTPException(status_code=503, detail="ATLAS_MASTER_KEY is not configured") from exc
except SecretsError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
return {"user_id": body.user_id, "provider": provider.lower(), "stored": True}
@router.get("/{provider}/status")
async def provider_credential_status(provider: str, request: Request, user_id: str = "local") -> dict[str, object]:
ip, ua = _request_meta(request)
row = await get_credential_status(user_id, provider)
await log_credential_access(user_id, provider, "status", ip, ua)
if not row:
return {"user_id": user_id, "provider": provider.lower(), "configured": False}
return {
"user_id": user_id,
"provider": provider.lower(),
"configured": True,
"created_at": row.get("created_at"),
"last_used_at": row.get("last_used_at"),
}
@router.delete("/{provider}")
async def delete_provider_credential(provider: str, request: Request, user_id: str = "local") -> dict[str, object]:
ip, ua = _request_meta(request)
deleted = await delete_credential(user_id, provider)
await log_credential_access(user_id, provider, "delete", ip, ua)
return {"user_id": user_id, "provider": provider.lower(), "deleted": deleted}
@@ -0,0 +1,40 @@
"""High-level credential storage helpers."""
from __future__ import annotations
from server.core.secrets import SecretsVault
from server.db.credentials_repo import (
get_credential_blob,
log_credential_access,
mark_credential_used,
upsert_credential,
)
async def store_credential(
user_id: str,
provider: str,
secret: str,
ip: str | None = None,
ua: str | None = None,
) -> None:
vault = SecretsVault.from_env()
encrypted = vault.encrypt(secret, user_id=user_id, provider=provider)
await upsert_credential(user_id, provider, encrypted)
await log_credential_access(user_id, provider, "upsert", ip, ua)
async def load_credential_secret(
user_id: str,
provider: str,
ip: str | None = None,
ua: str | None = None,
) -> str | None:
blob = await get_credential_blob(user_id, provider)
if blob is None:
await log_credential_access(user_id, provider, "miss", ip, ua)
return None
secret = SecretsVault.from_env().decrypt(blob, user_id=user_id, provider=provider)
await mark_credential_used(user_id, provider)
await log_credential_access(user_id, provider, "decrypt", ip, ua)
return secret