#!/usr/bin/env python3
"""
Combine 56variations dataset + curated questions from sorted.json
Streams 3.7GB file to avoid memory issues.
Adds curated questions as new conversation types.
"""

import json
import re
import os
from pathlib import Path

def load_sorted_questions(filepath):
    """Load sorted.json into memory (65MB)"""
    with open(filepath, 'r', encoding='utf-8') as f:
        data = json.load(f)

    question_map = {}
    for item in data:
        sid = item['shloka_id']
        questions = item.get('questions', [])
        metadata = {
            'sanskrit': item.get('sanskrit', ''),
            'transliteration': item.get('transliteration', ''),
            'translation_en': item.get('translation_english', ''),
            'translation_hi': item.get('translation_hindi', ''),
            'section': item.get('source', {}).get('section', ''),
            'chapter': item.get('source', {}).get('chapter', ''),
            'verse': item.get('source', {}).get('verse', ''),
        }
        if questions:
            question_map[sid] = {'questions': questions, 'metadata': metadata}

    return question_map

def build_system_prompt(lang, section):
    """Build system prompt"""
    if lang == "hi":
        return f"You are an expert Ayurvedic scholar fluent in Hindi. Answer based on Charaka Samhita {section} with proper citations."
    elif lang == "sa":
        return f"You are an expert Ayurvedic scholar fluent in Sanskrit. Answer based on Charaka Samhita {section} with classical citations."
    return f"You are an expert Ayurvedic scholar trained in Charaka Samhita. Provide accurate responses with precise citations from {section}."

def build_curated_answer(metadata, lang):
    """Build answer using curated metadata"""
    en = metadata.get('translation_en', '')
    hi = metadata.get('translation_hi', '')
    sanskrit = metadata.get('sanskrit', '')
    trans = metadata.get('transliteration', '')
    section = metadata.get('section', '')
    chapter = metadata.get('chapter', '')
    verse = metadata.get('verse', '')

    if lang == 'hi':
        ans = hi or en
        if sanskrit:
            ans += f"\n\n**संस्कृत:** {sanskrit}"
        ans += f"\n\n**सन्दर्भ:** {section}, अध्याय {chapter}, श्लोक {verse}"
    elif lang == 'sa':
        ans = en
        if sanskrit:
            ans += f"\n\n**मूलम्:** {sanskrit}"
        ans += f"\n\n**सन्दर्भः:** {section}, अध्यायः {chapter}, श्लोकः {verse}"
    else:
        ans = en
        if sanskrit:
            ans += f"\n\n**Sanskrit:** {sanskrit}"
        if trans:
            ans += f"\n**Transliteration:** {trans}"
        ans += f"\n\n**Reference:** {section}, Chapter {chapter}, Verse {verse}"
    return ans

def stream_jsonl(filepath):
    """Stream JSONL one object at a time"""
    buffer = ""
    brace_count = 0
    in_object = False

    with open(filepath, 'r', encoding='utf-8') as f:
        while True:
            chunk = f.read(8192)
            if not chunk:
                break
            for char in chunk:
                if char == '{':
                    if brace_count == 0:
                        in_object = True
                        buffer = "{"
                    else:
                        buffer += char
                    brace_count += 1
                elif char == '}':
                    brace_count -= 1
                    buffer += char
                    if brace_count == 0 and in_object:
                        try:
                            yield json.loads(buffer)
                        except:
                            pass
                        buffer = ""
                        in_object = False
                elif in_object:
                    buffer += char

def merge_datasets():
    input_variations = "gemma4_ayurveda_56variations.jsonl"
    input_sorted = "shlokas_export_2026-05-14_SORTED.json"
    output_file = "gemma4_ayurveda_final_combined.jsonl"

    print("Loading curated questions...")
    question_map = load_sorted_questions(input_sorted)
    print(f"Loaded {len(question_map)} shlokas with curated questions")

    # Get sample metadata from sorted.json instead
    print("Reading sample metadata...")
    sample_metadata = {}
    with open(input_sorted, 'r', encoding='utf-8') as f:
        first_item = json.load(f)[0]
        sample_metadata = {
            "sanskrit": first_item.get('sanskrit', ''),
            "transliteration": first_item.get('transliteration', ''),
            "section": first_item.get('source', {}).get('section', ''),
            "section_code": first_item.get('source', {}).get('section', ''),
            "chapter": first_item.get('source', {}).get('chapter', ''),
            "verse": first_item.get('source', {}).get('verse', '')
        }

    # Pre-build curated conversations for all shlokas
    print("Building curated conversations...")
    curated_conversations = []
    added_count = 0
    missing_count = 0

    for sid, info in question_map.items():
        questions = info['questions']
        metadata = info['metadata']

        # Determine section name for system prompt
        section_code = sid.split('_')[1] if '_' in sid else 'SU'
        section_names = {
            'SU': 'Sutrasthana', 'NI': 'Nidanasthana', 'VI': 'Vimanasthana',
            'SHA': 'Sharirasthana', 'IN': 'Indriyasthana', 'CHI': 'Cikitsasthana',
            'KAL': 'Kalpasthana', 'SID': 'Siddhisthana'
        }
        section_name = section_names.get(section_code, 'Charaka Samhita')

        for qidx, q in enumerate(questions, 1):
            for lang in ['en', 'hi', 'sa']:
                q_data = q.get(lang, {})
                if not q_data.get('q'):
                    continue

                conv_type = f"curated_q{qidx}_{lang}"
                answer = q_data.get('a', '') or build_curated_answer(metadata, lang)

                conv = {
                    "shloka_id": sid,
                    "conversation_type": conv_type,
                    "messages": [
                        {"role": "system", "content": build_system_prompt(lang, section_name)},
                        {"role": "user", "content": q_data['q']},
                        {"role": "model", "content": answer}
                    ],
                    "metadata": {
                        "sanskrit": metadata['sanskrit'],
                        "transliteration": metadata['transliteration'],
                        "section": section_name,
                        "section_code": section_code,
                        "chapter": metadata['chapter'],
                        "verse": metadata['verse'],
                        "translation_en": metadata['translation_en'],
                        "translation_hi": metadata['translation_hi'],
                        "source": "curated"
                    }
                }
                curated_conversations.append(conv)
                added_count += 1

    print(f"Built {added_count} curated conversations")

    # Stream original + append curated
    print(f"\nStreaming {input_variations} and writing to {output_file}...")
    existing_count = 0

    with open(output_file, 'w', encoding='utf-8') as f:
        for conv in stream_jsonl(input_variations):
            f.write(json.dumps(conv, ensure_ascii=False, indent=2) + '\n')
            existing_count += 1
            if existing_count % 50000 == 0:
                print(f"  Copied {existing_count}...")

        # Append curated
        for conv in curated_conversations:
            f.write(json.dumps(conv, ensure_ascii=False, indent=2) + '\n')

    size_mb = os.path.getsize(output_file) / (1024 * 1024)
    total = existing_count + len(curated_conversations)

    print(f"\n{'='*60}")
    print("ULTIMATE DATASET COMPLETE")
    print(f"{'='*60}")
    print(f"Existing (56variations): {existing_count:,}")
    print(f"Curated added: {len(curated_conversations):,}")
    print(f"Total: {total:,}")
    print(f"File: {output_file}")
    print(f"Size: {size_mb:.1f} MB")

    # Show new conversation type counts
    type_counts = {}
    for conv in curated_conversations:
        t = conv['conversation_type']
        base = t.rsplit('_v', 1)[0] if '_v' in t else t
        type_counts[base] = type_counts.get(base, 0) + 1

    print(f"\nTop curated types:")
    for t, c in sorted(type_counts.items(), key=lambda x: -x[1])[:10]:
        print(f"  {t}: {c}")

if __name__ == "__main__":
    merge_datasets()
