Skip to main content

stasis/infrastructure/runtime/
surreal_cluster_node_store.rs

1use 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}