Skip to main content

shared/models/
server_tunnel.rs

1use crate::{
2    models::{InsertQueryBuilder, UpdateQueryBuilder},
3    prelude::*,
4    tunnel::{MAX_SERVER_IDX, TunnelProtocol},
5};
6use garde::Validate;
7use serde::{Deserialize, Serialize};
8use sqlx::{Row, postgres::PgRow};
9use std::{
10    collections::{BTreeMap, HashMap, HashSet},
11    sync::{Arc, LazyLock},
12};
13use utoipa::ToSchema;
14
15const ADVISORY_LOCK_IDX: i64 = 0x7475_6e6e_656c_0001;
16
17fn validate_name(name: &compact_str::CompactString, _context: &()) -> Result<(), garde::Error> {
18    if name.is_empty() || name.len() > 63 {
19        return Err(garde::Error::new("name must be 1 to 63 characters"));
20    }
21
22    if !name
23        .bytes()
24        .all(|byte| byte.is_ascii_lowercase() || byte.is_ascii_digit() || byte == b'-')
25    {
26        return Err(garde::Error::new(
27            "name must only contain lowercase letters, digits and dashes",
28        ));
29    }
30
31    if name.starts_with('-') || name.ends_with('-') {
32        return Err(garde::Error::new("name must not start or end with a dash"));
33    }
34
35    if crate::tunnel::is_alias_shaped(name) {
36        return Err(garde::Error::new(
37            "name must not be eight hexadecimal characters, which is reserved for the address every server keeps",
38        ));
39    }
40
41    Ok(())
42}
43
44pub fn validate_optional_name(
45    name: &Option<compact_str::CompactString>,
46    context: &(),
47) -> Result<(), garde::Error> {
48    match name {
49        Some(name) => validate_name(name, context),
50        None => Ok(()),
51    }
52}
53
54fn validate_protocols(
55    protocols: &HashSet<TunnelProtocol>,
56    _context: &(),
57) -> Result<(), garde::Error> {
58    if protocols.is_empty() {
59        return Err(garde::Error::new("at least one protocol is required"));
60    }
61
62    Ok(())
63}
64
65#[derive(Serialize, Deserialize)]
66pub struct ServerTunnel {
67    pub server: Fetchable<super::server::Server>,
68
69    pub idx: u16,
70    pub name: compact_str::CompactString,
71
72    pub created: chrono::NaiveDateTime,
73
74    extension_data: super::ModelExtensionData,
75}
76
77impl BaseModel for ServerTunnel {
78    const NAME: &'static str = "server_tunnel";
79
80    fn get_extension_list() -> &'static super::ModelExtensionList {
81        static EXTENSIONS: LazyLock<super::ModelExtensionList> =
82            LazyLock::new(|| parking_lot::RwLock::new(Vec::new()));
83
84        &EXTENSIONS
85    }
86
87    fn get_extension_data(&self) -> &super::ModelExtensionData {
88        &self.extension_data
89    }
90
91    #[inline]
92    fn base_columns(prefix: Option<&str>) -> BTreeMap<&'static str, compact_str::CompactString> {
93        let prefix = prefix.unwrap_or_default();
94
95        BTreeMap::from([
96            (
97                "server_tunnels.server_uuid",
98                compact_str::format_compact!("{prefix}server_uuid"),
99            ),
100            (
101                "server_tunnels.idx",
102                compact_str::format_compact!("{prefix}idx"),
103            ),
104            (
105                "server_tunnels.name",
106                compact_str::format_compact!("{prefix}name"),
107            ),
108            (
109                "server_tunnels.created",
110                compact_str::format_compact!("{prefix}created"),
111            ),
112        ])
113    }
114
115    #[inline]
116    fn map(prefix: Option<&str>, row: &PgRow) -> Result<Self, crate::database::DatabaseError> {
117        let prefix = prefix.unwrap_or_default();
118
119        Ok(Self {
120            server: super::server::Server::get_fetchable(
121                row.try_get(compact_str::format_compact!("{prefix}server_uuid").as_str())?,
122            ),
123            idx: row.try_get::<i32, _>(compact_str::format_compact!("{prefix}idx").as_str())?
124                as u16,
125            name: row.try_get(compact_str::format_compact!("{prefix}name").as_str())?,
126            created: row.try_get(compact_str::format_compact!("{prefix}created").as_str())?,
127            extension_data: Self::map_extensions(prefix, row)?,
128        })
129    }
130}
131
132impl ServerTunnel {
133    pub async fn by_server_uuid(
134        database: &crate::database::Database,
135        server_uuid: uuid::Uuid,
136    ) -> Result<Option<Self>, crate::database::DatabaseError> {
137        let row = sqlx::query(sqlx::AssertSqlSafe(format!(
138            r#"
139            SELECT {}
140            FROM server_tunnels
141            WHERE server_tunnels.server_uuid = $1
142            "#,
143            Self::columns_sql(None)
144        )))
145        .bind(server_uuid)
146        .fetch_optional(database.read())
147        .await?;
148
149        row.try_map(|row| Self::map(None, &row))
150    }
151
152    async fn allocate_idx(
153        transaction: &mut sqlx::Transaction<'_, sqlx::Postgres>,
154    ) -> Result<i32, crate::database::DatabaseError> {
155        sqlx::query("SELECT pg_advisory_xact_lock($1)")
156            .bind(ADVISORY_LOCK_IDX)
157            .execute(&mut **transaction)
158            .await?;
159
160        let idx: Option<i32> = sqlx::query_scalar(
161            r#"
162            SELECT candidate.idx
163            FROM generate_series(0, $1) AS candidate(idx)
164            WHERE NOT EXISTS (
165                SELECT 1 FROM server_tunnels WHERE server_tunnels.idx = candidate.idx
166            )
167            ORDER BY candidate.idx
168            LIMIT 1
169            "#,
170        )
171        .bind(i32::from(MAX_SERVER_IDX))
172        .fetch_optional(&mut **transaction)
173        .await?;
174
175        idx.ok_or_else(|| {
176            crate::database::DatabaseError::Any(anyhow::anyhow!(
177                "the tunnel network is full, no frontend index is available"
178            ))
179        })
180    }
181
182    pub async fn ports(
183        &self,
184        database: &crate::database::Database,
185    ) -> Result<Vec<ServerTunnelPort>, crate::database::DatabaseError> {
186        ServerTunnelPort::by_server_uuid(database, self.server.uuid).await
187    }
188
189    pub fn suggest_name(server_name: &str) -> compact_str::CompactString {
190        let mut base = String::with_capacity(server_name.len());
191        for character in server_name.chars() {
192            match character {
193                'a'..='z' | '0'..='9' => base.push(character),
194                'A'..='Z' => base.push(character.to_ascii_lowercase()),
195                _ if !base.ends_with('-') => base.push('-'),
196                _ => {}
197            }
198        }
199
200        let base: String = base.trim_matches('-').chars().take(58).collect();
201        let base = base.trim_matches('-');
202        let base = if base.is_empty() { "server" } else { base };
203
204        if crate::tunnel::is_alias_shaped(base) {
205            compact_str::format_compact!("{base}-1")
206        } else {
207            base.into()
208        }
209    }
210}
211
212#[derive(ToSchema, Serialize, Deserialize)]
213pub struct ServerTunnelPort {
214    pub port: u16,
215    pub protocols: HashSet<TunnelProtocol>,
216
217    pub created: chrono::DateTime<chrono::Utc>,
218}
219
220impl ServerTunnelPort {
221    pub async fn by_server_uuid(
222        database: &crate::database::Database,
223        server_uuid: uuid::Uuid,
224    ) -> Result<Vec<Self>, crate::database::DatabaseError> {
225        sqlx::query(
226            r#"
227            SELECT server_tunnel_ports.port, server_tunnel_ports.protocols, server_tunnel_ports.created
228            FROM server_tunnel_ports
229            WHERE server_tunnel_ports.server_uuid = $1
230            ORDER BY server_tunnel_ports.port
231            "#,
232        )
233        .bind(server_uuid)
234        .fetch_all(database.read())
235        .await?
236        .into_iter()
237        .map(|row| {
238            Ok(Self {
239                port: row.try_get::<i32, _>("port")? as u16,
240                protocols: row
241                    .try_get::<Vec<TunnelProtocol>, _>("protocols")?
242                    .into_iter()
243                    .collect(),
244                created: row.try_get::<chrono::NaiveDateTime, _>("created")?.and_utc(),
245            })
246        })
247        .try_collect_vec()
248    }
249
250    pub async fn by_server_uuids(
251        database: &crate::database::Database,
252        server_uuids: &[uuid::Uuid],
253    ) -> Result<HashMap<uuid::Uuid, Vec<Self>>, crate::database::DatabaseError> {
254        let mut ports: HashMap<uuid::Uuid, Vec<Self>> = HashMap::new();
255
256        for row in sqlx::query(
257            r#"
258            SELECT server_tunnel_ports.server_uuid, server_tunnel_ports.port, server_tunnel_ports.protocols, server_tunnel_ports.created
259            FROM server_tunnel_ports
260            WHERE server_tunnel_ports.server_uuid = ANY($1)
261            ORDER BY server_tunnel_ports.port
262            "#,
263        )
264        .bind(server_uuids)
265        .fetch_all(database.read())
266        .await?
267        {
268            ports
269                .entry(row.try_get("server_uuid")?)
270                .or_default()
271                .push(Self {
272                    port: row.try_get::<i32, _>("port")? as u16,
273                    protocols: row
274                        .try_get::<Vec<TunnelProtocol>, _>("protocols")?
275                        .into_iter()
276                        .collect(),
277                    created: row.try_get::<chrono::NaiveDateTime, _>("created")?.and_utc(),
278                });
279        }
280
281        Ok(ports)
282    }
283
284    pub async fn replace(
285        database: &crate::database::Database,
286        server_uuid: uuid::Uuid,
287        ports: &[CreateServerTunnelPortOptions],
288    ) -> Result<(), crate::database::DatabaseError> {
289        for port in ports {
290            port.validate()?;
291        }
292
293        let mut transaction = database.write().begin().await?;
294
295        sqlx::query("DELETE FROM server_tunnel_ports WHERE server_tunnel_ports.server_uuid = $1")
296            .bind(server_uuid)
297            .execute(&mut *transaction)
298            .await?;
299
300        for port in ports {
301            sqlx::query(
302                r#"
303                INSERT INTO server_tunnel_ports (server_uuid, port, protocols)
304                VALUES ($1, $2, $3)
305                ON CONFLICT (server_uuid, port) DO UPDATE SET protocols = EXCLUDED.protocols
306                "#,
307            )
308            .bind(server_uuid)
309            .bind(i32::from(port.port))
310            .bind(port.protocols.iter().copied().collect::<Vec<_>>())
311            .execute(&mut *transaction)
312            .await?;
313        }
314
315        crate::tunnel::bump_epoch(&mut *transaction).await?;
316        transaction.commit().await?;
317
318        Ok(())
319    }
320}
321
322#[derive(ToSchema, Serialize)]
323#[schema(title = "ServerTunnelPeer")]
324pub struct ServerTunnelPeer {
325    pub server_uuid: uuid::Uuid,
326    pub server_name: compact_str::CompactString,
327
328    pub name: compact_str::CompactString,
329    pub alias: compact_str::CompactString,
330    pub address: Option<compact_str::CompactString>,
331    pub ports: Vec<ServerTunnelPort>,
332
333    pub created: chrono::DateTime<chrono::Utc>,
334}
335
336#[derive(ToSchema, Deserialize, Validate)]
337pub struct CreateServerTunnelPortOptions {
338    #[garde(range(min = 1))]
339    #[schema(minimum = 1)]
340    pub port: u16,
341    #[garde(custom(validate_protocols))]
342    pub protocols: HashSet<TunnelProtocol>,
343}
344
345pub struct ServerTunnelConnection;
346
347impl ServerTunnelConnection {
348    pub async fn by_src_server_uuid(
349        database: &crate::database::Database,
350        src_server_uuid: uuid::Uuid,
351    ) -> Result<Vec<uuid::Uuid>, crate::database::DatabaseError> {
352        Ok(sqlx::query_scalar(
353            r#"
354            SELECT server_tunnel_connections.dst_server_uuid
355            FROM server_tunnel_connections
356            WHERE server_tunnel_connections.src_server_uuid = $1
357            "#,
358        )
359        .bind(src_server_uuid)
360        .fetch_all(database.read())
361        .await?)
362    }
363
364    pub async fn peers(
365        database: &crate::database::Database,
366        server_uuid: uuid::Uuid,
367        incoming: bool,
368    ) -> Result<Vec<ServerTunnelPeer>, crate::database::DatabaseError> {
369        let (own, peer) = if incoming {
370            ("dst_server_uuid", "src_server_uuid")
371        } else {
372            ("src_server_uuid", "dst_server_uuid")
373        };
374
375        let rows = sqlx::query(sqlx::AssertSqlSafe(format!(
376            r#"
377            SELECT
378                server_tunnels.server_uuid,
379                server_tunnels.idx,
380                server_tunnels.name,
381                servers.name AS server_name,
382                servers.uuid_short,
383                server_tunnel_connections.created
384            FROM server_tunnel_connections
385            JOIN server_tunnels ON server_tunnels.server_uuid = server_tunnel_connections.{peer}
386            JOIN servers ON servers.uuid = server_tunnels.server_uuid
387            WHERE server_tunnel_connections.{own} = $1
388            ORDER BY server_tunnels.name
389            "#
390        )))
391        .bind(server_uuid)
392        .fetch_all(database.read())
393        .await?;
394
395        let peer_uuids = rows
396            .iter()
397            .map(|row| row.try_get("server_uuid"))
398            .collect::<Result<Vec<uuid::Uuid>, _>>()?;
399        let mut ports = ServerTunnelPort::by_server_uuids(database, &peer_uuids).await?;
400
401        let mut peers = Vec::with_capacity(rows.len());
402        for (row, server_uuid) in rows.into_iter().zip(peer_uuids) {
403            peers.push(ServerTunnelPeer {
404                server_uuid,
405                server_name: row.try_get("server_name")?,
406                name: row.try_get("name")?,
407                alias: crate::tunnel::alias_of(row.try_get("uuid_short")?),
408                address: crate::tunnel::frontend_address(row.try_get::<i32, _>("idx")? as u16),
409                ports: ports.remove(&server_uuid).unwrap_or_default(),
410                created: row
411                    .try_get::<chrono::NaiveDateTime, _>("created")?
412                    .and_utc(),
413            });
414        }
415
416        Ok(peers)
417    }
418
419    pub async fn colliding_ports(
420        database: &crate::database::Database,
421        src_server_uuid: uuid::Uuid,
422        dst_server_uuid: uuid::Uuid,
423    ) -> Result<Vec<u16>, crate::database::DatabaseError> {
424        let ports: Vec<i32> = sqlx::query_scalar(
425            r#"
426            SELECT DISTINCT node_allocations.port
427            FROM server_allocations
428            JOIN node_allocations ON node_allocations.uuid = server_allocations.allocation_uuid
429            WHERE server_allocations.server_uuid = $1
430                AND node_allocations.port IN (
431                    SELECT server_tunnel_ports.port
432                    FROM server_tunnel_ports
433                    WHERE server_tunnel_ports.server_uuid = $2
434                )
435            ORDER BY node_allocations.port
436            "#,
437        )
438        .bind(src_server_uuid)
439        .bind(dst_server_uuid)
440        .fetch_all(database.read())
441        .await?;
442
443        Ok(ports.into_iter().map(|port| port as u16).collect())
444    }
445
446    pub async fn count_by_src_server_uuid(
447        database: &crate::database::Database,
448        src_server_uuid: uuid::Uuid,
449    ) -> Result<i64, crate::database::DatabaseError> {
450        Ok(sqlx::query_scalar(
451            r#"
452            SELECT COUNT(*)
453            FROM server_tunnel_connections
454            WHERE server_tunnel_connections.src_server_uuid = $1
455            "#,
456        )
457        .bind(src_server_uuid)
458        .fetch_one(database.read())
459        .await?)
460    }
461
462    pub async fn create(
463        database: &crate::database::Database,
464        src_server_uuid: uuid::Uuid,
465        dst_server_uuid: uuid::Uuid,
466    ) -> Result<(), crate::database::DatabaseError> {
467        if src_server_uuid == dst_server_uuid {
468            return Err(crate::database::DatabaseError::Any(anyhow::anyhow!(
469                "a server cannot be connected to itself"
470            )));
471        }
472
473        let mut transaction = database.write().begin().await?;
474
475        let dst_name: compact_str::CompactString = sqlx::query_scalar(
476            r#"
477            SELECT server_tunnels.name
478            FROM server_tunnels
479            WHERE server_tunnels.server_uuid = $1
480            FOR KEY SHARE
481            "#,
482        )
483        .bind(dst_server_uuid)
484        .fetch_one(&mut *transaction)
485        .await?;
486
487        let affected = sqlx::query(
488            r#"
489            INSERT INTO server_tunnel_connections (src_server_uuid, dst_server_uuid, dst_name)
490            VALUES ($1, $2, $3)
491            ON CONFLICT (src_server_uuid, dst_server_uuid) DO NOTHING
492            "#,
493        )
494        .bind(src_server_uuid)
495        .bind(dst_server_uuid)
496        .bind(dst_name.as_str())
497        .execute(&mut *transaction)
498        .await?
499        .rows_affected();
500
501        if affected > 0 {
502            crate::tunnel::bump_epoch(&mut *transaction).await?;
503        }
504
505        transaction.commit().await?;
506
507        Ok(())
508    }
509
510    pub async fn delete(
511        database: &crate::database::Database,
512        src_server_uuid: uuid::Uuid,
513        dst_server_uuid: uuid::Uuid,
514    ) -> Result<bool, crate::database::DatabaseError> {
515        let mut transaction = database.write().begin().await?;
516
517        let affected = sqlx::query(
518            r#"
519            DELETE FROM server_tunnel_connections
520            WHERE server_tunnel_connections.src_server_uuid = $1
521                AND server_tunnel_connections.dst_server_uuid = $2
522            "#,
523        )
524        .bind(src_server_uuid)
525        .bind(dst_server_uuid)
526        .execute(&mut *transaction)
527        .await?
528        .rows_affected();
529
530        if affected == 0 {
531            return Ok(false);
532        }
533
534        crate::tunnel::bump_epoch(&mut *transaction).await?;
535        transaction.commit().await?;
536
537        Ok(true)
538    }
539}
540
541#[async_trait::async_trait]
542impl IntoApiObject for ServerTunnel {
543    type ApiObject = ApiServerTunnel;
544    type ExtraArgs<'a> = i32;
545
546    async fn into_api_object<'a>(
547        self,
548        state: &crate::State,
549        uuid_short: Self::ExtraArgs<'a>,
550    ) -> Result<Self::ApiObject, crate::database::DatabaseError> {
551        let api_object = ApiServerTunnel::init_hooks(&self, state).await?;
552
553        let api_object = finish_extendible!(
554            ApiServerTunnel {
555                name: self.name,
556                alias: crate::tunnel::alias_of(uuid_short),
557                address: crate::tunnel::frontend_address(self.idx),
558                created: self.created.and_utc(),
559            },
560            api_object,
561            state
562        )?;
563
564        Ok(api_object)
565    }
566}
567
568#[schema_extension_derive::extendible]
569#[init_args(ServerTunnel, crate::State)]
570#[hook_args(crate::State)]
571#[derive(ToSchema, Serialize)]
572#[schema(title = "ServerTunnel")]
573pub struct ApiServerTunnel {
574    pub name: compact_str::CompactString,
575    pub alias: compact_str::CompactString,
576    pub address: Option<compact_str::CompactString>,
577
578    pub created: chrono::DateTime<chrono::Utc>,
579}
580
581#[derive(ToSchema, Deserialize, Validate)]
582pub struct CreateServerTunnelOptions {
583    #[garde(skip)]
584    pub server_uuid: uuid::Uuid,
585
586    #[garde(custom(validate_name))]
587    #[schema(min_length = 1, max_length = 63)]
588    pub name: compact_str::CompactString,
589}
590
591#[async_trait::async_trait]
592impl CreatableModel for ServerTunnel {
593    type CreateOptions<'a> = CreateServerTunnelOptions;
594    type CreateResult = Self;
595
596    fn get_create_handlers() -> &'static LazyLock<CreateListenerList<Self>> {
597        static CREATE_LISTENERS: LazyLock<CreateListenerList<ServerTunnel>> =
598            LazyLock::new(|| Arc::new(ModelHandlerList::default()));
599
600        &CREATE_LISTENERS
601    }
602
603    async fn create_with_transaction(
604        state: &crate::State,
605        mut options: Self::CreateOptions<'_>,
606        transaction: &mut sqlx::Transaction<'_, sqlx::Postgres>,
607    ) -> Result<Self, crate::database::DatabaseError> {
608        options.validate()?;
609
610        let idx = Self::allocate_idx(transaction).await?;
611
612        let mut query_builder = InsertQueryBuilder::new("server_tunnels");
613
614        Self::run_create_handlers(&mut options, &mut query_builder, state, transaction).await?;
615
616        query_builder
617            .set("server_uuid", options.server_uuid)
618            .set("idx", idx)
619            .set("name", &options.name);
620
621        let row = query_builder
622            .returning(&Self::columns_sql(None))
623            .fetch_one(&mut **transaction)
624            .await?;
625        let mut server_tunnel = Self::map(None, &row)?;
626
627        crate::tunnel::bump_epoch(&mut **transaction).await?;
628
629        Self::run_after_create_handlers(&mut server_tunnel, &options, state, transaction).await?;
630
631        Ok(server_tunnel)
632    }
633}
634
635#[derive(ToSchema, Serialize, Deserialize, Validate, Default)]
636pub struct UpdateServerTunnelOptions {
637    #[garde(custom(validate_optional_name))]
638    #[schema(min_length = 1, max_length = 63)]
639    pub name: Option<compact_str::CompactString>,
640}
641
642#[async_trait::async_trait]
643impl UpdatableModel for ServerTunnel {
644    type UpdateOptions = UpdateServerTunnelOptions;
645
646    fn get_update_handlers() -> &'static LazyLock<UpdateHandlerList<Self>> {
647        static UPDATE_LISTENERS: LazyLock<UpdateHandlerList<ServerTunnel>> =
648            LazyLock::new(|| Arc::new(ModelHandlerList::default()));
649
650        &UPDATE_LISTENERS
651    }
652
653    async fn update_with_transaction(
654        &mut self,
655        state: &crate::State,
656        mut options: Self::UpdateOptions,
657        transaction: &mut sqlx::Transaction<'_, sqlx::Postgres>,
658    ) -> Result<(), crate::database::DatabaseError> {
659        options.validate()?;
660
661        let mut query_builder = UpdateQueryBuilder::new("server_tunnels");
662
663        self.run_update_handlers(&mut options, &mut query_builder, state, transaction)
664            .await?;
665
666        query_builder
667            .set("name", options.name.as_ref())
668            .where_eq("server_uuid", self.server.uuid);
669
670        query_builder.execute(&mut **transaction).await?;
671
672        crate::tunnel::bump_epoch(&mut **transaction).await?;
673
674        if let Some(name) = options.name {
675            self.name = name;
676        }
677
678        self.run_after_update_handlers(state, transaction).await?;
679
680        Ok(())
681    }
682}
683
684#[async_trait::async_trait]
685impl DeletableModel for ServerTunnel {
686    type DeleteOptions = ();
687
688    fn get_delete_handlers() -> &'static LazyLock<DeleteHandlerList<Self>> {
689        static DELETE_LISTENERS: LazyLock<DeleteHandlerList<ServerTunnel>> =
690            LazyLock::new(|| Arc::new(ModelHandlerList::default()));
691
692        &DELETE_LISTENERS
693    }
694
695    async fn delete_with_transaction(
696        &self,
697        state: &crate::State,
698        options: Self::DeleteOptions,
699        transaction: &mut sqlx::Transaction<'_, sqlx::Postgres>,
700    ) -> Result<(), anyhow::Error> {
701        self.run_delete_handlers(&options, state, transaction)
702            .await?;
703
704        sqlx::query(
705            r#"
706            DELETE FROM server_tunnels
707            WHERE server_tunnels.server_uuid = $1
708            "#,
709        )
710        .bind(self.server.uuid)
711        .execute(&mut **transaction)
712        .await?;
713
714        crate::tunnel::bump_epoch(&mut **transaction).await?;
715
716        self.run_after_delete_handlers(&options, state, transaction)
717            .await?;
718
719        Ok(())
720    }
721}