Spaces:
Sleeping
Sleeping
| """Standalone Athena client for dashboard. | |
| Uses requests + AWS SigV4 to query Athena directly. No dependency on | |
| app/ modules — compatible with HuggingFace Space deployment. | |
| Public API: | |
| - fetch_full_answer(answer_id) -> str | None | |
| - fetch_full_answers_batch(answer_ids) -> dict[int, str] | |
| - query_athena(sql, params) -> list[dict] | |
| - is_athena_configured() -> bool | |
| """ | |
| import hashlib | |
| import hmac | |
| import json | |
| import os | |
| import re | |
| import time | |
| import uuid | |
| from datetime import datetime, timezone, date as date_type | |
| from pathlib import Path | |
| import requests | |
| from dotenv import load_dotenv | |
| # Load environment variables (project root, APP_ENV-aware) | |
| _project_root = Path(__file__).parent.parent.parent.parent | |
| _app_env = os.environ.get("APP_ENV") | |
| if _app_env: | |
| _candidates = [f".env.{_app_env}", ".env.dev", ".env.prod"] | |
| else: | |
| _candidates = [".env.dev", ".env.prod"] | |
| for _env_name in _candidates: | |
| _env_path = _project_root / _env_name | |
| if _env_path.exists(): | |
| load_dotenv(_env_path, override=True) | |
| break | |
| # --------------------------------------------------------------------------- | |
| # Config from env vars (HuggingFace Secrets compatible) | |
| # --------------------------------------------------------------------------- | |
| _ACCESS_KEY_ID = os.environ.get("ATHENA_ACCESS_KEY_ID", "") | |
| _SECRET_ACCESS_KEY = os.environ.get("ATHENA_SECRET_ACCESS_KEY", "") | |
| _REGION = os.environ.get("ATHENA_REGION", "ap-northeast-2") | |
| _DATABASE = os.environ.get("ATHENA_DATABASE", "fde-chainshift-prod") | |
| _S3_OUTPUT = os.environ.get("ATHENA_S3_OUTPUT", "s3://chainshift-prod-rds-snapshots/athena-results/") | |
| # --------------------------------------------------------------------------- | |
| # AWS SigV4 signing | |
| # --------------------------------------------------------------------------- | |
| _ALGORITHM = "AWS4-HMAC-SHA256" | |
| def _sign(key: bytes, msg: str) -> bytes: | |
| return hmac.new(key, msg.encode("utf-8"), hashlib.sha256).digest() | |
| def _get_signature_key(secret: str, date_stamp: str, region: str, service: str) -> bytes: | |
| k_date = _sign(("AWS4" + secret).encode("utf-8"), date_stamp) | |
| k_region = _sign(k_date, region) | |
| k_service = _sign(k_region, service) | |
| return _sign(k_service, "aws4_request") | |
| def _sigv4_headers(action: str, body: str) -> dict[str, str]: | |
| """Build SigV4-signed headers for an Athena API call.""" | |
| now = datetime.now(timezone.utc) | |
| amz_date = now.strftime("%Y%m%dT%H%M%SZ") | |
| date_stamp = now.strftime("%Y%m%d") | |
| region = _REGION | |
| service = "athena" | |
| host = f"athena.{region}.amazonaws.com" | |
| payload_hash = hashlib.sha256(body.encode("utf-8")).hexdigest() | |
| canonical_headers = ( | |
| f"content-type:application/x-amz-json-1.1\n" | |
| f"host:{host}\n" | |
| f"x-amz-date:{amz_date}\n" | |
| f"x-amz-target:AmazonAthena.{action}\n" | |
| ) | |
| signed_headers = "content-type;host;x-amz-date;x-amz-target" | |
| canonical_request = ( | |
| f"POST\n/\n\n" | |
| f"{canonical_headers}\n{signed_headers}\n{payload_hash}" | |
| ) | |
| credential_scope = f"{date_stamp}/{region}/{service}/aws4_request" | |
| string_to_sign = ( | |
| f"{_ALGORITHM}\n{amz_date}\n{credential_scope}\n" | |
| f"{hashlib.sha256(canonical_request.encode('utf-8')).hexdigest()}" | |
| ) | |
| signing_key = _get_signature_key(_SECRET_ACCESS_KEY, date_stamp, region, service) | |
| signature = hmac.new(signing_key, string_to_sign.encode("utf-8"), hashlib.sha256).hexdigest() | |
| authorization = ( | |
| f"{_ALGORITHM} Credential={_ACCESS_KEY_ID}/{credential_scope}, " | |
| f"SignedHeaders={signed_headers}, Signature={signature}" | |
| ) | |
| return { | |
| "Content-Type": "application/x-amz-json-1.1", | |
| "Host": host, | |
| "X-Amz-Date": amz_date, | |
| "X-Amz-Target": f"AmazonAthena.{action}", | |
| "Authorization": authorization, | |
| } | |
| # --------------------------------------------------------------------------- | |
| # Athena API helpers | |
| # --------------------------------------------------------------------------- | |
| _POLL_INTERVAL = 1.0 | |
| _POLL_TIMEOUT = 120.0 | |
| def _athena_api_call(action: str, body: dict) -> dict: | |
| """POST to the Athena JSON API and return parsed response.""" | |
| payload = json.dumps(body) | |
| url = f"https://athena.{_REGION}.amazonaws.com/" | |
| headers = _sigv4_headers(action, payload) | |
| resp = requests.post(url, data=payload, headers=headers, timeout=30) | |
| if resp.status_code != 200: | |
| raise RuntimeError( | |
| f"Athena {action} failed ({resp.status_code}): {resp.text}" | |
| ) | |
| return resp.json() | |
| def _start_query(sql: str) -> str: | |
| """Submit a query and return the QueryExecutionId.""" | |
| body: dict = { | |
| "QueryString": sql, | |
| "ClientRequestToken": str(uuid.uuid4()), | |
| "QueryExecutionContext": {"Database": _DATABASE}, | |
| "ResultConfiguration": {"OutputLocation": _S3_OUTPUT}, | |
| } | |
| result = _athena_api_call("StartQueryExecution", body) | |
| return result["QueryExecutionId"] | |
| def _wait_for_query(query_id: str, timeout: float | None = None) -> None: | |
| """Poll until the query completes or times out.""" | |
| deadline = time.monotonic() + (timeout or _POLL_TIMEOUT) | |
| while time.monotonic() < deadline: | |
| result = _athena_api_call( | |
| "GetQueryExecution", {"QueryExecutionId": query_id}, | |
| ) | |
| state = result["QueryExecution"]["Status"]["State"] | |
| if state == "SUCCEEDED": | |
| return | |
| if state in ("FAILED", "CANCELLED"): | |
| reason = result["QueryExecution"]["Status"].get( | |
| "StateChangeReason", "unknown" | |
| ) | |
| raise RuntimeError(f"Athena query {state}: {reason}") | |
| time.sleep(_POLL_INTERVAL) | |
| raise TimeoutError(f"Athena query {query_id} timed out after {_POLL_TIMEOUT}s") | |
| def _convert_value(raw: str | None, athena_type: str) -> object: | |
| """Convert a string value from Athena to the appropriate Python type.""" | |
| if raw is None: | |
| return None | |
| athena_type = athena_type.lower() | |
| if athena_type in ("integer", "int", "bigint", "smallint", "tinyint"): | |
| return int(raw) | |
| if athena_type in ("double", "float", "decimal", "real"): | |
| return float(raw) | |
| if athena_type == "boolean": | |
| return raw.lower() == "true" | |
| if athena_type == "date": | |
| return date_type.fromisoformat(raw) | |
| return raw | |
| def _get_results(query_id: str) -> list[dict]: | |
| """Fetch all result pages and return as list[dict].""" | |
| rows: list[dict] = [] | |
| next_token: str | None = None | |
| while True: | |
| body: dict = {"QueryExecutionId": query_id, "MaxResults": 1000} | |
| if next_token: | |
| body["NextToken"] = next_token | |
| result = _athena_api_call("GetQueryResults", body) | |
| result_set = result["ResultSet"] | |
| columns = result_set["ResultSetMetadata"]["ColumnInfo"] | |
| col_names = [c["Name"] for c in columns] | |
| col_types = [c["Type"] for c in columns] | |
| data_rows = result_set.get("Rows", []) | |
| start = 1 if not next_token and data_rows else 0 | |
| for row in data_rows[start:]: | |
| values = row.get("Data", []) | |
| record: dict = {} | |
| for i, col_name in enumerate(col_names): | |
| if i < len(values): | |
| cell = values[i] | |
| raw = cell.get("VarCharValue") | |
| record[col_name] = _convert_value(raw, col_types[i]) | |
| else: | |
| record[col_name] = None | |
| rows.append(record) | |
| next_token = result.get("NextToken") | |
| if not next_token: | |
| break | |
| return rows | |
| # --------------------------------------------------------------------------- | |
| # Parameter substitution | |
| # --------------------------------------------------------------------------- | |
| def _substitute_params(sql: str, params: dict) -> str: | |
| """Replace %(name)s placeholders with escaped values.""" | |
| def _replace(match: re.Match) -> str: | |
| name = match.group(1) | |
| if name not in params: | |
| raise KeyError(f"Parameter '{name}' not found in params dict") | |
| val = params[name] | |
| if val is None: | |
| return "NULL" | |
| if isinstance(val, (int, float)): | |
| return str(val) | |
| escaped = str(val).replace("'", "''") | |
| return f"'{escaped}'" | |
| return re.sub(r"%\((\w+)\)s", _replace, sql) | |
| # --------------------------------------------------------------------------- | |
| # Public API | |
| # --------------------------------------------------------------------------- | |
| def is_athena_configured() -> bool: | |
| """Check if Athena credentials are configured.""" | |
| return bool(_ACCESS_KEY_ID and _SECRET_ACCESS_KEY and _S3_OUTPUT) | |
| def query_athena( | |
| sql: str, | |
| params: dict | None = None, | |
| timeout: float | None = None, | |
| ) -> list[dict]: | |
| """Execute Athena SQL and return results as list of dicts. | |
| Args: | |
| sql: SQL query string. Use %(name)s for parameter placeholders. | |
| params: Optional dict of parameters for the query. | |
| timeout: Max seconds to wait for query completion. | |
| Returns: | |
| List of dicts (one per row), column names as keys. | |
| """ | |
| if not is_athena_configured(): | |
| raise ValueError("Athena credentials not configured") | |
| if params: | |
| sql = _substitute_params(sql, params) | |
| query_id = _start_query(sql) | |
| _wait_for_query(query_id, timeout=timeout) | |
| return _get_results(query_id) | |
| # --------------------------------------------------------------------------- | |
| # Dashboard-specific functions | |
| # --------------------------------------------------------------------------- | |
| def fetch_full_answer(answer_id: int) -> str | None: | |
| """Fetch full answer content from Athena answers table. | |
| Args: | |
| answer_id: The answer ID to fetch | |
| Returns: | |
| Full answer content or None if not found | |
| """ | |
| if not is_athena_configured(): | |
| print(f"[Athena] Not configured, cannot fetch answer {answer_id}") | |
| return None | |
| try: | |
| rows = query_athena( | |
| "SELECT content FROM answers WHERE id = %(answer_id)s AND deleted_at IS NULL", | |
| {"answer_id": answer_id}, | |
| ) | |
| return rows[0]["content"] if rows else None | |
| except Exception as e: | |
| print(f"[Athena Error] Failed to fetch answer {answer_id}: {e}") | |
| return None | |
| def fetch_full_answers_batch(answer_ids: list[int]) -> dict[int, str]: | |
| """Fetch multiple full answers from Athena in a single query. | |
| Args: | |
| answer_ids: List of answer IDs to fetch | |
| Returns: | |
| Dict mapping answer_id to content | |
| """ | |
| if not answer_ids: | |
| return {} | |
| if not is_athena_configured(): | |
| print("[Athena] Not configured, cannot fetch answers batch") | |
| return {} | |
| try: | |
| # Athena doesn't support array params like PostgreSQL's ANY(%s). | |
| # Use IN clause with comma-separated IDs (all integers, safe). | |
| ids_str = ", ".join(str(int(aid)) for aid in answer_ids) | |
| rows = query_athena( | |
| f"SELECT id, content FROM answers WHERE id IN ({ids_str}) AND deleted_at IS NULL", | |
| ) | |
| return {row["id"]: row["content"] for row in rows if row.get("content")} | |
| except Exception as e: | |
| print(f"[Athena Error] Failed to fetch answers batch: {e}") | |
| return {} | |