Add Phase 1 frxd implementation with conformance test suite
This commit is contained in:
+434
@@ -0,0 +1,434 @@
|
||||
use std::collections::{HashMap, HashSet};
|
||||
use std::net::SocketAddr;
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use anyhow::{Context, Result, anyhow};
|
||||
use axum::extract::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;
|
||||
use crate::crypto::Keypair;
|
||||
use crate::index::{LocalIndex, SearchHit, response_items};
|
||||
use crate::message::{
|
||||
Envelope, QueryBody, ResponseBody, ResponseItem, TYPE_AGGREGATE, TYPE_QUERY, TYPE_RESPONSE,
|
||||
build_response,
|
||||
};
|
||||
|
||||
pub struct Node {
|
||||
pub config: Config,
|
||||
pub key: Keypair,
|
||||
index: LocalIndex,
|
||||
seen: Mutex<HashSet<String>>,
|
||||
pending: Mutex<HashMap<String, Vec<(String, ResponseBody)>>>,
|
||||
client: reqwest::Client,
|
||||
sent: AtomicU64,
|
||||
received: AtomicU64,
|
||||
}
|
||||
|
||||
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 index = LocalIndex::open(&config.index_dir())?;
|
||||
let client = reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(15))
|
||||
.build()
|
||||
.context("building http client")?;
|
||||
Ok(Arc::new(Self {
|
||||
config,
|
||||
key,
|
||||
index,
|
||||
seen: Mutex::new(HashSet::new()),
|
||||
pending: Mutex::new(HashMap::new()),
|
||||
client,
|
||||
sent: AtomicU64::new(0),
|
||||
received: AtomicU64::new(0),
|
||||
}))
|
||||
}
|
||||
|
||||
pub fn doc_count(&self) -> u64 {
|
||||
self.index.doc_count()
|
||||
}
|
||||
|
||||
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)> {
|
||||
self.index.search(text, limit, false)
|
||||
}
|
||||
|
||||
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();
|
||||
for relay in node.config.node.relays.clone() {
|
||||
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 (local_hits, local_total) = self.index.search(text, max, false)?;
|
||||
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 envelope = Envelope::new(&self.key, 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 {
|
||||
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
|
||||
}
|
||||
|
||||
async fn dispatch(self: &Arc<Self>, envelope: Envelope, relay: &str) {
|
||||
if envelope.verify().is_err() {
|
||||
return;
|
||||
}
|
||||
if !self.config.node.trusted_keys.is_empty()
|
||||
&& !self.config.node.trusted_keys.contains(&envelope.from)
|
||||
{
|
||||
return;
|
||||
}
|
||||
match envelope.msg_type.as_str() {
|
||||
TYPE_QUERY => {
|
||||
if envelope.from == self.key.public_hex() || !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.from.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 => {}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
async fn respond(&self, query: &QueryBody, querier: &str, relay: &str) -> Result<()> {
|
||||
let max = query.budget.max_results.clamp(1, 1000);
|
||||
let (hits, total) = self.index.search(&query.text, max, true)?;
|
||||
if hits.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
let body = build_response(&query.qid, response_items(&hits), total, max);
|
||||
let envelope = Envelope::new(&self.key, TYPE_RESPONSE, serde_json::to_value(&body)?);
|
||||
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('/'),
|
||||
querier
|
||||
);
|
||||
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 response: {}",
|
||||
response.status()
|
||||
))
|
||||
}
|
||||
Err(error) => last_error = Some(anyhow!("relay {candidate} unreachable: {error}")),
|
||||
}
|
||||
}
|
||||
Err(last_error.unwrap_or_else(|| anyhow!("no relays configured")))
|
||||
}
|
||||
}
|
||||
|
||||
async fn poll_relay(node: Arc<Node>, relay: String) {
|
||||
let base = relay.trim_end_matches('/').to_string();
|
||||
loop {
|
||||
let url = format!(
|
||||
"{}/v1/poll?member={}&timeout_ms=20000",
|
||||
base,
|
||||
node.key.public_hex()
|
||||
);
|
||||
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))
|
||||
.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 {
|
||||
Json(json!({
|
||||
"name": node.config.node.name,
|
||||
"pubkey": node.key.public_hex(),
|
||||
"listen": node.config.node.listen,
|
||||
"relays": node.config.node.relays,
|
||||
"responder": node.config.node.responder,
|
||||
"doc_count": node.doc_count(),
|
||||
"sent": node.sent(),
|
||||
"received": node.received(),
|
||||
}))
|
||||
.into_response()
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
Reference in New Issue
Block a user