use std::sync::Arc;
use tokio::time::{Duration, Instant};
use kitsune_p2p_types::{tx2::tx2_utils::ShareOpen, KAgent, KSpace};
use linked_hash_map::{Entry, LinkedHashMap};
use crate::{FetchContext, FetchKey, FetchPoolPush, RoughInt, TransferMethod};
mod pool_reader;
pub use pool_reader::*;
const NUM_ITEMS_PER_POLL: usize = 100;
#[derive(Clone)]
pub struct FetchPool {
config: FetchConfig,
state: ShareOpen<State>,
}
impl std::fmt::Debug for FetchPool {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
self.state
.share_ref(|state| f.debug_struct("FetchPool").field("state", state).finish())
}
}
pub type FetchConfig = Arc<dyn FetchPoolConfig>;
pub trait FetchPoolConfig: 'static + Send + Sync {
fn item_retry_delay(&self) -> Duration {
Duration::from_secs(90)
}
fn source_retry_delay(&self) -> Duration {
Duration::from_secs(5 * 60)
}
fn merge_fetch_contexts(&self, a: u32, b: u32) -> u32;
}
struct FetchPoolConfigBitwiseOr;
impl FetchPoolConfig for FetchPoolConfigBitwiseOr {
fn merge_fetch_contexts(&self, a: u32, b: u32) -> u32 {
a | b
}
}
#[derive(Debug, Default)]
pub struct State {
queue: LinkedHashMap<FetchKey, FetchPoolItem>,
}
impl FetchPool {
pub fn new(config: FetchConfig) -> Self {
Self {
config,
state: ShareOpen::new(State::default()),
}
}
pub fn new_bitwise_or() -> Self {
Self {
config: Arc::new(FetchPoolConfigBitwiseOr),
state: ShareOpen::new(State::default()),
}
}
pub fn push(&self, args: FetchPoolPush) {
self.state.share_mut(|s| {
tracing::debug!(
"FetchPool (size = {}) item added: {:?}",
s.queue.len() + 1,
args
);
s.push(&*self.config, args);
});
}
pub fn remove(&self, key: &FetchKey) -> Option<FetchPoolItem> {
self.state.share_mut(|s| {
let removed = s.remove(key);
tracing::debug!(
"FetchPool (size = {}) item removed: key={:?} val={:?}",
s.queue.len(),
key,
removed
);
removed
})
}
pub fn get_items_to_fetch(&self) -> Vec<(FetchKey, KSpace, FetchSource, Option<FetchContext>)> {
self.state
.share_mut(|s| s.iter_mut(&*self.config).collect())
}
pub fn len(&self) -> usize {
self.state.share_ref(|s| s.queue.len())
}
pub fn is_empty(&self) -> bool {
self.state.share_ref(|s| s.queue.is_empty())
}
}
impl State {
pub fn push(&mut self, config: &dyn FetchPoolConfig, args: FetchPoolPush) {
let FetchPoolPush {
key,
author,
context,
space,
source,
size,
transfer_method,
} = args;
match self.queue.entry(key) {
Entry::Vacant(e) => {
let sources = if let Some(author) = author {
Sources(
[
(source.clone(), SourceRecord::new(source, transfer_method)),
(
FetchSource::Agent(author.clone()),
SourceRecord::agent(author, transfer_method),
),
]
.into_iter()
.collect(),
)
} else {
Sources(
[(source.clone(), SourceRecord::new(source, transfer_method))]
.into_iter()
.collect(),
)
};
let item = FetchPoolItem {
sources,
space,
size,
context,
last_fetch: None,
};
e.insert(item);
}
Entry::Occupied(mut e) => {
let v = e.get_mut();
v.sources
.0
.insert(source.clone(), SourceRecord::new(source, transfer_method));
v.context = match (v.context.take(), context) {
(Some(a), Some(b)) => Some(config.merge_fetch_contexts(*a, *b).into()),
(Some(a), None) => Some(a),
(None, Some(b)) => Some(b),
(None, None) => None,
}
}
}
}
pub fn iter_mut<'a>(&'a mut self, config: &'a dyn FetchPoolConfig) -> StateIter {
StateIter {
state: self,
config,
}
}
pub fn remove(&mut self, key: &FetchKey) -> Option<FetchPoolItem> {
self.queue.remove(key)
}
#[cfg(any(test, feature = "test_utils"))]
pub fn summary(&self) -> String {
use human_repr::HumanCount;
let table = self
.queue
.iter()
.map(|(k, v)| {
let key = match k {
FetchKey::Op(hash) => {
let h = hash.to_string();
format!("{}..{}", &h[0..4], &h[h.len() - 4..])
}
};
let size = v.size.unwrap_or_default().get();
format!(
"{:10} {:^6} {:^6} {:>6}",
key,
v.sources.0.len(),
v.last_fetch
.map(|t| format!("{:?}", t.elapsed()))
.unwrap_or_else(|| "-".to_string()),
size.human_count_bytes(),
)
})
.collect::<Vec<_>>()
.join("\n");
format!("{}\n{} items total", table, self.queue.len())
}
#[cfg(any(test, feature = "test_utils"))]
pub fn summary_heading() -> String {
format!("{:10} {:>6} {:>6} {}", "key", "#src", "last", "size")
}
}
pub struct StateIter<'a> {
state: &'a mut State,
config: &'a dyn FetchPoolConfig,
}
impl<'a> Iterator for StateIter<'a> {
type Item = (FetchKey, KSpace, FetchSource, Option<FetchContext>);
fn next(&mut self) -> Option<Self::Item> {
let keys: Vec<_> = self
.state
.queue
.keys()
.take(NUM_ITEMS_PER_POLL)
.cloned()
.collect();
for key in keys {
let item = self.state.queue.get_refresh(&key)?;
let item_not_recently_fetched = item
.last_fetch
.map(|t| t.elapsed() >= self.config.item_retry_delay())
.unwrap_or(true); if item_not_recently_fetched {
if let Some(source) = item.sources.next(self.config.source_retry_delay()) {
let space = item.space.clone();
item.last_fetch = Some(Instant::now());
return Some((key, space, source, item.context));
}
}
}
None
}
}
#[derive(Debug, PartialEq, Eq)]
pub struct FetchPoolItem {
sources: Sources,
space: KSpace,
size: Option<RoughInt>,
pub context: Option<FetchContext>,
last_fetch: Option<Instant>,
}
#[derive(Debug, PartialEq, Eq)]
struct SourceRecord {
source: FetchSource,
transfer_method: TransferMethod,
last_request: Option<Instant>,
}
impl SourceRecord {
fn new(source: FetchSource, transfer_method: TransferMethod) -> Self {
Self {
source,
transfer_method,
last_request: None,
}
}
fn agent(agent: KAgent, transfer_method: TransferMethod) -> Self {
Self::new(FetchSource::Agent(agent), transfer_method)
}
}
#[derive(Debug, PartialEq, Eq)]
struct Sources(LinkedHashMap<FetchSource, SourceRecord>);
impl Sources {
fn next(&mut self, interval: Duration) -> Option<FetchSource> {
let source_keys: Vec<FetchSource> = self.0.keys().cloned().collect();
for source in source_keys {
if let Some(sr) = self.0.get_refresh(&source) {
if sr
.last_request
.map(|t| t.elapsed() >= interval)
.unwrap_or(true)
{
sr.last_request = Some(Instant::now());
return Some(source);
}
}
}
None
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub enum FetchSource {
Agent(KAgent),
}
#[cfg(test)]
mod tests {
use crate::test_utils::*;
use arbitrary::Arbitrary;
use arbitrary::Unstructured;
use pretty_assertions::assert_eq;
use rand::{Rng, RngCore};
use std::collections::HashSet;
use std::{sync::Arc, time::Duration};
use kitsune_p2p_types::bin_types::{KitsuneBinType, KitsuneSpace};
use super::*;
pub(super) struct Config(pub u32, pub u32);
impl FetchPoolConfig for Config {
fn item_retry_delay(&self) -> Duration {
Duration::from_secs(self.0 as u64)
}
fn source_retry_delay(&self) -> Duration {
Duration::from_secs(self.1 as u64)
}
fn merge_fetch_contexts(&self, a: u32, b: u32) -> u32 {
(a + b).min(1)
}
}
pub(super) fn item(
_cfg: &dyn FetchPoolConfig,
sources: Vec<FetchSource>,
context: Option<FetchContext>,
) -> FetchPoolItem {
FetchPoolItem {
sources: Sources(
sources
.into_iter()
.map(|s| (s.clone(), SourceRecord::new(s, TransferMethod::Gossip)))
.collect(),
),
space: Arc::new(KitsuneSpace::new(vec![0; 36])),
context,
size: None,
last_fetch: None,
}
}
fn arbitrary_test_sources(u: &mut Unstructured, count: usize) -> Vec<FetchSource> {
test_sources(std::iter::repeat_with(|| u8::arbitrary(u).unwrap()).take(count))
}
#[tokio::test(start_paused = true)]
async fn single_source() {
let source_delay = Duration::from_secs(10);
let mut sources = Sources(
[(
test_source(1),
SourceRecord {
source: test_source(1),
last_request: None,
transfer_method: TransferMethod::Gossip,
},
)]
.into_iter()
.collect(),
);
assert_eq!(sources.next(source_delay), Some(test_source(1)));
tokio::time::advance(source_delay).await;
assert_eq!(sources.next(source_delay), Some(test_source(1)));
assert_eq!(sources.next(source_delay), None);
}
#[tokio::test(start_paused = true)]
async fn source_rotation() {
let source_delay = Duration::from_secs(10);
let mut sources = Sources(
[
(
test_source(1),
SourceRecord {
source: test_source(1),
last_request: Some(Instant::now()),
transfer_method: TransferMethod::Gossip,
},
),
(
test_source(2),
SourceRecord {
source: test_source(2),
last_request: None,
transfer_method: TransferMethod::Gossip,
},
),
]
.into_iter()
.collect(),
);
tokio::time::advance(Duration::from_secs(1)).await;
assert_eq!(sources.next(source_delay), Some(test_source(2)));
assert_eq!(sources.next(source_delay), None);
tokio::time::advance(Duration::from_secs(9)).await;
assert_eq!(sources.next(source_delay), Some(test_source(1)));
tokio::time::advance(Duration::from_secs(1)).await;
assert_eq!(sources.next(source_delay), Some(test_source(2)));
assert_eq!(sources.next(source_delay), None);
tokio::time::advance(Duration::from_secs(10)).await;
assert_eq!(sources.next(source_delay), Some(test_source(1)));
assert_eq!(sources.next(source_delay), Some(test_source(2)));
assert_eq!(sources.next(source_delay), None);
}
#[tokio::test(start_paused = true)]
async fn source_rotation_prioritises_less_recently_tried() {
let source_delay = Duration::from_secs(10);
let mut sources = Sources(
[
(
test_source(1),
SourceRecord {
source: test_source(1),
last_request: Some(Instant::now()), transfer_method: TransferMethod::Gossip,
},
),
(
test_source(2),
SourceRecord {
source: test_source(2),
last_request: None, transfer_method: TransferMethod::Gossip,
},
),
(
test_source(3),
SourceRecord {
source: test_source(3),
last_request: None, transfer_method: TransferMethod::Gossip,
},
),
]
.into_iter()
.collect(),
);
assert_eq!(sources.next(source_delay), Some(test_source(2)));
tokio::time::advance(source_delay).await;
assert_eq!(sources.next(source_delay), Some(test_source(3)));
assert_eq!(sources.next(source_delay), Some(test_source(1)));
assert_eq!(sources.next(source_delay), Some(test_source(2)));
assert_eq!(sources.next(source_delay), None);
}
#[tokio::test(start_paused = true)]
async fn source_rotation_uses_all_sources() {
let mut noise = [0; 1_000];
rand::thread_rng().fill_bytes(&mut noise);
let mut u = Unstructured::new(&noise);
let source_delay = Duration::from_secs(rand::thread_rng().gen_range(1..50));
let v = std::iter::repeat_with(|| Duration::arbitrary(&mut u).unwrap())
.take(100)
.enumerate()
.map(|(i, duration)| {
(
test_source(i as u8),
SourceRecord {
source: test_source(i as u8),
last_request: Instant::now().checked_sub(duration),
transfer_method: TransferMethod::Gossip,
},
)
})
.collect();
let mut sources = Sources(v);
let mut seen_sources: HashSet<u8> = HashSet::new();
for _ in 0..100 {
if let Some(s) = sources.next(source_delay) {
match s {
FetchSource::Agent(a) => {
seen_sources.insert(a.0[0]);
}
}
}
tokio::time::advance(source_delay).await;
}
assert_eq!(100, seen_sources.len());
}
#[test]
fn state_keeps_context_on_merge_if_new_is_none() {
let mut q = State::default();
let cfg = Config(1, 1);
q.push(&cfg, test_req_op(1, test_ctx(1), test_source(1)));
assert_eq!(test_ctx(1), q.queue.front().unwrap().1.context);
q.push(&cfg, test_req_op(1, None, test_source(0)));
assert_eq!(test_ctx(1), q.queue.front().unwrap().1.context);
}
#[test]
fn state_adds_context_on_merge_if_current_is_none() {
let mut q = State::default();
let cfg = Config(1, 1);
q.push(&cfg, test_req_op(1, None, test_source(1)));
assert_eq!(None, q.queue.front().unwrap().1.context);
q.push(&cfg, test_req_op(1, test_ctx(1), test_source(0)));
assert_eq!(test_ctx(1), q.queue.front().unwrap().1.context);
}
#[test]
fn state_can_merge_two_items_without_contexts() {
let mut q = State::default();
let cfg = Config(1, 1);
q.push(&cfg, test_req_op(1, None, test_source(1)));
assert_eq!(None, q.queue.front().unwrap().1.context);
q.push(&cfg, test_req_op(1, None, test_source(0)));
assert_eq!(None, q.queue.front().unwrap().1.context);
assert_eq!(2, q.queue.front().unwrap().1.sources.0.len());
}
#[test]
fn state_ignores_duplicate_sources_on_merge() {
let mut q = State::default();
let cfg = Config(1, 1);
q.push(&cfg, test_req_op(1, test_ctx(1), test_source(1)));
assert_eq!(1, q.queue.front().unwrap().1.sources.0.len());
q.push(&cfg, test_req_op(1, test_ctx(2), test_source(1)));
assert_eq!(1, q.queue.front().unwrap().1.sources.0.len());
}
#[test]
fn queue_push() {
let mut q = State::default();
let c = Config(1, 1);
q.push(&c, test_req_op(1, test_ctx(0), test_source(0)));
q.push(&c, test_req_op(1, test_ctx(1), test_source(1)));
q.push(&c, test_req_op(2, test_ctx(0), test_source(0)));
let expected_ready = [
(test_key_op(1), item(&c, test_sources(0..=1), test_ctx(1))),
(test_key_op(2), item(&c, test_sources([0]), test_ctx(0))),
]
.into_iter()
.collect();
assert_eq!(q.queue, expected_ready);
}
#[tokio::test(start_paused = true)]
async fn queue_next() {
let cfg = Config(1, 10);
let mut q = {
let mut queue = [
(test_key_op(1), item(&cfg, test_sources(0..=2), test_ctx(1))),
(test_key_op(2), item(&cfg, test_sources(1..=3), test_ctx(1))),
(test_key_op(3), item(&cfg, test_sources(2..=4), test_ctx(1))),
];
queue[1]
.1
.sources
.0
.get_mut(&test_source(2))
.unwrap()
.last_request = Some(Instant::now() - Duration::from_secs(3));
let queue = queue.into_iter().collect();
State { queue }
};
assert_eq!(q.iter_mut(&cfg).count(), 3);
tokio::time::advance(Duration::from_secs(1)).await;
assert_eq!(q.iter_mut(&cfg).count(), 3);
tokio::time::advance(Duration::from_secs(1)).await;
assert_eq!(q.iter_mut(&cfg).count(), 2);
tokio::time::advance(Duration::from_secs(5)).await;
assert_eq!(
q.iter_mut(&cfg).collect::<Vec<_>>(),
vec![(test_key_op(2), test_space(0), test_source(2), test_ctx(1))]
);
assert_eq!(q.iter_mut(&cfg).count(), 0);
tokio::time::advance(Duration::from_secs(4)).await;
assert_eq!(q.iter_mut(&cfg).count(), 3);
}
#[tokio::test(start_paused = true)]
async fn state_iter_sees_all_items() {
let cfg = Config(1, 10);
let num_items = 2 * NUM_ITEMS_PER_POLL;
let mut q = {
let mut queue = vec![];
for i in 0..(num_items) {
queue.push((
test_key_op(i as u8),
item(&cfg, test_sources([(i % 100) as u8]), test_ctx(1)),
))
}
State {
queue: queue.into_iter().collect(),
}
};
assert_eq!(num_items, q.iter_mut(&cfg).count());
assert_eq!(0, q.iter_mut(&cfg).count());
tokio::time::advance(Duration::from_secs(30)).await;
assert_eq!(num_items, q.iter_mut(&cfg).count());
}
#[tokio::test(start_paused = true)]
async fn state_iter_uses_all_sources() {
let cfg = Config(1, 10);
let num_items = 10;
let mut q = {
let mut queue = vec![];
for i in 0..num_items {
queue.push((
test_key_op(i as u8),
item(
&cfg,
test_sources((i * num_items) as u8..(i * num_items + num_items) as u8),
test_ctx(1),
),
))
}
State {
queue: queue.into_iter().collect(),
}
};
let mut seen_sources = HashSet::new();
for _ in 0..num_items {
q.iter_mut(&cfg)
.map(|item| match item.2 {
FetchSource::Agent(a) => a.0.clone(),
})
.for_each(|source| {
seen_sources.insert(source);
});
tokio::time::advance(Duration::from_secs(30)).await;
}
assert_eq!(num_items * num_items, seen_sources.len());
}
#[tokio::test(start_paused = true)]
async fn remove_fetch_item() {
let cfg = Config(1, 10);
let mut q = {
let queue = [(test_key_op(1), item(&cfg, test_sources([1]), test_ctx(1)))];
let queue = queue.into_iter().collect();
State { queue }
};
assert_eq!(1, q.iter_mut(&cfg).count());
q.remove(&test_key_op(1));
tokio::time::advance(Duration::from_secs(30)).await;
assert_eq!(0, q.iter_mut(&cfg).count());
}
#[tokio::test(start_paused = true)]
async fn fetch_pool() {
struct TestFetchConfig {}
impl FetchPoolConfig for TestFetchConfig {
fn merge_fetch_contexts(&self, a: u32, b: u32) -> u32 {
a | b
}
}
let mut noise = [0; 1_000];
rand::thread_rng().fill_bytes(&mut noise);
let mut u = Unstructured::new(&noise);
let fetch_pool = FetchPool::new(Arc::new(TestFetchConfig {}));
let unavailable_sources: HashSet<FetchSource> =
arbitrary_test_sources(&mut u, 10).into_iter().collect();
fetch_pool.push(FetchPoolPush {
key: test_key_op(220),
space: test_space(u8::arbitrary(&mut u).unwrap()),
source: unavailable_sources.iter().last().cloned().unwrap(),
size: None, author: None, context: test_ctx(u32::arbitrary(&mut u).unwrap()),
transfer_method: TransferMethod::Gossip,
});
let mut failed_count = 0;
for i in (0..200).step_by(5) {
for _ in 0..5 {
fetch_pool.push(FetchPoolPush {
key: test_key_op(i),
space: test_space(u8::arbitrary(&mut u).unwrap()),
source: test_source(u8::arbitrary(&mut u).unwrap()),
size: None, author: None, context: test_ctx(u32::arbitrary(&mut u).unwrap()),
transfer_method: TransferMethod::Gossip,
});
}
let items = fetch_pool.get_items_to_fetch();
for item in items {
if !unavailable_sources.contains(&item.2) {
fetch_pool.remove(&item.0);
} else {
failed_count += 1;
}
}
tokio::time::advance(fetch_pool.config.item_retry_delay()).await;
}
assert!(
!fetch_pool.get_items_to_fetch().is_empty(),
"Pool should have had at least one item but got \n {}",
fetch_pool.state.share_ref(|s| format!(
"{}\n{}",
State::summary_heading(),
s.summary()
))
);
assert!(
failed_count >= 10,
"At least 10 items should have failed to be fetched but was {}",
failed_count
);
}
#[test]
fn default_fetch_context_merge_maintains_flags_from_both_contexts() {
const FLAG_1: u32 = 1 << 5;
const FLAG_2: u32 = 1 << 10;
let context_1 = FetchContext(FLAG_1);
let context_2 = FetchContext(FLAG_2);
let pool = FetchPool::new_bitwise_or();
let merged = pool.config.merge_fetch_contexts(*context_1, *context_2);
assert_eq!(FLAG_1, merged & FLAG_1);
assert_eq!(FLAG_2, merged & FLAG_2);
assert_eq!(0, merged ^ (FLAG_1 | FLAG_2)); }
}