AI 驱动的 Rust 代码生成:从接口描述自动生成样板代码的边界探索
AI 驱动的 Rust 代码生成:从接口描述自动生成样板代码的边界探索
一、CRUD CRUD 还是 CRUD
大家好,我是一铭。我相信每个后端程序员都有这样的经历:一个新项目启动,第一周的工作就是——建表、写 model、写 repository、写 service、写 controller、写测试。七八个文件,几百行代码,全是一个模子刻出来的。
Rust 更"痛苦"——因为它的类型系统严格,一个 model 定义要配上一个 FromRow 实现、一个 NewXxx 结构体、一个 UpdateXxx 结构体、还要写一堆的 impl、Serialize、Deserialize……手工写一遍,无聊又容易出错。
那能不能让 AI 来干这件事?给它一个 OpenAPI 或者 protobuf 的接口描述文件,让它自动生成 Rust 样板代码?
这篇文章,我来探索这个方向的可能性。
二、为什么不用宏(macro_rules / proc_macro)
你可能会问:Rust 不是有过程宏(proc_macro)吗?用 #[derive] 自动生成不好吗?
好问题。宏确实能解决一部分问题,但也有局限:
- 宏只能处理 Rust 内部结构,无法理解 OpenAPI、SQL schema、protobuf 这些外部描述文件。
-
宏的调试体验很差,
cargo expand才能看到生成的代码,编译器报错位置不直观。 - 宏无法生成"上下文感知"的代码——比如根据字段名推断业务逻辑、根据关联关系生成 JOIN 查询。
- AI 生成更灵活,可以处理非标准模式(比如"这个接口需要加缓存"、"这个字段需要脱敏")。
AI 代码生成不是替代宏,而是宏的补充——宏处理规则明确的重复代码,AI 处理需要语义理解的复杂场景。
三、方案一:基于 OpenAPI 生成 axum handler
3.1 输入:OpenAPI 描述
openapi: "3.0.0"
info:
title: 用户管理 API
version: "1.0.0"
paths:
/users:
get:
summary: 获取用户列表
parameters:
- name: page
in: query
schema:
type: integer
default: 1
- name: page_size
in: query
schema:
type: integer
default: 20
responses:
'200':
description: 用户列表
post:
summary: 创建用户
requestBody:
content:
application/json:
schema:
$ref: '#/components/schemas/CreateUser'
responses:
'201':
description: 创建成功
components:
schemas:
User:
type: object
properties:
id:
type: integer
name:
type: string
email:
type: string
created_at:
type: string
format: date-time
CreateUser:
type: object
required: [name, email]
properties:
name:
type: string
minLength: 1
maxLength: 50
email:
type: string
format: email
3.2 AI 生成的 Rust 代码
把上面的 OpenAPI 描述喂给 AI,让它生成 Rust 代码:
use axum::{
extract::{Query, Json, State},
http::StatusCode,
response::IntoResponse,
routing::{get, post},
Router,
};
use serde::{Deserialize, Serialize};
use sqlx::PgPool;
use validator::Validate; // 参数校验库
use chrono::NaiveDateTime;
// ====== 数据模型 ======
/// 用户实体(数据库映射)
#[derive(Debug, Serialize, sqlx::FromRow)]
pub struct User {
pub id: i64,
pub name: String,
pub email: String,
pub created_at: NaiveDateTime,
}
/// 创建用户请求体
#[derive(Debug, Deserialize, Validate)]
pub struct CreateUser {
#[validate(length(min = 1, max = 50, message = "姓名长度必须在1-50之间"))]
pub name: String,
#[validate(email(message = "邮箱格式不正确"))]
pub email: String,
}
/// 分页查询参数
#[derive(Debug, Deserialize)]
pub struct Pagination {
/// 页码,从 1 开始
#[serde(default = "default_page")]
pub page: i64,
/// 每页大小
#[serde(default = "default_page_size")]
pub page_size: i64,
}
fn default_page() -> i64 { 1 }
fn default_page_size() -> i64 { 20 }
// ====== Handler 实现 ======
/// GET /users - 获取用户列表(分页)
async fn list_users(
State(pool): State<PgPool>,
Query(pagination): Query<Pagination>,
) -> impl IntoResponse {
// 计算偏移量
let offset = (pagination.page - 1) * pagination.page_size;
// 查询用户列表 + 总数
let users = sqlx::query_as::<_, User>(
"SELECT id, name, email, created_at FROM users
ORDER BY created_at DESC LIMIT $1 OFFSET $2"
)
.bind(pagination.page_size)
.bind(offset)
.fetch_all(&pool)
.await;
match users {
Ok(users) => {
// 构造分页响应
let response = serde_json::json!({
"data": users,
"page": pagination.page,
"page_size": pagination.page_size,
"total": users.len() // 实际应单独查询总数
});
(StatusCode::OK, Json(response)).into_response()
}
Err(e) => {
let err = serde_json::json!({
"error": "查询失败",
"detail": e.to_string()
});
(StatusCode::INTERNAL_SERVER_ERROR, Json(err)).into_response()
}
}
}
/// POST /users - 创建新用户
async fn create_user(
State(pool): State<PgPool>,
// 使用 Json 提取器 + Validate 中间件自动校验参数
Json(payload): Json<CreateUser>,
) -> impl IntoResponse {
// 手动校验(生产环境建议用 axum-valid 中间件自动校验)
if let Err(errors) = payload.validate() {
let err = serde_json::json!({
"error": "参数校验失败",
"fields": errors.field_errors()
});
return (StatusCode::UNPROCESSABLE_ENTITY, Json(err)).into_response();
}
// 插入数据库
let result = sqlx::query_as::<_, User>(
"INSERT INTO users (name, email) VALUES ($1, $2)
RETURNING id, name, email, created_at"
)
.bind(&payload.name)
.bind(&payload.email)
.fetch_one(&pool)
.await;
match result {
Ok(user) => {
(StatusCode::CREATED, Json(user)).into_response()
}
Err(e) => {
let detail = e.to_string();
let status = if detail.contains("duplicate key") {
StatusCode::CONFLICT
} else {
StatusCode::INTERNAL_SERVER_ERROR
};
let err = serde_json::json!({
"error": "创建用户失败",
"detail": detail
});
(status, Json(err)).into_response()
}
}
}
/// 注册路由
pub fn user_routes() -> Router<PgPool> {
Router::new()
.route("/users", get(list_users).post(create_user))
}
四、方案二:基于 SQL Schema 生成 Repository 层
这是更常见的场景——已经有了数据库表结构,需要生成 Rust 的数据访问层。
4.1 输入:SQL Schema
CREATE TABLE orders (
id BIGSERIAL PRIMARY KEY,
user_id BIGINT NOT NULL REFERENCES users(id),
product_id BIGINT NOT NULL,
amount DECIMAL(10, 2) NOT NULL,
status VARCHAR(20) NOT NULL DEFAULT 'pending',
created_at TIMESTAMP NOT NULL DEFAULT NOW(),
updated_at TIMESTAMP NOT NULL DEFAULT NOW()
);
CREATE INDEX idx_orders_user_id ON orders(user_id);
CREATE INDEX idx_orders_status ON orders(status);
CREATE INDEX idx_orders_created_at ON orders(created_at);
4.2 AI 生成的 Repository
use sqlx::{PgPool, Postgres, Transaction};
use chrono::{DateTime, Utc};
use rust_decimal::Decimal;
/// 订单实体
#[derive(Debug, Clone, sqlx::FromRow, serde::Serialize, serde::Deserialize)]
pub struct Order {
pub id: i64,
pub user_id: i64,
pub product_id: i64,
pub amount: Decimal,
pub status: String,
pub created_at: DateTime<Utc>,
pub updated_at: DateTime<Utc>,
}
/// 创建订单的参数
#[derive(Debug, Deserialize)]
pub struct CreateOrderParams {
pub user_id: i64,
pub product_id: i64,
pub amount: Decimal,
}
/// 更新订单的参数(所有字段可选,只更新非空字段)
#[derive(Debug, Deserialize)]
pub struct UpdateOrderParams {
pub status: Option<String>,
pub amount: Option<Decimal>,
}
/// 订单查询条件
#[derive(Debug, Default)]
pub struct OrderFilter {
pub user_id: Option<i64>,
pub status: Option<String>,
pub start_date: Option<DateTime<Utc>>,
pub end_date: Option<DateTime<Utc>>,
pub limit: Option<i64>,
pub offset: Option<i64>,
}
/// 订单数据访问层
pub struct OrderRepo;
impl OrderRepo {
/// 创建订单并返回完整实体
pub async fn create(
pool: &PgPool,
params: CreateOrderParams,
) -> Result<Order, sqlx::Error> {
sqlx::query_as::<_, Order>(
"INSERT INTO orders (user_id, product_id, amount, status)
VALUES ($1, $2, $3, 'pending')
RETURNING id, user_id, product_id, amount, status, created_at, updated_at"
)
.bind(params.user_id)
.bind(params.product_id)
.bind(params.amount)
.fetch_one(pool)
.await
}
/// 根据 ID 查找订单
pub async fn find_by_id(
pool: &PgPool,
order_id: i64,
) -> Result<Option<Order>, sqlx::Error> {
sqlx::query_as::<_, Order>(
"SELECT id, user_id, product_id, amount, status, created_at, updated_at
FROM orders WHERE id = $1"
)
.bind(order_id)
.fetch_optional(pool)
.await
}
/// 按条件查询订单列表
/// 动态构建 WHERE 子句,只添加有值的条件
pub async fn find_by_filter(
pool: &PgPool,
filter: OrderFilter,
) -> Result<Vec<Order>, sqlx::Error> {
// 使用 query_builder 动态构建 SQL
let mut sql = String::from(
"SELECT id, user_id, product_id, amount, status, created_at, updated_at
FROM orders WHERE 1=1"
);
let mut params: Vec<String> = Vec::new();
let mut param_idx = 1;
// 动态追加条件
if let Some(ref user_id) = filter.user_id {
sql.push_str(&format!(" AND user_id = ${}", param_idx));
params.push(user_id.to_string());
param_idx += 1;
}
if let Some(ref status) = filter.status {
sql.push_str(&format!(" AND status = ${}", param_idx));
params.push(status.clone());
param_idx += 1;
}
if let Some(ref start) = filter.start_date {
sql.push_str(&format!(" AND created_at >= ${}", param_idx));
params.push(start.to_rfc3339());
param_idx += 1;
}
if let Some(ref end) = filter.end_date {
sql.push_str(&format!(" AND created_at <= ${}", param_idx));
params.push(end.to_rfc3339());
param_idx += 1;
}
sql.push_str(" ORDER BY created_at DESC");
if let Some(limit) = filter.limit {
sql.push_str(&format!(" LIMIT ${}", param_idx));
params.push(limit.to_string());
param_idx += 1;
}
if let Some(offset) = filter.offset {
sql.push_str(&format!(" OFFSET ${}", param_idx));
params.push(offset.to_string());
}
// 使用 sqlx::query 动态执行(注意:这里用了简单的字符串拼接
// 生产环境建议用 sqlx::QueryBuilder 防注入)
let mut query = sqlx::query_as::<_, Order>(&sql);
for p in ¶ms {
query = query.bind(p);
}
query.fetch_all(pool).await
}
/// 更新订单(只更新非空字段)
pub async fn update(
pool: &PgPool,
order_id: i64,
params: UpdateOrderParams,
) -> Result<Option<Order>, sqlx::Error> {
// 构建动态 UPDATE SET 子句
let mut sets = Vec::new();
let mut param_idx = 1;
// 用 query_builder 构建安全的参数化查询
if let Some(ref status) = params.status {
sets.push((format!("status = ${}", param_idx), status.clone()));
param_idx += 1;
}
if let Some(ref amount) = params.amount {
sets.push((format!("amount = ${}", param_idx), amount.to_string()));
param_idx += 1;
}
if sets.is_empty() {
// 无事可更新,直接返回原记录
return Self::find_by_id(pool, order_id).await;
}
// 总是更新 updated_at
sets.push((format!("updated_at = NOW()"), String::new()));
let set_clause: Vec<String> = sets.iter()
.map(|(s, _)| s.clone())
.collect();
let sql = format!(
"UPDATE orders SET {} WHERE id = ${} RETURNING id, user_id, product_id, amount, status, created_at, updated_at",
set_clause.join(", "),
param_idx
);
let mut query = sqlx::query_as::<_, Order>(&sql);
for (_, value) in &sets {
if !value.is_empty() {
query = query.bind(value);
}
}
query = query.bind(order_id);
query.fetch_optional(pool).await
}
}
AI 代码生成的边界与现实
我的实践心得
- AI 生成 + 人工审核是最佳组合。AI 生成 80% 的样板代码,工程师只需要关注 20% 的业务逻辑和安全细节。
- Prompt 的质量决定代码的质量。描述越具体(表结构、字段校验规则、错误处理策略),生成的代码越贴近需求。
- 把 AI 生成的代码当作"第一版草稿",而不是"最终代码"。编译器、clippy、集成测试是第二道防线。
-
Process Macro 和 AI 各司其职。能用
#[derive]搞定的(Serialize/Deserialize/FromRow),优先用宏。需要语义理解的(业务逻辑、SQL 拼接、错误处理),才交给 AI。
实际使用中的一个意外发现:AI 生成的 Repository 代码,在 find_by_filter 里用了字符串拼接 SQL 而不是 sqlx::QueryBuilder。虽然代码看起来"能用",但绑定参数的方式有 SQL 注入风险。这恰恰说明 AI 生成的"边界"——它能写出看起来正确的代码,但不能保证安全。编译通过只是及格线,clippy lint + 安全审计才是上线标准。
五、总结
- 从 OpenAPI 生成 axum handler:数据模型、参数校验、路由注册一气呵成。
- 从 SQL Schema 生成 Repository 层:完整的 CRUD、动态查询、分页排序。
AI 代码生成不是要取代程序员,而是要消灭那些浪费程序员生命的重复劳动。Rust 的类型系统严格,恰好是 AI 代码生成的好搭档——编译器会在编译期把所有类型错误揪出来,AI 写错的代码根本过不了编译。
这就是我目前探索到的边界——AI 可以帮我省掉 70% 的重复琐碎工作,但架构决策、业务建模、安全把控这些核心能力,暂时还是工程师的独特价值。
未来随着 AI Agent 能力增强,"从需求文档到可运行代码"的全自动流程也许真的不远了。但在那一天到来之前,我们要做的是:用好 AI 工具,把精力花在更有创造性的工作上。
有什么想法欢迎评论区交流!