use bytes::Bytes;
use reqwest::StatusCode;
use std::time::Instant;
use tracing::{error, warn};
use crate::balancer::LoadBalancer;
use crate::endpoint::{ErrorKind, LoadBalancerError, RequestOutcome, RpcEndpoint};
use crate::metrics::{
REQUEST_LATENCY_PER_ENDPOINT, REQUEST_TIMEOUTS, RPC_REQUESTS_FAILED, RPC_REQUESTS_METHOD,
RPC_REQUESTS_SUCCEEDED, RPC_REQUESTS_TOTAL, UPSTREAM_RATE_LIMITED_TOTAL,
};
pub fn response_has_error(bytes: &Bytes) -> bool {
let key = b"\"error\":";
if let Some(start) = bytes.windows(key.len()).position(|window| window == key) {
let mut idx = start + key.len();
while idx < bytes.len() && (bytes[idx] as char).is_whitespace() {
idx += 1;
}
let null_val = b"null";
if bytes.get(idx..idx + null_val.len()) == Some(null_val) {
return false;
}
return true;
}
false
}
pub async fn forward_request(
balancer: &LoadBalancer,
request_body: Bytes,
endpoint: &RpcEndpoint,
method: &str,
) -> Result<Bytes, LoadBalancerError> {
let timer = REQUEST_LATENCY_PER_ENDPOINT.with_label_values(&[&endpoint.name]).start_timer();
let start_time = Instant::now();
RPC_REQUESTS_TOTAL.inc();
RPC_REQUESTS_METHOD.with_label_values(&[method]).inc();
let client = balancer.client.read().clone();
let response_result = client
.post(&endpoint.url)
.header(reqwest::header::CONTENT_TYPE, "application/json")
.body(request_body)
.send()
.await;
let elapsed_ms = start_time.elapsed().as_millis() as u64;
timer.observe_duration();
let smoothing_factor = *balancer.latency_smoothing_factor.read();
let outcome: RequestOutcome;
match &response_result {
Ok(resp) => {
let status = resp.status();
if status.is_success() {
outcome = RequestOutcome::Success { latency_ms: elapsed_ms };
} else if status == StatusCode::TOO_MANY_REQUESTS {
outcome = RequestOutcome::Failure {
error_kind: ErrorKind::RateLimit,
latency_ms: elapsed_ms,
};
} else if status.is_server_error() {
outcome = RequestOutcome::Failure {
error_kind: ErrorKind::Http5xx,
latency_ms: elapsed_ms,
};
} else if status.is_client_error() {
outcome = RequestOutcome::Failure {
error_kind: ErrorKind::Http4xx,
latency_ms: elapsed_ms,
};
} else {
outcome = RequestOutcome::Failure {
error_kind: ErrorKind::ConnectionError,
latency_ms: elapsed_ms,
};
}
}
Err(e) => {
if e.is_timeout() {
let timeout_ms = (*balancer.timeout_secs.read()) * 1000;
outcome = RequestOutcome::Failure {
error_kind: ErrorKind::Timeout,
latency_ms: timeout_ms,
};
} else {
outcome = RequestOutcome::Failure {
error_kind: ErrorKind::ConnectionError,
latency_ms: elapsed_ms,
};
}
}
}
endpoint.metrics.update(&outcome, smoothing_factor);
let response = match response_result {
Ok(resp) => resp,
Err(e) => {
if e.is_timeout() {
REQUEST_TIMEOUTS.with_label_values(&[&endpoint.name]).inc();
}
balancer.mark_unhealthy(&endpoint.url).await;
RPC_REQUESTS_FAILED.with_label_values(&[&endpoint.name]).inc();
error!(endpoint = %endpoint.name, error = %e, "Network error from endpoint");
return Err(LoadBalancerError::UpstreamError(e.to_string()));
}
};
match response.status() {
reqwest::StatusCode::TOO_MANY_REQUESTS => {
UPSTREAM_RATE_LIMITED_TOTAL.with_label_values(&[&endpoint.name]).inc();
RPC_REQUESTS_FAILED.with_label_values(&[&endpoint.name]).inc();
balancer.mark_rate_limited(&endpoint.url).await;
error!(endpoint = %endpoint.name, "Upstream rate-limited (429)");
return Err(LoadBalancerError::RateLimited(endpoint.url.to_string()));
}
status if !status.is_success() => {
balancer.mark_unhealthy(&endpoint.url).await;
RPC_REQUESTS_FAILED.with_label_values(&[&endpoint.name]).inc();
warn!(endpoint = %endpoint.name, status = %status, "Non-success response");
return Err(LoadBalancerError::UpstreamError(format!(
"Upstream returned non-success status: {}",
status
)));
}
_ => {}
}
let response_bytes = match response.bytes().await {
Ok(bytes) => bytes,
Err(e) => {
balancer.mark_unhealthy(&endpoint.url).await;
RPC_REQUESTS_FAILED.with_label_values(&[&endpoint.name]).inc();
error!(endpoint = %endpoint.name, error = %e, "Failed to read response body");
return Err(LoadBalancerError::UpstreamError(e.to_string()));
}
};
if response_has_error(&response_bytes) {
RPC_REQUESTS_FAILED.with_label_values(&[&endpoint.name]).inc();
} else {
RPC_REQUESTS_SUCCEEDED.with_label_values(&[&endpoint.name]).inc();
}
Ok(response_bytes)
}