use super::{Memory, Resource};
use crate::validation::{report_of, rules, Validate, ValidationReport};
use acorn_core::prelude::alloc::{format, String, Vec};
use acorn_core::time::Milliseconds;
use acorn_core::validation::ValidationError;
use alloc::collections::BTreeSet;
use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct Device {
pub identifier: String,
pub index: Option<u32>,
pub name: String,
pub memory_total_mib: Option<u64>,
}
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(transparent)]
pub struct Inventory(Vec<Device>);
#[derive(Clone, Copy, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum RequirementState {
Satisfied,
Unsatisfied,
Unknown,
}
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct RequirementAssessment {
pub count: RequirementState,
pub memory: RequirementState,
pub architecture: RequirementState,
pub backend: RequirementState,
pub compute_capability: RequirementState,
pub name: RequirementState,
pub vendor: RequirementState,
}
impl RequirementAssessment {
#[must_use]
pub fn overall(&self) -> RequirementState {
let states = [
self.count,
self.memory,
self.architecture,
self.backend,
self.compute_capability,
self.name,
self.vendor,
];
if states.contains(&RequirementState::Unsatisfied) {
RequirementState::Unsatisfied
} else if states.contains(&RequirementState::Unknown) {
RequirementState::Unknown
} else {
RequirementState::Satisfied
}
}
}
impl From<bool> for RequirementState {
fn from(is_unobserved_requirement: bool) -> Self {
match is_unobserved_requirement {
| true => Self::Unknown,
| false => Self::Satisfied,
}
}
}
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct Reading {
pub identifier: String,
pub memory_used_mib: Option<u64>,
pub utilization_percent: Option<u8>,
}
#[derive(Clone, Debug, Deserialize, Eq, JsonSchema, PartialEq, Serialize)]
#[serde(deny_unknown_fields)]
pub struct Sample {
pub timestamp: Milliseconds,
pub readings: Vec<Reading>,
}
impl Inventory {
pub fn new(devices: Vec<Device>) -> Result<Self, ValidationReport> {
let inventory = Self(devices);
inventory.validate().map(|()| inventory)
}
pub fn devices(&self) -> &[Device] {
&self.0
}
pub fn into_devices(self) -> Vec<Device> {
self.0
}
fn assess_count(&self, required: u32) -> RequirementState {
usize::try_from(required).map_or(RequirementState::Unsatisfied, |required| {
if self.devices().len() >= required {
RequirementState::Satisfied
} else {
RequirementState::Unsatisfied
}
})
}
fn assess_memory(&self, required_count: u32, minimum: Option<&Memory>) -> RequirementState {
match minimum {
| None => RequirementState::Satisfied,
| Some(minimum) => match (minimum.checked_bytes(), usize::try_from(required_count)) {
| (None, _) => RequirementState::Unknown,
| (Some(_), Err(_)) => RequirementState::Unsatisfied,
| (Some(minimum_bytes), Ok(required_count)) => {
let (satisfying, unknown) =
self.devices()
.iter()
.fold((0_usize, 0_usize), |(satisfying, unknown), device| match device.memory_total_mib {
| Some(memory_mib) => match memory_mib.checked_mul(1_048_576) {
| Some(bytes) if bytes >= minimum_bytes => (satisfying.saturating_add(1), unknown),
| Some(_) => (satisfying, unknown),
| None => (satisfying, unknown.saturating_add(1)),
},
| None => (satisfying, unknown.saturating_add(1)),
});
match (satisfying >= required_count, satisfying.saturating_add(unknown) >= required_count) {
| (true, _) => RequirementState::Satisfied,
| (false, true) => RequirementState::Unknown,
| (false, false) => RequirementState::Unsatisfied,
}
}
},
}
}
}
impl Resource {
#[must_use]
pub fn assess_gpu_inventory(&self, inventory: &Inventory) -> Option<RequirementAssessment> {
match self {
| Self::GPU {
architecture,
backend,
compute_capability,
count,
memory,
name,
vendor,
..
} => {
let required_count = count.unwrap_or(1);
Some(RequirementAssessment {
count: inventory.assess_count(required_count),
memory: inventory.assess_memory(required_count, memory.as_ref()),
architecture: architecture.is_some().into(),
backend: backend.is_some().into(),
compute_capability: compute_capability.is_some().into(),
name: name.is_some().into(),
vendor: vendor.is_some().into(),
})
}
| _ => None,
}
}
}
impl Validate for Device {
fn validate(&self) -> Result<(), ValidationReport> {
field_report("identifier", rules::nonempty(&self.identifier))
.merge("", field_report("name", rules::nonempty(&self.name)))
.finish()
}
}
impl Validate for Inventory {
fn validate(&self) -> Result<(), ValidationReport> {
ValidationReport::empty("", self.0.is_empty())
.merge("", report_of(self.0.validate()))
.merge(
"",
unique_report(
self.0.iter().enumerate().map(|(index, device)| (index, device.identifier.as_str())),
"",
"identifier",
"GPU identifiers must be unique",
),
)
.merge(
"",
unique_report(
self.0
.iter()
.enumerate()
.filter_map(|(index, device)| device.index.map(|value| (index, value))),
"",
"index",
"GPU indices must be unique",
),
)
.finish()
}
}
impl Validate for Reading {
fn validate(&self) -> Result<(), ValidationReport> {
let utilization = self
.utilization_percent
.filter(|value| *value > 100)
.map_or_else(ValidationReport::new, |_| {
ValidationReport::from_error(
"utilization_percent",
ValidationError::new("range").with_message("Provide an integer from 0 through 100"),
)
});
field_report("identifier", rules::nonempty(&self.identifier))
.merge("", utilization)
.finish()
}
}
impl Validate for Sample {
fn validate(&self) -> Result<(), ValidationReport> {
ValidationReport::empty("readings", self.readings.is_empty())
.merge("readings", report_of(self.readings.validate()))
.merge(
"",
unique_report(
self.readings
.iter()
.enumerate()
.map(|(index, reading)| (index, reading.identifier.as_str())),
"readings",
"identifier",
"GPU identifiers must be unique",
),
)
.finish()
}
}
fn field_report(path: &str, result: Result<(), ValidationError>) -> ValidationReport {
result.map_or_else(|error| ValidationReport::from_error(path, error), |()| ValidationReport::new())
}
fn unique_report<T: Ord>(values: impl Iterator<Item = (usize, T)>, prefix: &str, field: &str, message: &'static str) -> ValidationReport {
values
.fold((BTreeSet::new(), ValidationReport::new()), |(mut seen, mut report), (index, value)| {
if !seen.insert(value) {
let path = match prefix.is_empty() {
| true => format!("[{index}].{field}"),
| false => format!("{prefix}[{index}].{field}"),
};
report.add(path, ValidationError::new("unique").with_message(message));
}
(seen, report)
})
.1
}
#[cfg(test)]
mod tests;