use crate::crypto_wrapper::CryptoWrapper;
use crate::errors::OpenIdError;
use crate::time_utils::time;
use crate::Res;
use std::net::IpAddr;
#[derive(Debug, rkyv::Archive, rkyv::Serialize, rkyv::Deserialize)]
struct State {
ip: IpAddr,
expire: u64,
}
impl State {
pub fn new(ip: IpAddr) -> Self {
Self {
ip,
expire: time() + 15 * 60,
}
}
}
pub struct BasicStateManager(CryptoWrapper);
impl BasicStateManager {
pub fn new() -> Self {
Self(CryptoWrapper::new_random())
}
pub fn new_with_wrapper(wrapper: CryptoWrapper) -> Self {
Self(wrapper)
}
pub fn gen_state(&self, ip: IpAddr) -> Res<String> {
let state = State::new(ip);
self.0.encrypt(&state)
}
pub fn validate_state(&self, ip: IpAddr, state: &str) -> Res {
let state: State = self.0.decrypt(state)?;
if state.ip != ip {
return Err(OpenIdError::StateErrorInvalidIP);
}
if state.expire < time() {
return Err(OpenIdError::StateErrorExpired);
}
Ok(())
}
}
impl Default for BasicStateManager {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod test {
use crate::basic_state_manager::BasicStateManager;
use std::net::{IpAddr, Ipv4Addr};
const IP_1: IpAddr = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1));
const IP_2: IpAddr = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 2));
#[test]
fn valid_state() {
let manager = BasicStateManager::new();
let state = manager.gen_state(IP_1).unwrap();
assert!(manager.validate_state(IP_1, &state).is_ok());
}
#[test]
fn invalid_ip() {
let manager = BasicStateManager::new();
let state = manager.gen_state(IP_1).unwrap();
assert!(manager.validate_state(IP_2, &state).is_err());
}
}