# GESP等级：四级 | GESP Python 四级考点
"""
==================================================
  算法复杂度可视化 —— turtle画复杂度对比曲线
  面向7-10岁少儿编程教学
==================================================

  功能：
  1. 画O(1), O(log n), O(n), O(n²) 对比曲线
  2. 排序速度实测对比
  3. 复杂度分析练习

  口诀：
  "O(1)神速一步到，O(n)稳步一步步"
  "O(n²)慢如蜗牛爬，O(log n)对半找"
==================================================
"""

import turtle
import time
import math
import random

# ==================================================
# 配置区
# ==================================================

# 最大n值（x轴范围）
MAX_N = 100

# 画图速度（1=最快，10=最慢）
DRAW_SPEED = 3

# 颜色配置
COLORS = {
    "O(1)": "#FF6B6B",      # 红色
    "O(log n)": "#4DABF7",  # 蓝色
    "O(n)": "#69DB7C",      # 绿色
    "O(n²)": "#FFA94D",     # 橙色
    "axis": "#333333",      # 坐标轴
    "grid": "#DDDDDD",      # 网格
}

# ==================================================
# turtle设置
# ==================================================

screen = turtle.Screen()
screen.setup(900, 700)
screen.title("算法复杂度对比 —— 谁跑得更快？")
screen.bgcolor("white")
screen.tracer(0)

# 主画笔
pen = turtle.Turtle()
pen.speed(DRAW_SPEED)
pen.hideturtle()

# 文字画笔
text = turtle.Turtle()
text.speed(0)
text.hideturtle()
text.penup()

# ==================================================
# 绘图区域设置
# ==================================================

# 绘图区域的坐标范围
ORIGIN_X = -350      # 原点x
ORIGIN_Y = -200      # 原点y
AXIS_LENGTH_X = 600  # x轴长度
AXIS_LENGTH_Y = 400  # y轴长度
SCALE_X = AXIS_LENGTH_X / MAX_N  # x轴比例
SCALE_Y = AXIS_LENGTH_Y / (MAX_N ** 2 / 10)  # y轴比例（适配n²）


def draw_axes():
    """画坐标轴"""
    pen.penup()
    pen.goto(ORIGIN_X, ORIGIN_Y)
    pen.pendown()
    pen.pencolor(COLORS["axis"])
    pen.width(2)

    # x轴
    pen.forward(AXIS_LENGTH_X)
    pen.stamp()

    # x轴箭头
    pen.goto(ORIGIN_X + AXIS_LENGTH_X, ORIGIN_Y)
    pen.write("n (数据量)", font=("Arial", 12, "normal"))

    # y轴
    pen.penup()
    pen.goto(ORIGIN_X, ORIGIN_Y)
    pen.pendown()
    pen.left(90)
    pen.forward(AXIS_LENGTH_Y)
    pen.goto(ORIGIN_X, ORIGIN_Y + AXIS_LENGTH_Y)
    pen.write("操作次数", font=("Arial", 12, "normal"))

    # 原点标签
    pen.penup()
    pen.goto(ORIGIN_X - 20, ORIGIN_Y - 15)
    pen.write("0", font=("Arial", 10, "normal"))

    # x轴刻度
    for n in range(0, MAX_N + 1, 10):
        x = ORIGIN_X + n * SCALE_X
        pen.goto(x, ORIGIN_Y - 5)
        pen.pendown()
        pen.goto(x, ORIGIN_Y + 5)
        pen.penup()
        if n > 0:
            pen.goto(x - 10, ORIGIN_Y - 20)
            pen.write(str(n), font=("Arial", 8, "normal"))

    # y轴刻度
    max_ops = int(MAX_N ** 2 / 10)
    for ops in range(0, max_ops + 1, max_ops // 5):
        y = ORIGIN_Y + ops * SCALE_Y
        if y > ORIGIN_Y + AXIS_LENGTH_Y:
            break
        pen.goto(ORIGIN_X - 5, y)
        pen.pendown()
        pen.goto(ORIGIN_X + 5, y)
        pen.penup()
        pen.goto(ORIGIN_X - 45, y - 7)
        pen.write(str(ops), font=("Arial", 8, "normal"))

    # 恢复方向
    pen.setheading(0)

    screen.update()


def draw_grid():
    """画网格"""
    pen.pencolor(COLORS["grid"])
    pen.width(1)

    # 竖网格
    for n in range(0, MAX_N + 1, 10):
        x = ORIGIN_X + n * SCALE_X
        pen.penup()
        pen.goto(x, ORIGIN_Y)
        pen.pendown()
        pen.goto(x, ORIGIN_Y + AXIS_LENGTH_Y)

    # 横网格
    max_ops = int(MAX_N ** 2 / 10)
    for ops in range(0, max_ops + 1, max_ops // 5):
        y = ORIGIN_Y + ops * SCALE_Y
        if y > ORIGIN_Y + AXIS_LENGTH_Y:
            break
        pen.penup()
        pen.goto(ORIGIN_X, y)
        pen.pendown()
        pen.goto(ORIGIN_X + AXIS_LENGTH_X, y)

    screen.update()


def plot_curve(name, func, color, label_y=0):
    """
    画函数曲线

    参数：
        name: 曲线名称（用于图例）
        func: 函数，接收n返回操作次数
        color: 颜色
        label_y: 图例的y坐标
    """
    pen.pencolor(color)
    pen.width(3)
    pen.penup()

    first = True
    for n in range(0, MAX_N + 1):
        ops = func(n)
        x = ORIGIN_X + n * SCALE_X
        y = ORIGIN_Y + ops * SCALE_Y

        if y > ORIGIN_Y + AXIS_LENGTH_Y:
            y = ORIGIN_Y + AXIS_LENGTH_Y  # 截断

        if first:
            pen.goto(x, y)
            pen.pendown()
            first = False
        else:
            pen.goto(x, y)

    # 画图例
    if label_y != 0:
        text.goto(ORIGIN_X + AXIS_LENGTH_X - 150, ORIGIN_Y + AXIS_LENGTH_Y - label_y)
        text.pencolor(color)
        text.write(f"━ {name}", font=("Arial", 12, "bold"))

    screen.update()


# ==================================================
# 复杂度函数
# ==================================================

def constant_time(n):
    """O(1) —— 常数时间"""
    return 1


def log_time(n):
    """O(log n) —— 对数时间"""
    if n <= 1:
        return 0
    return math.log2(n)


def linear_time(n):
    """O(n) —— 线性时间"""
    return n


def quadratic_time(n):
    """O(n²) —— 平方时间"""
    return n * n


# ==================================================
# 主绘图函数
# ==================================================

def draw_complexity_chart():
    """画完整复杂度对比图"""
    pen.clear()
    text.clear()

    # 标题
    text.goto(0, 300)
    text.pencolor("black")
    text.write("📊 算法复杂度对比图", align="center", font=("Arial", 20, "bold"))
    text.goto(0, 270)
    text.write("横轴：数据量(n)   |   纵轴：操作次数",
               align="center", font=("Arial", 12, "normal"))
    text.goto(0, 245)
    text.write("💡 曲线越平缓 → 算法越快！",
               align="center", font=("Arial", 12, "normal"))

    # 画网格
    draw_grid()

    # 画坐标轴
    draw_axes()

    # 画四条曲线
    print("画 O(1) 曲线...")
    plot_curve("O(1)", constant_time, COLORS["O(1)"], 30)

    print("画 O(log n) 曲线...")
    plot_curve("O(log n)", log_time, COLORS["O(log n)"], 55)

    print("画 O(n) 曲线...")
    plot_curve("O(n)", linear_time, COLORS["O(n)"], 80)

    print("画 O(n²) 曲线...")
    plot_curve("O(n²)", quadratic_time, COLORS["O(n²)"], 105)

    # 图例说明
    text.goto(ORIGIN_X + AXIS_LENGTH_X - 150, ORIGIN_Y + AXIS_LENGTH_Y - 20)
    text.pencolor("black")
    text.write("📌 图例：", font=("Arial", 12, "bold"))

    # 显示速度排名
    text.goto(0, -310)
    text.pencolor("black")
    text.write("⚡ 速度排名：O(1) > O(log n) > O(n) > O(n²)",
               align="center", font=("Arial", 14, "bold"))

    # 显示数据点示例
    text.goto(0, -335)
    text.write(f"n={MAX_N}时：O(1)=1次  O(log n)≈{math.log2(MAX_N):.0f}次  O(n)={MAX_N}次  O(n²)={MAX_N*MAX_N}次",
               align="center", font=("Arial", 11, "normal"))

    screen.update()
    print("✅ 复杂度对比图绘制完成！")


# ==================================================
# 排序速度实测对比
# ==================================================

def bubble_sort(arr):
    """冒泡排序"""
    n = len(arr)
    for i in range(n - 1):
        for j in range(n - 1 - i):
            if arr[j] > arr[j + 1]:
                arr[j], arr[j + 1] = arr[j + 1], arr[j]
    return arr


def selection_sort(arr):
    """选择排序"""
    n = len(arr)
    for i in range(n - 1):
        min_idx = i
        for j in range(i + 1, n):
            if arr[j] < arr[min_idx]:
                min_idx = j
        if min_idx != i:
            arr[i], arr[min_idx] = arr[min_idx], arr[i]
    return arr


def insertion_sort(arr):
    """插入排序"""
    n = len(arr)
    for i in range(1, n):
        key = arr[i]
        j = i - 1
        while j >= 0 and arr[j] > key:
            arr[j + 1] = arr[j]
            j -= 1
        arr[j + 1] = key
    return arr


def run_sort_comparison():
    """运行排序速度对比实验"""
    print("\n" + "=" * 50)
    print("🔬 排序速度实测对比")
    print("=" * 50)

    sizes = [10, 50, 100, 200, 500]

    print(f"\n{'n':>5} | {'冒泡排序(秒)':>15} | {'选择排序(秒)':>15} | {'插入排序(秒)':>15}")
    print("-" * 60)

    for size in sizes:
        data = [random.randint(1, 10000) for _ in range(size)]

        # 冒泡排序
        d1 = data.copy()
        start = time.time()
        bubble_sort(d1)
        t1 = time.time() - start

        # 选择排序
        d2 = data.copy()
        start = time.time()
        selection_sort(d2)
        t2 = time.time() - start

        # 插入排序
        d3 = data.copy()
        start = time.time()
        insertion_sort(d3)
        t3 = time.time() - start

        print(f"{size:>5} | {t1:>15.6f} | {t2:>15.6f} | {t3:>15.6f}")

    print("\n📊 结论：三种排序都是 O(n²) 时间复杂度")
    print("当 n 越来越大时，速度都会明显变慢！")


def run_best_worst_comparison():
    """对比最好和最坏情况"""
    print("\n" + "=" * 50)
    print("🔬 插入排序：最好 vs 最坏情况")
    print("=" * 50)

    size = 200

    # 最好情况：已经排好
    best_data = list(range(size))
    start = time.time()
    insertion_sort(best_data.copy())
    best_time = time.time() - start

    # 最坏情况：完全逆序
    worst_data = list(range(size, 0, -1))
    start = time.time()
    insertion_sort(worst_data.copy())
    worst_time = time.time() - start

    print(f"\n数据量：{size}")
    print(f"最好情况（已排好）：{best_time:.6f}秒  → O(n)")
    print(f"最坏情况（逆序） ：{worst_time:.6f}秒  → O(n²)")
    print(f"相差 {worst_time / best_time:.0f} 倍！")


# ==================================================
# 复杂度分析器
# ==================================================

def analyze_complexity(code_name, n):
    """
    分析代码的时间复杂度

    参数：
        code_name: 代码名称
        n: 数据量
    返回：
        估计的复杂度等级和操作次数
    """
    analyzers = {
        "sum_list": {
            "name": "列表求和",
            "complexity": "O(n)",
            "desc": "循环n次，每次做1个加法",
            "ops": n
        },
        "print_pairs": {
            "name": "打印所有数对",
            "complexity": "O(n²)",
            "desc": "两层循环，外层n次，内层n次",
            "ops": n * n
        },
        "get_mid": {
            "name": "取中间元素",
            "complexity": "O(1)",
            "desc": "直接索引，一步完成",
            "ops": 1
        },
        "binary_search": {
            "name": "二分查找",
            "complexity": "O(log n)",
            "desc": "每次砍掉一半数据",
            "ops": int(math.log2(n)) if n > 0 else 0
        },
        "bubble_sort_op": {
            "name": "冒泡排序",
            "complexity": "O(n²)",
            "desc": "两层循环，n(n-1)/2次比较",
            "ops": n * (n - 1) // 2
        }
    }

    if code_name in analyzers:
        info = analyzers[code_name]
        print(f"\n📝 代码：{info['name']}")
        print(f"   复杂度：{info['complexity']}")
        print(f"   说明：{info['desc']}")
        print(f"   n={n}时：约{info['ops']}次操作")
        return info
    else:
        print(f"未知代码：{code_name}")
        return None


# ==================================================
# 交互控制
# ==================================================

def on_click(x, y):
    """点击屏幕重新绘图"""
    print("\n🔄 重新绘制复杂度对比图...")
    text.clear()
    screen.update()
    draw_complexity_chart()


# ==================================================
# 主程序
# ==================================================

if __name__ == "__main__":
    print("🌟 欢迎来到算法复杂度世界！")
    print("=" * 50)
    print()
    print("📌 这里可以：")
    print("  1. 画复杂度对比曲线（turtle图形）")
    print("  2. 实测排序速度对比")
    print("  3. 分析不同代码的复杂度")
    print()

    # 画复杂度图
    draw_complexity_chart()

    # 排序速度对比
    run_sort_comparison()

    # 最好最坏对比
    run_best_worst_comparison()

    # 复杂度分析
    print("\n" + "=" * 50)
    print("📝 复杂度分析示例")
    print("=" * 50)
    for code_name in ["sum_list", "print_pairs", "get_mid", "binary_search", "bubble_sort_op"]:
        analyze_complexity(code_name, 100)

    # 更多示例
    print("\n" + "=" * 50)
    print("📊 n=1000时各复杂度操作次数对比")
    print("=" * 50)
    n = 1000
    print(f"  O(1)     = {constant_time(n)} 次")
    print(f"  O(log n) = {log_time(n):.0f} 次")
    print(f"  O(n)     = {linear_time(n)} 次")
    print(f"  O(n²)    = {quadratic_time(n)} 次")
    print()
    print(f"  O(n²) 是 O(n) 的 {quadratic_time(n) // linear_time(n)} 倍！")
    print(f"  O(n²) 是 O(log n) 的 {quadratic_time(n) // int(log_time(n))} 倍！")

    # 点击屏幕
    screen.listen()
    screen.onclick(on_click)

    print("\n💡 提示：点击屏幕可以重新绘制对比图！")
    print("💡 按 ESC 或关闭窗口退出程序")

    screen.mainloop()
