Files
Genarrative/apps/ai-game-creator-shell/src-tauri/src/http_client.rs
T
lhk229 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
AGC对主站的请求添加header (#241)
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>
2026-09-02 14:39:01 +08:00

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())
);
}
}