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 pub fn register(&self, entry: ReverseServerEntry) {
55 let id = entry.id.clone();
56 let mut guard = self.inner.write().expect("reverse registry poisoned");
57 guard.insert(id, entry);
58 }
59
60 pub fn unregister(&self, id: &ReverseServerId) {
63 let mut guard = self.inner.write().expect("reverse registry poisoned");
64 guard.remove(id);
65 }
66
67 pub fn snapshot(&self) -> Vec<ReverseServerEntrySnapshot> {
69 let guard = self.inner.read().expect("reverse registry poisoned");
70 guard
71 .values()
72 .map(|e| ReverseServerEntrySnapshot {
73 id: e.id.0.to_string(),
74 control_bind: e.control_bind.clone(),
75 state: e.state.snapshot(),
76 })
77 .collect()
78 }
79
80 pub fn is_empty(&self) -> bool {
82 let guard = self.inner.read().expect("reverse registry poisoned");
83 guard.is_empty()
84 }
85}
86
87#[derive(Debug, Clone, serde::Serialize)]
89pub struct ReverseServerEntrySnapshot {
90 pub id: String,
91 pub control_bind: String,
92 #[serde(flatten)]
93 pub state: eggress_protocol_reverse::server::ReverseServerStateSnapshot,
94}
95
96#[cfg(test)]
97mod tests {
98 use super::*;
99
100 #[test]
101 fn register_and_snapshot() {
102 let reg = ReverseRegistry::new();
103 assert!(reg.is_empty());
104 let state = Arc::new(ReverseServerState::default());
105 state
106 .active_control
107 .store(2, std::sync::atomic::Ordering::Relaxed);
108 reg.register(ReverseServerEntry {
109 id: ReverseServerId::from("rev-1"),
110 control_bind: "127.0.0.1:8080".to_string(),
111 state,
112 });
113 assert!(!reg.is_empty());
114 let snap = reg.snapshot();
115 assert_eq!(snap.len(), 1);
116 assert_eq!(snap[0].id, "rev-1");
117 assert_eq!(snap[0].state.active_control, 2);
118 }
119
120 #[test]
121 fn unregister_removes_entry() {
122 let reg = ReverseRegistry::new();
123 let state = Arc::new(ReverseServerState::default());
124 let id = ReverseServerId::from("rev-1");
125 reg.register(ReverseServerEntry {
126 id: id.clone(),
127 control_bind: "127.0.0.1:8080".to_string(),
128 state,
129 });
130 reg.unregister(&id);
131 assert!(reg.is_empty());
132 }
133
134 #[test]
135 fn register_replaces_existing_entry() {
136 let reg = ReverseRegistry::new();
137 let state1 = Arc::new(ReverseServerState::default());
138 state1
139 .active_control
140 .store(1, std::sync::atomic::Ordering::Relaxed);
141 reg.register(ReverseServerEntry {
142 id: ReverseServerId::from("rev-1"),
143 control_bind: "127.0.0.1:8080".to_string(),
144 state: state1,
145 });
146 let state2 = Arc::new(ReverseServerState::default());
147 state2
148 .active_control
149 .store(5, std::sync::atomic::Ordering::Relaxed);
150 reg.register(ReverseServerEntry {
151 id: ReverseServerId::from("rev-1"),
152 control_bind: "127.0.0.1:9090".to_string(),
153 state: state2,
154 });
155 let snap = reg.snapshot();
156 assert_eq!(snap.len(), 1);
157 assert_eq!(snap[0].control_bind, "127.0.0.1:9090");
158 assert_eq!(snap[0].state.active_control, 5);
159 }
160}