use crate::error::MoldError;
use anyhow::Result;
use reqwest::{Client, StatusCode};
use serde::{Deserialize, Serialize};
use std::collections::HashSet;
use std::fmt;
use std::time::Duration;
pub const DEFAULT_ENDPOINT: &str = "https://rest.runpod.io/v1";
pub const GRAPHQL_ENDPOINT: &str = "https://api.runpod.io/graphql";
pub const API_KEY_ENV: &str = "RUNPOD_API_KEY";
pub const NETWORK_VOLUME_MIN_GB: u32 = 10;
pub const NETWORK_VOLUME_MAX_GB: u32 = 3999;
pub fn valid_network_volume_size(size: u32) -> bool {
(NETWORK_VOLUME_MIN_GB..=NETWORK_VOLUME_MAX_GB).contains(&size)
}
#[derive(Debug, Clone, Deserialize, Serialize, Default)]
pub struct RunPodSettings {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub api_key: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub default_gpu: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub default_datacenter: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub default_network_volume_id: Option<String>,
#[serde(default)]
pub auto_teardown: bool,
#[serde(default = "default_auto_teardown_idle_mins")]
pub auto_teardown_idle_mins: u32,
#[serde(default)]
pub cost_alert_usd: f64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub endpoint: Option<String>,
}
fn default_auto_teardown_idle_mins() -> u32 {
20
}
impl RunPodSettings {
pub fn redacted_debug(&self) -> String {
format!(
"RunPodSettings {{ api_key: {}, default_gpu: {:?}, default_datacenter: {:?}, \
default_network_volume_id: {:?}, auto_teardown: {}, auto_teardown_idle_mins: {}, \
cost_alert_usd: {}, endpoint: {:?} }}",
if self.api_key.is_some() {
"Some(\"<redacted>\")"
} else {
"None"
},
self.default_gpu,
self.default_datacenter,
self.default_network_volume_id,
self.auto_teardown,
self.auto_teardown_idle_mins,
self.cost_alert_usd,
self.endpoint,
)
}
}
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct UserInfo {
pub id: String,
pub email: String,
#[serde(default)]
pub client_balance: f64,
#[serde(default)]
pub current_spend_per_hr: f64,
#[serde(default)]
pub spend_limit: Option<f64>,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct GpuType {
#[serde(default)]
pub id: Option<String>,
#[serde(rename = "displayName", default)]
pub display_name: String,
#[serde(rename = "gpuId", default)]
pub gpu_id: String,
#[serde(rename = "memoryInGb", default)]
pub memory_in_gb: Option<u32>,
#[serde(rename = "secureCloud", default)]
pub secure_cloud: bool,
#[serde(rename = "communityCloud", default)]
pub community_cloud: bool,
#[serde(rename = "stockStatus", default)]
pub stock_status: Option<String>,
#[serde(default)]
pub available: bool,
}
impl GpuType {
pub fn authoritative_type_id(&self) -> Option<&str> {
self.id
.as_deref()
.map(str::trim)
.filter(|id| !id.is_empty())
.or_else(|| {
let gpu_id = self.gpu_id.trim();
(!gpu_id.is_empty()).then_some(gpu_id)
})
}
pub fn allocation_type_id(&self) -> Option<&str> {
self.authoritative_type_id().or_else(|| {
let legacy_id = legacy_gpu_type_id_from_display_name(&self.display_name);
(!legacy_id.is_empty()).then_some(legacy_id)
})
}
}
pub fn legacy_gpu_type_id_from_display_name(display_name: &str) -> &str {
match display_name.trim() {
"RTX 4090" => "NVIDIA GeForce RTX 4090",
"RTX 5090" => "NVIDIA GeForce RTX 5090",
"RTX 3090" => "NVIDIA GeForce RTX 3090",
"L40S" => "NVIDIA L40S",
"L40" => "NVIDIA L40",
"A100 PCIe" => "NVIDIA A100 80GB PCIe",
"A100 SXM" => "NVIDIA A100-SXM4-80GB",
"H100 SXM" => "NVIDIA H100 80GB HBM3",
"H100 NVL" => "NVIDIA H100 NVL",
"RTX A6000" => "NVIDIA RTX A6000",
other => other,
}
}
pub fn normalized_gpu_type_identity(value: &str) -> String {
value
.split_whitespace()
.collect::<Vec<_>>()
.join(" ")
.to_ascii_lowercase()
}
pub fn canonical_supported_gpu_type_id<'a>(
supported_gpu_ids: &'a HashSet<String>,
candidate_id: &str,
) -> Option<&'a str> {
let candidate = normalized_gpu_type_identity(candidate_id);
if candidate.is_empty() {
return None;
}
supported_gpu_ids
.iter()
.filter(|supported| normalized_gpu_type_identity(supported) == candidate)
.map(String::as_str)
.min()
}
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct Datacenter {
pub id: String,
#[serde(default)]
pub name: String,
#[serde(default)]
pub location: Option<String>,
#[serde(rename = "gpuAvailability", default)]
pub gpu_availability: Vec<GpuAvailability>,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct GpuAvailability {
#[serde(rename = "displayName", default)]
pub display_name: String,
#[serde(rename = "gpuId", default)]
pub gpu_id: String,
#[serde(rename = "stockStatus", default)]
pub stock_status: Option<String>,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct Pod {
pub id: String,
#[serde(default)]
pub name: Option<String>,
#[serde(rename = "desiredStatus", default)]
pub desired_status: String,
#[serde(rename = "imageName", default)]
pub image_name: Option<String>,
#[serde(rename = "gpuCount", default)]
pub gpu_count: u32,
#[serde(rename = "costPerHr", default)]
pub cost_per_hr: f64,
#[serde(rename = "uptimeSeconds", default)]
pub uptime_seconds: u64,
#[serde(rename = "lastStatusChange", default)]
pub last_status_change: Option<String>,
#[serde(rename = "memoryInGb", default)]
pub memory_in_gb: u32,
#[serde(rename = "vcpuCount", default)]
pub vcpu_count: u32,
#[serde(rename = "volumeInGb", default)]
pub volume_in_gb: u32,
#[serde(rename = "volumeMountPath", default)]
pub volume_mount_path: Option<String>,
#[serde(default)]
pub ports: serde_json::Value,
#[serde(default)]
pub env: serde_json::Value,
#[serde(default)]
pub machine: Option<PodMachine>,
#[serde(default)]
pub gpu: Option<PodGpu>,
#[serde(default)]
pub runtime: Option<serde_json::Value>,
#[serde(rename = "networkVolume", default)]
pub network_volume: Option<NetworkVolume>,
#[serde(rename = "networkVolumeId", default)]
pub network_volume_id: Option<String>,
}
impl Pod {
pub fn attached_network_volume_id(&self) -> Option<&str> {
self.network_volume
.as_ref()
.map(|volume| volume.id.as_str())
.or(self.network_volume_id.as_deref())
}
pub fn gpu_name(&self) -> Option<&str> {
self.gpu
.as_ref()
.and_then(|gpu| gpu.display_name.as_deref())
.or_else(|| {
self.machine
.as_ref()
.and_then(|machine| machine.gpu_display_name.as_deref())
})
.or_else(|| {
self.machine
.as_ref()
.and_then(|machine| machine.gpu_type_id.as_deref())
})
.or_else(|| self.gpu.as_ref().and_then(|gpu| gpu.id.as_deref()))
}
pub fn datacenter_id(&self) -> Option<&str> {
self.machine
.as_ref()
.and_then(|machine| machine.data_center_id.as_deref())
}
}
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct PodMachine {
#[serde(rename = "gpuDisplayName", default)]
pub gpu_display_name: Option<String>,
#[serde(rename = "gpuTypeId", default)]
pub gpu_type_id: Option<String>,
#[serde(rename = "dataCenterId", default)]
pub data_center_id: Option<String>,
#[serde(default)]
pub location: Option<String>,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct PodGpu {
#[serde(default)]
pub id: Option<String>,
#[serde(rename = "displayName", default)]
pub display_name: Option<String>,
#[serde(default)]
pub count: Option<u32>,
}
#[derive(Debug, Clone, Serialize, Default)]
pub struct CreatePodRequest {
pub name: String,
#[serde(rename = "imageName")]
pub image_name: String,
#[serde(rename = "gpuTypeIds")]
pub gpu_type_ids: Vec<String>,
#[serde(rename = "cloudType")]
pub cloud_type: String,
#[serde(rename = "dataCenterIds", skip_serializing_if = "Option::is_none")]
pub data_center_ids: Option<Vec<String>>,
#[serde(rename = "gpuCount")]
pub gpu_count: u32,
#[serde(rename = "containerDiskInGb")]
pub container_disk_in_gb: u32,
#[serde(rename = "volumeInGb")]
pub volume_in_gb: u32,
#[serde(rename = "volumeMountPath")]
pub volume_mount_path: String,
pub ports: Vec<String>,
pub env: serde_json::Map<String, serde_json::Value>,
#[serde(rename = "networkVolumeId", skip_serializing_if = "Option::is_none")]
pub network_volume_id: Option<String>,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct NetworkVolume {
pub id: String,
pub name: String,
#[serde(rename = "dataCenterId", default)]
pub data_center_id: String,
pub size: u32,
}
#[derive(Debug, Clone, Serialize)]
pub struct CreateNetworkVolumeRequest {
pub name: String,
pub size: u32,
#[serde(rename = "dataCenterId")]
pub data_center_id: String,
}
#[derive(Debug, Clone, Serialize)]
pub struct UpdateNetworkVolumeRequest {
#[serde(skip_serializing_if = "Option::is_none")]
pub name: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub size: Option<u32>,
}
#[derive(Clone)]
pub struct RunPodClient {
endpoint: String,
graphql_endpoint: String,
api_key: String,
http: Client,
}
impl fmt::Debug for RunPodClient {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("RunPodClient")
.field("endpoint", &self.endpoint)
.field("api_key", &"<redacted>")
.finish()
}
}
impl RunPodClient {
pub fn new(endpoint: impl Into<String>, api_key: impl Into<String>) -> Self {
let rest = endpoint.into();
let graphql = if rest.starts_with(DEFAULT_ENDPOINT) {
GRAPHQL_ENDPOINT.to_string()
} else {
rest.clone()
};
Self::new_with_graphql(rest, graphql, api_key)
}
pub fn new_with_graphql(
endpoint: impl Into<String>,
graphql_endpoint: impl Into<String>,
api_key: impl Into<String>,
) -> Self {
let http = Client::builder()
.timeout(Duration::from_secs(30))
.build()
.unwrap_or_default();
Self {
endpoint: endpoint.into(),
graphql_endpoint: graphql_endpoint.into(),
api_key: api_key.into(),
http,
}
}
pub fn from_settings(settings: &RunPodSettings) -> std::result::Result<Self, MoldError> {
let key = std::env::var(API_KEY_ENV)
.ok()
.filter(|k| !k.is_empty())
.or_else(|| settings.api_key.clone())
.ok_or_else(|| {
MoldError::RunPodAuth(format!(
"RunPod API key not set — export {API_KEY_ENV} or run \
`mold config set runpod.api_key <key>`"
))
})?;
let endpoint = settings
.endpoint
.clone()
.unwrap_or_else(|| DEFAULT_ENDPOINT.to_string());
Ok(Self::new(endpoint, key))
}
fn url(&self, path: &str) -> String {
format!("{}{}", self.endpoint.trim_end_matches('/'), path)
}
async fn get_json<T: for<'de> Deserialize<'de>>(&self, path: &str) -> Result<T> {
let resp = self
.http
.get(self.url(path))
.bearer_auth(&self.api_key)
.send()
.await
.map_err(|e| MoldError::RunPod(format!("RunPod {path}: {e}")))?;
let status = resp.status();
if status.is_success() {
let body = resp
.text()
.await
.map_err(|e| MoldError::RunPod(format!("RunPod {path} body: {e}")))?;
serde_json::from_str(&body).map_err(|e| {
MoldError::RunPod(format!(
"RunPod {path}: failed to parse response: {e} — body: {}",
truncate_for_error(&body)
))
.into()
})
} else {
Err(http_error(path, status, resp).await.into())
}
}
async fn post_json<B: Serialize, T: for<'de> Deserialize<'de>>(
&self,
path: &str,
body: &B,
) -> Result<T> {
let resp = self
.http
.post(self.url(path))
.bearer_auth(&self.api_key)
.json(body)
.send()
.await
.map_err(|e| MoldError::RunPod(format!("RunPod {path}: {e}")))?;
let status = resp.status();
if status.is_success() {
let text = resp
.text()
.await
.map_err(|e| MoldError::RunPod(format!("RunPod {path} body: {e}")))?;
serde_json::from_str(&text).map_err(|e| {
MoldError::RunPod(format!(
"RunPod {path}: failed to parse response: {e} — body: {}",
truncate_for_error(&text)
))
.into()
})
} else {
Err(http_error(path, status, resp).await.into())
}
}
async fn post_empty(&self, path: &str) -> Result<()> {
let resp = self
.http
.post(self.url(path))
.bearer_auth(&self.api_key)
.send()
.await
.map_err(|e| MoldError::RunPod(format!("RunPod {path}: {e}")))?;
let status = resp.status();
if status.is_success() {
Ok(())
} else {
Err(http_error(path, status, resp).await.into())
}
}
async fn patch_json<B: Serialize, T: for<'de> Deserialize<'de>>(
&self,
path: &str,
body: &B,
) -> Result<T> {
let resp = self
.http
.patch(self.url(path))
.bearer_auth(&self.api_key)
.json(body)
.send()
.await
.map_err(|e| MoldError::RunPod(format!("RunPod {path}: {e}")))?;
let status = resp.status();
if status.is_success() {
let text = resp
.text()
.await
.map_err(|e| MoldError::RunPod(format!("RunPod {path} body: {e}")))?;
serde_json::from_str(&text).map_err(|e| {
MoldError::RunPod(format!(
"RunPod {path}: failed to parse response: {e} — body: {}",
truncate_for_error(&text)
))
.into()
})
} else {
Err(http_error(path, status, resp).await.into())
}
}
async fn delete(&self, path: &str) -> Result<()> {
let resp = self
.http
.delete(self.url(path))
.bearer_auth(&self.api_key)
.send()
.await
.map_err(|e| MoldError::RunPod(format!("RunPod {path}: {e}")))?;
let status = resp.status();
if status.is_success() {
Ok(())
} else {
Err(http_error(path, status, resp).await.into())
}
}
pub async fn user(&self) -> Result<UserInfo> {
let query = serde_json::json!({
"query": "query { myself { id email clientBalance currentSpendPerHr spendLimit } }"
});
let resp = self
.http
.post(&self.graphql_endpoint)
.bearer_auth(&self.api_key)
.json(&query)
.send()
.await
.map_err(|e| MoldError::RunPod(format!("RunPod graphql /user: {e}")))?;
let status = resp.status();
if !status.is_success() {
return Err(http_error("graphql /user", status, resp).await.into());
}
let body: serde_json::Value = resp
.json()
.await
.map_err(|e| MoldError::RunPod(format!("RunPod graphql /user json: {e}")))?;
if let Some(errs) = body.get("errors") {
return Err(MoldError::RunPod(format!("RunPod graphql errors: {errs}")).into());
}
let myself = body
.get("data")
.and_then(|d| d.get("myself"))
.ok_or_else(|| MoldError::RunPod("graphql: missing data.myself".into()))?;
let info = UserInfo {
id: myself
.get("id")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string(),
email: myself
.get("email")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string(),
client_balance: myself
.get("clientBalance")
.and_then(|v| v.as_f64())
.unwrap_or(0.0),
current_spend_per_hr: myself
.get("currentSpendPerHr")
.and_then(|v| v.as_f64())
.unwrap_or(0.0),
spend_limit: myself.get("spendLimit").and_then(|v| v.as_f64()),
};
Ok(info)
}
pub async fn gpu_types(&self) -> Result<Vec<GpuType>> {
let query = serde_json::json!({
"query": "query { gpuTypes { id displayName memoryInGb secureCloud communityCloud } dataCenters { gpuAvailability { displayName stockStatus } } }"
});
let body = self.graphql(&query).await?;
let data = body
.get("data")
.ok_or_else(|| MoldError::RunPod("graphql: missing data".into()))?;
let types: Vec<GpuType> = serde_json::from_value(
data.get("gpuTypes")
.cloned()
.unwrap_or(serde_json::Value::Array(vec![])),
)
.map_err(|e| MoldError::RunPod(format!("parse gpuTypes: {e}")))?;
let mut best_stock: std::collections::HashMap<String, String> =
std::collections::HashMap::new();
if let Some(dcs) = data.get("dataCenters").and_then(|v| v.as_array()) {
for dc in dcs {
if let Some(avail) = dc.get("gpuAvailability").and_then(|v| v.as_array()) {
for a in avail {
if let (Some(name), Some(stock)) = (
a.get("displayName").and_then(|v| v.as_str()),
a.get("stockStatus").and_then(|v| v.as_str()),
) {
let current = best_stock.get(name).cloned().unwrap_or_default();
if stock_rank(stock) > stock_rank(¤t) {
best_stock.insert(name.to_string(), stock.to_string());
}
}
}
}
}
}
let mut out = types;
for g in out.iter_mut() {
if let Some(s) = best_stock.get(&g.display_name) {
if !s.is_empty() {
g.stock_status = Some(s.clone());
}
}
g.available = g.stock_status.as_deref().is_some_and(|s| s != "None");
}
Ok(out)
}
pub async fn datacenters(&self) -> Result<Vec<Datacenter>> {
let query = serde_json::json!({
"query": "query { dataCenters { id name listed gpuAvailability { id displayName stockStatus } } }"
});
let body = self.graphql(&query).await?;
let arr = body
.get("data")
.and_then(|d| d.get("dataCenters"))
.cloned()
.unwrap_or(serde_json::Value::Array(vec![]));
let arr = match arr {
serde_json::Value::Array(mut dcs) => {
for dc in dcs.iter_mut() {
if let Some(avail) =
dc.get_mut("gpuAvailability").and_then(|v| v.as_array_mut())
{
for a in avail.iter_mut() {
if let Some(id) = a.get("id").and_then(|v| v.as_str()) {
let id = id.to_string();
if let Some(obj) = a.as_object_mut() {
obj.insert("gpuId".into(), serde_json::Value::String(id));
}
}
}
}
}
serde_json::Value::Array(dcs)
}
other => other,
};
let dcs: Vec<Datacenter> = serde_json::from_value(arr)
.map_err(|e| MoldError::RunPod(format!("parse dataCenters: {e}")))?;
Ok(dcs)
}
async fn graphql(&self, query: &serde_json::Value) -> Result<serde_json::Value> {
let resp = self
.http
.post(&self.graphql_endpoint)
.bearer_auth(&self.api_key)
.json(query)
.send()
.await
.map_err(|e| MoldError::RunPod(format!("RunPod graphql: {e}")))?;
let status = resp.status();
if !status.is_success() {
return Err(http_error("graphql", status, resp).await.into());
}
let body: serde_json::Value = resp
.json()
.await
.map_err(|e| MoldError::RunPod(format!("graphql body: {e}")))?;
if let Some(errs) = body
.get("errors")
.filter(|e| !e.as_array().map(|a| a.is_empty()).unwrap_or(true))
{
return Err(MoldError::RunPod(format!("graphql errors: {errs}")).into());
}
Ok(body)
}
pub async fn list_pods(&self) -> Result<Vec<Pod>> {
self.get_json("/pods?includeMachine=true").await
}
pub async fn get_pod(&self, id: &str) -> Result<Pod> {
self.get_json(&format!("/pods/{id}?includeMachine=true"))
.await
}
pub async fn create_pod(&self, req: &CreatePodRequest) -> Result<Pod> {
self.post_json("/pods", req).await
}
pub async fn supported_pod_gpu_type_ids(&self) -> Result<HashSet<String>> {
let spec: serde_json::Value = self.get_json("/openapi.json").await?;
parse_pod_gpu_type_ids(&spec)
}
pub async fn stop_pod(&self, id: &str) -> Result<()> {
self.post_empty(&format!("/pods/{id}/stop")).await
}
pub async fn start_pod(&self, id: &str) -> Result<()> {
self.post_empty(&format!("/pods/{id}/start")).await
}
pub async fn delete_pod(&self, id: &str) -> Result<()> {
self.delete(&format!("/pods/{id}")).await
}
pub async fn network_volumes(&self) -> Result<Vec<NetworkVolume>> {
self.get_json("/networkvolumes").await
}
pub async fn get_network_volume(&self, id: &str) -> Result<NetworkVolume> {
self.get_json(&format!("/networkvolumes/{id}")).await
}
pub async fn create_network_volume(
&self,
req: &CreateNetworkVolumeRequest,
) -> Result<NetworkVolume> {
self.post_json("/networkvolumes", req).await
}
pub async fn update_network_volume(
&self,
id: &str,
req: &UpdateNetworkVolumeRequest,
) -> Result<NetworkVolume> {
self.patch_json(&format!("/networkvolumes/{id}"), req).await
}
pub async fn delete_network_volume(&self, id: &str) -> Result<()> {
self.delete(&format!("/networkvolumes/{id}")).await
}
pub async fn delete_network_volume_if_detached(&self, id: &str) -> Result<()> {
let pods = self.list_pods().await?;
if let Some(pod) = pods
.iter()
.find(|pod| pod.attached_network_volume_id() == Some(id))
{
return Err(MoldError::RunPod(format!(
"delete pod {} before deleting its attached network volume",
pod.id
))
.into());
}
self.delete_network_volume(id).await
}
}
async fn http_error(path: &str, status: StatusCode, resp: reqwest::Response) -> MoldError {
let body = resp.text().await.unwrap_or_default();
let msg = truncate_for_error(&body);
match status {
StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN => {
MoldError::RunPodAuth(format!("RunPod {path} {status}: {msg}"))
}
StatusCode::NOT_FOUND => {
MoldError::RunPodNotFound(format!("RunPod {path} {status}: {msg}"))
}
StatusCode::CONFLICT
| StatusCode::SERVICE_UNAVAILABLE
| StatusCode::INTERNAL_SERVER_ERROR
if {
let lower = msg.to_lowercase();
lower.contains("does not have the resources")
|| lower.contains("no instances currently available")
} =>
{
MoldError::RunPodNoStock(format!("RunPod {path} {status}: {msg}"))
}
_ => MoldError::RunPod(format!("RunPod {path} {status}: {msg}")),
}
}
fn stock_rank(s: &str) -> u8 {
match s {
"High" => 3,
"Medium" => 2,
"Low" => 1,
_ => 0,
}
}
fn parse_pod_gpu_type_ids(spec: &serde_json::Value) -> Result<HashSet<String>> {
let values = spec
.pointer("/components/schemas/PodCreateInput/properties/gpuTypeIds/items/enum")
.and_then(serde_json::Value::as_array)
.ok_or_else(|| MoldError::RunPod("OpenAPI schema is missing Pod GPU types".into()))?;
Ok(values
.iter()
.filter_map(serde_json::Value::as_str)
.map(str::to_string)
.collect())
}
fn truncate_for_error(s: &str) -> String {
const MAX: usize = 400;
let s = s.trim();
if s.len() <= MAX {
s.to_string()
} else {
format!("{}…", &s[..MAX])
}
}
pub fn image_tag_for_gpu(
display_name: &str,
version: &str,
) -> Result<String, crate::cuda_distribution::UnsupportedPublishedImagePlatform> {
crate::cuda_distribution::image_tag_for_gpu_name(display_name, version)
}
pub const GPU_PREFERENCE: &[&str] = &[
"A100 PCIe",
"L40",
"L40S",
"RTX A6000",
"RTX 5090",
"RTX 4090",
];
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn gpu_type_authority_prefers_id_then_gpu_id_and_rejects_blank_values() {
let gpu: GpuType = serde_json::from_value(serde_json::json!({
"id": " primary-id ",
"gpuId": " alternate-id ",
"displayName": "Display label"
}))
.unwrap();
assert_eq!(gpu.authoritative_type_id(), Some("primary-id"));
let gpu: GpuType = serde_json::from_value(serde_json::json!({
"id": " ",
"gpuId": " alternate-id ",
"displayName": "Display label"
}))
.unwrap();
assert_eq!(gpu.authoritative_type_id(), Some("alternate-id"));
let gpu: GpuType = serde_json::from_value(serde_json::json!({
"id": " ",
"gpuId": "\t",
"displayName": "Display label"
}))
.unwrap();
assert_eq!(gpu.authoritative_type_id(), None);
}
#[test]
fn gpu_type_allocation_identity_uses_display_only_as_legacy_fallback() {
let gpu: GpuType = serde_json::from_value(serde_json::json!({
"id": " ",
"gpuId": "\t",
"displayName": " RTX 5090 "
}))
.unwrap();
assert_eq!(gpu.allocation_type_id(), Some("NVIDIA GeForce RTX 5090"));
let gpu: GpuType = serde_json::from_value(serde_json::json!({
"id": " provider-id ",
"gpuId": "alternate-id",
"displayName": "RTX 5090"
}))
.unwrap();
assert_eq!(gpu.allocation_type_id(), Some("provider-id"));
}
#[test]
fn gpu_type_memory_preserves_missing_and_present_zero() {
let missing: GpuType = serde_json::from_value(serde_json::json!({
"displayName": "RTX 5090"
}))
.unwrap();
let zero: GpuType = serde_json::from_value(serde_json::json!({
"displayName": "RTX 5090",
"memoryInGb": 0
}))
.unwrap();
assert_eq!(missing.memory_in_gb, None);
assert_eq!(zero.memory_in_gb, Some(0));
}
#[test]
fn supported_gpu_identity_matching_normalizes_case_and_whitespace() {
let supported = ["NVIDIA A100-SXM4-80GB".to_string()].into_iter().collect();
assert_eq!(
canonical_supported_gpu_type_id(&supported, " nvidia a100-sxm4-80gb "),
Some("NVIDIA A100-SXM4-80GB")
);
assert!(canonical_supported_gpu_type_id(&supported, " \t ").is_none());
}
#[test]
fn image_tag_mapping() {
assert_eq!(image_tag_for_gpu("RTX 4090", "latest").unwrap(), "latest");
assert_eq!(
image_tag_for_gpu("NVIDIA GeForce RTX 4090", "0.10.0").unwrap(),
"0.10.0"
);
assert_eq!(image_tag_for_gpu("L40S", "latest").unwrap(), "latest");
assert_eq!(
image_tag_for_gpu("RTX 5090", "latest").unwrap(),
"latest-sm120"
);
assert_eq!(
image_tag_for_gpu("NVIDIA GeForce RTX 5090", "0.10.0").unwrap(),
"0.10.0-sm120"
);
assert_eq!(
image_tag_for_gpu("RTX PRO 4500", "latest").unwrap(),
"latest-sm120"
);
assert_eq!(
image_tag_for_gpu("NVIDIA B200", "latest").unwrap(),
"latest-sm100"
);
assert!(image_tag_for_gpu("NVIDIA GB200", "0.10.0").is_err());
assert_eq!(
image_tag_for_gpu("A100 80GB", "latest").unwrap(),
"latest-sm80"
);
assert_eq!(
image_tag_for_gpu("A100 PCIe", "latest").unwrap(),
"latest-sm80"
);
assert_eq!(
image_tag_for_gpu("RTX 3090", "latest").unwrap(),
"latest-sm86"
);
assert_eq!(
image_tag_for_gpu("NVIDIA A40", "latest").unwrap(),
"latest-sm86"
);
assert_eq!(
image_tag_for_gpu("H100 SXM", "latest").unwrap(),
"latest-sm90"
);
assert_eq!(
image_tag_for_gpu("H200 SXM", "latest").unwrap(),
"latest-sm90"
);
assert_eq!(
image_tag_for_gpu("NVIDIA A10", "latest").unwrap(),
"latest-sm86"
);
assert_eq!(
image_tag_for_gpu("NVIDIA RTX A6000", "latest").unwrap(),
"latest-sm86"
);
assert_eq!(
image_tag_for_gpu("NVIDIA A16", "latest").unwrap(),
"latest-sm86"
);
assert_eq!(
image_tag_for_gpu("NVIDIA A2", "latest").unwrap(),
"latest-sm86"
);
assert_eq!(
image_tag_for_gpu("NVIDIA A30", "latest").unwrap(),
"latest-sm80"
);
assert_eq!(
image_tag_for_gpu("NVIDIA B300", "latest").unwrap(),
"latest-sm100"
);
assert!(image_tag_for_gpu("NVIDIA GB300", "latest").is_err());
assert_eq!(
image_tag_for_gpu("NVIDIA Blackwell", "latest").unwrap(),
"latest",
"generic Blackwell must not guess between incompatible targets"
);
}
#[test]
fn live_network_volume_size_bounds_are_enforced() {
assert!(!valid_network_volume_size(9));
assert!(valid_network_volume_size(10));
assert!(valid_network_volume_size(3999));
assert!(!valid_network_volume_size(4000));
}
#[test]
fn pod_reads_gpu_from_top_level_rest_shape() {
let pod: Pod = serde_json::from_value(serde_json::json!({
"id": "pod-1",
"desiredStatus": "RUNNING",
"gpu": { "id": "NVIDIA RTX 6000", "displayName": "RTX PRO 6000 Blackwell" }
}))
.unwrap();
assert_eq!(
pod.gpu.and_then(|gpu| gpu.display_name).as_deref(),
Some("RTX PRO 6000 Blackwell")
);
}
#[test]
fn pod_reads_current_machine_gpu_and_network_volume_id_shape() {
let pod: Pod = serde_json::from_value(serde_json::json!({
"id": "pod-1",
"desiredStatus": "RUNNING",
"networkVolumeId": "nv-1",
"machine": {
"gpuTypeId": "NVIDIA GeForce RTX 4090",
"dataCenterId": "EU-RO-1",
"location": "RO"
}
}))
.unwrap();
assert_eq!(pod.attached_network_volume_id(), Some("nv-1"));
assert_eq!(pod.gpu_name(), Some("NVIDIA GeForce RTX 4090"));
assert_eq!(pod.datacenter_id(), Some("EU-RO-1"));
assert_eq!(
pod.machine
.as_ref()
.and_then(|machine| machine.gpu_type_id.as_deref()),
Some("NVIDIA GeForce RTX 4090")
);
}
#[test]
fn pod_location_is_not_treated_as_an_exact_datacenter_id() {
let pod: Pod = serde_json::from_value(serde_json::json!({
"id": "pod-1",
"machine": { "location": "RO" }
}))
.unwrap();
assert_eq!(pod.datacenter_id(), None);
}
#[test]
fn parses_rest_pod_gpu_ids_from_openapi() {
let spec = serde_json::json!({
"components": { "schemas": { "PodCreateInput": { "properties": {
"gpuTypeIds": { "items": { "enum": ["NVIDIA GeForce RTX 5090", "NVIDIA L40S"] } }
} } } }
});
assert_eq!(
parse_pod_gpu_type_ids(&spec).unwrap(),
["NVIDIA GeForce RTX 5090", "NVIDIA L40S"]
.into_iter()
.map(str::to_string)
.collect()
);
}
#[test]
fn redacted_debug_hides_api_key() {
let s = RunPodSettings {
api_key: Some("secret-key".to_string()),
..Default::default()
};
let out = s.redacted_debug();
assert!(!out.contains("secret-key"));
assert!(out.contains("<redacted>"));
}
#[test]
fn from_settings_requires_key() {
std::env::remove_var(API_KEY_ENV);
let err = RunPodClient::from_settings(&RunPodSettings::default()).unwrap_err();
assert!(matches!(err, MoldError::RunPodAuth(_)));
}
#[test]
fn truncate_for_error_boundary() {
let short = "short";
assert_eq!(truncate_for_error(short), "short");
let long = "x".repeat(500);
let truncated = truncate_for_error(&long);
assert!(truncated.ends_with('…'));
assert!(truncated.chars().count() <= 401);
}
#[test]
fn runpod_settings_toml_roundtrip() {
let original = RunPodSettings {
api_key: Some("k".to_string()),
default_gpu: Some("RTX 5090".to_string()),
default_datacenter: Some("EUR-IS-2".to_string()),
default_network_volume_id: Some("nv-123".to_string()),
auto_teardown: true,
auto_teardown_idle_mins: 30,
cost_alert_usd: 3.5,
endpoint: None,
};
let toml_s = toml::to_string(&original).unwrap();
let round: RunPodSettings = toml::from_str(&toml_s).unwrap();
assert_eq!(round.api_key, original.api_key);
assert_eq!(round.default_gpu, original.default_gpu);
assert_eq!(round.default_datacenter, original.default_datacenter);
assert_eq!(
round.default_network_volume_id,
original.default_network_volume_id
);
assert_eq!(round.auto_teardown, original.auto_teardown);
assert_eq!(
round.auto_teardown_idle_mins,
original.auto_teardown_idle_mins
);
assert_eq!(round.cost_alert_usd, original.cost_alert_usd);
}
#[test]
fn default_auto_teardown_idle_mins_is_20() {
let s: RunPodSettings = toml::from_str("").unwrap();
assert_eq!(s.auto_teardown_idle_mins, 20);
}
}