257 lines
8.8 KiB
Python
257 lines
8.8 KiB
Python
"""Vault client for secrets management."""
|
|
import logging
|
|
import os
|
|
import re
|
|
from typing import Optional, Dict, Any
|
|
from urllib.parse import urlparse
|
|
import httpx
|
|
from config.settings import settings
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
class VaultClient:
|
|
"""Client for HashiCorp Vault integration."""
|
|
|
|
def __init__(self):
|
|
"""Initialize Vault client."""
|
|
self.vault_url = getattr(settings, 'VAULT_URL', None)
|
|
self.vault_token = getattr(settings, 'VAULT_TOKEN', None)
|
|
self.enabled = bool(self.vault_url and self.vault_token)
|
|
|
|
if not self.enabled:
|
|
logger.warning("Vault not configured - using mock mode")
|
|
|
|
def parse_vault_reference(self, ref: str) -> Optional[Dict[str, str]]:
|
|
"""Parse vault reference string.
|
|
|
|
Format: vault:path/to/secret#key
|
|
Example: vault:secret/data/model-gateway#api_key
|
|
|
|
Returns:
|
|
Dict with 'path' and 'key' if valid, None otherwise
|
|
"""
|
|
if not ref or not ref.startswith("vault:"):
|
|
return None
|
|
|
|
# Remove vault: prefix
|
|
ref = ref[6:]
|
|
|
|
# Split path and key
|
|
if '#' in ref:
|
|
path, key = ref.rsplit('#', 1)
|
|
else:
|
|
path = ref
|
|
key = None
|
|
|
|
return {
|
|
'path': path,
|
|
'key': key
|
|
}
|
|
|
|
def validate_vault_reference(self, ref: str) -> bool:
|
|
"""Validate vault reference format.
|
|
|
|
Args:
|
|
ref: Vault reference string (e.g., vault:secret/data/key#field)
|
|
|
|
Returns:
|
|
True if valid format, False otherwise
|
|
"""
|
|
if not ref or not isinstance(ref, str):
|
|
return False
|
|
|
|
if not ref.startswith("vault:"):
|
|
return False
|
|
|
|
parsed = self.parse_vault_reference(ref)
|
|
if not parsed:
|
|
return False
|
|
|
|
# Path must not be empty
|
|
if not parsed['path']:
|
|
return False
|
|
|
|
# Path should follow Vault conventions
|
|
# Must contain at least one /
|
|
if '/' not in parsed['path']:
|
|
return False
|
|
|
|
return True
|
|
|
|
def validate_azkv_reference(self, ref: str) -> bool:
|
|
"""Validate Azure Key Vault reference format.
|
|
|
|
Format: azkv://<vault>/secrets/<name>
|
|
Example: azkv://heicode-kv.vault.azure.net/secrets/user-123-repo-main
|
|
"""
|
|
if not ref or not isinstance(ref, str):
|
|
return False
|
|
parsed = urlparse(ref)
|
|
if parsed.scheme != "azkv":
|
|
return False
|
|
if not parsed.netloc:
|
|
return False
|
|
parts = [part for part in parsed.path.split("/") if part]
|
|
return len(parts) >= 2 and parts[0] == "secrets" and bool(parts[1])
|
|
|
|
def validate_secret_reference(self, ref: str) -> bool:
|
|
"""Validate supported secret reference formats."""
|
|
return self.validate_azkv_reference(ref) or self.validate_vault_reference(ref)
|
|
|
|
async def get_secret(self, ref: str) -> Optional[str]:
|
|
"""Fetch secret from Vault.
|
|
|
|
Args:
|
|
ref: Vault reference (e.g., vault:secret/data/model-gateway#api_key)
|
|
|
|
Returns:
|
|
Secret value if found, None otherwise
|
|
"""
|
|
if ref and ref.startswith("azkv://"):
|
|
if not self.validate_azkv_reference(ref):
|
|
logger.error(f"Invalid Azure Key Vault reference: {ref}")
|
|
return None
|
|
env_name = self._env_name_from_secret_ref(ref)
|
|
env_value = os.getenv(env_name)
|
|
if env_value:
|
|
return env_value
|
|
if not self.enabled:
|
|
logger.warning(f"Vault not configured, returning mock secret for {ref}")
|
|
parsed_azkv = urlparse(ref)
|
|
secret_name = parsed_azkv.path.rstrip("/").split("/")[-1]
|
|
return f"mock-secret-azkv-{secret_name}"
|
|
return await self._fetch_azkv_secret(ref)
|
|
|
|
parsed = self.parse_vault_reference(ref)
|
|
if not parsed:
|
|
logger.error(f"Invalid vault reference: {ref}")
|
|
return None
|
|
|
|
if not self.enabled:
|
|
# Mock mode - return placeholder
|
|
logger.warning(f"Vault not configured, returning mock secret for {ref}")
|
|
return f"mock-secret-{parsed['path'].replace('/', '-')}"
|
|
|
|
try:
|
|
# Fetch from Vault
|
|
url = f"{self.vault_url}/v1/{parsed['path']}"
|
|
headers = {
|
|
"X-Vault-Token": self.vault_token
|
|
}
|
|
|
|
async with httpx.AsyncClient() as client:
|
|
response = await client.get(url, headers=headers, timeout=5.0)
|
|
|
|
if response.status_code == 200:
|
|
data = response.json()
|
|
|
|
# Extract secret value
|
|
if 'data' in data:
|
|
secret_data = data['data']
|
|
|
|
# If key specified, get specific field
|
|
if parsed['key']:
|
|
if 'data' in secret_data:
|
|
# KV v2 format
|
|
return secret_data['data'].get(parsed['key'])
|
|
else:
|
|
# KV v1 format
|
|
return secret_data.get(parsed['key'])
|
|
else:
|
|
# Return entire secret
|
|
if 'data' in secret_data:
|
|
return secret_data['data']
|
|
else:
|
|
return secret_data
|
|
|
|
logger.error(f"Unexpected Vault response format for {ref}")
|
|
return None
|
|
|
|
elif response.status_code == 404:
|
|
logger.error(f"Secret not found in Vault: {ref}")
|
|
return None
|
|
|
|
else:
|
|
logger.error(f"Vault request failed: {response.status_code}")
|
|
return None
|
|
|
|
except Exception as e:
|
|
logger.error(f"Failed to fetch secret from Vault: {e}")
|
|
return None
|
|
|
|
async def get_secrets_batch(self, refs: list[str]) -> Dict[str, Optional[str]]:
|
|
"""Fetch multiple secrets from Vault.
|
|
|
|
Args:
|
|
refs: List of vault references
|
|
|
|
Returns:
|
|
Dict mapping reference to secret value
|
|
"""
|
|
results = {}
|
|
for ref in refs:
|
|
results[ref] = await self.get_secret(ref)
|
|
return results
|
|
|
|
def create_k8s_secret_data(self, secrets: Dict[str, str]) -> Dict[str, str]:
|
|
"""Create Kubernetes secret data from vault secrets.
|
|
|
|
Args:
|
|
secrets: Dict mapping env var name to secret value
|
|
|
|
Returns:
|
|
Dict suitable for K8s Secret data field
|
|
"""
|
|
# K8s secrets need base64 encoding, but the K8s client handles that
|
|
return secrets
|
|
|
|
def _env_name_from_secret_ref(self, ref: str) -> str:
|
|
"""Map an azkv secret ref to its conventional env var name."""
|
|
parsed = urlparse(ref)
|
|
secret_name = parsed.path.rstrip("/").split("/")[-1]
|
|
return secret_name.upper().replace("-", "_")
|
|
|
|
async def _fetch_azkv_secret(self, ref: str) -> Optional[str]:
|
|
"""Fetch azkv://<vault>/secrets/<name> using Azure app credentials."""
|
|
parsed = urlparse(ref)
|
|
parts = [part for part in parsed.path.split("/") if part]
|
|
if parsed.scheme != "azkv" or not parsed.netloc or len(parts) < 2 or parts[0] != "secrets":
|
|
logger.error(f"Invalid Azure Key Vault reference: {ref}")
|
|
return None
|
|
|
|
tenant_id = os.getenv("AZURE_TENANT_ID")
|
|
client_id = os.getenv("AZURE_CLIENT_ID")
|
|
client_secret = os.getenv("AZURE_CLIENT_SECRET")
|
|
if not (tenant_id and client_id and client_secret):
|
|
logger.warning("Azure credentials unavailable for azkv secret resolution")
|
|
return None
|
|
|
|
token_url = f"https://login.microsoftonline.com/{tenant_id}/oauth2/v2.0/token"
|
|
token_data = {
|
|
"grant_type": "client_credentials",
|
|
"client_id": client_id,
|
|
"client_secret": client_secret,
|
|
"scope": "https://vault.azure.net/.default",
|
|
}
|
|
secret_url = f"https://{parsed.netloc}/secrets/{parts[1]}?api-version=7.4"
|
|
|
|
try:
|
|
async with httpx.AsyncClient(timeout=10.0) as client:
|
|
token_response = await client.post(token_url, data=token_data)
|
|
token_response.raise_for_status()
|
|
access_token = token_response.json()["access_token"]
|
|
secret_response = await client.get(
|
|
secret_url,
|
|
headers={"Authorization": f"Bearer {access_token}"},
|
|
)
|
|
secret_response.raise_for_status()
|
|
return secret_response.json().get("value")
|
|
except Exception as e:
|
|
logger.error(f"Failed to fetch Azure Key Vault secret: {e}")
|
|
return None
|
|
|
|
|
|
# Global instance
|
|
vault_client = VaultClient()
|