mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-04 16:37:46 +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:
@@ -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