use std::sync::Arc;
use bytes::Bytes;
use pyo3::exceptions::PyTypeError;
use pyo3::prelude::*;
use pyo3::types::PyDict;
use eggfetch_core::trace::TraceObserver;
use eggfetch_core::TransportHints;
use crate::trace_bridge::{CallbackErrorSlot, PyTraceObserver};
pub(crate) struct ExtractedExtensions {
pub hints: TransportHints,
pub trace_error_slot: Option<CallbackErrorSlot>,
}
pub(crate) fn extract_native_extensions(
py: Python<'_>,
extensions: Option<&Bound<'_, PyAny>>,
) -> PyResult<ExtractedExtensions> {
let Some(ext) = extensions else {
return Ok(ExtractedExtensions {
hints: TransportHints::default(),
trace_error_slot: None,
});
};
if ext.is_none() {
return Ok(ExtractedExtensions {
hints: TransportHints::default(),
trace_error_slot: None,
});
}
let dict = ext.downcast::<PyDict>().map_err(|_| {
PyTypeError::new_err(
"extensions must be a dict containing supported keys: target, sni_hostname, trace",
)
})?;
let mut hints = TransportHints::default();
let mut trace_error_slot: Option<CallbackErrorSlot> = None;
for (key_obj, value_obj) in dict.iter() {
let key: String = key_obj.extract()?;
match key.as_str() {
"target" => {
if let Some(existing) = hints.target.as_ref() {
return Err(PyTypeError::new_err(format!(
"target extension already supplied ({existing:?}); multiple target values are not supported"
)));
}
let bytes = if let Ok(b) = value_obj.extract::<Vec<u8>>() {
Bytes::from(b)
} else if let Ok(s) = value_obj.extract::<String>() {
Bytes::from(s.into_bytes())
} else {
return Err(PyTypeError::new_err(
"target extension must be a str or bytes value",
));
};
if bytes.is_empty() {
return Err(PyTypeError::new_err("target extension must not be empty"));
}
if bytes.iter().any(|&b| b < 0x20 || b == 0x7f) {
return Err(PyTypeError::new_err(
"target extension contains forbidden characters (C0 controls/DEL; includes CR/LF/NUL)",
));
}
hints.target = Some(bytes);
}
"sni_hostname" => {
if let Some(existing) = hints.sni_hostname.as_ref() {
return Err(PyTypeError::new_err(format!(
"sni_hostname extension already supplied ({existing:?}); multiple values are not supported"
)));
}
let s: String = value_obj.extract()?;
if s.is_empty() {
return Err(PyTypeError::new_err(
"sni_hostname extension must not be empty",
));
}
hints.sni_hostname = Some(s);
}
"trace" => {
if trace_error_slot.is_some() {
return Err(PyTypeError::new_err(
"trace extension already supplied; only one trace callable is supported",
));
}
if value_obj.is_none() {
continue;
}
let (observer, slot) = PyTraceObserver::new(py, value_obj.clone());
let arc: Arc<dyn TraceObserver> = Arc::new(observer);
hints.trace = Some(arc);
trace_error_slot = Some(slot);
}
_ => {
}
}
}
Ok(ExtractedExtensions {
hints,
trace_error_slot,
})
}