from pathlib import Path

import matplotlib.pyplot as plt
import numpy as np


plt.rcParams["font.sans-serif"] = ["Microsoft YaHei", "SimHei", "DejaVu Sans"]
plt.rcParams["axes.unicode_minus"] = False

output_dir = Path(__file__).resolve().parent / "figures"
output_dir.mkdir(exist_ok=True)

nx = 81
nt = 40
c = 1.0
CFL = 0.5
domain_length = 2.0

x = np.linspace(0.0, domain_length, nx)
dx = x[1] - x[0]
dt = CFL * dx / c

u = np.ones(nx)
u[(x >= 0.5) & (x <= 1.0)] = 2.0
initial_u = u.copy()

snapshots = {0: u.copy()}
record_steps = [10, 20, 40]

for n in range(1, nt + 1):
    old_u = u.copy()
    u[1:] = old_u[1:] - c * dt / dx * (old_u[1:] - old_u[:-1])
    u[0] = 1.0

    if n in record_steps:
        snapshots[n] = u.copy()

shift_distance = c * nt * dt
exact_u = np.ones(nx)
exact_u[(x >= 0.5 + shift_distance) & (x <= 1.0 + shift_distance)] = 2.0

plt.figure(figsize=(7, 4))
plt.plot(x, initial_u, linewidth=2)
plt.title("初始条件")
plt.xlabel("x")
plt.ylabel("u")
plt.ylim(0.8, 2.2)
plt.grid(True, alpha=0.3)
plt.tight_layout()
plt.savefig(output_dir / "01_initial_condition.png", dpi=160)
plt.close()

plt.figure(figsize=(7, 4))
for step, values in snapshots.items():
    plt.plot(x, values, linewidth=2, label=f"第 {step} 步")
plt.title("一维线性对流的时间演化")
plt.xlabel("x")
plt.ylabel("u")
plt.ylim(0.8, 2.2)
plt.grid(True, alpha=0.3)
plt.legend()
plt.tight_layout()
plt.savefig(output_dir / "02_time_evolution.png", dpi=160)
plt.close()

plt.figure(figsize=(7, 4))
plt.plot(x, exact_u, "k--", linewidth=2, label="理论最终波形")
plt.plot(x, u, color="tab:red", linewidth=2, label="数值最终波形")
plt.title("数值解与理论解对比")
plt.xlabel("x")
plt.ylabel("u")
plt.ylim(0.8, 2.2)
plt.grid(True, alpha=0.3)
plt.legend()
plt.tight_layout()
plt.savefig(output_dir / "03_numerical_vs_exact.png", dpi=160)
plt.close()

print("一维线性对流程序运行完成。")
print(f"网格点数 nx = {nx}")
print(f"时间步数 nt = {nt}")
print(f"空间步长 dx = {dx:.6f}")
print(f"时间步长 dt = {dt:.6f}")
print(f"CFL = {c * dt / dx:.3f}")
print(f"理论传播距离 = {shift_distance:.6f}")
print(f"初始 u 范围 = [{initial_u.min():.3f}, {initial_u.max():.3f}]")
print(f"最终 u 范围 = [{u.min():.3f}, {u.max():.3f}]")
print(f"结果图保存在：{output_dir}")
