from PIL import Image
import tkinter as tk
from tkinter import filedialog, Scrollbar, Toplevel
import matplotlib.pyplot as plt
import numpy as np
import os
import math

#计算波长为lamda的光线横向位移
#K9玻璃，15度棱镜，735mm焦距
def calc_delta_d(lamda):
    # 定义常数
    A1 = 0.567853792
    A2 = 0.702269689
    A3 = 1.08102455
    C1 = 0.00238581579
    C2 = 0.0135348332
    C3 = 107.748228
    
    # 步骤1：计算折射率 n
    lamda_sq = lamda ** 2
    term1 = (A1 * lamda_sq) / (lamda_sq - C1)
    term2 = (A2 * lamda_sq) / (lamda_sq - C2)
    term3 = (A3 * lamda_sq) / (lamda_sq - C3)
    n = math.sqrt(1 + term1 + term2 + term3)
    
    # 步骤2：计算角度 alpha（弧度）
    theta = 0.261799  # 固定入射角（15度对应的弧度）
    sin_theta = math.sin(theta)
    alpha = math.asin(n * sin_theta) - theta
    
    # 步骤3：计算并返回 delta_d
    delta_d = 735 * math.tan(alpha)
    return delta_d

def process_and_save_curve(image_path):
    # 打开并处理图像
    img = Image.open(image_path)
    gray_img = img.convert('L')
    width, height = gray_img.size
    
    # 获取中间行像素
    middle_row = height // 2
    gray_values = [gray_img.getpixel((x, middle_row)) for x in range(width)]
    
    # 灰度值处理
    subtracted = [max(val - 70, 0) for val in gray_values]
    max_val = max(subtracted) if max(subtracted) > 0 else 1
    stretched = [int(val * 200 / max_val) for val in subtracted]

    #横坐标波长标记
    x_labels = ['350','400','450','500','550','600','650','700','750','800','850','900','950','1000']

    wavelengths = range(350, 1001, 50)  # 350, 400, 450,..., 900
    x_values = []

    for wl in wavelengths:
        try:
            delta_d = calc_delta_d(wl/1000)
            x_values.append((109.53-delta_d)*310.772)
        except ValueError as e:
            # 处理可能出现的数学错误（如负数的平方根）
            print(f"计算波长 {wl}nm 时出错: {str(e)}")
            x_values.append((wl, float('nan')))
    
    
    # 创建新图形（先不显示）
    fig, ax = plt.subplots(figsize=(24, 6), dpi=100)

    # 设置横坐标a 刻度和标签
    plt.xticks(x_values, x_labels)
    
    # 绘制曲线
    ax.plot(stretched, color='black', linewidth=1)
    
    # 1. 找到第一个 `_` 的位置，并截取前面的部分
    before_underscore = image_path.split('_', 1)[0]  # "path/to/some"

    # 2. 从右往左找最后一个 `/`（即前面最近的 `/`）
    last_slash_pos = before_underscore.rfind('/')  # 4

    # 3. 截取 `/` 之后到 `_` 之前的内容
    if last_slash_pos != -1:
        star_name = before_underscore[last_slash_pos + 1:]  # "some"
    else:
        star_name = before_underscore  # 如果没有 `/`，返回整个 `_` 前的部分
        
    ax.set_title(star_name+': Processed Grayscale Values (Stretched to Max 200)', fontsize=14)
    ax.set_xlabel('Horizontal Position (nm)', fontsize=12)
    ax.set_ylabel('Value (0-200)', fontsize=12)
    ax.grid(True, alpha=0.3)
    ax.set_ylim(0, 210)
    ax.set_xlim(0, width-1)
    
    # 先显示预览
    plt.show()
    
    # 用户确认后保存
    if input("是否保存曲线图？(y/n): ").lower() == 'y':
        output_path = f"{os.path.splitext(image_path)[0]}_specchart.png"
        fig.savefig(output_path, bbox_inches='tight', dpi=123.8)
        print(f"曲线图已保存为: {output_path}")
    
    plt.close()

# 使用示例
if __name__ == "__main__":
    root = tk.Tk()
    root.withdraw() 
    file_path = filedialog.askopenfilename(
        parent=root,
        title="选择一张图片",
        filetypes=[("Image Files", "*.jpg *.jpeg *.png *.bmp *.tif *.tiff")]
    )
    if file_path:
        process_and_save_curve(file_path)  # 替换为您的图片路径
    else:
        print("no file was selected")
        
