use super::ResourceRequest;
use crate::config::ResourceQuotas;
use crate::errors::ResourceError;
use crate::resource::ResourcePressure;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use tokio::sync::RwLock;
pub struct AdaptiveQuotas {
base: ResourceQuotas,
current: RwLock<ResourceQuotas>,
}
impl AdaptiveQuotas {
pub fn new(base: ResourceQuotas) -> Self {
Self {
base: base.clone(),
current: RwLock::new(base),
}
}
pub async fn adjust_for_pressure(&self, pressure: ResourcePressure) {
let mut current = self.current.write().await;
match pressure {
ResourcePressure::None => {
*current = self.base.clone();
}
ResourcePressure::Low => {
current.max_concurrent_requests =
self.base.max_concurrent_requests.saturating_sub(1);
}
ResourcePressure::Medium => {
current.max_concurrent_requests = self.base.max_concurrent_requests / 2;
current.max_context_tokens = self.base.max_context_tokens / 2;
}
ResourcePressure::High => {
current.max_concurrent_requests = 1;
current.max_context_tokens = self.base.max_context_tokens / 4;
current.max_queued_tasks = self.base.max_queued_tasks / 2;
}
ResourcePressure::Critical => {
current.max_concurrent_requests = 1;
current.max_context_tokens = 8192;
current.max_queued_tasks = 10;
current.max_gpu_memory_per_model = self.base.max_gpu_memory_per_model / 2;
}
}
}
pub async fn check(&self, request: &ResourceRequest) -> Result<(), ResourceError> {
let quotas = self.current.read().await;
if request.gpu_memory_bytes > quotas.max_gpu_memory_per_model {
return Err(ResourceError::QuotaExceeded {
resource: "gpu_memory_per_model".to_string(),
used: request.gpu_memory_bytes,
limit: quotas.max_gpu_memory_per_model,
});
}
if request.system_memory_bytes > quotas.max_context_tokens as u64 * 100 {
return Err(ResourceError::QuotaExceeded {
resource: "system_memory".to_string(),
used: request.system_memory_bytes,
limit: quotas.max_context_tokens as u64 * 100,
});
}
Ok(())
}
pub async fn current(&self) -> ResourceQuotas {
self.current.read().await.clone()
}
pub fn base(&self) -> &ResourceQuotas {
&self.base
}
}
pub struct ResourceLimitTracker {
quotas: ResourceQuotas,
current_gpu_memory: AtomicU64,
current_concurrent_requests: AtomicUsize,
current_queued_tasks: AtomicUsize,
}
impl ResourceLimitTracker {
pub fn new(quotas: ResourceQuotas) -> Self {
Self {
quotas,
current_gpu_memory: AtomicU64::new(0),
current_concurrent_requests: AtomicUsize::new(0),
current_queued_tasks: AtomicUsize::new(0),
}
}
pub fn allocate_gpu_memory(&self, bytes: u64) -> Result<GPUAllocationGuard<'_>, ResourceError> {
let mut current = self.current_gpu_memory.load(Ordering::SeqCst);
loop {
let new_total = current + bytes;
if new_total > self.quotas.max_gpu_memory_per_model {
return Err(ResourceError::QuotaExceeded {
resource: "gpu_memory".to_string(),
used: new_total,
limit: self.quotas.max_gpu_memory_per_model,
});
}
match self.current_gpu_memory.compare_exchange_weak(
current,
new_total,
Ordering::SeqCst,
Ordering::SeqCst,
) {
Ok(_) => break,
Err(c) => current = c,
}
}
Ok(GPUAllocationGuard {
tracker: self,
bytes,
})
}
pub fn start_request(&self) -> Result<RequestGuard<'_>, ResourceError> {
let mut current = self.current_concurrent_requests.load(Ordering::SeqCst);
loop {
if current >= self.quotas.max_concurrent_requests {
return Err(ResourceError::QuotaExceeded {
resource: "concurrent_requests".to_string(),
used: current as u64,
limit: self.quotas.max_concurrent_requests as u64,
});
}
match self.current_concurrent_requests.compare_exchange_weak(
current,
current + 1,
Ordering::SeqCst,
Ordering::SeqCst,
) {
Ok(_) => break,
Err(c) => current = c,
}
}
Ok(RequestGuard { tracker: self })
}
pub fn queue_task(&self) -> Result<TaskGuard<'_>, ResourceError> {
let mut current = self.current_queued_tasks.load(Ordering::SeqCst);
loop {
if current >= self.quotas.max_queued_tasks {
return Err(ResourceError::QuotaExceeded {
resource: "queued_tasks".to_string(),
used: current as u64,
limit: self.quotas.max_queued_tasks as u64,
});
}
match self.current_queued_tasks.compare_exchange_weak(
current,
current + 1,
Ordering::SeqCst,
Ordering::SeqCst,
) {
Ok(_) => break,
Err(c) => current = c,
}
}
Ok(TaskGuard { tracker: self })
}
fn release_gpu_memory(&self, bytes: u64) {
self.current_gpu_memory.fetch_sub(bytes, Ordering::SeqCst);
}
fn release_request(&self) {
self.current_concurrent_requests
.fetch_sub(1, Ordering::SeqCst);
}
fn release_task(&self) {
self.current_queued_tasks.fetch_sub(1, Ordering::SeqCst);
}
pub fn usage(&self) -> ResourceUsage {
ResourceUsage {
gpu_memory: self.current_gpu_memory.load(Ordering::SeqCst),
concurrent_requests: self.current_concurrent_requests.load(Ordering::SeqCst),
queued_tasks: self.current_queued_tasks.load(Ordering::SeqCst),
}
}
}
#[derive(Debug, Clone)]
pub struct ResourceUsage {
pub gpu_memory: u64,
pub concurrent_requests: usize,
pub queued_tasks: usize,
}
pub struct GPUAllocationGuard<'a> {
tracker: &'a ResourceLimitTracker,
bytes: u64,
}
impl<'a> Drop for GPUAllocationGuard<'a> {
fn drop(&mut self) {
self.tracker.release_gpu_memory(self.bytes);
}
}
pub struct RequestGuard<'a> {
tracker: &'a ResourceLimitTracker,
}
impl<'a> Drop for RequestGuard<'a> {
fn drop(&mut self) {
self.tracker.release_request();
}
}
pub struct TaskGuard<'a> {
tracker: &'a ResourceLimitTracker,
}
impl<'a> Drop for TaskGuard<'a> {
fn drop(&mut self) {
self.tracker.release_task();
}
}
#[cfg(test)]
#[path = "../../tests/unit/resource/quotas/quotas_test.rs"]
mod tests;