gossip_relay_picker/
lib.rs1use async_trait::async_trait;
14use dashmap::DashMap;
15pub use nostr_types::{PublicKeyHex, RelayUrl, Unixtime};
16use thiserror::Error;
17
18#[derive(Debug, Copy, Clone)]
20pub enum Direction {
21 Read,
22 Write,
23}
24
25#[derive(Error, Debug, Clone, PartialEq, Eq)]
27pub enum Error {
28 #[error("No relays to pick from")]
30 NoRelays,
31
32 #[error("All people accounted for.")]
34 NoPeopleLeft,
35
36 #[error("Unable to make further progress.")]
38 NoProgress,
39
40 #[error("Error: {0}")]
42 General(String),
43}
44
45#[derive(Debug, Clone)]
48pub struct RelayAssignment {
49 pub relay_url: RelayUrl,
51
52 pub pubkeys: Vec<PublicKeyHex>,
54}
55
56impl RelayAssignment {
57 pub fn merge_in(&mut self, other: RelayAssignment) -> Result<(), Error> {
58 if self.relay_url != other.relay_url {
59 return Err(Error::General(
60 "Attempted to merge relay assignments on different relays".to_owned(),
61 ));
62 }
63 self.pubkeys.extend(other.pubkeys);
64 Ok(())
65 }
66}
67
68#[async_trait]
70pub trait RelayPickerHooks: Send + Sync {
71 type Error: std::fmt::Display;
72
73 fn get_all_relays(&self) -> Vec<RelayUrl>;
75
76 async fn get_relays_for_pubkey(
82 &self,
83 pubkey: PublicKeyHex,
84 direction: Direction,
85 ) -> Result<Vec<(RelayUrl, u64)>, Self::Error>;
86
87 fn is_relay_connected(&self, relay: &RelayUrl) -> bool;
89
90 fn get_max_relays(&self) -> usize;
92
93 fn get_num_relays_per_person(&self) -> usize;
96
97 fn get_followed_pubkeys(&self) -> Vec<PublicKeyHex>;
99
100 fn adjust_score(&self, relay: RelayUrl, score: u64) -> u64;
102}
103
104#[derive(Debug, Default)]
109pub struct RelayPicker<H: RelayPickerHooks + Default> {
110 hooks: H,
112
113 person_relay_scores: DashMap<PublicKeyHex, Vec<(RelayUrl, u64)>>,
115
116 relay_assignments: DashMap<RelayUrl, RelayAssignment>,
119
120 excluded_relays: DashMap<RelayUrl, i64>,
124
125 pubkey_counts: DashMap<PublicKeyHex, usize>,
129}
130
131impl<H: RelayPickerHooks + Default> RelayPicker<H> {
132 pub async fn new(hooks: H) -> Result<RelayPicker<H>, Error> {
134 let rp = RelayPicker {
135 hooks,
136 ..Default::default()
137 };
138
139 rp.refresh_person_relay_scores_inner(true).await?;
140
141 Ok(rp)
142 }
143
144 pub async fn init(&self) -> Result<(), Error> {
147 self.relay_assignments.clear();
148 self.excluded_relays.clear();
149 self.pubkey_counts.clear();
150 self.person_relay_scores.clear();
151
152 self.refresh_person_relay_scores_inner(true).await?;
153
154 Ok(())
155 }
156
157 pub fn add_someone(&self, pubkey: PublicKeyHex) -> Result<(), Error> {
159 if self.pubkey_counts.get(&pubkey).is_some() {
161 return Ok(());
162 }
163 for elem in self.relay_assignments.iter() {
164 let assignment = elem.value();
165 if assignment.pubkeys.contains(&pubkey) {
166 return Ok(());
167 }
168 }
169
170 self.pubkey_counts
171 .insert(pubkey, self.hooks.get_num_relays_per_person());
172 Ok(())
173 }
174
175 pub fn remove_someone(&self, pubkey: PublicKeyHex) {
177 self.pubkey_counts.remove(&pubkey);
179
180 for mut elem in self.relay_assignments.iter_mut() {
182 let assignment = elem.value_mut();
183 if let Some(pos) = assignment.pubkeys.iter().position(|x| x == &pubkey) {
184 assignment.pubkeys.remove(pos);
185 }
186 }
187 }
188
189 pub async fn refresh_person_relay_scores(&self) -> Result<(), Error> {
191 self.refresh_person_relay_scores_inner(false).await
192 }
193
194 async fn refresh_person_relay_scores_inner(
196 &self,
197 initialize_counts: bool,
198 ) -> Result<(), Error> {
199 self.person_relay_scores.clear();
200
201 if initialize_counts {
202 self.pubkey_counts.clear();
203 }
204
205 let pubkeys: Vec<PublicKeyHex> = self
207 .hooks
208 .get_followed_pubkeys()
209 .iter()
210 .map(|p| p.to_owned())
211 .collect();
212
213 for pubkey in &pubkeys {
215 let best_relays: Vec<(RelayUrl, u64)> = self
216 .hooks
217 .get_relays_for_pubkey(pubkey.to_owned(), Direction::Write)
218 .await
219 .map_err(|e| Error::General(format!("{e}")))?;
220 self.person_relay_scores.insert(pubkey.clone(), best_relays);
221
222 if initialize_counts {
223 self.pubkey_counts
224 .insert(pubkey.clone(), self.hooks.get_num_relays_per_person());
225 }
226 }
227
228 Ok(())
229 }
230
231 pub fn relay_disconnected(&self, url: &RelayUrl) {
234 if let Some((_key, assignment)) = self.relay_assignments.remove(url) {
236 let hence = Unixtime::now().unwrap().0 + 30;
238 self.excluded_relays.insert(url.to_owned(), hence);
239 tracing::debug!("{} goes into the penalty box until {}", url, hence,);
240
241 for pubkey in assignment.pubkeys.iter() {
243 self.pubkey_counts
244 .entry(pubkey.to_owned())
245 .and_modify(|e| *e += 1)
246 .or_insert(1);
247 }
248 }
249 }
250
251 pub async fn pick(&self) -> Result<RelayUrl, Error> {
256 let at_max_relays = self.relay_assignments.len() >= self.hooks.get_max_relays();
259
260 let now = Unixtime::now().unwrap().0;
262 self.excluded_relays.retain(|_, v| *v > now);
263
264 if self.pubkey_counts.is_empty() {
265 return Err(Error::NoPeopleLeft);
266 }
267
268 let all_relays = self.hooks.get_all_relays();
269
270 if all_relays.is_empty() {
271 return Err(Error::NoRelays);
272 }
273
274 let scoreboard: DashMap<RelayUrl, u64> =
276 all_relays.iter().map(|x| (x.to_owned(), 0)).collect();
277
278 for elem in self.person_relay_scores.iter() {
280 let pubkeyhex = elem.key();
281 let relay_scores = elem.value();
282
283 if let Some(pkc) = self.pubkey_counts.get(pubkeyhex) {
285 if *pkc == 0 {
286 continue;
288 }
289 } else {
290 continue; }
292
293 for (relay, score) in relay_scores.iter() {
295 if self.excluded_relays.contains_key(relay) {
297 continue;
298 }
299
300 if at_max_relays && !self.hooks.is_relay_connected(relay) {
302 continue;
303 }
304
305 if let Some(assignment) = self.relay_assignments.get(relay) {
307 if assignment.pubkeys.contains(pubkeyhex) {
308 continue;
309 }
310 }
311
312 if let Some(mut entry) = scoreboard.get_mut(relay) {
314 *entry += score;
315 }
316 }
317 }
318
319 for mut score_entry in scoreboard.iter_mut() {
323 let url = score_entry.key().to_owned();
324 let score = score_entry.value_mut();
325 *score = self.hooks.adjust_score(url, *score);
326 }
327
328 let winner = scoreboard
329 .iter()
330 .max_by(|x, y| x.value().cmp(y.value()))
331 .unwrap();
332 let winning_url: RelayUrl = winner.key().to_owned();
333 let winning_score: u64 = *winner.value();
334
335 if winning_score == 0 {
336 return Err(Error::NoProgress);
337 }
338
339 let covered_public_keys = {
343 let pubkeys_seeking_relays: Vec<PublicKeyHex> = self
344 .pubkey_counts
345 .iter()
346 .filter(|e| *e.value() > 0)
347 .map(|e| e.key().to_owned())
348 .collect();
349
350 let mut covered_pubkeys: Vec<PublicKeyHex> = Vec::new();
351
352 for pubkey in pubkeys_seeking_relays.iter() {
353 if let Some(assignment) = self.relay_assignments.get(&winning_url) {
355 if assignment.pubkeys.contains(pubkey) {
356 continue;
357 }
358 }
359
360 if let Some(elem) = self.person_relay_scores.get(pubkey) {
361 let relay_scores = elem.value();
362
363 if relay_scores.iter().any(|e| e.0 == winning_url) {
364 covered_pubkeys.push(pubkey.to_owned());
365
366 if let Some(mut count) = self.pubkey_counts.get_mut(pubkey) {
367 if *count > 0 {
368 *count -= 1;
369 }
370 }
371 }
372 }
373 }
374
375 covered_pubkeys
376 };
377
378 if covered_public_keys.is_empty() {
379 return Err(Error::NoProgress);
380 }
381
382 self.pubkey_counts.retain(|_, count| *count > 0);
384
385 let assignment = RelayAssignment {
386 relay_url: winning_url.clone(),
387 pubkeys: covered_public_keys,
388 };
389
390 if let Some(mut maybe_elem) = self.relay_assignments.get_mut(&winning_url) {
392 maybe_elem.value_mut().merge_in(assignment).unwrap();
394 } else {
395 self.relay_assignments
396 .insert(winning_url.clone(), assignment);
397 }
398
399 Ok(winning_url)
400 }
401
402 pub fn get_relay_assignment(&self, relay_url: &RelayUrl) -> Option<RelayAssignment> {
404 self.relay_assignments
405 .get(relay_url)
406 .map(|elem| elem.value().to_owned())
407 }
408
409 pub fn relay_assignments_iter(&self) -> dashmap::iter::Iter<'_, RelayUrl, RelayAssignment> {
411 self.relay_assignments.iter()
412 }
413
414 pub fn excluded_relays_iter(&self) -> dashmap::iter::Iter<'_, RelayUrl, i64> {
417 self.excluded_relays.iter()
418 }
419
420 pub fn pubkey_counts_iter(&self) -> dashmap::iter::Iter<'_, PublicKeyHex, usize> {
423 self.pubkey_counts.iter()
424 }
425}