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),
}
#[async_trait]
pub trait PushStore: Send + Sync + Debug {
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 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 {
self.hosts.insert(
host.as_ref()
.trim()
.trim_end_matches('.')
.to_ascii_lowercase(),
);
self
}
pub fn check(&self, url: &str) -> Result<(), PushError> {
let parsed = reqwest::Url::parse(url).map_err(|e| PushError::Malformed(e.to_string()))?;
if parsed.scheme() != "https" {
return Err(PushError::NotHttps);
}
let host = parsed
.host_str()
.ok_or_else(|| PushError::Malformed("no host".to_owned()))?
.trim_end_matches('.')
.to_ascii_lowercase();
if !self.hosts.contains(&host) {
return Err(PushError::HostNotGranted(host));
}
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct PushSender {
policy: PushPolicy,
timeout: std::time::Duration,
}
#[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,
}
}
#[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
}
pub async fn deliver(
&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 resolved = tokio::net::lookup_host((host.as_str(), port))
.await
.map_err(|e| PushError::Unroutable(format!("DNS for '{host}': {e}")))?;
let addrs = 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 mut request = client
.post(url)
.header("Content-Type", "application/a2a+json")
.json(payload);
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> {
self.policy.check(&config.url)?;
if let Some(authentication) = &config.authentication {
authentication.validate()?;
}
Ok(())
}
async fn deliver(
&self,
config: &PushConfig,
payload: &serde_json::Value,
) -> Result<Delivered, PushError> {
PushSender::deliver(self, config, payload).await
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Delivered {
Accepted,
Rejected(u16),
Unreachable(String),
}