
import re
from pathlib import Path

def check_tex(tex_path):
    with open(tex_path, encoding='utf-8') as f:
        lines = f.readlines()

    content = ''.join(lines)
    figures = list(re.finditer(r'\\begin{figure}.*?\\end{figure}', content, re.DOTALL))
    tables = list(re.finditer(r'\\begin{table}.*?\\end{table}', content, re.DOTALL))
    refs = list(re.finditer(r'\\(?:auto)?ref{([^}]+)}', content))

    float_blocks = []
    label_counts = {}

    groups = {
        'main_figures': [],
        'supp_figures': [],
        'main_tables': [],
        'supp_tables': []
    }

    for match in figures + tables:
        block = match.group()
        label_match = re.search(r'\\label{([^}]+)}', block)
        caption_match = re.search(r'\\caption{([^}]+)}', block)

        label = label_match.group(1) if label_match else 'MISSING_LABEL'
        caption = caption_match.group(1) if caption_match else 'MISSING_CAPTION'
        float_type = 'figure' if 'figure' in block else 'table'
        line_number = content[:match.start()].count('\n') + 1

        float_blocks.append({
            'type': float_type,
            'label': label,
            'caption': caption,
            'line': line_number
        })

        if label != 'MISSING_LABEL':
            label_counts[label] = label_counts.get(label, 0) + 1
            if float_type == 'figure':
                if label.startswith("fgr:S"):
                    groups['supp_figures'].append(label)
                else:
                    groups['main_figures'].append(label)
            else:
                if label.startswith("tbl:S"):
                    groups['supp_tables'].append(label)
                else:
                    groups['main_tables'].append(label)

    # Record first citation of each label
    first_refs = {}
    for m in refs:
        label = m.group(1)
        if label not in first_refs:
            line_no = content[:m.start()].count('\n') + 1
            first_refs[label] = line_no

    # Sort references into the same four groups
    refs_by_group = {
        'main_figures': [],
        'supp_figures': [],
        'main_tables': [],
        'supp_tables': []
    }

    for label in first_refs:
        if label in groups['main_figures']:
            refs_by_group['main_figures'].append(label)
        elif label in groups['supp_figures']:
            refs_by_group['supp_figures'].append(label)
        elif label in groups['main_tables']:
            refs_by_group['main_tables'].append(label)
        elif label in groups['supp_tables']:
            refs_by_group['supp_tables'].append(label)

    def check_order(refs_list, defined_list, typename):
        errors = []
        last_index = -1
        for label in refs_list:
            if label in defined_list:
                idx = defined_list.index(label)
                if idx < last_index:
                    line_no = first_refs[label]
                    errors.append(f" {typename} label '{label}' is cited out of order (line {line_no})")
                else:
                    last_index = idx
        return errors

    ref_errors = []
    ref_errors += check_order(refs_by_group['main_figures'], groups['main_figures'], 'Main Figure')
    ref_errors += check_order(refs_by_group['supp_figures'], groups['supp_figures'], 'Supplementary Figure')
    ref_errors += check_order(refs_by_group['main_tables'], groups['main_tables'], 'Main Table')
    ref_errors += check_order(refs_by_group['supp_tables'], groups['supp_tables'], 'Supplementary Table')

    seen_labels = set(first_refs.keys())

    log_lines = []
    for i, f in enumerate(float_blocks, 1):
        status = []
        if f['label'] == 'MISSING_LABEL':
            status.append(' Missing label')
        elif label_counts[f['label']] > 1:
            status.append(' Duplicate label')
        if f['label'] not in seen_labels and f['label'] != 'MISSING_LABEL':
            status.append(' Not cited')

        log_lines.append(
            f"{i:3d}. Line {f['line']:4d} - {f['type'].capitalize()} - Label: {f['label']} - "
            f"Caption: {f['caption']} " + ("| " + ', '.join(status) if status else "")
        )

    log_lines.append("\nReference order issues:")
    if ref_errors:
        log_lines.extend(ref_errors)
    else:
        log_lines.append(" All references are in correct order.")

    log_path = Path(tex_path).with_suffix('.floatlog.txt')
    with open(log_path, 'w', encoding='utf-8') as f:
        f.write('\n'.join(log_lines))

    print(f"Log file written to: {log_path}")
    return log_path
