use core::fmt;
use std::time::Duration;
use serde::Deserialize;
use shep_core::barks::Bark;
use tokio::io::{AsyncBufReadExt, AsyncRead, AsyncWrite, AsyncWriteExt, BufReader};
use tokio::net::TcpStream;
use tokio_rustls::rustls::pki_types::ServerName;
use crate::fetch::{self, Target};
#[derive(Clone, PartialEq, Eq, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)]
pub enum Sink {
Discord {
url: String,
},
Slack {
url: String,
},
Json {
url: String,
body: Option<String>,
},
}
impl fmt::Debug for Sink {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let variant = match self {
Self::Discord { .. } => "Discord",
Self::Slack { .. } => "Slack",
Self::Json { .. } => "Json",
};
write!(f, "Sink::{variant} {{ url: <redacted> }}")
}
}
impl Sink {
fn url(&self) -> &str {
match self {
Self::Discord { url } | Self::Slack { url } | Self::Json { url, .. } => url,
}
}
fn https_only_kind(&self) -> Option<&'static str> {
match self {
Self::Discord { .. } => Some("discord"),
Self::Slack { .. } => Some("slack"),
Self::Json { .. } => None,
}
}
}
#[derive(Debug)]
pub enum SinkConfigError {
InsecureScheme {
name: String,
kind: &'static str,
},
}
impl fmt::Display for SinkConfigError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::InsecureScheme { name, kind } => write!(
f,
"sink \"{name}\" is a {kind} webhook configured with http://; \
{kind} only serves https://, and a {kind} webhook url is a \
bearer credential that must not travel in cleartext"
),
}
}
}
impl core::error::Error for SinkConfigError {}
pub fn require_secure_scheme(name: &str, sink: &Sink) -> Result<(), SinkConfigError> {
let Some(kind) = sink.https_only_kind() else {
return Ok(());
};
if sink.url().starts_with("https://") {
Ok(())
} else {
Err(SinkConfigError::InsecureScheme {
name: name.to_owned(),
kind,
})
}
}
#[derive(Debug)]
pub enum SinkError {
Template {
message: String,
},
Transport {
source: std::io::Error,
},
Status {
code: u16,
message: String,
},
}
impl fmt::Display for SinkError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Template { message } => {
write!(f, "templated sink body is not valid json: {message}")
}
Self::Transport { source } => write!(f, "sink delivery failed: {source}"),
Self::Status { code, message } => write!(f, "sink answered {code}: {message}"),
}
}
}
impl core::error::Error for SinkError {
fn source(&self) -> Option<&(dyn core::error::Error + 'static)> {
match self {
Self::Transport { source } => Some(source),
Self::Template { .. } | Self::Status { .. } => None,
}
}
}
impl From<std::io::Error> for SinkError {
fn from(source: std::io::Error) -> Self {
Self::Transport { source }
}
}
pub fn render_body(sink: &Sink, bark: &Bark) -> Result<String, SinkError> {
let body = match sink {
Sink::Discord { .. } => serde_json::json!({ "content": bark.message }).to_string(),
Sink::Slack { .. } => serde_json::json!({ "text": bark.message }).to_string(),
Sink::Json { body: None, .. } => serde_json::json!({
"subject": bark.subject,
"rule": bark.rule,
"message": bark.message,
"at_ms": bark.at_ms,
})
.to_string(),
Sink::Json {
body: Some(template),
..
} => {
let rendered = substitute(template, bark);
serde_json::from_str::<serde_json::Value>(&rendered).map_err(|source| {
SinkError::Template {
message: source.to_string(),
}
})?;
rendered
}
};
Ok(body)
}
fn substitute(template: &str, bark: &Bark) -> String {
let tokens: [(&str, String); 4] = [
("{subject}", json_escape(&bark.subject)),
("{rule}", json_escape(&bark.rule)),
("{message}", json_escape(&bark.message)),
("{at_ms}", bark.at_ms.to_string()),
];
let mut out = String::with_capacity(template.len());
let mut rest = template;
while let Some(brace) = rest.find('{') {
out.push_str(&rest[..brace]);
rest = &rest[brace..];
match tokens.iter().find(|(token, _)| rest.starts_with(token)) {
Some((token, value)) => {
out.push_str(value);
rest = &rest[token.len()..];
}
None => {
out.push('{');
rest = &rest[1..];
}
}
}
out.push_str(rest);
out
}
fn json_escape(s: &str) -> String {
let mut out = String::with_capacity(s.len());
for ch in s.chars() {
match ch {
'"' => out.push_str("\\\""),
'\\' => out.push_str("\\\\"),
'\n' => out.push_str("\\n"),
'\r' => out.push_str("\\r"),
'\t' => out.push_str("\\t"),
c if (c as u32) < 0x20 => out.push_str(&format!("\\u{:04x}", c as u32)),
c => out.push(c),
}
}
out
}
pub async fn deliver(sink: &Sink, bark: &Bark, timeout: Duration) -> Result<(), SinkError> {
let body = render_body(sink, bark)?;
let target = fetch::parse_url(sink.url()).map_err(|source| SinkError::Transport {
source: std::io::Error::other(source),
})?;
match tokio::time::timeout(timeout, deliver_inner(&target, &body)).await {
Ok(result) => result,
Err(_elapsed) => Err(SinkError::Transport {
source: std::io::Error::from(std::io::ErrorKind::TimedOut),
}),
}
}
fn build_request(target: &Target, body: &str) -> String {
let default_port = if target.https { 443 } else { 80 };
let host = if target.port == default_port {
target.host.clone()
} else {
format!("{}:{}", target.host, target.port)
};
format!(
"POST {path} HTTP/1.1\r\nHost: {host}\r\nContent-Type: application/json\r\nContent-Length: {len}\r\nConnection: close\r\n\r\n{body}",
path = target.path,
len = body.len(),
)
}
async fn deliver_inner(target: &Target, body: &str) -> Result<(), SinkError> {
let request = build_request(target, body);
let tcp = TcpStream::connect((target.host.as_str(), target.port)).await?;
if target.https {
let domain =
ServerName::try_from(target.host.clone()).map_err(|source| SinkError::Transport {
source: std::io::Error::other(source),
})?;
let tls = fetch::tls_connector().connect(domain, tcp).await?;
write_and_read(tls, &request).await
} else {
write_and_read(tcp, &request).await
}
}
async fn write_and_read<S: AsyncRead + AsyncWrite + Unpin>(
mut stream: S,
request: &str,
) -> Result<(), SinkError> {
stream.write_all(request.as_bytes()).await?;
stream.flush().await?;
read_response(stream).await
}
async fn read_response<S: AsyncRead + Unpin>(stream: S) -> Result<(), SinkError> {
let mut reader = BufReader::new(stream);
let mut status_line = String::new();
reader.read_line(&mut status_line).await?;
let code = parse_status_code(&status_line)?;
if (200..300).contains(&code) {
return Ok(());
}
loop {
let mut line = String::new();
let read = reader.read_line(&mut line).await?;
if read == 0 || line == "\r\n" || line == "\n" {
break;
}
}
let mut diagnostic = String::new();
reader.read_line(&mut diagnostic).await?;
Err(SinkError::Status {
code,
message: diagnostic.trim_end().to_string(),
})
}
fn parse_status_code(status_line: &str) -> Result<u16, SinkError> {
status_line
.split_whitespace()
.nth(1)
.and_then(|code| code.parse().ok())
.ok_or_else(|| SinkError::Transport {
source: std::io::Error::other("malformed http status line"),
})
}
#[cfg(test)]
mod tests {
use tokio::sync::oneshot;
use super::*;
use crate::http::{HttpRequest, read_request, write_response};
fn bark_for(subject: &str, message: &str) -> Bark {
Bark {
at_ms: 1_700_000_000_000,
rule: "watchdog".to_string(),
subject: subject.to_string(),
message: message.to_string(),
sinks: Vec::new(),
}
}
fn discord_sink() -> Sink {
Sink::Discord {
url: "https://discord.com/api/webhooks/1/super-secret-token".to_string(),
}
}
fn slack_sink() -> Sink {
Sink::Slack {
url: "https://hooks.slack.com/services/T0/B0/super-secret-token".to_string(),
}
}
async fn one_shot_sink(
status: u16,
body: &str,
) -> (std::net::SocketAddr, oneshot::Receiver<HttpRequest>) {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let (tx, rx) = oneshot::channel();
let body = body.to_string();
tokio::spawn(async move {
let (mut stream, _peer) = listener.accept().await.unwrap();
let req = read_request(&mut stream, Duration::from_secs(5))
.await
.unwrap();
write_response(&mut stream, status, "application/json", body.as_bytes())
.await
.unwrap();
let _ = tx.send(req);
});
(addr, rx)
}
#[test]
fn each_webhook_gets_the_body_its_own_endpoint_expects() {
let bark = bark_for("web", "the shepherd gave up on web");
let discord: serde_json::Value =
serde_json::from_str(&render_body(&discord_sink(), &bark).unwrap()).unwrap();
assert_eq!(discord["content"], "the shepherd gave up on web");
assert!(discord.get("text").is_none());
let slack: serde_json::Value =
serde_json::from_str(&render_body(&slack_sink(), &bark).unwrap()).unwrap();
assert_eq!(slack["text"], "the shepherd gave up on web");
assert!(slack.get("content").is_none());
}
#[test]
fn a_template_that_does_not_render_json_is_refused_before_it_is_sent() {
let sink = Sink::Json {
url: "http://127.0.0.1:1/".to_string(),
body: Some(r#"{"text": "{message}"#.to_string()),
};
assert!(matches!(
render_body(&sink, &bark_for("web", "x")),
Err(SinkError::Template { .. })
));
}
#[test]
fn a_substituted_value_is_json_escaped_into_the_template() {
let sink = Sink::Json {
url: "http://127.0.0.1:1/".to_string(),
body: Some(r#"{"text": "{message}"}"#.to_string()),
};
let bark = bark_for("web", r#"app "we"b" crashed"#);
let rendered = render_body(&sink, &bark).unwrap();
let value: serde_json::Value = serde_json::from_str(&rendered).unwrap();
assert_eq!(value["text"], bark.message);
}
#[test]
fn a_placeholder_inside_a_substituted_value_survives_later_passes() {
let bark = Bark {
at_ms: 12_345,
rule: "gave_up".to_string(),
subject: "web".to_string(),
message: "{at_ms} gave up: restart budget exhausted".to_string(),
sinks: Vec::new(),
};
let sink = Sink::Json {
url: "http://127.0.0.1:1/".to_string(),
body: Some(r#"{"text": "{message}", "stamp": {at_ms}}"#.to_string()),
};
let rendered = render_body(&sink, &bark).unwrap();
let value: serde_json::Value = serde_json::from_str(&rendered).unwrap();
assert_eq!(
value["text"], "{at_ms} gave up: restart budget exhausted",
"the literal {{at_ms}} carried inside the message must not be rewritten"
);
assert_eq!(value["stamp"], 12_345);
}
#[tokio::test]
async fn a_delivery_posts_json_to_the_url_it_was_given() {
let (addr, captured) = one_shot_sink(200, "").await;
let sink = Sink::Json {
url: format!("http://{addr}/hook"),
body: None,
};
deliver(&sink, &bark_for("web", "x"), Duration::from_secs(5))
.await
.unwrap();
let req = tokio::time::timeout(Duration::from_secs(5), captured)
.await
.expect("the sink server must receive a request")
.unwrap();
assert_eq!(req.method, "POST");
assert_eq!(req.target, "/hook");
assert_eq!(
req.headers.get("content-type").map(String::as_str),
Some("application/json")
);
assert_eq!(
serde_json::from_slice::<serde_json::Value>(&req.body).unwrap()["subject"],
"web"
);
}
#[tokio::test]
async fn a_refused_delivery_is_a_failure_carrying_the_status() {
let (addr, _captured) = one_shot_sink(429, "rate limited").await;
let err = deliver(
&Sink::Json {
url: format!("http://{addr}/"),
body: None,
},
&bark_for("web", "x"),
Duration::from_secs(5),
)
.await
.unwrap_err();
assert!(matches!(err, SinkError::Status { code: 429, .. }));
}
#[test]
fn a_sinks_debug_never_prints_its_webhook() {
let rendered = format!("{:?}", discord_sink());
assert_eq!(rendered, "Sink::Discord { url: <redacted> }");
assert!(!rendered.contains("discord.com"));
}
#[test]
fn a_discord_sink_over_http_is_refused() {
let sink = Sink::Discord {
url: "http://discord.com/api/webhooks/1/super-secret-token".to_string(),
};
let err = require_secure_scheme("ops", &sink).unwrap_err();
assert!(matches!(
err,
SinkConfigError::InsecureScheme {
kind: "discord",
..
}
));
assert!(!err.to_string().contains("discord.com"));
}
#[test]
fn a_slack_sink_over_http_is_refused() {
let sink = Sink::Slack {
url: "http://hooks.slack.com/services/T0/B0/super-secret-token".to_string(),
};
let err = require_secure_scheme("ops", &sink).unwrap_err();
assert!(matches!(
err,
SinkConfigError::InsecureScheme { kind: "slack", .. }
));
assert!(!err.to_string().contains("hooks.slack.com"));
}
#[test]
fn a_json_sink_over_http_is_accepted() {
let sink = Sink::Json {
url: "http://127.0.0.1:8080/hook".to_string(),
body: None,
};
require_secure_scheme("ops", &sink).unwrap();
}
}