use sqlx::Row;
use sqlx::sqlite::SqliteRow;
use tracing::debug;
use crate::sqlite::db::Database;
use crate::sqlite::nonce::now_secs;
pub(crate) const MAX_PROCESSING_BATCH: usize = 500;
#[derive(Debug, Clone)]
pub struct UpstreamOrder {
pub order_id: String,
pub upstream_order_url: String,
pub upstream_finalize_url: Option<String>,
pub upstream_certificate_url: Option<String>,
pub csr_der: Vec<u8>,
pub status: String,
pub error: Option<String>,
pub created_at: i64,
pub updated_at: i64,
pub client_ip: Option<String>,
pub client_ptr: Option<String>,
pub user_agent: Option<String>,
pub request_id: Option<String>,
}
impl UpstreamOrder {
fn from_row(row: SqliteRow) -> Result<Self, sqlx::Error> {
Ok(UpstreamOrder {
order_id: row.try_get("order_id")?,
upstream_order_url: row.try_get("upstream_order_url")?,
upstream_finalize_url: row.try_get("upstream_finalize_url")?,
upstream_certificate_url: row.try_get("upstream_certificate_url")?,
csr_der: row.try_get("csr_der")?,
status: row.try_get("status")?,
error: row.try_get("error")?,
created_at: row.try_get("created_at")?,
updated_at: row.try_get("updated_at")?,
client_ip: row.try_get("client_ip")?,
client_ptr: row.try_get("client_ptr")?,
user_agent: row.try_get("user_agent")?,
request_id: row.try_get("request_id")?,
})
}
#[must_use]
pub fn client(&self) -> crate::audit::ClientContext {
crate::audit::ClientContext {
ip: self.client_ip.clone(),
ptr: self.client_ptr.clone(),
user_agent: self.user_agent.clone(),
request_id: self.request_id.clone(),
}
}
pub async fn set_client(
order_id: &str,
client: &crate::audit::ClientContext,
database: &Database,
) -> Result<(), sqlx::Error> {
sqlx::query(
"UPDATE upstream_orders \
SET client_ip = ?, client_ptr = ?, user_agent = ?, request_id = ? \
WHERE order_id = ?;",
)
.bind(&client.ip)
.bind(&client.ptr)
.bind(&client.user_agent)
.bind(&client.request_id)
.bind(order_id)
.execute(&database.pool)
.await?;
Ok(())
}
pub async fn create(
order_id: &str,
upstream_order_url: &str,
upstream_finalize_url: Option<&str>,
csr_der: &[u8],
database: &Database,
) -> Result<Option<UpstreamOrder>, sqlx::Error> {
let now = now_secs();
let record = UpstreamOrder {
order_id: order_id.to_string(),
upstream_order_url: upstream_order_url.to_string(),
upstream_finalize_url: upstream_finalize_url.map(str::to_string),
upstream_certificate_url: None,
csr_der: csr_der.to_vec(),
status: "processing".to_string(),
error: None,
created_at: now,
updated_at: now,
client_ip: None,
client_ptr: None,
user_agent: None,
request_id: None,
};
debug!(event = "db_upstream_order_create_started", outcome = "progress", order_id = ?order_id);
let result = sqlx::query(
"INSERT OR IGNORE INTO upstream_orders \
(order_id, upstream_order_url, upstream_finalize_url, csr_der, status, \
created_at, updated_at) \
VALUES (?, ?, ?, ?, ?, ?, ?);",
)
.bind(&record.order_id)
.bind(&record.upstream_order_url)
.bind(&record.upstream_finalize_url)
.bind(&record.csr_der)
.bind(&record.status)
.bind(record.created_at)
.bind(record.updated_at)
.execute(&database.pool)
.await?;
if result.rows_affected() == 0 {
return Ok(None);
}
Ok(Some(record))
}
pub async fn find_by_order_id(
order_id: &str,
database: &Database,
) -> Result<Option<UpstreamOrder>, sqlx::Error> {
let row = sqlx::query("SELECT * FROM upstream_orders WHERE order_id = ?;")
.bind(order_id)
.fetch_optional(&database.pool)
.await?;
row.map(UpstreamOrder::from_row).transpose()
}
pub async fn mark_valid(
order_id: &str,
certificate_url: Option<&str>,
database: &Database,
) -> Result<(), sqlx::Error> {
sqlx::query(
"UPDATE upstream_orders \
SET status = 'valid', upstream_certificate_url = ?, updated_at = ? \
WHERE order_id = ?;",
)
.bind(certificate_url)
.bind(now_secs())
.bind(order_id)
.execute(&database.pool)
.await?;
Ok(())
}
pub async fn mark_invalid(
order_id: &str,
error: &str,
database: &Database,
) -> Result<(), sqlx::Error> {
sqlx::query(
"UPDATE upstream_orders SET status = 'invalid', error = ?, updated_at = ? \
WHERE order_id = ?;",
)
.bind(error)
.bind(now_secs())
.bind(order_id)
.execute(&database.pool)
.await?;
Ok(())
}
pub async fn list_processing(
profiles: &[String],
database: &Database,
) -> Result<Vec<UpstreamOrder>, sqlx::Error> {
if profiles.is_empty() {
return Ok(Vec::new());
}
let placeholders = std::iter::repeat_n("?", profiles.len())
.collect::<Vec<_>>()
.join(", ");
let limit = MAX_PROCESSING_BATCH;
let sql = format!(
"SELECT u.* FROM upstream_orders u \
JOIN orders o ON o.id = u.order_id \
WHERE u.status = 'processing' AND o.profile IN ({placeholders}) \
ORDER BY u.created_at ASC LIMIT {limit};"
);
let mut query = sqlx::query(sqlx::AssertSqlSafe(sql));
for profile in profiles {
query = query.bind(profile);
}
let rows = query.fetch_all(&database.pool).await?;
rows.into_iter().map(UpstreamOrder::from_row).collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::audit::ClientContext;
use crate::sqlite::account::Account;
use crate::sqlite::order::{Identifier, Order};
use std::sync::Arc;
async fn order(database: &Database) -> Order {
let (account, _) = Account::find_or_create(
"default",
uuid::Uuid::new_v4().as_bytes(),
Vec::new(),
&ClientContext::default(),
database,
)
.await
.unwrap();
Order::create(
"default",
&account.id,
vec![Identifier::dns("example.com")],
now_secs() + 3600,
None,
None,
database,
)
.await
.unwrap()
}
async fn database() -> Arc<Database> {
Arc::new(Database::connect_in_memory().await.unwrap())
}
#[tokio::test]
async fn the_finalize_context_is_stored_and_handed_back_for_the_audit_row() {
let database = database().await;
let order = order(&database).await;
UpstreamOrder::create(
&order.id,
"https://up.example/order/1",
None,
b"csr",
&database,
)
.await
.unwrap()
.unwrap();
let mapping = UpstreamOrder::find_by_order_id(&order.id, &database)
.await
.unwrap()
.unwrap();
assert_eq!(mapping.client(), ClientContext::default());
let client = ClientContext {
ip: Some("203.0.113.7".to_string()),
ptr: Some("host.example.com".to_string()),
user_agent: Some("lego".to_string()),
request_id: Some("req-1".to_string()),
};
UpstreamOrder::set_client(&order.id, &client, &database)
.await
.unwrap();
let mapping = UpstreamOrder::find_by_order_id(&order.id, &database)
.await
.unwrap()
.unwrap();
assert_eq!(mapping.client(), client);
UpstreamOrder::set_client("no-such-order", &client, &database)
.await
.unwrap();
}
#[tokio::test]
async fn create_then_find_round_trips() {
let db = database().await;
let order = order(&db).await;
let created = UpstreamOrder::create(
&order.id,
"https://up.example/order/1",
Some("https://up.example/order/1/finalize"),
b"csr-bytes",
&db,
)
.await
.unwrap()
.expect("a first insert must succeed");
assert_eq!(created.status, "processing");
let found = UpstreamOrder::find_by_order_id(&order.id, &db)
.await
.unwrap()
.expect("the row just written must be found");
assert_eq!(found.upstream_order_url, "https://up.example/order/1");
assert_eq!(
found.upstream_finalize_url.as_deref(),
Some("https://up.example/order/1/finalize")
);
assert!(found.upstream_certificate_url.is_none());
}
#[tokio::test]
async fn a_second_create_for_one_order_is_refused() {
let db = database().await;
let order = order(&db).await;
assert!(
UpstreamOrder::create(&order.id, "https://up.example/order/1", None, b"csr", &db)
.await
.unwrap()
.is_some()
);
assert!(
UpstreamOrder::create(&order.id, "https://up.example/order/2", None, b"csr", &db)
.await
.unwrap()
.is_none(),
"a duplicate must be reported as already-in-flight, not inserted"
);
let found = UpstreamOrder::find_by_order_id(&order.id, &db)
.await
.unwrap()
.unwrap();
assert_eq!(found.upstream_order_url, "https://up.example/order/1");
}
#[tokio::test]
async fn mark_valid_records_the_certificate_url() {
let db = database().await;
let order = order(&db).await;
UpstreamOrder::create(&order.id, "https://up.example/order/1", None, b"csr", &db)
.await
.unwrap();
UpstreamOrder::mark_valid(&order.id, Some("https://up.example/cert/1"), &db)
.await
.unwrap();
let found = UpstreamOrder::find_by_order_id(&order.id, &db)
.await
.unwrap()
.unwrap();
assert_eq!(found.status, "valid");
assert_eq!(
found.upstream_certificate_url.as_deref(),
Some("https://up.example/cert/1")
);
}
#[tokio::test]
async fn mark_invalid_records_the_reason() {
let db = database().await;
let order = order(&db).await;
UpstreamOrder::create(&order.id, "https://up.example/order/1", None, b"csr", &db)
.await
.unwrap();
UpstreamOrder::mark_invalid(&order.id, "upstream said no", &db)
.await
.unwrap();
let found = UpstreamOrder::find_by_order_id(&order.id, &db)
.await
.unwrap()
.unwrap();
assert_eq!(found.status, "invalid");
assert_eq!(found.error.as_deref(), Some("upstream said no"));
}
#[tokio::test]
async fn list_processing_skips_settled_rows() {
let db = database().await;
let still_running = order(&db).await;
let finished = order(&db).await;
UpstreamOrder::create(&still_running.id, "https://up.example/a", None, b"csr", &db)
.await
.unwrap();
UpstreamOrder::create(&finished.id, "https://up.example/b", None, b"csr", &db)
.await
.unwrap();
UpstreamOrder::mark_valid(&finished.id, None, &db)
.await
.unwrap();
let processing = UpstreamOrder::list_processing(&["default".to_string()], &db)
.await
.unwrap();
assert_eq!(processing.len(), 1);
assert_eq!(processing[0].order_id, still_running.id);
}
#[tokio::test]
async fn find_by_order_id_is_none_for_an_unknown_order() {
let db = database().await;
assert!(
UpstreamOrder::find_by_order_id("nope", &db)
.await
.unwrap()
.is_none()
);
}
}