use std::time::Duration;
use crate::protocol::FrontendMessage;
use crate::auth::Codec;
use crate::config::Config;
use crate::error::{Error, PgError, Result};
use crate::transport::{AsyncTransport, SslMode};
#[cfg(feature = "tracing")]
use crate::tracing_ext::TARGET_CANCEL;
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct CancelToken {
pub(crate) host: String,
pub(crate) port: u16,
pub(crate) process_id: i32,
pub(crate) secret_key: i32,
pub(crate) ssl_mode: SslMode,
pub(crate) accept_invalid_certs: bool,
}
impl CancelToken {
#[cfg(feature = "tls")]
fn build_tls_config(&self) -> crate::transport::TlsConfig {
crate::transport::TlsConfig::new(self.ssl_mode, &self.host)
.accept_invalid_certs(self.accept_invalid_certs)
}
#[must_use = "cancel errors should be checked"]
pub async fn cancel(&self) -> Result<()> {
self.cancel_with_timeout(None).await
}
#[must_use = "cancel errors should be checked"]
pub async fn cancel_with_timeout(&self, timeout: Option<Duration>) -> Result<()> {
#[cfg(feature = "tracing")]
tracing::debug!(target: TARGET_CANCEL, process_id = self.process_id, "Sending cancel request");
let mut config = Config::new()
.host(&self.host)
.port(self.port)
.user("cancel");
config = config.ssl_mode(self.ssl_mode);
if self.accept_invalid_certs {
config = config.accept_invalid_certs(true);
}
let raw_transport = build_cancel_transport(&config, timeout).await?;
#[cfg(feature = "tls")]
let mut transport = if self.ssl_mode != SslMode::Disable {
let tls_config = self.build_tls_config();
crate::transport::negotiate_tls(raw_transport, &tls_config)
.await
.map_err(PgError::Transport)?
} else {
crate::transport::PgTransport::Plain(crate::transport::BufferedTransport::new(
raw_transport,
))
};
#[cfg(not(feature = "tls"))]
let mut transport = crate::transport::PgTransport::Plain(
crate::transport::BufferedTransport::new(raw_transport),
);
let mut codec = Codec::new();
codec
.send(
&mut transport,
&FrontendMessage::CancelRequest {
process_id: self.process_id,
secret_key: self.secret_key,
},
)
.await
.map_err(Error::from)?;
let _ = transport.shutdown().await;
#[cfg(feature = "tracing")]
tracing::info!(target: TARGET_CANCEL, process_id = self.process_id, "Cancel request sent");
Ok(())
}
pub fn process_id(&self) -> i32 {
self.process_id
}
pub fn secret_key(&self) -> i32 {
self.secret_key
}
}
#[cfg(target_arch = "wasm32")]
async fn build_cancel_transport(
config: &Config,
timeout: Option<Duration>,
) -> Result<crate::transport::ClientTransport> {
use crate::transport::{connect_with_timeout, ClientTransport};
let tcp = connect_with_timeout(config.get_host(), config.get_port(), timeout)
.await
.map_err(PgError::Transport)?;
Ok(ClientTransport::Wasi(tcp))
}
#[cfg(all(not(target_arch = "wasm32"), feature = "tokio-transport"))]
async fn build_cancel_transport(
config: &Config,
timeout: Option<Duration>,
) -> Result<crate::transport::ClientTransport> {
use crate::transport::{connect_with_timeout, ClientTransport};
let tcp = connect_with_timeout(config.get_host(), config.get_port(), timeout)
.await
.map_err(PgError::Transport)?;
Ok(ClientTransport::Tokio(tcp))
}
#[cfg(all(
not(target_arch = "wasm32"),
not(feature = "tokio-transport"),
feature = "test-native"
))]
async fn build_cancel_transport(
config: &Config,
timeout: Option<Duration>,
) -> Result<crate::transport::ClientTransport> {
use crate::transport::{ClientTransport, NativeTcpTransport};
let tcp =
NativeTcpTransport::connect_with_timeout(config.get_host(), config.get_port(), timeout)
.map_err(PgError::Transport)?;
Ok(ClientTransport::Native(tcp))
}
#[cfg(all(
not(target_arch = "wasm32"),
not(feature = "tokio-transport"),
not(feature = "test-native")
))]
async fn build_cancel_transport(
_config: &Config,
_timeout: Option<Duration>,
) -> Result<crate::transport::ClientTransport> {
Err(PgError::Unsupported(
"no transport available for cancellation. Enable the 'tokio-transport' feature (recommended) or 'test-native' feature, or compile for wasm32-wasip2".into(),
))
}
#[cfg(test)]
mod tests {
use super::*;
fn make_token() -> CancelToken {
CancelToken {
host: "localhost".to_string(),
port: 5432,
process_id: 12345,
secret_key: 67890,
ssl_mode: SslMode::Disable,
accept_invalid_certs: false,
}
}
#[test]
fn test_cancel_token_clone() {
let token = make_token();
let cloned = token.clone();
assert_eq!(cloned.host, "localhost");
assert_eq!(cloned.port, 5432);
assert_eq!(cloned.process_id, 12345);
assert_eq!(cloned.secret_key, 67890);
}
#[test]
fn test_cancel_token_accessors() {
let token = CancelToken {
process_id: 42,
secret_key: 99,
..make_token()
};
assert_eq!(token.process_id(), 42);
assert_eq!(token.secret_key(), 99);
}
#[cfg(feature = "tls")]
#[test]
fn test_cancel_token_tls_config_preserves_hostname_verification() {
let token = CancelToken {
ssl_mode: SslMode::VerifyFull,
accept_invalid_certs: true,
..make_token()
};
let tls_config = token.build_tls_config();
assert_eq!(tls_config.mode, SslMode::VerifyFull);
assert_eq!(tls_config.server_name, "localhost");
assert!(tls_config.accept_invalid_certs);
}
}