use serde::{Deserialize, Serialize};
use std::fmt;
use url::Url;
pub use http::{HeaderMap, Method};
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Request {
pub id: RequestId,
pub url: Url,
pub loaded_url: Option<Url>,
pub unique_key: String,
#[serde(with = "method_serde")]
pub method: Method,
#[serde(with = "header_map_serde")]
pub headers: HeaderMap,
pub body: Option<RequestBody>,
pub user_data: UserData,
pub label: Option<String>,
pub retry_count: u32,
pub session_rotation_count: u32,
pub max_retries: Option<u32>,
pub no_retry: bool,
pub error_messages: Vec<String>,
#[serde(with = "time::serde::rfc3339::option")]
pub handled_at: Option<time::OffsetDateTime>,
pub state: RequestState,
pub crawl_depth: u32,
pub skip_navigation: bool,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
pub enum RequestState {
#[default]
Unprocessed,
BeforeNav,
AfterNav,
RequestHandler,
Done,
ErrorHandler,
Error,
Skipped,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub struct RequestId(String);
impl RequestId {
pub fn from_unique_key(unique_key: &str) -> Self {
let h1 = fnv1a64(unique_key.as_bytes());
let h2 = fnv1a64_seeded(unique_key.as_bytes(), h1);
Self(format!("{h1:016x}{h2:016x}"))
}
pub fn as_str(&self) -> &str {
&self.0
}
}
impl fmt::Display for RequestId {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str(&self.0)
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum RequestBody {
Bytes(Vec<u8>),
Form(Vec<(String, String)>),
Json(serde_json::Value),
}
impl RequestBody {
pub fn canonical_bytes(&self) -> Vec<u8> {
match self {
Self::Bytes(bytes) => bytes.clone(),
Self::Form(pairs) => pairs
.iter()
.map(|(key, value)| format!("{key}={value}"))
.collect::<Vec<_>>()
.join("&")
.into_bytes(),
Self::Json(value) => serde_json::to_vec(value).expect(
"a serde_json::Value is always serializable",
),
}
}
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub struct UserData(pub serde_json::Map<String, serde_json::Value>);
impl UserData {
pub fn get_typed<T: serde::de::DeserializeOwned>(
&self,
key: &str,
) -> Option<Result<T, serde_json::Error>> {
self.0.get(key).cloned().map(serde_json::from_value)
}
pub fn set_typed<T: Serialize>(
&mut self,
key: &str,
value: &T,
) -> Result<(), serde_json::Error> {
self.0.insert(key.to_owned(), serde_json::to_value(value)?);
Ok(())
}
pub fn is_empty(&self) -> bool {
self.0.is_empty()
}
pub fn get(&self, key: &str) -> Option<&serde_json::Value> {
self.0.get(key)
}
}
impl Request {
pub fn get(url: impl IntoUrl) -> RequestBuilder {
Self::builder().url(url).method(Method::GET)
}
pub fn post(url: impl IntoUrl) -> RequestBuilder {
Self::builder().url(url).method(Method::POST)
}
pub fn builder() -> RequestBuilder {
RequestBuilder::default()
}
pub fn compute_unique_key(url: &Url, method: &Method, body: Option<&RequestBody>) -> String {
let mut normalized = url.clone();
normalized.set_fragment(None);
if *method == Method::GET && body.is_none() {
normalized.into()
} else {
let body_bytes = body.map(RequestBody::canonical_bytes).unwrap_or_default();
format!(
"{}({:016x}):{}",
method.as_str(),
fnv1a64(&body_bytes),
normalized
)
}
}
}
pub trait IntoUrl {
fn into_url(self) -> Result<Url, url::ParseError>;
}
impl IntoUrl for Url {
fn into_url(self) -> Result<Url, url::ParseError> {
Ok(self)
}
}
impl IntoUrl for &str {
fn into_url(self) -> Result<Url, url::ParseError> {
Url::parse(self)
}
}
impl IntoUrl for String {
fn into_url(self) -> Result<Url, url::ParseError> {
Url::parse(&self)
}
}
impl IntoUrl for &String {
fn into_url(self) -> Result<Url, url::ParseError> {
Url::parse(self)
}
}
#[derive(Debug)]
struct PendingHeader {
name: String,
value: String,
}
#[derive(Debug, Default)]
#[must_use = "builders do nothing unless consumed by build"]
pub struct RequestBuilder {
url: Option<Result<Url, url::ParseError>>,
method: Option<Method>,
headers: HeaderMap,
pending_headers: Vec<PendingHeader>,
body: Option<RequestBody>,
serialization_error: Option<serde_json::Error>,
user_data: UserData,
label: Option<String>,
max_retries: Option<u32>,
no_retry: bool,
skip_navigation: bool,
unique_key: Option<String>,
pub(crate) forefront: bool,
crawl_depth: u32,
}
impl RequestBuilder {
pub async fn enqueue(
self,
queue: &dyn crate::storage::RequestQueue,
) -> Result<crate::storage::QueueOpInfo, crate::errors::CrawlError> {
let forefront = self.is_forefront();
let request = self.build()?;
Ok(queue
.add(
request,
crate::storage::AddOptions {
forefront,
..Default::default()
},
)
.await?)
}
pub fn url(mut self, url: impl IntoUrl) -> Self {
self.url = Some(url.into_url());
self
}
pub fn method(mut self, method: Method) -> Self {
self.method = Some(method);
self
}
pub fn header(mut self, name: &str, value: &str) -> Self {
self.pending_headers.push(PendingHeader {
name: name.to_owned(),
value: value.to_owned(),
});
self
}
pub fn headers(mut self, headers: HeaderMap) -> Self {
self.headers = headers;
self.pending_headers.clear();
self
}
pub fn body(mut self, body: RequestBody) -> Self {
self.body = Some(body);
self
}
pub fn json<T: Serialize>(mut self, value: &T) -> Self {
match serde_json::to_value(value) {
Ok(value) => self.body = Some(RequestBody::Json(value)),
Err(error) => self.serialization_error = Some(error),
}
self
}
pub fn form(mut self, pairs: impl IntoIterator<Item = (String, String)>) -> Self {
self.body = Some(RequestBody::Form(pairs.into_iter().collect()));
self
}
pub fn label(mut self, label: impl Into<String>) -> Self {
self.label = Some(label.into());
self
}
pub fn user_data(mut self, user_data: UserData) -> Self {
self.user_data = user_data;
self
}
pub fn user_data_entry<T: Serialize>(mut self, key: impl Into<String>, value: &T) -> Self {
match serde_json::to_value(value) {
Ok(value) => {
self.user_data.0.insert(key.into(), value);
}
Err(error) => self.serialization_error = Some(error),
}
self
}
pub fn max_retries(mut self, max_retries: u32) -> Self {
self.max_retries = Some(max_retries);
self
}
pub fn no_retry(mut self, no_retry: bool) -> Self {
self.no_retry = no_retry;
self
}
pub fn skip_navigation(mut self, skip_navigation: bool) -> Self {
self.skip_navigation = skip_navigation;
self
}
pub fn unique_key(mut self, unique_key: impl Into<String>) -> Self {
self.unique_key = Some(unique_key.into());
self
}
pub fn forefront(mut self, forefront: bool) -> Self {
self.forefront = forefront;
self
}
pub fn is_forefront(&self) -> bool {
self.forefront
}
pub fn crawl_depth(mut self, crawl_depth: u32) -> Self {
self.crawl_depth = crawl_depth;
self
}
pub fn build(mut self) -> Result<Request, RequestBuildError> {
if let Some(error) = self.serialization_error {
return Err(error.into());
}
let url = self.url.ok_or(RequestBuildError::MissingUrl)??;
for pending in self.pending_headers {
let name =
http::header::HeaderName::from_bytes(pending.name.as_bytes()).map_err(|error| {
RequestBuildError::InvalidHeader {
name: pending.name.clone(),
message: error.to_string(),
}
})?;
let value =
http::HeaderValue::from_bytes(pending.value.as_bytes()).map_err(|error| {
RequestBuildError::InvalidHeader {
name: pending.name.clone(),
message: error.to_string(),
}
})?;
self.headers.append(name, value);
}
let method = self.method.unwrap_or(Method::GET);
let unique_key = self
.unique_key
.unwrap_or_else(|| Request::compute_unique_key(&url, &method, self.body.as_ref()));
let id = RequestId::from_unique_key(&unique_key);
Ok(Request {
id,
url,
loaded_url: None,
unique_key,
method,
headers: self.headers,
body: self.body,
user_data: self.user_data,
label: self.label,
retry_count: 0,
session_rotation_count: 0,
max_retries: self.max_retries,
no_retry: self.no_retry,
error_messages: Vec::new(),
handled_at: None,
state: RequestState::Unprocessed,
crawl_depth: self.crawl_depth,
skip_navigation: self.skip_navigation,
})
}
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum RequestBuildError {
#[error("request URL is missing")]
MissingUrl,
#[error("invalid URL: {0}")]
InvalidUrl(#[from] url::ParseError),
#[error("invalid header {name}: {message}")]
InvalidHeader {
name: String,
message: String,
},
#[error("user data serialization failed: {0}")]
Serialization(#[from] serde_json::Error),
}
fn fnv1a64(bytes: &[u8]) -> u64 {
fnv1a64_seeded(bytes, 0xcbf29ce484222325)
}
fn fnv1a64_seeded(bytes: &[u8], seed: u64) -> u64 {
bytes.iter().fold(seed, |hash, byte| {
(hash ^ u64::from(*byte)).wrapping_mul(0x100000001b3)
})
}
mod method_serde {
use http::Method;
use serde::{Deserialize, Deserializer, Serializer, de::Error as _};
pub fn serialize<S: Serializer>(method: &Method, serializer: S) -> Result<S::Ok, S::Error> {
serializer.serialize_str(method.as_str())
}
pub fn deserialize<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Method, D::Error> {
let value = String::deserialize(deserializer)?;
Method::from_bytes(value.as_bytes()).map_err(D::Error::custom)
}
}
mod header_map_serde {
use http::{HeaderMap, HeaderName, HeaderValue};
use serde::{
Deserialize, Deserializer, Serialize, Serializer, de::Error as _, ser::Error as _,
};
pub fn serialize<S: Serializer>(headers: &HeaderMap, serializer: S) -> Result<S::Ok, S::Error> {
headers
.iter()
.map(|(name, value)| {
Ok((
name.as_str().to_owned(),
value.to_str().map_err(S::Error::custom)?.to_owned(),
))
})
.collect::<Result<Vec<(String, String)>, S::Error>>()?
.serialize(serializer)
}
pub fn deserialize<'de, D: Deserializer<'de>>(deserializer: D) -> Result<HeaderMap, D::Error> {
let pairs = Vec::<(String, String)>::deserialize(deserializer)?;
let mut headers = HeaderMap::new();
for (name, value) in pairs {
let name = HeaderName::from_bytes(name.as_bytes()).map_err(D::Error::custom)?;
let value = HeaderValue::from_bytes(value.as_bytes()).map_err(D::Error::custom)?;
headers.append(name, value);
}
Ok(headers)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn builder_happy_path() {
let request = Request::get("https://example.com/a?x=1")
.label("l")
.build()
.unwrap();
assert_eq!(request.label.as_deref(), Some("l"));
assert_eq!(request.method, Method::GET);
}
#[test]
fn unique_key_override_is_respected() {
let request = Request::get("https://example.com")
.unique_key("mine")
.build()
.unwrap();
assert_eq!(request.unique_key, "mine");
}
#[test]
fn get_key_is_fragmentless_url() {
let request = Request::get("https://example.com/a#frag").build().unwrap();
assert_eq!(request.unique_key, "https://example.com/a");
}
#[test]
fn fragment_does_not_affect_key() {
let first = Request::get("https://example.com/a#frag").build().unwrap();
let second = Request::get("https://example.com/a").build().unwrap();
assert_eq!(first.unique_key, second.unique_key);
}
#[test]
fn post_body_affects_key_deterministically() {
let build = |bytes| {
Request::post("https://example.com/a")
.body(RequestBody::Bytes(bytes))
.build()
.unwrap()
};
assert_eq!(build(vec![1]).unique_key, build(vec![1]).unique_key);
assert_ne!(build(vec![1]).unique_key, build(vec![2]).unique_key);
}
#[test]
fn request_id_is_deterministic() {
assert_eq!(
RequestId::from_unique_key("key"),
RequestId::from_unique_key("key")
);
}
#[test]
fn user_data_typed_roundtrip() {
#[derive(Debug, PartialEq, Serialize, Deserialize)]
struct Example {
number: u32,
}
let mut data = UserData::default();
data.set_typed("example", &Example { number: 7 }).unwrap();
assert_eq!(
data.get_typed::<Example>("example").unwrap().unwrap(),
Example { number: 7 }
);
}
#[test]
fn invalid_header_is_reported_at_build() {
let result = Request::get("https://example.com")
.header("bad header", "value")
.build();
assert!(matches!(
result,
Err(RequestBuildError::InvalidHeader { .. })
));
}
}