stasis/infrastructure/runtime/
surreal_cluster_node_store.rs1use async_trait::async_trait;
2use chrono::{DateTime, Duration, Utc};
3use serde::{Deserialize, Serialize};
4use surrealdb::{engine::any::Any, Surreal};
5use surrealdb_types::SurrealValue;
6
7use crate::domain::errors::{Result, StasisError};
8use crate::domain::runtime::cluster_node::{
9 ClusterNode, ClusterNodeHeartbeat, ClusterNodeRole, NewClusterNode,
10};
11use crate::ports::outbound::runtime::cluster_node_store::ClusterNodeStore;
12
13#[derive(Clone)]
14pub struct SurrealClusterNodeStore {
15 db: Surreal<Any>,
16 table: String,
17}
18
19impl SurrealClusterNodeStore {
20 pub fn new(db: Surreal<Any>) -> Self {
21 Self {
22 db,
23 table: "cluster_node".to_string(),
24 }
25 }
26
27 fn port_err(prefix: &str, err: impl std::fmt::Display) -> StasisError {
28 StasisError::PortFailure(format!("{prefix}: {err}"))
29 }
30
31 async fn load_record(&self, node_id: &str) -> Result<Option<ClusterNodeRecord>> {
32 let mut response = match self
33 .db
34 .query("SELECT * FROM type::record($table, $id)")
35 .bind(("table", self.table.clone()))
36 .bind(("id", node_id.to_string()))
37 .await
38 {
39 Ok(response) => response,
40 Err(err) => {
41 let message = err.to_string();
42 if message.contains("does not exist") && message.contains(&self.table) {
43 return Ok(None);
44 }
45 return Err(Self::port_err("load cluster node", err));
46 }
47 };
48
49 let row: Option<ClusterNodeRecord> = match response.take(0) {
50 Ok(row) => row,
51 Err(err) => {
52 let message = err.to_string();
53 if message.contains("does not exist") && message.contains(&self.table) {
54 return Ok(None);
55 }
56 return Err(Self::port_err("decode cluster node", err));
57 }
58 };
59
60 Ok(row)
61 }
62
63 async fn save_record(&self, record: ClusterNodeRecord) -> Result<()> {
64 self.db
65 .query("UPSERT type::record($table, $id) CONTENT $data")
66 .bind(("table", self.table.clone()))
67 .bind(("id", record.node_id.clone()))
68 .bind(("data", record))
69 .await
70 .map_err(|e| Self::port_err("save cluster node", e))?;
71 Ok(())
72 }
73}
74
75#[derive(Clone, Debug, Deserialize, Serialize, SurrealValue)]
76struct ClusterNodeRecord {
77 node_id: String,
78 role: String,
79 region: String,
80 queue_ownership: Vec<String>,
81 capability_tags: Vec<String>,
82 heartbeat_at: DateTime<Utc>,
83 lease_expires_at: DateTime<Utc>,
84 metadata: Option<String>,
85 created_at: DateTime<Utc>,
86 updated_at: DateTime<Utc>,
87}
88
89impl TryFrom<ClusterNodeRecord> for ClusterNode {
90 type Error = StasisError;
91
92 fn try_from(value: ClusterNodeRecord) -> std::result::Result<Self, Self::Error> {
93 let role = match value.role.as_str() {
94 "coordinator" => ClusterNodeRole::Coordinator,
95 "scheduler" => ClusterNodeRole::Scheduler,
96 "worker" => ClusterNodeRole::Worker,
97 other => {
98 return Err(StasisError::PortFailure(format!(
99 "invalid cluster node role: {other}"
100 )));
101 }
102 };
103
104 Ok(Self {
105 node_id: value.node_id,
106 role,
107 region: value.region,
108 queue_ownership: value.queue_ownership,
109 capability_tags: value.capability_tags,
110 heartbeat_at: value.heartbeat_at,
111 lease_expires_at: value.lease_expires_at,
112 metadata: value.metadata,
113 created_at: value.created_at,
114 updated_at: value.updated_at,
115 })
116 }
117}
118
119impl From<ClusterNode> for ClusterNodeRecord {
120 fn from(value: ClusterNode) -> Self {
121 let role = match value.role {
122 ClusterNodeRole::Coordinator => "coordinator".to_string(),
123 ClusterNodeRole::Scheduler => "scheduler".to_string(),
124 ClusterNodeRole::Worker => "worker".to_string(),
125 };
126
127 Self {
128 node_id: value.node_id,
129 role,
130 region: value.region,
131 queue_ownership: value.queue_ownership,
132 capability_tags: value.capability_tags,
133 heartbeat_at: value.heartbeat_at,
134 lease_expires_at: value.lease_expires_at,
135 metadata: value.metadata,
136 created_at: value.created_at,
137 updated_at: value.updated_at,
138 }
139 }
140}
141
142impl From<NewClusterNode> for ClusterNodeRecord {
143 fn from(value: NewClusterNode) -> Self {
144 ClusterNodeRecord::from(value.into_record())
145 }
146}
147
148#[async_trait]
149impl ClusterNodeStore for SurrealClusterNodeStore {
150 async fn register(&self, node: NewClusterNode) -> Result<ClusterNode> {
151 let record: ClusterNodeRecord = node.into();
152 if self.load_record(&record.node_id).await?.is_some() {
153 return Err(StasisError::PortFailure(format!(
154 "cluster node already exists: {}",
155 record.node_id
156 )));
157 }
158
159 self.save_record(record.clone()).await?;
160 ClusterNode::try_from(record)
161 }
162
163 async fn heartbeat(&self, heartbeat: ClusterNodeHeartbeat) -> Result<Option<ClusterNode>> {
164 let Some(existing) = self.load_record(&heartbeat.node_id).await? else {
165 return Ok(None);
166 };
167
168 let mut node = ClusterNode::try_from(existing)?;
169 node.heartbeat_at = heartbeat.heartbeat_at;
170 node.lease_expires_at =
171 heartbeat.heartbeat_at + Duration::seconds(heartbeat.lease_ttl_seconds.max(1));
172 if let Some(queue_ownership) = heartbeat.queue_ownership {
173 node.queue_ownership = queue_ownership;
174 }
175 if let Some(capability_tags) = heartbeat.capability_tags {
176 node.capability_tags = capability_tags;
177 }
178 if heartbeat.metadata.is_some() {
179 node.metadata = heartbeat.metadata;
180 }
181 node.updated_at = heartbeat.heartbeat_at;
182
183 self.save_record(node.clone().into()).await?;
184 Ok(Some(node))
185 }
186
187 async fn get(&self, node_id: &str) -> Result<Option<ClusterNode>> {
188 self.load_record(node_id)
189 .await?
190 .map(ClusterNode::try_from)
191 .transpose()
192 }
193
194 async fn list(&self) -> Result<Vec<ClusterNode>> {
195 let mut response = match self
196 .db
197 .query("SELECT * FROM type::table($table)")
198 .bind(("table", self.table.clone()))
199 .await
200 {
201 Ok(response) => response,
202 Err(err) => {
203 let message = err.to_string();
204 if message.contains("does not exist") && message.contains(&self.table) {
205 return Ok(Vec::new());
206 }
207 return Err(Self::port_err("list cluster nodes", err));
208 }
209 };
210
211 let rows: Vec<ClusterNodeRecord> = match response.take(0) {
212 Ok(rows) => rows,
213 Err(err) => {
214 let message = err.to_string();
215 if message.contains("does not exist") && message.contains(&self.table) {
216 return Ok(Vec::new());
217 }
218 return Err(Self::port_err("decode cluster nodes", err));
219 }
220 };
221
222 let mut nodes = Vec::with_capacity(rows.len());
223 for row in rows {
224 nodes.push(ClusterNode::try_from(row)?);
225 }
226 nodes.sort_by(|a, b| a.node_id.cmp(&b.node_id));
227 Ok(nodes)
228 }
229
230 async fn remove(&self, node_id: &str) -> Result<bool> {
231 let mut response = match self
232 .db
233 .query("DELETE type::record($table, $id) RETURN BEFORE")
234 .bind(("table", self.table.clone()))
235 .bind(("id", node_id.to_string()))
236 .await
237 {
238 Ok(response) => response,
239 Err(err) => {
240 let message = err.to_string();
241 if message.contains("does not exist") && message.contains(&self.table) {
242 return Ok(false);
243 }
244 return Err(Self::port_err("remove cluster node", err));
245 }
246 };
247
248 let deleted: Option<ClusterNodeRecord> = match response.take(0) {
249 Ok(value) => value,
250 Err(err) => {
251 let message = err.to_string();
252 if message.contains("does not exist") && message.contains(&self.table) {
253 return Ok(false);
254 }
255 return Err(Self::port_err("decode removed cluster node", err));
256 }
257 };
258
259 Ok(deleted.is_some())
260 }
261
262 async fn prune_expired(&self, now: DateTime<Utc>) -> Result<u64> {
263 let mut response = match self
264 .db
265 .query("DELETE type::table($table) WHERE lease_expires_at < $now RETURN BEFORE")
266 .bind(("table", self.table.clone()))
267 .bind(("now", now))
268 .await
269 {
270 Ok(response) => response,
271 Err(err) => {
272 let message = err.to_string();
273 if message.contains("does not exist") && message.contains(&self.table) {
274 return Ok(0);
275 }
276 return Err(Self::port_err("prune expired cluster nodes", err));
277 }
278 };
279
280 let deleted: Vec<ClusterNodeRecord> = match response.take(0) {
281 Ok(rows) => rows,
282 Err(err) => {
283 let message = err.to_string();
284 if message.contains("does not exist") && message.contains(&self.table) {
285 return Ok(0);
286 }
287 return Err(Self::port_err("decode pruned cluster nodes", err));
288 }
289 };
290
291 Ok(deleted.len() as u64)
292 }
293}