简单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 strlenRust 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 或消息传递
□ 文档是否描述不变量?→ 用类型系统编码不变量

相关链接