import os
import json
import logging
import requests
import functions_framework
from google.cloud import bigquery
from google.api_core.exceptions import NotFound
from datetime import datetime, timedelta
from urllib.parse import urljoin

# --- Configuration ---
DX_API_URL = os.environ.get("DX_API_URL")
DX_API_TOKEN = os.environ.get("DX_API_KEY").strip() # Should be mounted from Secret Manager
PROJECT_ID = os.environ.get("GCP_PROJECT")
DATASET_ID = os.environ.get("DATASET_ID", "dx_gemini_observability")

# Configure Logging
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)

# --- SQL CTE Definitions ---
# Query logic: Fetch data starting from Yesterday (Midnight). 
# This captures the full previous day (finalizing it) AND the current partial day.
IDE_CTE = f"""
ide_activity AS (
  SELECT
    labels.user_id as email,
    DATE(timestamp) as date,
    "Gemini Code Assist" as tool,
    
    -- 1. Count acceptances (nested inside jsonpayload_v1_metadatalog)
    COUNTIF(jsonpayload_v1_metadatalog.codeacceptance IS NOT NULL) as acceptances,
    
    -- 2. Sum lines accepted 
    -- (Navigates the wrapper -> casts string "16.0" to FLOAT -> then INT)
    SUM(CAST(SAFE_CAST(jsonpayload_v1_metadatalog.codeacceptance.linescount AS FLOAT64) AS INT64)) as lines_accepted,
    
    -- 3. Count chat interactions
    COUNTIF(labels.method = 'GenerateCode') as chat_interactions,
    
    CAST(NULL as INT64) as total_tokens,
    CAST(NULL as INT64) as input_tokens,
    CAST(NULL as INT64) as output_tokens,
    CAST(NULL as INT64) as cached_tokens,
    CAST(NULL as STRING) as model
    
  FROM `{PROJECT_ID}.{DATASET_ID}.cloudaicompanion_googleapis_com_metadata`
  WHERE timestamp >= TIMESTAMP(DATE_SUB(CURRENT_DATE(), INTERVAL 1 DAY))
  AND labels.user_id IS NOT NULL
  GROUP BY 1, 2, 3
)"""

CLI_CTE = f"""
cli_activity AS (
  SELECT
    -- Directly access the flattened columns
    jsonPayload.user_email as email,
    
    DATE(timestamp) as date,
    "Gemini CLI" as tool,
    
    COUNTIF(jsonPayload.decision = 'accept') as acceptances,
    
    SUM(SAFE_CAST(jsonPayload.lines AS INT64)) as lines_accepted,
    
    COUNTIF(jsonPayload.event_name = 'gemini_cli.user_prompt') as chat_interactions,

    SUM(SAFE_CAST(jsonPayload.total_token_count AS INT64)) as total_tokens,
    SUM(SAFE_CAST(jsonPayload.input_token_count AS INT64)) as input_tokens,
    SUM(SAFE_CAST(jsonPayload.output_token_count AS INT64)) as output_tokens,
    SUM(SAFE_CAST(jsonPayload.cached_content_token_count AS INT64)) as cached_tokens,

    MAX(jsonPayload.model) as model
    
  FROM `{PROJECT_ID}.{DATASET_ID}.gemini_cli`
  WHERE timestamp >= TIMESTAMP(DATE_SUB(CURRENT_DATE(), INTERVAL 1 DAY))
  AND jsonPayload.user_email IS NOT NULL
  GROUP BY 1, 2, 3
)"""

@functions_framework.http
def export_gemini_metrics(request):
    """
    Cloud Function triggered by Cloud Scheduler.
    Executes BigQuery analysis and pushes results to DX.
    """
    if not DX_API_TOKEN:
        logger.error("DX_API_KEY environment variable is missing.")
        return ("Configuration Error: Missing API Key", 500)

    try:
        bq_client = bigquery.Client()
        
        # 1. Dynamic Table Check
        # BigQuery will throw 400/404 if we query a table that doesn't exist yet.
        active_ctes = []
        active_selects = []
        
        # Check IDE Table
        ide_table_name = "cloudaicompanion_googleapis_com_metadata"
        try:
            bq_client.get_table(f"{PROJECT_ID}.{DATASET_ID}.{ide_table_name}")
            active_ctes.append(IDE_CTE)
            active_selects.append("SELECT * FROM ide_activity")
        except NotFound:
            logger.info(f"IDE table {ide_table_name} not found. Skipping.")

        # Check CLI Table
        cli_table_name = "gemini_cli"
        try:
            bq_client.get_table(f"{PROJECT_ID}.{DATASET_ID}.{cli_table_name}")
            active_ctes.append(CLI_CTE)
            active_selects.append("SELECT * FROM cli_activity")
        except NotFound:
            logger.info(f"CLI table {cli_table_name} not found. Skipping.")

        if not active_ctes:
            logger.info("No Gemini log tables found yet. Waiting for data.")
            return ("No tables found", 200)

        # 2. Construct Dynamic Query
        query = f"""
        WITH 
        {','.join(active_ctes)}
        
        {' UNION ALL '.join(active_selects)}
        """
        
        # 3. Execute Query
        query_job = bq_client.query(query)
        results = query_job.result()
        logger.info("BigQuery query completed successfully.")
        
        # 4. Transform Data
        payload_records = []
        for row in results:
            record = {
                "email": row.email,
                "date": row.date.isoformat(),
                "tool": row.tool, # "Gemini CLI" or "Gemini Code Assist"
                "is_active": True, 
                "metrics": {
                    "model": row.model,
                    "chat_interactions": row.chat_interactions,
                    "lines_accepted": row.lines_accepted if row.lines_accepted else 0,
                    "acceptances": row.acceptances,
                    "total_tokens": row.total_tokens if row.total_tokens else 0,
                    "input_tokens": row.input_tokens if row.input_tokens else 0,
                    "output_tokens": row.output_tokens if row.output_tokens else 0,
                    "cached_tokens": row.cached_tokens if row.cached_tokens else 0,
                }
            }
            payload_records.append(record)

        if not payload_records:
            return ("No Data Found for Yesterday", 200)

        # 5. Push to DX
        headers = {
            "Authorization": f"Bearer {DX_API_TOKEN}",
            "Content-Type": "application/json",
            "Accept": "application/json"
        }
        body = { "data": payload_records }
        
        logger.info(f"Pushing {len(payload_records)} records to DX...")

        combined_url = urljoin(DX_API_URL, "api/aiToolMetrics.pushAll")
        response = requests.post(combined_url, json=body, headers=headers, timeout=60)
        response.raise_for_status()

        logger.info(f"DX data payload: {body}")
        
        return (f"Success: {len(payload_records)} records exported", 200)

    except Exception as e:
        logger.error(f"Export failed: {str(e)}")
        return (f"Error: {str(e)}", 500)