1use std::collections::HashMap;
12use std::sync::Arc;
13
14use eggress_protocol_reverse::server::ReverseServerState;
15
16#[derive(Debug, Clone, Hash, Eq, PartialEq)]
19pub struct ReverseServerId(pub Arc<str>);
20
21impl std::fmt::Display for ReverseServerId {
22 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
23 f.write_str(&self.0)
24 }
25}
26
27impl From<&str> for ReverseServerId {
28 fn from(s: &str) -> Self {
29 Self(Arc::from(s))
30 }
31}
32
33pub struct ReverseServerEntry {
35 pub id: ReverseServerId,
36 pub control_bind: String,
37 pub state: Arc<ReverseServerState>,
38}
39
40#[derive(Default, Clone)]
43pub struct ReverseRegistry {
44 inner: Arc<std::sync::RwLock<HashMap<ReverseServerId, ReverseServerEntry>>>,
45}
46
47impl ReverseRegistry {
48 pub fn new() -> Self {
49 Self::default()
50 }
51
52 fn read_lock(
53 &self,
54 ) -> std::sync::RwLockReadGuard<'_, HashMap<ReverseServerId, ReverseServerEntry>> {
55 self.inner.read().unwrap_or_else(|e| {
56 tracing::warn!("reverse registry read lock was poisoned; recovering: {e}");
57 e.into_inner()
58 })
59 }
60
61 fn write_lock(
62 &self,
63 ) -> std::sync::RwLockWriteGuard<'_, HashMap<ReverseServerId, ReverseServerEntry>> {
64 self.inner.write().unwrap_or_else(|e| {
65 tracing::warn!("reverse registry write lock was poisoned; recovering: {e}");
66 e.into_inner()
67 })
68 }
69
70 pub fn register(&self, entry: ReverseServerEntry) {
73 let id = entry.id.clone();
74 let mut guard = self.write_lock();
75 guard.insert(id, entry);
76 }
77
78 pub fn unregister(&self, id: &ReverseServerId) {
81 let mut guard = self.write_lock();
82 guard.remove(id);
83 }
84
85 pub fn snapshot(&self) -> Vec<ReverseServerEntrySnapshot> {
87 let guard = self.read_lock();
88 guard
89 .values()
90 .map(|e| ReverseServerEntrySnapshot {
91 id: e.id.0.to_string(),
92 control_bind: e.control_bind.clone(),
93 state: e.state.snapshot(),
94 })
95 .collect()
96 }
97
98 pub fn is_empty(&self) -> bool {
100 let guard = self.read_lock();
101 guard.is_empty()
102 }
103}
104
105#[derive(Debug, Clone, serde::Serialize)]
107pub struct ReverseServerEntrySnapshot {
108 pub id: String,
109 pub control_bind: String,
110 #[serde(flatten)]
111 pub state: eggress_protocol_reverse::server::ReverseServerStateSnapshot,
112}
113
114#[cfg(test)]
115mod tests {
116 use super::*;
117
118 #[test]
119 fn register_and_snapshot() {
120 let reg = ReverseRegistry::new();
121 assert!(reg.is_empty());
122 let state = Arc::new(ReverseServerState::default());
123 state
124 .active_control
125 .store(2, std::sync::atomic::Ordering::Relaxed);
126 reg.register(ReverseServerEntry {
127 id: ReverseServerId::from("rev-1"),
128 control_bind: "127.0.0.1:8080".to_string(),
129 state,
130 });
131 assert!(!reg.is_empty());
132 let snap = reg.snapshot();
133 assert_eq!(snap.len(), 1);
134 assert_eq!(snap[0].id, "rev-1");
135 assert_eq!(snap[0].state.active_control, 2);
136 }
137
138 #[test]
139 fn unregister_removes_entry() {
140 let reg = ReverseRegistry::new();
141 let state = Arc::new(ReverseServerState::default());
142 let id = ReverseServerId::from("rev-1");
143 reg.register(ReverseServerEntry {
144 id: id.clone(),
145 control_bind: "127.0.0.1:8080".to_string(),
146 state,
147 });
148 reg.unregister(&id);
149 assert!(reg.is_empty());
150 }
151
152 #[test]
153 fn register_replaces_existing_entry() {
154 let reg = ReverseRegistry::new();
155 let state1 = Arc::new(ReverseServerState::default());
156 state1
157 .active_control
158 .store(1, std::sync::atomic::Ordering::Relaxed);
159 reg.register(ReverseServerEntry {
160 id: ReverseServerId::from("rev-1"),
161 control_bind: "127.0.0.1:8080".to_string(),
162 state: state1,
163 });
164 let state2 = Arc::new(ReverseServerState::default());
165 state2
166 .active_control
167 .store(5, std::sync::atomic::Ordering::Relaxed);
168 reg.register(ReverseServerEntry {
169 id: ReverseServerId::from("rev-1"),
170 control_bind: "127.0.0.1:9090".to_string(),
171 state: state2,
172 });
173 let snap = reg.snapshot();
174 assert_eq!(snap.len(), 1);
175 assert_eq!(snap[0].control_bind, "127.0.0.1:9090");
176 assert_eq!(snap[0].state.active_control, 5);
177 }
178}