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;
#[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 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 && is_loopback_name(&host)) {
return Err(PushError::NotHttps);
}
if !self.hosts.contains(&host) {
return Err(PushError::HostNotGranted(host));
}
Ok(())
}
}
fn is_loopback_name(host: &str) -> bool {
if host == "localhost" {
return true;
}
let bare = host
.strip_prefix('[')
.and_then(|h| h.strip_suffix(']'))
.unwrap_or(host);
bare.parse::<std::net::IpAddr>()
.is_ok_and(|ip| ip.is_loopback())
}
#[derive(Debug, Clone)]
pub struct PushSender {
policy: PushPolicy,
timeout: std::time::Duration,
#[cfg(feature = "testkit")]
plaintext_loopback: bool,
operator: bool,
signing: std::collections::BTreeMap<String, BodySigning>,
}
#[async_trait]
pub trait PushTransport: Send + Sync + Debug {
fn validate(&self, config: &PushConfig) -> Result<(), PushError>;
async fn deliver(
&self,
config: &PushConfig,
payload: &serde_json::Value,
) -> 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,
#[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,
#[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
}
async fn deliver_validated(
&self,
config: &PushConfig,
payload: &serde_json::Value,
) -> 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}")))?;
let addrs = if self.operator {
let addrs: Vec<_> = resolved.collect();
if addrs.is_empty() {
return Err(PushError::Unroutable(format!(
"DNS for '{host}' returned no addresses"
)));
}
addrs
} else if self.loopback_allowed() && is_loopback_name(&host) {
let addrs: Vec<_> = resolved.collect();
if addrs.is_empty() {
return Err(PushError::Unroutable(format!(
"DNS for '{host}' returned no addresses"
)));
}
addrs
} else {
crate::netguard::all_public(&host, resolved)
.map_err(|e| PushError::Unroutable(e.to_string()))?
};
let mut client = reqwest::Client::builder()
.timeout(self.timeout)
.no_proxy()
.redirect(reqwest::redirect::Policy::none());
for addr in &addrs {
client = client.resolve(&host, *addr);
}
let client = client
.build()
.map_err(|e| PushError::Unroutable(e.to_string()))?;
let body = crate::core::canon::value_bytes(payload);
let mut request = client
.post(url)
.header("Content-Type", "application/a2a+json");
if let Some(signing) = self.signing.get(&config.id) {
request = request.header(signing.header_name().clone(), signing.value_for(&body));
}
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(response.status().as_u16()),
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,
payload: &serde_json::Value,
) -> Result<Delivered, PushError> {
self.deliver_validated(config, payload).await
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Delivered {
Accepted,
Rejected(u16),
Unreachable(String),
}