import sqlite3
import os
import sys
import json
import random
from pathlib import Path
from datetime import datetime, timedelta

SCRIPTS_DIR = os.path.dirname(os.path.abspath(__file__))
ROOT_DIR = os.path.dirname(SCRIPTS_DIR)
DB_DIR = Path(ROOT_DIR) / "database"
EXPORT_DIR = Path(ROOT_DIR) / "export" / "wordpress_xml"


def get_databases():
    if not DB_DIR.exists():
        print(f"Database folder not found: {DB_DIR}")
        sys.exit(1)
    dbs = sorted(DB_DIR.glob("*.sqlite"))
    if not dbs:
        print("No .sqlite files found in database/")
        sys.exit(1)
    return dbs


def count_exportable(db_path):
    conn = sqlite3.connect(str(db_path), timeout=30)
    count = conn.execute(
        "SELECT COUNT(*) FROM posts WHERE images IS NOT NULL AND images != '' AND images != '[]' AND ai_content IS NOT NULL AND TRIM(ai_content) != ''"
    ).fetchone()[0]
    conn.close()
    return count


def stream_posts(db_path, offset, limit):
    conn = sqlite3.connect(str(db_path), timeout=30)
    conn.row_factory = sqlite3.Row
    rows = conn.execute(
        """SELECT id, keyword, slug, images, ai_title, ai_content, created_at 
           FROM posts 
           WHERE images IS NOT NULL AND images != '' AND images != '[]'
             AND ai_content IS NOT NULL AND TRIM(ai_content) != ''
           ORDER BY id ASC
           LIMIT ? OFFSET ?""",
        (limit, offset)
    ).fetchall()
    conn.close()
    return rows


def escape_xml(s):
    if s is None:
        return ""
    s = str(s)
    s = s.replace("&", "&amp;")
    s = s.replace("<", "&lt;")
    s = s.replace(">", "&gt;")
    s = s.replace('"', "&quot;")
    s = s.replace("'", "&apos;")
    return s


def cdata(s):
    if s is None:
        return "<![CDATA[]]>"
    return f"<![CDATA[{s}]]>"


def parse_images(images_json):
    try:
        images = json.loads(images_json)
        if isinstance(images, list):
            return images
    except (json.JSONDecodeError, TypeError):
        pass
    return []


def format_wp_date(dt):
    return dt.strftime("%Y-%m-%d %H:%M:%S")


def parse_date_input(s, default=None):
    for fmt in ("%Y-%m-%d", "%d-%m-%Y", "%d/%m/%Y", "%Y/%m/%d"):
        try:
            return datetime.strptime(s.strip(), fmt)
        except ValueError:
            continue
    return default


def generate_dates_even(total_posts, start_dt, end_dt):
    total_days = (end_dt - start_dt).days
    if total_days <= 0:
        return [start_dt] * total_posts

    posts_per_day = total_posts / total_days
    dates = []
    current_day = 0
    for i in range(total_posts):
        day_offset = int(i / posts_per_day)
        if day_offset >= total_days:
            day_offset = total_days - 1
        base = start_dt + timedelta(days=day_offset)
        hour = random.randint(6, 22)
        minute = random.randint(0, 59)
        second = random.randint(0, 59)
        dates.append(base.replace(hour=hour, minute=minute, second=second))
    return dates


def generate_dates_random(total_posts, start_dt, end_dt):
    total_days = (end_dt - start_dt).days
    if total_days <= 0:
        return [start_dt] * total_posts

    dates = []
    for _ in range(total_posts):
        day_offset = random.randint(0, total_days)
        base = start_dt + timedelta(days=day_offset)
        hour = random.randint(6, 22)
        minute = random.randint(0, 59)
        second = random.randint(0, 59)
        dates.append(base.replace(hour=hour, minute=minute, second=second))
    dates.sort()
    return dates


IMAGE_INSERT_POSITIONS = [2, 6, 9, 13, 17]


def inject_images_into_content(content, images):
    if not images or not content:
        return content

    img_urls = []
    for img in images:
        url = img.get("image_url", "") if isinstance(img, dict) else ""
        alt = img.get("title", "") if isinstance(img, dict) else ""
        if url:
            img_urls.append((url, alt))

    if not img_urls:
        return content

    parts = content.split("</p>")
    if len(parts) <= 1:
        return content

    img_idx = 0
    result = []
    for i, part in enumerate(parts):
        result.append(part)
        para_num = i + 1

        if img_idx < len(img_urls) and para_num in IMAGE_INSERT_POSITIONS:
            url, alt = img_urls[img_idx]
            result.append(f'\n<p><img src="{url}" alt="{alt}" /></p>')
            img_idx += 1

        if i < len(parts) - 1:
            result.append("</p>")

    while img_idx < len(img_urls):
        url, alt = img_urls[img_idx]
        result.append(f'\n<p><img src="{url}" alt="{alt}" /></p>')
        img_idx += 1

    return "".join(result)


def build_wxr_channel_header(site_url, site_title):
    return f"""<?xml version="1.0" encoding="UTF-8" ?>
<rss version="2.0"
  xmlns:wp="http://wordpress.org/export/1.2/"
  xmlns:content="http://purl.org/rss/1.0/modules/content/"
  xmlns:dc="http://purl.org/dc/elements/1.1/"
  xmlns:excerpt="http://wordpress.org/export/1.2/excerpt/"
  xmlns:wfw="http://wellformedweb.org/CommentAPI/"
  xmlns:slash="http://purl.org/rss/1.0/modules/slash/"
  xmlns:media="http://search.yahoo.com/mrss/">
<channel>
  <title>{escape_xml(site_title)}</title>
  <link>{escape_xml(site_url)}</link>
  <description></description>
  <language>en-US</language>
  <wp:wxr_version>1.2</wp:wxr_version>
  <wp:base_site_url>{escape_xml(site_url)}</wp:base_site_url>
  <wp:base_blog_url>{escape_xml(site_url)}</wp:base_blog_url>
"""


def build_item(post_id, title, slug, content, post_date, author, images, site_url):
    post_date_str = format_wp_date(post_date)
    link = f"{site_url}/{slug}.html"

    full_content = inject_images_into_content(content or "", images)

    item_xml = f"""  <item>
    <title>{escape_xml(title)}</title>
    <link>{escape_xml(link)}</link>
    <pubDate>{escape_xml(post_date_str)}</pubDate>
    <dc:creator>{cdata(author)}</dc:creator>
    <guid isPermaLink="false">{escape_xml(link)}</guid>
    <description></description>
    <content:encoded>{cdata(full_content)}</content:encoded>
    <excerpt:encoded>{cdata('')}</excerpt:encoded>
    <wp:post_id>{post_id}</wp:post_id>
    <wp:post_date>{cdata(post_date_str)}</wp:post_date>
    <wp:post_date_gmt>{cdata(post_date_str)}</wp:post_date_gmt>
    <wp:post_name>{cdata(slug)}</wp:post_name>
    <wp:status>{cdata('publish')}</wp:status>
    <wp:post_type>{cdata('post')}</wp:post_type>
    <wp:is_sticky>0</wp:is_sticky>
  </item>
"""
    return item_xml


def export_chunk(posts, dates, file_index, site_url, site_title, author):
    os.makedirs(EXPORT_DIR, exist_ok=True)
    filename = f"wordpress_export_{file_index:03d}.xml"
    filepath = EXPORT_DIR / filename

    with open(filepath, 'w', encoding='utf-8') as f:
        f.write(build_wxr_channel_header(site_url, site_title))

        for i, row in enumerate(posts):
            post_id = row['id']
            keyword = row['keyword'] or ''
            slug = row['slug'] or ''
            images_json = row['images'] or '[]'
            ai_title = row['ai_title'] or ''
            ai_content = row['ai_content'] or ''

            title = ai_title if ai_title else keyword.title()
            images = parse_images(images_json)
            post_date = dates[i] if i < len(dates) else datetime.now()

            f.write(build_item(post_id, title, slug, ai_content, post_date, author, images, site_url))

        f.write("</channel>\n</rss>\n")

    return filepath


def main():
    print("=" * 50)
    print("   WordPress XML Exporter")
    print("=" * 50)
    print()

    dbs = get_databases()

    total_exportable = 0
    db_stats = []
    for db_path in dbs:
        count = count_exportable(db_path)
        db_stats.append((db_path, count))
        total_exportable += count

    if total_exportable == 0:
        print("No exportable posts found (need images + ai_content).")
        return

    print(f"Found {total_exportable} exportable posts across {len(dbs)} database(s):")
    for db_path, count in db_stats:
        print(f"  {db_path.stem}: {count} posts")
    print()

    posts_per_xml = input(f"Posts per XML file [{total_exportable}]: ").strip()
    if posts_per_xml.isdigit() and int(posts_per_xml) > 0:
        posts_per_xml = int(posts_per_xml)
    else:
        posts_per_xml = total_exportable

    max_xml_files = input(f"Max XML files to create (0=all) [0]: ").strip()
    if max_xml_files.isdigit() and int(max_xml_files) > 0:
        max_xml_files = int(max_xml_files)
    else:
        max_xml_files = 0

    max_posts = posts_per_xml * max_xml_files if max_xml_files > 0 else total_exportable
    if max_posts > total_exportable:
        max_posts = total_exportable

    total_xml_files = (max_posts + posts_per_xml - 1) // posts_per_xml

    print()
    print(f"Posts per XML  : {posts_per_xml}")
    print(f"Total XML files: {total_xml_files}")
    print(f"Total posts    : {max_posts}")
    print()

    site_url = input("Site URL (e.g. https://storage.googleapis.com/bucket-name) [https://example.com]: ").strip()
    if not site_url:
        site_url = "https://example.com"
    site_url = site_url.rstrip('/')

    site_title = input("Site title [My Blog]: ").strip()
    if not site_title:
        site_title = "My Blog"

    author = input("Author name [admin]: ").strip()
    if not author:
        author = "admin"

    print()
    print("-" * 50)
    print("   Date Range Settings")
    print("-" * 50)
    print()

    today = datetime.now()
    default_start = (today - timedelta(days=365)).strftime("%Y-%m-%d")
    default_end = today.strftime("%Y-%m-%d")

    start_input = input(f"Start date (YYYY-MM-DD) [{default_start}]: ").strip()
    start_dt = parse_date_input(start_input)
    if not start_dt:
        start_dt = datetime.strptime(default_start, "%Y-%m-%d")

    end_input = input(f"End date (YYYY-MM-DD) [{default_end}]: ").strip()
    end_dt = parse_date_input(end_input)
    if not end_dt:
        end_dt = datetime.strptime(default_end, "%Y-%m-%d")

    if end_dt < start_dt:
        start_dt, end_dt = end_dt, start_dt
        print("  (dates swapped: start must be before end)")

    total_days = (end_dt - start_dt).days

    print()
    print("  Distribution mode:")
    print("    1. Even  - posts spread equally across all days")
    print("    2. Random - posts randomly assigned to days")
    print()
    dist_input = input("Distribution [1]: ").strip()
    dist_mode = "random" if dist_input == "2" else "even"

    print()
    print(f"Start date     : {start_dt.strftime('%Y-%m-%d')}")
    print(f"End date       : {end_dt.strftime('%Y-%m-%d')}")
    print(f"Total days     : {total_days}")
    print(f"Distribution   : {dist_mode}")
    print()

    print("=" * 50)
    print("   Exporting...")
    print("=" * 50)
    print()

    all_posts = []
    for db_path, db_count in db_stats:
        if len(all_posts) >= max_posts:
            break
        remaining = max_posts - len(all_posts)
        fetch_count = min(db_count, remaining)
        posts = stream_posts(db_path, 0, fetch_count)
        all_posts.extend(posts)

    if dist_mode == "even":
        all_dates = generate_dates_even(len(all_posts), start_dt, end_dt)
    else:
        all_dates = generate_dates_random(len(all_posts), start_dt, end_dt)

    file_index = 1
    for chunk_start in range(0, len(all_posts), posts_per_xml):
        chunk_end = min(chunk_start + posts_per_xml, len(all_posts))
        chunk_posts = all_posts[chunk_start:chunk_end]
        chunk_dates = all_dates[chunk_start:chunk_end]

        filepath = export_chunk(chunk_posts, chunk_dates, file_index, site_url, site_title, author)
        print(f"  [{file_index:03d}] {filepath.name} ({len(chunk_posts)} posts)")
        file_index += 1

    total_files = file_index - 1
    print()
    print(f"Exported {len(all_posts)} posts into {total_files} XML file(s)")
    print(f"Output: {EXPORT_DIR}")
    print()
    print("=" * 50)
    print("   Export completed!")
    print("=" * 50)


if __name__ == "__main__":
    main()
