os/arch/riscv/kernel/
task.rs

1//! RISC-V 架构的任务管理相关功能
2use core::mem::size_of;
3use core::ptr;
4
5use alloc::vec::Vec;
6use riscv::register::sstatus;
7
8use crate::arch::constant::STACK_ALIGN_MASK;
9
10/// 为新任务设置用户栈布局,包含命令行参数和环境变量
11/// 返回新的栈指针位置,以及 argc, argv, envp 的地址
12pub fn setup_stack_layout(
13    sp: usize,
14    argv: &[&str],
15    envp: &[&str],
16    phdr_addr: usize,
17    phnum: usize,
18    phent: usize,
19    entry_point: usize,
20) -> (usize, usize, usize, usize) {
21    let mut sp = sp;
22    let mut arg_ptrs: Vec<usize> = Vec::with_capacity(argv.len());
23    let mut env_ptrs: Vec<usize> = Vec::with_capacity(envp.len());
24    unsafe {
25        sstatus::set_sum();
26    }
27
28    for &env in envp.iter().rev() {
29        let bytes = env.as_bytes();
30        sp -= bytes.len() + 1; // 预留 NUL
31        unsafe {
32            ptr::copy_nonoverlapping(bytes.as_ptr(), sp as *mut u8, bytes.len());
33            (sp as *mut u8).add(bytes.len()).write(0); // NUL 终止符
34        }
35        env_ptrs.push(sp); // 存储字符串的地址
36    }
37
38    // 命令行参数 (argv)
39    for &arg in argv.iter().rev() {
40        let bytes = arg.as_bytes();
41        sp -= bytes.len() + 1; // 预留 NUL
42        unsafe {
43            ptr::copy_nonoverlapping(bytes.as_ptr(), sp as *mut u8, bytes.len());
44            (sp as *mut u8).add(bytes.len()).write(0); // NUL 终止符
45        }
46        arg_ptrs.push(sp); // 存储字符串的地址
47    }
48
49    // --- 对齐到字大小 (确保指针数组从对齐的地址开始) ---
50    sp &= !(size_of::<usize>() - 1);
51
52    // --- 构建 argc, argv, envp 数组 (ABI 标准布局: [argc] -> [argv] -> [NULL] -> [envp] -> [NULL]) ---
53    // 注意:栈向下增长,所以压栈顺序是从 envp NULL 往回压到 argc
54
55    // 0. 写入 auxv (Auxiliary Vector)
56    // 必须位于 envp NULL 之后(高地址),但在 envp 数组之前。
57    // 常见的 auxv 条目:AT_PAGESZ(6), AT_NULL(0), AT_RANDOM(25)
58    let random_bytes = [
59        0x89, 0xab, 0xcd, 0xef, 0x01, 0x23, 0x45, 0x67, 0x89, 0xab, 0xcd, 0xef, 0x01, 0x23, 0x45,
60        0x67,
61    ];
62    let random_ptr = sp - 16;
63    unsafe { ptr::copy_nonoverlapping(random_bytes.as_ptr(), random_ptr as *mut u8, 16) };
64    sp = random_ptr;
65
66    // 2. Platform string "riscv64\0" (8 bytes)
67    let platform = "riscv64\0";
68    let platform_len = platform.len();
69    sp -= platform_len;
70    unsafe { ptr::copy_nonoverlapping(platform.as_ptr(), sp as *mut u8, platform_len) };
71    let platform_ptr = sp;
72
73    // 3. Align sp to 16 bytes (auxv requirement)
74    sp &= !(size_of::<usize>() * 2 - 1); // Align to 16 bytes
75
76    // 4. AT_EXECFN (use argv[0] if available)
77    let execfn = if !arg_ptrs.is_empty() { arg_ptrs[0] } else { 0 };
78
79    let auxv = [
80        (3, phdr_addr),     // AT_PHDR
81        (4, phent),         // AT_PHENT
82        (5, phnum),         // AT_PHNUM
83        (6, 4096),          // AT_PAGESZ
84        (7, 0),             // AT_BASE
85        (8, 0),             // AT_FLAGS
86        (9, entry_point),   // AT_ENTRY
87        (11, 0),            // AT_UID
88        (12, 0),            // AT_EUID
89        (13, 0),            // AT_GID
90        (14, 0),            // AT_EGID
91        (15, platform_ptr), // AT_PLATFORM
92        (16, 0),            // AT_HWCAP
93        (17, 100),          // AT_CLKTCK
94        (23, 0),            // AT_SECURE
95        (25, random_ptr),   // AT_RANDOM
96        (31, execfn),       // AT_EXECFN
97        (0, 0),             // AT_NULL
98    ];
99
100    // Debug print auxv
101    for (i, (k, v)) in auxv.iter().enumerate() {
102        crate::pr_debug!("auxv[{}]: type={}, val={:#x}", i, k, v);
103    }
104    crate::pr_debug!(
105        "setup_stack_layout: sp={:#x}, random_ptr={:#x}, phdr_addr={:#x}, entry={:#x}",
106        sp,
107        random_ptr,
108        phdr_addr,
109        entry_point
110    );
111
112    // Calculate total size of the pointer block to ensure final sp is 16-byte aligned
113    // Block includes: auxv[], padding, envp NULL, envp[], argv NULL, argv[], argc
114    let total_size = auxv.len() * 2 * size_of::<usize>()
115        + size_of::<usize>() // envp NULL
116        + env_ptrs.len() * size_of::<usize>()
117        + size_of::<usize>() // argv NULL
118        + arg_ptrs.len() * size_of::<usize>()
119        + size_of::<usize>(); // argc
120
121    // Align the final stack pointer
122    let sp_final = (sp - total_size) & !STACK_ALIGN_MASK;
123    sp = sp_final + total_size;
124
125    for (type_, val) in auxv.iter().rev() {
126        sp -= size_of::<usize>();
127        unsafe { ptr::write(sp as *mut usize, *val) };
128        sp -= size_of::<usize>();
129        unsafe { ptr::write(sp as *mut usize, *type_) };
130    }
131
132    // 1. 写入 envp NULL 终止符
133    sp -= size_of::<usize>();
134    unsafe {
135        ptr::write(sp as *mut usize, 0);
136    }
137
138    // 2. 写入 envp 指针数组(逆序写入,使 envp[0] 处于最低地址)
139    // env_ptrs 已经是逆序 (envp[n-1] ... envp[0])
140    for &p in env_ptrs.iter() {
141        sp -= size_of::<usize>();
142        unsafe {
143            ptr::write(sp as *mut usize, p);
144        }
145    }
146    let envp_vec_ptr = sp; // envp 数组的起始地址 (envp[0] 的地址)
147
148    // 3. 写入 argv NULL 终止符
149    sp -= size_of::<usize>();
150    unsafe {
151        ptr::write(sp as *mut usize, 0);
152    }
153
154    // 4. 写入 argv 指针数组(逆序写入,使 argv[0] 处于最低地址)
155    // arg_ptrs 已经是逆序 (argv[n-1] ... argv[0])
156    for &p in arg_ptrs.iter() {
157        sp -= size_of::<usize>();
158        unsafe {
159            ptr::write(sp as *mut usize, p);
160        }
161    }
162    let argv_vec_ptr = sp; // argv 数组的起始地址 (argv[0] 的地址)
163
164    // 5. 写入 argc
165    let argc = argv.len();
166    sp -= size_of::<usize>();
167    unsafe {
168        ptr::write(sp as *mut usize, argc);
169    }
170
171    // 拷贝完成,恢复 SUM
172    unsafe {
173        sstatus::clear_sum();
174    }
175
176    // 6. 最终 sp 应该已经是 16 字节对齐的
177    // sp &= !STACK_ALIGN_MASK;
178    (sp, argc, argv_vec_ptr, envp_vec_ptr)
179}