use compact_str::CompactString;
use serde::{Deserialize, Serialize};
use crate::config::QuotaLimits;
#[derive(Deserialize, Serialize, Debug, Clone, Copy, PartialEq, Eq, Hash, strum::Display)]
#[serde(rename_all = "snake_case")]
#[strum(serialize_all = "snake_case")]
pub enum QuotaScope {
User,
Project,
DefaultUser,
DefaultProject,
}
impl QuotaScope {
pub fn is_named(self) -> bool {
matches!(self, QuotaScope::User | QuotaScope::Project)
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct QuotaUsage {
pub jobs: usize,
pub gpus: u32,
}
impl QuotaUsage {
fn add(&mut self, gpus: u32) {
self.jobs += 1;
self.gpus += gpus;
}
fn remove(&mut self, gpus: u32) {
self.jobs = self.jobs.saturating_sub(1);
self.gpus = self.gpus.saturating_sub(gpus);
}
fn is_zero(&self) -> bool {
self.jobs == 0 && self.gpus == 0
}
}
#[derive(Debug, Default)]
pub struct QuotaUsageIndex {
pub users: std::collections::HashMap<CompactString, QuotaUsage>,
pub projects: std::collections::HashMap<CompactString, QuotaUsage>,
}
impl QuotaUsageIndex {
pub fn clear(&mut self) {
self.users.clear();
self.projects.clear();
}
pub fn record_running(
&mut self,
user: &CompactString,
project: Option<&CompactString>,
gpus: u32,
) {
self.users.entry(user.clone()).or_default().add(gpus);
if let Some(project) = project {
self.projects.entry(project.clone()).or_default().add(gpus);
}
}
pub fn release_running(
&mut self,
user: &CompactString,
project: Option<&CompactString>,
gpus: u32,
) {
if let Some(usage) = self.users.get_mut(user) {
usage.remove(gpus);
if usage.is_zero() {
self.users.remove(user);
}
}
if let Some(project) = project {
if let Some(usage) = self.projects.get_mut(project) {
usage.remove(gpus);
if usage.is_zero() {
self.projects.remove(project);
}
}
}
}
pub fn user(&self, user: &str) -> QuotaUsage {
self.users.get(user).copied().unwrap_or_default()
}
pub fn project(&self, project: &str) -> QuotaUsage {
self.projects.get(project).copied().unwrap_or_default()
}
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq, Eq)]
pub struct QuotaStatusEntry {
pub scope: QuotaScope,
pub name: String,
pub limits: QuotaLimits,
pub running_jobs: usize,
pub running_gpus: u32,
pub queued_jobs: usize,
}