952 lines
32 KiB
Rust
952 lines
32 KiB
Rust
use std::collections::{HashMap, HashSet};
|
|
use std::fs;
|
|
use std::net::SocketAddr;
|
|
use std::sync::atomic::{AtomicU64, Ordering};
|
|
use std::sync::{Arc, Mutex, RwLock};
|
|
use std::time::{Duration, Instant, SystemTime};
|
|
|
|
use anyhow::{Context, Result, anyhow};
|
|
use axum::extract::{Query, State};
|
|
use axum::http::StatusCode;
|
|
use axum::response::{IntoResponse, Response};
|
|
use axum::routing::{get, post};
|
|
use axum::{Json, Router};
|
|
use serde::{Deserialize, Serialize};
|
|
use serde_json::{Value, json};
|
|
use tokio::net::TcpListener;
|
|
use tokio::task::JoinHandle;
|
|
|
|
use crate::config::{Config, Member, load_members};
|
|
use crate::crypto::{Keypair, now_ts, poll_signing_bytes};
|
|
use crate::engine::{SearchEngine, TantivyEngine};
|
|
use crate::index::{LocalIndex, SearchHit, response_items};
|
|
use crate::message::{
|
|
AggregateBody, Envelope, QueryBody, ResponseBody, ResponseItem, TYPE_AGGREGATE, TYPE_QUERY,
|
|
TYPE_RESPONSE, build_response, timestamp_is_fresh,
|
|
};
|
|
use crate::registry::Watcher;
|
|
|
|
#[derive(Default)]
|
|
struct Aggregates {
|
|
sent: HashMap<String, u64>,
|
|
passed: HashMap<(String, String), u64>,
|
|
}
|
|
|
|
pub struct Node {
|
|
pub config: Config,
|
|
pub key: Keypair,
|
|
engine: Arc<dyn SearchEngine>,
|
|
members: RwLock<Vec<Member>>,
|
|
members_mtime: Mutex<Option<SystemTime>>,
|
|
registry: Option<Arc<Watcher>>,
|
|
enc_secret: Option<String>,
|
|
enc_public: Option<String>,
|
|
aggregates: Mutex<Aggregates>,
|
|
seen: Mutex<HashSet<String>>,
|
|
pending: Mutex<HashMap<String, Vec<(String, ResponseBody)>>>,
|
|
pending_aggregates: Mutex<HashMap<(String, String), AggregateBody>>,
|
|
client: reqwest::Client,
|
|
sent: AtomicU64,
|
|
received: AtomicU64,
|
|
}
|
|
|
|
pub fn current_period() -> String {
|
|
chrono::Utc::now().format("%Y-%m").to_string()
|
|
}
|
|
|
|
pub fn is_valid_period(period: &str) -> bool {
|
|
if period.len() == 4 {
|
|
return period.chars().all(|c| c.is_ascii_digit());
|
|
}
|
|
if period.len() != 7 || period.as_bytes()[4] != b'-' {
|
|
return false;
|
|
}
|
|
let digits = period[0..4].chars().all(|c| c.is_ascii_digit());
|
|
let month: Result<u32, _> = period[5..7].parse();
|
|
digits && month.map(|m| (1..=12).contains(&m)).unwrap_or(false)
|
|
}
|
|
|
|
pub struct NodeHandle {
|
|
pub addr: SocketAddr,
|
|
pub pubkey: String,
|
|
pub node: Arc<Node>,
|
|
tasks: Vec<JoinHandle<()>>,
|
|
}
|
|
|
|
impl Drop for NodeHandle {
|
|
fn drop(&mut self) {
|
|
for task in &self.tasks {
|
|
task.abort();
|
|
}
|
|
}
|
|
}
|
|
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct RemoteResponse {
|
|
pub member: String,
|
|
pub results: Vec<ResponseItem>,
|
|
pub truncated: bool,
|
|
pub more_available: u64,
|
|
}
|
|
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct MergedItem {
|
|
pub provenance: String,
|
|
#[serde(flatten)]
|
|
pub item: ResponseItem,
|
|
}
|
|
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct LocalQueryOutcome {
|
|
pub qid: String,
|
|
pub text: String,
|
|
pub local: LocalPart,
|
|
pub responses: Vec<RemoteResponse>,
|
|
pub merged: Vec<MergedItem>,
|
|
}
|
|
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct LocalPart {
|
|
pub results: Vec<ResponseItem>,
|
|
pub total: u64,
|
|
}
|
|
|
|
impl Node {
|
|
pub fn open(config: Config) -> Result<Arc<Self>> {
|
|
let key = config.load_key()?;
|
|
let engine: Arc<dyn SearchEngine> = Arc::new(TantivyEngine {
|
|
index: LocalIndex::open(&config.index_dir())?,
|
|
min_coverage: config.matching.min_coverage,
|
|
});
|
|
let members_path = config.members_path();
|
|
let members = load_members(&members_path)?;
|
|
let members_mtime = fs::metadata(&members_path)
|
|
.and_then(|metadata| metadata.modified())
|
|
.ok();
|
|
if !config.node.allow_insecure {
|
|
let insecure = config.insecure_endpoints();
|
|
if !insecure.is_empty() {
|
|
return Err(anyhow!(
|
|
"refusing plain http endpoints (use https, configure ca_cert, or set allow_insecure for private networks): {}",
|
|
insecure.join(", ")
|
|
));
|
|
}
|
|
}
|
|
let client = crate::net::build_client(
|
|
config.node.ca_cert.as_deref().map(std::path::Path::new),
|
|
Duration::from_secs(15),
|
|
)?;
|
|
let registry = match (
|
|
config.node.registry.as_deref(),
|
|
config.node.ma_key.as_deref(),
|
|
) {
|
|
(Some(source), Some(ma_key)) => Some(Watcher::new(
|
|
source,
|
|
ma_key,
|
|
Some(config.registry_cache_path()),
|
|
)),
|
|
(Some(_), None) => return Err(anyhow!("registry configured without ma_key")),
|
|
_ => None,
|
|
};
|
|
if let Some(watcher) = ®istry {
|
|
watcher.load_initial();
|
|
}
|
|
let enc_secret = config.load_enc_key()?;
|
|
let enc_public = match &enc_secret {
|
|
Some(secret) => Some(crate::crypto::enc_public_from_secret(secret)?),
|
|
None => None,
|
|
};
|
|
Ok(Arc::new(Self {
|
|
config,
|
|
key,
|
|
members: RwLock::new(members),
|
|
members_mtime: Mutex::new(members_mtime),
|
|
registry,
|
|
engine,
|
|
enc_secret,
|
|
enc_public,
|
|
aggregates: Mutex::new(Aggregates::default()),
|
|
seen: Mutex::new(HashSet::new()),
|
|
pending: Mutex::new(HashMap::new()),
|
|
pending_aggregates: Mutex::new(HashMap::new()),
|
|
client,
|
|
sent: AtomicU64::new(0),
|
|
received: AtomicU64::new(0),
|
|
}))
|
|
}
|
|
|
|
fn refresh_members(&self) {
|
|
let path = self.config.members_path();
|
|
let mtime = fs::metadata(&path)
|
|
.and_then(|metadata| metadata.modified())
|
|
.ok();
|
|
{
|
|
let last = self.members_mtime.lock().expect("members mtime lock");
|
|
if *last == mtime {
|
|
return;
|
|
}
|
|
}
|
|
let Ok(members) = load_members(&path) else {
|
|
return;
|
|
};
|
|
*self.members.write().expect("members lock") = members;
|
|
*self.members_mtime.lock().expect("members mtime lock") = mtime;
|
|
}
|
|
|
|
fn member_known(&self, members: &[Member], pubkey: &str) -> bool {
|
|
members
|
|
.iter()
|
|
.any(|member| member.pubkey == pubkey || member.previous.iter().any(|key| key == pubkey))
|
|
}
|
|
|
|
pub fn aggregate_for(&self, period: &str, member: Option<&str>) -> AggregateBody {
|
|
let aggregates = self.aggregates.lock().expect("aggregates lock");
|
|
let in_period = |candidate: &str| {
|
|
if period.len() == 4 {
|
|
candidate.starts_with(period)
|
|
} else {
|
|
candidate == period
|
|
}
|
|
};
|
|
let sent = aggregates
|
|
.sent
|
|
.iter()
|
|
.filter(|(key, _)| in_period(key))
|
|
.map(|(_, count)| *count)
|
|
.sum();
|
|
let passed = aggregates
|
|
.passed
|
|
.iter()
|
|
.filter(|((key, key_period), _)| {
|
|
in_period(key_period) && member.map(|m| m == key).unwrap_or(true)
|
|
})
|
|
.map(|(_, count)| *count)
|
|
.sum();
|
|
AggregateBody {
|
|
period: period.to_string(),
|
|
sent,
|
|
passed,
|
|
}
|
|
}
|
|
|
|
pub fn doc_count(&self) -> u64 {
|
|
self.engine.doc_count()
|
|
}
|
|
|
|
pub fn identifier(&self) -> String {
|
|
self.config
|
|
.node
|
|
.id
|
|
.clone()
|
|
.unwrap_or_else(|| self.key.public_hex())
|
|
}
|
|
|
|
pub fn sent(&self) -> u64 {
|
|
self.sent.load(Ordering::SeqCst)
|
|
}
|
|
|
|
pub fn received(&self) -> u64 {
|
|
self.received.load(Ordering::SeqCst)
|
|
}
|
|
|
|
pub fn local_search(&self, text: &str, limit: usize) -> Result<(Vec<SearchHit>, u64)> {
|
|
let output = self.engine.search(text, limit, false)?;
|
|
Ok((output.hits, output.total.unwrap_or(0)))
|
|
}
|
|
|
|
pub async fn start(config: Config) -> Result<NodeHandle> {
|
|
let node = Node::open(config)?;
|
|
let listener = TcpListener::bind(&node.config.node.listen)
|
|
.await
|
|
.with_context(|| format!("binding {}", node.config.node.listen))?;
|
|
let addr = listener.local_addr()?;
|
|
let mut tasks = Vec::new();
|
|
if let Some(watcher) = node.registry.clone() {
|
|
if watcher.is_url() {
|
|
if let Err(error) = watcher.fetch(&node.client).await {
|
|
eprintln!("initial registry fetch failed: {error}");
|
|
}
|
|
let client = node.client.clone();
|
|
tasks.push(tokio::spawn(async move {
|
|
loop {
|
|
tokio::time::sleep(Duration::from_secs(60)).await;
|
|
if let Err(error) = watcher.fetch(&client).await {
|
|
eprintln!("registry fetch failed: {error}");
|
|
}
|
|
}
|
|
}));
|
|
}
|
|
}
|
|
let relays = if node.config.node.relays.is_empty() {
|
|
node.registry
|
|
.as_ref()
|
|
.map(|watcher| watcher.relays())
|
|
.unwrap_or_default()
|
|
} else {
|
|
node.config.node.relays.clone()
|
|
};
|
|
for relay in relays {
|
|
tasks.push(tokio::spawn(poll_relay(node.clone(), relay)));
|
|
}
|
|
let app = router(node.clone());
|
|
tasks.push(tokio::spawn(async move {
|
|
if let Err(error) = axum::serve(listener, app).await {
|
|
eprintln!("control server stopped: {error}");
|
|
}
|
|
}));
|
|
Ok(NodeHandle {
|
|
addr,
|
|
pubkey: node.key.public_hex(),
|
|
node,
|
|
tasks,
|
|
})
|
|
}
|
|
|
|
pub async fn local_query(
|
|
&self,
|
|
text: &str,
|
|
max_results: Option<usize>,
|
|
timeout_ms: Option<u64>,
|
|
network: bool,
|
|
) -> Result<LocalQueryOutcome> {
|
|
let max = max_results.unwrap_or(self.config.query.max_results).max(1);
|
|
let query = QueryBody::new(text, max);
|
|
let qid = query.qid.clone();
|
|
let output = self.engine.search(text, max, false)?;
|
|
let local_hits = output.hits;
|
|
let local_total = output.total.unwrap_or(0);
|
|
let local_items = response_items(&local_hits);
|
|
let mut responses = Vec::new();
|
|
if network {
|
|
self.pending
|
|
.lock()
|
|
.expect("pending lock")
|
|
.insert(qid.clone(), Vec::new());
|
|
self.sent.fetch_add(1, Ordering::SeqCst);
|
|
{
|
|
let mut aggregates = self.aggregates.lock().expect("aggregates lock");
|
|
*aggregates.sent.entry(current_period()).or_default() += 1;
|
|
}
|
|
let envelope = Envelope::new(
|
|
&self.key,
|
|
&self.identifier(),
|
|
TYPE_QUERY,
|
|
serde_json::to_value(&query)?,
|
|
);
|
|
self.publish(&envelope).await;
|
|
let timeout = Duration::from_millis(timeout_ms.unwrap_or(self.config.query.timeout_ms));
|
|
let deadline = Instant::now() + timeout;
|
|
while Instant::now() < deadline {
|
|
tokio::time::sleep(Duration::from_millis(25)).await;
|
|
}
|
|
let collected = self
|
|
.pending
|
|
.lock()
|
|
.expect("pending lock")
|
|
.remove(&qid)
|
|
.unwrap_or_default();
|
|
self.received
|
|
.fetch_add(collected.len() as u64, Ordering::SeqCst);
|
|
for (member, body) in collected {
|
|
{
|
|
let mut aggregates = self.aggregates.lock().expect("aggregates lock");
|
|
*aggregates
|
|
.passed
|
|
.entry((member.clone(), current_period()))
|
|
.or_default() += 1;
|
|
}
|
|
responses.push(RemoteResponse {
|
|
member,
|
|
results: body.results,
|
|
truncated: body.truncated,
|
|
more_available: body.more_available,
|
|
});
|
|
}
|
|
}
|
|
let mut merged = Vec::new();
|
|
let mut seen_urls = HashSet::new();
|
|
for item in &local_items {
|
|
if seen_urls.insert(item.url.clone()) {
|
|
merged.push(MergedItem {
|
|
provenance: "local".to_string(),
|
|
item: item.clone(),
|
|
});
|
|
}
|
|
}
|
|
for response in &responses {
|
|
for item in &response.results {
|
|
if seen_urls.insert(item.url.clone()) {
|
|
merged.push(MergedItem {
|
|
provenance: response.member.clone(),
|
|
item: item.clone(),
|
|
});
|
|
}
|
|
}
|
|
}
|
|
Ok(LocalQueryOutcome {
|
|
qid,
|
|
text: text.to_string(),
|
|
local: LocalPart {
|
|
results: local_items,
|
|
total: local_total,
|
|
},
|
|
responses,
|
|
merged,
|
|
})
|
|
}
|
|
|
|
async fn publish(&self, envelope: &Envelope) -> usize {
|
|
let mut delivered = 0;
|
|
for relay in &self.config.node.relays {
|
|
let url = format!("{}/v1/publish", relay.trim_end_matches('/'));
|
|
match self.client.post(&url).json(envelope).send().await {
|
|
Ok(response) if response.status().is_success() => delivered += 1,
|
|
Ok(response) => eprintln!("relay {relay} rejected query: {}", response.status()),
|
|
Err(error) => eprintln!("relay {relay} unreachable: {error}"),
|
|
}
|
|
}
|
|
delivered
|
|
}
|
|
|
|
fn encrypt_for(&self, recipient_key: &str, body: &Value) -> Value {
|
|
let Some(watcher) = &self.registry else {
|
|
return body.clone();
|
|
};
|
|
let Some(recipient_enc) = watcher.enc_key(recipient_key) else {
|
|
return body.clone();
|
|
};
|
|
let context = crate::crypto::unicast_context(&self.identifier(), &recipient_enc);
|
|
match serde_json::to_vec(body) {
|
|
Ok(plaintext) => crate::crypto::encrypt_unicast(&recipient_enc, &context, &plaintext)
|
|
.unwrap_or_else(|_| body.clone()),
|
|
Err(_) => body.clone(),
|
|
}
|
|
}
|
|
|
|
fn decrypt_body(&self, envelope: &mut Envelope) -> Result<()> {
|
|
let Some(enc) = envelope.body.get("enc").cloned() else {
|
|
return Ok(());
|
|
};
|
|
let secret = self
|
|
.enc_secret
|
|
.as_deref()
|
|
.ok_or_else(|| anyhow!("encrypted message received but no enc key configured"))?;
|
|
let public = self
|
|
.enc_public
|
|
.as_deref()
|
|
.ok_or_else(|| anyhow!("encrypted message received but no enc key configured"))?;
|
|
let context = crate::crypto::unicast_context(&envelope.from, public);
|
|
let plaintext = crate::crypto::decrypt_unicast(secret, &context, &enc)?;
|
|
envelope.body = serde_json::from_slice(&plaintext).context("decrypted body is not json")?;
|
|
Ok(())
|
|
}
|
|
|
|
async fn send_payload(
|
|
&self,
|
|
to_key: &str,
|
|
msg_type: &str,
|
|
body: Value,
|
|
relay: &str,
|
|
) -> Result<()> {
|
|
let body = self.encrypt_for(to_key, &body);
|
|
let envelope = Envelope::new(&self.key, &self.identifier(), msg_type, body);
|
|
self.send_unicast(to_key, envelope, relay).await
|
|
}
|
|
|
|
async fn dispatch(self: &Arc<Self>, mut envelope: Envelope, relay: &str) {
|
|
if envelope.verify().is_err() {
|
|
return;
|
|
}
|
|
if !timestamp_is_fresh(envelope.ts, crate::crypto::now_ts()) {
|
|
return;
|
|
}
|
|
self.refresh_members();
|
|
let listed = if let Some(watcher) = &self.registry {
|
|
watcher.refresh_if_changed();
|
|
match watcher.authorized(&envelope.key, now_ts()) {
|
|
Some(id) if id == envelope.from => true,
|
|
_ => false,
|
|
}
|
|
} else {
|
|
let members = self.members.read().expect("members lock");
|
|
if members.is_empty() {
|
|
self.config.node.dev_bootstrap
|
|
} else {
|
|
self.member_known(&members, &envelope.key) && envelope.from == envelope.key
|
|
}
|
|
};
|
|
if !listed {
|
|
return;
|
|
}
|
|
if let Err(error) = self.decrypt_body(&mut envelope) {
|
|
eprintln!("dropped encrypted message: {error}");
|
|
return;
|
|
}
|
|
match envelope.msg_type.as_str() {
|
|
TYPE_QUERY => {
|
|
if envelope.from == self.identifier() || !self.config.node.responder {
|
|
return;
|
|
}
|
|
let Ok(query) = envelope.parse_body::<QueryBody>() else {
|
|
return;
|
|
};
|
|
if query.text.trim().is_empty() || query.qid.is_empty() {
|
|
return;
|
|
}
|
|
{
|
|
let mut seen = self.seen.lock().expect("seen lock");
|
|
if !seen.insert(query.qid.clone()) {
|
|
return;
|
|
}
|
|
if seen.len() > 10_000 {
|
|
seen.clear();
|
|
}
|
|
}
|
|
let node = self.clone();
|
|
let relay = relay.to_string();
|
|
let querier = envelope.key.clone();
|
|
tokio::spawn(async move {
|
|
if let Err(error) = node.respond(&query, &querier, &relay).await {
|
|
eprintln!("responder failed: {error}");
|
|
}
|
|
});
|
|
}
|
|
TYPE_RESPONSE => {
|
|
let Ok(body) = envelope.parse_body::<ResponseBody>() else {
|
|
return;
|
|
};
|
|
let mut pending = self.pending.lock().expect("pending lock");
|
|
if let Some(list) = pending.get_mut(&body.qid) {
|
|
list.push((envelope.from.clone(), body));
|
|
}
|
|
}
|
|
TYPE_AGGREGATE => {
|
|
let Ok(value) = envelope.parse_body::<Value>() else {
|
|
return;
|
|
};
|
|
let Some(period) = value.get("period").and_then(Value::as_str) else {
|
|
return;
|
|
};
|
|
if !is_valid_period(period) {
|
|
return;
|
|
}
|
|
if value.get("sent").is_some() {
|
|
let Ok(body) = envelope.parse_body::<AggregateBody>() else {
|
|
return;
|
|
};
|
|
let mut pending = self
|
|
.pending_aggregates
|
|
.lock()
|
|
.expect("aggregate reply lock");
|
|
pending.insert((envelope.from.clone(), body.period.clone()), body);
|
|
return;
|
|
}
|
|
let node = self.clone();
|
|
let period = period.to_string();
|
|
let requester_id = envelope.from.clone();
|
|
let requester_key = envelope.key.clone();
|
|
let relay = relay.to_string();
|
|
tokio::spawn(async move {
|
|
node.serve_aggregate(&period, &requester_id, &requester_key, &relay)
|
|
.await;
|
|
});
|
|
}
|
|
_ => {}
|
|
}
|
|
}
|
|
|
|
async fn respond(&self, query: &QueryBody, querier: &str, relay: &str) -> Result<()> {
|
|
let max = query.budget.max_results.clamp(1, 1000);
|
|
let output = self.engine.search(&query.text, max, true)?;
|
|
if output.hits.is_empty() {
|
|
return Ok(());
|
|
}
|
|
let body = build_response(&query.qid, response_items(&output.hits), output.total, max);
|
|
self.send_payload(querier, TYPE_RESPONSE, serde_json::to_value(&body)?, relay)
|
|
.await
|
|
}
|
|
|
|
pub async fn request_aggregate(
|
|
&self,
|
|
to: &str,
|
|
period: &str,
|
|
timeout_ms: Option<u64>,
|
|
) -> Result<AggregateBody> {
|
|
if !is_valid_period(period) {
|
|
return Err(anyhow!("period must be YYYY or YYYY-MM (monthly floor)"));
|
|
}
|
|
let relay = self
|
|
.config
|
|
.node
|
|
.relays
|
|
.first()
|
|
.ok_or_else(|| anyhow!("no relays configured"))?
|
|
.clone();
|
|
let envelope = Envelope::new(
|
|
&self.key,
|
|
&self.identifier(),
|
|
TYPE_AGGREGATE,
|
|
json!({ "period": period }),
|
|
);
|
|
self.send_unicast(to, envelope, &relay).await?;
|
|
let key = (to.to_string(), period.to_string());
|
|
let deadline = Instant::now()
|
|
+ Duration::from_millis(timeout_ms.unwrap_or(self.config.query.timeout_ms));
|
|
loop {
|
|
{
|
|
let pending = self
|
|
.pending_aggregates
|
|
.lock()
|
|
.expect("aggregate reply lock");
|
|
if let Some(body) = pending.get(&key) {
|
|
return Ok(body.clone());
|
|
}
|
|
}
|
|
if Instant::now() >= deadline {
|
|
return Err(anyhow!("aggregate request timed out"));
|
|
}
|
|
tokio::time::sleep(Duration::from_millis(25)).await;
|
|
}
|
|
}
|
|
|
|
async fn serve_aggregate(
|
|
&self,
|
|
period: &str,
|
|
requester_id: &str,
|
|
requester_key: &str,
|
|
relay: &str,
|
|
) {
|
|
let body = self.aggregate_for(period, Some(requester_id));
|
|
let Ok(value) = serde_json::to_value(&body) else {
|
|
return;
|
|
};
|
|
if let Err(error) = self
|
|
.send_payload(requester_key, TYPE_AGGREGATE, value, relay)
|
|
.await
|
|
{
|
|
eprintln!("aggregate reply failed: {error}");
|
|
}
|
|
}
|
|
|
|
async fn send_unicast(&self, to: &str, envelope: Envelope, relay: &str) -> Result<()> {
|
|
let mut relays = vec![relay.to_string()];
|
|
for configured in &self.config.node.relays {
|
|
if !relays.contains(configured) {
|
|
relays.push(configured.clone());
|
|
}
|
|
}
|
|
let mut last_error: Option<anyhow::Error> = None;
|
|
for candidate in relays {
|
|
let url = format!("{}/v1/unicast?to={}", candidate.trim_end_matches('/'), to);
|
|
match self.client.post(&url).json(&envelope).send().await {
|
|
Ok(response) if response.status().is_success() => return Ok(()),
|
|
Ok(response) => {
|
|
last_error = Some(anyhow!(
|
|
"relay {candidate} rejected message: {}",
|
|
response.status()
|
|
))
|
|
}
|
|
Err(error) => last_error = Some(anyhow!("relay {candidate} unreachable: {error}")),
|
|
}
|
|
}
|
|
Err(last_error.unwrap_or_else(|| anyhow!("no relays configured")))
|
|
}
|
|
}
|
|
|
|
async fn fetch_challenge(node: &Node, base: &str, member: &str) -> Option<String> {
|
|
let response = node
|
|
.client
|
|
.get(format!("{base}/v1/challenge?member={member}"))
|
|
.send()
|
|
.await
|
|
.ok()?;
|
|
if !response.status().is_success() {
|
|
return None;
|
|
}
|
|
let payload: Value = response.json().await.ok()?;
|
|
payload
|
|
.get("nonce")
|
|
.and_then(Value::as_str)
|
|
.map(str::to_string)
|
|
}
|
|
|
|
async fn poll_relay(node: Arc<Node>, relay: String) {
|
|
let base = relay.trim_end_matches('/').to_string();
|
|
let member = node.key.public_hex();
|
|
loop {
|
|
let Some(nonce) = fetch_challenge(&node, &base, &member).await else {
|
|
tokio::time::sleep(Duration::from_secs(1)).await;
|
|
continue;
|
|
};
|
|
let signature = node.key.sign(&poll_signing_bytes(&member, &nonce));
|
|
let url = format!(
|
|
"{}/v1/stream?member={}&nonce={}&sig={}",
|
|
base, member, nonce, signature
|
|
);
|
|
match node.client.get(&url).send().await {
|
|
Ok(response) if response.status().is_success() => {
|
|
if let Err(error) = consume_stream(&node, &base, response).await {
|
|
eprintln!("stream from {base} ended: {error}");
|
|
}
|
|
tokio::time::sleep(Duration::from_millis(200)).await;
|
|
}
|
|
Ok(response)
|
|
if response.status() == reqwest::StatusCode::NOT_FOUND
|
|
|| response.status() == reqwest::StatusCode::METHOD_NOT_ALLOWED =>
|
|
{
|
|
long_poll_relay(&node, &base).await;
|
|
return;
|
|
}
|
|
Ok(response) => {
|
|
eprintln!("stream from {base} rejected: {}", response.status());
|
|
tokio::time::sleep(Duration::from_secs(1)).await;
|
|
}
|
|
Err(_) => {
|
|
tokio::time::sleep(Duration::from_secs(1)).await;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
async fn consume_stream(
|
|
node: &Arc<Node>,
|
|
base: &str,
|
|
response: reqwest::Response,
|
|
) -> anyhow::Result<()> {
|
|
use futures_util::StreamExt;
|
|
let mut stream = response.bytes_stream();
|
|
let mut buffer = String::new();
|
|
let mut event = String::new();
|
|
let mut data = String::new();
|
|
while let Some(chunk) = stream.next().await {
|
|
let chunk = chunk?;
|
|
buffer.push_str(&String::from_utf8_lossy(&chunk));
|
|
while let Some(newline) = buffer.find('\n') {
|
|
let line = buffer[..newline].trim_end_matches('\r').to_string();
|
|
buffer.drain(..=newline);
|
|
if line.is_empty() {
|
|
if event == "envelope" && !data.is_empty() {
|
|
if let Ok(envelope) = serde_json::from_str::<Envelope>(&data) {
|
|
node.dispatch(envelope, base).await;
|
|
}
|
|
} else if event == "lag" {
|
|
eprintln!("relay {base} reports lag: {data}");
|
|
}
|
|
event.clear();
|
|
data.clear();
|
|
} else if let Some(rest) = line.strip_prefix("event:") {
|
|
event = rest.trim().to_string();
|
|
} else if let Some(rest) = line.strip_prefix("data:") {
|
|
if !data.is_empty() {
|
|
data.push('\n');
|
|
}
|
|
data.push_str(rest.strip_prefix(' ').unwrap_or(rest));
|
|
}
|
|
}
|
|
}
|
|
Ok(())
|
|
}
|
|
|
|
async fn long_poll_relay(node: &Arc<Node>, base: &str) {
|
|
let member = node.key.public_hex();
|
|
loop {
|
|
let Some(nonce) = fetch_challenge(node, base, &member).await else {
|
|
tokio::time::sleep(Duration::from_secs(1)).await;
|
|
continue;
|
|
};
|
|
let signature = node.key.sign(&poll_signing_bytes(&member, &nonce));
|
|
let url = format!(
|
|
"{}/v1/poll?member={}&nonce={}&sig={}&timeout_ms=20000",
|
|
base, member, nonce, signature
|
|
);
|
|
match node.client.get(&url).send().await {
|
|
Ok(response) if response.status() == reqwest::StatusCode::NO_CONTENT => continue,
|
|
Ok(response) if response.status().is_success() => {
|
|
match response.json::<Value>().await {
|
|
Ok(payload) => {
|
|
if let Some(messages) = payload.get("messages").and_then(Value::as_array) {
|
|
for message in messages {
|
|
if let Ok(envelope) =
|
|
serde_json::from_value::<Envelope>(message.clone())
|
|
{
|
|
node.dispatch(envelope, base).await;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
Err(_) => tokio::time::sleep(Duration::from_secs(1)).await,
|
|
}
|
|
}
|
|
Ok(_) => tokio::time::sleep(Duration::from_secs(1)).await,
|
|
Err(_) => tokio::time::sleep(Duration::from_secs(1)).await,
|
|
}
|
|
}
|
|
}
|
|
|
|
pub fn router(node: Arc<Node>) -> Router {
|
|
Router::new()
|
|
.route("/v1/local/query", post(local_query))
|
|
.route("/v1/local/status", get(local_status))
|
|
.route("/v1/local/aggregates", get(local_aggregates))
|
|
.route("/v1/local/aggregate/request", post(local_aggregate_request))
|
|
.with_state(node)
|
|
}
|
|
|
|
#[derive(Deserialize)]
|
|
pub struct LocalQueryRequest {
|
|
pub text: String,
|
|
pub max_results: Option<usize>,
|
|
pub timeout_ms: Option<u64>,
|
|
#[serde(default = "default_network")]
|
|
pub network: bool,
|
|
}
|
|
|
|
fn default_network() -> bool {
|
|
true
|
|
}
|
|
|
|
async fn local_query(
|
|
State(node): State<Arc<Node>>,
|
|
Json(request): Json<LocalQueryRequest>,
|
|
) -> Response {
|
|
match node
|
|
.local_query(
|
|
&request.text,
|
|
request.max_results,
|
|
request.timeout_ms,
|
|
request.network,
|
|
)
|
|
.await
|
|
{
|
|
Ok(outcome) => (StatusCode::OK, Json(outcome)).into_response(),
|
|
Err(error) => (
|
|
StatusCode::INTERNAL_SERVER_ERROR,
|
|
Json(json!({ "error": error.to_string() })),
|
|
)
|
|
.into_response(),
|
|
}
|
|
}
|
|
|
|
async fn local_status(State(node): State<Arc<Node>>) -> Response {
|
|
let aggregates = node.aggregate_for(¤t_period(), None);
|
|
let registry_version = node.registry.as_ref().and_then(|watcher| watcher.version());
|
|
Json(json!({
|
|
"name": node.config.node.name,
|
|
"id": node.config.node.id,
|
|
"enc_key": node.enc_public,
|
|
"registry_version": registry_version,
|
|
"pubkey": node.key.public_hex(),
|
|
"listen": node.config.node.listen,
|
|
"relays": node.config.node.relays,
|
|
"responder": node.config.node.responder,
|
|
"members": node.members.read().expect("members lock").len(),
|
|
"doc_count": node.doc_count(),
|
|
"sent": node.sent(),
|
|
"received": node.received(),
|
|
"aggregates": aggregates,
|
|
}))
|
|
.into_response()
|
|
}
|
|
|
|
#[derive(Deserialize)]
|
|
struct AggregateParams {
|
|
period: Option<String>,
|
|
}
|
|
|
|
async fn local_aggregates(
|
|
State(node): State<Arc<Node>>,
|
|
Query(params): Query<AggregateParams>,
|
|
) -> Response {
|
|
let period = params.period.unwrap_or_else(current_period);
|
|
if !is_valid_period(&period) {
|
|
return (
|
|
StatusCode::BAD_REQUEST,
|
|
Json(json!({ "error": "period must be YYYY or YYYY-MM (monthly floor)" })),
|
|
)
|
|
.into_response();
|
|
}
|
|
Json(node.aggregate_for(&period, None)).into_response()
|
|
}
|
|
|
|
#[derive(Deserialize)]
|
|
struct AggregateRequest {
|
|
to: String,
|
|
period: String,
|
|
timeout_ms: Option<u64>,
|
|
}
|
|
|
|
async fn local_aggregate_request(
|
|
State(node): State<Arc<Node>>,
|
|
Json(request): Json<AggregateRequest>,
|
|
) -> Response {
|
|
match node
|
|
.request_aggregate(&request.to, &request.period, request.timeout_ms)
|
|
.await
|
|
{
|
|
Ok(body) => Json(body).into_response(),
|
|
Err(error) => (
|
|
StatusCode::BAD_GATEWAY,
|
|
Json(json!({ "error": error.to_string() })),
|
|
)
|
|
.into_response(),
|
|
}
|
|
}
|
|
|
|
pub async fn control_aggregate_request(
|
|
base: &str,
|
|
to: &str,
|
|
period: &str,
|
|
timeout_ms: Option<u64>,
|
|
) -> Result<Value> {
|
|
let client = reqwest::Client::builder()
|
|
.timeout(Duration::from_secs(30))
|
|
.build()?;
|
|
let url = format!("{}/v1/local/aggregate/request", base.trim_end_matches('/'));
|
|
let response = client
|
|
.post(&url)
|
|
.json(&json!({
|
|
"to": to,
|
|
"period": period,
|
|
"timeout_ms": timeout_ms,
|
|
}))
|
|
.send()
|
|
.await
|
|
.with_context(|| format!("calling local node at {url} (is `frxd serve` running?)"))?;
|
|
let status = response.status();
|
|
let value: Value = response.json().await.context("parsing node response")?;
|
|
if !status.is_success() {
|
|
return Err(anyhow!("local node error {status}: {value}"));
|
|
}
|
|
Ok(value)
|
|
}
|
|
|
|
pub async fn control_query(
|
|
base: &str,
|
|
text: &str,
|
|
max_results: Option<usize>,
|
|
timeout_ms: Option<u64>,
|
|
network: bool,
|
|
) -> Result<Value> {
|
|
let client = reqwest::Client::builder()
|
|
.timeout(Duration::from_secs(30))
|
|
.build()?;
|
|
let url = format!("{}/v1/local/query", base.trim_end_matches('/'));
|
|
let response = client
|
|
.post(&url)
|
|
.json(&json!({
|
|
"text": text,
|
|
"max_results": max_results,
|
|
"timeout_ms": timeout_ms,
|
|
"network": network,
|
|
}))
|
|
.send()
|
|
.await
|
|
.with_context(|| format!("calling local node at {url} (is `frxd serve` running?)"))?;
|
|
let status = response.status();
|
|
let value: Value = response.json().await.context("parsing node response")?;
|
|
if !status.is_success() {
|
|
return Err(anyhow!("local node error {status}: {value}"));
|
|
}
|
|
Ok(value)
|
|
}
|