Skip to content

Commit f6d3b3d

Browse files
fix(provider-utils,providers): validate and pin provider-supplied download URLs
Providers fetch URLs taken from response bodies and headers — generated assets, polling and result URLs, the Google Files upload URL — through the shared auto-redirecting client with no target validation, so a compromised or spoofed provider response could steer authenticated requests at internal services (cloud metadata, loopback, RFC1918). The new send_validated_download applies AI SDK-aligned SSRF protection: - URL validation against AI SDK's validateUrl blocklists: http/https only, localhost/.local/.localhost rejected, literal-IP checks including the IPv4-embedded IPv6 forms (mapped, compatible, SIIT, NAT64), so [::ffff:169.254.169.254] is caught. - Every DNS answer is validated and the connection pinned to exactly those addresses (resolve overrides), defeating TTL-0 rebinding. The proxy configuration is applied to the pinned client so reqwest makes the per-URL routing decision: proxied requests are resolved by that trusted transport, direct ones (NO_PROXY match, no proxy for the scheme) stay pinned. - Redirects are followed manually (RedirectMode::Manual), at most MAX_REDIRECTS hops, each hop re-validated and re-pinned; redirects to non-HTTP(S) schemes are rejected, matching fetch. - Caller headers are sanitized (hop-by-hop, forwarding, and metadata-service headers dropped) and sent only when the target is strictly same-origin with the configured trusted origin — AI SDK's credentialedOrigin semantics — checked before the first request and on every redirect hop. - URLs same-origin with the configured base_url skip address validation so self-hosted deployments serving assets from their own (possibly private) origin keep working. Wired into every response-supplied-URL fetch: Black Forest Labs (polling + download), Gladia (result polling), Luma, Recraft, Replicate, and fal downloads, and the Google Files upload.
1 parent 42b0fec commit f6d3b3d

12 files changed

Lines changed: 1178 additions & 33 deletions

File tree

‎aimux-provider-utils/src/download_guard.rs‎

Lines changed: 419 additions & 0 deletions
Large diffs are not rendered by default.

‎aimux-provider-utils/src/http.rs‎

Lines changed: 337 additions & 7 deletions
Large diffs are not rendered by default.

‎aimux-provider-utils/src/lib.rs‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66
//! and retry logic — the Rust equivalents of `@ai-sdk/provider-utils`.
77
88
pub mod api_key;
9+
mod download_guard;
910
pub mod headers;
1011
pub mod http;
1112
pub mod logging;
@@ -19,11 +20,12 @@ pub mod url;
1920
pub mod ws;
2021

2122
pub use api_key::load_api_key;
23+
pub use download_guard::same_origin;
2224
pub use headers::with_user_agent_suffix;
2325
pub use http::{
2426
HttpBody, HttpMethod, HttpRequest, HttpResponse, HttpStreamResponse, PoolConfig, ProxyConfig,
2527
RequestTimeout, TimeoutConfig, init_proxy, send, send_stream, send_stream_timed, send_timed,
26-
shared_client, shared_streaming_client, sleep_or_abort,
28+
send_validated, shared_client, shared_streaming_client, sleep_or_abort,
2729
};
2830
pub use logging::{body_logging_enabled, init_logging, redact_body};
2931
pub use multipart::{MultipartForm, media_type_to_extension};
Lines changed: 232 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,232 @@
1+
//! SSRF guard wiring tests for `send_validated`.
2+
//!
3+
//! The guard must reject provider-supplied URLs that point at
4+
//! private/loopback/link-local space (before any connection is attempted),
5+
//! re-validate every redirect hop, and keep same-origin traffic against the
6+
//! configured `trusted_origin` working — mirroring AI SDK's `validateUrl` +
7+
//! `trustedOrigin` semantics. The mock server runs on loopback, so it is only
8+
//! reachable through the trusted-origin exemption; anything the guard treats
9+
//! as foreign is rejected as a non-public literal.
10+
11+
use serde_json::json;
12+
use wiremock::matchers::{header, method, path};
13+
use wiremock::{Mock, MockServer, ResponseTemplate};
14+
15+
use aimux_core::error::AiMuxError;
16+
use aimux_provider_utils::response::DEFAULT_ERROR_STRUCTURE;
17+
use aimux_provider_utils::{HttpBody, HttpMethod, HttpRequest, RetryConfig, send_validated};
18+
19+
fn request(url: String) -> HttpRequest {
20+
HttpRequest {
21+
method: HttpMethod::Get,
22+
url,
23+
headers: vec![("authorization".into(), "Bearer test".into())],
24+
body: HttpBody::Empty,
25+
abort_signal: None,
26+
call_id: None,
27+
recording_context: None,
28+
}
29+
}
30+
31+
#[tokio::test]
32+
async fn rejects_non_public_literal_urls_before_connecting() {
33+
for url in [
34+
"http://169.254.169.254/latest/meta-data",
35+
"http://127.0.0.1:9/file",
36+
"http://[::ffff:169.254.169.254]/meta",
37+
"http://10.1.2.3/file",
38+
] {
39+
let error = send_validated(
40+
request(url.into()),
41+
None,
42+
None,
43+
RetryConfig::default(),
44+
&DEFAULT_ERROR_STRUCTURE,
45+
)
46+
.await
47+
.expect_err("non-public literal must be rejected");
48+
assert!(
49+
matches!(error, AiMuxError::InvalidArgument(ref m) if m.contains("non-public")),
50+
"unexpected error for {url}: {error:?}"
51+
);
52+
}
53+
}
54+
55+
#[tokio::test]
56+
async fn trusted_origin_download_succeeds_and_keeps_auth_headers() {
57+
let server = MockServer::start().await;
58+
Mock::given(method("GET"))
59+
.and(path("/generated/file"))
60+
.and(header("authorization", "Bearer test"))
61+
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"value": "ok"})))
62+
.expect(1)
63+
.mount(&server)
64+
.await;
65+
66+
let response = send_validated(
67+
request(format!("{}/generated/file", server.uri())),
68+
Some(&server.uri()),
69+
Some(&server.uri()),
70+
RetryConfig::default(),
71+
&DEFAULT_ERROR_STRUCTURE,
72+
)
73+
.await
74+
.unwrap();
75+
assert_eq!(response.status, 200);
76+
assert_eq!(
77+
serde_json::from_slice::<serde_json::Value>(&response.body).unwrap(),
78+
json!({"value": "ok"})
79+
);
80+
}
81+
82+
#[tokio::test]
83+
async fn follows_a_relative_redirect_within_the_trusted_origin() {
84+
let server = MockServer::start().await;
85+
Mock::given(method("GET"))
86+
.and(path("/start"))
87+
.respond_with(ResponseTemplate::new(302).insert_header("location", "/final"))
88+
.expect(1)
89+
.mount(&server)
90+
.await;
91+
Mock::given(method("GET"))
92+
.and(path("/final"))
93+
.respond_with(ResponseTemplate::new(200).set_body_bytes(b"done"))
94+
.expect(1)
95+
.mount(&server)
96+
.await;
97+
98+
let response = send_validated(
99+
request(format!("{}/start", server.uri())),
100+
Some(&server.uri()),
101+
Some(&server.uri()),
102+
RetryConfig::default(),
103+
&DEFAULT_ERROR_STRUCTURE,
104+
)
105+
.await
106+
.unwrap();
107+
assert_eq!(response.status, 200);
108+
assert_eq!(response.body.as_ref(), b"done");
109+
}
110+
111+
#[tokio::test]
112+
async fn rejects_a_redirect_to_a_non_public_target() {
113+
let server = MockServer::start().await;
114+
Mock::given(method("GET"))
115+
.and(path("/start"))
116+
.respond_with(
117+
ResponseTemplate::new(302).insert_header("location", "http://169.254.169.254/meta"),
118+
)
119+
.expect(1)
120+
.mount(&server)
121+
.await;
122+
123+
let error = send_validated(
124+
request(format!("{}/start", server.uri())),
125+
Some(&server.uri()),
126+
Some(&server.uri()),
127+
RetryConfig::default(),
128+
&DEFAULT_ERROR_STRUCTURE,
129+
)
130+
.await
131+
.expect_err("redirect to metadata IP must be rejected");
132+
assert!(
133+
matches!(error, AiMuxError::InvalidArgument(ref m) if m.contains("non-public")),
134+
"unexpected error: {error:?}"
135+
);
136+
}
137+
138+
#[tokio::test]
139+
async fn rejects_a_redirect_onto_a_foreign_loopback_origin() {
140+
// The trusted origin is one loopback server; a redirect to a DIFFERENT
141+
// loopback origin must not inherit the exemption.
142+
let server = MockServer::start().await;
143+
let other = MockServer::start().await;
144+
Mock::given(method("GET"))
145+
.and(path("/start"))
146+
.respond_with(
147+
ResponseTemplate::new(302)
148+
.insert_header("location", format!("{}/private", other.uri()).as_str()),
149+
)
150+
.expect(1)
151+
.mount(&server)
152+
.await;
153+
154+
let error = send_validated(
155+
request(format!("{}/start", server.uri())),
156+
Some(&server.uri()),
157+
Some(&server.uri()),
158+
RetryConfig::default(),
159+
&DEFAULT_ERROR_STRUCTURE,
160+
)
161+
.await
162+
.expect_err("foreign loopback origin must be rejected");
163+
assert!(
164+
matches!(error, AiMuxError::InvalidArgument(ref m) if m.contains("non-public")),
165+
"unexpected error: {error:?}"
166+
);
167+
}
168+
169+
#[tokio::test]
170+
async fn rejects_a_redirect_to_a_data_url() {
171+
// Fetch treats a redirect to a non-HTTP(S) scheme as a network error; a
172+
// server must not be able to fabricate a response via Location: data:.
173+
let server = MockServer::start().await;
174+
Mock::given(method("GET"))
175+
.and(path("/start"))
176+
.respond_with(
177+
ResponseTemplate::new(302).insert_header("location", "data:text/plain,forged"),
178+
)
179+
.expect(1)
180+
.mount(&server)
181+
.await;
182+
183+
let error = send_validated(
184+
request(format!("{}/start", server.uri())),
185+
Some(&server.uri()),
186+
Some(&server.uri()),
187+
RetryConfig::default(),
188+
&DEFAULT_ERROR_STRUCTURE,
189+
)
190+
.await
191+
.expect_err("redirect to a data: URL must be rejected");
192+
assert!(
193+
matches!(error, AiMuxError::ApiCall(ref e) if e.message.contains("non-HTTP scheme")),
194+
"unexpected error: {error:?}"
195+
);
196+
}
197+
198+
#[tokio::test]
199+
async fn sanitizes_metadata_and_cookie_headers_from_download_requests() {
200+
let server = MockServer::start().await;
201+
Mock::given(method("GET"))
202+
.and(path("/file"))
203+
.respond_with(ResponseTemplate::new(200).set_body_bytes(b"ok"))
204+
.expect(1)
205+
.mount(&server)
206+
.await;
207+
208+
let mut req = request(format!("{}/file", server.uri()));
209+
req.headers.push(("Cookie".into(), "session=secret".into()));
210+
req.headers
211+
.push(("Metadata-Flavor".into(), "Google".into()));
212+
send_validated(
213+
req,
214+
Some(&server.uri()),
215+
Some(&server.uri()),
216+
RetryConfig::default(),
217+
&DEFAULT_ERROR_STRUCTURE,
218+
)
219+
.await
220+
.unwrap();
221+
222+
let received = server.received_requests().await.unwrap();
223+
assert_eq!(received.len(), 1);
224+
let names: Vec<String> = received[0]
225+
.headers
226+
.keys()
227+
.map(|name| name.as_str().to_ascii_lowercase())
228+
.collect();
229+
assert!(!names.contains(&"cookie".to_string()));
230+
assert!(!names.contains(&"metadata-flavor".to_string()));
231+
assert!(names.contains(&"authorization".to_string()));
232+
}

‎aimux-providers/src/black_forest_labs.rs‎

Lines changed: 63 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -20,8 +20,8 @@ use aimux_core::image_model::{
2020
use aimux_core::shared::Warning;
2121
use aimux_provider_utils::response::DEFAULT_ERROR_STRUCTURE;
2222
use aimux_provider_utils::{
23-
HttpBody, HttpMethod, HttpRequest, RetryConfig, load_api_key, send, sleep_or_abort,
24-
without_trailing_slash,
23+
HttpBody, HttpMethod, HttpRequest, RetryConfig, load_api_key, send, send_validated,
24+
sleep_or_abort, without_trailing_slash,
2525
};
2626

2727
const DEFAULT_POLL_INTERVAL_MS: u64 = 500;
@@ -140,6 +140,23 @@ fn gcd(a: u32, b: u32) -> u32 {
140140
if b == 0 { a } else { gcd(b, a % b) }
141141
}
142142

143+
/// AI SDK's `isTrustedUrl` (black-forest-labs-api.ts): credentials may go to
144+
/// the configured origin, or over HTTPS to `bfl.ai` and its subdomains — BFL
145+
/// serves polling and asset URLs from regional clusters. Allowlisted hosts
146+
/// still get full URL/DNS validation; this gates only the headers.
147+
fn bfl_trusted_url(url: &str, base_url: &str) -> bool {
148+
if aimux_provider_utils::same_origin(url, base_url) {
149+
return true;
150+
}
151+
let Ok(parsed) = url::Url::parse(url) else {
152+
return false;
153+
};
154+
parsed.scheme() == "https"
155+
&& parsed
156+
.host_str()
157+
.is_some_and(|host| host == "bfl.ai" || host.ends_with(".bfl.ai"))
158+
}
159+
143160
#[async_trait]
144161
impl ImageModel for BlackForestLabsImageModel {
145162
fn provider(&self) -> &str {
@@ -285,6 +302,15 @@ impl ImageModel for BlackForestLabsImageModel {
285302

286303
let headers = self.build_headers(options.headers.as_ref());
287304
let header_list: Vec<(String, String)> = headers.into_iter().collect();
305+
// AI SDK gates headers per URL via isTrustedUrl; response-supplied
306+
// targets outside the BFL allowlist get none.
307+
let gated_headers = |url: &str| -> Vec<(String, String)> {
308+
if bfl_trusted_url(url, &self.config.base_url) {
309+
header_list.clone()
310+
} else {
311+
vec![]
312+
}
313+
};
288314

289315
// Submit
290316
let resp = send(
@@ -343,17 +369,23 @@ impl ImageModel for BlackForestLabsImageModel {
343369
let mut result_duration = None;
344370

345371
for _ in 0..max_attempts {
346-
let pr = send(
372+
// AI SDK polls polling_url with validateUrl: true and gates the
373+
// headers itself via isTrustedUrl (base_url origin or HTTPS
374+
// *.bfl.ai), so credentialed_origin is None here.
375+
let poll_url = poll_url_with_id.to_string();
376+
let pr = send_validated(
347377
HttpRequest {
348378
method: HttpMethod::Get,
349-
url: poll_url_with_id.to_string(),
350-
headers: header_list.clone(),
379+
headers: gated_headers(&poll_url),
380+
url: poll_url,
351381
body: HttpBody::Empty,
352382

353383
abort_signal: options.abort_signal.clone(),
354384
call_id: None,
355385
recording_context: None,
356386
},
387+
Some(&self.config.base_url),
388+
None,
357389
RetryConfig::default(),
358390
&DEFAULT_ERROR_STRUCTURE,
359391
)
@@ -405,18 +437,22 @@ impl ImageModel for BlackForestLabsImageModel {
405437
))
406438
})?;
407439

408-
// Download image
409-
let ir = send(
440+
// Download image; result.sample is a URL from the poll response body.
441+
// AI SDK sends its headers to trusted BFL hosts on the download too
442+
// (isTrustedUrl-gated), so mirror the poll's header policy.
443+
let ir = send_validated(
410444
HttpRequest {
411445
method: HttpMethod::Get,
446+
headers: gated_headers(&image_url),
412447
url: image_url,
413-
headers: vec![],
414448
body: HttpBody::Empty,
415449

416450
abort_signal: options.abort_signal.clone(),
417451
call_id: None,
418452
recording_context: None,
419453
},
454+
Some(&self.config.base_url),
455+
None,
420456
RetryConfig::default(),
421457
&DEFAULT_ERROR_STRUCTURE,
422458
)
@@ -465,3 +501,22 @@ impl ImageModel for BlackForestLabsImageModel {
465501
})
466502
}
467503
}
504+
505+
#[cfg(test)]
506+
mod tests {
507+
use super::bfl_trusted_url;
508+
509+
#[test]
510+
fn credentials_go_to_the_base_url_and_bfl_hosts_only() {
511+
let base = "https://api.bfl.ai";
512+
assert!(bfl_trusted_url("https://api.bfl.ai/v1/get_result", base));
513+
assert!(bfl_trusted_url(
514+
"https://api.us1.bfl.ai/v1/get_result",
515+
base
516+
));
517+
assert!(bfl_trusted_url("https://bfl.ai/x", base));
518+
assert!(!bfl_trusted_url("http://api.us1.bfl.ai/x", base));
519+
assert!(!bfl_trusted_url("https://evil-bfl.ai/x", base));
520+
assert!(!bfl_trusted_url("https://attacker.example/x", base));
521+
}
522+
}

0 commit comments

Comments
 (0)