use candid::Principal;
use url::Url;
const DEFAULT_SCHEME: &str = "https";
const DEFAULT_STATEMENT: &str = "SIWE Fields:";
const DEFAULT_CHAIN_ID: u32 = 1; const DEFAULT_SIGN_IN_EXPIRES_IN: u64 = 60 * 5 * 1_000_000_000; const DEFAULT_SESSION_EXPIRES_IN: u64 = 30 * 60 * 1_000_000_000;
#[derive(Debug, Clone, PartialEq)]
pub enum RuntimeFeature {
IncludeUriInSeed,
}
#[derive(Default, Debug, Clone)]
pub struct Settings {
pub domain: String,
pub uri: String,
pub salt: String,
pub chain_id: u32,
pub scheme: String,
pub statement: String,
pub sign_in_expires_in: u64,
pub session_expires_in: u64,
pub targets: Option<Vec<Principal>>,
pub runtime_features: Option<Vec<RuntimeFeature>>,
}
pub struct SettingsBuilder {
settings: Settings,
}
impl SettingsBuilder {
pub fn new<S: Into<String>, T: Into<String>, U: Into<String>>(
domain: S,
uri: T,
salt: U,
) -> Self {
SettingsBuilder {
settings: Settings {
domain: domain.into(),
uri: uri.into(),
salt: salt.into(),
chain_id: DEFAULT_CHAIN_ID,
scheme: DEFAULT_SCHEME.to_string(),
statement: DEFAULT_STATEMENT.to_string(),
sign_in_expires_in: DEFAULT_SIGN_IN_EXPIRES_IN,
session_expires_in: DEFAULT_SESSION_EXPIRES_IN,
targets: None,
runtime_features: None,
},
}
}
pub fn chain_id(mut self, chain_id: u32) -> Self {
self.settings.chain_id = chain_id;
self
}
pub fn scheme<S: Into<String>>(mut self, scheme: S) -> Self {
self.settings.scheme = scheme.into();
self
}
pub fn statement<S: Into<String>>(mut self, statement: S) -> Self {
self.settings.statement = statement.into();
self
}
pub fn sign_in_expires_in(mut self, expires_in: u64) -> Self {
self.settings.sign_in_expires_in = expires_in;
self
}
pub fn session_expires_in(mut self, expires_in: u64) -> Self {
self.settings.session_expires_in = expires_in;
self
}
pub fn targets(mut self, targets: Vec<Principal>) -> Self {
self.settings.targets = Some(targets);
self
}
pub fn runtime_features(mut self, features: Vec<RuntimeFeature>) -> Self {
self.settings.runtime_features = Some(features);
self
}
pub fn build(self) -> Result<Settings, String> {
validate_domain(&self.settings.scheme, &self.settings.domain)?;
validate_uri(&self.settings.uri)?;
validate_salt(&self.settings.salt)?;
validate_chain_id(self.settings.chain_id)?;
validate_scheme(&self.settings.scheme)?;
validate_statement(&self.settings.statement)?;
validate_sign_in_expires_in(self.settings.sign_in_expires_in)?;
validate_session_expires_in(self.settings.session_expires_in)?;
validate_targets(&self.settings.targets)?;
Ok(self.settings)
}
}
fn validate_domain(scheme: &str, domain: &str) -> Result<String, String> {
let url_str = format!("{}://{}", scheme, domain);
let parsed_url = Url::parse(&url_str).map_err(|_| String::from("Invalid domain"))?;
if !parsed_url.has_authority() {
Err(String::from("Invalid domain"))
} else {
Ok(parsed_url.host_str().unwrap().to_string())
}
}
fn validate_uri(uri: &str) -> Result<String, String> {
let parsed_uri = Url::parse(uri).map_err(|_| String::from("Invalid URI"))?;
if !parsed_uri.has_host() {
Err(String::from("Invalid URI"))
} else {
Ok(uri.to_string())
}
}
fn validate_salt(salt: &str) -> Result<String, String> {
if salt.is_empty() {
return Err(String::from("Salt cannot be empty"));
}
if salt.chars().any(|c| !c.is_ascii() || !c.is_ascii_graphic()) {
return Err(String::from("Invalid salt"));
}
Ok(salt.to_string())
}
fn validate_chain_id(chain_id: u32) -> Result<u32, String> {
if chain_id == 0 {
return Err(String::from("Chain ID must be greater than 0"));
}
Ok(chain_id)
}
fn validate_scheme(scheme: &str) -> Result<String, String> {
if scheme == "http" || scheme == "https" {
return Ok(scheme.to_string());
}
Err(String::from("Invalid scheme"))
}
fn validate_statement(statement: &str) -> Result<String, String> {
if statement.contains('\n') {
return Err(String::from("Invalid statement"));
}
Ok(statement.to_string())
}
fn validate_sign_in_expires_in(expires_in: u64) -> Result<u64, String> {
if expires_in == 0 {
return Err(String::from("Sign in expires in must be greater than 0"));
}
Ok(expires_in)
}
fn validate_session_expires_in(expires_in: u64) -> Result<u64, String> {
if expires_in == 0 {
return Err(String::from("Session expires in must be greater than 0"));
}
Ok(expires_in)
}
fn validate_targets(targets: &Option<Vec<Principal>>) -> Result<Option<Vec<Principal>>, String> {
if let Some(targets) = targets {
if targets.is_empty() {
return Err(String::from("Targets cannot be empty"));
}
if targets.len() > 1000 {
return Err(String::from("Too many targets"));
}
let mut targets_clone = targets.clone();
targets_clone.sort();
targets_clone.dedup();
if targets_clone.len() != targets.len() {
return Err(String::from("Duplicate targets are not allowed"));
}
}
Ok(targets.clone())
}
#[cfg(test)]
mod tests {
use super::*;
use candid::Principal;
#[test]
fn test_successful_settings_creation_defaults() {
let builder = SettingsBuilder::new("example.com", "http://example.com", "some_salt");
let settings = builder
.build()
.expect("Failed to create settings with defaults");
assert_eq!(settings.domain, "example.com");
assert_eq!(settings.uri, "http://example.com");
assert_eq!(settings.salt, "some_salt");
assert_eq!(settings.chain_id, DEFAULT_CHAIN_ID);
assert_eq!(settings.scheme, DEFAULT_SCHEME);
assert_eq!(settings.statement, DEFAULT_STATEMENT);
assert_eq!(settings.sign_in_expires_in, DEFAULT_SIGN_IN_EXPIRES_IN);
assert_eq!(settings.session_expires_in, DEFAULT_SESSION_EXPIRES_IN);
assert!(settings.targets.is_none());
}
#[test]
fn test_successful_settings_creation_custom() {
let targets = vec![Principal::anonymous()];
let builder = SettingsBuilder::new("example.com", "http://example.com", "some_salt")
.chain_id(3)
.scheme("http")
.statement("Custom statement")
.sign_in_expires_in(10_000_000_000)
.session_expires_in(20_000_000_000)
.targets(targets.clone());
let settings = builder
.build()
.expect("Failed to create settings with custom values");
assert_eq!(settings.chain_id, 3);
assert_eq!(settings.scheme, "http");
assert_eq!(settings.statement, "Custom statement");
assert_eq!(settings.sign_in_expires_in, 10_000_000_000);
assert_eq!(settings.session_expires_in, 20_000_000_000);
assert_eq!(settings.targets, Some(targets));
}
#[test]
fn test_empty_salt() {
let builder = SettingsBuilder::new("example.com", "http://example.com", "");
assert!(builder.build().is_err());
}
#[test]
fn test_invalid_chain_id() {
let builder =
SettingsBuilder::new("example.com", "http://example.com", "some_salt").chain_id(0);
assert!(builder.build().is_err());
}
#[test]
fn test_invalid_scheme() {
let builder =
SettingsBuilder::new("example.com", "http://example.com", "some_salt").scheme("ftp");
assert!(builder.build().is_err());
}
#[test]
fn test_invalid_statement() {
let builder = SettingsBuilder::new("example.com", "http://example.com", "some_salt")
.statement("Invalid\nStatement");
assert!(builder.build().is_err());
}
#[test]
fn test_sign_in_expires_in_zero() {
let builder = SettingsBuilder::new("example.com", "http://example.com", "some_salt")
.sign_in_expires_in(0);
assert!(builder.build().is_err());
}
#[test]
fn test_session_expires_in_zero() {
let builder = SettingsBuilder::new("example.com", "http://example.com", "some_salt")
.session_expires_in(0);
assert!(builder.build().is_err());
}
#[test]
fn test_empty_targets() {
let builder =
SettingsBuilder::new("example.com", "http://example.com", "some_salt").targets(vec![]);
assert!(builder.build().is_err());
}
#[test]
fn test_too_many_targets() {
let builder = SettingsBuilder::new("example.com", "http://example.com", "some_salt")
.targets(vec![Principal::anonymous(); 1001]);
assert!(builder.build().is_err());
}
#[test]
fn test_duplicate_targets() {
let builder = SettingsBuilder::new("example.com", "http://example.com", "some_salt")
.targets(vec![Principal::anonymous(), Principal::anonymous()]);
assert!(builder.build().is_err());
}
#[test]
fn test_valid_domain_formats() {
let domains = vec!["example.com", "sub.domain.com", "example.co.uk"];
for domain in domains {
let builder = SettingsBuilder::new(domain, "http://example.com", "some_salt");
assert!(builder.build().is_ok(), "Failed with domain: {}", domain);
}
}
#[test]
fn test_invalid_domain_formats() {
let domains = vec![""];
for domain in domains {
let builder = SettingsBuilder::new(domain, "http://example.com", "some_salt");
assert!(
builder.build().is_err(),
"Should fail with domain: {}",
domain
);
}
}
#[test]
fn test_valid_uri_formats() {
let uris = vec!["http://example.com", "https://example.com:8080/path"];
for uri in uris {
let builder = SettingsBuilder::new("example.com", uri, "some_salt");
assert!(builder.build().is_ok(), "Failed with URI: {}", uri);
}
}
#[test]
fn test_invalid_uris() {
let uris = vec!["", "just_string"];
for uri in uris {
let builder = SettingsBuilder::new("example.com", uri, "some_salt");
assert!(builder.build().is_err(), "Should fail with URI: {}", uri);
}
}
#[test]
fn test_chain_id_zero() {
let builder =
SettingsBuilder::new("example.com", "http://example.com", "some_salt").chain_id(0);
assert!(builder.build().is_err(), "Chain ID zero should fail");
}
#[test]
fn test_uri_with_port_numbers() {
let builder = SettingsBuilder::new("example.com", "http://example.com:8080", "some_salt");
assert!(builder.build().is_ok());
}
#[test]
fn test_valid_salt_lengths() {
for len in [1, 10, 100].iter() {
let salt = "a".repeat(*len);
let builder = SettingsBuilder::new("example.com", "http://example.com", &salt);
assert!(builder.build().is_ok(), "Failed with salt length: {}", len);
}
}
#[test]
fn test_chain_id_boundary_values() {
let max_value = u32::MAX;
let builder = SettingsBuilder::new("example.com", "http://example.com", "some_salt")
.chain_id(max_value);
assert!(builder.build().is_ok());
}
#[test]
fn test_scheme_case_sensitivity() {
let builder =
SettingsBuilder::new("example.com", "http://example.com", "some_salt").scheme("HTTP");
assert!(builder.build().is_err());
}
#[test]
fn test_statement_length_and_content() {
let long_statement = "a".repeat(1000);
let builder = SettingsBuilder::new("example.com", "http://example.com", "some_salt")
.statement(long_statement);
assert!(builder.build().is_ok());
}
#[test]
fn test_extreme_expiration_values() {
let builder = SettingsBuilder::new("example.com", "http://example.com", "some_salt")
.sign_in_expires_in(1)
.session_expires_in(u64::MAX);
assert!(builder.build().is_ok());
}
#[test]
fn test_targets_with_various_principal_formats() {
let targets = vec![
Principal::anonymous(),
Principal::from_text("aaaaa-aa").unwrap(),
];
let builder =
SettingsBuilder::new("example.com", "http://example.com", "some_salt").targets(targets);
assert!(builder.build().is_ok());
}
#[test]
fn test_multiple_valid_and_invalid_combinations() {
let builder = SettingsBuilder::new("", "invalid_uri", "")
.chain_id(0)
.scheme("ftp");
assert!(builder.build().is_err());
}
#[test]
fn test_partially_initialized_builder() {
let builder =
SettingsBuilder::new("example.com", "http://example.com", "some_salt").scheme("http");
assert!(builder.build().is_ok());
}
#[test]
fn test_overwriting_default_values() {
let builder = SettingsBuilder::new("example.com", "http://example.com", "some_salt")
.scheme(DEFAULT_SCHEME)
.chain_id(DEFAULT_CHAIN_ID);
assert!(builder.build().is_ok());
}
#[test]
fn test_malformed_uris() {
let builder = SettingsBuilder::new("example.com", "://missing_protocol.com", "some_salt");
assert!(builder.build().is_err());
}
#[test]
fn test_invalid_salt_content() {
let builder = SettingsBuilder::new("example.com", "http://example.com", "\0invalid_salt");
assert!(builder.build().is_err());
}
#[test]
fn test_invalid_statement_formats() {
let builder = SettingsBuilder::new("example.com", "http://example.com", "some_salt")
.statement("Invalid\nStatement");
assert!(builder.build().is_err());
}
#[test]
fn test_validating_an_empty_settingsbuilder() {
let builder = SettingsBuilder::new("", "", "");
assert!(builder.build().is_err());
}
#[test]
fn test_domain_with_international_characters() {
let builder = SettingsBuilder::new("xn--exmple-cua.com", "http://example.com", "some_salt");
assert!(builder.build().is_ok());
}
}