diff --git a/Cargo.lock b/Cargo.lock index 83b93ccb..4b14055d 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1970,6 +1970,7 @@ dependencies = [ "non-empty-string", "rand 0.9.5", "rmp-serde", + "secrecy", "semver", "serde", "serde_json", diff --git a/components/spider-client/src/client.rs b/components/spider-client/src/client.rs index 2a66ccb0..ea1bc601 100644 --- a/components/spider-client/src/client.rs +++ b/components/spider-client/src/client.rs @@ -9,6 +9,7 @@ use spider_core::types::id::JobId; use spider_core::types::id::ResourceGroupId; use spider_core::types::io::TaskInput; use spider_core::types::io::TaskOutput; +use spider_core::types::resource_group::ExternalResourceGroupCredentials; use spider_utils::grpc::retry::RetryConfig; use tonic::transport::Endpoint; @@ -181,12 +182,9 @@ impl SpiderClient { /// * [`ClientError::Server`] for any other server-reported error. pub async fn add_resource_group( &self, - external_resource_group_id: String, - password: Vec, + credentials: ExternalResourceGroupCredentials, ) -> Result { - self.resource_group - .add_resource_group(external_resource_group_id, password) - .await + self.resource_group.add_resource_group(credentials).await } /// Verifies a resource group's password. @@ -308,6 +306,7 @@ fn assert_client_futures_send( resource_group_id: ResourceGroupId, job_id: JobId, task_graph: &TaskGraph, + credentials: ExternalResourceGroupCredentials, ) { const fn assert_send(_: &FutureType) {} assert_send(&client.submit_job(resource_group_id, task_graph, Vec::new())); @@ -316,6 +315,6 @@ fn assert_client_futures_send( assert_send(&client.get_job_state(job_id)); assert_send(&client.get_job_outputs(job_id)); assert_send(&client.get_job_error(job_id)); - assert_send(&client.add_resource_group(String::new(), Vec::new())); + assert_send(&client.add_resource_group(credentials)); assert_send(&client.verify_resource_group(resource_group_id, Vec::new())); } diff --git a/components/spider-client/src/grpc/resource_group.rs b/components/spider-client/src/grpc/resource_group.rs index 7159e676..934e2f3f 100644 --- a/components/spider-client/src/grpc/resource_group.rs +++ b/components/spider-client/src/grpc/resource_group.rs @@ -3,6 +3,7 @@ use std::num::NonZeroUsize; use spider_core::types::id::ResourceGroupId; +use spider_core::types::resource_group::ExternalResourceGroupCredentials; use spider_proto_rust::storage::ResourceGroupManagementServiceClient; use spider_proto_rust::storage::{self}; use spider_utils::grpc::client::ConnectionPool; @@ -66,15 +67,13 @@ impl ResourceGroupManagementClient { /// * Forwards [`ResourceGroupManagementServiceClient::add_resource_group`]'s status on failure. pub async fn add_resource_group( &self, - external_resource_group_id: String, - password: Vec, + credentials: ExternalResourceGroupCredentials, ) -> Result { let pool = self.connection_pool.clone(); let response = call_with_retry(self.retry_config, move || { let mut client = pool.get_client(); let request = storage::AddResourceGroupRequest { - external_resource_group_id: external_resource_group_id.clone(), - password: password.clone(), + credentials: Some((&credentials).into()), }; async move { client.add_resource_group(request).await } }) diff --git a/components/spider-core/Cargo.toml b/components/spider-core/Cargo.toml index 012eba27..4f133302 100644 --- a/components/spider-core/Cargo.toml +++ b/components/spider-core/Cargo.toml @@ -11,6 +11,7 @@ path = "src/lib.rs" non-empty-string = { workspace = true } rand = { workspace = true } rmp-serde = { workspace = true } +secrecy = { workspace = true } semver = { workspace = true } serde = { workspace = true } serde_json = { workspace = true } diff --git a/components/spider-core/src/types/mod.rs b/components/spider-core/src/types/mod.rs index d6f3baa5..0945e1bb 100644 --- a/components/spider-core/src/types/mod.rs +++ b/components/spider-core/src/types/mod.rs @@ -1,3 +1,4 @@ pub mod id; pub mod io; +pub mod resource_group; pub mod scheduler; diff --git a/components/spider-core/src/types/resource_group.rs b/components/spider-core/src/types/resource_group.rs new file mode 100644 index 00000000..d2f98d16 --- /dev/null +++ b/components/spider-core/src/types/resource_group.rs @@ -0,0 +1,87 @@ +//! Resource group types shared across Spider components. + +use secrecy::ExposeSecret; +use secrecy::SecretSlice; + +/// Environment variable that supplies the external resource group ID. +pub const EXTERNAL_RESOURCE_GROUP_ID_ENV: &str = "SPIDER_EXTERNAL_RESOURCE_GROUP_ID"; + +/// Environment variable that supplies the external resource group password. +pub const EXTERNAL_RESOURCE_GROUP_PASSWORD_ENV: &str = "SPIDER_EXTERNAL_RESOURCE_GROUP_PASSWORD"; + +/// Credentials identifying and authenticating an external resource group. +#[derive(Debug, Clone)] +pub struct ExternalResourceGroupCredentials { + /// The external resource group ID. + external_resource_group_id: String, + + /// The resource group password. + password: SecretSlice, +} + +impl ExternalResourceGroupCredentials { + /// Factory function. + /// + /// # Returns + /// + /// The newly created external resource group credentials. + #[must_use] + pub fn new(external_resource_group_id: String, password: Vec) -> Self { + Self { + external_resource_group_id, + password: password.into(), + } + } + + /// Factory function. + /// + /// Reads the external resource group credentials from the [`EXTERNAL_RESOURCE_GROUP_ID_ENV`] + /// and [`EXTERNAL_RESOURCE_GROUP_PASSWORD_ENV`] environment variables. + /// + /// # Returns + /// + /// The credentials on success. + /// + /// # Errors + /// + /// Returns an error if: + /// + /// * [`ExternalResourceGroupCredentialsError::MissingEnvVar`] if either environment variable is + /// unset. + pub fn from_env() -> Result { + let external_resource_group_id = + std::env::var(EXTERNAL_RESOURCE_GROUP_ID_ENV).map_err(|_| { + ExternalResourceGroupCredentialsError::MissingEnvVar(EXTERNAL_RESOURCE_GROUP_ID_ENV) + })?; + let password = std::env::var(EXTERNAL_RESOURCE_GROUP_PASSWORD_ENV).map_err(|_| { + ExternalResourceGroupCredentialsError::MissingEnvVar( + EXTERNAL_RESOURCE_GROUP_PASSWORD_ENV, + ) + })?; + Ok(Self::new(external_resource_group_id, password.into_bytes())) + } + + /// # Returns + /// + /// The external resource group ID. + #[must_use] + pub fn get_external_resource_group_id(&self) -> &str { + &self.external_resource_group_id + } + + /// # Returns + /// + /// The resource group password. + #[must_use] + pub fn get_password(&self) -> &[u8] { + self.password.expose_secret() + } +} + +/// An error returned while reading external resource group credentials from the environment. +#[derive(Debug, thiserror::Error)] +pub enum ExternalResourceGroupCredentialsError { + /// A required environment variable is unavailable. + #[error("required environment variable `{0}` is not set")] + MissingEnvVar(&'static str), +} diff --git a/components/spider-proto-rust/src/generated/storage.rs b/components/spider-proto-rust/src/generated/storage.rs index 41f50f6e..81a82de2 100644 --- a/components/spider-proto-rust/src/generated/storage.rs +++ b/components/spider-proto-rust/src/generated/storage.rs @@ -134,10 +134,8 @@ pub struct ReportTaskFailureRequest { } #[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)] pub struct AddResourceGroupRequest { - #[prost(string, tag = "1")] - pub external_resource_group_id: ::prost::alloc::string::String, - #[prost(bytes = "vec", tag = "2")] - pub password: ::prost::alloc::vec::Vec, + #[prost(message, optional, tag = "1")] + pub credentials: ::core::option::Option, } #[derive(Clone, Copy, PartialEq, Eq, Hash, ::prost::Message)] pub struct ResourceGroupIdResponse { diff --git a/components/spider-proto-rust/src/lib.rs b/components/spider-proto-rust/src/lib.rs index 1383ef4c..3d82e9a1 100644 --- a/components/spider-proto-rust/src/lib.rs +++ b/components/spider-proto-rust/src/lib.rs @@ -5,6 +5,7 @@ pub mod error; pub mod id; pub mod io; pub mod job; +pub mod resource_group; pub mod scheduler_registration; pub mod unpack; diff --git a/components/spider-proto-rust/src/resource_group.rs b/components/spider-proto-rust/src/resource_group.rs new file mode 100644 index 00000000..5098204c --- /dev/null +++ b/components/spider-proto-rust/src/resource_group.rs @@ -0,0 +1,48 @@ +//! Conversions between protobuf and Spider core resource group types. + +use spider_core::types::resource_group::ExternalResourceGroupCredentials; + +use crate::storage; + +impl From<&ExternalResourceGroupCredentials> for storage::ExternalResourceGroupCredentials { + fn from(credentials: &ExternalResourceGroupCredentials) -> Self { + Self { + external_resource_group_id: credentials.get_external_resource_group_id().to_owned(), + password: credentials.get_password().to_vec(), + } + } +} + +impl From for ExternalResourceGroupCredentials { + fn from(credentials: storage::ExternalResourceGroupCredentials) -> Self { + Self::new(credentials.external_resource_group_id, credentials.password) + } +} + +#[cfg(test)] +mod tests { + use prost::Message; + use spider_core::types::resource_group::ExternalResourceGroupCredentials; + + use crate::storage; + + #[test] + fn test_external_resource_group_credentials_protocol_round_trip() { + let credentials = ExternalResourceGroupCredentials::new( + "external-resource-group".to_owned(), + vec![0, 1, 2, 255], + ); + + let encoded = storage::ExternalResourceGroupCredentials::from(&credentials).encode_to_vec(); + let decoded = ExternalResourceGroupCredentials::from( + storage::ExternalResourceGroupCredentials::decode(encoded.as_slice()) + .expect("external resource group credentials should decode"), + ); + + assert_eq!( + decoded.get_external_resource_group_id(), + credentials.get_external_resource_group_id() + ); + assert_eq!(decoded.get_password(), credentials.get_password()); + } +} diff --git a/components/spider-proto-rust/src/unpack/storage.rs b/components/spider-proto-rust/src/unpack/storage.rs index dcea51c5..bbd73610 100644 --- a/components/spider-proto-rust/src/unpack/storage.rs +++ b/components/spider-proto-rust/src/unpack/storage.rs @@ -9,6 +9,7 @@ use spider_core::types::id::ResourceGroupId; use spider_core::types::id::SessionId; use spider_core::types::id::TaskId; use spider_core::types::id::TaskInstanceId; +use spider_core::types::resource_group::ExternalResourceGroupCredentials; use spider_utils::config::Host; use tonic::Code; @@ -142,15 +143,14 @@ impl RequestUnpack for ReportTaskFailureRequest { } } -/// Unpacks [`AddResourceGroupRequest`] into a tuple containing: -/// -/// * The external resource group ID. -/// * The password. +/// Unpacks [`AddResourceGroupRequest`] into external resource group credentials. impl RequestUnpack for AddResourceGroupRequest { - type Unpacked = (String, Vec); + type Unpacked = ExternalResourceGroupCredentials; fn unpack(self) -> Result { - Ok((self.external_resource_group_id, self.password)) + self.credentials + .map(Into::into) + .ok_or_else(|| invalid_argument("resource group credentials are missing".to_owned())) } } @@ -171,16 +171,14 @@ impl RequestUnpack for VerifyResourceGroupRequest { /// * The execution manager's IP address. /// * The external resource group credentials, if present. impl RequestUnpack for RegisterExecutionManagerRequest { - type Unpacked = (IpAddr, Option<(String, Vec)>); + type Unpacked = (IpAddr, Option); fn unpack(self) -> Result { let ip_address = self .ip_address .parse::() .map_err(|error| invalid_argument(format!("invalid IP address: {error}")))?; - let resource_group_credentials = self - .resource_group_credentials - .map(|credentials| (credentials.external_resource_group_id, credentials.password)); + let resource_group_credentials = self.resource_group_credentials.map(Into::into); Ok((ip_address, resource_group_credentials)) } } diff --git a/components/spider-proto/storage/storage.proto b/components/spider-proto/storage/storage.proto index b74ff4e5..012fe3f2 100644 --- a/components/spider-proto/storage/storage.proto +++ b/components/spider-proto/storage/storage.proto @@ -139,8 +139,7 @@ message ReportTaskFailureRequest { } message AddResourceGroupRequest { - string external_resource_group_id = 1; - bytes password = 2; + ExternalResourceGroupCredentials credentials = 1; } message ResourceGroupIdResponse { diff --git a/components/spider-storage/src/db/mariadb.rs b/components/spider-storage/src/db/mariadb.rs index 77810c5a..c4b6dd32 100644 --- a/components/spider-storage/src/db/mariadb.rs +++ b/components/spider-storage/src/db/mariadb.rs @@ -11,6 +11,7 @@ use spider_core::types::id::SchedulerId; use spider_core::types::id::SessionId; use spider_core::types::io::SerializedTaskOutputs; use spider_core::types::io::TaskOutput; +use spider_core::types::resource_group::ExternalResourceGroupCredentials; use spider_core::types::scheduler::RegisteredScheduler; use spider_derive::MySqlEnum; use spider_utils::config::Host; @@ -23,7 +24,6 @@ use crate::db::DbError; use crate::db::DbStorage; use crate::db::ExecutionManagerLivenessManagement; use crate::db::ExternalJobOrchestration; -use crate::db::ExternalResourceGroupCredentials; use crate::db::InternalJobOrchestration; use crate::db::RecoverableJobContext; use crate::db::ResourceGroupManagement; @@ -439,13 +439,9 @@ impl ResourceGroupManagement for MariaDbStorageConnector { table = RESOURCE_GROUPS_TABLE_NAME, ); - let ExternalResourceGroupCredentials { - external_resource_group_id, - password, - } = credentials; let resource_group_id = sqlx::query(QUERY) - .bind(&external_resource_group_id) - .bind(password) + .bind(credentials.get_external_resource_group_id()) + .bind(credentials.get_password()) .execute(&self.pool) .await .map_err(|e| match e { @@ -453,7 +449,9 @@ impl ResourceGroupManagement for MariaDbStorageConnector { if e.try_downcast_ref::() .is_some_and(|mysql_err| mysql_err.number() == MYSQL_ER_DUP_ENTRY) => { - DbError::ResourceGroupAlreadyExists(external_resource_group_id) + DbError::ResourceGroupAlreadyExists( + credentials.get_external_resource_group_id().to_owned(), + ) } e => e.into(), })? @@ -516,20 +514,19 @@ impl ExecutionManagerLivenessManagement for MariaDbStorageConnector { table = RESOURCE_GROUPS_TABLE_NAME, ); - let resource_group_id = if let Some(ExternalResourceGroupCredentials { - external_resource_group_id, - password, - }) = resource_group_credentials - { + let resource_group_id = if let Some(credentials) = resource_group_credentials { let resource_group_id = sqlx::query_scalar::<_, ResourceGroupId>(SELECT_RESOURCE_GROUP_QUERY) - .bind(&external_resource_group_id) + .bind(credentials.get_external_resource_group_id()) .fetch_optional(&self.pool) .await? .ok_or_else(|| { - DbError::ExternalResourceGroupNotFound(external_resource_group_id) + DbError::ExternalResourceGroupNotFound( + credentials.get_external_resource_group_id().to_owned(), + ) })?; - ResourceGroupManagement::verify(self, resource_group_id, &password).await?; + ResourceGroupManagement::verify(self, resource_group_id, credentials.get_password()) + .await?; Some(resource_group_id) } else { None diff --git a/components/spider-storage/src/db/mod.rs b/components/spider-storage/src/db/mod.rs index b65a0ddc..62bff60a 100644 --- a/components/spider-storage/src/db/mod.rs +++ b/components/spider-storage/src/db/mod.rs @@ -7,9 +7,9 @@ pub use mariadb::MariaDbStorageConnector; pub use protocol::DbStorage; pub use protocol::ExecutionManagerLivenessManagement; pub use protocol::ExternalJobOrchestration; -pub use protocol::ExternalResourceGroupCredentials; pub use protocol::InternalJobOrchestration; pub use protocol::RecoverableJobContext; pub use protocol::ResourceGroupManagement; pub use protocol::SchedulerRegistrationManagement; pub use protocol::SessionManagement; +pub use spider_core::types::resource_group::ExternalResourceGroupCredentials; diff --git a/components/spider-storage/src/db/protocol.rs b/components/spider-storage/src/db/protocol.rs index 02cd4698..8babe5ec 100644 --- a/components/spider-storage/src/db/protocol.rs +++ b/components/spider-storage/src/db/protocol.rs @@ -8,6 +8,7 @@ use spider_core::types::id::ResourceGroupId; use spider_core::types::id::SchedulerId; use spider_core::types::id::SessionId; use spider_core::types::io::TaskOutput; +use spider_core::types::resource_group::ExternalResourceGroupCredentials; use spider_core::types::scheduler::RegisteredScheduler; use crate::db::error::DbError; @@ -24,15 +25,6 @@ pub struct RecoverableJobContext { pub outputs: Option>, } -/// Credentials identifying and authenticating an external resource group. -pub struct ExternalResourceGroupCredentials { - /// The external resource group ID. - pub external_resource_group_id: String, - - /// The resource group password. - pub password: Vec, -} - /// The database storage interface. A database storage must implement the following traits: /// /// * [`ExternalJobOrchestration`] diff --git a/components/spider-storage/src/grpc.rs b/components/spider-storage/src/grpc.rs index 5cad7bb4..334a6912 100644 --- a/components/spider-storage/src/grpc.rs +++ b/components/spider-storage/src/grpc.rs @@ -22,7 +22,6 @@ use tonic::Status; use crate::cache::error::CacheError; use crate::db::DbError; use crate::db::DbStorage; -use crate::db::ExternalResourceGroupCredentials; use crate::inbound_queue::InboundQueueEntry; use crate::inbound_queue::InboundQueueSender; use crate::state::ServiceState; @@ -770,14 +769,14 @@ impl< &self, request: Request, ) -> Result, Status> { - let (external_id, password) = request.into_inner().unpack()?; - tracing::info!(external_id = % external_id, "Add resource group request received."); + let credentials = request.into_inner().unpack()?; + tracing::info!( + external_id = % credentials.get_external_resource_group_id(), + "Add resource group request received." + ); let rg_id = self .inner - .add_resource_group(ExternalResourceGroupCredentials { - external_resource_group_id: external_id, - password, - }) + .add_resource_group(credentials) .await .map_err(|error| { self.resource_group_management_service_error_handler(error, "add_resource_group") @@ -822,15 +821,7 @@ impl< tracing::info!(% ip_address, "Execution manager registration request received."); let (em_id, resource_group_id) = self .inner - .register_execution_manager( - ip_address, - resource_group_credentials.map(|(external_resource_group_id, password)| { - ExternalResourceGroupCredentials { - external_resource_group_id, - password, - } - }), - ) + .register_execution_manager(ip_address, resource_group_credentials) .await .map_err(|error| { self.execution_manager_liveness_service_error_handler( diff --git a/components/spider-storage/src/state/service.rs b/components/spider-storage/src/state/service.rs index 460acfd9..bec21db4 100644 --- a/components/spider-storage/src/state/service.rs +++ b/components/spider-storage/src/state/service.rs @@ -14,6 +14,7 @@ use spider_core::types::id::TaskInstanceId; use spider_core::types::io::ExecutionContext; use spider_core::types::io::TaskOutput; use spider_core::types::io::TaskOutputsSerializer; +use spider_core::types::resource_group::ExternalResourceGroupCredentials; use spider_core::types::scheduler::RegisteredScheduler; use spider_tdl::error::TdlError; use spider_utils::config::Host; @@ -23,7 +24,6 @@ use crate::cache::error::CacheError; use crate::cache::error::InternalError; use crate::cache::job::SharedJobControlBlock; use crate::db::DbStorage; -use crate::db::ExternalResourceGroupCredentials; use crate::inbound_queue::CleanupTaskMarker; use crate::inbound_queue::CommitTaskMarker; use crate::inbound_queue::InboundQueueEntry; @@ -1644,10 +1644,10 @@ mod tests { let service = create_test_service(); assert!( service - .add_resource_group(ExternalResourceGroupCredentials { - external_resource_group_id: "external_123".to_owned(), - password: vec![1, 2, 3], - }) + .add_resource_group(ExternalResourceGroupCredentials::new( + "external_123".to_owned(), + vec![1, 2, 3], + )) .await .is_ok() ); @@ -1659,10 +1659,10 @@ mod tests { let service = create_test_service(); let password = vec![1, 2, 3]; let rg_id = service - .add_resource_group(ExternalResourceGroupCredentials { - external_resource_group_id: "external_123".to_owned(), - password: password.clone(), - }) + .add_resource_group(ExternalResourceGroupCredentials::new( + "external_123".to_owned(), + password.clone(), + )) .await?; service.verify_resource_group(rg_id, &password).await?; Ok(()) @@ -1672,10 +1672,10 @@ mod tests { async fn verify_resource_group_fails_for_wrong_password() -> anyhow::Result<()> { let service = create_test_service(); let rg_id = service - .add_resource_group(ExternalResourceGroupCredentials { - external_resource_group_id: "external_123".to_owned(), - password: vec![1, 2, 3], - }) + .add_resource_group(ExternalResourceGroupCredentials::new( + "external_123".to_owned(), + vec![1, 2, 3], + )) .await?; let result = service.verify_resource_group(rg_id, &[4, 5, 6]).await; assert!(result.is_err(), "verify should fail for wrong password"); @@ -1771,19 +1771,19 @@ mod tests { let service = create_test_service(); let password = b"password"; let resource_group_id = service - .add_resource_group(ExternalResourceGroupCredentials { - external_resource_group_id: "external_123".to_owned(), - password: password.to_vec(), - }) + .add_resource_group(ExternalResourceGroupCredentials::new( + "external_123".to_owned(), + password.to_vec(), + )) .await?; let (_, registered_resource_group_id) = service .register_execution_manager( "127.0.0.1".parse()?, - Some(ExternalResourceGroupCredentials { - external_resource_group_id: "external_123".to_owned(), - password: password.to_vec(), - }), + Some(ExternalResourceGroupCredentials::new( + "external_123".to_owned(), + password.to_vec(), + )), ) .await?; @@ -1798,10 +1798,10 @@ mod tests { let result = service .register_execution_manager( "127.0.0.1".parse()?, - Some(ExternalResourceGroupCredentials { - external_resource_group_id: "unknown".to_owned(), - password: b"password".to_vec(), - }), + Some(ExternalResourceGroupCredentials::new( + "unknown".to_owned(), + b"password".to_vec(), + )), ) .await; @@ -1818,19 +1818,19 @@ mod tests { async fn register_execution_manager_rejects_invalid_password() -> anyhow::Result<()> { let service = create_test_service(); service - .add_resource_group(ExternalResourceGroupCredentials { - external_resource_group_id: "external_123".to_owned(), - password: b"password".to_vec(), - }) + .add_resource_group(ExternalResourceGroupCredentials::new( + "external_123".to_owned(), + b"password".to_vec(), + )) .await?; let result = service .register_execution_manager( "127.0.0.1".parse()?, - Some(ExternalResourceGroupCredentials { - external_resource_group_id: "external_123".to_owned(), - password: b"wrong-password".to_vec(), - }), + Some(ExternalResourceGroupCredentials::new( + "external_123".to_owned(), + b"wrong-password".to_vec(), + )), ) .await; diff --git a/components/spider-storage/src/state/test_utils.rs b/components/spider-storage/src/state/test_utils.rs index cd4c11dc..204aee53 100644 --- a/components/spider-storage/src/state/test_utils.rs +++ b/components/spider-storage/src/state/test_utils.rs @@ -12,6 +12,7 @@ use spider_core::types::id::SchedulerId; use spider_core::types::id::SessionId; use spider_core::types::id::TaskInstanceId; use spider_core::types::io::TaskOutput; +use spider_core::types::resource_group::ExternalResourceGroupCredentials; use spider_core::types::scheduler::RegisteredScheduler; use crate::cache::error::InternalError; @@ -21,7 +22,6 @@ use crate::db::DbError; use crate::db::DbStorage; use crate::db::ExecutionManagerLivenessManagement; use crate::db::ExternalJobOrchestration; -use crate::db::ExternalResourceGroupCredentials; use crate::db::InternalJobOrchestration; use crate::db::RecoverableJobContext; use crate::db::ResourceGroupManagement; @@ -70,7 +70,7 @@ pub struct MockDbConnector { pub states: Arc>, pub errors: Arc>, pub outputs: Arc>>, - pub resource_groups: Arc>>, + pub resource_groups: Arc>, pub resource_group_ids: Arc>, pub next_resource_group_id: Arc, pub execution_managers: Arc)>>, @@ -180,13 +180,10 @@ impl ResourceGroupManagement for MockDbConnector { &self, credentials: ExternalResourceGroupCredentials, ) -> Result { - let ExternalResourceGroupCredentials { - external_resource_group_id, - password, - } = credentials; let counter = self.next_resource_group_id.fetch_add(1, Ordering::Relaxed); let id = ResourceGroupId::from(counter as u64); - self.resource_groups.insert(id, password); + let external_resource_group_id = credentials.get_external_resource_group_id().to_owned(); + self.resource_groups.insert(id, credentials); self.resource_group_ids .insert(external_resource_group_id, id); Ok(id) @@ -201,7 +198,7 @@ impl ResourceGroupManagement for MockDbConnector { .resource_groups .get(&resource_group_id) .ok_or(DbError::ResourceGroupNotFound(resource_group_id))?; - let matches = stored.as_slice() == password; + let matches = stored.get_password() == password; drop(stored); if !matches { return Err(DbError::InvalidPassword(resource_group_id)); @@ -227,18 +224,22 @@ impl ExecutionManagerLivenessManagement for MockDbConnector { resource_group_credentials: Option, ) -> Result<(ExecutionManagerId, Option), DbError> { let resource_group_id = match resource_group_credentials { - Some(ExternalResourceGroupCredentials { - external_resource_group_id, - password, - }) => { + Some(credentials) => { let resource_group_id = self .resource_group_ids - .get(&external_resource_group_id) + .get(credentials.get_external_resource_group_id()) .map(|entry| *entry.value()) .ok_or_else(|| { - DbError::ExternalResourceGroupNotFound(external_resource_group_id) + DbError::ExternalResourceGroupNotFound( + credentials.get_external_resource_group_id().to_owned(), + ) })?; - ResourceGroupManagement::verify(self, resource_group_id, &password).await?; + ResourceGroupManagement::verify( + self, + resource_group_id, + credentials.get_password(), + ) + .await?; Some(resource_group_id) } None => None, diff --git a/components/spider-storage/src/task_instance_pool.rs b/components/spider-storage/src/task_instance_pool.rs index d4c40433..e732b166 100644 --- a/components/spider-storage/src/task_instance_pool.rs +++ b/components/spider-storage/src/task_instance_pool.rs @@ -560,11 +560,11 @@ mod tests { use spider_core::types::id::JobId; use spider_core::types::id::ResourceGroupId; use spider_core::types::io::TaskInput; + use spider_core::types::resource_group::ExternalResourceGroupCredentials; use tokio::sync::Mutex; use super::*; use crate::db::DbError; - use crate::db::ExternalResourceGroupCredentials; use crate::job_submission::create_validated_submission; const DEFAULT_CHANNEL_SIZE: usize = 128; diff --git a/components/spider-storage/tests/mariadb_infra.rs b/components/spider-storage/tests/mariadb_infra.rs index 46791dc0..3a663dc4 100644 --- a/components/spider-storage/tests/mariadb_infra.rs +++ b/components/spider-storage/tests/mariadb_infra.rs @@ -1,7 +1,7 @@ use spider_core::types::id::ResourceGroupId; +use spider_core::types::resource_group::ExternalResourceGroupCredentials; use spider_storage::DatabaseConfig; use spider_storage::DatabaseCredentials; -use spider_storage::db::ExternalResourceGroupCredentials; use spider_storage::db::MariaDbStorageConnector; use spider_storage::db::ResourceGroupManagement; @@ -64,10 +64,10 @@ pub fn create_mariadb_config() -> DatabaseConfig { pub async fn create_test_resource_group(storage: &MariaDbStorageConnector) -> ResourceGroupId { let external_id = format!("test-resource-group-{}", rand::random::()); storage - .add(ExternalResourceGroupCredentials { - external_resource_group_id: external_id, - password: b"test-password".to_vec(), - }) + .add(ExternalResourceGroupCredentials::new( + external_id, + b"test-password".to_vec(), + )) .await .expect("add should succeed") } diff --git a/components/spider-storage/tests/mariadb_test.rs b/components/spider-storage/tests/mariadb_test.rs index c5d436b9..292f488a 100644 --- a/components/spider-storage/tests/mariadb_test.rs +++ b/components/spider-storage/tests/mariadb_test.rs @@ -8,10 +8,10 @@ use spider_core::types::id::JobId; use spider_core::types::id::ResourceGroupId; use spider_core::types::id::SchedulerId; use spider_core::types::io::TaskInput; +use spider_core::types::resource_group::ExternalResourceGroupCredentials; use spider_storage::db::DbError; use spider_storage::db::ExecutionManagerLivenessManagement; use spider_storage::db::ExternalJobOrchestration; -use spider_storage::db::ExternalResourceGroupCredentials; use spider_storage::db::InternalJobOrchestration; use spider_storage::db::MariaDbStorageConnector; use spider_storage::db::ResourceGroupManagement; @@ -553,18 +553,18 @@ async fn test_add_duplicate_resource_group() { let external_id = format!("test-resource-group-{}", rand::random::()); storage - .add(ExternalResourceGroupCredentials { - external_resource_group_id: external_id.clone(), - password: b"password".to_vec(), - }) + .add(ExternalResourceGroupCredentials::new( + external_id.clone(), + b"password".to_vec(), + )) .await .expect("first add should succeed"); let result = storage - .add(ExternalResourceGroupCredentials { - external_resource_group_id: external_id, - password: b"password".to_vec(), - }) + .add(ExternalResourceGroupCredentials::new( + external_id, + b"password".to_vec(), + )) .await; assert!( matches!(result, Err(DbError::ResourceGroupAlreadyExists(_))), @@ -578,10 +578,10 @@ async fn test_verify_correct_password() { let storage = create_mariadb_connector().await; let rg_id = storage - .add(ExternalResourceGroupCredentials { - external_resource_group_id: format!("test-resource-group-{}", rand::random::()), - password: b"correct-password".to_vec(), - }) + .add(ExternalResourceGroupCredentials::new( + format!("test-resource-group-{}", rand::random::()), + b"correct-password".to_vec(), + )) .await .expect("add should succeed"); @@ -597,10 +597,10 @@ async fn test_verify_wrong_password() { let storage = create_mariadb_connector().await; let rg_id = storage - .add(ExternalResourceGroupCredentials { - external_resource_group_id: format!("test-resource-group-{}", rand::random::()), - password: b"correct-password".to_vec(), - }) + .add(ExternalResourceGroupCredentials::new( + format!("test-resource-group-{}", rand::random::()), + b"correct-password".to_vec(), + )) .await .expect("add should succeed"); @@ -815,20 +815,20 @@ async fn test_register_execution_manager_with_resource_group() { let external_resource_group_id = format!("test-resource-group-{}", rand::random::()); let password = b"password"; let resource_group_id = storage - .add(ExternalResourceGroupCredentials { - external_resource_group_id: external_resource_group_id.clone(), - password: password.to_vec(), - }) + .add(ExternalResourceGroupCredentials::new( + external_resource_group_id.clone(), + password.to_vec(), + )) .await .expect("add should succeed"); let (_, registered_resource_group_id) = storage .register_execution_manager( IpAddr::V4(Ipv4Addr::LOCALHOST), - Some(ExternalResourceGroupCredentials { + Some(ExternalResourceGroupCredentials::new( external_resource_group_id, - password: password.to_vec(), - }), + password.to_vec(), + )), ) .await .expect("register_execution_manager should succeed"); @@ -845,10 +845,10 @@ async fn test_register_execution_manager_with_unknown_resource_group() { let result = storage .register_execution_manager( IpAddr::V4(Ipv4Addr::LOCALHOST), - Some(ExternalResourceGroupCredentials { + Some(ExternalResourceGroupCredentials::new( external_resource_group_id, - password: b"password".to_vec(), - }), + b"password".to_vec(), + )), ) .await; @@ -864,20 +864,20 @@ async fn test_register_execution_manager_with_invalid_password() { let storage = create_mariadb_connector().await; let external_resource_group_id = format!("test-resource-group-{}", rand::random::()); let resource_group_id = storage - .add(ExternalResourceGroupCredentials { - external_resource_group_id: external_resource_group_id.clone(), - password: b"password".to_vec(), - }) + .add(ExternalResourceGroupCredentials::new( + external_resource_group_id.clone(), + b"password".to_vec(), + )) .await .expect("add should succeed"); let result = storage .register_execution_manager( IpAddr::V4(Ipv4Addr::LOCALHOST), - Some(ExternalResourceGroupCredentials { + Some(ExternalResourceGroupCredentials::new( external_resource_group_id, - password: b"wrong-password".to_vec(), - }), + b"wrong-password".to_vec(), + )), ) .await; diff --git a/components/spider-storage/tests/runtime_recovery_test.rs b/components/spider-storage/tests/runtime_recovery_test.rs index 97198d74..359f0649 100644 --- a/components/spider-storage/tests/runtime_recovery_test.rs +++ b/components/spider-storage/tests/runtime_recovery_test.rs @@ -8,10 +8,10 @@ use spider_core::types::id::JobId; use spider_core::types::id::TaskInstanceId; use spider_core::types::io::TaskOutput; use spider_core::types::io::TaskOutputsSerializer; +use spider_core::types::resource_group::ExternalResourceGroupCredentials; use spider_storage::cache::error::CacheError; use spider_storage::cache::error::StaleStateError; use spider_storage::db::ExternalJobOrchestration; -use spider_storage::db::ExternalResourceGroupCredentials; use spider_storage::inbound_queue::CleanupTaskMarker; use spider_storage::inbound_queue::CommitTaskMarker; use spider_storage::inbound_queue::InboundQueueConfig; @@ -393,10 +393,10 @@ async fn register_job< with_cleanup: bool, ) -> anyhow::Result { let rg_id = service - .add_resource_group(ExternalResourceGroupCredentials { - external_resource_group_id: format!("recovery-test-{}", rand::random::()), - password: b"test-password".to_vec(), - }) + .add_resource_group(ExternalResourceGroupCredentials::new( + format!("recovery-test-{}", rand::random::()), + b"test-password".to_vec(), + )) .await?; let (task_graph, inputs) = build_flat_task_graph(1, 4, with_commit, with_cleanup); let compressed_task_graph = compress_task_graph(&task_graph)?; diff --git a/tests/huntsman/e2e/src/test_driver.rs b/tests/huntsman/e2e/src/test_driver.rs index 48dc0864..5f41d7d3 100644 --- a/tests/huntsman/e2e/src/test_driver.rs +++ b/tests/huntsman/e2e/src/test_driver.rs @@ -15,6 +15,7 @@ use spider_client::SpiderClient; use spider_core::job::JobState; use spider_core::types::id::JobId; use spider_core::types::id::ResourceGroupId; +use spider_core::types::resource_group::ExternalResourceGroupCredentials; use tokio::sync::Mutex; use tokio::sync::OnceCell; use tokio::sync::RwLock; @@ -190,7 +191,10 @@ impl SpiderTestDriver { return Ok(*resource_group_id); } let resource_group_id = client - .add_resource_group(external_resource_group_id.to_owned(), Vec::new()) + .add_resource_group(ExternalResourceGroupCredentials::new( + external_resource_group_id.to_owned(), + Vec::new(), + )) .await?; resource_groups.insert(external_resource_group_id.to_owned(), resource_group_id); drop(resource_groups);