use crate::sql::Row;
use tracing::debug;
use uuid::Uuid;
use crate::db::Database;
use crate::nonce::now_secs;
use crate::status::{self, OrderStatus, UpstreamOrderStatus};
use acme_proxy_core::identifier::Identifier;
pub const MAX_PROCESSING_BATCH: usize = 500;
#[derive(Debug, Clone)]
pub struct UpstreamOrder {
pub order_id: Uuid,
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: Row) -> 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) -> acme_proxy_core::audit::ClientContext {
acme_proxy_core::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: &acme_proxy_core::audit::ClientContext,
database: &Database,
) -> Result<(), sqlx::Error> {
let Some(order_id) = crate::id::parse(order_id) else {
return Ok(());
};
crate::sql::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)
.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 Some(order_id) = crate::id::parse(order_id) else {
return Ok(None);
};
let now = now_secs();
let record = UpstreamOrder {
order_id,
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 = crate::sql::query(
"INSERT INTO upstream_orders \
(order_id, upstream_order_url, upstream_finalize_url, csr_der, status, \
created_at, updated_at) \
VALUES (?, ?, ?, ?, ?, ?, ?) \
ON CONFLICT DO NOTHING;",
)
.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)
.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 Some(order_id) = crate::id::parse(order_id) else {
return Ok(None);
};
let row = crate::sql::query("SELECT * FROM upstream_orders WHERE order_id = ?;")
.bind(order_id)
.fetch_optional(database)
.await?;
row.map(UpstreamOrder::from_row).transpose()
}
pub async fn mark_valid(
order_id: &str,
certificate_url: Option<&str>,
database: &Database,
) -> Result<(), sqlx::Error> {
let Some(order_id) = crate::id::parse(order_id) else {
return Ok(());
};
crate::sql::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)
.await?;
Ok(())
}
pub async fn mark_invalid(
order_id: &str,
error: &str,
database: &Database,
) -> Result<(), sqlx::Error> {
let Some(order_id) = crate::id::parse(order_id) else {
return Ok(());
};
crate::sql::query(
"UPDATE upstream_orders SET status = 'invalid', error = ?, updated_at = ? \
WHERE order_id = ?;",
)
.bind(error)
.bind(now_secs())
.bind(order_id)
.execute(database)
.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 = crate::sql::query(sqlx::AssertSqlSafe(sql));
for profile in profiles {
query = query.bind(profile);
}
let rows = query.fetch_all(database).await?;
rows.into_iter().map(UpstreamOrder::from_row).collect()
}
}
#[derive(Debug, Clone)]
pub struct UpstreamOrderRow {
pub order_id: Uuid,
pub upstream_order_url: String,
pub upstream_finalize_url: Option<String>,
pub upstream_certificate_url: Option<String>,
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>,
pub profile: String,
pub account_id: Uuid,
pub identifiers: Vec<Identifier>,
pub local_status: OrderStatus,
pub local_expires: i64,
}
impl UpstreamOrderRow {
const COLUMNS: &'static str = "u.order_id, u.upstream_order_url, u.upstream_finalize_url, \
u.upstream_certificate_url, u.status, u.error, u.created_at, u.updated_at, \
u.client_ip, u.client_ptr, u.user_agent, u.request_id, \
o.profile, o.account_id, o.identifiers, o.status AS local_status, \
o.expires AS local_expires";
fn from_joined_row(row: Row) -> Result<Self, sqlx::Error> {
let identifiers_json: String = row.try_get("identifiers")?;
let identifiers: Vec<Identifier> = serde_json::from_str(&identifiers_json)
.map_err(|e| sqlx::Error::Decode(Box::new(e)))?;
Ok(UpstreamOrderRow {
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")?,
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")?,
profile: row.try_get("profile")?,
account_id: row.try_get("account_id")?,
identifiers,
local_status: status::from_column(row.try_get::<String>("local_status")?.as_str())?,
local_expires: row.try_get("local_expires")?,
})
}
#[must_use]
pub fn client(&self) -> acme_proxy_core::audit::ClientContext {
acme_proxy_core::audit::ClientContext {
ip: self.client_ip.clone(),
ptr: self.client_ptr.clone(),
user_agent: self.user_agent.clone(),
request_id: self.request_id.clone(),
}
}
}
#[derive(Debug, Clone, Default)]
pub struct UpstreamOrderQuery {
pub profile: Option<String>,
pub status: Option<UpstreamOrderStatus>,
pub limit: i64,
pub offset: i64,
}
impl UpstreamOrderQuery {
fn push_predicates(&self, builder: &mut crate::sql::Builder) {
crate::query::push_equalities(
builder,
crate::query::WHERE,
&[
("o.profile = ", self.profile.as_deref()),
("u.status = ", self.status.map(UpstreamOrderStatus::as_str)),
],
);
}
}
impl UpstreamOrder {
pub async fn search(
query: &UpstreamOrderQuery,
database: &Database,
) -> Result<(Vec<UpstreamOrderRow>, i64), sqlx::Error> {
debug!(
event = "db_upstream_order_search_started",
outcome = "progress",
profile = ?query.profile,
status = ?query.status,
limit = query.limit,
offset = query.offset,
);
let mut page = crate::sql::Builder::new(
database.dialect(),
format!(
"SELECT {} FROM upstream_orders u JOIN orders o ON o.id = u.order_id",
UpstreamOrderRow::COLUMNS
),
);
query.push_predicates(&mut page);
page.push(" ORDER BY u.created_at DESC, u.order_id DESC LIMIT ");
page.push_bind(query.limit);
page.push(" OFFSET ");
page.push_bind(query.offset);
let rows = page.build().fetch_all(database).await?;
let items: Vec<UpstreamOrderRow> = rows
.into_iter()
.map(UpstreamOrderRow::from_joined_row)
.collect::<Result<_, _>>()?;
let mut count = crate::sql::Builder::new(
database.dialect(),
"SELECT COUNT(*) FROM upstream_orders u JOIN orders o ON o.id = u.order_id",
);
query.push_predicates(&mut count);
let total: i64 = count.build().fetch_one(database).await?.try_get::<i64>(0)?;
Ok((items, total))
}
pub async fn find_row_by_order_id(
order_id: &str,
database: &Database,
) -> Result<Option<UpstreamOrderRow>, sqlx::Error> {
let Some(order_id) = crate::id::parse(order_id) else {
return Ok(None);
};
let mut query = crate::sql::Builder::new(
database.dialect(),
format!(
"SELECT {} FROM upstream_orders u JOIN orders o ON o.id = u.order_id \
WHERE u.order_id = ",
UpstreamOrderRow::COLUMNS
),
);
query.push_bind(order_id);
let row = query.build().fetch_optional(database).await?;
row.map(UpstreamOrderRow::from_joined_row).transpose()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::account::Account;
use crate::order::Order;
use acme_proxy_core::audit::ClientContext;
use acme_proxy_core::identifier::Identifier;
use std::sync::Arc;
async fn order(database: &Database) -> Order {
let (account, _) = Account::find_or_create(
"default",
&acme_proxy_core::random::random_bytes::<16>(),
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_for_test().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.to_string().as_str(),
"https://up.example/order/1",
None,
b"csr",
&database,
)
.await
.unwrap()
.unwrap();
let mapping = UpstreamOrder::find_by_order_id(order.id.to_string().as_str(), &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.to_string().as_str(), &client, &database)
.await
.unwrap();
let mapping = UpstreamOrder::find_by_order_id(order.id.to_string().as_str(), &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.to_string().as_str(),
"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.to_string().as_str(), &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.to_string().as_str(),
"https://up.example/order/1",
None,
b"csr",
&db
)
.await
.unwrap()
.is_some()
);
assert!(
UpstreamOrder::create(
order.id.to_string().as_str(),
"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.to_string().as_str(), &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.to_string().as_str(),
"https://up.example/order/1",
None,
b"csr",
&db,
)
.await
.unwrap();
UpstreamOrder::mark_valid(
order.id.to_string().as_str(),
Some("https://up.example/cert/1"),
&db,
)
.await
.unwrap();
let found = UpstreamOrder::find_by_order_id(order.id.to_string().as_str(), &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.to_string().as_str(),
"https://up.example/order/1",
None,
b"csr",
&db,
)
.await
.unwrap();
UpstreamOrder::mark_invalid(order.id.to_string().as_str(), "upstream said no", &db)
.await
.unwrap();
let found = UpstreamOrder::find_by_order_id(order.id.to_string().as_str(), &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.to_string().as_str(),
"https://up.example/a",
None,
b"csr",
&db,
)
.await
.unwrap();
UpstreamOrder::create(
finished.id.to_string().as_str(),
"https://up.example/b",
None,
b"csr",
&db,
)
.await
.unwrap();
UpstreamOrder::mark_valid(finished.id.to_string().as_str(), 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()
);
}
use crate::status::{OrderStatus, UpstreamOrderStatus};
async fn order_on(profile: &str, names: &[&str], database: &Database) -> Order {
let (account, _) = Account::find_or_create(
profile,
&acme_proxy_core::random::random_bytes::<16>(),
Vec::new(),
&ClientContext::default(),
database,
)
.await
.unwrap();
Order::create(
profile,
account.id,
names.iter().map(|n| Identifier::dns(*n)).collect(),
now_secs() + 3600,
None,
None,
database,
)
.await
.unwrap()
}
#[tokio::test]
async fn search_joins_the_local_order_for_identifiers_and_status() {
let db = database().await;
let ord = order_on("default", &["a.example.com", "b.example.com"], &db).await;
UpstreamOrder::create(
ord.id.to_string().as_str(),
"https://up.example/order/1",
Some("https://up.example/order/1/finalize"),
b"secret-csr-bytes",
&db,
)
.await
.unwrap();
let (rows, total) = UpstreamOrder::search(
&UpstreamOrderQuery {
limit: 50,
..UpstreamOrderQuery::default()
},
&db,
)
.await
.unwrap();
assert_eq!(total, 1);
let row = &rows[0];
assert_eq!(row.order_id, ord.id);
assert_eq!(row.profile, "default");
assert_eq!(row.local_status, OrderStatus::Pending);
let names: Vec<&str> = row.identifiers.iter().map(|i| i.value.as_str()).collect();
assert_eq!(names, vec!["a.example.com", "b.example.com"]);
}
#[tokio::test]
async fn search_filters_profile_and_status_together_and_pages() {
let db = database().await;
let a = order_on("default", &["a.example.com"], &db).await;
let b = order_on("default", &["b.example.com"], &db).await;
let c = order_on("other", &["c.example.com"], &db).await;
for ord in [&a, &b, &c] {
UpstreamOrder::create(
ord.id.to_string().as_str(),
"https://up.example/o",
None,
b"csr",
&db,
)
.await
.unwrap();
}
UpstreamOrder::mark_invalid(b.id.to_string().as_str(), "upstream said no", &db)
.await
.unwrap();
let (rows, total) = UpstreamOrder::search(
&UpstreamOrderQuery {
profile: Some("default".to_string()),
limit: 50,
..UpstreamOrderQuery::default()
},
&db,
)
.await
.unwrap();
assert_eq!(total, 2);
assert_eq!(rows.len(), 2);
let (rows, total) = UpstreamOrder::search(
&UpstreamOrderQuery {
profile: Some("default".to_string()),
status: Some(UpstreamOrderStatus::Invalid),
limit: 50,
offset: 0,
},
&db,
)
.await
.unwrap();
assert_eq!(total, 1);
assert_eq!(rows[0].order_id, b.id);
assert_eq!(rows[0].error.as_deref(), Some("upstream said no"));
let (rows, total) = UpstreamOrder::search(
&UpstreamOrderQuery {
profile: Some("' OR 1=1 --".to_string()),
limit: 50,
..UpstreamOrderQuery::default()
},
&db,
)
.await
.unwrap();
assert_eq!(total, 0);
assert!(rows.is_empty());
let mut seen = std::collections::BTreeSet::new();
for offset in 0..3 {
let (rows, total) = UpstreamOrder::search(
&UpstreamOrderQuery {
limit: 1,
offset,
..UpstreamOrderQuery::default()
},
&db,
)
.await
.unwrap();
assert_eq!(total, 3);
assert!(seen.insert(rows[0].order_id));
}
assert_eq!(seen.len(), 3);
}
#[tokio::test]
async fn find_row_by_order_id_is_none_for_junk_and_for_an_unrelayed_order() {
let db = database().await;
assert!(
UpstreamOrder::find_row_by_order_id("nope", &db)
.await
.unwrap()
.is_none()
);
let ord = order_on("default", &["a.example.com"], &db).await;
assert!(
UpstreamOrder::find_row_by_order_id(ord.id.to_string().as_str(), &db)
.await
.unwrap()
.is_none()
);
}
}