use serde_json::Value;
use sqlx::Row;
use sqlx::sqlite::SqliteRow;
use tracing::{debug, info};
use uuid::Uuid;
use crate::audit::ClientContext;
use crate::sqlite::db::Database;
use crate::sqlite::nonce::now_secs;
#[derive(Debug)]
pub struct Account {
pub id: Uuid,
pub profile: String,
pub pubkey: Vec<u8>,
pub contact: Vec<String>,
pub status: String,
pub created_at: i64,
pub eab_kid: Option<Uuid>,
pub terms_of_service_agreed: Option<bool>,
pub created_ip: Option<String>,
pub created_ptr: Option<String>,
pub last_seen_at: Option<i64>,
pub last_seen_ip: Option<String>,
pub last_seen_ptr: Option<String>,
}
pub const ACCOUNT_TOUCH_INTERVAL: i64 = 60;
pub(crate) fn pubkey_fingerprint(pubkey: &[u8]) -> String {
let digest = ring::digest::digest(&ring::digest::SHA256, pubkey);
hex::encode(&digest.as_ref()[..8])
}
macro_rules! columns {
() => {
"id, profile, pubkey, contact, status, created_at, eab_kid, \
terms_of_service_agreed, created_ip, created_ptr, last_seen_at, \
last_seen_ip, last_seen_ptr"
};
}
impl Account {
fn from_row(row: SqliteRow) -> Result<Self, sqlx::Error> {
let contact_json: String = row.try_get("contact")?;
let contact: Vec<String> =
serde_json::from_str(&contact_json).map_err(|e| sqlx::Error::Decode(Box::new(e)))?;
Ok(Account {
id: row.try_get("id")?,
profile: row.try_get("profile")?,
pubkey: row.try_get("pubkey")?,
contact,
status: row.try_get("status")?,
created_at: row.try_get("created_at")?,
eab_kid: row.try_get("eab_kid")?,
terms_of_service_agreed: row.try_get("terms_of_service_agreed")?,
created_ip: row.try_get("created_ip")?,
created_ptr: row.try_get("created_ptr")?,
last_seen_at: row.try_get("last_seen_at")?,
last_seen_ip: row.try_get("last_seen_ip")?,
last_seen_ptr: row.try_get("last_seen_ptr")?,
})
}
#[tracing::instrument(name = "Account::find_by_pubkey", skip(pubkey, database))]
pub async fn find_by_pubkey(
profile: &str,
pubkey: &[u8],
database: &Database,
) -> Result<Option<Account>, sqlx::Error> {
debug!(event = "db_account_find_by_pubkey_started", outcome = "progress", profile = %profile, pubkey_fp = %pubkey_fingerprint(pubkey));
let row = sqlx::query(concat!(
"SELECT ",
columns!(),
" FROM accounts WHERE profile = ? AND pubkey = ?;"
))
.bind(profile)
.bind(pubkey)
.fetch_optional(&database.pool)
.await?;
let result = row.map(Account::from_row).transpose()?;
if let Some(ref account) = result {
debug!(event = "db_account_found_by_pubkey", outcome = "success", account_id = %account.id, pubkey_fp = %pubkey_fingerprint(pubkey));
} else {
debug!(event = "db_account_not_found_by_pubkey", outcome = "failure", pubkey_fp = %pubkey_fingerprint(pubkey));
}
Ok(result)
}
#[tracing::instrument(name = "Account::find_by_id", skip(database), fields(account_id = %id))]
pub async fn find_by_id(
profile: &str,
id: &str,
database: &Database,
) -> Result<Option<Account>, sqlx::Error> {
debug!(event = "db_account_find_by_id_started", outcome = "progress", profile = %profile, account_id = %id);
let Some(id) = crate::sqlite::id::parse(id) else {
return Ok(None);
};
let row = sqlx::query(concat!(
"SELECT ",
columns!(),
" FROM accounts WHERE profile = ? AND id = ?;"
))
.bind(profile)
.bind(id)
.fetch_optional(&database.pool)
.await?;
let result = row.map(Account::from_row).transpose()?;
if let Some(ref account) = result {
debug!(event = "db_account_found_by_id", outcome = "success", account_id = %account.id);
} else {
debug!(event = "db_account_not_found_by_id", outcome = "failure", account_id = %id);
}
Ok(result)
}
#[tracing::instrument(name = "Account::find_or_create", skip(pubkey, client, database))]
pub async fn find_or_create(
profile: &str,
pubkey: &[u8],
contact: Vec<String>,
client: &ClientContext,
database: &Database,
) -> Result<(Account, bool), sqlx::Error> {
debug!(event = "db_account_find_or_create_started", outcome = "progress", profile = %profile, pubkey_fp = %pubkey_fingerprint(pubkey));
if let Some(account) = Account::find_by_pubkey(profile, pubkey, database).await? {
debug!(event = "db_account_found_existing", outcome = "success", account_id = %account.id, pubkey_fp = %pubkey_fingerprint(pubkey));
return Ok((account, false));
}
let account = Account {
id: crate::sqlite::id::mint(),
profile: profile.to_string(),
pubkey: pubkey.to_vec(),
contact,
status: "valid".to_string(),
created_at: now_secs(),
eab_kid: None,
terms_of_service_agreed: None,
created_ip: client.ip.clone(),
created_ptr: client.ptr.clone(),
last_seen_at: Some(now_secs()),
last_seen_ip: client.ip.clone(),
last_seen_ptr: client.ptr.clone(),
};
let contact_json = Value::from(account.contact.clone()).to_string();
debug!(event = "db_account_create_started", outcome = "progress", account_id = %account.id);
let inserted = sqlx::query(
"INSERT INTO accounts (id, profile, pubkey, contact, status, created_at, created_ip, \
created_ptr, last_seen_at, last_seen_ip, last_seen_ptr) \
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?);",
)
.bind(account.id)
.bind(&account.profile)
.bind(&account.pubkey)
.bind(contact_json)
.bind(&account.status)
.bind(account.created_at)
.bind(&account.created_ip)
.bind(&account.created_ptr)
.bind(account.last_seen_at)
.bind(&account.last_seen_ip)
.bind(&account.last_seen_ptr)
.execute(&database.pool)
.await;
if let Err(error) = inserted {
if is_pubkey_conflict(&error)
&& let Some(existing) = Account::find_by_pubkey(profile, pubkey, database).await?
{
debug!(event = "db_account_create_lost_race", outcome = "advisory", account_id = %existing.id, pubkey_fp = %pubkey_fingerprint(pubkey));
return Ok((existing, false));
}
return Err(error);
}
debug!(event = "db_account_created", outcome = "success", account_id = %account.id, pubkey_fp = %pubkey_fingerprint(pubkey));
Ok((account, true))
}
#[must_use]
pub fn needs_touch(&self, now: i64, ip: Option<&str>) -> bool {
match self.last_seen_at {
None => true,
Some(last) => {
now.saturating_sub(last) >= ACCOUNT_TOUCH_INTERVAL
|| self.last_seen_ip.as_deref() != ip
}
}
}
#[tracing::instrument(name = "Account::touch", skip(self, client, database), fields(account_id = %self.id))]
pub async fn touch(
&mut self,
client: &ClientContext,
database: &Database,
) -> Result<(), sqlx::Error> {
let now = now_secs();
sqlx::query(
"UPDATE accounts SET last_seen_at = ?, last_seen_ip = ?, last_seen_ptr = ? \
WHERE id = ?;",
)
.bind(now)
.bind(&client.ip)
.bind(&client.ptr)
.bind(self.id)
.execute(&database.pool)
.await?;
self.last_seen_at = Some(now);
self.last_seen_ip = client.ip.clone();
self.last_seen_ptr = client.ptr.clone();
debug!(event = "db_account_touched", outcome = "success", account_id = %self.id);
Ok(())
}
#[tracing::instrument(name = "Account::update_contact", skip(self, database), fields(account_id = %self.id))]
pub async fn update_contact(
&mut self,
contact: Vec<String>,
database: &Database,
) -> Result<(), sqlx::Error> {
debug!(event = "db_account_contact_update_started", outcome = "progress", account_id = %self.id);
let contact_json = Value::from(contact.clone()).to_string();
sqlx::query("UPDATE accounts SET contact = ? WHERE id = ?;")
.bind(contact_json)
.bind(self.id)
.execute(&database.pool)
.await?;
self.contact = contact;
debug!(event = "db_account_contact_updated", outcome = "success", account_id = %self.id);
Ok(())
}
#[tracing::instrument(name = "Account::deactivate", skip(self, database), fields(account_id = %self.id))]
pub async fn deactivate(&mut self, database: &Database) -> Result<(), sqlx::Error> {
debug!(event = "db_account_deactivation_started", outcome = "progress", account_id = %self.id);
sqlx::query("UPDATE accounts SET status = 'deactivated' WHERE id = ?;")
.bind(self.id)
.execute(&database.pool)
.await?;
self.status = "deactivated".to_string();
debug!(event = "db_account_deactivated", outcome = "success", account_id = %self.id);
Ok(())
}
pub async fn update_pubkey(
&mut self,
pubkey: &[u8],
database: &Database,
) -> Result<(), sqlx::Error> {
debug!(event = "db_account_pubkey_update_started", outcome = "progress", account_id = ?self.id);
sqlx::query("UPDATE accounts SET pubkey = ? WHERE id = ?;")
.bind(pubkey)
.bind(self.id)
.execute(&database.pool)
.await?;
self.pubkey = pubkey.to_vec();
info!(event = "db_account_pubkey_updated", outcome = "success", account_id = ?self.id);
Ok(())
}
pub async fn set_eab_kid(
&mut self,
eab_kid: Uuid,
database: &Database,
) -> Result<(), sqlx::Error> {
debug!(event = "db_account_eab_kid_set_started", outcome = "progress", account_id = ?self.id, eab_kid = ?eab_kid);
sqlx::query("UPDATE accounts SET eab_kid = ? WHERE id = ?;")
.bind(eab_kid)
.bind(self.id)
.execute(&database.pool)
.await?;
self.eab_kid = Some(eab_kid);
info!(event = "db_account_eab_kid_set", outcome = "success", account_id = ?self.id, eab_kid = ?eab_kid);
Ok(())
}
pub async fn set_terms_agreed(&mut self, database: &Database) -> Result<(), sqlx::Error> {
debug!(event = "db_account_terms_agreed_started", outcome = "progress", account_id = ?self.id);
sqlx::query("UPDATE accounts SET terms_of_service_agreed = 1 WHERE id = ?;")
.bind(self.id)
.execute(&database.pool)
.await?;
self.terms_of_service_agreed = Some(true);
info!(event = "db_account_terms_agreed", outcome = "success", account_id = ?self.id);
Ok(())
}
pub async fn find_any_by_id(
id: &str,
database: &Database,
) -> Result<Option<Account>, sqlx::Error> {
debug!(event = "db_account_find_any_by_id_started", outcome = "progress", account_id = %id);
let Some(id) = crate::sqlite::id::parse(id) else {
return Ok(None);
};
let row = sqlx::query(concat!(
"SELECT ",
columns!(),
" FROM accounts WHERE id = ?;"
))
.bind(id)
.fetch_optional(&database.pool)
.await?;
row.map(Account::from_row).transpose()
}
pub async fn search(
profile: Option<&str>,
limit: i64,
offset: i64,
database: &Database,
) -> Result<(Vec<Account>, i64), sqlx::Error> {
debug!(event = "db_account_search_started", outcome = "progress", profile = ?profile, limit = limit, offset = offset);
let (rows, total) = match profile {
Some(profile) => {
let rows = sqlx::query(concat!(
"SELECT ",
columns!(),
" FROM accounts WHERE profile = ? \
ORDER BY created_at DESC, id DESC LIMIT ? OFFSET ?;"
))
.bind(profile)
.bind(limit)
.bind(offset)
.fetch_all(&database.pool)
.await?;
let total: i64 = sqlx::query("SELECT COUNT(*) FROM accounts WHERE profile = ?;")
.bind(profile)
.fetch_one(&database.pool)
.await?
.try_get(0)?;
(rows, total)
}
None => {
let rows = sqlx::query(concat!(
"SELECT ",
columns!(),
" FROM accounts ORDER BY created_at DESC, id DESC LIMIT ? OFFSET ?;"
))
.bind(limit)
.bind(offset)
.fetch_all(&database.pool)
.await?;
let total: i64 = sqlx::query("SELECT COUNT(*) FROM accounts;")
.fetch_one(&database.pool)
.await?
.try_get(0)?;
(rows, total)
}
};
let accounts = rows
.into_iter()
.map(Account::from_row)
.collect::<Result<_, _>>()?;
Ok((accounts, total))
}
pub async fn delete(id: &str, database: &Database) -> Result<bool, sqlx::Error> {
debug!(event = "db_account_delete_started", outcome = "progress", account_id = ?id);
let Some(id) = crate::sqlite::id::parse(id) else {
return Ok(false);
};
let result = sqlx::query("DELETE FROM accounts WHERE id = ?;")
.bind(id)
.execute(&database.pool)
.await?;
let deleted = result.rows_affected() > 0;
if deleted {
info!(event = "db_account_deleted", outcome = "success", account_id = ?id);
} else {
debug!(event = "db_account_delete_missing", outcome = "success", account_id = ?id);
}
Ok(deleted)
}
#[must_use]
pub fn to_json(&self, base_url: &str) -> Value {
let mut object = serde_json::Map::new();
object.insert("status".to_string(), Value::String(self.status.clone()));
if !self.contact.is_empty() {
object.insert("contact".to_string(), Value::from(self.contact.clone()));
}
object.insert(
"orders".to_string(),
Value::String(format!("{base_url}/acct/{}/orders", self.id)),
);
if let Some(agreed) = self.terms_of_service_agreed {
object.insert("termsOfServiceAgreed".to_string(), Value::Bool(agreed));
}
Value::Object(object)
}
}
pub(crate) fn is_pubkey_conflict(error: &sqlx::Error) -> bool {
matches!(error, sqlx::Error::Database(db) if db.is_unique_violation()
&& db.message().contains("accounts.pubkey"))
}
#[cfg(test)]
mod tests {
#[test]
fn needs_touch_yields_to_the_interval_but_never_to_a_changed_address() {
let mut account = Account {
id: crate::sqlite::id::mint(),
profile: "default".to_string(),
pubkey: vec![1],
contact: vec![],
status: "valid".to_string(),
created_at: 0,
eab_kid: None,
terms_of_service_agreed: None,
created_ip: None,
created_ptr: None,
last_seen_at: None,
last_seen_ip: None,
last_seen_ptr: None,
};
assert!(account.needs_touch(1_000, Some("203.0.113.7")));
account.last_seen_at = Some(1_000);
account.last_seen_ip = Some("203.0.113.7".to_string());
assert!(!account.needs_touch(1_000, Some("203.0.113.7")));
assert!(!account.needs_touch(1_000 + ACCOUNT_TOUCH_INTERVAL - 1, Some("203.0.113.7")));
assert!(account.needs_touch(1_000 + ACCOUNT_TOUCH_INTERVAL, Some("203.0.113.7")));
assert!(account.needs_touch(1_000, Some("198.51.100.4")));
assert!(account.needs_touch(1_000, None));
assert!(!account.needs_touch(0, Some("203.0.113.7")));
}
#[tokio::test]
async fn creation_stamps_the_address_once_and_touch_moves_only_the_last_seen_columns() {
let db = Database::connect_in_memory().await.unwrap();
let first = ClientContext {
ip: Some("203.0.113.7".to_string()),
ptr: Some("first.example.com".to_string()),
user_agent: Some("certbot".to_string()),
request_id: Some("req-1".to_string()),
};
let (created, is_new) = Account::find_or_create("default", &[42u8], vec![], &first, &db)
.await
.unwrap();
assert!(is_new);
assert_eq!(created.created_ip.as_deref(), Some("203.0.113.7"));
assert_eq!(created.created_ptr.as_deref(), Some("first.example.com"));
assert_eq!(created.last_seen_ip.as_deref(), Some("203.0.113.7"));
assert!(created.last_seen_at.is_some());
let second = ClientContext {
ip: Some("198.51.100.4".to_string()),
ptr: Some("second.example.com".to_string()),
..ClientContext::default()
};
let (mut found, is_new) = Account::find_or_create("default", &[42u8], vec![], &second, &db)
.await
.unwrap();
assert!(!is_new);
assert_eq!(found.created_ip.as_deref(), Some("203.0.113.7"));
assert_eq!(found.created_ptr.as_deref(), Some("first.example.com"));
found.touch(&second, &db).await.unwrap();
assert_eq!(found.last_seen_ip.as_deref(), Some("198.51.100.4"));
assert_eq!(found.last_seen_ptr.as_deref(), Some("second.example.com"));
let reloaded = Account::find_by_id("default", found.id.to_string().as_str(), &db)
.await
.unwrap()
.unwrap();
assert_eq!(reloaded.created_ip.as_deref(), Some("203.0.113.7"));
assert_eq!(reloaded.last_seen_ip.as_deref(), Some("198.51.100.4"));
assert_eq!(
reloaded.last_seen_ptr.as_deref(),
Some("second.example.com")
);
assert!(reloaded.last_seen_at >= reloaded.created_at.into());
let nameless = ClientContext {
ip: Some("198.51.100.4".to_string()),
..ClientContext::default()
};
found.touch(&nameless, &db).await.unwrap();
let reloaded = Account::find_by_id("default", found.id.to_string().as_str(), &db)
.await
.unwrap()
.unwrap();
assert_eq!(reloaded.last_seen_ptr, None);
}
#[tokio::test]
async fn to_json_exposes_none_of_the_traceability_columns() {
let db = Database::connect_in_memory().await.unwrap();
let client = ClientContext {
ip: Some("203.0.113.7".to_string()),
ptr: Some("host.example.com".to_string()),
..ClientContext::default()
};
let (account, _) = Account::find_or_create("default", &[7u8], vec![], &client, &db)
.await
.unwrap();
let json = account.to_json("http://localhost:3000");
let object = json.as_object().unwrap();
for absent in [
"createdIp",
"created_ip",
"createdPtr",
"lastSeenAt",
"lastSeenIp",
"lastSeenPtr",
] {
assert!(!object.contains_key(absent), "{absent} leaked into to_json");
}
assert!(!json.to_string().contains("203.0.113.7"));
}
use super::*;
use std::sync::Arc;
#[tokio::test]
async fn find_or_create_creates_then_returns_existing() {
let db = Arc::new(Database::connect_in_memory().await.unwrap());
let pubkey = vec![1u8, 2, 3, 4];
let contact = vec!["mailto:a@example.com".to_string()];
let (created, is_new) = Account::find_or_create(
"default",
&pubkey,
contact.clone(),
&ClientContext::default(),
&db,
)
.await
.unwrap();
assert!(is_new);
assert_eq!(created.status, "valid");
assert_eq!(created.contact, contact);
let (existing, is_new) =
Account::find_or_create("default", &pubkey, vec![], &ClientContext::default(), &db)
.await
.unwrap();
assert!(!is_new);
assert_eq!(existing.id, created.id);
assert_eq!(existing.contact, contact);
}
#[tokio::test]
async fn find_by_id_and_pubkey_round_trip() {
let db = Arc::new(Database::connect_in_memory().await.unwrap());
let pubkey = vec![9u8; 16];
let (account, _) =
Account::find_or_create("default", &pubkey, vec![], &ClientContext::default(), &db)
.await
.unwrap();
let by_id = Account::find_by_id("default", account.id.to_string().as_str(), &db)
.await
.unwrap()
.unwrap();
assert_eq!(by_id.pubkey, pubkey);
let by_key = Account::find_by_pubkey("default", &pubkey, &db)
.await
.unwrap()
.unwrap();
assert_eq!(by_key.id, account.id);
}
#[tokio::test]
async fn absent_lookups_return_none() {
let db = Arc::new(Database::connect_in_memory().await.unwrap());
assert!(
Account::find_by_id("default", "nope", &db)
.await
.unwrap()
.is_none()
);
assert!(
Account::find_by_pubkey("default", &[0u8; 4], &db)
.await
.unwrap()
.is_none()
);
}
#[tokio::test]
async fn update_contact_persists_and_syncs() {
let db = Arc::new(Database::connect_in_memory().await.unwrap());
let pubkey = vec![7u8; 8];
let (mut account, _) = Account::find_or_create(
"default",
&pubkey,
vec!["mailto:old@example.com".to_string()],
&ClientContext::default(),
&db,
)
.await
.unwrap();
let new_contact = vec!["mailto:new@example.com".to_string()];
account
.update_contact(new_contact.clone(), &db)
.await
.unwrap();
assert_eq!(account.contact, new_contact);
let reloaded = Account::find_by_id("default", account.id.to_string().as_str(), &db)
.await
.unwrap()
.unwrap();
assert_eq!(reloaded.contact, new_contact);
}
#[tokio::test]
async fn deactivate_persists_and_syncs() {
let db = Arc::new(Database::connect_in_memory().await.unwrap());
let pubkey = vec![8u8; 8];
let (mut account, _) =
Account::find_or_create("default", &pubkey, vec![], &ClientContext::default(), &db)
.await
.unwrap();
assert_eq!(account.status, "valid");
account.deactivate(&db).await.unwrap();
assert_eq!(account.status, "deactivated");
let reloaded = Account::find_by_id("default", account.id.to_string().as_str(), &db)
.await
.unwrap()
.unwrap();
assert_eq!(reloaded.status, "deactivated");
}
#[tokio::test]
async fn update_pubkey_persists_and_syncs() {
let db = Arc::new(Database::connect_in_memory().await.unwrap());
let (mut account, _) =
Account::find_or_create("default", &[9u8; 8], vec![], &ClientContext::default(), &db)
.await
.unwrap();
let new_pubkey = vec![10u8; 8];
account.update_pubkey(&new_pubkey, &db).await.unwrap();
assert_eq!(account.pubkey, new_pubkey);
let reloaded = Account::find_by_id("default", account.id.to_string().as_str(), &db)
.await
.unwrap()
.unwrap();
assert_eq!(reloaded.pubkey, new_pubkey);
assert!(
Account::find_by_pubkey("default", &new_pubkey, &db)
.await
.unwrap()
.is_some()
);
}
#[tokio::test]
async fn update_pubkey_to_a_key_owned_by_another_account_is_rejected() {
let db = Arc::new(Database::connect_in_memory().await.unwrap());
let (_first, _) = Account::find_or_create(
"default",
&[11u8; 8],
vec![],
&ClientContext::default(),
&db,
)
.await
.unwrap();
let (mut second, _) = Account::find_or_create(
"default",
&[12u8; 8],
vec![],
&ClientContext::default(),
&db,
)
.await
.unwrap();
let error = second
.update_pubkey(&[11u8; 8], &db)
.await
.expect_err("taking another account's key must not succeed");
assert!(
is_pubkey_conflict(&error),
"the unique violation must be recognisable as a pubkey conflict: {error}"
);
}
#[tokio::test]
async fn set_eab_kid_persists_and_syncs() {
let db = Arc::new(Database::connect_in_memory().await.unwrap());
let kid = crate::sqlite::id::mint();
let (mut account, _) =
Account::find_or_create("default", &[5u8], vec![], &ClientContext::default(), &db)
.await
.unwrap();
assert!(account.eab_kid.is_none());
account.set_eab_kid(kid, &db).await.unwrap();
assert_eq!(account.eab_kid, Some(kid));
let reloaded = Account::find_by_id("default", account.id.to_string().as_str(), &db)
.await
.unwrap()
.unwrap();
assert_eq!(reloaded.eab_kid, Some(kid));
}
#[tokio::test]
async fn delete_removes_the_row_and_reports_true() {
let db = Arc::new(Database::connect_in_memory().await.unwrap());
let (account, _) =
Account::find_or_create("default", &[3u8], vec![], &ClientContext::default(), &db)
.await
.unwrap();
assert!(
Account::delete(account.id.to_string().as_str(), &db)
.await
.unwrap()
);
assert!(
Account::find_by_id("default", account.id.to_string().as_str(), &db)
.await
.unwrap()
.is_none()
);
}
#[tokio::test]
async fn delete_of_unknown_id_reports_false() {
let db = Arc::new(Database::connect_in_memory().await.unwrap());
assert!(!Account::delete("nope", &db).await.unwrap());
}
#[tokio::test]
async fn delete_cascades_to_the_accounts_orders() {
let db = Arc::new(Database::connect_in_memory().await.unwrap());
let (account, _) =
Account::find_or_create("default", &[4u8], vec![], &ClientContext::default(), &db)
.await
.unwrap();
crate::sqlite::order::Order::create(
"default",
account.id,
vec![],
now_secs() + 3600,
None,
None,
&db,
)
.await
.unwrap();
Account::delete(account.id.to_string().as_str(), &db)
.await
.unwrap();
let remaining = crate::sqlite::order::Order::find_by_account(account.id, &db)
.await
.unwrap();
assert!(remaining.is_empty());
}
async fn seed_accounts(db: &Arc<Database>, profile: &str, count: usize) -> Vec<String> {
let base = now_secs();
let mut ids = Vec::new();
for index in 0..count {
let (account, _) = Account::find_or_create(
profile,
&[profile.len() as u8, index as u8],
vec![],
&ClientContext::default(),
db,
)
.await
.unwrap();
sqlx::query("UPDATE accounts SET created_at = ? WHERE id = ?;")
.bind(base - index as i64)
.bind(account.id)
.execute(&db.pool)
.await
.unwrap();
ids.push(account.id);
}
ids.into_iter().map(|v| v.to_string()).collect()
}
#[tokio::test]
async fn search_pages_newest_first_and_reports_the_unpaged_total() {
let db = Arc::new(Database::connect_in_memory().await.unwrap());
let ids = seed_accounts(&db, "default", 5).await;
let (page, total) = Account::search(None, 2, 0, &db).await.unwrap();
assert_eq!(total, 5, "the total must ignore the page window");
assert_eq!(
page.iter().map(|a| a.id.to_string()).collect::<Vec<_>>(),
ids[..2]
);
let (second, _) = Account::search(None, 2, 2, &db).await.unwrap();
assert_eq!(
second.iter().map(|a| a.id.to_string()).collect::<Vec<_>>(),
ids[2..4]
);
let (beyond, total) = Account::search(None, 2, 99, &db).await.unwrap();
assert!(beyond.is_empty());
assert_eq!(total, 5);
}
#[tokio::test]
async fn search_scopes_by_profile_and_counts_only_that_profile() {
let db = Arc::new(Database::connect_in_memory().await.unwrap());
seed_accounts(&db, "default", 2).await;
seed_accounts(&db, "other", 3).await;
let (rows, total) = Account::search(Some("other"), 50, 0, &db).await.unwrap();
assert_eq!(total, 3);
assert_eq!(rows.len(), 3);
assert!(rows.iter().all(|a| a.profile == "other"));
let (_, total) = Account::search(None, 50, 0, &db).await.unwrap();
assert_eq!(total, 5, "no profile means every endpoint");
let (rows, total) = Account::search(Some("nope"), 50, 0, &db).await.unwrap();
assert!(rows.is_empty());
assert_eq!(total, 0);
}
#[tokio::test]
async fn search_on_an_empty_table_is_empty_rather_than_an_error() {
let db = Arc::new(Database::connect_in_memory().await.unwrap());
let (rows, total) = Account::search(None, 50, 0, &db).await.unwrap();
assert!(rows.is_empty());
assert_eq!(total, 0);
}
#[tokio::test]
async fn concurrent_find_or_create_for_one_key_yields_one_account() {
let file =
std::env::temp_dir().join(format!("acme-proxy-test-{}.db", uuid::Uuid::now_v7()));
let url = format!("sqlite://{}", file.display());
let db = Arc::new(Database::connect(&url).await.unwrap());
const RACERS: usize = 8;
let barrier = Arc::new(tokio::sync::Barrier::new(RACERS));
let mut tasks = Vec::with_capacity(RACERS);
for _ in 0..RACERS {
let db = db.clone();
let barrier = barrier.clone();
tasks.push(tokio::spawn(async move {
barrier.wait().await;
Account::find_or_create(
"default",
&[7u8; 32],
vec![],
&ClientContext::default(),
&db,
)
.await
}));
}
let mut ids = Vec::with_capacity(RACERS);
let mut created = 0;
for task in tasks {
let (account, is_new) = task
.await
.unwrap()
.expect("losing the insert race is not an error");
if is_new {
created += 1;
}
ids.push(account.id);
}
assert_eq!(created, 1, "exactly one caller may create the account");
assert!(
ids.windows(2).all(|pair| pair[0] == pair[1]),
"every caller must be handed the same account: {ids:?}"
);
db.pool.close().await;
for suffix in ["", "-wal", "-shm"] {
let _ = std::fs::remove_file(format!("{}{suffix}", file.display()));
}
}
}