use spark_connect_proto as proto;
use std::collections::HashMap;
#[derive(Debug, Clone, Default)]
pub struct ExecutorResourceRequests {
resources: HashMap<String, proto::ExecutorResourceRequest>,
}
impl ExecutorResourceRequests {
pub fn new() -> Self {
ExecutorResourceRequests {
resources: HashMap::new(),
}
}
pub fn memory(mut self, memory_mb: i64) -> Self {
self.resources.insert(
"memory".to_string(),
proto::ExecutorResourceRequest {
resource_name: "memory".to_string(),
amount: memory_mb,
discovery_script: None,
vendor: None,
},
);
self
}
pub fn off_heap_memory(mut self, memory_mb: i64) -> Self {
self.resources.insert(
"offHeap".to_string(),
proto::ExecutorResourceRequest {
resource_name: "offHeap".to_string(),
amount: memory_mb,
discovery_script: None,
vendor: None,
},
);
self
}
pub fn cores(mut self, num_cores: i64) -> Self {
self.resources.insert(
"cores".to_string(),
proto::ExecutorResourceRequest {
resource_name: "cores".to_string(),
amount: num_cores,
discovery_script: None,
vendor: None,
},
);
self
}
pub fn resource(
mut self,
name: &str,
amount: i64,
discovery_script: Option<String>,
vendor: Option<String>,
) -> Self {
self.resources.insert(
name.to_string(),
proto::ExecutorResourceRequest {
resource_name: name.to_string(),
amount,
discovery_script,
vendor,
},
);
self
}
}
#[derive(Debug, Clone, Default)]
pub struct TaskResourceRequests {
resources: HashMap<String, proto::TaskResourceRequest>,
}
impl TaskResourceRequests {
pub fn new() -> Self {
TaskResourceRequests {
resources: HashMap::new(),
}
}
pub fn cpus(mut self, num_cpus: f64) -> Self {
self.resources.insert(
"cpus".to_string(),
proto::TaskResourceRequest {
resource_name: "cpus".to_string(),
amount: num_cpus,
},
);
self
}
pub fn resource(mut self, name: &str, amount: f64) -> Self {
self.resources.insert(
name.to_string(),
proto::TaskResourceRequest {
resource_name: name.to_string(),
amount,
},
);
self
}
}
#[derive(Debug, Clone, Default)]
pub struct ResourceProfileBuilder {
executor_requests: ExecutorResourceRequests,
task_requests: TaskResourceRequests,
}
impl ResourceProfileBuilder {
pub fn new() -> Self {
ResourceProfileBuilder {
executor_requests: ExecutorResourceRequests::new(),
task_requests: TaskResourceRequests::new(),
}
}
pub fn executor_resources(mut self, requests: ExecutorResourceRequests) -> Self {
self.executor_requests = requests;
self
}
pub fn task_resources(mut self, requests: TaskResourceRequests) -> Self {
self.task_requests = requests;
self
}
pub fn build(self) -> ResourceProfile {
ResourceProfile {
proto_profile: proto::ResourceProfile {
executor_resources: self.executor_requests.resources.clone(),
task_resources: self.task_requests.resources.clone(),
},
profile_id: None,
}
}
}
#[derive(Debug, Clone)]
pub struct ResourceProfile {
pub(crate) proto_profile: proto::ResourceProfile,
pub(crate) profile_id: Option<i32>,
}
impl ResourceProfile {
pub fn id(&self) -> Option<i32> {
self.profile_id
}
pub(crate) fn proto(&self) -> &proto::ResourceProfile {
&self.proto_profile
}
}
#[cfg(test)]
mod tests {
use super::*;
use prost::Message;
#[test]
fn executor_resource_requests_encode_decode() {
let reqs = ExecutorResourceRequests::new()
.memory(2048)
.cores(4)
.resource("gpu", 2, None, Some("nvidia".to_string()));
let profile = proto::ResourceProfile {
executor_resources: reqs.resources.clone(),
task_resources: HashMap::new(),
};
let encoded = profile.encode_to_vec();
let decoded = proto::ResourceProfile::decode(encoded.as_slice()).unwrap();
assert_eq!(decoded.executor_resources.len(), 3);
assert_eq!(
decoded.executor_resources.get("memory").unwrap().amount,
2048
);
assert_eq!(decoded.executor_resources.get("cores").unwrap().amount, 4);
assert_eq!(decoded.executor_resources.get("gpu").unwrap().amount, 2);
assert_eq!(
decoded.executor_resources.get("gpu").unwrap().vendor,
Some("nvidia".to_string())
);
}
#[test]
fn task_resource_requests_encode_decode() {
let reqs = TaskResourceRequests::new().cpus(0.5).resource("gpu", 0.25);
let profile = proto::ResourceProfile {
executor_resources: HashMap::new(),
task_resources: reqs.resources.clone(),
};
let encoded = profile.encode_to_vec();
let decoded = proto::ResourceProfile::decode(encoded.as_slice()).unwrap();
assert_eq!(decoded.task_resources.len(), 2);
assert_eq!(decoded.task_resources.get("cpus").unwrap().amount, 0.5);
assert_eq!(decoded.task_resources.get("gpu").unwrap().amount, 0.25);
}
#[test]
fn resource_profile_builder() {
let executor_reqs = ExecutorResourceRequests::new().memory(4096).cores(8);
let task_reqs = TaskResourceRequests::new().cpus(1.0);
let profile = ResourceProfileBuilder::new()
.executor_resources(executor_reqs)
.task_resources(task_reqs)
.build();
assert!(profile.profile_id.is_none());
assert_eq!(profile.proto_profile.executor_resources.len(), 2);
assert_eq!(profile.proto_profile.task_resources.len(), 1);
let executor_mem = profile
.proto_profile
.executor_resources
.get("memory")
.unwrap();
assert_eq!(executor_mem.amount, 4096);
let task_cpus = profile.proto_profile.task_resources.get("cpus").unwrap();
assert_eq!(task_cpus.amount, 1.0);
}
#[test]
fn off_heap_memory_and_profile_accessors() {
let reqs = ExecutorResourceRequests::new()
.off_heap_memory(1024)
.cores(2);
assert_eq!(reqs.resources.get("offHeap").unwrap().amount, 1024);
let profile = ResourceProfileBuilder::new()
.executor_resources(reqs)
.build();
assert!(profile.id().is_none());
assert_eq!(profile.proto().executor_resources.len(), 2);
}
}