前言#
昨天我们做了一些基础配置,今天我们正式开始来写路由了,我们先来实现用户相关的内容
项目地址:
https://github.com/1714080902120/rust_rocket_crud_demo由于我是先实现完再写的这篇文章,如果有些地方无法运行,可以看下我项目里的代码。
目前我还在往全栈的方向学习,所以如果看的不顺眼,请多多包涵。如果觉得那里可以改进,麻烦评论区说下,谢谢~
token#
我们的服务器在请求之前必须要先登录才行,那么就需要有一个字段即token用于我们的验证。
token生成和解析#
在我们开始实现校验部分之前,我们需要了解token的创建和解析。
这里我用的jsonwebtoken[1]这个crate
简单的看下它的用法
encode
let mut header = Header::new(Algorithm::HS512);
header.kid = Some("blabla".to_owned());
let token = encode(&header, &my_claims, &EncodingKey::from_secret("secret".as_ref()))?;
let mut header = Header::new(Algorithm::HS512);
header.kid = Some("blabla".to_owned());
let token = encode(&header, &my_claims, &EncodingKey::from_secret("secret".as_ref()))?;
decode
// `token` is a struct with 2 fields: `header` and `claims` where `claims` is your own struct.
let token = decode::<Claims>(&token, &DecodingKey::from_secret("secret".as_ref()), &Validation::default())?;你可以自定义插入到header里的数据
#[derive(Debug, Serialize, Deserialize)]
struct Claims {
aud: String, // Optional. Audience
exp: usize, // Required (validate_exp defaults to true in validation). Expiration time (as UTC timestamp)
iat: usize, // Optional. Issued at (as UTC timestamp)
iss: String, // Optional. Issuer
nbf: usize, // Optional. Not Before (as UTC timestamp)
sub: String, // Optional. Subject (whom token refers to)
} 其中exp是强制要求的字段,即你自定义的struct里必须要有一个exp字段,用于表示token的有效时间。
我们先在src下创建一个auth的mod,然后在里面分别创建token.rs和mod.rs
mod.rs中我们创建一个UserToken的struct,它是我们插入到header里的东西。
#[derive(Serialize, Deserialize, Clone, PartialEq, Eq, Debug)]
pub struct UserToken {
pub id: String,
pub exp: u64,
} id:用户的id,为了确保安全,这个信息只会放在header的_token字段中,并且是加密过后的。exp:这个则是用于校验token是否过期的。
token.rs 存放我们的生成/解析token的方法
use jsonwebtoken::{
decode, encode, errors::Error, get_current_timestamp, DecodingKey, EncodingKey, Header,
TokenData, Validation,
};
use rocket::Request;
use super::{UserToken};
use crate::config::MyConfig;
pub fn decode_token(token: &str, key: &str) -> Result<TokenData<UserToken>, Error> {
decode::<UserToken>(
token,
&DecodingKey::from_secret(key.as_ref()),
&Validation::default(),
)
}
pub fn encode_token(user_id: String, exp: u64, key: &str) -> String {
let user_token = UserToken {
id: user_id,
exp,
};
encode::<UserToken>(
&Header::default(),
&user_token,
&EncodingKey::from_secret(key.as_ref()),
)
.expect("encode token error")
}
pub fn set_token<'r>(req: &'r Request<'_>, user_id: &str) -> (String, String) {
let my_config = req
.rocket()
.state::<MyConfig>()
.expect("get global state error when response in login");
let token_field = my_config.token_field.as_str();
let token_key = my_config.token_key.as_str();
let exp = my_config.exp + get_current_timestamp();
let token = encode_token(user_id.to_string(), exp, token_key);
(token_field.to_string(), token)
}
encode_token就是用来生成token,注意这里的Header是jsonwebtoken自己的Header,而不是rocket的Header,而user_token则是我们的数据,它有一个强制的字段要求exp,表示过期时间。key则是我们加密的key,加密的方式我们就用它默认的即可。decode_token自然就是解析set_token:这个用于返回我们生成的token信息,在这里我们用到了之前的MyConfig。get_current_timestamp[2]:获取当前时间的时间戳,加上我们在MyConfig中配置的exp时间间隔,这样就可以用来表示过期时间了。req.rocket().state::():我们可以通过这种方式获取全局的state,注意这里必须用::来声明指定的类型。
校验#
由于我这定义了请求前必须要登入,任何接口都是。所以这个时候我们就需要给每个路由都加上对于token的校验,根据我们学的知识,我们可以把这个放到全局的fairing,相当于一个路由守卫,能参与一次请求的任意生命周期。
回到auth文件夹,在mod.rs文件夹中我们新建一个名为AuthMsg的struct
#[derive(Debug)]
pub struct AuthMsg {
pub is_valid_token: bool,
}这个struct用于临时数据,等会会用到。
然后我们新建一个fairing.rs的文件
use crate::types::rt_type::Rt;
use jsonwebtoken::{get_current_timestamp};
use rocket::{
fairing::{Fairing, Info, Kind},
http::{ContentType, Status},
Data, Request, Response,
};
use std::{cmp::Ordering, io::Cursor};
use crate::{types::{RtData, FailureData}, config::MyConfig};
use crate::auth::{ UserToken, AuthMsg};
use super::token::decode_token;
const EXCEPT_LIST: [Status; 3] = [Status::NotFound, Status::InternalServerError, Status::BadRequest];
#[rocket::async_trait]
impl Fairing for UserToken {
fn info(&self) -> Info {
Info {
name: "user authorized",
kind: Kind::Request | Kind::Response,
}
}
async fn on_request(&self, request: &mut Request<'_>, _: &mut Data<'_>) {
let my_config = request.rocket().state::<MyConfig>().expect("get global custom config error in fairing");
let token_field = my_config.token_field.as_str();
let token_key = my_config.token_key.as_str();
let header = request.headers();
let token_data = header.get(token_field).next();
let token = match token_data {
Some(token) => {
dbg!(token);
match decode_token(token, token_key) {
Ok(user_token) => user_token,
Err(err) => {
dbg!(err);
return;
}
}
}
None => {
dbg!("has none token");
return;
}
};
// validate time
let exp: u64 = token.claims.exp;
match get_current_timestamp().cmp(&exp) {
Ordering::Less => {
request.local_cache(|| AuthMsg {
is_valid_token: true,
});
}
_ => {}
}
}
async fn on_response<'r>(&self, req: &'r Request<'_>, res: &mut Response<'r>) {
let auth_state = req.local_cache(|| AuthMsg {
is_valid_token: false,
});
dbg!(auth_state);
if !auth_state.is_valid_token && !EXCEPT_LIST.contains(&res.status()) {
res.set_status(Status::NonAuthoritativeInformation);
res.set_header(ContentType::JSON);
let data = RtData {
success: false,
rt: Rt::Fail,
data: FailureData(()),
msg: String::from("user not login or expired token !")
};
let data_str = data.to_string();
res.set_sized_body(data_str.len(), Cursor::new(data_str));
}
}
}实现Fairing[3]必须实现info这个method,目的是告知rocket哪些生命周期需要监听,这里我监听的是Request和Response。
在request进入到路由之前就会触发这个on_request函数,而response准备发送给前端的时候即执行完处理器才会触发on_response。
这里就用到了我们前面实现的decode_token方法,用于获取request的Header里面的token信息。
on_request我通过判断是否存在token,token的exp是否已经过期这两步来实现校验。
如果通过校验,那么请求的暂时缓存local_cache[4]就会放一个AuthMsg的数据,里面存放了一个标志位is_valid_token ,这个标志位用于response的时候判断是否符合要求,不符合直接重写response(这么做其实不太好,因为处理器里面的逻辑实际上还是执行了,浪费了很多性能,但是我找不到较好的方案。。。如果你有好的方案,请务必评论区里说下,谢谢!)
另外吐槽下这个local_cache,它并不能改动数据,是的。。。也可能是我姿势不对,我尝试了一段时间后就放弃了。
虽然说每个路由都需要校验token,但是实际上有些是不需要的,比如404、500、400等。 所以这里搞了个白名单EXCEPT_LIST由于过滤掉上面的几种错误场景
on_response中重写数据时,我用到了前面我们实现的RtData类型,但是如果直接用是会报错的,因为RtData并不符合类型约束,没有实现Responder[5] 。
那么我们回到type/mod.rs中,给RtData实现这个trait。
use std::io::Cursor;
use rocket::{
data::{self, FromData},
http::{ContentType, Status},
request::{FromRequest, Outcome},
response,
response::Responder,
Data, Request, Response,
};
use serde::{Deserialize, Serialize};
use crate::{
auth::{token::decode_token, AuthMsg},
config::MyConfig,
};
#[derive(Debug, PartialEq, Eq, Clone, Copy, Serialize, Deserialize)]
pub struct FailureData(pub ());
impl<'r> Responder<'r, 'static> for RtData<FailureData> {
fn respond_to(self, _: &'r Request<'_>) -> response::Result<'static> {
let data = self.to_string();
Response::build()
.header(ContentType::JSON)
.sized_body(data.len(), Cursor::new(data))
.ok()
}
}给它实现Responder,支持我们的类型作为response。
这样就可以了。
另外提一下这个std::io::Cursor::new,这个用来指向数据的起始位置

那么token校验相关的就都完成了
最后我们在main.rs中引入,不过在这之前,我们还需要在auth/mod.rs中引入,包括前面提到的一些文件,不要忘了在mod中引入以及pub出去,如果只是文件夹里面跨文件使用,那么可以不需要pub出去。
mod fairing; 然后我们新建一个state的文件夹,在里面新建mod.rs文件。这个文件夹作为存放我们全局state的地方。
use crate::auth::{UserToken};
use jsonwebtoken::get_current_timestamp;
use rocket::http::Status;
pub fn get_default_user_token() -> UserToken {
UserToken {
id: uuid::Uuid::new_v4().to_string(),
exp: get_current_timestamp(),
}
}
最后我们在mian.rs中引入
// ...
mod state;
// ...
use state::{get_default_user_token};
// ...
#[launch]
fn rocket() -> _ {
rocket::custom(get_custom_figment())
// ...
.attach(get_default_user_token())
// ...
} 这样就行啦
下一章我们来开始编写user相关的路由了。
其它#
另外有一点就是这个log,错误的收集,实际开发我们是需要把log都收集到.log文件中的,这样方便日志跟踪排错,我这里图省事就没这么做。我逛了一圈,发现有一个log4rs[6]的,显然是rust版本的log4j,但是我这里并没有采用,大家有需要可以自行引入。
当然你也可以自己实现,比如用channel收集错误信息,用队列存储,达到一定的量或者每隔一小时就批量写一次进文件中。。。
补充优化#
前面我们用的是fairing监听on_request和on_response来限制反馈,但是这里有个大问题和bug。
bug是路由的function实际上还是会执行,也就是说如果不在路由函数那边加上判断是否可执行和return的逻辑,那么数据还是可以修改成功。
另外有些路由是没必要走这个校验逻辑的,所以要么我们加上白名单,要么我们最好就是换一个方案。。
实现一个内部参数:user_msg#
还记得我们在给路由参数实现的Request等类型么?实际上我们可以把token 校验放在参数中,这个参数不需要前端传递,而是一个内部参数。这里命名为user_msg,里面存放的是token解密之后的参数id,而另一个数据exp就没必要了,仅做为token校验用的。
#[derive(Debug, Serialize, Deserialize, PartialEq, Eq)]
pub struct UserMsg {
pub id: String,
}
#[rocket::async_trait]
impl<'r> FromRequest<'r> for UserMsg {
type Error = String;
async fn from_request(req: &'r Request<'_>) -> Outcome<Self, Self::Error> {
let my_config = req
.rocket()
.state::<MyConfig>()
.expect("get global custom config error in fairing");
let token_field = my_config.token_field.as_str();
let token_key = my_config.token_key.as_str();
let header = req.headers();
let token_data = header.get(token_field).next();
if let Some(token) = token_data {
let token = decode_token(token, token_key).unwrap();
let id = token.claims.id;
let exp = token.claims.exp;
return match (get_current_timestamp() * 1000).cmp(&exp) {
Ordering::Less => Outcome::Success(UserMsg { id }),
_ => Outcome::Failure((Status::Unauthorized, String::from("not auth"))),
};
} else {
return Outcome::Failure((
Status::Unauthorized,
String::from("user no login or token expired"),
));
}
}
}代码很简单,和前面fairings中做的一样,不同点在于这里校验失败直接通过Outcome::Failure的方式直接去到error catcher中,这里就不会继续往下执行路由的逻辑了。
那么该如何使用这个UserMsg来限制一些借口呢?简单,如果有需要检验token的地方,就在参数中加上这个UserMsg即可。
比如:更新用户数据的,这里就必须要要有token才行。
#[post("/update_user_data", data = "<update_user_data>")]
pub async fn update_user_data(db: BlogDBC, user_msg: UserMsg, update_user_data: Form<UpdateUserData>) -> Result<RtData<DataType>, Status> {
match try_update_user_data(db, &user_msg.id, update_user_data).await {
Ok(state) if state => {
return Ok(RtData { success: true, rt: Rt::Fail, data: DataType::DefaultSuccessData(()), msg: String::from("email had been register") });
}
Ok(_) => {
return Ok(RtData { success: false, rt: Rt::Fail, data: DataType::Fail, msg: String::from("email had been register") });
},
Err(err) => {
dbg!(&err);
return Err(Status::InternalServerError)
}
}
}这样也不需要使用白名单的方式来过滤路由
参考#
- ^jsonwebtoken https://crates.io/crates/jsonwebtoken
- ^get_current_timestamp https://docs.rs/jsonwebtoken/8.3.0/jsonwebtoken/fn.get_current_timestamp.html
- ^Fairing https://api.rocket.rs/v0.5-rc/rocket/fairing/trait.Fairing.html
- ^local_cache https://api.rocket.rs/v0.5-rc/rocket/request/macro.local_cache.html
- ^Responder https://api.rocket.rs/v0.5-rc/rocket/response/trait.Responder.html
- ^log4rs https://crates.io/crates/log4rs
编辑于 2023-09-04 21:03・IP 属地广东
