Skip to main content

chio_api_protect/
spec_discovery.rs

1//! OpenAPI spec auto-discovery and loading.
2
3use std::collections::BTreeSet;
4use std::net::Ipv4Addr;
5
6use chio_http_core::{client_builder_with_contract, send_with_contract, HttpEgressContract};
7use url::{Host, Url};
8
9use crate::error::ProtectError;
10
11const DEFAULT_UPSTREAM_REDIRECT_LIMIT: u8 = 4;
12const DEFAULT_UPSTREAM_MAX_RESPONSE_BYTES: u64 = 64 * 1024 * 1024;
13
14/// Load an OpenAPI spec from a file path or URL.
15pub fn load_spec_from_file(path: &str) -> Result<String, ProtectError> {
16    std::fs::read_to_string(path)
17        .map_err(|e| ProtectError::SpecLoad(format!("cannot read {path}: {e}")))
18}
19
20pub(crate) fn default_upstream_egress_contract(
21    upstream: &str,
22) -> Result<HttpEgressContract, ProtectError> {
23    let url = Url::parse(upstream)
24        .map_err(|error| ProtectError::Config(format!("invalid upstream URL: {error}")))?;
25    let host = url
26        .host()
27        .ok_or_else(|| ProtectError::Config("upstream URL must include a host".to_string()))?;
28    let mut allowed_schemes = BTreeSet::new();
29    allowed_schemes.insert(url.scheme().to_ascii_lowercase());
30    let mut allowed_authority_set = BTreeSet::new();
31    allowed_authority_set.insert(normalized_authority(&url)?);
32    let contract = HttpEgressContract {
33        tenant_egress_namespace: "chio-api-protect-upstream".to_string(),
34        allowed_schemes,
35        allowed_authority_set,
36        deny_loopback: !host_is_loopback(&host),
37        deny_link_local: true,
38        deny_ipv6_ula: true,
39        max_redirect_chain: DEFAULT_UPSTREAM_REDIRECT_LIMIT,
40        max_response_bytes: DEFAULT_UPSTREAM_MAX_RESPONSE_BYTES,
41    };
42    contract
43        .validate_dispatchable_with_pinned_dns()
44        .map_err(|error| {
45            ProtectError::Config(format!("invalid upstream egress contract: {error}"))
46        })?;
47    Ok(contract)
48}
49
50/// Try to discover the OpenAPI spec from the upstream server.
51///
52/// Probes well-known paths (`/openapi.json`, `/openapi.yaml`,
53/// `/swagger.json`, `/api-docs`) in order, returning the first
54/// non-empty successful response.
55pub async fn discover_spec(upstream: &str) -> Result<String, ProtectError> {
56    let contract = default_upstream_egress_contract(upstream)?;
57    let client = client_builder_with_contract(&contract).build()?;
58    let well_known_paths = [
59        "/openapi.json",
60        "/openapi.yaml",
61        "/swagger.json",
62        "/api-docs",
63    ];
64
65    for path in &well_known_paths {
66        let url = format!("{}{}", upstream.trim_end_matches('/'), path);
67        let request = match client.get(&url).build() {
68            Ok(request) => request,
69            Err(_) => continue,
70        };
71        match send_with_contract(&contract, &client, request).await {
72            Ok(resp) if resp.status().is_success() => match resp.text().await {
73                Ok(body) if !body.is_empty() => return Ok(body),
74                _ => continue,
75            },
76            _ => continue,
77        }
78    }
79
80    Err(ProtectError::SpecLoad(
81        "could not auto-discover OpenAPI spec from upstream; use --spec to provide one".to_string(),
82    ))
83}
84
85fn normalized_authority(url: &Url) -> Result<String, ProtectError> {
86    let host = url
87        .host()
88        .ok_or_else(|| ProtectError::Config("upstream URL must include a host".to_string()))?;
89    let host = match host {
90        Host::Domain(domain) => domain.trim_end_matches('.').to_ascii_lowercase(),
91        Host::Ipv4(address) => address.to_string(),
92        Host::Ipv6(address) => format!("[{address}]"),
93    };
94    Ok(match url.port() {
95        Some(port) => format!("{host}:{port}"),
96        None => host,
97    })
98}
99
100fn host_is_loopback(host: &Host<&str>) -> bool {
101    match host {
102        Host::Domain(domain) => matches!(
103            domain.trim_end_matches('.').to_ascii_lowercase().as_str(),
104            "localhost" | "localhost.localdomain"
105        ),
106        Host::Ipv4(address) => address.is_loopback(),
107        Host::Ipv6(address) => {
108            address.is_loopback()
109                || address
110                    .to_ipv4_mapped()
111                    .is_some_and(|mapped: Ipv4Addr| mapped.is_loopback())
112        }
113    }
114}
115
116#[cfg(test)]
117mod tests {
118    use super::*;
119    use std::time::{SystemTime, UNIX_EPOCH};
120
121    #[test]
122    fn load_spec_from_existing_file() -> Result<(), Box<dyn std::error::Error>> {
123        let suffix = SystemTime::now().duration_since(UNIX_EPOCH)?.as_nanos();
124        let dir = std::env::temp_dir().join(format!("chio-api-protect-test-{suffix}"));
125        std::fs::create_dir_all(&dir)?;
126        let path = dir.join("spec.json");
127        std::fs::write(&path, r#"{"openapi":"3.1.0"}"#)?;
128        let spec = load_spec_from_file(&path.to_string_lossy())?;
129        assert!(spec.contains("3.1.0"));
130        let _ = std::fs::remove_dir_all(&dir);
131        Ok(())
132    }
133
134    #[test]
135    fn load_spec_from_missing_file_fails() {
136        let result = load_spec_from_file("/nonexistent/path/openapi.json");
137        assert!(result.is_err());
138    }
139
140    #[test]
141    fn default_upstream_egress_contract_accepts_production_hostname() {
142        let contract = match default_upstream_egress_contract("https://api.example.com") {
143            Ok(contract) => contract,
144            Err(error) => panic!("production hostname should build an egress contract: {error}"),
145        };
146        assert!(contract.allowed_authority_set.contains("api.example.com"));
147        assert!(contract.deny_loopback);
148    }
149
150    #[test]
151    fn default_upstream_egress_contract_accepts_loopback_ip_for_tests() {
152        let contract = match default_upstream_egress_contract("http://127.0.0.1:8080") {
153            Ok(contract) => contract,
154            Err(error) => panic!("loopback IP remains available for local proxy tests: {error}"),
155        };
156        assert!(contract.allowed_authority_set.contains("127.0.0.1:8080"));
157    }
158}