mod delivery;
mod outbox;
mod sign;
pub use delivery::{DeliveryWorker, Projection, PushSweepReport};
pub use outbox::{Destination, OPERATOR_PREFIX, Outbox, RunCompleted, is_operator_id};
pub use sign::{
BodySigning, DEFAULT_TOLERANCE, HEADER_A2A_TOKEN, HEADER_ID, HEADER_SIGNATURE,
HEADER_TIMESTAMP, MIN_KEY_BYTES, SCHEME, SigningKeyError, VerifiedDelivery, WebhookRejected,
WebhookVerifier,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PushNamespace {
Operator,
Caller,
}
impl PushNamespace {
#[must_use]
pub fn owns_id(self, id: &str) -> bool {
match self {
Self::Operator => is_operator_id(id),
Self::Caller => !is_operator_id(id),
}
}
}
#[derive(Debug, Clone, Default)]
pub struct DueBatch {
pub rows: Vec<PushRegistration>,
pub unserved: usize,
}
use std::fmt::Debug;
use async_trait::async_trait;
use crate::core::{RunId, Secret, Seq, StoreError};
#[derive(Debug, Clone)]
pub struct PushConfig {
pub id: String,
pub task: RunId,
pub url: String,
pub token: Option<Secret>,
pub authentication: Option<PushAuthentication>,
}
#[derive(Clone)]
pub struct PushAuthentication {
pub scheme: String,
pub credentials: Secret,
}
impl Debug for PushAuthentication {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PushAuthentication")
.field("scheme", &self.scheme)
.field("credentials", &"<redacted>")
.finish()
}
}
impl PushAuthentication {
pub fn validate(&self) -> Result<(), PushError> {
if self.scheme.is_empty()
|| !self.scheme.bytes().all(|byte| {
byte.is_ascii_alphanumeric()
|| matches!(
byte,
b'!' | b'#'
| b'$'
| b'%'
| b'&'
| b'\''
| b'*'
| b'+'
| b'-'
| b'.'
| b'^'
| b'_'
| b'`'
| b'|'
| b'~'
)
})
{
return Err(PushError::Malformed(
"authentication.scheme must be an HTTP authentication token".to_owned(),
));
}
let value = format!("{} {}", self.scheme, self.credentials.expose());
reqwest::header::HeaderValue::from_str(&value)
.map(|_| ())
.map_err(|error| PushError::Malformed(format!("invalid authentication: {error}")))
}
}
#[derive(Debug, Clone)]
pub struct PushRegistration {
pub config: PushConfig,
pub next_seq: Seq,
pub attempts: u32,
pub next_attempt_at: u64,
pub last_error: Option<String>,
}
impl PushConfig {
#[must_use]
pub fn redacted(&self) -> serde_json::Value {
serde_json::json!({
"id": self.id,
"taskId": self.task.to_string(),
"url": self.url,
"authentication": self.authentication.as_ref().map(|auth| serde_json::json!({
"scheme": auth.scheme,
})),
})
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum PushError {
#[error(
"a webhook URL must be https — the payload describes a task, and \
sending it in clear to an address the recipient chose is a disclosure"
)]
NotHttps,
#[error("this deployment does not permit webhooks to '{0}'")]
HostNotGranted(String),
#[error("the webhook URL is not a URL: {0}")]
Malformed(String),
#[error("'{0}'")]
Unroutable(String),
}
impl PushError {
#[must_use]
pub const fn is_permanent(&self) -> bool {
matches!(
self,
Self::NotHttps | Self::HostNotGranted(_) | Self::Malformed(_)
)
}
}
#[async_trait]
pub trait PushStore: Send + Sync + Debug {
fn tenant(&self) -> &str {
crate::core::TenantId::DEFAULT
}
async fn put(&self, config: &PushConfig, next_seq: Seq) -> Result<(), StoreError>;
async fn get(&self, task: RunId, id: &str) -> Result<Option<PushConfig>, StoreError>;
async fn list(&self, task: RunId) -> Result<Vec<PushConfig>, StoreError>;
async fn due(&self, at: u64, limit: usize) -> Result<Vec<PushRegistration>, StoreError>;
async fn due_in(
&self,
at: u64,
limit: usize,
namespace: PushNamespace,
) -> Result<DueBatch, StoreError> {
let mut window = limit.max(1);
loop {
let all = self.due(at, window).await?;
let exhausted = all.len() < window;
let mut batch = DueBatch::default();
for registration in all {
if namespace.owns_id(®istration.config.id) {
batch.rows.push(registration);
} else {
batch.unserved = batch.unserved.saturating_add(1);
}
}
if batch.rows.len() >= limit || exhausted {
batch.rows.truncate(limit);
return Ok(batch);
}
window = window.saturating_mul(2);
}
}
async fn advance(&self, task: RunId, id: &str, next_seq: Seq) -> Result<(), StoreError>;
async fn retry(
&self,
task: RunId,
id: &str,
next_attempt_at: u64,
error: &str,
) -> Result<(), StoreError>;
async fn park(&self, task: RunId, id: &str, error: &str) -> Result<(), StoreError>;
async fn parked(&self, limit: usize) -> Result<Vec<PushRegistration>, StoreError>;
async fn unpark(&self, task: RunId, id: &str, at: u64) -> Result<bool, StoreError>;
async fn delete(&self, task: RunId, id: &str) -> Result<(), StoreError>;
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct PushPolicy {
hosts: std::collections::BTreeSet<String>,
}
impl PushPolicy {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn allow_host(mut self, host: impl AsRef<str>) -> Self {
let raw = host.as_ref();
let host = crate::netguard::canonical_host(raw).unwrap_or_else(|| {
panic!(
"push host grant '{raw}' is not a host a URL can name — give an \
internationalised host in the form the URL parser accepts, or it \
would silently never match a webhook"
)
});
self.hosts.insert(host);
self
}
pub fn check(&self, url: &str) -> Result<(), PushError> {
self.check_allowing_loopback(url, false)
}
fn check_allowing_loopback(&self, url: &str, allow_loopback: bool) -> Result<(), PushError> {
let parsed = reqwest::Url::parse(url).map_err(|e| PushError::Malformed(e.to_string()))?;
let host = parsed
.host_str()
.ok_or_else(|| PushError::Malformed("no host".to_owned()))?
.trim_end_matches('.')
.to_ascii_lowercase();
if parsed.scheme() != "https"
&& !(allow_loopback && crate::netguard::is_loopback_name(&host))
{
return Err(PushError::NotHttps);
}
if !self.hosts.contains(&host) {
return Err(PushError::HostNotGranted(host));
}
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct PushSender {
policy: PushPolicy,
timeout: std::time::Duration,
http: std::sync::OnceLock<reqwest::Client>,
#[cfg(feature = "testkit")]
plaintext_loopback: bool,
operator: bool,
signing: std::collections::BTreeMap<String, BodySigning>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PushMessage {
pub id: String,
pub content_type: String,
pub payload: serde_json::Value,
}
impl PushMessage {
#[must_use]
pub fn json(id: impl Into<String>, payload: serde_json::Value) -> Self {
Self {
id: id.into(),
content_type: "application/json".to_owned(),
payload,
}
}
#[must_use]
pub fn cloudevent(event: &crate::core::CloudEvent) -> Self {
Self {
id: event.id().to_owned(),
content_type: crate::core::CLOUDEVENT_CONTENT_TYPE.to_owned(),
payload: event.to_value(),
}
}
#[must_use]
pub fn typed(mut self, content_type: impl Into<String>) -> Self {
self.content_type = content_type.into();
self
}
}
#[async_trait]
pub trait PushTransport: Send + Sync + Debug {
fn validate(&self, config: &PushConfig) -> Result<(), PushError>;
async fn deliver(
&self,
config: &PushConfig,
message: &PushMessage,
at: u64,
) -> Result<Delivered, PushError>;
}
impl PushSender {
pub const DEFAULT_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(15);
#[must_use]
pub fn new(policy: PushPolicy) -> Self {
Self {
policy,
timeout: Self::DEFAULT_TIMEOUT,
http: std::sync::OnceLock::new(),
#[cfg(feature = "testkit")]
plaintext_loopback: false,
operator: false,
signing: std::collections::BTreeMap::new(),
}
}
#[must_use]
pub fn for_operator_destinations(destinations: &[Destination]) -> Self {
Self {
policy: PushPolicy::new(),
timeout: Self::DEFAULT_TIMEOUT,
http: std::sync::OnceLock::new(),
#[cfg(feature = "testkit")]
plaintext_loopback: false,
operator: true,
signing: destinations
.iter()
.filter_map(|destination| {
destination
.signing
.clone()
.map(|signing| (destination.registration_id(), signing))
})
.collect(),
}
}
#[cfg(feature = "testkit")]
#[must_use]
pub const fn allow_plaintext_loopback(mut self) -> Self {
self.plaintext_loopback = true;
self
}
#[allow(clippy::unused_self)]
const fn loopback_allowed(&self) -> bool {
#[cfg(feature = "testkit")]
{
self.plaintext_loopback
}
#[cfg(not(feature = "testkit"))]
{
false
}
}
#[must_use]
pub const fn timeout(mut self, d: std::time::Duration) -> Self {
self.timeout = d;
self
}
#[must_use]
pub const fn policy(&self) -> &PushPolicy {
&self.policy
}
fn reach(&self) -> crate::netguard::Reach {
if self.operator {
crate::netguard::Reach::Configured
} else if self.loopback_allowed() {
crate::netguard::Reach::PublicOrLoopbackName
} else {
crate::netguard::Reach::Public
}
}
fn http(&self) -> Result<&reqwest::Client, PushError> {
if let Some(client) = self.http.get() {
return Ok(client);
}
let client = crate::netguard::guarded_client(self.reach())
.timeout(self.timeout)
.build()
.map_err(|e| PushError::Unroutable(e.to_string()))?;
Ok(self.http.get_or_init(|| client))
}
async fn deliver_validated(
&self,
config: &PushConfig,
message: &PushMessage,
at: u64,
) -> Result<Delivered, PushError> {
<Self as PushTransport>::validate(self, config)?;
let url =
reqwest::Url::parse(&config.url).map_err(|e| PushError::Malformed(e.to_string()))?;
let host = url
.host_str()
.ok_or_else(|| PushError::Malformed("no host".to_owned()))?
.to_owned();
let port = url.port_or_known_default().unwrap_or(443);
let lookup = host
.strip_prefix('[')
.and_then(|inner| inner.strip_suffix(']'))
.unwrap_or(&host)
.to_owned();
let resolved = tokio::net::lookup_host((lookup.as_str(), port))
.await
.map_err(|e| PushError::Unroutable(format!("DNS for '{host}': {e}")))?;
crate::netguard::judge(self.reach(), &host, resolved)
.map_err(|e| PushError::Unroutable(e.to_string()))?;
let client = self.http()?;
let body = crate::core::canon::value_bytes(&message.payload);
let content_type = reqwest::header::HeaderValue::from_str(&message.content_type)
.map_err(|error| PushError::Malformed(format!("invalid content type: {error}")))?;
let id = reqwest::header::HeaderValue::from_str(&message.id)
.map_err(|error| PushError::Malformed(format!("invalid message id: {error}")))?;
let mut request = client
.post(url)
.header(reqwest::header::CONTENT_TYPE, content_type)
.header(HEADER_ID, id)
.header(HEADER_TIMESTAMP, at);
if let Some(signing) = self.signing.get(&config.id) {
request = request.header(HEADER_SIGNATURE, signing.value_for(&message.id, at, &body));
}
if let Some(token) = &config.token {
let value = reqwest::header::HeaderValue::from_str(token.expose())
.map_err(|error| PushError::Malformed(format!("invalid push token: {error}")))?;
request = request.header(HEADER_A2A_TOKEN, value);
}
let mut request = request.body(body);
if let Some(authentication) = &config.authentication {
let value = format!(
"{} {}",
authentication.scheme,
authentication.credentials.expose()
);
let value = reqwest::header::HeaderValue::from_str(&value).map_err(|error| {
PushError::Malformed(format!("invalid authentication: {error}"))
})?;
request = request.header(reqwest::header::AUTHORIZATION, value);
}
Ok(match request.send().await {
Ok(response) if response.status().is_success() => Delivered::Accepted,
Ok(response) => Delivered::Rejected {
status: response.status().as_u16(),
retry_after: retry_after_seconds(response.headers()),
},
Err(e) => Delivered::Unreachable(e.to_string()),
})
}
}
#[async_trait]
impl PushTransport for PushSender {
fn validate(&self, config: &PushConfig) -> Result<(), PushError> {
if self.operator {
reqwest::Url::parse(&config.url).map_err(|e| PushError::Malformed(e.to_string()))?;
} else {
self.policy
.check_allowing_loopback(&config.url, self.loopback_allowed())?;
}
if let Some(authentication) = &config.authentication {
authentication.validate()?;
}
Ok(())
}
async fn deliver(
&self,
config: &PushConfig,
message: &PushMessage,
at: u64,
) -> Result<Delivered, PushError> {
self.deliver_validated(config, message, at).await
}
}
fn retry_after_seconds(headers: &reqwest::header::HeaderMap) -> Option<u64> {
crate::core::retry_after_seconds(headers.get(reqwest::header::RETRY_AFTER)?.to_str().ok()?)
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Delivered {
Accepted,
Rejected {
status: u16,
retry_after: Option<u64>,
},
Unreachable(String),
}
impl Delivered {
#[must_use]
pub const fn is_permanent(&self) -> bool {
matches!(self, Self::Rejected { status: 410, .. })
}
#[must_use]
pub const fn retry_after(&self) -> Option<u64> {
match self {
Self::Rejected { retry_after, .. } => *retry_after,
_ => None,
}
}
}