use std::collections::HashMap;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex, MutexGuard, Weak};
use std::time::Duration;
use tokio::sync::watch;
use tokio::time::Instant;
use crate::cbor::Value;
use crate::frame::StreamMode;
use crate::manifest::{self, block_mcid, chunk_mcid, Manifest, Mcid, DEFAULT_CHUNK_SIZE};
use crate::record::{
self, new_content_announcement, ContentAnnouncementOptions, Reason, Record, RecordType,
TombstoneOptions,
};
use crate::station_link::{
self, stream_handler, Link, LinkError, Stream, StreamEvent, DEFAULT_CALL_TIMEOUT,
};
use super::{shuffle, Offer, Pool, PoolError, PoolInner, Served};
pub const CONTENT_PROCEDURE: &str = "content_v1";
const ANNOUNCEMENT_TTL: Duration = Duration::from_secs(60 * 60);
const MAX_BLOCK_BYTES: u64 = DEFAULT_CHUNK_SIZE;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ContentOptions {
pub max_bytes: u64,
pub max_chunks: u64,
pub parallel: usize,
pub chunk_timeout: Duration,
}
impl Default for ContentOptions {
fn default() -> Self {
ContentOptions {
max_bytes: 256 << 20,
max_chunks: 16_384,
parallel: 4,
chunk_timeout: Duration::from_secs(15),
}
}
}
#[derive(Default)]
pub(super) struct Sharer {
realms: Mutex<HashMap<[u8; 32], SharedRealm>>,
serving: tokio::sync::Mutex<()>,
}
struct SharedRealm {
_served: Served,
roots: HashMap<Mcid, Root>,
chunks: HashMap<Mcid, Vec<u8>>,
announcements: HashMap<Mcid, Announced>,
}
#[derive(Clone)]
enum Root {
Block(Vec<u8>),
Chunked(Arc<Manifest>),
}
struct Announced {
latest: Arc<Mutex<Record>>,
stop: watch::Sender<bool>,
}
impl Sharer {
fn lock(&self) -> MutexGuard<'_, HashMap<[u8; 32], SharedRealm>> {
self.realms.lock().unwrap_or_else(|p| p.into_inner())
}
}
pub fn content_procedure_bound(procedure: &str, node: &[u8; 32]) -> bool {
let node_hex: String = node.iter().map(|b| format!("{b:02x}")).collect();
if procedure == record::own_procedure(node, CONTENT_PROCEDURE) {
return true;
}
match procedure.split_once('/') {
Some((org, name)) => {
!org.is_empty()
&& org != "_"
&& !org.starts_with(record::OWN_NAMESPACE_PREFIX)
&& name == format!("{CONTENT_PROCEDURE}_{node_hex}")
}
None => false,
}
}
impl Pool {
pub async fn share_content(
&self,
realm: &[u8; 32],
data: &[u8],
name: &str,
) -> Result<Mcid, PoolError> {
let inner = &self.inner;
let (mcid, root, chunks) = if data.len() as u64 <= MAX_BLOCK_BYTES {
(block_mcid(data), Root::Block(data.to_vec()), HashMap::new())
} else {
let (m, parts) = manifest::create(data, name, DEFAULT_CHUNK_SIZE)
.map_err(|e| PoolError::ContentMismatch(e.to_string()))?;
let chunks: HashMap<Mcid, Vec<u8>> = parts
.into_iter()
.enumerate()
.filter_map(|(i, part)| chunk_mcid(&m, i).map(|c| (c, part)))
.collect();
(m.mcid, Root::Chunked(Arc::new(m)), chunks)
};
inner.shared_realm(realm).await?;
let already = {
let mut realms = inner.content.lock();
let shared = realms.get_mut(realm).ok_or(PoolError::Closed)?;
shared.roots.insert(mcid, root.clone());
shared.chunks.extend(chunks);
shared.announcements.contains_key(&mcid)
};
if already {
return Ok(mcid);
}
let latest = match inner.announce_content(realm, &mcid, &root).await {
Ok(latest) => latest,
Err(e) => {
inner.forget_shared(realm, &mcid);
return Err(e);
}
};
let latest = Arc::new(Mutex::new(latest));
let (stop, stopped) = watch::channel(false);
if let Some(shared) = inner.content.lock().get_mut(realm) {
shared.announcements.insert(
mcid,
Announced {
latest: latest.clone(),
stop,
},
);
}
tokio::spawn(renew_announcement(
Arc::downgrade(inner),
*realm,
mcid,
root,
latest,
stopped,
));
Ok(mcid)
}
pub async fn unshare_content(&self, realm: &[u8; 32], mcid: &Mcid) -> Result<(), PoolError> {
let Some(announced) = self.inner.forget_shared(realm, mcid) else {
return Ok(());
};
announced.stop.send_replace(true);
let latest = announced
.latest
.lock()
.unwrap_or_else(|p| p.into_inner())
.clone();
let tombstone =
record::new_tombstone(&latest, Reason::Shutdown, &TombstoneOptions::default())
.map_err(LinkError::from)?;
let signed =
record::sign(&tombstone, &self.inner.opts.identity).map_err(LinkError::from)?;
let wire = record::encode(&signed).map_err(LinkError::from)?;
self.put_record(&wire).await
}
pub async fn get_content(
&self,
realm: &[u8; 32],
mcid: &Mcid,
opts: ContentOptions,
) -> Result<Vec<u8>, PoolError> {
let sharers = self.inner.content_sharers(realm, mcid).await?;
if sharers.is_empty() {
return Err(PoolError::NotShared);
}
let mut tried = Vec::new();
for sharer in sharers {
match self.inner.fetch_from(realm, &sharer, mcid, &opts).await {
Ok(data) => return Ok(data),
Err(e) => tried.push((sharer.node, e)),
}
}
Err(PoolError::ContentUnavailable(tried))
}
}
#[derive(Debug, Clone)]
struct Sharing {
node: [u8; 32],
station: [u8; 32],
procedure: String,
}
impl PoolInner {
async fn shared_realm(self: &Arc<Self>, realm: &[u8; 32]) -> Result<(), PoolError> {
let _serving = self.content.serving.lock().await;
if self.content.lock().contains_key(realm) {
return Ok(());
}
let pool = Arc::downgrade(self);
let r = *realm;
let answer = stream_handler(move |s| {
let pool = pool.clone();
async move {
match pool.upgrade() {
Some(inner) => inner.answer_fetch(&r, &s).await,
None => {
s.abort("not_shared", "the node no longer shares content")
.await
}
}
.map_err(|e| e.to_string())
}
});
let procedure = record::own_procedure(&self.self_id, CONTENT_PROCEDURE);
let served = Pool {
inner: self.clone(),
}
.serve(Offer::stream(
*realm,
&procedure,
StreamMode::ServerStream,
answer,
))
.await?;
self.content.lock().insert(
*realm,
SharedRealm {
_served: served,
roots: HashMap::new(),
chunks: HashMap::new(),
announcements: HashMap::new(),
},
);
Ok(())
}
async fn answer_fetch(&self, realm: &[u8; 32], s: &Stream) -> Result<(), LinkError> {
let args = &s.request().payload;
let mcid: Option<Mcid> = match wire_bytes(args, "mcid") {
Some(b) if b.len() == 50 && b[0] == 2 => b.as_slice().try_into().ok(),
_ => None,
};
let want = wire_text(args, "want");
let (Some(mcid), Some(want @ ("root" | "block"))) = (mcid, want.as_deref()) else {
return s
.abort(
"malformed",
"a fetch names one content id and wants root or block",
)
.await;
};
let body = self
.content
.lock()
.get(realm)
.and_then(|shared| shared.body(want, &mcid));
match body {
None => {
s.abort("not_shared", "this node does not share that content")
.await
}
Some(body) => {
s.send_value(body).await?;
s.close().await
}
}
}
fn forget_shared(&self, realm: &[u8; 32], mcid: &Mcid) -> Option<Announced> {
let mut realms = self.content.lock();
let shared = realms.get_mut(realm)?;
if let Some(Root::Chunked(m)) = shared.roots.remove(mcid) {
for i in 0..m.chunks.len() {
if let Some(c) = chunk_mcid(&m, i) {
shared.chunks.remove(&c);
}
}
}
shared.announcements.remove(mcid)
}
async fn announce_content(
self: &Arc<Self>,
realm: &[u8; 32],
mcid: &Mcid,
root: &Root,
) -> Result<Record, PoolError> {
let station = self
.links()
.first()
.map(Link::station_node_id)
.ok_or(PoolError::NoLink(Vec::new()))?;
let mut opts = ContentAnnouncementOptions {
realm_id: *realm,
serving_station: station,
procedure: record::own_procedure(&self.self_id, CONTENT_PROCEDURE),
ttl_ms: ANNOUNCEMENT_TTL.as_millis() as u64,
..ContentAnnouncementOptions::default()
};
match root {
Root::Chunked(m) => {
opts.name = String::from_utf8_lossy(&m.name).into_owned();
opts.size = Some(m.size);
opts.chunk_count = Some(m.chunk_count);
}
Root::Block(b) => opts.size = Some(b.len() as u64),
}
let unsigned =
new_content_announcement(&self.self_id, mcid, &opts).map_err(LinkError::from)?;
let signed = record::sign(&unsigned, &self.opts.identity).map_err(LinkError::from)?;
let wire = record::encode(&signed).map_err(LinkError::from)?;
Pool {
inner: self.clone(),
}
.put_record(&wire)
.await?;
Ok(signed)
}
async fn content_sharers(
&self,
realm: &[u8; 32],
mcid: &Mcid,
) -> Result<Vec<Sharing>, PoolError> {
let key = record::content_key(mcid).map_err(LinkError::from)?;
let found = match self
.first_answer(|l| async move { l.find_records(&key).await })
.await
{
Ok((found, _)) => found,
Err(PoolError::Link(LinkError::RecordNotFound)) => Vec::new(),
Err(e) => return Err(e),
};
let mut out: Vec<Sharing> = found
.iter()
.filter(|v| v.record().record_type == RecordType::CONTENT_ANNOUNCEMENT)
.filter_map(|v| record::read_content_announcement(v.record()).ok())
.filter(|a| {
a.mcid.as_slice() == mcid.as_slice()
&& a.realm_id == *realm
&& a.serving_station != [0; 32]
&& content_procedure_bound(&a.procedure, &a.announcer_node)
})
.map(|a| Sharing {
node: a.announcer_node,
station: a.serving_station,
procedure: a.procedure,
})
.collect();
shuffle(&mut out);
Ok(out)
}
async fn fetch_from(
self: &Arc<Self>,
realm: &[u8; 32],
s: &Sharing,
mcid: &Mcid,
opts: &ContentOptions,
) -> Result<Vec<u8>, PoolError> {
let link = self
.link_to(&s.station, Instant::now() + DEFAULT_CALL_TIMEOUT)
.await?;
let (kind, body) = fetch_one(&link, realm, s, mcid, "root", opts.chunk_timeout).await?;
if kind == "block" {
let bytes = wire_bytes(&body, "bytes").unwrap_or_default();
if block_mcid(&bytes) != *mcid {
return Err(PoolError::ContentMismatch("the block".into()));
}
return Ok(bytes);
}
let m = manifest::from_wire(body.get("manifest").unwrap_or(&Value::Null))
.map_err(|e| PoolError::ContentReply(format!("the manifest: {e}")))?;
manifest::verify_mcid(&m, mcid)
.map_err(|e| PoolError::ContentMismatch(format!("the manifest: {e}")))?;
if m.size > opts.max_bytes || m.chunk_count > opts.max_chunks {
return Err(PoolError::ContentTooLarge(format!(
"{} bytes in {} chunks",
m.size, m.chunk_count
)));
}
manifest::check_whole(&m)
.map_err(|e| PoolError::ContentMismatch(format!("the manifest: {e}")))?;
fetch_chunks(link, *realm, s.clone(), Arc::new(m), opts).await
}
}
impl SharedRealm {
fn body(&self, want: &str, mcid: &Mcid) -> Option<Value> {
let block = |b: &[u8]| {
Value::Map(vec![
(Value::text("kind"), Value::text("block")),
(Value::text("mcid"), Value::Bytes(mcid.to_vec())),
(Value::text("bytes"), Value::Bytes(b.to_vec())),
])
};
if want == "block" {
return self.chunks.get(mcid).map(|b| block(b));
}
match self.roots.get(mcid)? {
Root::Block(b) => Some(block(b)),
Root::Chunked(m) => Some(Value::Map(vec![
(Value::text("kind"), Value::text("manifest")),
(Value::text("mcid"), Value::Bytes(mcid.to_vec())),
(Value::text("manifest"), manifest::to_wire(m)),
])),
}
}
}
async fn renew_announcement(
pool: Weak<PoolInner>,
realm: [u8; 32],
mcid: Mcid,
root: Root,
latest: Arc<Mutex<Record>>,
mut stopped: watch::Receiver<bool>,
) {
loop {
tokio::select! {
_ = stopped.wait_for(|s| *s) => return,
_ = tokio::time::sleep(ANNOUNCEMENT_TTL / 2) => {}
}
let Some(inner) = pool.upgrade() else { return };
if inner.lock().closed {
return;
}
let renewed = tokio::time::timeout(
DEFAULT_CALL_TIMEOUT,
inner.announce_content(&realm, &mcid, &root),
)
.await;
if let Ok(Ok(record)) = renewed {
*latest.lock().unwrap_or_else(|p| p.into_inner()) = record;
}
}
}
async fn fetch_chunks(
link: Link,
realm: [u8; 32],
s: Sharing,
m: Arc<Manifest>,
opts: &ContentOptions,
) -> Result<Vec<u8>, PoolError> {
let count = m.chunks.len();
let parts: Arc<Mutex<Vec<Option<Vec<u8>>>>> = Arc::new(Mutex::new(vec![None; count]));
let next = Arc::new(AtomicUsize::new(0));
let (failed_tx, failed) = watch::channel::<Option<PoolError>>(None);
let failed_tx = Arc::new(failed_tx);
let mut workers = tokio::task::JoinSet::new();
for _ in 0..opts.parallel.max(1).min(count) {
let (link, s, m, parts, next, failed_tx) = (
link.clone(),
s.clone(),
m.clone(),
parts.clone(),
next.clone(),
failed_tx.clone(),
);
let timeout = opts.chunk_timeout;
workers.spawn(async move {
loop {
if failed_tx.borrow().is_some() {
return;
}
let i = next.fetch_add(1, Ordering::SeqCst);
let Some(want) = chunk_mcid(&m, i) else {
return;
};
let fetched = fetch_one(&link, &realm, &s, &want, "block", timeout).await;
let outcome = fetched.and_then(|(_, body)| {
let bytes = wire_bytes(&body, "bytes").unwrap_or_default();
if block_mcid(&bytes) == want {
Ok(bytes)
} else {
Err(PoolError::ContentMismatch(format!("chunk {i}")))
}
});
match outcome {
Ok(bytes) => parts.lock().unwrap_or_else(|p| p.into_inner())[i] = Some(bytes),
Err(e) => {
failed_tx.send_if_modified(|f| {
let first = f.is_none();
if first {
*f = Some(e);
}
first
});
return;
}
}
}
});
}
while workers.join_next().await.is_some() {}
if let Some(e) = failed.borrow().clone() {
return Err(e);
}
let parts = std::mem::take(&mut *parts.lock().unwrap_or_else(|p| p.into_inner()));
let mut whole = Vec::with_capacity(m.size as usize);
for part in parts {
whole.extend(part.ok_or_else(|| PoolError::ContentReply("a chunk never arrived".into()))?);
}
manifest::verify(&m, &whole)
.map_err(|e| PoolError::ContentMismatch(format!("the whole: {e}")))?;
Ok(whole)
}
async fn fetch_one(
link: &Link,
realm: &[u8; 32],
s: &Sharing,
mcid: &Mcid,
want: &str,
timeout: Duration,
) -> Result<(String, Value), PoolError> {
let asked = async {
let stream = link
.open_stream(station_link::StreamCall {
realm: *realm,
procedure: s.procedure.clone(),
target: s.node,
mode: StreamMode::ServerStream,
payload: Value::Map(vec![
(Value::text("mcid"), Value::Bytes(mcid.to_vec())),
(Value::text("want"), Value::text(want)),
]),
deadline: timeout,
seal: Some(station_link::Seal::Clear),
..station_link::StreamCall::default()
})
.await?;
let event = stream.recv().await;
let _ = stream.close().await;
read_body(event, mcid, want)
};
tokio::time::timeout(timeout, asked)
.await
.unwrap_or(Err(PoolError::Link(LinkError::CallTimeout)))
}
fn read_body(
event: Result<StreamEvent, LinkError>,
mcid: &Mcid,
want: &str,
) -> Result<(String, Value), PoolError> {
let body = match event {
Err(LinkError::Stream { code, .. }) if code == "not_shared" => {
return Err(PoolError::NotShared)
}
Err(LinkError::EndOfStream) => {
return Err(PoolError::ContentReply(
"the stream ended with no body".into(),
))
}
Err(e) => return Err(e.into()),
Ok(StreamEvent::Data { body, .. }) => body,
Ok(_) => return Err(PoolError::ContentReply("a frame that is not DATA".into())),
};
let kind = wire_text(&body, "kind").unwrap_or_default();
if wire_bytes(&body, "mcid").as_deref() != Some(mcid.as_slice()) {
return Err(PoolError::ContentReply(
"a body for another content id".into(),
));
}
let bytes = wire_bytes(&body, "bytes");
match (kind.as_str(), bytes) {
("block", Some(b)) if b.len() as u64 > MAX_BLOCK_BYTES => Err(PoolError::ContentTooLarge(
format!("a block of {} bytes", b.len()),
)),
("block", Some(_)) => Ok((kind, body)),
("manifest", _) if want == "root" => Ok((kind, body)),
_ => Err(PoolError::ContentReply(format!("kind {kind:?}"))),
}
}
fn wire_text(v: &Value, name: &str) -> Option<String> {
match v.get(name) {
Some(Value::Text(t)) => Some(t.clone()),
Some(Value::Bytes(b)) => String::from_utf8(b.clone()).ok(),
_ => None,
}
}
fn wire_bytes(v: &Value, name: &str) -> Option<Vec<u8>> {
match v.get(name) {
Some(Value::Bytes(b)) => Some(b.clone()),
_ => None,
}
}