from flask import Flask, render_template, request, redirect, url_for, jsonify, session
import mysql.connector
import requests
import re
import logging
import tiktoken
from functools import wraps
from werkzeug.middleware.proxy_fix import ProxyFix
from config import (
    MYSQL_CONFIG, MYSQL_CONNECTION_STRING, WEBHOOK_URL, WEBHOOK_TIMEOUT, 
    APP_PORT, APP_PASSWORD, APP_VERSION, PROXY_PREFIX
)

# Configure logging
logging.basicConfig(level=logging.DEBUG)
logger = logging.getLogger(__name__)

app = Flask(__name__)
app.secret_key = APP_PASSWORD  # Using APP_PASSWORD as secret key for simplicity

# Configure application for proxy
app.wsgi_app = ProxyFix(app.wsgi_app, x_prefix=1)

# Update url generation to include proxy prefix
class ReverseProxied(object):
    def __init__(self, app, script_name=None):
        self.app = app
        self.script_name = script_name

    def __call__(self, environ, start_response):
        script_name = self.script_name
        if script_name:
            environ['SCRIPT_NAME'] = script_name
            path_info = environ['PATH_INFO']
            if path_info.startswith(script_name):
                environ['PATH_INFO'] = path_info[len(script_name):]
        return self.app(environ, start_response)

app.wsgi_app = ReverseProxied(app.wsgi_app, script_name=PROXY_PREFIX)

# Configure Flask to generate URLs with proxy prefix
def url_for_with_prefix(*args, **kwargs):
    url = url_for(*args, **kwargs)
    if not url.startswith(PROXY_PREFIX):
        url = PROXY_PREFIX + url
    return url

# Override url_for in templates
app.jinja_env.globals['url_for'] = url_for_with_prefix

def count_tokens(text):
    """Count tokens using tiktoken"""
    try:
        encoding = tiktoken.encoding_for_model("gpt-3.5-turbo")
        token_count = len(encoding.encode(text))
        return token_count
    except Exception as e:
        logger.error(f"Error counting tokens: {str(e)}")
        return 0

# Login decorator
def login_required(f):
    @wraps(f)
    def decorated_function(*args, **kwargs):
        if 'logged_in' not in session:
            return redirect(url_for('login'))
        return f(*args, **kwargs)
    return decorated_function

def get_db_connection():
    logger.debug("Attempting to establish database connection")
    try:
        conn = mysql.connector.connect(
            user=MYSQL_CONFIG['user'],
            password=MYSQL_CONFIG['password'],
            host=MYSQL_CONFIG['host'],
            port=MYSQL_CONFIG['port'],
            database=MYSQL_CONFIG['database'],
            ssl_verify_cert=False
        )
        logger.info("Database connection established successfully")
        return conn
    except mysql.connector.Error as err:
        logger.error(f"Error connecting to MySQL: {err}")
        return None

def get_database_views():
    conn = get_db_connection()
    if conn:
        try:
            cursor = conn.cursor()
            cursor.execute("SHOW FULL TABLES WHERE Table_type = 'VIEW'")
            views = [view[0] for view in cursor.fetchall()]
            logger.debug(f"Retrieved views: {views}")
            return views
        except mysql.connector.Error as err:
            logger.error(f"Error fetching views: {err}")
            return []
        finally:
            cursor.close()
            conn.close()
    return []

def get_view_columns(view_name):
    logger.debug(f"Fetching columns for view: {view_name}")
    conn = get_db_connection()
    if conn:
        try:
            cursor = conn.cursor()
            cursor.execute(f"SHOW FULL COLUMNS FROM {view_name}")
            columns = [{'name': column[0], 'description': column[8] or ''} for column in cursor.fetchall()]
            logger.debug(f"Retrieved columns with descriptions: {columns}")
            return columns
        except mysql.connector.Error as err:
            logger.error(f"Error fetching columns: {err}")
            return []
        finally:
            cursor.close()
            conn.close()
    return []

def is_valid_sql(sql):
    logger.debug(f"Validating SQL query: {sql}")
    sql_lower = sql.lower().strip()
    
    # Check if it's a SELECT statement
    if not sql_lower.startswith('select'):
        logger.warning("SQL validation failed: Not a SELECT statement")
        return False
    
    # Check for dangerous keywords
    dangerous_keywords = ['drop', 'delete', 'update', 'insert', 'alter', 'truncate', 'create', 'replace']
    for keyword in dangerous_keywords:
        if keyword in sql_lower:
            logger.warning(f"SQL validation failed: Contains dangerous keyword '{keyword}'")
            return False
    
    logger.info("SQL validation passed")
    return True

def execute_sql_query(sql):
    logger.debug(f"Attempting to execute SQL query: {sql}")
    
    if not is_valid_sql(sql):
        logger.error("Invalid SQL query rejected")
        return "Invalid SQL query. Only SELECT statements are allowed."
    
    conn = get_db_connection()
    if conn:
        try:
            cursor = conn.cursor(dictionary=True)
            logger.info("Executing SQL query")
            cursor.execute(sql)
            results = cursor.fetchall()
            logger.info(f"Query executed successfully. Retrieved {len(results)} rows")
            logger.debug(f"Query results: {results}")
            return results
        except mysql.connector.Error as err:
            logger.error(f"Error executing query: {err}")
            return f"Error executing query: {err}"
        finally:
            cursor.close()
            conn.close()
            logger.debug("Database connection closed")
    return "Failed to connect to database"

@app.route('/login', methods=['GET', 'POST'])
def login():
    if request.method == 'POST':
        if request.form['password'] == APP_PASSWORD:
            session['logged_in'] = True
            return redirect(url_for('index'))
        return render_template('login.html', error='Invalid password', version=APP_VERSION)
    return render_template('login.html', version=APP_VERSION)

@app.route('/logout')
def logout():
    session.pop('logged_in', None)
    return redirect(url_for('login'))

@app.route('/')
@login_required
def index():
    views = get_database_views()
    return render_template('index.html', views=views, version=APP_VERSION)

@app.route('/question/<view_name>')
@login_required
def question_form(view_name):
    columns = get_view_columns(view_name)
    return render_template('question.html', view_name=view_name, columns=columns, version=APP_VERSION)

@app.route('/count_tokens', methods=['POST'])
@login_required
def token_counter():
    text = request.json.get('text', '')
    token_count = count_tokens(text)
    return jsonify({'count': token_count})

@app.route('/submit', methods=['POST'])
@login_required
def submit():
    logger.info("Received form submission")
    view_name = request.form['view_name']
    question = request.form['question']
    analysis_type = request.form['analysis_type']
    
    logger.debug(f"Form data - View: {view_name}, Question: {question}, Analysis Type: {analysis_type}")
    
    # Send to webhook
    webhook_data = {
        'view_name': view_name,
        'question': question,
        'analysis_type': analysis_type
    }
    
    try:
        logger.info(f"Sending request to webhook: {WEBHOOK_URL}")
        response = requests.post(WEBHOOK_URL, json=webhook_data, timeout=WEBHOOK_TIMEOUT)
        logger.debug(f"Webhook response status: {response.status_code}")
        logger.debug(f"Webhook response content: {response.text}")
        
        response.raise_for_status()
        
        # Try to parse as JSON first
        try:
            webhook_response = response.json()
            sql_query = webhook_response.get('sql')
        except ValueError:
            # If not JSON, treat the entire response as SQL
            sql_query = response.text.strip()
        
        logger.debug(f"Extracted SQL query: {sql_query}")
        
        if sql_query:
            logger.info("SQL query received from webhook")
            # Execute the SQL query
            results = execute_sql_query(sql_query)
            
            # Format results for display
            if isinstance(results, list):
                formatted_results = '\n'.join([str(row) for row in results])
            else:
                formatted_results = str(results)
            
            # Count tokens in the formatted results
            token_count = count_tokens(formatted_results)
            
            logger.debug(f"Formatted results: {formatted_results}")
            
            # Split results if token count exceeds 100k
            if token_count > 100000:
                logger.info("Token count exceeds 100k, splitting results")
                split_results = []
                current_tokens = 0
                current_part = []
                
                for row in results:
                    row_str = str(row)
                    row_tokens = count_tokens(row_str)
                    
                    if current_tokens + row_tokens > 100000:
                        split_results.append('\n'.join(current_part))
                        current_part = [row_str]
                        current_tokens = row_tokens
                    else:
                        current_part.append(row_str)
                        current_tokens += row_tokens
                
                if current_part:
                    split_results.append('\n'.join(current_part))
                
                response_data = {
                    'sql': sql_query,
                    'results': split_results,
                    'token_count': token_count
                }
            else:
                response_data = {
                    'sql': sql_query,
                    'results': [formatted_results],
                    'token_count': token_count
                }
            
            logger.info("Sending successful response back to client")
            logger.debug(f"Response data: {response_data}")
            return jsonify(response_data)
        
        logger.warning("No SQL query found in webhook response")
        return jsonify({
            'error': 'No SQL query received from webhook'
        }), 400
        
    except requests.exceptions.RequestException as e:
        logger.error(f"Webhook request failed: {str(e)}")
        return jsonify({
            'error': f'Error processing request: {str(e)}'
        }), 500
    except Exception as e:
        logger.error(f"Unexpected error: {str(e)}", exc_info=True)
        return jsonify({
            'error': f'Unexpected error: {str(e)}'
        }), 500

if __name__ == '__main__':
    app.run(host='0.0.0.0', port=APP_PORT, debug=True)
