use std::collections::HashMap;
use std::sync::Arc;
use base64::Engine;
use futures::StreamExt;
use futures::SinkExt;
use futures::channel::mpsc;
use serde::Deserialize;
use tokio::sync::Mutex;
use chromiumoxide::error::CdpError;
use chromiumoxide::page::Page;
use chromiumoxide::cdp::browser_protocol::network::{
EnableParams, EventResponseReceived, EventLoadingFinished, GetResponseBodyParams,
};
#[derive(thiserror::Error, Debug)]
pub enum Error {
#[error("enable_network: {0}")]
EnableNetwork(CdpError),
#[error("get_response_body: {0}")]
GetResponseBody(CdpError),
#[error("base64_decode: {0}")]
Base64Decode(base64::DecodeError),
}
#[derive(Clone, Debug)]
pub struct EventStreamConfig {
pub url_substring_filter: Option<String>,
pub content_type_substring_filter: Option<String>,
}
impl Default for EventStreamConfig {
fn default() -> Self {
Self {
url_substring_filter: None,
content_type_substring_filter: None,
}
}
}
#[derive(Clone, Debug, Deserialize)]
pub struct Event {
pub url: String,
#[serde(rename = "contentType", default)]
pub content_type: Option<String>,
#[serde(default)]
pub status: Option<u16>,
pub body: String,
}
#[derive(Clone, Debug)]
struct PendingResponse {
url: String,
content_type: Option<String>,
status: Option<u16>,
}
fn should_capture(
config: &EventStreamConfig,
url: &str,
content_type: Option<&str>,
) -> bool {
let url_ok = config
.url_substring_filter
.as_ref()
.map(|filter| url.contains(filter))
.unwrap_or(true);
let ct_ok = config
.content_type_substring_filter
.as_ref()
.map(|filter| {
content_type
.map(|ct| ct.contains(filter))
.unwrap_or(false)
})
.unwrap_or(true);
url_ok && ct_ok
}
pub async fn start_event_stream(
page: Page,
config: EventStreamConfig,
) -> Result<mpsc::UnboundedReceiver<Event>, Error> {
page.execute(EnableParams::default())
.await
.map_err(Error::EnableNetwork)?;
let (mut tx, rx) = mpsc::unbounded();
let pending_responses: Arc<Mutex<HashMap<String, PendingResponse>>> =
Arc::new(Mutex::new(HashMap::new()));
let pending_clone = pending_responses.clone();
let page_response = page.clone();
tokio::spawn(async move {
let mut events = match page_response.event_listener::<EventResponseReceived>().await {
Ok(e) => e,
Err(_) => return, };
while let Some(event) = events.next().await {
let url = event.response.url.clone();
let status = Some(event.response.status as u16);
let headers = &event.response.headers;
let headers_value = headers.inner();
let content_type = headers_value
.get("content-type")
.or_else(|| headers_value.get("Content-Type"))
.and_then(|v| v.as_str())
.map(|s| s.to_string());
if should_capture(&config, &url, content_type.as_deref()) {
let pending = PendingResponse {
url,
content_type,
status,
};
pending_clone
.lock()
.await
.insert(event.request_id.inner().clone(), pending);
}
}
});
tokio::spawn(async move {
let mut events = match page.event_listener::<EventLoadingFinished>().await {
Ok(e) => e,
Err(_) => return, };
while let Some(event) = events.next().await {
let request_id_str = event.request_id.inner().clone();
let pending = pending_responses.lock().await.remove(&request_id_str);
let pending = match pending {
Some(p) => p,
None => continue, };
let body_result = page
.execute(GetResponseBodyParams {
request_id: event.request_id.clone(),
})
.await;
let body = match body_result {
Ok(result) => {
if result.base64_encoded {
match base64::engine::general_purpose::STANDARD.decode(&result.body) {
Ok(bytes) => String::from_utf8_lossy(&bytes).to_string(),
Err(e) => {
eprintln!("Failed to decode base64 body: {}", e);
continue;
}
}
} else {
result.body.clone()
}
}
Err(_) => {
continue;
}
};
let event = Event {
url: pending.url,
content_type: pending.content_type,
status: pending.status,
body,
};
if tx.send(event).await.is_err() {
return; }
}
});
Ok(rx)
}