use crate::progress::Progress;
use anyhow::{anyhow, bail, Context, Result};
use std::process::Command;
const CACHE_DISABLED: &str = "4135ea2d-6df8-44a3-9df3-4b5a84be39ad";
const CACHING_OPTIMIZED: &str = "658327ea-f89d-4fab-a63d-7e88639e58f6";
const ALL_VIEWER_EXCEPT_HOST: &str = "b689b0a8-53d0-40ab-baf2-68738e2966ac";
const PAGE_REWRITE_JS: &str = r#"function handler(event) {
var req = event.request;
if (req.uri.startsWith('/d/')) { req.uri = '/d/p'; }
if (req.uri.startsWith('/r/')) { req.uri = '/r/p'; }
return req;
}
"#;
pub struct Front {
pub distribution_id: String,
pub domain: String,
}
pub fn front_gate(
profile: Option<&str>,
account: &str,
origin_host: &str,
existing_distribution: Option<&str>,
progress: &dyn Progress,
) -> Result<Front> {
if let Some(dist_id) = existing_distribution {
progress.step("cloudfront (reuse)");
let page_fn_arn = ensure_page_function(profile, account)?;
let out = aws(
profile,
&[
"cloudfront",
"get-distribution-config",
"--id",
dist_id,
"--output",
"json",
],
)?;
if !out.status.success() {
bail!(
"get-distribution-config: {}",
String::from_utf8_lossy(&out.stderr).trim()
);
}
let v: serde_json::Value = serde_json::from_slice(&out.stdout)?;
let etag = v["ETag"]
.as_str()
.ok_or_else(|| anyhow!("no ETag"))?
.to_string();
let mut cfg = v["DistributionConfig"].clone();
merge_request_behaviors(&mut cfg, &page_fn_arn);
let tmp = std::env::temp_dir().join(format!("dove-dist-{dist_id}.json"));
std::fs::write(&tmp, cfg.to_string())?;
let cfg_arg = format!("file://{}", tmp.display());
let out = aws(
profile,
&[
"cloudfront",
"update-distribution",
"--id",
dist_id,
"--distribution-config",
&cfg_arg,
"--if-match",
&etag,
"--output",
"json",
],
);
let _ = std::fs::remove_file(&tmp);
let out = out?;
if !out.status.success() {
bail!(
"update-distribution: {}",
String::from_utf8_lossy(&out.stderr).trim()
);
}
let domain = distribution_domain(profile, dist_id)?;
progress.done("cloudfront (reuse)");
return Ok(Front {
distribution_id: dist_id.to_string(),
domain,
});
}
progress.step("page cache function");
let fn_arn = ensure_page_function(profile, account);
if fn_arn.is_ok() {
progress.done("page cache function");
}
let fn_arn = fn_arn?;
progress.step("cloudfront distribution");
let front = create_distribution(profile, origin_host, &fn_arn);
if front.is_ok() {
progress.done("cloudfront distribution");
}
front
}
fn ensure_page_function(profile: Option<&str>, account: &str) -> Result<String> {
let arn = format!("arn:aws:cloudfront::{account}:function/dove-page-rewrite");
let tmp = std::env::temp_dir().join("dove-page-rewrite.js");
std::fs::write(&tmp, PAGE_REWRITE_JS)?;
let code_arg = format!("fileb://{}", tmp.display());
let out = aws(
profile,
&[
"cloudfront",
"create-function",
"--name",
"dove-page-rewrite",
"--function-config",
"Comment=dove page cache-key rewrite,Runtime=cloudfront-js-2.0",
"--function-code",
&code_arg,
"--output",
"json",
],
);
let _ = std::fs::remove_file(&tmp);
let out = out?;
if out.status.success() {
let v: serde_json::Value = serde_json::from_slice(&out.stdout)?;
let etag = v["ETag"]
.as_str()
.ok_or_else(|| anyhow!("no ETag from create-function"))?;
let pub_out = aws(
profile,
&[
"cloudfront",
"publish-function",
"--name",
"dove-page-rewrite",
"--if-match",
etag,
],
)?;
if !pub_out.status.success() {
bail!(
"publish-function: {}",
String::from_utf8_lossy(&pub_out.stderr).trim()
);
}
} else if !String::from_utf8_lossy(&out.stderr).contains("FunctionAlreadyExists") {
bail!(
"create-function: {}",
String::from_utf8_lossy(&out.stderr).trim()
);
}
Ok(arn)
}
fn create_distribution(
profile: Option<&str>,
origin_host: &str,
page_fn_arn: &str,
) -> Result<Front> {
let caller_ref = format!("dove-{origin_host}");
let config = distribution_config(&caller_ref, origin_host, page_fn_arn);
let out = aws(
profile,
&[
"cloudfront",
"create-distribution",
"--distribution-config",
&config,
"--output",
"json",
],
)?;
if !out.status.success() {
bail!(
"create-distribution: {}",
String::from_utf8_lossy(&out.stderr).trim()
);
}
let v: serde_json::Value = serde_json::from_slice(&out.stdout)?;
Ok(Front {
distribution_id: v["Distribution"]["Id"]
.as_str()
.ok_or_else(|| anyhow!("no distribution Id"))?
.to_string(),
domain: v["Distribution"]["DomainName"]
.as_str()
.ok_or_else(|| anyhow!("no distribution DomainName"))?
.to_string(),
})
}
pub fn add_alias(
profile: Option<&str>,
dist_id: &str,
domain: &str,
cert_arn: &str,
) -> Result<String> {
let out = aws(
profile,
&[
"cloudfront",
"get-distribution-config",
"--id",
dist_id,
"--output",
"json",
],
)?;
if !out.status.success() {
bail!(
"get-distribution-config: {}",
String::from_utf8_lossy(&out.stderr).trim()
);
}
let v: serde_json::Value = serde_json::from_slice(&out.stdout)?;
let etag = v["ETag"]
.as_str()
.ok_or_else(|| anyhow!("no ETag"))?
.to_string();
let mut cfg = v["DistributionConfig"].clone();
cfg["Aliases"] = serde_json::json!({"Quantity": 1, "Items": [domain]});
cfg["ViewerCertificate"] = serde_json::json!({
"ACMCertificateArn": cert_arn,
"SSLSupportMethod": "sni-only",
"MinimumProtocolVersion": "TLSv1.2_2021",
"CloudFrontDefaultCertificate": false,
});
let domain_name = cfg["DomainName"].as_str().unwrap_or_default().to_string();
let tmp = std::env::temp_dir().join(format!("dove-dist-{dist_id}.json"));
std::fs::write(&tmp, cfg.to_string())?;
let cfg_arg = format!("file://{}", tmp.display());
let out = aws(
profile,
&[
"cloudfront",
"update-distribution",
"--id",
dist_id,
"--distribution-config",
&cfg_arg,
"--if-match",
&etag,
"--output",
"json",
],
);
let _ = std::fs::remove_file(&tmp);
if !out?.status.success() {
bail!("update-distribution failed");
}
Ok(domain_name)
}
fn distribution_domain(profile: Option<&str>, dist_id: &str) -> Result<String> {
let out = aws(
profile,
&[
"cloudfront",
"get-distribution",
"--id",
dist_id,
"--output",
"json",
],
)?;
if !out.status.success() {
bail!(
"get-distribution: {}",
String::from_utf8_lossy(&out.stderr).trim()
);
}
let v: serde_json::Value = serde_json::from_slice(&out.stdout)?;
v["Distribution"]["DomainName"]
.as_str()
.map(str::to_string)
.ok_or_else(|| anyhow!("no DomainName"))
}
fn cache_behavior(path: &str, cache_policy: &str, fn_arn: Option<&str>) -> serde_json::Value {
let fns = match fn_arn {
Some(arn) => serde_json::json!({
"Quantity": 1,
"Items": [{"EventType": "viewer-request", "FunctionARN": arn}]
}),
None => serde_json::json!({"Quantity": 0}),
};
serde_json::json!({
"PathPattern": path,
"TargetOriginId": "gate",
"ViewerProtocolPolicy": "redirect-to-https",
"CachePolicyId": cache_policy,
"Compress": true,
"AllowedMethods": {
"Quantity": 2, "Items": ["GET", "HEAD"],
"CachedMethods": {"Quantity": 2, "Items": ["GET", "HEAD"]}
},
"FunctionAssociations": fns,
"SmoothStreaming": false,
"FieldLevelEncryptionId": "",
"LambdaFunctionAssociations": {"Quantity": 0},
"TrustedSigners": {"Enabled": false, "Quantity": 0},
"TrustedKeyGroups": {"Enabled": false, "Quantity": 0}
})
}
fn merge_request_behaviors(config: &mut serde_json::Value, page_fn_arn: &str) {
config["DefaultCacheBehavior"]["AllowedMethods"] = serde_json::json!({
"Quantity": 7,
"Items": ["GET", "HEAD", "OPTIONS", "PUT", "POST", "PATCH", "DELETE"],
"CachedMethods": {"Quantity": 2, "Items": ["GET", "HEAD"]}
});
if !config["CacheBehaviors"]["Items"].is_array() {
config["CacheBehaviors"] = serde_json::json!({"Quantity": 0, "Items": []});
}
let items = config["CacheBehaviors"]["Items"]
.as_array_mut()
.expect("CacheBehaviors.Items normalized to an array above");
if !items.iter().any(|b| b["PathPattern"] == "/r/*") {
items.push(cache_behavior("/r/*", CACHING_OPTIMIZED, Some(page_fn_arn)));
}
let len = items.len();
config["CacheBehaviors"]["Quantity"] = serde_json::json!(len);
}
pub fn distribution_config(caller_ref: &str, origin_host: &str, page_fn_arn: &str) -> String {
serde_json::json!({
"CallerReference": caller_ref,
"Comment": "dove gate",
"Enabled": true,
"Origins": {"Quantity": 1, "Items": [{
"Id": "gate",
"DomainName": origin_host,
"CustomOriginConfig": {
"HTTPPort": 80,
"HTTPSPort": 443,
"OriginProtocolPolicy": "https-only",
"OriginSslProtocols": {"Quantity": 1, "Items": ["TLSv1.2"]}
}
}]},
"DefaultCacheBehavior": {
"TargetOriginId": "gate",
"ViewerProtocolPolicy": "redirect-to-https",
"AllowedMethods": {
"Quantity": 7, "Items": ["GET", "HEAD", "OPTIONS", "PUT", "POST", "PATCH", "DELETE"],
"CachedMethods": {"Quantity": 2, "Items": ["GET", "HEAD"]}
},
"CachePolicyId": CACHE_DISABLED,
"OriginRequestPolicyId": ALL_VIEWER_EXCEPT_HOST,
"Compress": true
},
"CacheBehaviors": {"Quantity": 3, "Items": [
cache_behavior("/d/*", CACHING_OPTIMIZED, Some(page_fn_arn)),
cache_behavior("/r/*", CACHING_OPTIMIZED, Some(page_fn_arn)),
cache_behavior("/og.png", CACHING_OPTIMIZED, None)
]},
"ViewerCertificate": {"CloudFrontDefaultCertificate": true}
})
.to_string()
}
fn aws(profile: Option<&str>, args: &[&str]) -> Result<std::process::Output> {
let mut cmd = Command::new("aws");
if let Some(p) = profile {
cmd.args(["--profile", p]);
}
cmd.args(args)
.output()
.with_context(|| format!("running aws {}", args.join(" ")))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn distribution_config_fronts_the_api_gateway_origin() {
let c = distribution_config(
"dove-x",
"abc.execute-api.us-east-1.amazonaws.com",
"arn:aws:cloudfront::1:function/dove-page-rewrite",
);
assert!(c.contains("\"DomainName\":\"abc.execute-api.us-east-1.amazonaws.com\""));
assert!(c.contains(CACHE_DISABLED)); assert!(c.contains(CACHING_OPTIMIZED)); assert!(c.contains("/d/*") && c.contains("/r/*") && c.contains("/og.png"));
assert!(c.contains("\"Quantity\":3"));
let parsed: serde_json::Value = serde_json::from_str(&c).unwrap();
let behaviors = parsed["CacheBehaviors"]["Items"].as_array().unwrap();
let r_behavior = behaviors
.iter()
.find(|b| b["PathPattern"] == "/r/*")
.expect("/r/* behavior present");
assert_eq!(
r_behavior["FunctionAssociations"]["Items"][0]["FunctionARN"],
"arn:aws:cloudfront::1:function/dove-page-rewrite"
);
assert_eq!(r_behavior["CachePolicyId"], CACHING_OPTIMIZED);
assert!(c.contains(ALL_VIEWER_EXCEPT_HOST));
assert!(!c.contains("OriginAccessControlId"));
let default_methods = &parsed["DefaultCacheBehavior"]["AllowedMethods"];
let allowed: Vec<&str> = default_methods["Items"]
.as_array()
.unwrap()
.iter()
.map(|v| v.as_str().unwrap())
.collect();
assert!(
allowed.contains(&"POST"),
"default behavior must allow POST for /done"
);
assert!(allowed.contains(&"GET") && allowed.contains(&"HEAD"));
let cached: Vec<&str> = default_methods["CachedMethods"]["Items"]
.as_array()
.unwrap()
.iter()
.map(|v| v.as_str().unwrap())
.collect();
assert_eq!(cached, vec!["GET", "HEAD"]);
}
const TEST_FN_ARN: &str = "arn:aws:cloudfront::000000000000:function/dove-page-rewrite";
fn live_config_with_custom_domain() -> serde_json::Value {
serde_json::json!({
"CallerReference": "dove-abc.execute-api.us-east-1.amazonaws.com",
"Comment": "dove gate",
"Enabled": true,
"Aliases": {"Quantity": 1, "Items": ["share.example.com"]},
"Origins": {"Quantity": 1, "Items": [{
"Id": "gate",
"DomainName": "abc.execute-api.us-east-1.amazonaws.com"
}]},
"DefaultCacheBehavior": {
"TargetOriginId": "gate",
"ViewerProtocolPolicy": "redirect-to-https",
"AllowedMethods": {
"Quantity": 2, "Items": ["GET", "HEAD"],
"CachedMethods": {"Quantity": 2, "Items": ["GET", "HEAD"]}
},
"CachePolicyId": CACHE_DISABLED,
"OriginRequestPolicyId": ALL_VIEWER_EXCEPT_HOST,
"Compress": true
},
"CacheBehaviors": {"Quantity": 2, "Items": [
cache_behavior("/d/*", CACHING_OPTIMIZED, Some(TEST_FN_ARN)),
cache_behavior("/og.png", CACHING_OPTIMIZED, None)
]},
"ViewerCertificate": {
"ACMCertificateArn": "arn:aws:acm:us-east-1:000000000000:certificate/test",
"SSLSupportMethod": "sni-only",
"CloudFrontDefaultCertificate": false
}
})
}
fn path_patterns(config: &serde_json::Value) -> Vec<String> {
config["CacheBehaviors"]["Items"]
.as_array()
.unwrap()
.iter()
.map(|b| b["PathPattern"].as_str().unwrap().to_string())
.collect()
}
#[test]
fn merge_preserves_alias_and_cert() {
let mut cfg = live_config_with_custom_domain();
merge_request_behaviors(&mut cfg, TEST_FN_ARN);
assert_eq!(
cfg["Aliases"]["Items"].as_array().unwrap(),
&vec![serde_json::json!("share.example.com")]
);
assert_eq!(
cfg["ViewerCertificate"]["ACMCertificateArn"],
"arn:aws:acm:us-east-1:000000000000:certificate/test"
);
assert_eq!(
cfg["ViewerCertificate"]["CloudFrontDefaultCertificate"],
false
);
assert_eq!(cfg["ViewerCertificate"]["SSLSupportMethod"], "sni-only");
assert_eq!(cfg["Comment"], "dove gate");
assert_eq!(cfg["Origins"]["Items"][0]["Id"], "gate");
}
#[test]
fn merge_allows_post_but_caches_only_get_head() {
let mut cfg = live_config_with_custom_domain();
merge_request_behaviors(&mut cfg, TEST_FN_ARN);
let dm = &cfg["DefaultCacheBehavior"]["AllowedMethods"];
let allowed: Vec<&str> = dm["Items"]
.as_array()
.unwrap()
.iter()
.map(|v| v.as_str().unwrap())
.collect();
assert!(
allowed.contains(&"POST"),
"default behavior must allow POST for /done"
);
let cached: Vec<&str> = dm["CachedMethods"]["Items"]
.as_array()
.unwrap()
.iter()
.map(|v| v.as_str().unwrap())
.collect();
assert_eq!(cached, vec!["GET", "HEAD"]);
}
#[test]
fn merge_adds_r_behavior() {
let mut cfg = live_config_with_custom_domain();
merge_request_behaviors(&mut cfg, TEST_FN_ARN);
let patterns = path_patterns(&cfg);
assert!(patterns.contains(&"/r/*".to_string()));
assert!(patterns.contains(&"/d/*".to_string()));
assert!(patterns.contains(&"/og.png".to_string()));
let items = cfg["CacheBehaviors"]["Items"].as_array().unwrap();
let r = items
.iter()
.find(|b| b["PathPattern"] == "/r/*")
.expect("/r/* behavior present");
assert_eq!(
r["FunctionAssociations"]["Items"][0]["FunctionARN"],
TEST_FN_ARN
);
assert_eq!(r["CachePolicyId"], CACHING_OPTIMIZED);
assert_eq!(
cfg["CacheBehaviors"]["Quantity"].as_u64().unwrap() as usize,
items.len()
);
}
#[test]
fn merge_is_idempotent() {
let mut cfg = live_config_with_custom_domain();
merge_request_behaviors(&mut cfg, TEST_FN_ARN);
let len_after_first = cfg["CacheBehaviors"]["Items"].as_array().unwrap().len();
let quantity_after_first = cfg["CacheBehaviors"]["Quantity"].as_u64().unwrap();
merge_request_behaviors(&mut cfg, TEST_FN_ARN);
let items = cfg["CacheBehaviors"]["Items"].as_array().unwrap();
let r_count = items.iter().filter(|b| b["PathPattern"] == "/r/*").count();
assert_eq!(r_count, 1, "exactly one /r/* behavior after two merges");
assert_eq!(items.len(), len_after_first);
assert_eq!(
cfg["CacheBehaviors"]["Quantity"].as_u64().unwrap(),
quantity_after_first
);
assert_eq!(
cfg["CacheBehaviors"]["Quantity"].as_u64().unwrap() as usize,
items.len()
);
}
}