import pandas as pd
import matplotlib.pyplot as plt
import numpy as np
from pathlib import Path
import warnings

# Suppress matplotlib font warnings
warnings.filterwarnings('ignore', category=UserWarning, module='matplotlib')

class GradeVisualizer:
    def __init__(self):
        """初始化成绩可视化器"""
        # 使用默认字体设置，避免中文字体问题
        plt.rcParams['font.family'] = 'sans-serif'
        plt.style.use('default')
        
        # 定义列名显示映射（图表用英文，终端输出用中文）
        self.column_display_names = {
            '平时作业1': 'sum 1',
            '平时作业2': 'sum 2',
            'work1': 'Work 1',
            'work2': 'Work 2',
            'work3': 'Work 3',
            'work4': 'Work 4',
            'work5': 'Work 5'
        }
    
    def load_data(self, csv_file='data/数据结构成绩记录表_完整.csv'):
        """加载成绩数据"""
        try:
            self.df = pd.read_csv(csv_file)
            print(f"✓ 数据加载成功：{len(self.df)} 条记录")
            return True
        except Exception as e:
            print(f"✗ 数据加载失败：{e}")
            return False
    
    def plot_group_score_distribution(self, score_column, save_path='assets/'):
        """绘制各分组的成绩分布图，包含统计信息"""
        Path(save_path).mkdir(exist_ok=True)
        
        # 获取列的显示名称（英文，避免字体问题）
        display_name = self.column_display_names.get(score_column, score_column)
        
        # 获取所有分组
        groups = self.df['分组'].unique()
        
        # 创建分组名称映射（避免中文字符）
        group_mapping = {}
        for idx, group in enumerate(groups):
            group_mapping[group] = f'Group {idx + 1}'
        
        # 创建子图
        fig, axes = plt.subplots(2, 2, figsize=(15, 10))
        fig.suptitle(f'{display_name} Score Distribution by Group', fontsize=16, fontweight='bold')
        
        # 展平axes数组便于迭代
        axes = axes.flatten()
        
        for i, group in enumerate(groups):
            if i >= 4:  # 最多支持4个分组
                break
                
            # 获取分组数据
            group_data = self.df[self.df['分组'] == group]
            scores = pd.to_numeric(group_data[score_column], errors='coerce').dropna()
            
            # 使用英文分组名称
            group_english = group_mapping[group]
            
            if len(scores) == 0:
                axes[i].text(0.5, 0.5, f'{group_english}\nNo Valid Data', 
                           ha='center', va='center', transform=axes[i].transAxes)
                axes[i].set_title(f'{group_english}')
                continue
            
            # 计算统计数据
            mean_score = scores.mean()
            max_score = scores.max()
            min_score = scores.min()
            std_score = scores.std()
            
            # 绘制直方图
            axes[i].hist(scores, bins=min(15, len(scores)//2+1), alpha=0.7, 
                        color='skyblue', edgecolor='black')
            
            # 添加统计线（使用英文标签）
            axes[i].axvline(mean_score, color='red', linestyle='--', linewidth=2, 
                          label=f'Mean: {mean_score:.1f}')
            axes[i].axvline(max_score, color='green', linestyle=':', linewidth=2, 
                          label=f'Max: {max_score:.1f}')
            axes[i].axvline(min_score, color='orange', linestyle=':', linewidth=2, 
                          label=f'Min: {min_score:.1f}')
            
            # 设置标题和标签（使用英文）
            axes[i].set_title(f'{group_english} (Std={std_score:.1f})')
            axes[i].set_xlabel('Score')
            axes[i].set_ylabel('Count')
            axes[i].legend(fontsize=8)
            axes[i].grid(True, alpha=0.3)
        
        # 隐藏多余的子图
        for j in range(len(groups), 4):
            axes[j].set_visible(False)
        
        plt.tight_layout()
        # 使用安全的文件名
        safe_filename = display_name.replace(' ', '_').lower()
        filename = f'{save_path}{safe_filename}_group_distribution.png'
        plt.savefig(filename, dpi=300, bbox_inches='tight')
        plt.close()
        print(f"✓ 已生成 {score_column} 分组分布图")
    
    def generate_all_charts(self, save_path='assets/'):
        """生成所有必要的图表"""
        if not hasattr(self, 'df'):
            print("请先加载数据")
            return False
        
        print("正在生成分组成绩分布图...")
        
        # 确保目录存在
        Path(save_path).mkdir(exist_ok=True)
        
        try:
            # 为主要成绩列生成分组分布图
            score_columns = ['平时作业1', '平时作业2']
            
            for col in score_columns:
                if col in self.df.columns:
                    self.plot_group_score_distribution(col, save_path)
            
            print(f"✓ 图表已保存到 {save_path} 目录")
            return True
            
        except Exception as e:
            print(f"✗ 生成图表时出错：{e}")
            return False
    
    def create_summary_report(self, save_path='assets/'):
        """创建简单的统计汇总"""
        Path(save_path).mkdir(exist_ok=True)
        
        try:
            print("✓ 分组成绩统计汇总：")
            
            # 为每个成绩列生成统计信息
            score_columns = ['平时作业1', '平时作业2']
            
            for col in score_columns:
                if col in self.df.columns:
                    display_name = self.column_display_names.get(col, col)
                    print(f"\n{display_name}:")
                    groups = self.df['分组'].unique()
                    
                    for group in groups:
                        group_data = self.df[self.df['分组'] == group]
                        scores = pd.to_numeric(group_data[col], errors='coerce').dropna()
                        
                        if len(scores) > 0:
                            mean_score = scores.mean()
                            max_score = scores.max()
                            min_score = scores.min()
                            std_score = scores.std()
                            
                            print(f"  {group}组: 均值 {mean_score:.1f} 最高 {max_score:.1f} 最低 {min_score:.1f} 标准差 {std_score:.1f}")
                        else:
                            print(f"  {group}组: 无有效数据")
            
            return True
            
        except Exception as e:
            print(f"✗ 生成统计汇总时出错：{e}")
            return False

# 使用示例
if __name__ == "__main__":
    visualizer = GradeVisualizer()
    
    # 加载数据
    if visualizer.load_data():
        # 生成分组分布图表
        visualizer.generate_all_charts()
        # 显示统计汇总
        visualizer.create_summary_report()
        print("✓ 可视化完成！")
