use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use pyo3::prelude::*;
use pyo3::types::PyBytes;
use pyo3_async_runtimes::tokio::future_into_py;
use super::to_py_err;
use crate::documents::SessionType;
use crate::{
DocumentSpec, OutputStream, PortForwardConfig, PortForwarder, Session, SessionConfig,
SessionManager, ShutdownSignal,
};
#[pyclass(name = "SessionManager", frozen)]
#[derive(Debug)]
pub struct PySessionManager {
inner: SessionManager,
}
#[pymethods]
impl PySessionManager {
#[staticmethod]
#[pyo3(signature = (region = None))]
#[allow(clippy::new_ret_no_self)]
fn new(py: Python<'_>, region: Option<String>) -> PyResult<Bound<'_, PyAny>> {
future_into_py(py, async move {
let inner = match region {
Some(region) => SessionManager::for_region(region).await,
None => SessionManager::new().await,
}
.map_err(to_py_err)?;
Ok(PySessionManager { inner })
})
}
#[pyo3(signature = (
target,
document_name = None,
parameters = None,
reason = None,
ready_timeout = 30.0,
))]
fn start_session<'py>(
&self,
py: Python<'py>,
target: String,
document_name: Option<String>,
parameters: Option<HashMap<String, Vec<String>>>,
reason: Option<String>,
ready_timeout: f64,
) -> PyResult<Bound<'py, PyAny>> {
let manager = self.inner.clone();
let config = build_config(target, document_name, parameters, reason, ready_timeout)?;
future_into_py(py, async move {
let session = manager.start_session(config).await.map_err(to_py_err)?;
Ok(PySession {
inner: Arc::new(session),
})
})
}
#[pyo3(signature = (target, remote_port, reason = None))]
fn start_port_forward<'py>(
&self,
py: Python<'py>,
target: String,
remote_port: u16,
reason: Option<String>,
) -> PyResult<Bound<'py, PyAny>> {
let manager = self.inner.clone();
let config = SessionConfig {
document: Some(DocumentSpec::new(
&crate::documents::PortForwardingSession::new(remote_port),
)),
reason,
..SessionConfig::new(target)
};
future_into_py(py, async move {
let session = manager.start_session(config).await.map_err(to_py_err)?;
Ok(PySession {
inner: Arc::new(session),
})
})
}
#[pyo3(signature = (target, host, remote_port, reason = None))]
fn start_remote_port_forward<'py>(
&self,
py: Python<'py>,
target: String,
host: String,
remote_port: u16,
reason: Option<String>,
) -> PyResult<Bound<'py, PyAny>> {
let manager = self.inner.clone();
let config = SessionConfig {
document: Some(DocumentSpec::new(
&crate::documents::PortForwardingToRemoteHost::new(host, remote_port),
)),
reason,
..SessionConfig::new(target)
};
future_into_py(py, async move {
let session = manager.start_session(config).await.map_err(to_py_err)?;
Ok(PySession {
inner: Arc::new(session),
})
})
}
fn terminate_session<'py>(
&self,
py: Python<'py>,
session_id: String,
) -> PyResult<Bound<'py, PyAny>> {
let manager = self.inner.clone();
future_into_py(py, async move {
manager
.terminate_session(&session_id)
.await
.map_err(to_py_err)
})
}
fn __repr__(&self) -> String {
"SessionManager()".to_owned()
}
}
fn build_config(
target: String,
document_name: Option<String>,
parameters: Option<HashMap<String, Vec<String>>>,
reason: Option<String>,
ready_timeout: f64,
) -> PyResult<SessionConfig> {
let parameters = parameters.unwrap_or_default();
let document = document_name.map(|name| {
let session_type = if name.starts_with("AWS-StartPortForwardingSession") {
SessionType::Port
} else {
SessionType::StandardStream
};
DocumentSpec {
name,
parameters,
session_type,
}
});
Ok(SessionConfig {
document,
reason,
ready_timeout: seconds(ready_timeout, "ready_timeout")?,
..SessionConfig::new(target)
})
}
fn seconds(value: f64, name: &str) -> PyResult<Duration> {
if !value.is_finite() || value < 0.0 || value > Duration::MAX.as_secs_f64() {
return Err(pyo3::exceptions::PyValueError::new_err(format!(
"{name} must be a finite, non-negative number of seconds, got {value}"
)));
}
Ok(Duration::from_secs_f64(value))
}
#[pyclass(name = "Session", frozen)]
#[derive(Debug)]
pub struct PySession {
pub(crate) inner: Arc<Session>,
}
#[pymethods]
impl PySession {
#[getter]
fn id(&self) -> &str {
self.inner.id()
}
#[getter]
fn target(&self) -> &str {
&self.inner.config().target
}
#[getter]
fn agent_version(&self) -> Option<&str> {
self.inner.agent_version()
}
#[getter]
fn banner(&self) -> Option<String> {
self.inner.banner()
}
#[getter]
fn exit_code(&self) -> Option<i32> {
self.inner.exit_code()
}
#[getter]
fn is_encrypted(&self) -> bool {
self.inner.is_encrypted()
}
#[getter]
fn is_ready(&self) -> bool {
self.inner.is_ready()
}
#[getter]
fn is_closed(&self) -> bool {
self.inner.is_closed()
}
#[getter]
fn close_reason(&self) -> Option<String> {
self.inner.close_reason().map(|r| r.to_string())
}
fn wait_ready<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
let session = Arc::clone(&self.inner);
future_into_py(
py,
async move { session.wait_ready().await.map_err(to_py_err) },
)
}
fn wait_closed<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
let session = Arc::clone(&self.inner);
future_into_py(py, async move {
session.closed().await;
Ok(())
})
}
fn send<'py>(&self, py: Python<'py>, data: Vec<u8>) -> PyResult<Bound<'py, PyAny>> {
let session = Arc::clone(&self.inner);
future_into_py(py, async move {
session
.send(bytes::Bytes::from(data))
.await
.map_err(to_py_err)
})
}
fn send_terminal_size<'py>(
&self,
py: Python<'py>,
cols: u16,
rows: u16,
) -> PyResult<Bound<'py, PyAny>> {
let session = Arc::clone(&self.inner);
future_into_py(py, async move {
session
.send_terminal_size(cols, rows)
.await
.map_err(to_py_err)
})
}
fn output(&self) -> PyOutputStream {
PyOutputStream {
inner: Arc::new(tokio::sync::Mutex::new(self.inner.output())),
}
}
fn terminate<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
let session = Arc::clone(&self.inner);
future_into_py(
py,
async move { session.terminate().await.map_err(to_py_err) },
)
}
fn __aenter__<'py>(slf: PyRef<'py, Self>, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
let session = Arc::clone(&slf.inner);
let handle: Py<PySession> = slf.into();
future_into_py(py, async move {
session.wait_ready().await.map_err(to_py_err)?;
Ok(handle)
})
}
#[pyo3(signature = (_exc_type = None, _exc_value = None, _traceback = None))]
fn __aexit__<'py>(
&self,
py: Python<'py>,
_exc_type: Option<Bound<'_, PyAny>>,
_exc_value: Option<Bound<'_, PyAny>>,
_traceback: Option<Bound<'_, PyAny>>,
) -> PyResult<Bound<'py, PyAny>> {
let session = Arc::clone(&self.inner);
future_into_py(py, async move {
let _ = session.terminate().await;
Ok(false) })
}
fn __repr__(&self) -> String {
format!(
"Session(id={:?}, target={:?}, ready={}, closed={})",
self.inner.id(),
self.inner.config().target,
self.inner.is_ready(),
self.inner.is_closed(),
)
}
}
#[pyclass(name = "OutputStream", frozen)]
#[derive(Debug)]
pub struct PyOutputStream {
inner: Arc<tokio::sync::Mutex<OutputStream>>,
}
#[pymethods]
impl PyOutputStream {
fn __aiter__(slf: PyRef<'_, Self>) -> PyRef<'_, Self> {
slf
}
fn __anext__<'py>(&self, py: Python<'py>) -> PyResult<Option<Bound<'py, PyAny>>> {
use pyo3::exceptions::PyStopAsyncIteration;
let stream = Arc::clone(&self.inner);
let future = future_into_py(py, async move {
match stream.lock().await.recv().await {
Some(chunk) => Ok(Python::attach(|py| {
PyBytes::new(py, &chunk).unbind().into_any()
})),
None => Err(PyStopAsyncIteration::new_err(())),
}
})?;
Ok(Some(future))
}
fn __repr__(&self) -> String {
"OutputStream()".to_owned()
}
}
#[pyclass(name = "PortForwarder")]
#[derive(Debug)]
pub struct PyPortForwarder {
forwarder: Option<PortForwarder>,
local_addr: std::net::SocketAddr,
shutdown: ShutdownSignal,
}
#[pymethods]
impl PyPortForwarder {
#[staticmethod]
#[pyo3(signature = (local_addr = "127.0.0.1:0", max_connections = 100))]
fn bind<'py>(
py: Python<'py>,
local_addr: &str,
max_connections: usize,
) -> PyResult<Bound<'py, PyAny>> {
let addr: std::net::SocketAddr = local_addr.parse().map_err(|e| {
pyo3::exceptions::PyValueError::new_err(format!(
"{local_addr:?} is not a valid address: {e}"
))
})?;
future_into_py(py, async move {
let forwarder = PortForwarder::bind(PortForwardConfig {
local_addr: addr,
max_connections,
..Default::default()
})
.await
.map_err(to_py_err)?;
Ok(PyPortForwarder {
local_addr: forwarder.local_addr(),
forwarder: Some(forwarder),
shutdown: ShutdownSignal::new(),
})
})
}
#[getter]
fn address(&self) -> String {
self.local_addr.to_string()
}
#[getter]
fn port(&self) -> u16 {
self.local_addr.port()
}
fn forward<'py>(
&mut self,
py: Python<'py>,
session: PyRef<'_, PySession>,
) -> PyResult<Bound<'py, PyAny>> {
let forwarder = self.forwarder.take().ok_or_else(|| {
pyo3::exceptions::PyRuntimeError::new_err(
"this PortForwarder has already been used; bind a new one",
)
})?;
let session = Arc::clone(&session.inner);
let shutdown = self.shutdown.clone();
future_into_py(py, async move {
forwarder
.forward(session, shutdown)
.await
.map_err(to_py_err)
})
}
fn stop(&self) {
self.shutdown.shutdown();
}
fn __repr__(&self) -> String {
format!("PortForwarder(address={:?})", self.local_addr.to_string())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn invalid_timeouts_raise_instead_of_panicking() {
for bad in [f64::NAN, f64::INFINITY, -1.0, f64::MAX] {
assert!(
seconds(bad, "ready_timeout").is_err(),
"{bad} must be rejected"
);
}
assert_eq!(seconds(1.5, "t").unwrap(), Duration::from_millis(1500));
assert_eq!(seconds(0.0, "t").unwrap(), Duration::ZERO);
}
#[test]
fn document_names_map_to_the_right_session_type() {
let port = build_config(
"i-abc".into(),
Some("AWS-StartPortForwardingSessionToRemoteHost".into()),
None,
None,
30.0,
)
.unwrap();
assert_eq!(port.session_type(), SessionType::Port);
let ssh = build_config(
"i-abc".into(),
Some("AWS-StartSSHSession".into()),
None,
None,
30.0,
)
.unwrap();
assert_eq!(ssh.session_type(), SessionType::StandardStream);
let shell = build_config("i-abc".into(), None, None, None, 30.0).unwrap();
assert!(shell.document.is_none());
}
}