diff --git a/signal-gateway-bin/src/main.rs b/signal-gateway-bin/src/main.rs index f658f2a..1e422f8 100644 --- a/signal-gateway-bin/src/main.rs +++ b/signal-gateway-bin/src/main.rs @@ -244,7 +244,7 @@ limits = [ Duration::from_secs(10) ); assert_eq!(config.gateway.admin_signal_uuids.len(), 2); - assert!(config.gateway.admin_signal_uuids.contains("abc-123-uuid")); + assert!(config.gateway.admin_signal_uuids.get("abc-123-uuid").is_some()); let syslog = config.syslog.expect("syslog should be present"); assert_eq!(syslog.listen_addr, "0.0.0.0:1514".parse().unwrap()); diff --git a/signal-gateway/src/gateway/mod.rs b/signal-gateway/src/gateway/mod.rs index cf9f9b8..746292a 100644 --- a/signal-gateway/src/gateway/mod.rs +++ b/signal-gateway/src/gateway/mod.rs @@ -468,7 +468,7 @@ impl Gateway { let from_group = data_message.group_info.as_ref().map(|g| g.group_id.clone()); // Check if sender is an admin - if !self.config.admin_signal_uuids.contains(&msg.envelope.source_uuid) { + if !self.config.admin_signal_uuids.is_trusted(&msg.envelope) { warn!("Ignoring message from non-admin: {msg:?}"); continue; } diff --git a/signal-gateway/src/gateway/signal_trust_set.rs b/signal-gateway/src/gateway/signal_trust_set.rs index f325b4c..7844ac6 100644 --- a/signal-gateway/src/gateway/signal_trust_set.rs +++ b/signal-gateway/src/gateway/signal_trust_set.rs @@ -4,6 +4,7 @@ //! - Map: `{"uuid1": ["safety1", "safety2"], "uuid2": []}` - UUIDs with safety numbers //! - Sequence: `["uuid1", "uuid2"]` - UUIDs with no safety numbers (simpler) +use crate::signal_jsonrpc::Envelope; use serde::de::{MapAccess, SeqAccess, Visitor}; use serde::{Deserialize, Deserializer}; use std::collections::HashMap; @@ -25,9 +26,12 @@ impl SignalTrustSet { Self::default() } - /// Check if a UUID is a registered admin. - pub fn contains(&self, uuid: &str) -> bool { - self.map.contains_key(uuid) + /// Check if the sender of an envelope is trusted. + /// + /// Currently checks if the source UUID is in the trust set. + /// Will eventually also verify safety numbers. + pub fn is_trusted(&self, envelope: &Envelope) -> bool { + self.map.contains_key(&envelope.source_uuid) } /// Get all UUIDs as an iterator. @@ -129,28 +133,28 @@ mod tests { #[test] fn test_deserialize_map() { let json = r#"{"uuid1": ["safety1", "safety2"], "uuid2": []}"#; - let uuids: SignalTrustSet = serde_json::from_str(json).unwrap(); + let trust_set: SignalTrustSet = serde_json::from_str(json).unwrap(); - assert_eq!(uuids.len(), 2); - assert!(uuids.contains("uuid1")); - assert!(uuids.contains("uuid2")); - assert_eq!(uuids.get("uuid1").unwrap(), &vec!["safety1", "safety2"]); - assert_eq!(uuids.get("uuid2").unwrap(), &Vec::::new()); + assert_eq!(trust_set.len(), 2); + assert!(trust_set.get("uuid1").is_some()); + assert!(trust_set.get("uuid2").is_some()); + assert_eq!(trust_set.get("uuid1").unwrap(), &vec!["safety1", "safety2"]); + assert_eq!(trust_set.get("uuid2").unwrap(), &Vec::::new()); } #[test] fn test_deserialize_seq() { let json = r#"["uuid1", "uuid2", "uuid3"]"#; - let uuids: SignalTrustSet = serde_json::from_str(json).unwrap(); + let trust_set: SignalTrustSet = serde_json::from_str(json).unwrap(); - assert_eq!(uuids.len(), 3); - assert!(uuids.contains("uuid1")); - assert!(uuids.contains("uuid2")); - assert!(uuids.contains("uuid3")); + assert_eq!(trust_set.len(), 3); + assert!(trust_set.get("uuid1").is_some()); + assert!(trust_set.get("uuid2").is_some()); + assert!(trust_set.get("uuid3").is_some()); // All should have empty safety numbers - assert!(uuids.get("uuid1").unwrap().is_empty()); - assert!(uuids.get("uuid2").unwrap().is_empty()); - assert!(uuids.get("uuid3").unwrap().is_empty()); + assert!(trust_set.get("uuid1").unwrap().is_empty()); + assert!(trust_set.get("uuid2").unwrap().is_empty()); + assert!(trust_set.get("uuid3").unwrap().is_empty()); } #[test]