pool.go raw

   1  package relay
   2  
   3  import (
   4  	"context"
   5  	"fmt"
   6  	"log"
   7  	"sort"
   8  	"strings"
   9  	"sync"
  10  	"time"
  11  
  12  	"git.mleku.dev/mleku/dendrite/pkg/nostr"
  13  )
  14  
  15  // Pool tracks known relays and their popularity across the network.
  16  // Each time a relay appears in a user's kind-10002 relay list, its encounter
  17  // count increments. Relays that appear in many users' lists are tried first
  18  // when searching for content.
  19  type Pool struct {
  20  	mu      sync.Mutex
  21  	relays  map[string]*PoolRelayInfo
  22  	primary string
  23  }
  24  
  25  // PoolRelayInfo tracks a single relay's statistics.
  26  type PoolRelayInfo struct {
  27  	URL         string        `json:"url"`
  28  	Encounters  int           `json:"encounters"`
  29  	Hits        int           `json:"hits"`
  30  	Misses      int           `json:"misses"`
  31  	RateLimited int           `json:"rate_limited"`
  32  	MinInterval time.Duration `json:"min_interval"`
  33  	LastRequest time.Time     `json:"-"`
  34  }
  35  
  36  // NewPool creates a pool with the given primary relay.
  37  func NewPool(primary string) *Pool {
  38  	primary = normalizeRelayURL(primary)
  39  	return &Pool{
  40  		relays:  map[string]*PoolRelayInfo{primary: {URL: primary}},
  41  		primary: primary,
  42  	}
  43  }
  44  
  45  // Add registers a relay URL, incrementing its encounter count.
  46  func (rp *Pool) Add(url string) {
  47  	url = normalizeRelayURL(url)
  48  	if url == "" {
  49  		return
  50  	}
  51  	rp.mu.Lock()
  52  	defer rp.mu.Unlock()
  53  	if ri, ok := rp.relays[url]; ok {
  54  		ri.Encounters++
  55  	} else {
  56  		rp.relays[url] = &PoolRelayInfo{URL: url, Encounters: 1}
  57  	}
  58  }
  59  
  60  // RecordHit marks a successful fetch from a relay.
  61  func (rp *Pool) RecordHit(url string) {
  62  	url = normalizeRelayURL(url)
  63  	rp.mu.Lock()
  64  	defer rp.mu.Unlock()
  65  	if ri, ok := rp.relays[url]; ok {
  66  		ri.Hits++
  67  	}
  68  }
  69  
  70  // RecordMiss marks a failed/empty fetch from a relay.
  71  func (rp *Pool) RecordMiss(url string) {
  72  	url = normalizeRelayURL(url)
  73  	rp.mu.Lock()
  74  	defer rp.mu.Unlock()
  75  	if ri, ok := rp.relays[url]; ok {
  76  		ri.Misses++
  77  	}
  78  }
  79  
  80  // Ranked returns all relay URLs sorted by encounter count descending,
  81  // excluding the primary relay.
  82  func (rp *Pool) Ranked() []string {
  83  	rp.mu.Lock()
  84  	defer rp.mu.Unlock()
  85  
  86  	type entry struct {
  87  		url   string
  88  		count int
  89  	}
  90  	var entries []entry
  91  	for url, ri := range rp.relays {
  92  		if url == rp.primary {
  93  			continue
  94  		}
  95  		entries = append(entries, entry{url, ri.Encounters})
  96  	}
  97  	sort.Slice(entries, func(i, j int) bool {
  98  		return entries[i].count > entries[j].count
  99  	})
 100  	urls := make([]string, len(entries))
 101  	for i, e := range entries {
 102  		urls[i] = e.url
 103  	}
 104  	return urls
 105  }
 106  
 107  // Size returns the number of known relays.
 108  func (rp *Pool) Size() int {
 109  	rp.mu.Lock()
 110  	defer rp.mu.Unlock()
 111  	return len(rp.relays)
 112  }
 113  
 114  // Stats returns a summary of the relay pool.
 115  func (rp *Pool) Stats() string {
 116  	rp.mu.Lock()
 117  	defer rp.mu.Unlock()
 118  
 119  	var b strings.Builder
 120  	fmt.Fprintf(&b, "relay pool: %d relays\n", len(rp.relays))
 121  
 122  	type entry struct {
 123  		url        string
 124  		encounters int
 125  		hits       int
 126  		misses     int
 127  	}
 128  	var entries []entry
 129  	for _, ri := range rp.relays {
 130  		entries = append(entries, entry{ri.URL, ri.Encounters, ri.Hits, ri.Misses})
 131  	}
 132  	sort.Slice(entries, func(i, j int) bool {
 133  		return entries[i].encounters > entries[j].encounters
 134  	})
 135  
 136  	limit := 20
 137  	if len(entries) < limit {
 138  		limit = len(entries)
 139  	}
 140  	for i := 0; i < limit; i++ {
 141  		e := entries[i]
 142  		fmt.Fprintf(&b, "  %3d encounters, %3d hits, %3d misses: %s\n",
 143  			e.encounters, e.hits, e.misses, e.url)
 144  	}
 145  	if len(entries) > limit {
 146  		fmt.Fprintf(&b, "  ... and %d more\n", len(entries)-limit)
 147  	}
 148  	return b.String()
 149  }
 150  
 151  // FetchRelayList gets the kind-10002 (NIP-65) relay list for a pubkey.
 152  func (rp *Pool) FetchRelayList(ctx context.Context, pubkey string, timeout time.Duration) []string {
 153  	relays := rp.fetchRelayListFrom(ctx, rp.primary, pubkey, timeout)
 154  	if len(relays) == 0 {
 155  		for _, url := range rp.topN(5) {
 156  			relays = rp.fetchRelayListFrom(ctx, url, pubkey, timeout)
 157  			if len(relays) > 0 {
 158  				break
 159  			}
 160  		}
 161  	}
 162  	for _, r := range relays {
 163  		rp.Add(r)
 164  	}
 165  	return relays
 166  }
 167  
 168  func (rp *Pool) topN(n int) []string {
 169  	ranked := rp.Ranked()
 170  	if len(ranked) > n {
 171  		ranked = ranked[:n]
 172  	}
 173  	return ranked
 174  }
 175  
 176  // FetchEventsMultiRelay tries to fetch kind-1 events for a pubkey,
 177  // starting with the primary relay, then the pubkey's own relay list,
 178  // then popular relays from the pool.
 179  func (rp *Pool) FetchEventsMultiRelay(ctx context.Context, pubkey string, limit int, minContent int, timeout time.Duration) ([]*nostr.Event, string) {
 180  	// 1. Try primary relay.
 181  	events, rateLimited := rp.fetchEventsFrom(ctx, rp.primary, pubkey, limit, timeout)
 182  	filtered := filterByContentLen(events, minContent)
 183  	if len(filtered) > 0 {
 184  		rp.RecordHit(rp.primary)
 185  		return filtered, rp.primary
 186  	}
 187  	if !rateLimited {
 188  		rp.RecordMiss(rp.primary)
 189  	}
 190  
 191  	// 2. Get this pubkey's relay list and try those.
 192  	userRelays := rp.FetchRelayList(ctx, pubkey, timeout)
 193  	for _, url := range userRelays {
 194  		if url == rp.primary {
 195  			continue
 196  		}
 197  		events, rateLimited = rp.fetchEventsFrom(ctx, url, pubkey, limit, timeout)
 198  		filtered = filterByContentLen(events, minContent)
 199  		if len(filtered) > 0 {
 200  			rp.RecordHit(url)
 201  			return filtered, url
 202  		}
 203  		if !rateLimited {
 204  			rp.RecordMiss(url)
 205  		}
 206  	}
 207  
 208  	// 3. Fall back to popular relays (max 3 attempts).
 209  	fallbackTries := 0
 210  	for _, url := range rp.Ranked() {
 211  		if ctx.Err() != nil || fallbackTries >= 3 {
 212  			break
 213  		}
 214  		if containsStr(userRelays, url) || url == rp.primary {
 215  			continue
 216  		}
 217  		fallbackTries++
 218  		events, rateLimited = rp.fetchEventsFrom(ctx, url, pubkey, limit, timeout)
 219  		filtered = filterByContentLen(events, minContent)
 220  		if len(filtered) > 0 {
 221  			rp.RecordHit(url)
 222  			return filtered, url
 223  		}
 224  		if !rateLimited {
 225  			rp.RecordMiss(url)
 226  		}
 227  	}
 228  
 229  	return nil, ""
 230  }
 231  
 232  // WaitForRelay blocks until the rate limit interval has elapsed for a relay.
 233  func (rp *Pool) WaitForRelay(ctx context.Context, url string) error {
 234  	url = normalizeRelayURL(url)
 235  	rp.mu.Lock()
 236  	ri, ok := rp.relays[url]
 237  	if !ok {
 238  		rp.mu.Unlock()
 239  		return nil
 240  	}
 241  	interval := ri.MinInterval
 242  	last := ri.LastRequest
 243  	rp.mu.Unlock()
 244  
 245  	if interval > 0 && !last.IsZero() {
 246  		wait := time.Until(last.Add(interval))
 247  		if wait > 0 {
 248  			select {
 249  			case <-time.After(wait):
 250  			case <-ctx.Done():
 251  				return ctx.Err()
 252  			}
 253  		}
 254  	}
 255  
 256  	rp.mu.Lock()
 257  	if ri, ok := rp.relays[url]; ok {
 258  		ri.LastRequest = time.Now()
 259  	}
 260  	rp.mu.Unlock()
 261  	return nil
 262  }
 263  
 264  // RecordRateLimit marks that a relay rate-limited us. Starts at 2s,
 265  // doubles each time, caps at 60s.
 266  func (rp *Pool) RecordRateLimit(url string) {
 267  	url = normalizeRelayURL(url)
 268  	rp.mu.Lock()
 269  	defer rp.mu.Unlock()
 270  	ri, ok := rp.relays[url]
 271  	if !ok {
 272  		return
 273  	}
 274  	ri.RateLimited++
 275  	if ri.MinInterval == 0 {
 276  		ri.MinInterval = 2 * time.Second
 277  	} else {
 278  		ri.MinInterval *= 2
 279  		if ri.MinInterval > 60*time.Second {
 280  			ri.MinInterval = 60 * time.Second
 281  		}
 282  	}
 283  	log.Printf("  rate limited by %s, backing off to %s", url, ri.MinInterval)
 284  }
 285  
 286  func isRateLimitError(err error) bool {
 287  	if err == nil {
 288  		return false
 289  	}
 290  	s := strings.ToLower(err.Error())
 291  	return strings.Contains(s, "rate") || strings.Contains(s, "429") ||
 292  		strings.Contains(s, "too many") || strings.Contains(s, "slow down")
 293  }
 294  
 295  func isRateLimitNotice(notice string) bool {
 296  	s := strings.ToLower(notice)
 297  	return strings.Contains(s, "rate") || strings.Contains(s, "too many") ||
 298  		strings.Contains(s, "slow down") || strings.Contains(s, "throttl")
 299  }
 300  
 301  // --- low-level fetch ---
 302  
 303  func (rp *Pool) fetchEventsFrom(ctx context.Context, relayURL, pubkey string, limit int, timeout time.Duration) ([]*nostr.Event, bool) {
 304  	if err := rp.WaitForRelay(ctx, relayURL); err != nil {
 305  		return nil, false
 306  	}
 307  
 308  	connCtx, connCancel := context.WithTimeout(ctx, 5*time.Second)
 309  	defer connCancel()
 310  
 311  	client, err := nostr.Connect(connCtx, relayURL)
 312  	if err != nil {
 313  		if isRateLimitError(err) {
 314  			rp.RecordRateLimit(relayURL)
 315  			return nil, true
 316  		}
 317  		return nil, false
 318  	}
 319  	defer client.Disconnect()
 320  
 321  	listenDone := make(chan error, 1)
 322  	go func() { listenDone <- client.Listen(ctx) }()
 323  
 324  	if err := client.Subscribe(ctx, "sc-ev", nostr.Filter{
 325  		Authors: []string{pubkey},
 326  		Kinds:   []int{1},
 327  		Limit:   &limit,
 328  	}); err != nil {
 329  		if isRateLimitError(err) {
 330  			rp.RecordRateLimit(relayURL)
 331  			return nil, true
 332  		}
 333  		return nil, false
 334  	}
 335  
 336  	var events []*nostr.Event
 337  	rateLimited := false
 338  	timer := time.NewTimer(timeout)
 339  	defer timer.Stop()
 340  
 341  	for {
 342  		select {
 343  		case ev := <-client.Events:
 344  			if ev != nil {
 345  				events = append(events, ev)
 346  				if len(events) >= limit {
 347  					return events, false
 348  				}
 349  			}
 350  		case notice := <-client.Notices:
 351  			if isRateLimitNotice(notice) {
 352  				rp.RecordRateLimit(relayURL)
 353  				rateLimited = true
 354  			}
 355  		case <-timer.C:
 356  			return events, rateLimited
 357  		case <-listenDone:
 358  			return events, rateLimited
 359  		case <-ctx.Done():
 360  			return events, rateLimited
 361  		}
 362  	}
 363  }
 364  
 365  func (rp *Pool) fetchRelayListFrom(ctx context.Context, relayURL, pubkey string, timeout time.Duration) []string {
 366  	if err := rp.WaitForRelay(ctx, relayURL); err != nil {
 367  		return nil
 368  	}
 369  
 370  	connCtx, connCancel := context.WithTimeout(ctx, 5*time.Second)
 371  	defer connCancel()
 372  
 373  	client, err := nostr.Connect(connCtx, relayURL)
 374  	if err != nil {
 375  		if isRateLimitError(err) {
 376  			rp.RecordRateLimit(relayURL)
 377  		}
 378  		return nil
 379  	}
 380  	defer client.Disconnect()
 381  
 382  	listenDone := make(chan error, 1)
 383  	go func() { listenDone <- client.Listen(ctx) }()
 384  
 385  	fetchLimit := 1
 386  	if err := client.Subscribe(ctx, "sc-rl", nostr.Filter{
 387  		Authors: []string{pubkey},
 388  		Kinds:   []int{10002},
 389  		Limit:   &fetchLimit,
 390  	}); err != nil {
 391  		return nil
 392  	}
 393  
 394  	timer := time.NewTimer(timeout)
 395  	defer timer.Stop()
 396  
 397  	for {
 398  		select {
 399  		case ev := <-client.Events:
 400  			if ev == nil {
 401  				continue
 402  			}
 403  			var relays []string
 404  			for _, tag := range ev.Tags {
 405  				if len(tag) >= 2 && tag[0] == "r" {
 406  					url := normalizeRelayURL(tag[1])
 407  					if url != "" {
 408  						relays = append(relays, url)
 409  					}
 410  				}
 411  			}
 412  			return relays
 413  		case notice := <-client.Notices:
 414  			if isRateLimitNotice(notice) {
 415  				rp.RecordRateLimit(relayURL)
 416  			}
 417  		case <-timer.C:
 418  			return nil
 419  		case <-listenDone:
 420  			return nil
 421  		case <-ctx.Done():
 422  			return nil
 423  		}
 424  	}
 425  }
 426  
 427  // --- helpers ---
 428  
 429  func normalizeRelayURL(url string) string {
 430  	url = strings.TrimSpace(url)
 431  	url = strings.TrimRight(url, "/")
 432  	if url == "" {
 433  		return ""
 434  	}
 435  	if !strings.HasPrefix(url, "wss://") && !strings.HasPrefix(url, "ws://") {
 436  		return ""
 437  	}
 438  	return url
 439  }
 440  
 441  func filterByContentLen(events []*nostr.Event, minContent int) []*nostr.Event {
 442  	var filtered []*nostr.Event
 443  	for _, ev := range events {
 444  		if len(ev.Content) >= minContent {
 445  			filtered = append(filtered, ev)
 446  		}
 447  	}
 448  	return filtered
 449  }
 450  
 451  func containsStr(ss []string, s string) bool {
 452  	for _, v := range ss {
 453  		if v == s {
 454  			return true
 455  		}
 456  	}
 457  	return false
 458  }
 459