use std::sync::{Arc, Mutex};
use pyo3::prelude::*;
use pyo3::types::{PyBytes, PyDict};
use crate::errors::map_err;
use crate::streaming::RuntimeLease;
fn validated_timeout(timeout: Option<f64>) -> PyResult<Option<std::time::Duration>> {
match timeout {
Some(secs) => {
if !secs.is_finite() || secs < 0.0 {
Err(pyo3::exceptions::PyValueError::new_err(
"timeout must be a finite, non-negative number",
))
} else {
Ok(Some(std::time::Duration::from_secs_f64(secs)))
}
}
None => Ok(None),
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum StreamVariant {
Tcp,
Tls,
Adapter,
}
impl From<eggfetch_core::network_stream::UpgradedStreamVariant> for StreamVariant {
fn from(value: eggfetch_core::network_stream::UpgradedStreamVariant) -> Self {
match value {
eggfetch_core::network_stream::UpgradedStreamVariant::Tcp => StreamVariant::Tcp,
eggfetch_core::network_stream::UpgradedStreamVariant::Tls => StreamVariant::Tls,
eggfetch_core::network_stream::UpgradedStreamVariant::Adapter => StreamVariant::Adapter,
}
}
}
#[allow(dead_code)]
pub(crate) type SharedStreamInner =
Arc<Mutex<Option<eggfetch_core::network_stream::UpgradedStream>>>;
#[pyclass(name = "NetworkStream")]
pub struct PyNetworkStream {
inner: SharedStreamInner,
metadata: Option<Arc<eggfetch_core::network_stream::ConnectionMetadata>>,
variant: Option<StreamVariant>,
runtime_handle: tokio::runtime::Handle,
runtime_lease: Option<RuntimeLease>,
}
impl std::fmt::Debug for PyNetworkStream {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let is_upgraded = self.inner.lock().is_ok_and(|g| g.is_some());
f.debug_struct("PyNetworkStream")
.field("is_upgraded", &is_upgraded)
.field("metadata", &self.metadata)
.field("variant", &self.variant)
.field("runtime_handle", &"<elided>")
.field("runtime_lease", &self.runtime_lease.is_some())
.finish()
}
}
impl Clone for PyNetworkStream {
fn clone(&self) -> Self {
Self {
inner: Arc::clone(&self.inner),
metadata: self.metadata.clone(),
variant: self.variant,
runtime_handle: self.runtime_handle.clone(),
runtime_lease: self.runtime_lease.clone(),
}
}
}
impl PyNetworkStream {
pub fn from_metadata(
metadata: Arc<eggfetch_core::network_stream::ConnectionMetadata>,
) -> PyResult<Self> {
let runtime_handle = tokio::runtime::Handle::try_current().map_err(|_| {
PyErr::new::<pyo3::exceptions::PyRuntimeError, _>(
"network stream metadata requires a Tokio runtime",
)
})?;
Ok(Self {
inner: Arc::new(Mutex::new(None)),
metadata: Some(metadata),
variant: None,
runtime_handle,
runtime_lease: None,
})
}
pub fn from_upgraded(upgraded: eggfetch_core::network_stream::UpgradedStream) -> Self {
let variant = upgraded.variant().into();
let metadata = upgraded.metadata().clone();
let runtime_handle = tokio::runtime::Handle::current();
Self {
inner: Arc::new(Mutex::new(Some(upgraded))),
metadata: Some(metadata),
variant: Some(variant),
runtime_handle,
runtime_lease: None,
}
}
pub(crate) fn from_upgraded_with_handle(
upgraded: eggfetch_core::network_stream::UpgradedStream,
runtime_handle: tokio::runtime::Handle,
runtime_lease: Option<RuntimeLease>,
) -> Self {
let variant = upgraded.variant().into();
let metadata = upgraded.metadata().clone();
Self {
inner: Arc::new(Mutex::new(Some(upgraded))),
metadata: Some(metadata),
variant: Some(variant),
runtime_handle,
runtime_lease,
}
}
#[allow(dead_code)]
pub(crate) fn from_shared_inner(
inner: SharedStreamInner,
variant: StreamVariant,
metadata: Arc<eggfetch_core::network_stream::ConnectionMetadata>,
runtime_handle: tokio::runtime::Handle,
runtime_lease: Option<RuntimeLease>,
) -> Self {
Self {
inner,
metadata: Some(metadata),
variant: Some(variant),
runtime_handle,
runtime_lease,
}
}
}
#[pymethods]
impl PyNetworkStream {
#[pyo3(signature = (max_bytes=65536, timeout=None))]
fn read<'py>(
&self,
py: Python<'py>,
max_bytes: usize,
timeout: Option<f64>,
) -> PyResult<Bound<'py, PyBytes>> {
let dur = validated_timeout(timeout)?;
let handle = self.runtime_handle.clone();
let result: PyResult<
Result<Result<bytes::Bytes, eggfetch_core::Error>, tokio::time::error::Elapsed>,
> = py.allow_threads(|| {
let mut guard = self.inner.lock().map_err(|e| {
pyo3::exceptions::PyRuntimeError::new_err(format!("lock poisoned: {e}"))
})?;
let inner = guard.as_mut().ok_or_else(|| {
pyo3::exceptions::PyValueError::new_err(
"cannot read from a metadata-only network stream",
)
})?;
if let Some(dur) = dur {
Ok(handle
.block_on(async { tokio::time::timeout(dur, inner.read(max_bytes)).await }))
} else {
Ok(Ok(handle.block_on(inner.read(max_bytes))))
}
});
match result? {
Ok(Ok(data)) => Ok(PyBytes::new(py, &data)),
Ok(Err(e)) => Err(map_err(e)),
Err(_) => Err(pyo3::exceptions::PyTimeoutError::new_err("read timed out")),
}
}
#[pyo3(signature = (data, timeout=None))]
fn write(
&self,
py: Python<'_>,
data: &Bound<'_, PyBytes>,
timeout: Option<f64>,
) -> PyResult<()> {
let dur = validated_timeout(timeout)?;
let buf = data.as_bytes().to_vec();
let handle = self.runtime_handle.clone();
let result: PyResult<
Result<Result<(), eggfetch_core::Error>, tokio::time::error::Elapsed>,
> = py.allow_threads(|| {
let mut guard = self.inner.lock().map_err(|e| {
pyo3::exceptions::PyRuntimeError::new_err(format!("lock poisoned: {e}"))
})?;
let inner = guard.as_mut().ok_or_else(|| {
pyo3::exceptions::PyValueError::new_err(
"cannot write to a metadata-only network stream",
)
})?;
if let Some(dur) = dur {
Ok(handle
.block_on(async { tokio::time::timeout(dur, inner.write_all(&buf)).await }))
} else {
Ok(Ok(handle.block_on(inner.write_all(&buf))))
}
});
match result? {
Ok(Ok(())) => Ok(()),
Ok(Err(e)) => Err(map_err(e)),
Err(_) => Err(pyo3::exceptions::PyTimeoutError::new_err("write timed out")),
}
}
fn close(&self, py: Python<'_>) {
let Ok(mut guard) = self.inner.lock() else {
return;
};
let Some(inner) = guard.as_mut() else {
return;
};
let handle = self.runtime_handle.clone();
let _ = py.allow_threads(|| handle.block_on(inner.close()));
}
fn get_extra_info(&self, key: &str) -> Option<String> {
let meta = self.metadata.as_ref()?;
match key {
"client_addr" => meta.local_addr.map(|a| a.to_string()),
"server_addr" => meta.peer_addr.map(|a| a.to_string()),
"ssl_version" => meta.tls_info.as_ref()?.tls_version.clone(),
"ssl_cipher" => meta.tls_info.as_ref()?.cipher_suite.clone(),
"ssl_alpn" => meta.tls_info.as_ref()?.alpn_protocol.clone(),
"ssl_server_name" => meta.tls_info.as_ref()?.server_name.clone(),
_ => None,
}
}
#[getter]
pub fn is_upgraded(&self) -> bool {
self.inner.lock().is_ok_and(|g| g.is_some())
}
#[pyo3(signature = (ssl_context, server_hostname, *, timeout=None))]
fn start_tls(
&self,
py: Python<'_>,
ssl_context: Option<&Bound<'_, PyAny>>,
server_hostname: &str,
timeout: Option<f64>,
) -> PyResult<PyNetworkStream> {
match self.variant {
Some(StreamVariant::Tcp) => {}
Some(StreamVariant::Tls) => {
return Err(pyo3::exceptions::PyValueError::new_err(
"cannot start TLS on an already-TLS-wrapped stream",
));
}
Some(StreamVariant::Adapter) => {
return Err(pyo3::exceptions::PyValueError::new_err(
"cannot start TLS on a Hyper adapter-backed stream; the underlying \
TCP socket is not recoverable from the upgrade future",
));
}
None => {
return Err(pyo3::exceptions::PyValueError::new_err(
"cannot start TLS on a metadata-only network stream",
));
}
}
let tls_config = match ssl_context {
Some(ctx) if !ctx.is_none() => {
let cfg = crate::tls::ssl_context_to_tls_config(py, Some(ctx))?;
cfg.ok_or_else(|| {
pyo3::exceptions::PyTypeError::new_err(
"eggfetch cannot safely translate this ssl.SSLContext; \
use eggfetch.compat.httpx.create_ssl_context() or pass \
verify/cert kwargs directly",
)
})?
}
_ => eggfetch_core::TlsConfig::builder().build(),
};
let connector = tls_config.tls_connector().map_err(map_err)?;
let mut guard = self.inner.lock().map_err(|e| {
PyErr::new::<pyo3::exceptions::PyRuntimeError, _>(format!("lock poisoned: {e}"))
})?;
let inner = guard.take().ok_or_else(|| {
pyo3::exceptions::PyValueError::new_err(
"cannot start TLS on an already-closed, metadata-only, or handshake-consumed network stream (a prior start_tls consumed the stream, including failed handshakes)",
)
})?;
let server_name_owned = server_hostname.to_owned();
let handle = self.runtime_handle.clone();
let dur = validated_timeout(timeout)?;
let result = py.allow_threads(|| {
let handshake = inner.start_tls(&connector, &server_name_owned);
match dur {
Some(d) => handle.block_on(async { tokio::time::timeout(d, handshake).await }),
None => Ok(handle.block_on(handshake)),
}
});
match result {
Ok(Ok(upgraded)) => Ok(PyNetworkStream::from_upgraded_with_handle(
upgraded, handle, None,
)),
Ok(Err(e)) => Err(crate::errors::map_err(e)),
Err(_) => Err(pyo3::exceptions::PyTimeoutError::new_err(
"TLS handshake timed out",
)),
}
}
fn __repr__(&self) -> String {
let is_upgraded = self.inner.lock().is_ok_and(|g| g.is_some());
if is_upgraded {
"<NetworkStream (upgraded)>".to_string()
} else {
"<NetworkStream (metadata-only)>".to_string()
}
}
}
#[allow(dead_code)]
#[pyfunction]
fn extra_info_dict<'py>(
py: Python<'py>,
metadata: &PyNetworkStream,
) -> PyResult<Bound<'py, PyDict>> {
let dict = PyDict::new(py);
if let Some(ref meta) = metadata.metadata {
if let Some(addr) = meta.local_addr {
dict.set_item("client_addr", addr.to_string())?;
}
if let Some(addr) = meta.peer_addr {
dict.set_item("server_addr", addr.to_string())?;
}
dict.set_item("transport_kind", format!("{:?}", meta.transport_kind))?;
if let Some(ref tls) = meta.tls_info {
let tls_dict = PyDict::new(py);
if let Some(ref proto) = tls.alpn_protocol {
tls_dict.set_item("alpn_protocol", proto)?;
}
if let Some(ref ver) = tls.tls_version {
tls_dict.set_item("tls_version", ver)?;
}
if let Some(ref cipher) = tls.cipher_suite {
tls_dict.set_item("cipher_suite", cipher)?;
}
if let Some(ref name) = tls.server_name {
tls_dict.set_item("server_name", name)?;
}
dict.set_item("tls_info", tls_dict)?;
}
}
Ok(dict)
}
#[derive(Debug, Clone)]
pub(crate) enum EitherNetworkStream {
Sync(PyNetworkStream),
Async(PyAsyncNetworkStream),
}
impl EitherNetworkStream {
pub fn is_upgraded(&self) -> bool {
match self {
Self::Sync(s) => s.is_upgraded(),
Self::Async(s) => s.is_upgraded(),
}
}
pub fn insert_into_dict<'py>(
&self,
py: Python<'py>,
dict: &Bound<'py, PyDict>,
) -> PyResult<()> {
match self {
Self::Sync(s) => {
if s.is_upgraded() {
dict.set_item("network_stream", s.clone())?;
} else {
dict.set_item("network_stream", py.None())?;
}
}
Self::Async(s) => {
if s.is_upgraded() {
dict.set_item("network_stream", s.clone())?;
} else {
dict.set_item("network_stream", py.None())?;
}
}
}
Ok(())
}
}
type AsyncStreamLock =
Arc<tokio::sync::Mutex<Option<eggfetch_core::network_stream::UpgradedStream>>>;
#[pyclass(name = "AsyncNetworkStream")]
#[derive(Clone)]
pub struct PyAsyncNetworkStream {
inner: AsyncStreamLock,
metadata: Option<Arc<eggfetch_core::network_stream::ConnectionMetadata>>,
variant: Option<StreamVariant>,
#[allow(dead_code)]
runtime_handle: tokio::runtime::Handle,
}
impl std::fmt::Debug for PyAsyncNetworkStream {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let is_upgraded = self.inner.try_lock().is_ok_and(|g| g.is_some());
f.debug_struct("PyAsyncNetworkStream")
.field("is_upgraded", &is_upgraded)
.field("metadata", &self.metadata)
.field("variant", &self.variant)
.field("runtime_handle", &"<elided>")
.finish()
}
}
impl PyAsyncNetworkStream {
pub fn from_metadata(
metadata: Arc<eggfetch_core::network_stream::ConnectionMetadata>,
) -> PyResult<Self> {
let runtime_handle = tokio::runtime::Handle::try_current().map_err(|_| {
PyErr::new::<pyo3::exceptions::PyRuntimeError, _>(
"network stream metadata requires a Tokio runtime",
)
})?;
Ok(Self {
inner: Arc::new(tokio::sync::Mutex::new(None)),
metadata: Some(metadata),
variant: None,
runtime_handle,
})
}
pub fn from_upgraded(upgraded: eggfetch_core::network_stream::UpgradedStream) -> Self {
let variant = upgraded.variant().into();
let metadata = upgraded.metadata().clone();
let runtime_handle = tokio::runtime::Handle::current();
Self {
inner: Arc::new(tokio::sync::Mutex::new(Some(upgraded))),
metadata: Some(metadata),
variant: Some(variant),
runtime_handle,
}
}
}
#[pymethods]
impl PyAsyncNetworkStream {
#[pyo3(signature = (max_bytes=65536, timeout=None))]
fn read<'py>(
&self,
py: Python<'py>,
max_bytes: usize,
timeout: Option<f64>,
) -> PyResult<Bound<'py, pyo3::PyAny>> {
let inner = self.inner.clone();
let dur = validated_timeout(timeout)?;
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let mut guard = inner.lock().await;
let stream = guard.as_mut().ok_or_else(|| {
pyo3::exceptions::PyValueError::new_err(
"cannot read from a metadata-only or already-closed network stream",
)
})?;
let result = match dur {
Some(d) => tokio::time::timeout(d, stream.read(max_bytes))
.await
.map_err(|_| eggfetch_core::Error::Timeout {
phase: eggfetch_core::TimeoutPhase::Read,
elapsed: d,
})
.and_then(|r| r),
None => stream.read(max_bytes).await,
};
let data = result.map_err(crate::errors::map_err)?;
Ok::<_, PyErr>(data.to_vec())
})
}
#[pyo3(signature = (data, timeout=None))]
fn write<'py>(
&self,
py: Python<'py>,
data: &Bound<'_, PyBytes>,
timeout: Option<f64>,
) -> PyResult<Bound<'py, pyo3::PyAny>> {
let buf = data.as_bytes().to_vec();
let inner = self.inner.clone();
let dur = validated_timeout(timeout)?;
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let mut guard = inner.lock().await;
let stream = guard.as_mut().ok_or_else(|| {
pyo3::exceptions::PyValueError::new_err(
"cannot write to a metadata-only or already-closed network stream",
)
})?;
let result = match dur {
Some(d) => tokio::time::timeout(d, stream.write_all(&buf))
.await
.map_err(|_| eggfetch_core::Error::Timeout {
phase: eggfetch_core::TimeoutPhase::Write,
elapsed: d,
})
.and_then(|r| r),
None => stream.write_all(&buf).await,
};
result.map_err(crate::errors::map_err)?;
Ok::<_, PyErr>(())
})
}
fn aclose<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, pyo3::PyAny>> {
let inner = self.inner.clone();
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let mut guard = inner.lock().await;
if let Some(mut stream) = guard.take() {
let _ = stream.close().await;
}
Ok::<_, PyErr>(())
})
}
fn get_extra_info(&self, key: &str) -> Option<String> {
let meta = self.metadata.as_ref()?;
match key {
"client_addr" => meta.local_addr.map(|a| a.to_string()),
"server_addr" => meta.peer_addr.map(|a| a.to_string()),
"ssl_version" => meta.tls_info.as_ref()?.tls_version.clone(),
"ssl_cipher" => meta.tls_info.as_ref()?.cipher_suite.clone(),
"ssl_alpn" => meta.tls_info.as_ref()?.alpn_protocol.clone(),
"ssl_server_name" => meta.tls_info.as_ref()?.server_name.clone(),
_ => None,
}
}
#[getter]
pub fn is_upgraded(&self) -> bool {
self.inner.try_lock().is_ok_and(|g| g.is_some())
}
#[pyo3(signature = (ssl_context, server_hostname, *, timeout=None))]
fn start_tls<'py>(
&self,
py: Python<'py>,
ssl_context: Option<&Bound<'_, PyAny>>,
server_hostname: &str,
timeout: Option<f64>,
) -> PyResult<Bound<'py, pyo3::PyAny>> {
match self.variant {
Some(StreamVariant::Tcp) => {}
Some(StreamVariant::Tls) => {
return Err(pyo3::exceptions::PyValueError::new_err(
"cannot start TLS on an already-TLS-wrapped stream",
));
}
Some(StreamVariant::Adapter) => {
return Err(pyo3::exceptions::PyValueError::new_err(
"cannot start TLS on a Hyper adapter-backed stream",
));
}
None => {
return Err(pyo3::exceptions::PyValueError::new_err(
"cannot start TLS on a metadata-only network stream",
));
}
}
let tls_config = match ssl_context {
Some(ctx) if !ctx.is_none() => {
let cfg = crate::tls::ssl_context_to_tls_config(py, Some(ctx))?;
cfg.ok_or_else(|| {
pyo3::exceptions::PyTypeError::new_err(
"eggfetch cannot safely translate this ssl.SSLContext",
)
})?
}
_ => eggfetch_core::TlsConfig::builder().build(),
};
let connector = tls_config.tls_connector().map_err(map_err)?;
let inner = self.inner.clone();
let server_name_owned = server_hostname.to_owned();
let dur = validated_timeout(timeout)?;
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let mut guard = inner.lock().await;
let stream = guard.take().ok_or_else(|| {
pyo3::exceptions::PyValueError::new_err(
"cannot start TLS on an already-closed, metadata-only, or handshake-consumed network stream (a prior start_tls consumed the stream, including failed handshakes)",
)
})?;
let handshake = stream.start_tls(&connector, &server_name_owned);
let result = match dur {
Some(d) => tokio::time::timeout(d, handshake)
.await
.map_err(|_| eggfetch_core::Error::Timeout {
phase: eggfetch_core::TimeoutPhase::Connect,
elapsed: d,
})
.and_then(|r| r),
None => handshake.await,
};
let upgraded = result.map_err(crate::errors::map_err)?;
Python::with_gil(|py| {
let new_stream = PyAsyncNetworkStream::from_upgraded(upgraded);
Bound::new(py, new_stream).map(|b| b.into_any().unbind())
})
})
}
fn __repr__(&self) -> String {
let is_upgraded = self.inner.try_lock().is_ok_and(|g| g.is_some());
if is_upgraded {
"<AsyncNetworkStream (upgraded)>".to_string()
} else {
"<AsyncNetworkStream (metadata-only)>".to_string()
}
}
}