最小化的上下文切换

一、核心思想

上下文切换 = 换掉 CPU 的「现场」,三步:

  1. 把当前 CPU 现场存下来(存到 old);
  2. 把另一个任务的现场装回 CPU(从 new 读);
  3. 把 PC 跳到新任务该执行的地方(ret)。

「上下文」就是任务执行到一半时,CPU 上所有寄存器的值。

二、寄存器 vs 内存(关键认知)

寄存器 内存
位置 CPU 芯片内部 CPU 芯片外部
数量 极少(RISC-V 通用寄存器 32 个) 海量
速度 极快 慢
访问方式 靠名字(sp、ra…) 靠地址
有没有地址 没有地址 有地址

核心结论:寄存器没有内存地址,只能按名字访问。 要保存寄存器的值,必须把「值」拷贝到内存的某个地址。

三、为什么要加 8 字节偏移

偏移  0  → sp
偏移  8  → ra
偏移 16  → s0
偏移 24  → s1
...
偏移104  → s11

「存」和「取」靠同一套顺序 + 偏移约定来对应,这就是 #[repr(C)] 结构体字段顺序必须和汇编偏移量一致的原因。

四、为什么只保存 sp、ra、s0~s11

RISC-V ABI 把寄存器分两类:

类型 寄存器 谁负责恢复
被调用者保存(callee-saved) sp、ra、s0~s11 被调函数负责
调用者保存(caller-saved) a0~a7、t0~t6 等 调用者负责

switch_context 在编译器眼里是一次普通函数调用:

所以 TaskContext 只存这 14 个寄存器。

五、a0/a1 哪来的

RISC-V 调用约定:第 1 个参数在 a0,第 2 个参数在 a1。

所以 switch_context_naked(old, new) 调用时:

extern "C" 保证 Rust 遵循这套约定。

六、汇编逐条讲解

1. 保存阶段(sd = store doubleword)

sd sp,  0(a0)     # 把当前 sp 存到 old.sp
sd ra,  8(a0)     # 把当前 ra 存到 old.ra
sd s0,  16(a0)    # 存 s0
...
sd s11, 104(a0)   # 存 s11

把「当前 CPU 现场」完整记录到 old 指向的内存。

2. 加载阶段(ld = load doubleword)

ld sp,  0(a1)     # new.sp → 装入 CPU 的 sp
ld ra,  8(a1)     # new.ra → 装入 CPU 的 ra
ld s0,  16(a1)    # 装入 s0
...
ld s11, 104(a1)   # 装入 s11

把新任务现场搬到 CPU。从这开始 sp 换成了新任务的栈,ra 换成了新任务的返回地址。

3. 清零 a0/a1

mv a0, zero
mv a1, zero

目的:防止指针泄漏。 切换完成后 a0/a1 还残留 old/new 两个指针,这是调度器的内部信息,不应泄漏给新任务。必须在 ret 之前清掉。

4. ret

ret          # ≡ jalr zero, 0(ra)

跳转到 ra 里存的地址。 这是上下文切换的关键一跳:

七、ra(寄存器)vs ret(指令)

名字 是什么 作用
ra 寄存器 存「返回地址」这个数据
ret 指令 执行「跳到 ra」这个动作

ret 不是寄存器,是一条跳转指令,它读取 ra 的值改变 PC(程序计数器),让 CPU 从新位置执行。

八、为什么必须 #[unsafe(naked)]

普通函数编译器会自动加 prologue/epilogue(保存/恢复 sp、ra 等),会干扰我们手动保存恢复寄存器的逻辑。

#[unsafe(naked)] 告诉编译器:「函数体只有我手写的汇编,不要加任何前置/后置代码」。

九、两层函数结构

pub unsafe fn switch_context(old: &mut TaskContext, new: &TaskContext) {
    unsafe { switch_context_naked(old, new); }
}

#[unsafe(naked)]
unsafe extern "C" fn switch_context_naked(_old: &mut TaskContext, _new: &TaskContext) {
    unsafe { core::arch::naked_asm!(...) }
}

十、完整心智模型(用测试场景串起来)

  1. main 执行到 switch_context(&mut main_ctx, &task_ctx);
  2. 汇编把 main 现场(含返回地址)存入 main_ctx;
  3. 汇编把 task_ctx 现场装入 CPU(sp=协程栈,ra=入口);
  4. ret 跳进 cooperative_task,设 COUNTER=99;
  5. 协程再调 switch_context,把协程现场存入 task_ctx,装回 main 现场;
  6. ret 跳回 main,继续执行,断言 COUNTER==99。

一句话总结

保存「我现在的现场」→ 加载「别人的现场」→ ret 跳到「别人上次停下的地方」。 只处理被调用者保存寄存器(其余编译器已处理),靠「顺序 + 8 字节偏移」的约定让存和取对得上。

# 基础04.01内联汇编最小化的上下文切换代码片段
pub unsafe fn switch_context(old: &mut TaskContext, new: &TaskContext) {

    unsafe {

        switch_context_naked(old, new);

    }

}

  
  

/// 真正的上下文切换:裸汇编函数,不能有编译器生成的 prologue/epilogue。

#[unsafe(naked)]

unsafe extern "C" fn switch_context_naked(_old: &mut TaskContext, _new: &TaskContext) {

    unsafe {

        core::arch::naked_asm!(

            // 保存当前被调用者保存寄存器到 old(a0 指向的地址)

            "sd sp,  0(a0)",

            "sd ra,  8(a0)",

            "sd s0,  16(a0)",

            "sd s1,  24(a0)",

            "sd s2,  32(a0)",

            "sd s3,  40(a0)",

            "sd s4,  48(a0)",

            "sd s5,  56(a0)",

            "sd s6,  64(a0)",

            "sd s7,  72(a0)",

            "sd s8,  80(a0)",

            "sd s9,  88(a0)",

            "sd s10, 96(a0)",

            "sd s11, 104(a0)",

  

            // 从 new(a1 指向的地址)加载被调用者保存寄存器

            "ld sp,  0(a1)",

            "ld ra,  8(a1)",

            "ld s0,  16(a1)",

            "ld s1,  24(a1)",

            "ld s2,  32(a1)",

            "ld s3,  40(a1)",

            "ld s4,  48(a1)",

            "ld s5,  56(a1)",

            "ld s6,  64(a1)",

            "ld s7,  72(a1)",

            "ld s8,  80(a1)",

            "ld s9,  88(a1)",

            "ld s10, 96(a1)",

            "ld s11, 104(a1)",

  

            // 清零 a0/a1,避免把指针泄漏进新上下文

            "mv a0, zero",

            "mv a1, zero",

  

            // ret = jalr zero, 0(ra),跳转到 new.ra

            "ret",

        );

    }

}