use crate::session::ClientSession;
use crate::types::TransferMetrics;
use anyhow::{Context, Result};
use tokio::io::BufReader;
use tracing::info;
mod error {
pub const ROUTER_REQUIRED: &str = "Hybrid mode requires a router";
pub const BACKEND_NOT_FOUND: &str = "Backend not found";
}
struct StatefulSessionGuard<'a> {
metrics: &'a crate::metrics::MetricsCollector,
ended: bool,
}
impl<'a> StatefulSessionGuard<'a> {
fn start(metrics: &'a crate::metrics::MetricsCollector) -> Self {
metrics.stateful_session_started();
Self {
metrics,
ended: false,
}
}
}
impl Drop for StatefulSessionGuard<'_> {
fn drop(&mut self) {
if !self.ended {
self.metrics.stateful_session_ended();
}
}
}
fn stateful_initial_client_bytes(
carried_client_to_backend_bytes: u64,
initial_request: &crate::protocol::RequestContext,
) -> u64 {
carried_client_to_backend_bytes + initial_request.request_wire_len().as_u64()
}
impl ClientSession {
pub(super) async fn switch_to_stateful_mode<R, W>(
&self,
client_reader: BufReader<R>,
client_write: W,
initial_request: &crate::protocol::RequestContext,
client_to_backend_bytes: u64,
backend_to_client_bytes: u64,
) -> Result<TransferMetrics, crate::session::SessionError>
where
R: tokio::io::AsyncRead + Unpin,
W: tokio::io::AsyncWrite + Unpin,
{
self.mode_state.switch_to_stateful();
let (pooled_conn, backend_id, _pending_guard, provider) = self
.acquire_stateful_backend()
.await
.context("Failed to acquire backend for stateful mode")?;
let mut conn_guard = crate::pool::ConnectionGuard::new(pooled_conn, provider);
let _session_guard = StatefulSessionGuard::start(&self.metrics);
info!(
client = %self.client_addr,
backend = ?backend_id,
"Switched to stateful mode"
);
initial_request
.write_wire_to(&mut **conn_guard)
.await
.context("Failed to send initial request to backend")?;
let initial_bytes = stateful_initial_client_bytes(client_to_backend_bytes, initial_request);
let mut state = crate::session::state::SessionLoopState::from_initial_bytes(
initial_bytes,
backend_to_client_bytes,
self.auth_handler.is_enabled(),
);
state.mark_backend_request_sent(initial_request.kind());
let (backend_read, backend_write) = tokio::io::split(&mut **conn_guard);
let result = self
.run_stateful_proxy_loop(
client_reader,
client_write,
backend_read,
backend_write,
state,
backend_id,
)
.await;
if result.is_ok() {
let _conn = conn_guard.release();
}
result.map_err(crate::session::SessionError::from)
}
async fn acquire_stateful_backend(
&self,
) -> Result<(
deadpool::managed::Object<crate::pool::deadpool_connection::TcpManager>,
crate::types::BackendId,
crate::router::CommandGuard,
crate::pool::DeadpoolConnectionProvider,
)> {
let router = self
.router
.as_ref()
.ok_or_else(|| anyhow::anyhow!(error::ROUTER_REQUIRED))?;
let backend_id = router.route_without_availability(self.client_id)?;
let pending_guard = crate::router::CommandGuard::new(router.clone(), backend_id);
let provider = router
.backend_provider(backend_id)
.ok_or_else(|| anyhow::anyhow!("{}: {:?}", error::BACKEND_NOT_FOUND, backend_id))?;
let provider = provider.clone();
let conn = provider.get_pooled_connection().await?;
Ok((conn, backend_id, pending_guard, provider))
}
}
#[cfg(test)]
mod tests {
use crate::protocol::RequestContext;
#[test]
fn test_error_messages_are_descriptive() {
use super::error::*;
assert!(ROUTER_REQUIRED.contains("router"));
assert!(BACKEND_NOT_FOUND.contains("Backend"));
}
#[test]
fn stateful_initial_client_bytes_uses_typed_wire_len() {
let request = RequestContext::parse(b"group alt.test\r\n").expect("valid request line");
assert_eq!(
super::stateful_initial_client_bytes(10, &request),
10 + "group alt.test\r\n".len() as u64
);
}
}