mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-10 11:19:50 +08:00
feat(proxy,failover,transport): Hyper 上游客户端精细计时、连续失败退避与连接泄漏修复
Proxy: - 将上游 HTTP 客户端从 reqwest 替换为 hyper,新增 InstrumentedConnector 实现 TCP 连接/TLS 握手级别的独立计时,上报 connection_reused 等指标 - 前端展示细粒度代理计时(连接复用、等待响应头等) Failover: - 引入连续失败退避机制,每 10 次失败递增退避间隔 - 检测 H2 max outbound streams 错误并触发上游客户端重建 - 新增 HTTPClientPool.reset_upstream_client 支持按需重建缓存客户端 连接泄漏修复: - Handler 异常路径确保 response_ctx 被正确关闭 - HubResponseStream 迭代结束后在 finally 块中清理 stream_id - HubTunnelTransport.handle_request 捕获所有异常并清理流状态
This commit is contained in:
+6
-23
@@ -13,8 +13,8 @@ use crate::config::{Config, ServerEntry};
|
||||
use crate::net;
|
||||
use crate::registration::client::AetherClient;
|
||||
use crate::runtime::{self, DynamicConfig};
|
||||
use crate::safe_dns::SafeDnsResolver;
|
||||
use crate::state::{AppState, ProxyMetrics, ServerContext};
|
||||
use crate::upstream_client;
|
||||
use crate::{hardware, target_filter, tunnel};
|
||||
|
||||
/// Run the full application lifecycle after config has been parsed.
|
||||
@@ -67,27 +67,10 @@ pub async fn run(mut config: Config, servers: Vec<ServerEntry>) -> anyhow::Resul
|
||||
config.dns_cache_capacity,
|
||||
));
|
||||
|
||||
// Build reqwest client for tunnel upstream requests (shared).
|
||||
// Inject SafeDnsResolver so reqwest only connects to addresses that were
|
||||
// validated by validate_target() — this eliminates the DNS rebinding
|
||||
// TOCTTOU gap where a second DNS lookup could return a private IP.
|
||||
let safe_resolver = SafeDnsResolver::new(Arc::clone(&dns_cache));
|
||||
let mut reqwest_builder = reqwest::Client::builder()
|
||||
.dns_resolver(Arc::new(safe_resolver))
|
||||
.pool_max_idle_per_host(config.upstream_pool_max_idle_per_host)
|
||||
.pool_idle_timeout(Duration::from_secs(config.upstream_pool_idle_timeout_secs))
|
||||
.connect_timeout(Duration::from_secs(config.upstream_connect_timeout_secs))
|
||||
.tcp_nodelay(config.upstream_tcp_nodelay);
|
||||
|
||||
if config.upstream_tcp_keepalive_secs > 0 {
|
||||
reqwest_builder = reqwest_builder.tcp_keepalive(Some(Duration::from_secs(
|
||||
config.upstream_tcp_keepalive_secs,
|
||||
)));
|
||||
}
|
||||
|
||||
let reqwest_client = reqwest_builder
|
||||
.build()
|
||||
.expect("failed to build reqwest client");
|
||||
// Build Hyper client for tunnel upstream requests (shared).
|
||||
// DNS still flows through validated addresses from DnsCache, while the
|
||||
// custom connector exposes per-request connect/TLS timing when available.
|
||||
let upstream_client = upstream_client::build_upstream_client(&config, Arc::clone(&dns_cache));
|
||||
|
||||
// Register with each Aether server and build per-server contexts.
|
||||
// Wrapped in Arc<Mutex> so retry_failed_registrations can append later.
|
||||
@@ -160,7 +143,7 @@ pub async fn run(mut config: Config, servers: Vec<ServerEntry>) -> anyhow::Resul
|
||||
let state = Arc::new(AppState {
|
||||
config: Arc::new(config),
|
||||
dns_cache,
|
||||
reqwest_client,
|
||||
upstream_client,
|
||||
tunnel_tls_config,
|
||||
});
|
||||
|
||||
|
||||
@@ -4,11 +4,11 @@ mod hardware;
|
||||
mod net;
|
||||
mod registration;
|
||||
mod runtime;
|
||||
mod safe_dns;
|
||||
mod setup;
|
||||
mod state;
|
||||
mod target_filter;
|
||||
mod tunnel;
|
||||
mod upstream_client;
|
||||
|
||||
use std::path::PathBuf;
|
||||
|
||||
|
||||
@@ -8,14 +8,15 @@ use crate::config::Config;
|
||||
use crate::registration::client::AetherClient;
|
||||
use crate::runtime::SharedDynamicConfig;
|
||||
use crate::target_filter::DnsCache;
|
||||
use crate::upstream_client::UpstreamClient;
|
||||
|
||||
/// Central application state shared across all servers/tunnels.
|
||||
pub struct AppState {
|
||||
pub config: Arc<Config>,
|
||||
/// DNS cache for upstream target resolution (shared).
|
||||
pub dns_cache: Arc<DnsCache>,
|
||||
/// Reqwest client for tunnel upstream requests (shared).
|
||||
pub reqwest_client: reqwest::Client,
|
||||
/// Hyper client for tunnel upstream requests with validated DNS and connection timing.
|
||||
pub upstream_client: UpstreamClient,
|
||||
/// Shared TLS config for tunnel WebSocket connections (avoids re-parsing root CAs on each reconnect).
|
||||
pub tunnel_tls_config: Arc<rustls::ClientConfig>,
|
||||
}
|
||||
|
||||
@@ -9,11 +9,13 @@ use std::time::{Duration, Instant};
|
||||
|
||||
use bytes::Bytes;
|
||||
use futures_util::StreamExt;
|
||||
use http_body_util::BodyExt;
|
||||
use tokio::sync::mpsc;
|
||||
use tracing::{debug, warn};
|
||||
|
||||
use crate::state::{AppState, ServerContext};
|
||||
use crate::target_filter;
|
||||
use crate::upstream_client::{self, UpstreamRequestBody};
|
||||
|
||||
use super::protocol::{
|
||||
compress_payload, decompress_if_gzip, flags, Frame, MsgType, RequestMeta, ResponseMeta,
|
||||
@@ -203,44 +205,61 @@ async fn handle_stream_inner(
|
||||
let dns_ms = connect_start.elapsed().as_millis() as u64;
|
||||
|
||||
// Execute upstream request
|
||||
let client = &state.reqwest_client;
|
||||
let client = &state.upstream_client;
|
||||
let timeout = Duration::from_secs(meta.timeout.clamp(MIN_TIMEOUT_SECS, MAX_TIMEOUT_SECS));
|
||||
|
||||
let method: reqwest::Method = meta.method.parse().unwrap_or(reqwest::Method::GET);
|
||||
// Build a complete HeaderMap from tunnel headers, then set it all at once
|
||||
// via .headers() which *replaces* reqwest defaults (e.g. Accept: */*),
|
||||
// ensuring upstream sees exactly what Aether server intended.
|
||||
let mut header_map = reqwest::header::HeaderMap::with_capacity(meta.headers.len());
|
||||
let method: hyper::Method = meta.method.parse().unwrap_or(hyper::Method::GET);
|
||||
let mut request = match hyper::Request::builder()
|
||||
.method(method)
|
||||
.uri(meta.url.as_str())
|
||||
.body(UpstreamRequestBody::new(body.clone()))
|
||||
{
|
||||
Ok(request) => request,
|
||||
Err(e) => {
|
||||
send_error(
|
||||
frame_tx,
|
||||
stream_id,
|
||||
&format!("invalid upstream request: {e}"),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
}
|
||||
};
|
||||
|
||||
let headers = request.headers_mut();
|
||||
for (k, v) in &meta.headers {
|
||||
let k_lower = k.to_ascii_lowercase();
|
||||
if BLOCKED_HEADERS.contains(&k_lower.as_str()) {
|
||||
continue;
|
||||
}
|
||||
if let (Ok(name), Ok(value)) = (
|
||||
reqwest::header::HeaderName::from_bytes(k.as_bytes()),
|
||||
reqwest::header::HeaderValue::from_str(v),
|
||||
hyper::header::HeaderName::from_bytes(k.as_bytes()),
|
||||
hyper::header::HeaderValue::from_str(v),
|
||||
) {
|
||||
header_map.insert(name, value);
|
||||
headers.insert(name, value);
|
||||
}
|
||||
}
|
||||
let mut req = client.request(method, &meta.url).headers(header_map);
|
||||
|
||||
let body_size = body.len();
|
||||
if !body.is_empty() {
|
||||
req = req.body(body);
|
||||
}
|
||||
req = req.timeout(timeout);
|
||||
let mut captured_connection = upstream_client::capture_connection(&mut request);
|
||||
let connection_start = Instant::now();
|
||||
let connection_capture = tokio::spawn(async move {
|
||||
let connected = captured_connection.wait_for_connection_metadata().await;
|
||||
connected
|
||||
.as_ref()
|
||||
.map(|_| connection_start.elapsed().as_millis() as u64)
|
||||
});
|
||||
|
||||
let upstream_start = Instant::now();
|
||||
let response = match req.send().await {
|
||||
Ok(r) => r,
|
||||
Err(e) => {
|
||||
let response = match tokio::time::timeout(timeout, client.request(request)).await {
|
||||
Ok(Ok(response)) => response,
|
||||
Ok(Err(e)) => {
|
||||
connection_capture.abort();
|
||||
server
|
||||
.metrics
|
||||
.failed_requests
|
||||
.fetch_add(1, Ordering::Release);
|
||||
let msg = if e.is_timeout() {
|
||||
"upstream timeout".to_string()
|
||||
} else if e.is_connect() {
|
||||
let msg = if e.is_connect() {
|
||||
format!("upstream connect error: {e}")
|
||||
} else {
|
||||
format!("upstream error: {e}")
|
||||
@@ -248,6 +267,15 @@ async fn handle_stream_inner(
|
||||
send_error(frame_tx, stream_id, &msg).await;
|
||||
return None;
|
||||
}
|
||||
Err(_) => {
|
||||
connection_capture.abort();
|
||||
server
|
||||
.metrics
|
||||
.failed_requests
|
||||
.fetch_add(1, Ordering::Release);
|
||||
send_error(frame_tx, stream_id, "upstream timeout").await;
|
||||
return None;
|
||||
}
|
||||
};
|
||||
|
||||
// Capture connection-establishment duration (DNS + TCP/TLS + TTFB)
|
||||
@@ -257,18 +285,34 @@ async fn handle_stream_inner(
|
||||
// Send RESPONSE_HEADERS
|
||||
let status = response.status().as_u16();
|
||||
let ttfb_ms = upstream_start.elapsed().as_millis() as u64;
|
||||
// Short timeout: on connection reuse hyper may never fire the connect
|
||||
// callback, so avoid blocking indefinitely.
|
||||
let connection_acquire_ms =
|
||||
match tokio::time::timeout(Duration::from_millis(100), connection_capture).await {
|
||||
Ok(Ok(ms)) => ms,
|
||||
Ok(Err(_)) => None, // JoinError (task panicked / cancelled)
|
||||
Err(_) => None, // timeout -- task is detached but lightweight
|
||||
};
|
||||
let request_timing =
|
||||
upstream_client::resolve_request_timing(&response, connection_acquire_ms, ttfb_ms);
|
||||
let mut resp_headers: Vec<(String, String)> = Vec::with_capacity(response.headers().len() + 1);
|
||||
for (k, v) in response.headers() {
|
||||
if let Ok(vs) = v.to_str() {
|
||||
resp_headers.push((k.as_str().to_string(), vs.to_string()));
|
||||
}
|
||||
}
|
||||
// Inject proxy timing (same format as delegate mode)
|
||||
let timing = serde_json::json!({
|
||||
"dns_ms": dns_ms,
|
||||
"connection_acquire_ms": request_timing.connection_acquire_ms,
|
||||
"connection_reused": request_timing.connection_reused,
|
||||
"connect_ms": request_timing.connect_ms,
|
||||
"tls_ms": request_timing.tls_ms,
|
||||
"ttfb_ms": ttfb_ms,
|
||||
"upstream_ms": ttfb_ms,
|
||||
"upstream_processing_ms": ttfb_ms.saturating_sub(dns_ms),
|
||||
"response_wait_ms": request_timing.response_wait_ms,
|
||||
"upstream_processing_ms": request_timing.response_wait_ms,
|
||||
"timing_source": "instrumented_connector",
|
||||
"total_ms": connect_elapsed.as_millis() as u64,
|
||||
"body_size": body_size,
|
||||
"mode": "tunnel",
|
||||
});
|
||||
@@ -298,7 +342,7 @@ async fn handle_stream_inner(
|
||||
// (e.g. uncompressed SSE text). Already-compressed data (gzip/br from
|
||||
// upstream Content-Encoding) won't shrink further and will be sent as-is
|
||||
// thanks to the size check in compress_payload().
|
||||
let mut stream = response.bytes_stream();
|
||||
let mut stream = response.into_body().into_data_stream();
|
||||
while let Some(chunk_result) = stream.next().await {
|
||||
match chunk_result {
|
||||
Ok(chunk) => {
|
||||
|
||||
@@ -0,0 +1,434 @@
|
||||
use std::future::Future;
|
||||
use std::io;
|
||||
use std::net::IpAddr;
|
||||
use std::pin::Pin;
|
||||
use std::sync::Arc;
|
||||
use std::task::{Context, Poll};
|
||||
use std::time::Duration;
|
||||
|
||||
use bytes::Bytes;
|
||||
use http_body_util::Full;
|
||||
use hyper::rt;
|
||||
use hyper::Response;
|
||||
use hyper::Uri;
|
||||
pub use hyper_util::client::legacy::connect::capture_connection;
|
||||
use hyper_util::client::legacy::connect::dns::Name;
|
||||
use hyper_util::client::legacy::connect::{Connected, Connection, HttpConnector};
|
||||
use hyper_util::client::legacy::Client;
|
||||
use hyper_util::rt::{TokioExecutor, TokioIo, TokioTimer};
|
||||
use rustls::pki_types::ServerName;
|
||||
use rustls::ClientConfig;
|
||||
use tokio::net::TcpStream;
|
||||
use tokio_rustls::TlsConnector;
|
||||
use tower_service::Service;
|
||||
|
||||
use crate::config::Config;
|
||||
use crate::target_filter::{self, DnsCache};
|
||||
|
||||
type BoxError = Box<dyn std::error::Error + Send + Sync>;
|
||||
|
||||
type PlainStream = TokioIo<TcpStream>;
|
||||
type TlsStream = TokioIo<tokio_rustls::client::TlsStream<TcpStream>>;
|
||||
|
||||
pub type UpstreamRequestBody = Full<Bytes>;
|
||||
pub type UpstreamClient = Client<InstrumentedConnector, UpstreamRequestBody>;
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default)]
|
||||
pub struct ConnectTiming {
|
||||
pub connect_ms: u64,
|
||||
pub tls_ms: u64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default)]
|
||||
pub struct RequestTiming {
|
||||
pub connection_acquire_ms: u64,
|
||||
pub connect_ms: u64,
|
||||
pub tls_ms: u64,
|
||||
pub response_wait_ms: u64,
|
||||
pub connection_reused: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ValidatedResolver {
|
||||
dns_cache: Arc<DnsCache>,
|
||||
}
|
||||
|
||||
impl ValidatedResolver {
|
||||
pub fn new(dns_cache: Arc<DnsCache>) -> Self {
|
||||
Self { dns_cache }
|
||||
}
|
||||
}
|
||||
|
||||
pub struct ValidatedAddrs {
|
||||
inner: std::vec::IntoIter<std::net::SocketAddr>,
|
||||
}
|
||||
|
||||
impl Iterator for ValidatedAddrs {
|
||||
type Item = std::net::SocketAddr;
|
||||
|
||||
fn next(&mut self) -> Option<Self::Item> {
|
||||
self.inner.next()
|
||||
}
|
||||
}
|
||||
|
||||
impl Service<Name> for ValidatedResolver {
|
||||
type Response = ValidatedAddrs;
|
||||
type Error = io::Error;
|
||||
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
|
||||
|
||||
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
|
||||
Poll::Ready(Ok(()))
|
||||
}
|
||||
|
||||
fn call(&mut self, name: Name) -> Self::Future {
|
||||
let dns_cache = Arc::clone(&self.dns_cache);
|
||||
let host = name.as_str().to_string();
|
||||
Box::pin(async move {
|
||||
if let Some(addrs) = dns_cache.get_by_host(&host).await {
|
||||
return Ok(ValidatedAddrs {
|
||||
inner: (*addrs).clone().into_iter(),
|
||||
});
|
||||
}
|
||||
|
||||
let resolved = target_filter::resolve_public_addrs(&host, 0, dns_cache.as_ref())
|
||||
.await
|
||||
.map_err(|err| io::Error::other(err.to_string()))?;
|
||||
Ok(ValidatedAddrs {
|
||||
inner: resolved.into_iter(),
|
||||
})
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct InstrumentedConnector {
|
||||
http: HttpConnector<ValidatedResolver>,
|
||||
tls_config: Arc<ClientConfig>,
|
||||
}
|
||||
|
||||
impl Service<Uri> for InstrumentedConnector {
|
||||
type Response = TimedConn;
|
||||
type Error = BoxError;
|
||||
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
|
||||
|
||||
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
|
||||
self.http.poll_ready(cx).map_err(Into::into)
|
||||
}
|
||||
|
||||
fn call(&mut self, dst: Uri) -> Self::Future {
|
||||
let scheme = dst.scheme_str().map(|value| value.to_ascii_lowercase());
|
||||
let tls_config = Arc::clone(&self.tls_config);
|
||||
let connecting = self.http.call(dst.clone());
|
||||
let connect_start = std::time::Instant::now();
|
||||
|
||||
Box::pin(async move {
|
||||
match scheme.as_deref() {
|
||||
Some("http") => {
|
||||
let tcp = connecting.await.map_err(|err| Box::new(err) as BoxError)?;
|
||||
let connect_ms = connect_start.elapsed().as_millis() as u64;
|
||||
Ok(TimedConn::new(
|
||||
MaybeHttpsStream::Http(tcp),
|
||||
ConnectTiming {
|
||||
connect_ms,
|
||||
tls_ms: 0,
|
||||
},
|
||||
))
|
||||
}
|
||||
Some("https") => {
|
||||
let server_name = resolve_server_name(&dst)?;
|
||||
let tcp = connecting.await.map_err(|err| Box::new(err) as BoxError)?;
|
||||
let connect_ms = connect_start.elapsed().as_millis() as u64;
|
||||
|
||||
let tls_start = std::time::Instant::now();
|
||||
let tls_stream = TlsConnector::from(tls_config)
|
||||
.connect(server_name, tcp.into_inner())
|
||||
.await
|
||||
.map_err(io::Error::other)?;
|
||||
let tls_ms = tls_start.elapsed().as_millis() as u64;
|
||||
|
||||
Ok(TimedConn::new(
|
||||
MaybeHttpsStream::Https(TokioIo::new(tls_stream)),
|
||||
ConnectTiming { connect_ms, tls_ms },
|
||||
))
|
||||
}
|
||||
Some(other) => Err(io::Error::other(format!("unsupported scheme {other}")).into()),
|
||||
None => Err(io::Error::other("missing scheme").into()),
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub fn build_upstream_client(config: &Config, dns_cache: Arc<DnsCache>) -> UpstreamClient {
|
||||
let mut http = HttpConnector::new_with_resolver(ValidatedResolver::new(dns_cache));
|
||||
http.enforce_http(false);
|
||||
http.set_connect_timeout(Some(Duration::from_secs(
|
||||
config.upstream_connect_timeout_secs,
|
||||
)));
|
||||
http.set_nodelay(config.upstream_tcp_nodelay);
|
||||
if config.upstream_tcp_keepalive_secs > 0 {
|
||||
http.set_keepalive(Some(Duration::from_secs(
|
||||
config.upstream_tcp_keepalive_secs,
|
||||
)));
|
||||
} else {
|
||||
http.set_keepalive(None);
|
||||
}
|
||||
|
||||
let connector = InstrumentedConnector {
|
||||
http,
|
||||
tls_config: build_tls_config(),
|
||||
};
|
||||
|
||||
let mut builder = Client::builder(TokioExecutor::new());
|
||||
builder.pool_max_idle_per_host(config.upstream_pool_max_idle_per_host);
|
||||
builder.pool_idle_timeout(Duration::from_secs(config.upstream_pool_idle_timeout_secs));
|
||||
builder.pool_timer(TokioTimer::new());
|
||||
builder.build(connector)
|
||||
}
|
||||
|
||||
pub fn resolve_request_timing<B>(
|
||||
response: &Response<B>,
|
||||
connection_acquire_ms: Option<u64>,
|
||||
ttfb_ms: u64,
|
||||
) -> RequestTiming {
|
||||
let raw = response
|
||||
.extensions()
|
||||
.get::<ConnectTiming>()
|
||||
.copied()
|
||||
.unwrap_or_default();
|
||||
|
||||
let raw_connection_ms = raw.connect_ms.saturating_add(raw.tls_ms);
|
||||
let measured_acquire_ms = connection_acquire_ms.unwrap_or(raw_connection_ms.min(ttfb_ms));
|
||||
let likely_reused = measured_acquire_ms <= 5 && raw_connection_ms > 0;
|
||||
let connector_matches_request = raw_connection_ms <= measured_acquire_ms.saturating_add(25);
|
||||
|
||||
let (connect_ms, tls_ms) = if likely_reused || !connector_matches_request {
|
||||
(0, 0)
|
||||
} else {
|
||||
(raw.connect_ms, raw.tls_ms)
|
||||
};
|
||||
|
||||
RequestTiming {
|
||||
connection_acquire_ms: measured_acquire_ms,
|
||||
connect_ms,
|
||||
tls_ms,
|
||||
response_wait_ms: ttfb_ms.saturating_sub(measured_acquire_ms),
|
||||
connection_reused: likely_reused,
|
||||
}
|
||||
}
|
||||
|
||||
fn build_tls_config() -> Arc<ClientConfig> {
|
||||
let root_store =
|
||||
rustls::RootCertStore::from_iter(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
|
||||
let mut config = ClientConfig::builder()
|
||||
.with_root_certificates(root_store)
|
||||
.with_no_client_auth();
|
||||
config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()];
|
||||
Arc::new(config)
|
||||
}
|
||||
|
||||
fn resolve_server_name(uri: &Uri) -> Result<ServerName<'static>, BoxError> {
|
||||
let host = uri.host().ok_or_else(|| io::Error::other("missing host"))?;
|
||||
let host = host.trim_start_matches('[').trim_end_matches(']');
|
||||
|
||||
if let Ok(ip) = host.parse::<IpAddr>() {
|
||||
return Ok(ServerName::from(ip));
|
||||
}
|
||||
|
||||
Ok(ServerName::try_from(host.to_string())?)
|
||||
}
|
||||
|
||||
pub struct TimedConn {
|
||||
inner: MaybeHttpsStream,
|
||||
timing: ConnectTiming,
|
||||
}
|
||||
|
||||
impl TimedConn {
|
||||
fn new(inner: MaybeHttpsStream, timing: ConnectTiming) -> Self {
|
||||
Self { inner, timing }
|
||||
}
|
||||
}
|
||||
|
||||
impl Connection for TimedConn {
|
||||
fn connected(&self) -> Connected {
|
||||
self.inner.connected().extra(self.timing)
|
||||
}
|
||||
}
|
||||
|
||||
impl rt::Read for TimedConn {
|
||||
fn poll_read(
|
||||
mut self: Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
buf: rt::ReadBufCursor<'_>,
|
||||
) -> Poll<Result<(), io::Error>> {
|
||||
Pin::new(&mut self.inner).poll_read(cx, buf)
|
||||
}
|
||||
}
|
||||
|
||||
impl rt::Write for TimedConn {
|
||||
fn poll_write(
|
||||
mut self: Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
buf: &[u8],
|
||||
) -> Poll<Result<usize, io::Error>> {
|
||||
Pin::new(&mut self.inner).poll_write(cx, buf)
|
||||
}
|
||||
|
||||
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
|
||||
Pin::new(&mut self.inner).poll_flush(cx)
|
||||
}
|
||||
|
||||
fn poll_shutdown(
|
||||
mut self: Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
) -> Poll<Result<(), io::Error>> {
|
||||
Pin::new(&mut self.inner).poll_shutdown(cx)
|
||||
}
|
||||
|
||||
fn is_write_vectored(&self) -> bool {
|
||||
self.inner.is_write_vectored()
|
||||
}
|
||||
|
||||
fn poll_write_vectored(
|
||||
mut self: Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
bufs: &[std::io::IoSlice<'_>],
|
||||
) -> Poll<Result<usize, io::Error>> {
|
||||
Pin::new(&mut self.inner).poll_write_vectored(cx, bufs)
|
||||
}
|
||||
}
|
||||
|
||||
pub enum MaybeHttpsStream {
|
||||
Http(PlainStream),
|
||||
Https(TlsStream),
|
||||
}
|
||||
|
||||
impl Connection for MaybeHttpsStream {
|
||||
fn connected(&self) -> Connected {
|
||||
match self {
|
||||
Self::Http(stream) => stream.connected(),
|
||||
Self::Https(stream) => {
|
||||
let (tcp, tls) = stream.inner().get_ref();
|
||||
if tls.alpn_protocol() == Some(b"h2") {
|
||||
tcp.connected().negotiated_h2()
|
||||
} else {
|
||||
tcp.connected()
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl rt::Read for MaybeHttpsStream {
|
||||
fn poll_read(
|
||||
self: Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
buf: rt::ReadBufCursor<'_>,
|
||||
) -> Poll<Result<(), io::Error>> {
|
||||
match Pin::get_mut(self) {
|
||||
Self::Http(stream) => Pin::new(stream).poll_read(cx, buf),
|
||||
Self::Https(stream) => Pin::new(stream).poll_read(cx, buf),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl rt::Write for MaybeHttpsStream {
|
||||
fn poll_write(
|
||||
self: Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
buf: &[u8],
|
||||
) -> Poll<Result<usize, io::Error>> {
|
||||
match Pin::get_mut(self) {
|
||||
Self::Http(stream) => Pin::new(stream).poll_write(cx, buf),
|
||||
Self::Https(stream) => Pin::new(stream).poll_write(cx, buf),
|
||||
}
|
||||
}
|
||||
|
||||
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
|
||||
match Pin::get_mut(self) {
|
||||
Self::Http(stream) => Pin::new(stream).poll_flush(cx),
|
||||
Self::Https(stream) => Pin::new(stream).poll_flush(cx),
|
||||
}
|
||||
}
|
||||
|
||||
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
|
||||
match Pin::get_mut(self) {
|
||||
Self::Http(stream) => Pin::new(stream).poll_shutdown(cx),
|
||||
Self::Https(stream) => Pin::new(stream).poll_shutdown(cx),
|
||||
}
|
||||
}
|
||||
|
||||
fn is_write_vectored(&self) -> bool {
|
||||
match self {
|
||||
Self::Http(stream) => stream.is_write_vectored(),
|
||||
Self::Https(stream) => stream.is_write_vectored(),
|
||||
}
|
||||
}
|
||||
|
||||
fn poll_write_vectored(
|
||||
self: Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
bufs: &[std::io::IoSlice<'_>],
|
||||
) -> Poll<Result<usize, io::Error>> {
|
||||
match Pin::get_mut(self) {
|
||||
Self::Http(stream) => Pin::new(stream).poll_write_vectored(cx, bufs),
|
||||
Self::Https(stream) => Pin::new(stream).poll_write_vectored(cx, bufs),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use hyper::Response;
|
||||
|
||||
#[test]
|
||||
fn fresh_connection_uses_connector_breakdown() {
|
||||
let mut response = Response::new(());
|
||||
response.extensions_mut().insert(ConnectTiming {
|
||||
connect_ms: 80,
|
||||
tls_ms: 40,
|
||||
});
|
||||
|
||||
let timing = resolve_request_timing(&response, Some(125), 600);
|
||||
|
||||
assert_eq!(timing.connection_acquire_ms, 125);
|
||||
assert_eq!(timing.connect_ms, 80);
|
||||
assert_eq!(timing.tls_ms, 40);
|
||||
assert_eq!(timing.response_wait_ms, 475);
|
||||
assert!(!timing.connection_reused);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn reused_connection_zeroes_stale_connect_timings() {
|
||||
let mut response = Response::new(());
|
||||
response.extensions_mut().insert(ConnectTiming {
|
||||
connect_ms: 70,
|
||||
tls_ms: 30,
|
||||
});
|
||||
|
||||
let timing = resolve_request_timing(&response, Some(0), 310);
|
||||
|
||||
assert_eq!(timing.connection_acquire_ms, 0);
|
||||
assert_eq!(timing.connect_ms, 0);
|
||||
assert_eq!(timing.tls_ms, 0);
|
||||
assert_eq!(timing.response_wait_ms, 310);
|
||||
assert!(timing.connection_reused);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn falls_back_to_connector_timings_when_capture_missing() {
|
||||
let mut response = Response::new(());
|
||||
response.extensions_mut().insert(ConnectTiming {
|
||||
connect_ms: 55,
|
||||
tls_ms: 25,
|
||||
});
|
||||
|
||||
let timing = resolve_request_timing(&response, None, 400);
|
||||
|
||||
assert_eq!(timing.connection_acquire_ms, 80);
|
||||
assert_eq!(timing.connect_ms, 55);
|
||||
assert_eq!(timing.tls_ms, 25);
|
||||
assert_eq!(timing.response_wait_ms, 320);
|
||||
assert!(!timing.connection_reused);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user