"""Snowflake Cortex Search REST API client.""" import base64 import hashlib import json import time import threading import requests from cryptography.hazmat.primitives import hashes, serialization from cryptography.hazmat.primitives.asymmetric import padding from charts import config _jwt_lock = threading.Lock() _cached_jwt: str | None = None _cached_jwt_expiry: float = 0 def _b64url(data: bytes) -> str: return base64.urlsafe_b64encode(data).rstrip(b'=').decode() def _generate_jwt() -> str: private_key_der: bytes = config.SNOWFLAKE_CONNECT_ARGS['private_key'] private_key = serialization.load_der_private_key(private_key_der, password=None) public_key_der = private_key.public_key().public_bytes( encoding=serialization.Encoding.DER, format=serialization.PublicFormat.SubjectPublicKeyInfo, ) fingerprint = 'SHA256:' + base64.b64encode(hashlib.sha256(public_key_der).digest()).decode() account_id = config.SNOWFLAKE_ACCOUNT.upper().replace('.', '-') username = config.SNOWFLAKE_USER.upper() now = int(time.time()) header = _b64url(json.dumps({'alg': 'RS256', 'typ': 'JWT'}).encode()) payload = _b64url(json.dumps({ 'iss': f'{account_id}.{username}.{fingerprint}', 'sub': f'{account_id}.{username}', 'iat': now, 'exp': now + 59 * 60, }).encode()) signature = private_key.sign( f'{header}.{payload}'.encode(), padding.PKCS1v15(), hashes.SHA256(), ) return f'{header}.{payload}.{_b64url(signature)}' def _get_jwt() -> str: global _cached_jwt, _cached_jwt_expiry now = time.time() with _jwt_lock: if _cached_jwt and _cached_jwt_expiry > now + 5 * 60: return _cached_jwt _cached_jwt = _generate_jwt() _cached_jwt_expiry = now + 59 * 60 return _cached_jwt def query(payload: dict) -> list[dict]: schema = config.SNOWFLAKE_SCHEMA.lower() url = ( config.SNOWFLAKE_CORTEX_SEARCH_API_URL.rstrip('/') + f'/databases/facts/schemas/{schema}' + '/cortex-search-services/nmf_chart_search:query' ) response = requests.post( url, json=payload, headers={ 'Authorization': f'Bearer {_get_jwt()}', 'X-Snowflake-Authorization-Token-Type': 'KEYPAIR_JWT', 'Content-Type': 'application/json', }, timeout=config.SNOWFLAKE_CORTEX_SEARCH_TIMEOUT_MS / 1000, ) response.raise_for_status() return response.json().get('results', [])