use std::{
fmt,
sync::{Arc, Mutex},
time::Duration,
};
use serde::{Deserialize, Serialize};
use time::OffsetDateTime;
use crate::{
cookies::CookieJar,
errors::CrawlError,
request::UserData,
storage::{KeyValueStore, KeyValueStoreExt},
};
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub struct SessionId(String);
impl SessionId {
pub fn generate() -> Self {
Self(format!("session-{:016x}", crate::util::rand_u64()))
}
pub fn as_str(&self) -> &str {
&self.0
}
}
impl fmt::Display for SessionId {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str(&self.0)
}
}
impl From<String> for SessionId {
fn from(value: String) -> Self {
Self(value)
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub struct SessionToken(String);
impl SessionToken {
pub fn new(value: impl Into<String>) -> Self {
Self(value.into())
}
pub fn as_str(&self) -> &str {
&self.0
}
}
impl fmt::Display for SessionToken {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str(&self.0)
}
}
impl From<String> for SessionToken {
fn from(value: String) -> Self {
Self(value)
}
}
impl From<&str> for SessionToken {
fn from(value: &str) -> Self {
Self(value.to_owned())
}
}
impl From<&SessionId> for SessionToken {
fn from(id: &SessionId) -> Self {
Self(id.as_str().to_owned())
}
}
impl From<SessionId> for SessionToken {
fn from(id: SessionId) -> Self {
Self(id.as_str().to_owned())
}
}
#[cfg(test)]
mod session_token_tests {
use super::{SessionId, SessionToken};
#[test]
fn session_token_from_session_id_preserves_text() {
let id = SessionId::from("session-stable".to_owned());
assert_eq!(SessionToken::from(&id).as_str(), id.as_str());
assert_eq!(SessionToken::from(id).as_str(), "session-stable");
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
#[must_use = "session configuration does nothing unless used to create a session"]
pub struct SessionConfig {
pub max_error_score_scaled: u32,
pub error_score_decrement_scaled: u32,
pub max_usage_count: u32,
pub max_age: Duration,
}
impl Default for SessionConfig {
fn default() -> Self {
Self {
max_error_score_scaled: 3_000,
error_score_decrement_scaled: 500,
max_usage_count: 50,
max_age: Duration::from_secs(3_000),
}
}
}
impl SessionConfig {
pub fn with_max_error_score_scaled(mut self, value: u32) -> Self {
self.max_error_score_scaled = value;
self
}
pub fn with_error_score_decrement_scaled(mut self, value: u32) -> Self {
self.error_score_decrement_scaled = value;
self
}
pub fn with_max_usage_count(mut self, value: u32) -> Self {
self.max_usage_count = value;
self
}
pub fn with_max_age(mut self, value: Duration) -> Self {
self.max_age = value;
self
}
}
struct SessionState {
user_data: UserData,
error_score_scaled: u32,
usage_count: u32,
retired: bool,
}
pub struct Session {
id: SessionId,
cookies: Arc<CookieJar>,
state: tokio::sync::Mutex<SessionState>,
expires_at: OffsetDateTime,
config: SessionConfig,
}
impl Session {
pub fn new(config: SessionConfig) -> Self {
let expires_at = OffsetDateTime::now_utc() + config.max_age;
Self {
id: SessionId::generate(),
cookies: Arc::new(CookieJar::new()),
state: tokio::sync::Mutex::new(SessionState {
user_data: UserData::default(),
error_score_scaled: 0,
usage_count: 0,
retired: false,
}),
expires_at,
config,
}
}
fn restored(value: PersistedSession, config: SessionConfig) -> Result<Self, CrawlError> {
let cookies = CookieJar::from_json(&value.cookies).map_err(CrawlError::non_retryable)?;
Ok(Self {
id: value.id.into(),
cookies: Arc::new(cookies),
state: tokio::sync::Mutex::new(SessionState {
user_data: UserData::default(),
error_score_scaled: value.error_score_scaled,
usage_count: value.usage_count,
retired: value.retired,
}),
expires_at: value.expires_at,
config,
})
}
pub fn id(&self) -> &SessionId {
&self.id
}
pub fn cookie_jar(&self) -> &Arc<CookieJar> {
&self.cookies
}
pub async fn with_user_data<R>(&self, f: impl FnOnce(&UserData) -> R) -> R {
f(&self.state.lock().await.user_data)
}
pub async fn update_user_data(&self, f: impl FnOnce(&mut UserData)) {
f(&mut self.state.lock().await.user_data);
}
pub async fn error_score(&self) -> f32 {
self.state.lock().await.error_score_scaled as f32 / 1_000.0
}
pub async fn usage_count(&self) -> u32 {
self.state.lock().await.usage_count
}
pub async fn record_usage(&self) {
let mut state = self.state.lock().await;
state.usage_count = state.usage_count.saturating_add(1);
}
pub async fn is_blocked(&self) -> bool {
self.state.lock().await.error_score_scaled >= self.config.max_error_score_scaled
}
pub fn is_expired(&self) -> bool {
OffsetDateTime::now_utc() >= self.expires_at
}
pub async fn is_retired(&self) -> bool {
self.state.lock().await.retired
}
pub async fn is_usable(&self) -> bool {
let state = self.state.lock().await;
!state.retired
&& !self.is_expired()
&& state.error_score_scaled < self.config.max_error_score_scaled
&& state.usage_count < self.config.max_usage_count
}
pub async fn mark_good(&self) {
let mut state = self.state.lock().await;
state.error_score_scaled = state
.error_score_scaled
.saturating_sub(self.config.error_score_decrement_scaled);
}
pub async fn mark_bad(&self) {
let mut state = self.state.lock().await;
state.error_score_scaled = state.error_score_scaled.saturating_add(1_000);
}
pub async fn retire(&self) {
self.state.lock().await.retired = true;
}
pub fn set_cookies_from_response(&self, response: &crate::http_client::HttpResponse) {
self.cookies
.store_response_cookies(&response.url, &response.headers);
}
}
impl fmt::Debug for Session {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("Session")
.field("id", &self.id)
.field("expires_at", &self.expires_at)
.finish()
}
}
pub const SESSION_POOL_PERSIST_KEY: &str = "SDK_SESSION_POOL_STATE";
#[derive(Debug, Clone)]
#[non_exhaustive]
#[must_use = "pool options do nothing unless passed to SessionPool::new"]
pub struct SessionPoolOptions {
pub max_pool_size: usize,
pub session_config: SessionConfig,
pub persist_state_key: String,
}
impl Default for SessionPoolOptions {
fn default() -> Self {
Self {
max_pool_size: 1_000,
session_config: SessionConfig::default(),
persist_state_key: SESSION_POOL_PERSIST_KEY.into(),
}
}
}
impl SessionPoolOptions {
pub fn with_max_pool_size(mut self, value: usize) -> Self {
self.max_pool_size = value;
self
}
pub fn with_session_config(mut self, value: SessionConfig) -> Self {
self.session_config = value;
self
}
pub fn with_persist_state_key(mut self, value: impl Into<String>) -> Self {
self.persist_state_key = value.into();
self
}
}
pub struct SessionPool {
sessions: tokio::sync::Mutex<Vec<Arc<Session>>>,
options: SessionPoolOptions,
kvs: Mutex<Option<Arc<dyn KeyValueStore>>>,
}
impl SessionPool {
pub fn new(options: SessionPoolOptions) -> Self {
Self {
sessions: tokio::sync::Mutex::new(Vec::new()),
options,
kvs: Mutex::new(None),
}
}
pub fn attach_persistence(&self, kvs: Arc<dyn KeyValueStore>) {
*self.kvs.lock().unwrap_or_else(|e| e.into_inner()) = Some(kvs);
}
pub async fn session(&self, sticky: Option<&SessionId>) -> Arc<Session> {
let mut sessions = self.sessions.lock().await;
if let Some(session) = sticky
.and_then(|id| sessions.iter().find(|session| session.id() == id))
.cloned()
{
if session.is_usable().await {
session.record_usage().await;
return session;
}
}
let mut usable = Vec::with_capacity(sessions.len());
for session in sessions.iter() {
usable.push(session.is_usable().await);
}
let mut index = 0;
sessions.retain(|_| {
let keep = usable[index];
index += 1;
keep
});
if sessions.len() < self.options.max_pool_size {
let session = Arc::new(Session::new(self.options.session_config.clone()));
session.record_usage().await;
sessions.push(Arc::clone(&session));
return session;
}
if sessions.is_empty() {
let session = Arc::new(Session::new(self.options.session_config.clone()));
session.record_usage().await;
return session;
}
let index = crate::util::rand_u64() as usize % sessions.len();
let session = Arc::clone(&sessions[index]);
session.record_usage().await;
session
}
pub async fn retire_session(&self, id: &SessionId) {
if let Some(session) = self
.sessions
.lock()
.await
.iter()
.find(|s| s.id() == id)
.cloned()
{
session.retire().await;
}
}
pub async fn session_count(&self) -> usize {
self.sessions.lock().await.len()
}
pub async fn persist(&self) -> Result<(), CrawlError> {
let kvs = self.kvs.lock().unwrap_or_else(|e| e.into_inner()).clone();
let Some(kvs) = kvs else {
return Ok(());
};
let sessions = self.sessions.lock().await;
let mut persisted = Vec::with_capacity(sessions.len());
for session in sessions.iter() {
let state = session.state.lock().await;
persisted.push(PersistedSession {
id: session.id.to_string(),
cookies: session
.cookies
.to_json()
.map_err(CrawlError::non_retryable)?,
error_score_scaled: state.error_score_scaled,
usage_count: state.usage_count,
retired: state.retired,
expires_at: session.expires_at,
});
}
drop(sessions);
kvs.set(
&self.options.persist_state_key,
&SessionPoolState {
sessions: persisted,
},
)
.await
.map_err(CrawlError::retry)
}
pub async fn restore(&self) -> Result<(), CrawlError> {
let kvs = self.kvs.lock().unwrap_or_else(|e| e.into_inner()).clone();
let Some(kvs) = kvs else {
return Ok(());
};
let Some(state) = kvs
.get::<SessionPoolState>(&self.options.persist_state_key)
.await
.map_err(CrawlError::retry)?
else {
return Ok(());
};
let mut restored = Vec::with_capacity(state.sessions.len());
for persisted in state.sessions {
match Session::restored(persisted, self.options.session_config.clone()) {
Ok(session) => restored.push(Arc::new(session)),
Err(error) => tracing::warn!(%error, "skipping corrupt persisted session"),
}
}
*self.sessions.lock().await = restored;
Ok(())
}
}
#[derive(Serialize, Deserialize)]
struct SessionPoolState {
sessions: Vec<PersistedSession>,
}
#[derive(Serialize, Deserialize)]
struct PersistedSession {
id: String,
cookies: String,
error_score_scaled: u32,
usage_count: u32,
retired: bool,
#[serde(with = "time::serde::rfc3339")]
expires_at: OffsetDateTime,
}