简单C函数重构实战
一、重构原则
原则1: 先理解,后改写 — 不理解C原意绝不写Rust
原则2: 保持语义等价 — Rust版本必须产生相同结果
原则3: 利用类型系统 — 将运行时约束转化为编译期约束
原则4: 测试驱动 — 先写测试,C和Rust共享测试用例
原则5: 逐步安全化 — 无法一次安全时用unsafe封装再渐进重构
二、函数1: strlen
C版本:
size_t strlen(const char* s) {
const char* p = s;
while (*p != '\0') p++;
return (size_t)(p - s);
}不安全点:s可为NULL → 段错误;无终止符 → 无限循环越界。
Rust重构:
fn strlen(s: &str) -> usize {
s.bytes().take_while(|&b| b != b'\0').count()
}| 方面 | C strlen | Rust strlen |
|---|---|---|
| 参数 | const char* (可为NULL) | &str (编译期非空) |
| 边界 | 依赖\0 | 切片自带长度 |
| 无终止符 | 未定义行为 | 最多读到切片末尾 |
Luogu练习:P1035 级数求和 — 练习字符串长度计算的边界条件。
三、函数2: strdup
C版本:
char* strdup(const char* s) {
if (s == NULL) return NULL;
size_t len = strlen(s);
char* dup = (char*)malloc(len + 1);
if (dup == NULL) return NULL;
memcpy(dup, s, len + 1);
return dup;
}问题场景:忘记free → 泄漏;free后继续使用 → UAF;free两次 → double-free。
Rust重构:
fn strdup(s: &str) -> String { s.to_string() }无需手动释放,编译器跟踪所有权。C版本的5种内存错误在Rust中被编译期消除。
四、函数3: create_array
C版本:
int* create_array(int n) {
if (n <= 0) return NULL;
int* arr = (int*)malloc((size_t)n * sizeof(int));
if (arr == NULL) return NULL;
for (int i = 0; i < n; i++) arr[i] = 0;
return arr;
}不安全点:n可能为负;n*sizeof(int)溢出;malloc失败返回NULL被调用者忽略。
Rust重构:
fn create_array(n: usize) -> Vec<i32> { vec![0i32; n] }
fn create_array_with_value(n: usize, val: i32) -> Vec<i32> { vec![val; n] }| 特性 | C版本 | Rust版本 |
|---|---|---|
| 创建 | malloc + 手动循环 | vec![]宏一行 |
| 大小跟踪 | 无内置机制 | .len()方法 |
| 边界检查 | 无 | .get()返回 Option |
| 释放 | 手动free | 自动Drop |
五、函数4: sort_array
C版本:
void sort_array(int* arr, int n) {
if (arr == NULL || n <= 1) return;
for (int i = 0; i < n - 1; i++) {
int swapped = 0;
for (int j = 0; j < n - 1 - i; j++) {
if (arr[j] > arr[j + 1]) {
int temp = arr[j];
arr[j] = arr[j + 1];
arr[j + 1] = temp;
swapped = 1;
}
}
if (!swapped) break;
}
}问题:arr和n无关联验证;只支持int;O(n)冒泡效率低。
Rust重构:
fn sort_array_generic<T: Ord>(arr: &mut [T]) {
let n = arr.len();
if n <= 1 { return; }
for i in 0..n-1 {
let mut swapped = false;
for j in 0..n-1-i {
if arr[j] > arr[j+1] { arr.swap(j, j+1); swapped = true; }
}
if !swapped { break; }
}
}
// 生产代码直接用 arr.sort() (timsort, O(n log n))切片自带长度,arr和n永不分离;泛型支持所有Ord类型;标准库sort更快且稳定。
Luogu练习:P1177 【模板】排序 — 用Rust排序解决经典排序问题。
六、函数5: read_file
C版本:
char* read_file(const char* path) {
if (path == NULL) return NULL;
FILE* file = fopen(path, "rb");
if (!file) return NULL;
fseek(file, 0, SEEK_END);
long size = ftell(file); rewind(file);
char* buffer = malloc((size_t)size + 1);
if (!buffer) { fclose(file); return NULL; }
size_t n = fread(buffer, 1, size, file);
if (ferror(file)) { free(buffer); fclose(file); return NULL; }
buffer[n] = '\0'; fclose(file);
return buffer;
}C版本错误路径:所有错误返回NULL(丢失原因);ftell对管道不可靠;TOCTOU问题。
Rust重构:
fn read_file(path: &str) -> Result<String, io::Error> {
fs::read_to_string(path)
}
// 流式读取大文件:
fn read_file_streaming(path: &str) -> Result<String, io::Error> {
let file = File::open(path)?;
let reader = BufReader::new(file);
let mut result = String::new();
for line in reader.lines() { result.push_str(&line?); result.push('\n'); }
Ok(result)
}Rust的io::Error携带精确错误信息(NotFound、PermissionDenied等),?运算符一行完成错误传播。
七、函数6: split_string — 字符串分割
C版本(40行):
char** split_string(const char* str, char delimiter, int* count) {
if (str == NULL || count == NULL) return NULL;
int delim_count = 0;
for (const char* p = str; *p; p++) if (*p == delimiter) delim_count++;
char** result = (char**)malloc((delim_count + 2) * sizeof(char*));
int idx = 0; const char* start = str;
for (const char* p = str; ; p++) {
if (*p == delimiter || *p == '\0') {
size_t len = (size_t)(p - start);
result[idx] = (char*)malloc(len + 1);
memcpy(result[idx], start, len);
result[idx][len] = '\0'; idx++;
if (*p == '\0') break;
start = p + 1;
}
}
result[idx] = NULL;
*count = idx;
return result;
}问题:delim_count可能溢出;嵌套malloc极易泄漏;调用者必须同时释放数组和每个子串。
Rust重构(3行):
fn split_string(s: &str, delimiter: char) -> Vec<String> {
s.split(delimiter).map(|part| part.to_string()).collect()
}C版本40行代码的复杂手动内存管理,Rust 3行完成且完全内存安全。
八、filter_positive — 迭代器实现
C版本两遍遍历(统计+复制):
int* filter_positive(const int* arr, int len, int* out_len) {
int count = 0;
for (int i = 0; i < len; i++) if (arr[i] > 0) count++;
int* result = (int*)malloc(count * sizeof(int));
int idx = 0;
for (int i = 0; i < len; i++) if (arr[i] > 0) result[idx++] = arr[i];
*out_len = count; return result;
}Rust重构(1行迭代器):
fn filter_positive(arr: &[i32]) -> Vec<i32> {
arr.iter().copied().filter(|&x| x > 0).collect()
}安全改进:无需NULL检查;无需out_len输出参数;Vec自动管理内存;迭代器惰性求值一次完成过滤和收集;空数组返回空Vec而非NULL。
Luogu练习:P1428 小鱼比可爱 — 用迭代器方法解决数组过滤问题。
九、综合重构检查清单
□ 参数是否可为NULL?→ Option<T> 或禁止
□ 返回值是否可为NULL?→ Option<T> 或 Result<T,E>
□ 是否有输出参数?→ 用返回值或 &mut
□ 谁拥有分配的内存?→ 明确所有权,利用Drop
□ 是否有数组+长度分离?→ 切片 &[T] 或 &mut [T]
□ 错误信息是否被整型错误码吞没?→ Result携带详细信息
□ 是否有goto清理?→ RAII + Drop + ?
□ 是否有隐式类型转换?→ From/Into trait显式转换
□ 是否有多线程共享?→ Arc/Mutex 或消息传递
□ 文档是否描述不变量?→ 用类型系统编码不变量