mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-04 00:17:45 +08:00
fix(dns): unify provider resolution and bound SMTP and tunnel egress
Share provider DNS policy across WebSocket and connection probes, handle bracketed IPv6 literals, and preserve bounded address sets for outbound clients. Bound SMTP DNS and TCP setup with multi-address fallback. Add opt-in trusted proxy DNS for tunnel upstreams while retaining default IP ACLs and origin isolation. Document DNS policy boundaries and verify 809 gateway, tunnel, and HTTP regression tests.
This commit is contained in:
@@ -23,7 +23,6 @@ const MAX_BARK_TEMPLATE_BYTES: usize = 256 * 1024;
|
||||
const MAX_BARK_TITLE_BYTES: usize = 512;
|
||||
const MAX_BARK_BODY_BYTES: usize = 2 * 1024 * 1024;
|
||||
const MAX_BARK_RENDERED_BODY_BYTES: usize = 2 * 1024 * 1024;
|
||||
const MAX_BARK_RESOLVED_ADDRESSES: usize = 32;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct BarkPushConfig {
|
||||
@@ -208,19 +207,20 @@ async fn build_bark_push_client_and_url(
|
||||
let port = push_url
|
||||
.port_or_known_default()
|
||||
.ok_or_else(|| GatewayError::Internal("Bark 服务器地址缺少端口".to_string()))?;
|
||||
let addresses = if let Ok(ip) = host.parse::<IpAddr>() {
|
||||
vec![SocketAddr::new(ip, port)]
|
||||
} else {
|
||||
tokio::time::timeout(
|
||||
std::time::Duration::from_millis(BARK_CONNECT_TIMEOUT_MS),
|
||||
tokio::net::lookup_host((host.as_str(), port)),
|
||||
)
|
||||
.await
|
||||
.map_err(|_| GatewayError::Internal("Bark 服务器 DNS 解析超时".to_string()))?
|
||||
.map_err(|_| GatewayError::Internal("Bark 服务器 DNS 解析失败".to_string()))?
|
||||
.take(MAX_BARK_RESOLVED_ADDRESSES)
|
||||
.collect::<Vec<_>>()
|
||||
};
|
||||
let addresses = aether_http::lookup_host_with_limits(
|
||||
host.as_str(),
|
||||
port,
|
||||
std::time::Duration::from_millis(BARK_CONNECT_TIMEOUT_MS),
|
||||
)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
let message = match error.kind() {
|
||||
std::io::ErrorKind::TimedOut => "Bark 服务器 DNS 解析超时",
|
||||
std::io::ErrorKind::InvalidData => "Bark 服务器 DNS 解析返回过多地址",
|
||||
_ => "Bark 服务器 DNS 解析失败",
|
||||
};
|
||||
GatewayError::Internal(message.to_string())
|
||||
})?;
|
||||
let allow_benchmarking_ip = push_url.scheme() == "https"
|
||||
&& push_url.port_or_known_default() == Some(443)
|
||||
&& host.eq_ignore_ascii_case("api.day.app");
|
||||
|
||||
@@ -136,14 +136,16 @@ pub(crate) async fn send_smtp_email(
|
||||
email: ComposedEmail,
|
||||
) -> Result<(), GatewayError> {
|
||||
validate_smtp_delivery_inputs(&config, &email)?;
|
||||
tokio::task::spawn_blocking(move || send_smtp_email_blocking(config, email))
|
||||
let stream = connect_tcp_stream(&config).await?;
|
||||
tokio::task::spawn_blocking(move || send_smtp_email_blocking(config, email, stream))
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?
|
||||
}
|
||||
|
||||
pub(crate) async fn probe_smtp_connection(config: SmtpDeliveryConfig) -> Result<(), GatewayError> {
|
||||
validate_smtp_config(&config)?;
|
||||
tokio::task::spawn_blocking(move || probe_smtp_connection_blocking(config))
|
||||
let stream = connect_tcp_stream(&config).await?;
|
||||
tokio::task::spawn_blocking(move || probe_smtp_connection_blocking(config, stream))
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?
|
||||
}
|
||||
@@ -328,43 +330,58 @@ fn resolve_server_name(host: &str) -> Result<rustls::pki_types::ServerName<'stat
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
}
|
||||
|
||||
fn connect_tcp_stream(config: &SmtpDeliveryConfig) -> Result<std::net::TcpStream, GatewayError> {
|
||||
use std::net::ToSocketAddrs;
|
||||
let addresses = (config.host.as_str(), config.port)
|
||||
.to_socket_addrs()
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?
|
||||
.take(16)
|
||||
.collect::<Vec<_>>();
|
||||
if addresses.is_empty() {
|
||||
return Err(GatewayError::Internal(
|
||||
"smtp host did not resolve to an address".to_string(),
|
||||
));
|
||||
}
|
||||
let deadline = std::time::Instant::now()
|
||||
.checked_add(std::time::Duration::from_secs(SMTP_TIMEOUT_SECS))
|
||||
.unwrap_or_else(std::time::Instant::now);
|
||||
let mut last_error = None;
|
||||
let mut stream = None;
|
||||
for address in addresses {
|
||||
let remaining = deadline.saturating_duration_since(std::time::Instant::now());
|
||||
if remaining.is_zero() {
|
||||
break;
|
||||
async fn connect_tcp_stream(
|
||||
config: &SmtpDeliveryConfig,
|
||||
) -> Result<std::net::TcpStream, GatewayError> {
|
||||
connect_tcp_stream_with_dns(
|
||||
aether_http::lookup_host_with_limits(
|
||||
&config.host,
|
||||
config.port,
|
||||
aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT,
|
||||
),
|
||||
std::time::Duration::from_secs(SMTP_TIMEOUT_SECS),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn connect_tcp_stream_with_dns(
|
||||
lookup: impl std::future::Future<Output = std::io::Result<Vec<std::net::SocketAddr>>>,
|
||||
timeout: std::time::Duration,
|
||||
) -> Result<std::net::TcpStream, GatewayError> {
|
||||
let stream = tokio::time::timeout(timeout, async {
|
||||
let addresses = lookup.await.map_err(|error| {
|
||||
let message = match error.kind() {
|
||||
std::io::ErrorKind::TimedOut => "smtp DNS resolution timed out",
|
||||
std::io::ErrorKind::InvalidData => {
|
||||
"smtp DNS resolution returned too many addresses"
|
||||
}
|
||||
_ => "smtp DNS resolution failed",
|
||||
};
|
||||
GatewayError::Internal(message.to_string())
|
||||
})?;
|
||||
if addresses.is_empty() {
|
||||
return Err(GatewayError::Internal(
|
||||
"smtp host did not resolve to an address".to_string(),
|
||||
));
|
||||
}
|
||||
match std::net::TcpStream::connect_timeout(&address, remaining) {
|
||||
Ok(candidate) => {
|
||||
stream = Some(candidate);
|
||||
break;
|
||||
}
|
||||
Err(err) => last_error = Some(err),
|
||||
}
|
||||
}
|
||||
let stream = stream.ok_or_else(|| {
|
||||
GatewayError::Internal(
|
||||
last_error
|
||||
.map(|err| err.to_string())
|
||||
.unwrap_or_else(|| "smtp connection timed out".to_string()),
|
||||
)
|
||||
})?;
|
||||
let attempts = addresses
|
||||
.into_iter()
|
||||
.map(|address| Box::pin(tokio::net::TcpStream::connect(address)));
|
||||
futures_util::future::select_ok(attempts)
|
||||
.await
|
||||
.map(|(stream, _)| stream)
|
||||
.map_err(|error| {
|
||||
GatewayError::Internal(format!("smtp connection failed ({})", error.kind()))
|
||||
})
|
||||
})
|
||||
.await
|
||||
.map_err(|_| GatewayError::Internal("smtp DNS or TCP connection timed out".to_string()))??;
|
||||
let stream = stream
|
||||
.into_std()
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
stream
|
||||
.set_nonblocking(false)
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
stream
|
||||
.set_read_timeout(Some(std::time::Duration::from_secs(SMTP_TIMEOUT_SECS)))
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
@@ -680,16 +697,15 @@ fn smtp_probe_connection<S: std::io::Read + std::io::Write>(
|
||||
fn send_smtp_email_blocking(
|
||||
config: SmtpDeliveryConfig,
|
||||
email: ComposedEmail,
|
||||
stream: std::net::TcpStream,
|
||||
) -> Result<(), GatewayError> {
|
||||
if config.use_ssl {
|
||||
let stream = connect_tcp_stream(&config)?;
|
||||
let tls_stream = wrap_tls_stream(stream, &config.host)?;
|
||||
let mut reader = std::io::BufReader::new(tls_stream);
|
||||
let _ = smtp_expect(&mut reader, &[220])?;
|
||||
return smtp_send_message(&mut reader, &config, &email);
|
||||
}
|
||||
|
||||
let stream = connect_tcp_stream(&config)?;
|
||||
let mut reader = std::io::BufReader::new(stream);
|
||||
let _ = smtp_expect(&mut reader, &[220])?;
|
||||
let _ = smtp_send_command(&mut reader, "EHLO aether.local", &[250])?;
|
||||
@@ -705,16 +721,17 @@ fn send_smtp_email_blocking(
|
||||
smtp_deliver_message(&mut reader, &config, &email)
|
||||
}
|
||||
|
||||
fn probe_smtp_connection_blocking(config: SmtpDeliveryConfig) -> Result<(), GatewayError> {
|
||||
fn probe_smtp_connection_blocking(
|
||||
config: SmtpDeliveryConfig,
|
||||
stream: std::net::TcpStream,
|
||||
) -> Result<(), GatewayError> {
|
||||
if config.use_ssl {
|
||||
let stream = connect_tcp_stream(&config)?;
|
||||
let tls_stream = wrap_tls_stream(stream, &config.host)?;
|
||||
let mut reader = std::io::BufReader::new(tls_stream);
|
||||
let _ = smtp_expect(&mut reader, &[220])?;
|
||||
return smtp_probe_connection(&mut reader, &config);
|
||||
}
|
||||
|
||||
let stream = connect_tcp_stream(&config)?;
|
||||
let mut reader = std::io::BufReader::new(stream);
|
||||
let _ = smtp_expect(&mut reader, &[220])?;
|
||||
let _ = smtp_send_command(&mut reader, "EHLO aether.local", &[250])?;
|
||||
@@ -778,6 +795,140 @@ mod tests {
|
||||
assert!(validate_smtp_delivery_inputs(&config(), &email()).is_ok());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn smtp_connection_deadline_includes_a_stalled_dns_lookup() {
|
||||
let error = connect_tcp_stream_with_dns(
|
||||
std::future::pending(),
|
||||
std::time::Duration::from_millis(5),
|
||||
)
|
||||
.await
|
||||
.expect_err("DNS must not outlive the connection deadline");
|
||||
assert!(format!("{error:?}").contains("smtp DNS or TCP connection timed out"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn smtp_dns_errors_and_empty_answers_fail_without_connecting() {
|
||||
for (addresses, expected) in [
|
||||
(Ok(Vec::new()), "smtp host did not resolve to an address"),
|
||||
(
|
||||
Err(std::io::Error::other("sensitive-dns-detail")),
|
||||
"smtp DNS resolution failed",
|
||||
),
|
||||
(
|
||||
Err(std::io::Error::from(std::io::ErrorKind::InvalidData)),
|
||||
"smtp DNS resolution returned too many addresses",
|
||||
),
|
||||
(
|
||||
Err(std::io::Error::from(std::io::ErrorKind::TimedOut)),
|
||||
"smtp DNS resolution timed out",
|
||||
),
|
||||
] {
|
||||
let error = connect_tcp_stream_with_dns(
|
||||
std::future::ready(addresses),
|
||||
std::time::Duration::from_secs(1),
|
||||
)
|
||||
.await
|
||||
.expect_err("invalid DNS answers must fail before TCP connect");
|
||||
assert!(format!("{error:?}").contains(expected));
|
||||
assert!(!format!("{error:?}").contains("sensitive-dns-detail"));
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn smtp_connection_tries_answers_beyond_the_old_sixteen_address_limit() {
|
||||
let unavailable = tokio::net::TcpSocket::new_v4().unwrap();
|
||||
unavailable.bind("127.0.0.1:0".parse().unwrap()).unwrap();
|
||||
let listener = crate::test_support::bind_loopback_listener().await.unwrap();
|
||||
let available = listener.local_addr().unwrap();
|
||||
let mut addresses = vec![unavailable.local_addr().unwrap(); 16];
|
||||
addresses.push(available);
|
||||
let stream = connect_tcp_stream_with_dns(
|
||||
std::future::ready(Ok(addresses)),
|
||||
std::time::Duration::from_secs(5),
|
||||
)
|
||||
.await
|
||||
.expect("later DNS answers should remain available for fallback");
|
||||
assert_eq!(stream.peer_addr().unwrap(), available);
|
||||
assert_eq!(
|
||||
stream.read_timeout().unwrap(),
|
||||
Some(std::time::Duration::from_secs(SMTP_TIMEOUT_SECS))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn smtp_probe_and_delivery_use_the_preconnected_stream() {
|
||||
use tokio::io::{AsyncBufReadExt, AsyncWriteExt};
|
||||
|
||||
for deliver in [false, true] {
|
||||
let listener = crate::test_support::bind_loopback_listener().await.unwrap();
|
||||
let port = listener.local_addr().unwrap().port();
|
||||
let server = tokio::spawn(async move {
|
||||
let (stream, _) = listener.accept().await.unwrap();
|
||||
let mut reader = tokio::io::BufReader::new(stream);
|
||||
reader
|
||||
.get_mut()
|
||||
.write_all(b"220 mock SMTP ready\r\n")
|
||||
.await
|
||||
.unwrap();
|
||||
let mut delivered = false;
|
||||
loop {
|
||||
let mut line = String::new();
|
||||
assert!(reader.read_line(&mut line).await.unwrap() > 0);
|
||||
let response = if line.starts_with("EHLO ")
|
||||
|| line.starts_with("MAIL FROM:")
|
||||
|| line.starts_with("RCPT TO:")
|
||||
{
|
||||
&b"250 OK\r\n"[..]
|
||||
} else if line == "DATA\r\n" {
|
||||
reader
|
||||
.get_mut()
|
||||
.write_all(b"354 End with dot\r\n")
|
||||
.await
|
||||
.unwrap();
|
||||
loop {
|
||||
line.clear();
|
||||
assert!(reader.read_line(&mut line).await.unwrap() > 0);
|
||||
if line == ".\r\n" {
|
||||
break;
|
||||
}
|
||||
}
|
||||
delivered = true;
|
||||
&b"250 Accepted\r\n"[..]
|
||||
} else {
|
||||
assert_eq!(line, "QUIT\r\n");
|
||||
reader
|
||||
.get_mut()
|
||||
.write_all(b"221 Goodbye\r\n")
|
||||
.await
|
||||
.unwrap();
|
||||
break;
|
||||
};
|
||||
reader.get_mut().write_all(response).await.unwrap();
|
||||
}
|
||||
assert_eq!(delivered, deliver);
|
||||
});
|
||||
let config = SmtpDeliveryConfig {
|
||||
host: "127.0.0.1".to_string(),
|
||||
port,
|
||||
user: None,
|
||||
password: None,
|
||||
use_tls: false,
|
||||
use_ssl: false,
|
||||
..config()
|
||||
};
|
||||
tokio::time::timeout(std::time::Duration::from_secs(5), async {
|
||||
if deliver {
|
||||
send_smtp_email(config, email()).await.unwrap();
|
||||
} else {
|
||||
probe_smtp_connection(config).await.unwrap();
|
||||
}
|
||||
server.await.unwrap();
|
||||
})
|
||||
.await
|
||||
.expect("local SMTP probe and delivery should complete");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_authentication_over_plaintext_smtp() {
|
||||
let mut insecure = config();
|
||||
|
||||
@@ -68,7 +68,6 @@ const CHATGPT_WEB_IMAGE_PUBLIC_CONNECT_TIMEOUT_MS: u64 = 10_000;
|
||||
const CHATGPT_WEB_IMAGE_PUBLIC_READ_TIMEOUT_MS: u64 = 30_000;
|
||||
const CHATGPT_WEB_IMAGE_PUBLIC_TOTAL_TIMEOUT_MS: u64 = 300_000;
|
||||
const CHATGPT_WEB_OPAQUE_ID_MAX_BYTES: usize = 256;
|
||||
const CHATGPT_WEB_IMAGE_MAX_RESOLVED_ADDRESSES: usize = 32;
|
||||
const CHATGPT_WEB_IMAGE_MAX_UPLOAD_URL_BYTES: usize = 64 * 1024;
|
||||
const CHATGPT_WEB_IMAGE_UPLOAD_RESPONSE_LIMIT_BYTES: usize = 64 * 1024;
|
||||
const CHATGPT_WEB_IMAGE_MAX_PROMPT_BYTES: usize = 32 * 1024;
|
||||
@@ -1334,24 +1333,18 @@ async fn resolve_public_web_image_addrs(
|
||||
"ChatGPT-Web image URL is missing a port".to_string(),
|
||||
)
|
||||
})?;
|
||||
let resolved = if let Ok(ip) = host.parse::<IpAddr>() {
|
||||
vec![SocketAddr::new(ip, port)]
|
||||
} else {
|
||||
tokio::time::timeout(lookup_timeout, tokio::net::lookup_host((host, port)))
|
||||
.await
|
||||
.map_err(|_| {
|
||||
ExecutionRuntimeTransportError::UpstreamRequest(
|
||||
"ChatGPT-Web image URL DNS resolution timed out".to_string(),
|
||||
)
|
||||
})?
|
||||
.map_err(|err| {
|
||||
ExecutionRuntimeTransportError::UpstreamRequest(format!(
|
||||
"ChatGPT-Web image URL DNS resolution failed: {err}"
|
||||
))
|
||||
})?
|
||||
.take(CHATGPT_WEB_IMAGE_MAX_RESOLVED_ADDRESSES)
|
||||
.collect::<Vec<_>>()
|
||||
};
|
||||
let resolved = aether_http::lookup_host_with_limits(host, port, lookup_timeout)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
let message = match error.kind() {
|
||||
std::io::ErrorKind::TimedOut => "ChatGPT-Web image URL DNS resolution timed out",
|
||||
std::io::ErrorKind::InvalidData => {
|
||||
"ChatGPT-Web image URL DNS resolution returned too many addresses"
|
||||
}
|
||||
_ => "ChatGPT-Web image URL DNS resolution failed",
|
||||
};
|
||||
ExecutionRuntimeTransportError::UpstreamRequest(message.to_string())
|
||||
})?;
|
||||
if resolved.is_empty() {
|
||||
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
|
||||
"ChatGPT-Web image URL DNS resolution returned no addresses".to_string(),
|
||||
|
||||
@@ -1826,12 +1826,12 @@ async fn fetch_grok_attachment_url(
|
||||
// a fragment from the previous URL, while an absolute Location can
|
||||
// introduce either explicitly.
|
||||
validate_grok_attachment_url(&url)?;
|
||||
let public_addr = public_socket_addr_for_url(&url).await?;
|
||||
let public_addrs = public_socket_addrs_for_url(&url).await?;
|
||||
let response = reqwest::Client::builder()
|
||||
.no_proxy()
|
||||
.timeout(Duration::from_secs(30))
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.resolve_to_addrs(url.host_str().unwrap_or_default(), &[public_addr])
|
||||
.resolve_to_addrs(url.host_str().unwrap_or_default(), &public_addrs)
|
||||
.build()
|
||||
.map_err(ExecutionRuntimeTransportError::ClientBuild)?
|
||||
.get(url.clone())
|
||||
@@ -1897,10 +1897,10 @@ fn validate_grok_attachment_url(url: &reqwest::Url) -> Result<(), ExecutionRunti
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn public_socket_addr_for_url(
|
||||
async fn public_socket_addrs_for_url(
|
||||
url: &reqwest::Url,
|
||||
) -> Result<std::net::SocketAddr, ExecutionRuntimeTransportError> {
|
||||
let host = url.host().ok_or_else(|| {
|
||||
) -> Result<Vec<std::net::SocketAddr>, ExecutionRuntimeTransportError> {
|
||||
let host = url.host_str().ok_or_else(|| {
|
||||
ExecutionRuntimeTransportError::UpstreamRequest(
|
||||
"Grok attachment URL is missing a host".to_string(),
|
||||
)
|
||||
@@ -1910,64 +1910,34 @@ async fn public_socket_addr_for_url(
|
||||
"Grok attachment URL is missing a port".to_string(),
|
||||
)
|
||||
})?;
|
||||
let host = match host {
|
||||
url::Host::Ipv4(ip) => {
|
||||
let ip = IpAddr::V4(ip);
|
||||
if !grok_attachment_ip_is_public(ip) {
|
||||
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
|
||||
"Grok attachment URL resolves to a non-public address".to_string(),
|
||||
));
|
||||
}
|
||||
return Ok(std::net::SocketAddr::new(ip, port));
|
||||
}
|
||||
url::Host::Ipv6(ip) => {
|
||||
let ip = IpAddr::V6(ip);
|
||||
if !grok_attachment_ip_is_public(ip) {
|
||||
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
|
||||
"Grok attachment URL resolves to a non-public address".to_string(),
|
||||
));
|
||||
}
|
||||
return Ok(std::net::SocketAddr::new(ip, port));
|
||||
}
|
||||
url::Host::Domain(host) => host,
|
||||
};
|
||||
if let Ok(ip) = host.parse::<IpAddr>() {
|
||||
if !grok_attachment_ip_is_public(ip) {
|
||||
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
|
||||
"Grok attachment URL resolves to a non-public address".to_string(),
|
||||
));
|
||||
}
|
||||
return Ok(std::net::SocketAddr::new(ip, port));
|
||||
}
|
||||
let mut public_addr = None;
|
||||
let mut resolved_any = false;
|
||||
for addr in
|
||||
let addresses =
|
||||
aether_http::lookup_host_with_limits(host, port, aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT)
|
||||
.await
|
||||
.map_err(|err| {
|
||||
ExecutionRuntimeTransportError::UpstreamRequest(format!(
|
||||
"Grok attachment URL DNS resolution failed: {err}"
|
||||
))
|
||||
})?
|
||||
{
|
||||
resolved_any = true;
|
||||
if !grok_attachment_ip_is_public(addr.ip()) {
|
||||
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
|
||||
"Grok attachment URL resolves to a non-public address".to_string(),
|
||||
));
|
||||
}
|
||||
public_addr.get_or_insert(addr);
|
||||
}
|
||||
if !resolved_any {
|
||||
})?;
|
||||
validate_grok_attachment_addresses(addresses)
|
||||
}
|
||||
|
||||
fn validate_grok_attachment_addresses(
|
||||
addresses: Vec<std::net::SocketAddr>,
|
||||
) -> Result<Vec<std::net::SocketAddr>, ExecutionRuntimeTransportError> {
|
||||
if addresses.is_empty() {
|
||||
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
|
||||
"Grok attachment URL DNS resolution returned no addresses".to_string(),
|
||||
));
|
||||
}
|
||||
public_addr.ok_or_else(|| {
|
||||
ExecutionRuntimeTransportError::UpstreamRequest(
|
||||
"Grok attachment URL has no public address".to_string(),
|
||||
)
|
||||
})
|
||||
if addresses
|
||||
.iter()
|
||||
.any(|address| !grok_attachment_ip_is_public(address.ip()))
|
||||
{
|
||||
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
|
||||
"Grok attachment URL resolves to a non-public address".to_string(),
|
||||
));
|
||||
}
|
||||
Ok(addresses)
|
||||
}
|
||||
|
||||
fn grok_attachment_ip_is_public(ip: IpAddr) -> bool {
|
||||
@@ -3898,7 +3868,7 @@ mod tests {
|
||||
grok_should_use_imagine_websocket, grok_success_frame_stream, grok_upload_url,
|
||||
grok_upstream_model_name, grok_usage_estimate, grok_user_id_from_cookie_header,
|
||||
materialize_grok_image_assets, maximum_base64_len_for_decoded_limit, openai_chat_body,
|
||||
openai_image_body, openai_image_sse, openai_responses_body, public_socket_addr_for_url,
|
||||
openai_image_body, openai_image_sse, openai_responses_body, public_socket_addrs_for_url,
|
||||
set_grok_image_edit_config, validate_grok_attachment_url, GrokAttachmentInput,
|
||||
GrokCollected, GrokImagineImage, GrokStreamAdapter,
|
||||
};
|
||||
@@ -4520,7 +4490,7 @@ mod tests {
|
||||
] {
|
||||
let url = reqwest::Url::parse(raw_url).expect("URL should parse");
|
||||
assert!(
|
||||
public_socket_addr_for_url(&url).await.is_err(),
|
||||
public_socket_addrs_for_url(&url).await.is_err(),
|
||||
"private IPv6 literal should be rejected: {raw_url}"
|
||||
);
|
||||
}
|
||||
@@ -4528,13 +4498,31 @@ mod tests {
|
||||
let url = reqwest::Url::parse("https://[2606:4700:4700::1111]/attachment")
|
||||
.expect("URL should parse");
|
||||
assert_eq!(
|
||||
public_socket_addr_for_url(&url)
|
||||
public_socket_addrs_for_url(&url)
|
||||
.await
|
||||
.expect("public IPv6 literal should pass"),
|
||||
"[2606:4700:4700::1111]:443".parse().unwrap()
|
||||
vec!["[2606:4700:4700::1111]:443".parse().unwrap()]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grok_attachment_dns_keeps_all_safe_addresses_for_connection_fallback() {
|
||||
let addresses = vec![
|
||||
"[2606:4700:4700::1111]:443".parse().unwrap(),
|
||||
"8.8.8.8:443".parse().unwrap(),
|
||||
];
|
||||
assert_eq!(
|
||||
super::validate_grok_attachment_addresses(addresses.clone()).unwrap(),
|
||||
addresses
|
||||
);
|
||||
assert!(super::validate_grok_attachment_addresses(Vec::new()).is_err());
|
||||
for blocked in ["198.18.0.1:443", "127.0.0.1:443", "[fd00::1]:443"] {
|
||||
let mut mixed = addresses.clone();
|
||||
mixed.push(blocked.parse().unwrap());
|
||||
assert!(super::validate_grok_attachment_addresses(mixed).is_err());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grok_attachment_url_rejects_credentials_and_fragments_on_every_hop() {
|
||||
for raw_url in [
|
||||
|
||||
@@ -438,7 +438,7 @@ static DIRECT_H2C_SENDER_CACHE_METRICS: LazyLock<DirectHyperH2cSenderCacheMetric
|
||||
LazyLock::new(DirectHyperH2cSenderCacheMetrics::default);
|
||||
|
||||
#[derive(Debug, Clone, Copy, Default)]
|
||||
struct ExecutionSafeDnsResolver;
|
||||
pub(crate) struct ExecutionSafeDnsResolver;
|
||||
|
||||
#[derive(Debug, Clone, Copy, Default)]
|
||||
struct ExecutionSafeHyperDnsResolver;
|
||||
@@ -446,10 +446,7 @@ struct ExecutionSafeHyperDnsResolver;
|
||||
fn dns_host_explicitly_allows_loopback(host: &str) -> bool {
|
||||
let host = host.trim_end_matches('.');
|
||||
host.eq_ignore_ascii_case("localhost")
|
||||
|| host
|
||||
.parse::<IpAddr>()
|
||||
.map(|ip| ip.is_loopback())
|
||||
.unwrap_or(false)
|
||||
|| aether_http::parse_ip_literal_host(host).is_some_and(|ip| ip.is_loopback())
|
||||
}
|
||||
|
||||
fn validate_resolved_execution_addresses(
|
||||
@@ -491,12 +488,9 @@ async fn resolve_execution_target_addresses_with_policy(
|
||||
port: u16,
|
||||
provider_execution: bool,
|
||||
) -> Result<Vec<SocketAddr>, std::io::Error> {
|
||||
let addresses = if let Ok(ip) = host.parse::<IpAddr>() {
|
||||
vec![SocketAddr::new(ip, port)]
|
||||
} else {
|
||||
let addresses =
|
||||
aether_http::lookup_host_with_limits(host, port, aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT)
|
||||
.await?
|
||||
};
|
||||
.await?;
|
||||
validate_resolved_execution_addresses(host, addresses, provider_execution)
|
||||
}
|
||||
|
||||
@@ -5149,7 +5143,7 @@ fn execution_log_url_host(url: &str) -> String {
|
||||
.unwrap_or_else(|| "-".to_string())
|
||||
}
|
||||
|
||||
fn validate_execution_upstream_url(
|
||||
pub(crate) fn validate_execution_upstream_url(
|
||||
raw_url: &str,
|
||||
) -> Result<url::Url, ExecutionRuntimeTransportError> {
|
||||
let url = url::Url::parse(raw_url).map_err(|_| {
|
||||
@@ -5440,6 +5434,8 @@ mod tests {
|
||||
"93.184.216.34:443".parse().unwrap(),
|
||||
];
|
||||
for host in [
|
||||
"chatgpt.com",
|
||||
"api.openai.com",
|
||||
"oauth2.googleapis.com",
|
||||
"www.googleapis.com",
|
||||
"custom.example.test",
|
||||
@@ -5452,6 +5448,46 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn execution_dns_handles_url_ipv6_without_weakening_relay_filtering() {
|
||||
for provider_execution in [false, true] {
|
||||
let addresses = super::resolve_execution_target_addresses_with_policy(
|
||||
"[::1]",
|
||||
8443,
|
||||
provider_execution,
|
||||
)
|
||||
.await
|
||||
.expect("literal IPv6 loopback should resolve without DNS");
|
||||
assert_eq!(addresses, vec!["[::1]:8443".parse().unwrap()]);
|
||||
}
|
||||
let error = super::resolve_execution_target_addresses_with_policy("[fd00::1]", 443, false)
|
||||
.await
|
||||
.expect_err("private IPv6 must remain blocked for relay traffic");
|
||||
assert_eq!(error.kind(), std::io::ErrorKind::PermissionDenied);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn execution_dns_resolvers_preserve_provider_fake_ip_answers() {
|
||||
for host in ["198.18.78.41", "198.19.1.2"] {
|
||||
let expected = vec![format!("{host}:0").parse::<std::net::SocketAddr>().unwrap()];
|
||||
let reqwest_addresses = reqwest::dns::Resolve::resolve(
|
||||
&super::ExecutionSafeDnsResolver,
|
||||
host.parse().unwrap(),
|
||||
)
|
||||
.await
|
||||
.expect("HTTP provider DNS must accept Fake-IP answers")
|
||||
.collect::<Vec<_>>();
|
||||
let wreq_addresses =
|
||||
wreq::dns::Resolve::resolve(&super::ExecutionSafeDnsResolver, host.into())
|
||||
.await
|
||||
.expect("WebSocket provider DNS must accept Fake-IP answers")
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
assert_eq!(reqwest_addresses, expected);
|
||||
assert_eq!(wreq_addresses, expected);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn execution_dns_answers_keep_relay_address_filtering() {
|
||||
let public = "93.184.216.34:443".parse().unwrap();
|
||||
|
||||
@@ -24,7 +24,7 @@ use wreq::ws::message::{CloseFrame as WreqCloseFrame, Message as WreqWsMessage};
|
||||
use crate::ai_serving::AiExecutionDecision;
|
||||
use crate::execution_runtime::transport::{
|
||||
build_browser_wreq_client, build_request_headers, normalize_execution_proxy_url,
|
||||
ExecutionTransportControls,
|
||||
validate_execution_upstream_url, ExecutionSafeDnsResolver, ExecutionTransportControls,
|
||||
};
|
||||
use crate::frontdoor_loop_guard::gateway_frontdoor_self_loop_guard_error;
|
||||
use crate::handlers::proxy::websocket::session::{
|
||||
@@ -66,7 +66,7 @@ pub(crate) async fn connect_upstream_websocket(
|
||||
)?;
|
||||
let headers =
|
||||
websocket_handshake_headers(&decision.provider_request_headers, errors.headers_invalid)?;
|
||||
let client = build_websocket_client(decision, &upstream_url, errors).await?;
|
||||
let client = build_websocket_client(decision, errors)?;
|
||||
let response = client
|
||||
.websocket(upstream_url.as_str())
|
||||
.headers(headers)
|
||||
@@ -149,31 +149,14 @@ pub(crate) fn websocket_upstream_url(
|
||||
invalid_code: &'static str,
|
||||
) -> Result<Url, &'static str> {
|
||||
let mut url = Url::parse(raw).map_err(|_| invalid_code)?;
|
||||
if url.host_str().is_none()
|
||||
|| !url.username().is_empty()
|
||||
|| url.password().is_some()
|
||||
|| url.fragment().is_some()
|
||||
{
|
||||
return Err(invalid_code);
|
||||
}
|
||||
let websocket_scheme = match url.scheme() {
|
||||
"https" | "wss" => "wss",
|
||||
"http" | "ws" => "ws",
|
||||
let (http_scheme, websocket_scheme) = match url.scheme() {
|
||||
"https" | "wss" => ("https", "wss"),
|
||||
"http" | "ws" => ("http", "ws"),
|
||||
_ => return Err(invalid_code),
|
||||
};
|
||||
url.set_scheme(http_scheme).map_err(|_| invalid_code)?;
|
||||
let mut url = validate_execution_upstream_url(url.as_str()).map_err(|_| invalid_code)?;
|
||||
url.set_scheme(websocket_scheme).map_err(|_| invalid_code)?;
|
||||
if url.scheme() == "ws" {
|
||||
let literal_ip = match url.host() {
|
||||
Some(url::Host::Ipv4(address)) => Some(std::net::IpAddr::V4(address)),
|
||||
Some(url::Host::Ipv6(address)) => Some(std::net::IpAddr::V6(address)),
|
||||
_ => None,
|
||||
};
|
||||
if literal_ip.is_some_and(|address| {
|
||||
aether_http::is_private_or_reserved_ip(address) && !address.is_loopback()
|
||||
}) {
|
||||
return Err(invalid_code);
|
||||
}
|
||||
}
|
||||
Ok(url)
|
||||
}
|
||||
|
||||
@@ -229,9 +212,8 @@ pub(crate) fn websocket_handshake_headers(
|
||||
Ok(headers)
|
||||
}
|
||||
|
||||
async fn build_websocket_client(
|
||||
fn build_websocket_client(
|
||||
decision: &AiExecutionDecision,
|
||||
upstream_url: &Url,
|
||||
errors: UpstreamWebSocketErrorCodes,
|
||||
) -> Result<wreq::Client, &'static str> {
|
||||
let timeouts = websocket_timeouts(decision);
|
||||
@@ -255,41 +237,7 @@ async fn build_websocket_client(
|
||||
let proxy = wreq::Proxy::all(proxy_url).map_err(|_| errors.proxy_invalid)?;
|
||||
builder = builder.proxy(proxy);
|
||||
} else {
|
||||
// Pin every direct WebSocket connection to the DNS answers validated
|
||||
// here. This also covers the explicitly permitted loopback `ws://`
|
||||
// form; otherwise the client would perform a second lookup and a
|
||||
// rebinding could escape the loopback-only policy.
|
||||
let host = upstream_url.host_str().ok_or(errors.upstream_url_invalid)?;
|
||||
let port = upstream_url
|
||||
.port_or_known_default()
|
||||
.ok_or(errors.upstream_url_invalid)?;
|
||||
let addresses = if let Ok(ip) = host.parse::<std::net::IpAddr>() {
|
||||
vec![std::net::SocketAddr::new(ip, port)]
|
||||
} else {
|
||||
aether_http::lookup_host_with_limits(
|
||||
host,
|
||||
port,
|
||||
aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT,
|
||||
)
|
||||
.await
|
||||
.map_err(|_| errors.upstream_url_invalid)?
|
||||
};
|
||||
let allows_loopback = host.trim_end_matches('.').eq_ignore_ascii_case("localhost")
|
||||
|| host
|
||||
.parse::<std::net::IpAddr>()
|
||||
.map(|ip| ip.is_loopback())
|
||||
.unwrap_or(false);
|
||||
let unsafe_answer = if allows_loopback {
|
||||
addresses.iter().any(|address| !address.ip().is_loopback())
|
||||
} else {
|
||||
addresses
|
||||
.iter()
|
||||
.any(|address| aether_http::is_private_or_reserved_ip(address.ip()))
|
||||
};
|
||||
if addresses.is_empty() || unsafe_answer {
|
||||
return Err(errors.upstream_url_invalid);
|
||||
}
|
||||
builder = builder.resolve_to_addrs(host.to_string(), addresses.iter().copied());
|
||||
builder = builder.dns_resolver(ExecutionSafeDnsResolver);
|
||||
}
|
||||
builder.build().map_err(|_| errors.client_build_failed)
|
||||
}
|
||||
@@ -684,15 +632,17 @@ pub(crate) async fn close_client_socket(client_socket: &mut WebSocket, code: u16
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
bounded_send, guarded_websocket_upstream_url, resolve_websocket_proxy_url,
|
||||
responses_websocket_error_event, responses_websocket_error_event_with_stream_id,
|
||||
websocket_handshake_headers, websocket_relay_frame_queue, websocket_response_headers,
|
||||
websocket_upstream_url, UpstreamWebSocketErrorCodes, WebSocketRelayPumpControl,
|
||||
WebSocketRelayQueueError, WebSocketWriteError, RELAY_FRAME_QUEUE_CAPACITY,
|
||||
RELAY_WRITE_TIMEOUT, TEARDOWN_WRITE_TIMEOUT,
|
||||
bounded_send, build_websocket_client, guarded_websocket_upstream_url,
|
||||
resolve_websocket_proxy_url, responses_websocket_error_event,
|
||||
responses_websocket_error_event_with_stream_id, websocket_handshake_headers,
|
||||
websocket_relay_frame_queue, websocket_response_headers, websocket_upstream_url,
|
||||
UpstreamWebSocketErrorCodes, WebSocketRelayPumpControl, WebSocketRelayQueueError,
|
||||
WebSocketWriteError, RELAY_FRAME_QUEUE_CAPACITY, RELAY_WRITE_TIMEOUT,
|
||||
TEARDOWN_WRITE_TIMEOUT,
|
||||
};
|
||||
use crate::ai_serving::AiExecutionDecision;
|
||||
use crate::frontdoor_loop_guard::configured_gateway_frontdoor_base_url;
|
||||
use aether_contracts::ProxySnapshot;
|
||||
use aether_contracts::{ProxySnapshot, ResolvedTransportProfile};
|
||||
use axum::http::HeaderMap;
|
||||
use std::collections::BTreeMap;
|
||||
use std::time::Duration;
|
||||
@@ -876,7 +826,9 @@ mod tests {
|
||||
"ws://example.test:8080/v1/responses",
|
||||
"http://example.test:8080/v1/responses",
|
||||
"http://8.8.8.8:8080/v1/responses",
|
||||
"wss://8.8.8.8/v1/responses",
|
||||
"ws://[2606:4700:4700::1111]:8080/v1/responses",
|
||||
"wss://[2606:4700:4700::1111]/v1/responses",
|
||||
"ws://localhost:8080/v1/responses",
|
||||
"http://127.42.0.1:8080/v1/responses",
|
||||
"ws://[::1]:8080/v1/responses",
|
||||
@@ -888,6 +840,14 @@ mod tests {
|
||||
}
|
||||
for rejected in [
|
||||
"http://10.0.0.1/v1/responses",
|
||||
"wss://10.0.0.1/v1/responses",
|
||||
"wss://127.0.0.1/v1/responses",
|
||||
"wss://[::1]/v1/responses",
|
||||
"wss://[fd00::1]/v1/responses",
|
||||
"wss://[::ffff:127.0.0.1]/v1/responses",
|
||||
"wss://169.254.169.254/v1/responses",
|
||||
"wss://198.18.78.41/v1/responses",
|
||||
"wss://198.19.1.2/v1/responses",
|
||||
"ws://0.0.0.0:8080/v1/responses",
|
||||
"ws://[::ffff:127.0.0.1]:8080/v1/responses",
|
||||
"wss://example.test/v1/responses#secret",
|
||||
@@ -903,6 +863,60 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn websocket_client_build_defers_provider_dns_for_all_transport_profiles() {
|
||||
let errors = UpstreamWebSocketErrorCodes {
|
||||
upstream_url_missing: "missing",
|
||||
upstream_url_invalid: "upstream_invalid",
|
||||
frontdoor_self_loop: "frontdoor_self_loop",
|
||||
headers_invalid: "headers_invalid",
|
||||
client_build_failed: "client_build_failed",
|
||||
proxy_invalid: "proxy_invalid",
|
||||
tunnel_proxy_unsupported: "tunnel_unsupported",
|
||||
handshake_failed: "handshake_failed",
|
||||
upgrade_rejected: "upgrade_rejected",
|
||||
upgrade_failed: "upgrade_failed",
|
||||
};
|
||||
for profile in [
|
||||
None,
|
||||
Some(ResolvedTransportProfile {
|
||||
profile_id: "chrome136".to_string(),
|
||||
backend: aether_contracts::TRANSPORT_BACKEND_BROWSER_WREQ.to_string(),
|
||||
..Default::default()
|
||||
}),
|
||||
] {
|
||||
for proxy in [
|
||||
None,
|
||||
Some(ProxySnapshot {
|
||||
enabled: Some(false),
|
||||
url: Some("http://proxy.invalid:8080".to_string()),
|
||||
..Default::default()
|
||||
}),
|
||||
Some(ProxySnapshot {
|
||||
enabled: Some(true),
|
||||
url: Some("http://proxy.invalid:8080".to_string()),
|
||||
..Default::default()
|
||||
}),
|
||||
Some(ProxySnapshot {
|
||||
enabled: Some(true),
|
||||
url: Some("socks5h://proxy.invalid:1080".to_string()),
|
||||
..Default::default()
|
||||
}),
|
||||
] {
|
||||
let mut decision: AiExecutionDecision = serde_json::from_value(serde_json::json!({
|
||||
"action": "proxy",
|
||||
"upstream_url": "wss://upstream.invalid/v1/responses"
|
||||
}))
|
||||
.expect("minimal provider decision should deserialize");
|
||||
decision.transport_profile = profile.clone();
|
||||
decision.proxy = proxy;
|
||||
|
||||
build_websocket_client(&decision, errors)
|
||||
.expect("building a client must not resolve the provider or proxy hostname");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn active_websocket_proxy_without_a_target_fails_closed() {
|
||||
let errors = UpstreamWebSocketErrorCodes {
|
||||
@@ -937,6 +951,94 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn websocket_handshake_keeps_provider_dns_remote_for_http_and_socks_proxies() {
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
|
||||
let errors = UpstreamWebSocketErrorCodes {
|
||||
upstream_url_missing: "missing",
|
||||
upstream_url_invalid: "upstream_invalid",
|
||||
frontdoor_self_loop: "frontdoor_self_loop",
|
||||
headers_invalid: "headers_invalid",
|
||||
client_build_failed: "client_build_failed",
|
||||
proxy_invalid: "proxy_invalid",
|
||||
tunnel_proxy_unsupported: "tunnel_unsupported",
|
||||
handshake_failed: "handshake_failed",
|
||||
upgrade_rejected: "upgrade_rejected",
|
||||
upgrade_failed: "upgrade_failed",
|
||||
};
|
||||
for profile in [
|
||||
None,
|
||||
Some(ResolvedTransportProfile {
|
||||
profile_id: "chrome136".to_string(),
|
||||
backend: aether_contracts::TRANSPORT_BACKEND_BROWSER_WREQ.to_string(),
|
||||
..Default::default()
|
||||
}),
|
||||
] {
|
||||
for scheme in ["http", "socks5", "socks5h"] {
|
||||
let listener = crate::test_support::bind_loopback_listener().await.unwrap();
|
||||
let proxy_addr = listener.local_addr().unwrap();
|
||||
let (release, released) = tokio::sync::oneshot::channel::<()>();
|
||||
let server = tokio::spawn(async move {
|
||||
let (mut stream, _) = listener.accept().await.unwrap();
|
||||
if scheme != "http" {
|
||||
let mut greeting = [0; 2];
|
||||
stream.read_exact(&mut greeting).await.unwrap();
|
||||
assert_eq!(greeting[0], 5);
|
||||
let mut methods = vec![0; greeting[1] as usize];
|
||||
stream.read_exact(&mut methods).await.unwrap();
|
||||
assert!(methods.contains(&0));
|
||||
stream.write_all(&[5, 0]).await.unwrap();
|
||||
|
||||
let mut request = [0; 4];
|
||||
stream.read_exact(&mut request).await.unwrap();
|
||||
assert_eq!(
|
||||
request,
|
||||
[5, 1, 0, 3],
|
||||
"proxy must receive a domain, not an IP"
|
||||
);
|
||||
let host_len = stream.read_u8().await.unwrap();
|
||||
let mut host = vec![0; host_len as usize];
|
||||
stream.read_exact(&mut host).await.unwrap();
|
||||
assert_eq!(host, b"provider-dns.invalid");
|
||||
assert_eq!(stream.read_u16().await.unwrap(), 80);
|
||||
stream
|
||||
.write_all(&[5, 0, 0, 1, 127, 0, 0, 1, 0, 80])
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
let socket = tokio_tungstenite::accept_async(stream).await.unwrap();
|
||||
let _ = released.await;
|
||||
drop(socket);
|
||||
});
|
||||
let mut decision: AiExecutionDecision = serde_json::from_value(serde_json::json!({
|
||||
"action": "proxy",
|
||||
"upstream_url": "ws://provider-dns.invalid/v1/responses",
|
||||
"proxy": {"enabled": true, "url": format!("{scheme}://{proxy_addr}")}
|
||||
}))
|
||||
.unwrap();
|
||||
decision.transport_profile = profile.clone();
|
||||
let connection = tokio::time::timeout(
|
||||
Duration::from_secs(5),
|
||||
super::connect_upstream_websocket(
|
||||
&decision,
|
||||
crate::handlers::proxy::websocket::session::RESPONSES_WEBSOCKET_SESSION_LIMITS,
|
||||
errors,
|
||||
),
|
||||
)
|
||||
.await
|
||||
.expect("proxied handshake must not wait for local provider DNS")
|
||||
.unwrap_or_else(|error| panic!("{scheme} handshake failed: {error}"));
|
||||
release.send(()).unwrap();
|
||||
tokio::time::timeout(Duration::from_secs(5), server)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
drop(connection);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_responses_websocket_frontdoor_self_loop_before_connecting() {
|
||||
let base_url = configured_gateway_frontdoor_base_url();
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::net::{IpAddr, SocketAddr};
|
||||
use std::sync::Arc;
|
||||
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
||||
|
||||
use axum::{
|
||||
@@ -10,6 +10,10 @@ use axum::{
|
||||
};
|
||||
use serde_json::json;
|
||||
|
||||
use crate::execution_runtime::transport::{
|
||||
validate_execution_upstream_url, ExecutionSafeDnsResolver,
|
||||
};
|
||||
|
||||
use super::test_connection_shared::select_test_connection_provider;
|
||||
use super::{
|
||||
provider_catalog_key_supports_format, query_param_value, AppState, GatewayPublicRequestContext,
|
||||
@@ -18,98 +22,16 @@ use super::{
|
||||
const DEFAULT_NON_STREAM_TOTAL_TIMEOUT_MS: u64 = 300_000;
|
||||
const MAX_TEST_CONNECTION_RESPONSE_BYTES: usize = 256 * 1024;
|
||||
|
||||
#[cfg(test)]
|
||||
fn build_test_connection_client() -> Result<reqwest::Client, reqwest::Error> {
|
||||
reqwest::Client::builder()
|
||||
.no_proxy()
|
||||
.dns_resolver(Arc::new(ExecutionSafeDnsResolver))
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.connect_timeout(Duration::from_secs(10))
|
||||
.http2_adaptive_window(true)
|
||||
.build()
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct ResolvedTestConnectionTarget {
|
||||
url: reqwest::Url,
|
||||
host: String,
|
||||
addresses: Vec<SocketAddr>,
|
||||
}
|
||||
|
||||
/// Resolve the provider endpoint once and pin reqwest to that answer. The
|
||||
/// test-connection route is reachable through the public front door, so it
|
||||
/// must not perform an unbounded DNS lookup on every connect (which would
|
||||
/// permit DNS rebinding into private/reserved networks).
|
||||
async fn resolve_test_connection_target(
|
||||
raw_url: &str,
|
||||
allow_private_targets: bool,
|
||||
) -> Result<ResolvedTestConnectionTarget, &'static str> {
|
||||
let url = reqwest::Url::parse(raw_url).map_err(|_| "provider endpoint URL is invalid")?;
|
||||
let literal_loopback = aether_http::url_has_literal_loopback_host(&url);
|
||||
if !matches!(url.scheme(), "http" | "https")
|
||||
|| url.host_str().is_none()
|
||||
|| !url.username().is_empty()
|
||||
|| url.password().is_some()
|
||||
|| url.fragment().is_some()
|
||||
{
|
||||
return Err("provider endpoint must be an HTTP(S) URL without credentials or fragment");
|
||||
}
|
||||
let host = url
|
||||
.host_str()
|
||||
.ok_or("provider endpoint is missing a host")?
|
||||
.to_string();
|
||||
let literal_ip = host.parse::<IpAddr>().ok();
|
||||
let port = url
|
||||
.port_or_known_default()
|
||||
.ok_or("provider endpoint is missing a port")?;
|
||||
let addresses = if let Some(ip) = literal_ip {
|
||||
vec![SocketAddr::new(ip, port)]
|
||||
} else {
|
||||
aether_http::lookup_host_with_limits(
|
||||
host.as_str(),
|
||||
port,
|
||||
aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT,
|
||||
)
|
||||
.await
|
||||
.map_err(|_| "provider endpoint DNS resolution failed")?
|
||||
};
|
||||
if addresses.is_empty() {
|
||||
return Err("provider endpoint DNS resolution returned no addresses");
|
||||
}
|
||||
let has_private_answer = addresses
|
||||
.iter()
|
||||
.any(|address| aether_http::is_private_or_reserved_ip(address.ip()));
|
||||
// `allow_private_targets` is only enabled for in-process test fixtures.
|
||||
// Keep that escape hatch narrowly scoped to literal loopback URLs whose
|
||||
// every DNS answer is loopback; otherwise a test-only build (or an
|
||||
// accidentally reused helper) could turn this public route into a
|
||||
// private-network HTTP client.
|
||||
let test_loopback_target = allow_private_targets
|
||||
&& literal_loopback
|
||||
&& addresses.iter().all(|address| address.ip().is_loopback());
|
||||
if has_private_answer && !test_loopback_target {
|
||||
return Err("provider endpoint resolves to a private or reserved address");
|
||||
}
|
||||
Ok(ResolvedTestConnectionTarget {
|
||||
url,
|
||||
host,
|
||||
addresses,
|
||||
})
|
||||
}
|
||||
|
||||
fn build_pinned_test_connection_client(
|
||||
target: &ResolvedTestConnectionTarget,
|
||||
) -> Result<reqwest::Client, reqwest::Error> {
|
||||
let mut builder = reqwest::Client::builder()
|
||||
.no_proxy()
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.connect_timeout(Duration::from_secs(10))
|
||||
.http2_adaptive_window(true);
|
||||
if target.host.parse::<IpAddr>().is_err() {
|
||||
builder = builder.resolve_to_addrs(&target.host, &target.addresses);
|
||||
}
|
||||
builder.build()
|
||||
}
|
||||
|
||||
pub(super) async fn maybe_build_local_test_connection_route_response(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
@@ -384,18 +306,14 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
|
||||
);
|
||||
}
|
||||
|
||||
// Resolve and pin the endpoint before constructing the request. This
|
||||
// keeps the public health-check route subject to the same DNS/SSRF
|
||||
// boundary as the main execution transport. Unit-test fixtures may use
|
||||
// loopback listeners; production requests never opt into private targets.
|
||||
let target = match resolve_test_connection_target(&upstream_url, cfg!(test)).await {
|
||||
Ok(target) => target,
|
||||
let upstream_url = match validate_execution_upstream_url(&upstream_url) {
|
||||
Ok(url) => url,
|
||||
Err(reason) => {
|
||||
tracing::warn!(
|
||||
event_name = "provider_test_connection_target_rejected",
|
||||
provider_id = %provider.id,
|
||||
endpoint_id = %endpoint.id,
|
||||
reason,
|
||||
reason = %reason,
|
||||
"provider connection test target was rejected"
|
||||
);
|
||||
return Some(
|
||||
@@ -407,7 +325,7 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
|
||||
);
|
||||
}
|
||||
};
|
||||
let test_client = match build_pinned_test_connection_client(&target) {
|
||||
let test_client = match build_test_connection_client() {
|
||||
Ok(client) => client,
|
||||
Err(_) => {
|
||||
return Some(
|
||||
@@ -419,7 +337,7 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
|
||||
);
|
||||
}
|
||||
};
|
||||
let mut upstream_request = test_client.post(target.url);
|
||||
let mut upstream_request = test_client.post(upstream_url);
|
||||
for (name, value) in &provider_request_headers {
|
||||
upstream_request = upstream_request.header(name, value);
|
||||
}
|
||||
@@ -495,7 +413,7 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{build_test_connection_client, resolve_test_connection_target};
|
||||
use super::{build_test_connection_client, validate_execution_upstream_url};
|
||||
use axum::{
|
||||
body::Body,
|
||||
http::{header, Request, StatusCode},
|
||||
@@ -567,74 +485,59 @@ mod tests {
|
||||
redirected_server.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_connection_target_rejects_private_addresses_in_production_mode() {
|
||||
#[test]
|
||||
fn test_connection_target_rejects_private_literals_like_provider_requests() {
|
||||
for raw_url in [
|
||||
"http://127.0.0.1:8080/v1/chat/completions",
|
||||
"http://10.0.0.1/v1/chat/completions",
|
||||
"http://169.254.169.254/v1/chat/completions",
|
||||
"https://10.0.0.1/v1/chat/completions",
|
||||
"https://127.0.0.1/v1/chat/completions",
|
||||
"https://[::1]/v1/chat/completions",
|
||||
"https://localhost/v1/chat/completions",
|
||||
"https://198.18.78.41/v1/chat/completions",
|
||||
] {
|
||||
assert!(
|
||||
resolve_test_connection_target(raw_url, false)
|
||||
.await
|
||||
.is_err(),
|
||||
validate_execution_upstream_url(raw_url).is_err(),
|
||||
"private provider target should be rejected: {raw_url}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_connection_target_accepts_public_http_and_https_addresses() {
|
||||
for allow_private_targets in [false, true] {
|
||||
for (raw_url, expected_port) in [
|
||||
("http://8.8.8.8/v1/chat", 80),
|
||||
("http://8.8.8.8:8080/v1/chat", 8080),
|
||||
("https://8.8.8.8/v1/chat", 443),
|
||||
] {
|
||||
let target = resolve_test_connection_target(raw_url, allow_private_targets)
|
||||
.await
|
||||
.expect("public HTTP(S) provider target should resolve");
|
||||
assert_eq!(target.url.as_str(), raw_url);
|
||||
assert_eq!(target.host, "8.8.8.8");
|
||||
assert_eq!(target.addresses.len(), 1);
|
||||
assert_eq!(target.addresses[0].ip().to_string(), "8.8.8.8");
|
||||
assert_eq!(target.addresses[0].port(), expected_port);
|
||||
}
|
||||
#[test]
|
||||
fn test_connection_target_accepts_public_http_and_https_addresses() {
|
||||
for (raw_url, expected_port) in [
|
||||
("http://8.8.8.8/v1/chat", 80),
|
||||
("http://8.8.8.8:8080/v1/chat", 8080),
|
||||
("https://8.8.8.8/v1/chat", 443),
|
||||
("https://[2606:4700:4700::1111]/v1/chat", 443),
|
||||
] {
|
||||
let url = validate_execution_upstream_url(raw_url)
|
||||
.expect("public HTTP(S) provider target should be valid");
|
||||
assert_eq!(url.as_str(), raw_url);
|
||||
assert_eq!(url.port_or_known_default(), Some(expected_port));
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_connection_target_allows_loopback_only_for_test_fixtures() {
|
||||
let target = resolve_test_connection_target("http://127.0.0.1:8080/v1/chat", true)
|
||||
.await
|
||||
.expect("test fixture target should resolve");
|
||||
assert_eq!(target.host, "127.0.0.1");
|
||||
assert_eq!(target.addresses.len(), 1);
|
||||
assert!(
|
||||
resolve_test_connection_target("http://10.0.0.1/v1/chat", true)
|
||||
.await
|
||||
.is_err(),
|
||||
"test mode must not make private non-loopback HTTP endpoints acceptable"
|
||||
);
|
||||
assert!(
|
||||
resolve_test_connection_target("https://10.0.0.1/v1/chat", true)
|
||||
.await
|
||||
.is_err(),
|
||||
"test mode must not make private non-loopback endpoints acceptable"
|
||||
);
|
||||
assert!(
|
||||
resolve_test_connection_target("http://localhost:8080/v1/chat", true)
|
||||
.await
|
||||
.is_ok(),
|
||||
"literal localhost should remain available for local fixtures"
|
||||
);
|
||||
async fn test_connection_target_defers_dns_and_accepts_provider_loopback_urls() {
|
||||
for raw_url in [
|
||||
"http://127.0.0.1:8080/v1/chat",
|
||||
"http://[::1]:8080/v1/chat",
|
||||
"http://localhost:8080/v1/chat",
|
||||
"https://provider-dns.invalid/v1/chat",
|
||||
] {
|
||||
let url = validate_execution_upstream_url(raw_url)
|
||||
.expect("target validation must not depend on the current DNS answer");
|
||||
let request = build_test_connection_client()
|
||||
.expect("client should build without DNS")
|
||||
.post(url)
|
||||
.build()
|
||||
.expect("provider request should build without DNS");
|
||||
assert_eq!(request.url().as_str(), raw_url);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_connection_target_rejects_url_credentials_and_fragments() {
|
||||
#[test]
|
||||
fn test_connection_target_rejects_url_credentials_and_fragments() {
|
||||
for raw_url in [
|
||||
"https://user:[email protected]/v1/chat",
|
||||
"https://example.com/v1/chat#fragment",
|
||||
@@ -643,9 +546,7 @@ mod tests {
|
||||
"ftp://example.com/v1/chat",
|
||||
] {
|
||||
assert!(
|
||||
resolve_test_connection_target(raw_url, false)
|
||||
.await
|
||||
.is_err(),
|
||||
validate_execution_upstream_url(raw_url).is_err(),
|
||||
"unsafe provider target should be rejected: {raw_url}"
|
||||
);
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user