65d7a57eb7
Project CI / Frontend tests (push) Successful in 4m0s
Project CI / Repository checks (push) Successful in 4m4s
Project CI / Backend tests (push) Successful in 8m40s
Project CI / Native shell tests (push) Successful in 17m56s
Project CI / Frontend tests (pull_request) Successful in 5m16s
Project CI / Repository checks (pull_request) Successful in 6m57s
Project CI / Backend tests (pull_request) Successful in 6m31s
Project CI / Native shell tests (pull_request) Successful in 16m57s
Reviewed-on: http://192.168.35.82/git/GenarrativeAI/Genarrative/pulls/241 Co-authored-by: Linghong <ink29535@proton.me> Co-committed-by: Linghong <ink29535@proton.me>
415 lines
16 KiB
Rust
415 lines
16 KiB
Rust
use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
|
|
|
|
const AGC_CLIENT_MARKER_HEADER: &str = "x-genarrative-client";
|
|
const AGC_CLIENT_MARKER_VALUE: &str = "agc";
|
|
|
|
fn same_origin(initial: &reqwest::Url, next: &reqwest::Url) -> bool {
|
|
initial.origin() == next.origin()
|
|
}
|
|
|
|
fn agc_main_site_redirect_policy() -> reqwest::redirect::Policy {
|
|
// Keep reqwest's default same-origin behavior, but fail closed before a
|
|
// custom marker can be copied to a different origin.
|
|
let default_policy = reqwest::redirect::Policy::default();
|
|
reqwest::redirect::Policy::custom(move |attempt| {
|
|
let origin_matches = attempt
|
|
.previous()
|
|
.first()
|
|
.map(|initial| same_origin(initial, attempt.url()))
|
|
.unwrap_or(false);
|
|
|
|
if origin_matches {
|
|
default_policy.redirect(attempt)
|
|
} else {
|
|
attempt.stop()
|
|
}
|
|
})
|
|
}
|
|
|
|
fn agc_main_site_marker_headers() -> HeaderMap {
|
|
let mut headers = HeaderMap::new();
|
|
headers.insert(
|
|
HeaderName::from_static(AGC_CLIENT_MARKER_HEADER),
|
|
HeaderValue::from_static(AGC_CLIENT_MARKER_VALUE),
|
|
);
|
|
headers
|
|
}
|
|
|
|
pub(crate) fn agc_main_site_client_builder() -> reqwest::ClientBuilder {
|
|
reqwest::Client::builder()
|
|
.default_headers(agc_main_site_marker_headers())
|
|
.redirect(agc_main_site_redirect_policy())
|
|
}
|
|
|
|
/// Finalize a request sent through the AGC main-site client.
|
|
///
|
|
/// `ClientBuilder::default_headers` only fills a missing request header. A
|
|
/// request-level header with the same name would otherwise win, so use
|
|
/// `RequestBuilder::headers` here to replace any caller-provided value with
|
|
/// the reserved AGC marker after all business headers have been configured.
|
|
pub(crate) fn with_agc_main_site_marker(
|
|
request: reqwest::RequestBuilder,
|
|
) -> reqwest::RequestBuilder {
|
|
request.headers(agc_main_site_marker_headers())
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::{
|
|
agc_main_site_client_builder, same_origin, AGC_CLIENT_MARKER_HEADER,
|
|
AGC_CLIENT_MARKER_VALUE,
|
|
};
|
|
use reqwest::header::AUTHORIZATION;
|
|
use std::io::{Read, Write};
|
|
use std::net::{TcpListener, TcpStream};
|
|
use std::time::Duration;
|
|
|
|
fn read_http_request(stream: &mut TcpStream) -> String {
|
|
stream
|
|
.set_read_timeout(Some(Duration::from_secs(5)))
|
|
.expect("set HTTP fixture read timeout");
|
|
let mut bytes = Vec::new();
|
|
let mut buffer = [0_u8; 4096];
|
|
loop {
|
|
let read = stream.read(&mut buffer).expect("read HTTP fixture request");
|
|
assert!(read > 0, "HTTP fixture request closed before headers");
|
|
bytes.extend_from_slice(&buffer[..read]);
|
|
if bytes.windows(4).any(|value| value == b"\r\n\r\n") {
|
|
break;
|
|
}
|
|
}
|
|
String::from_utf8_lossy(&bytes).into_owned()
|
|
}
|
|
|
|
fn request_header(request: &str, expected_name: &str) -> Option<String> {
|
|
request.lines().find_map(|line| {
|
|
let (name, value) = line.split_once(':')?;
|
|
name.eq_ignore_ascii_case(expected_name)
|
|
.then(|| value.trim().to_string())
|
|
})
|
|
}
|
|
|
|
fn spawn_http_fixture(listener: TcpListener) -> std::thread::JoinHandle<String> {
|
|
std::thread::spawn(move || {
|
|
let (mut stream, _) = listener.accept().expect("accept HTTP fixture request");
|
|
let request = read_http_request(&mut stream);
|
|
stream
|
|
.write_all(
|
|
b"HTTP/1.1 204 No Content\r\nContent-Length: 0\r\nConnection: close\r\n\r\n",
|
|
)
|
|
.expect("write HTTP fixture response");
|
|
request
|
|
})
|
|
}
|
|
|
|
fn spawn_redirect_fixture(
|
|
listener: TcpListener,
|
|
location: String,
|
|
) -> std::thread::JoinHandle<String> {
|
|
std::thread::spawn(move || {
|
|
let (mut stream, _) = listener.accept().expect("accept HTTP redirect request");
|
|
let request = read_http_request(&mut stream);
|
|
let response = format!(
|
|
"HTTP/1.1 302 Found\r\nLocation: {location}\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
|
|
);
|
|
stream
|
|
.write_all(response.as_bytes())
|
|
.expect("write HTTP redirect response");
|
|
request
|
|
})
|
|
}
|
|
|
|
#[test]
|
|
fn main_site_redirect_policy_compares_full_origin() {
|
|
let origin_matches = |initial: &str, next: &str| {
|
|
let initial = reqwest::Url::parse(initial).expect("parse initial URL");
|
|
let next = reqwest::Url::parse(next).expect("parse next URL");
|
|
same_origin(&initial, &next)
|
|
};
|
|
|
|
assert!(origin_matches(
|
|
"https://main.example/api/first",
|
|
"https://main.example/api/second"
|
|
));
|
|
assert!(origin_matches(
|
|
"https://main.example",
|
|
"https://main.example:443/api/second"
|
|
));
|
|
assert!(!origin_matches(
|
|
"https://main.example",
|
|
"http://main.example/api/second"
|
|
));
|
|
assert!(!origin_matches(
|
|
"https://main.example",
|
|
"https://cdn.example/api/second"
|
|
));
|
|
assert!(!origin_matches(
|
|
"https://main.example:8443",
|
|
"https://main.example:9443/api/second"
|
|
));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn factory_sets_the_agc_marker_as_a_default_header() {
|
|
let listener = TcpListener::bind("127.0.0.1:0").expect("bind HTTP fixture");
|
|
let address = listener.local_addr().expect("read HTTP fixture address");
|
|
let fixture = spawn_http_fixture(listener);
|
|
let client = agc_main_site_client_builder()
|
|
.build()
|
|
.expect("build AGC main-site client");
|
|
let response = client
|
|
.get(format!("http://{address}/api/auth/me"))
|
|
.send()
|
|
.await
|
|
.expect("send request");
|
|
let request = fixture.join().expect("join HTTP fixture");
|
|
|
|
assert_eq!(response.status(), reqwest::StatusCode::NO_CONTENT);
|
|
assert_eq!(
|
|
request_header(&request, AGC_CLIENT_MARKER_HEADER),
|
|
Some(AGC_CLIENT_MARKER_VALUE.to_string())
|
|
);
|
|
assert!(request_header(&request, AUTHORIZATION.as_str()).is_none());
|
|
assert!(request_header(&request, "idempotency-key").is_none());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn request_finalizer_overrides_a_caller_provided_marker() {
|
|
let listener = TcpListener::bind("127.0.0.1:0").expect("bind HTTP fixture");
|
|
let address = listener.local_addr().expect("read HTTP fixture address");
|
|
let fixture = spawn_http_fixture(listener);
|
|
let client = agc_main_site_client_builder()
|
|
.build()
|
|
.expect("build AGC main-site client");
|
|
let request = client
|
|
.get(format!("http://{address}/api/auth/me"))
|
|
.header(AGC_CLIENT_MARKER_HEADER, "spoofed-value");
|
|
let response = super::with_agc_main_site_marker(request)
|
|
.send()
|
|
.await
|
|
.expect("send finalized request");
|
|
let request = fixture.join().expect("join HTTP fixture");
|
|
|
|
assert_eq!(response.status(), reqwest::StatusCode::NO_CONTENT);
|
|
assert_eq!(
|
|
request_header(&request, AGC_CLIENT_MARKER_HEADER),
|
|
Some(AGC_CLIENT_MARKER_VALUE.to_string())
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn factory_keeps_request_headers_and_transport_options_configurable() {
|
|
let listener = TcpListener::bind("127.0.0.1:0").expect("bind HTTP fixture");
|
|
let address = listener.local_addr().expect("read HTTP fixture address");
|
|
let fixture = spawn_http_fixture(listener);
|
|
let client = agc_main_site_client_builder()
|
|
.connect_timeout(Duration::from_secs(10))
|
|
.timeout(Duration::from_secs(60))
|
|
.redirect(reqwest::redirect::Policy::none())
|
|
.no_proxy()
|
|
.build()
|
|
.expect("build configured AGC main-site client");
|
|
let response = client
|
|
.post(format!("http://{address}/api/editor/images/generations"))
|
|
.bearer_auth("fixture-token")
|
|
.header("Idempotency-Key", "fixture-id")
|
|
.send()
|
|
.await
|
|
.expect("send configured request");
|
|
let request = fixture.join().expect("join HTTP fixture");
|
|
|
|
assert_eq!(response.status(), reqwest::StatusCode::NO_CONTENT);
|
|
assert_eq!(
|
|
request_header(&request, AGC_CLIENT_MARKER_HEADER),
|
|
Some(AGC_CLIENT_MARKER_VALUE.to_string())
|
|
);
|
|
assert_eq!(
|
|
request_header(&request, AUTHORIZATION.as_str()),
|
|
Some("Bearer fixture-token".to_string())
|
|
);
|
|
assert_eq!(
|
|
request_header(&request, "idempotency-key"),
|
|
Some("fixture-id".to_string())
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn factory_follows_same_origin_redirects_with_the_agc_marker() {
|
|
let listener = TcpListener::bind("127.0.0.1:0").expect("bind HTTP fixture");
|
|
let address = listener.local_addr().expect("read HTTP fixture address");
|
|
let fixture = std::thread::spawn(move || {
|
|
let (mut first_stream, _) = listener.accept().expect("accept first HTTP request");
|
|
let first_request = read_http_request(&mut first_stream);
|
|
first_stream
|
|
.write_all(
|
|
b"HTTP/1.1 302 Found\r\nLocation: /api/auth/me/final\r\nContent-Length: 0\r\nConnection: close\r\n\r\n",
|
|
)
|
|
.expect("write same-origin redirect response");
|
|
|
|
let (mut second_stream, _) = listener.accept().expect("accept redirected HTTP request");
|
|
let second_request = read_http_request(&mut second_stream);
|
|
second_stream
|
|
.write_all(
|
|
b"HTTP/1.1 204 No Content\r\nContent-Length: 0\r\nConnection: close\r\n\r\n",
|
|
)
|
|
.expect("write final HTTP response");
|
|
|
|
(first_request, second_request)
|
|
});
|
|
let client = agc_main_site_client_builder()
|
|
.timeout(Duration::from_secs(2))
|
|
.no_proxy()
|
|
.build()
|
|
.expect("build AGC main-site client");
|
|
|
|
let response = client
|
|
.get(format!("http://{address}/api/auth/me"))
|
|
.send()
|
|
.await
|
|
.expect("send request");
|
|
let (first_request, second_request) = fixture.join().expect("join HTTP fixture");
|
|
|
|
assert_eq!(response.status(), reqwest::StatusCode::NO_CONTENT);
|
|
assert_eq!(
|
|
request_header(&first_request, AGC_CLIENT_MARKER_HEADER),
|
|
Some(AGC_CLIENT_MARKER_VALUE.to_string())
|
|
);
|
|
assert_eq!(
|
|
request_header(&second_request, AGC_CLIENT_MARKER_HEADER),
|
|
Some(AGC_CLIENT_MARKER_VALUE.to_string())
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn factory_stops_cross_origin_redirects_before_sending_the_marker() {
|
|
let source_listener = TcpListener::bind("127.0.0.1:0").expect("bind source fixture");
|
|
let source_address = source_listener
|
|
.local_addr()
|
|
.expect("read source fixture address");
|
|
let target_listener = TcpListener::bind("127.0.0.1:0").expect("bind target fixture");
|
|
target_listener
|
|
.set_nonblocking(true)
|
|
.expect("configure target fixture");
|
|
let target_address = target_listener
|
|
.local_addr()
|
|
.expect("read target fixture address");
|
|
let source_fixture = spawn_redirect_fixture(
|
|
source_listener,
|
|
format!("http://{target_address}/oss/object"),
|
|
);
|
|
let client = agc_main_site_client_builder()
|
|
.timeout(Duration::from_secs(2))
|
|
.no_proxy()
|
|
.build()
|
|
.expect("build AGC main-site client");
|
|
|
|
let response = client
|
|
.get(format!("http://{source_address}/api/assets/read-url"))
|
|
.send()
|
|
.await
|
|
.expect("send request");
|
|
let source_request = source_fixture.join().expect("join source fixture");
|
|
|
|
assert_eq!(response.status(), reqwest::StatusCode::FOUND);
|
|
assert_eq!(
|
|
request_header(&source_request, AGC_CLIENT_MARKER_HEADER),
|
|
Some(AGC_CLIENT_MARKER_VALUE.to_string())
|
|
);
|
|
assert!(matches!(
|
|
target_listener.accept(),
|
|
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock
|
|
));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn factory_allows_same_origin_then_blocks_cross_origin_redirect_chain() {
|
|
let source_listener = TcpListener::bind("127.0.0.1:0").expect("bind source fixture");
|
|
let source_address = source_listener
|
|
.local_addr()
|
|
.expect("read source fixture address");
|
|
let target_listener = TcpListener::bind("127.0.0.1:0").expect("bind target fixture");
|
|
target_listener
|
|
.set_nonblocking(true)
|
|
.expect("configure target fixture");
|
|
let target_address = target_listener
|
|
.local_addr()
|
|
.expect("read target fixture address");
|
|
let fixture = std::thread::spawn(move || {
|
|
let (mut first_stream, _) = source_listener
|
|
.accept()
|
|
.expect("accept first source request");
|
|
let first_request = read_http_request(&mut first_stream);
|
|
first_stream
|
|
.write_all(
|
|
b"HTTP/1.1 302 Found\r\nLocation: /api/auth/me/second\r\nContent-Length: 0\r\nConnection: close\r\n\r\n",
|
|
)
|
|
.expect("write same-origin redirect response");
|
|
|
|
let (mut second_stream, _) = source_listener
|
|
.accept()
|
|
.expect("accept second source request");
|
|
let second_request = read_http_request(&mut second_stream);
|
|
let response = format!(
|
|
"HTTP/1.1 302 Found\r\nLocation: http://{target_address}/oss/object\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
|
|
);
|
|
second_stream
|
|
.write_all(response.as_bytes())
|
|
.expect("write cross-origin redirect response");
|
|
|
|
(first_request, second_request)
|
|
});
|
|
let client = agc_main_site_client_builder()
|
|
.timeout(Duration::from_secs(2))
|
|
.no_proxy()
|
|
.build()
|
|
.expect("build AGC main-site client");
|
|
|
|
let response = client
|
|
.get(format!("http://{source_address}/api/auth/me"))
|
|
.send()
|
|
.await
|
|
.expect("send request");
|
|
let (first_request, second_request) = fixture.join().expect("join source fixture");
|
|
|
|
assert_eq!(response.status(), reqwest::StatusCode::FOUND);
|
|
assert_eq!(
|
|
request_header(&first_request, AGC_CLIENT_MARKER_HEADER),
|
|
Some(AGC_CLIENT_MARKER_VALUE.to_string())
|
|
);
|
|
assert_eq!(
|
|
request_header(&second_request, AGC_CLIENT_MARKER_HEADER),
|
|
Some(AGC_CLIENT_MARKER_VALUE.to_string())
|
|
);
|
|
assert!(matches!(
|
|
target_listener.accept(),
|
|
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock
|
|
));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn explicit_no_redirect_policy_still_overrides_the_factory_policy() {
|
|
let listener = TcpListener::bind("127.0.0.1:0").expect("bind HTTP fixture");
|
|
let address = listener.local_addr().expect("read HTTP fixture address");
|
|
let fixture =
|
|
spawn_redirect_fixture(listener, format!("http://{address}/api/auth/me/final"));
|
|
let client = agc_main_site_client_builder()
|
|
.redirect(reqwest::redirect::Policy::none())
|
|
.no_proxy()
|
|
.build()
|
|
.expect("build no-redirect AGC main-site client");
|
|
|
|
let response = client
|
|
.get(format!("http://{address}/api/auth/me"))
|
|
.send()
|
|
.await
|
|
.expect("send request");
|
|
let request = fixture.join().expect("join HTTP fixture");
|
|
|
|
assert_eq!(response.status(), reqwest::StatusCode::FOUND);
|
|
assert_eq!(
|
|
request_header(&request, AGC_CLIENT_MARKER_HEADER),
|
|
Some(AGC_CLIENT_MARKER_VALUE.to_string())
|
|
);
|
|
}
|
|
}
|