"""
月次レポート生成サービス - データベース接続・クエリ

DB_TYPE設定値:
  - 'sqlite'    : ローカル開発用SQLite
  - 'neon'      : Neon PostgreSQL（外部ホスト、SSL必須）
  - 'lightsail' : Lightsail内PostgreSQL（localhost、SSL不要）
  - 'pgsql'     : 後方互換（neonと同じ動作）
"""
import sqlite3
import os
import sys
from datetime import datetime, timedelta
from calendar import monthrange

# 親ディレクトリのconfigをインポート
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from config import (
    DB_TYPE, SQLITE_PATH, PG_HOST, PG_PORT, PG_DBNAME, PG_USER, PG_PASSWORD,
    UNANSWERED_PATTERNS
)

# PG_SSLMODEが設定されている場合はインポート
try:
    from config import PG_SSLMODE
except ImportError:
    PG_SSLMODE = ''


def _is_postgres():
    """PostgreSQL系かどうかを判定"""
    return DB_TYPE in ('neon', 'lightsail', 'pgsql')


def _get_sslmode():
    """SSLモードを取得（DB_TYPEに応じたデフォルト値）"""
    if PG_SSLMODE:
        return PG_SSLMODE
    if DB_TYPE == 'neon' or DB_TYPE == 'pgsql':
        return 'require'
    if DB_TYPE == 'lightsail':
        return 'disable'
    return 'prefer'


def get_connection():
    """DB接続を取得（読み取り専用）"""
    if _is_postgres():
        import psycopg2
        sslmode = _get_sslmode()
        conn = psycopg2.connect(
            host=PG_HOST,
            port=PG_PORT,
            dbname=PG_DBNAME,
            user=PG_USER,
            password=PG_PASSWORD,
            sslmode=sslmode
        )
        return conn
    else:
        # SQLiteの相対パスを解決
        db_path = SQLITE_PATH
        if not os.path.isabs(db_path):
            db_path = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), db_path)
        conn = sqlite3.connect(db_path)
        conn.row_factory = sqlite3.Row
        return conn


def get_cursor(conn):
    """カーソルを取得（PostgreSQLはRealDictCursor）"""
    if _is_postgres():
        import psycopg2.extras
        return conn.cursor(cursor_factory=psycopg2.extras.RealDictCursor)
    return conn.cursor()


def _get_date_cast():
    """DBタイプに応じた日付キャストを返す"""
    if _is_postgres():
        return "timestamp::date"
    return "DATE(timestamp)"


def _get_hour_extract():
    """DBタイプに応じた時間抽出を返す"""
    if _is_postgres():
        return "EXTRACT(HOUR FROM timestamp)::INTEGER"
    return "CAST(strftime('%H', timestamp) AS INTEGER)"


def _get_day_of_week_extract():
    """DBタイプに応じた曜日抽出を返す（0=日曜, 6=土曜 → 0=月曜, 6=日曜に変換）"""
    if _is_postgres():
        # PostgreSQL: EXTRACT(DOW ...)は0=日曜なので、月曜始まりに変換
        # %%はpsycopg2で%リテラルを表す
        return "(EXTRACT(DOW FROM timestamp)::INTEGER + 6) %% 7"
    # SQLite: strftime('%w', ...)は0=日曜なので、月曜始まりに変換
    return "(CAST(strftime('%w', timestamp) AS INTEGER) + 6) % 7"


def _build_unanswered_condition():
    """未回答判定のSQL条件を構築"""
    conditions = []
    for pattern in UNANSWERED_PATTERNS:
        # PostgreSQLでは%をエスケープする必要がある
        if _is_postgres():
            conditions.append(f"assistant_message LIKE '%%{pattern}%%'")
        else:
            conditions.append(f"assistant_message LIKE '%{pattern}%'")
    return " OR ".join(conditions)


def _prepare_query(query: str):
    """DBタイプに応じてプレースホルダを変換"""
    if _is_postgres():
        # SQLiteの ? を PostgreSQLの %s に変換
        return query.replace('?', '%s')
    return query


def _execute(cursor, query: str, params: tuple = ()):
    """クエリを実行（プレースホルダ変換付き）"""
    converted_query = _prepare_query(query)
    cursor.execute(converted_query, params)


def get_conversations(start_date: str, end_date: str):
    """期間内の会話を取得"""
    conn = get_connection()
    cursor = get_cursor(conn)
    date_cast = _get_date_cast()

    query = f"""
        SELECT id, session_id, user_message, assistant_message, timestamp, ip_address, user_agent
        FROM conversations
        WHERE {date_cast} BETWEEN ? AND ?
        ORDER BY timestamp ASC
    """
    _execute(cursor, query, (start_date, end_date))
    rows = cursor.fetchall()
    conn.close()
    return [dict(row) for row in rows]


def get_feedbacks(start_date: str, end_date: str):
    """期間内のフィードバックを取得"""
    conn = get_connection()
    cursor = get_cursor(conn)
    date_cast = _get_date_cast()

    query = f"""
        SELECT id, session_id, feedback, question, answer, timestamp, ip_address
        FROM feedbacks
        WHERE {date_cast} BETWEEN ? AND ?
        ORDER BY timestamp ASC
    """
    _execute(cursor, query, (start_date, end_date))
    rows = cursor.fetchall()
    conn.close()
    return [dict(row) for row in rows]


def get_statistics_for_period(start_date: str, end_date: str):
    """期間の統計サマリー"""
    conn = get_connection()
    cursor = get_cursor(conn)
    date_cast = _get_date_cast()
    unanswered_cond = _build_unanswered_condition()

    stats = {}

    # 総会話数
    query = f"""
        SELECT COUNT(*) as count FROM conversations
        WHERE {date_cast} BETWEEN ? AND ?
    """
    _execute(cursor, query, (start_date, end_date))
    stats['total_conversations'] = cursor.fetchone()['count']

    # 回答できた数
    query = f"""
        SELECT COUNT(*) as count FROM conversations
        WHERE {date_cast} BETWEEN ? AND ?
        AND NOT ({unanswered_cond})
    """
    _execute(cursor, query, (start_date, end_date))
    stats['answered_count'] = cursor.fetchone()['count']

    # 回答できなかった数
    query = f"""
        SELECT COUNT(*) as count FROM conversations
        WHERE {date_cast} BETWEEN ? AND ?
        AND ({unanswered_cond})
    """
    _execute(cursor, query, (start_date, end_date))
    stats['unanswered_count'] = cursor.fetchone()['count']

    # ユニークセッション数
    query = f"""
        SELECT COUNT(DISTINCT session_id) as count FROM conversations
        WHERE {date_cast} BETWEEN ? AND ?
    """
    _execute(cursor, query, (start_date, end_date))
    stats['unique_sessions'] = cursor.fetchone()['count']

    # いいね数
    query = f"""
        SELECT COUNT(*) as count FROM feedbacks
        WHERE {date_cast} BETWEEN ? AND ?
        AND feedback = 'like'
    """
    _execute(cursor, query, (start_date, end_date))
    stats['likes_count'] = cursor.fetchone()['count']

    # よくない数
    query = f"""
        SELECT COUNT(*) as count FROM feedbacks
        WHERE {date_cast} BETWEEN ? AND ?
        AND feedback = 'dislike'
    """
    _execute(cursor, query, (start_date, end_date))
    stats['dislikes_count'] = cursor.fetchone()['count']

    # ユニークIP数（ユーザー数の参考値）
    query = f"""
        SELECT COUNT(DISTINCT ip_address) as count FROM conversations
        WHERE {date_cast} BETWEEN ? AND ?
        AND ip_address IS NOT NULL AND ip_address != ''
    """
    _execute(cursor, query, (start_date, end_date))
    stats['unique_ips'] = cursor.fetchone()['count']

    # 総フィードバック数
    stats['total_feedbacks'] = stats['likes_count'] + stats['dislikes_count']

    conn.close()
    return stats


def get_monthly_comparison(year: int, month: int):
    """前月比較データ"""
    # 当月の期間
    _, last_day = monthrange(year, month)
    current_start = f"{year}-{month:02d}-01"
    current_end = f"{year}-{month:02d}-{last_day:02d}"

    # 前月の期間
    if month == 1:
        prev_year = year - 1
        prev_month = 12
    else:
        prev_year = year
        prev_month = month - 1
    _, prev_last_day = monthrange(prev_year, prev_month)
    prev_start = f"{prev_year}-{prev_month:02d}-01"
    prev_end = f"{prev_year}-{prev_month:02d}-{prev_last_day:02d}"

    current_stats = get_statistics_for_period(current_start, current_end)
    prev_stats = get_statistics_for_period(prev_start, prev_end)

    def calc_change(current, previous):
        if previous == 0:
            return 100 if current > 0 else 0
        return round((current - previous) / previous * 100, 1)

    return {
        'current': current_stats,
        'previous': prev_stats,
        'changes': {
            'total_conversations': calc_change(current_stats['total_conversations'], prev_stats['total_conversations']),
            'answered_count': calc_change(current_stats['answered_count'], prev_stats['answered_count']),
            'unanswered_count': calc_change(current_stats['unanswered_count'], prev_stats['unanswered_count']),
            'likes_count': calc_change(current_stats['likes_count'], prev_stats['likes_count']),
            'dislikes_count': calc_change(current_stats['dislikes_count'], prev_stats['dislikes_count']),
        }
    }


def get_daily_statistics(start_date: str, end_date: str):
    """日別統計"""
    conn = get_connection()
    cursor = get_cursor(conn)
    date_cast = _get_date_cast()
    unanswered_cond = _build_unanswered_condition()

    query = f"""
        SELECT
            {date_cast} as date,
            COUNT(*) as total_conversations,
            COUNT(DISTINCT session_id) as unique_sessions,
            SUM(CASE WHEN ({unanswered_cond}) THEN 1 ELSE 0 END) as unanswered_count,
            SUM(CASE WHEN NOT ({unanswered_cond}) THEN 1 ELSE 0 END) as answered_count
        FROM conversations
        WHERE {date_cast} BETWEEN ? AND ?
        GROUP BY {date_cast}
        ORDER BY {date_cast} ASC
    """
    _execute(cursor, query, (start_date, end_date))
    rows = cursor.fetchall()
    conn.close()

    result = []
    for row in rows:
        d = dict(row)
        # 日付を文字列に変換
        if hasattr(d['date'], 'strftime'):
            d['date'] = d['date'].strftime('%Y-%m-%d')
        result.append(d)
    return result


def get_hourly_statistics(start_date: str, end_date: str):
    """時間帯別統計"""
    conn = get_connection()
    cursor = get_cursor(conn)
    date_cast = _get_date_cast()
    hour_extract = _get_hour_extract()
    unanswered_cond = _build_unanswered_condition()

    query = f"""
        SELECT
            {hour_extract} as hour,
            COUNT(*) as total_conversations,
            SUM(CASE WHEN ({unanswered_cond}) THEN 1 ELSE 0 END) as unanswered_count
        FROM conversations
        WHERE {date_cast} BETWEEN ? AND ?
        GROUP BY {hour_extract}
        ORDER BY hour ASC
    """
    _execute(cursor, query, (start_date, end_date))
    rows = cursor.fetchall()
    conn.close()
    return [dict(row) for row in rows]


def get_day_of_week_statistics(start_date: str, end_date: str):
    """曜日別統計（月曜=0, 日曜=6）"""
    conn = get_connection()
    cursor = get_cursor(conn)
    date_cast = _get_date_cast()
    dow_extract = _get_day_of_week_extract()
    unanswered_cond = _build_unanswered_condition()

    query = f"""
        SELECT
            {dow_extract} as day_of_week,
            COUNT(*) as total_conversations,
            SUM(CASE WHEN ({unanswered_cond}) THEN 1 ELSE 0 END) as unanswered_count
        FROM conversations
        WHERE {date_cast} BETWEEN ? AND ?
        GROUP BY {dow_extract}
        ORDER BY day_of_week ASC
    """
    _execute(cursor, query, (start_date, end_date))
    rows = cursor.fetchall()
    conn.close()
    return [dict(row) for row in rows]


def get_hour_by_day_of_week(start_date: str, end_date: str):
    """曜日×時間帯分布（ヒートマップ用）"""
    conn = get_connection()
    cursor = get_cursor(conn)
    date_cast = _get_date_cast()
    hour_extract = _get_hour_extract()
    dow_extract = _get_day_of_week_extract()

    query = f"""
        SELECT
            {dow_extract} as day_of_week,
            {hour_extract} as hour,
            COUNT(*) as count
        FROM conversations
        WHERE {date_cast} BETWEEN ? AND ?
        GROUP BY {dow_extract}, {hour_extract}
        ORDER BY day_of_week, hour
    """
    _execute(cursor, query, (start_date, end_date))
    rows = cursor.fetchall()
    conn.close()
    return [dict(row) for row in rows]


def get_average_turns(start_date: str, end_date: str):
    """平均会話ターン数（セッションあたり）"""
    conn = get_connection()
    cursor = get_cursor(conn)
    date_cast = _get_date_cast()

    query = f"""
        SELECT
            session_id,
            COUNT(*) as turn_count
        FROM conversations
        WHERE {date_cast} BETWEEN ? AND ?
        GROUP BY session_id
    """
    _execute(cursor, query, (start_date, end_date))
    rows = cursor.fetchall()
    conn.close()

    if not rows:
        return {'average': 0, 'max': 0, 'min': 0, 'total_sessions': 0}

    turns = [row['turn_count'] for row in rows]
    return {
        'average': round(sum(turns) / len(turns), 2),
        'max': max(turns),
        'min': min(turns),
        'total_sessions': len(turns)
    }


def get_unresolved_by_hour(start_date: str, end_date: str):
    """未解決が多い時間帯"""
    conn = get_connection()
    cursor = get_cursor(conn)
    date_cast = _get_date_cast()
    hour_extract = _get_hour_extract()
    unanswered_cond = _build_unanswered_condition()

    query = f"""
        SELECT
            {hour_extract} as hour,
            COUNT(*) as unanswered_count
        FROM conversations
        WHERE {date_cast} BETWEEN ? AND ?
        AND ({unanswered_cond})
        GROUP BY {hour_extract}
        ORDER BY unanswered_count DESC
    """
    _execute(cursor, query, (start_date, end_date))
    rows = cursor.fetchall()
    conn.close()
    return [dict(row) for row in rows]


def get_keywords(start_date: str, end_date: str, limit: int = 15):
    """頻出キーワード"""
    conn = get_connection()
    cursor = get_cursor(conn)
    date_cast = _get_date_cast()

    query = f"""
        SELECT user_message FROM conversations
        WHERE {date_cast} BETWEEN ? AND ?
    """
    _execute(cursor, query, (start_date, end_date))
    rows = cursor.fetchall()
    conn.close()

    return _extract_keywords([row['user_message'] for row in rows], limit)


def get_unresolved_keywords(start_date: str, end_date: str, limit: int = 15):
    """未解決に多いキーワード"""
    conn = get_connection()
    cursor = get_cursor(conn)
    date_cast = _get_date_cast()
    unanswered_cond = _build_unanswered_condition()

    query = f"""
        SELECT user_message FROM conversations
        WHERE {date_cast} BETWEEN ? AND ?
        AND ({unanswered_cond})
    """
    _execute(cursor, query, (start_date, end_date))
    rows = cursor.fetchall()
    conn.close()

    return _extract_keywords([row['user_message'] for row in rows], limit)


def _extract_keywords(messages: list, limit: int = 15):
    """メッセージからキーワードを抽出"""
    import re

    stop_words = {'は', 'が', 'の', 'に', 'を', 'で', 'と', 'です', 'ます', 'か', 'ですか',
                  'ますか', 'について', 'とは', 'なん', 'どう', 'どの', 'この', 'その',
                  'いる', 'ある', 'する', 'なる', 'できる', 'れる', 'られる', 'ない',
                  'たい', 'ほしい', 'ください', 'おねがい', 'しまう', 'てる', 'てい'}

    keywords = {}
    for message in messages:
        if not message:
            continue
        # 日本語の単語を抽出（2文字以上）
        words = re.findall(r'[ぁ-んァ-ヶー一-龠々〆ヵヶ]{2,}', message)
        for word in words:
            if word not in stop_words and len(word) >= 2:
                keywords[word] = keywords.get(word, 0) + 1

    # 出現回数でソートして上位を返す
    sorted_keywords = sorted(keywords.items(), key=lambda x: x[1], reverse=True)
    return sorted_keywords[:limit]


def get_unanswered_questions(start_date: str, end_date: str):
    """回答できなかった質問一覧"""
    conn = get_connection()
    cursor = get_cursor(conn)
    date_cast = _get_date_cast()
    unanswered_cond = _build_unanswered_condition()

    query = f"""
        SELECT user_message, timestamp, session_id
        FROM conversations
        WHERE {date_cast} BETWEEN ? AND ?
        AND ({unanswered_cond})
        ORDER BY timestamp DESC
    """
    _execute(cursor, query, (start_date, end_date))
    rows = cursor.fetchall()
    conn.close()
    return [dict(row) for row in rows]


def get_recent_disliked(start_date: str, end_date: str, limit: int = 10):
    """低評価フィードバック"""
    conn = get_connection()
    cursor = get_cursor(conn)
    date_cast = _get_date_cast()

    query = f"""
        SELECT question, answer, timestamp
        FROM feedbacks
        WHERE {date_cast} BETWEEN ? AND ?
        AND feedback = 'dislike'
        ORDER BY timestamp DESC
        LIMIT ?
    """
    _execute(cursor, query, (start_date, end_date, limit))
    rows = cursor.fetchall()
    conn.close()
    return [dict(row) for row in rows]


def get_feedback_daily_stats(start_date: str, end_date: str):
    """日別フィードバック統計"""
    conn = get_connection()
    cursor = get_cursor(conn)
    date_cast = _get_date_cast()

    query = f"""
        SELECT
            {date_cast} as date,
            SUM(CASE WHEN feedback = 'like' THEN 1 ELSE 0 END) as likes,
            SUM(CASE WHEN feedback = 'dislike' THEN 1 ELSE 0 END) as dislikes,
            COUNT(*) as total
        FROM feedbacks
        WHERE {date_cast} BETWEEN ? AND ?
        GROUP BY {date_cast}
        ORDER BY {date_cast} ASC
    """
    _execute(cursor, query, (start_date, end_date))
    rows = cursor.fetchall()
    conn.close()

    result = []
    for row in rows:
        d = dict(row)
        if hasattr(d['date'], 'strftime'):
            d['date'] = d['date'].strftime('%Y-%m-%d')
        result.append(d)
    return result


def load_category_mapping():
    """category.txtからカテゴリマッピングを読み込み"""
    category_path = os.path.join(os.path.dirname(os.path.dirname(__file__)), 'category.txt')
    categories = {}

    if not os.path.exists(category_path):
        return categories

    try:
        with open(category_path, 'r', encoding='utf-8') as f:
            for line in f:
                line = line.strip()
                if ':' in line:
                    cat_id, cat_name = line.split(':', 1)
                    # カテゴリIDの最初の数字部分を取得（例: "6-1" -> "6"）
                    main_id = cat_id.split('-')[0]
                    if main_id not in categories:
                        categories[main_id] = cat_name.strip().strip('"')
    except Exception as e:
        print(f"カテゴリファイル読み込みエラー: {e}")

    return categories


def get_qa_usage_ranking(start_date: str, end_date: str, limit: int = 15):
    """よく参照されるQ&Aランキング"""
    import re

    conn = get_connection()
    cursor = get_cursor(conn)
    date_cast = _get_date_cast()

    query = f"""
        SELECT assistant_message FROM conversations
        WHERE {date_cast} BETWEEN ? AND ?
    """
    _execute(cursor, query, (start_date, end_date))
    rows = cursor.fetchall()
    conn.close()

    # Q&A IDを抽出してカウント
    qa_counts = {}
    for row in rows:
        msg = row['assistant_message'] or ''
        # **Q7:**, **Q1-1:**, **Q123:** などの形式を抽出（コロン付き）
        ids = re.findall(r'\*\*Q(\d+(?:-\d+)?):', msg)
        for qa_id in ids:
            qa_counts[qa_id] = qa_counts.get(qa_id, 0) + 1

    # 出現回数でソート
    sorted_qa = sorted(qa_counts.items(), key=lambda x: x[1], reverse=True)
    return sorted_qa[:limit]


def get_category_usage_ranking(start_date: str, end_date: str, limit: int = 10):
    """カテゴリ別Q&A参照ランキング"""
    import re

    conn = get_connection()
    cursor = get_cursor(conn)
    date_cast = _get_date_cast()

    query = f"""
        SELECT assistant_message FROM conversations
        WHERE {date_cast} BETWEEN ? AND ?
    """
    _execute(cursor, query, (start_date, end_date))
    rows = cursor.fetchall()
    conn.close()

    # カテゴリマッピングを読み込み
    categories = load_category_mapping()

    # カテゴリ別にカウント
    category_counts = {}
    for row in rows:
        msg = row['assistant_message'] or ''
        # **Q7:**, **Q1-1:**, **Q123:** などの形式を抽出
        ids = re.findall(r'\*\*Q(\d+(?:-\d+)?):', msg)
        for qa_id in ids:
            # カテゴリID（最初の数字部分）を取得
            cat_id = qa_id.split('-')[0]
            cat_name = categories.get(cat_id, f'カテゴリ{cat_id}')

            if cat_id not in category_counts:
                category_counts[cat_id] = {'name': cat_name, 'count': 0}
            category_counts[cat_id]['count'] += 1

    # 出現回数でソート
    sorted_categories = sorted(
        category_counts.items(),
        key=lambda x: x[1]['count'],
        reverse=True
    )

    # [(cat_id, cat_name, count), ...] の形式で返す
    result = [(cat_id, data['name'], data['count']) for cat_id, data in sorted_categories]
    return result[:limit]
