use serde_json::Value;
use crate::GpuInfo;
use crate::report::{GpuSurvey, NoteKind, SurveyNote};
use crate::vendor::{GpuArch, GpuVendor};
pub const ENV_TESTING_GPU_JSON: &str = "FLODL_TESTING_GPU_JSON";
pub(crate) fn spoofed_survey() -> Option<GpuSurvey> {
let raw = std::env::var(ENV_TESTING_GPU_JSON).ok()?;
if raw.trim().is_empty() {
return None;
}
match parse_survey(&raw) {
Ok(s) => Some(s),
Err(e) => panic!(
"{ENV_TESTING_GPU_JSON} is set but could not be parsed: {e}\n\
Expected a JSON array of devices, e.g.\n \
[{{\"vendor\":\"amd\",\"arch\":\"gfx1030\",\"vram_mb\":16384}}]\n\
or an envelope {{\"gpus\":[...],\"notes\":[...]}}.\n\
Unset the variable to use real hardware detection."
),
}
}
fn parse_survey(raw: &str) -> Result<GpuSurvey, String> {
let root: Value = serde_json::from_str(raw).map_err(|e| e.to_string())?;
let empty: Vec<Value> = Vec::new();
let (devices, notes) = match &root {
Value::Array(items) => (items, &empty),
Value::Object(_) => {
let devices = match field(&root, "gpus") {
Some(Value::Array(items)) => items,
Some(other) => {
return Err(format!("`gpus` must be an array, found {}", kind(other)));
}
None => &empty,
};
let notes = match field(&root, "notes") {
Some(Value::Array(items)) => items,
Some(other) => {
return Err(format!("`notes` must be an array, found {}", kind(other)));
}
None => &empty,
};
(devices, notes)
}
other => {
return Err(format!(
"expected an array of devices or a {{\"gpus\":[...]}} envelope, found {}",
kind(other)
));
}
};
let mut out = GpuSurvey::default();
for (pos, d) in devices.iter().enumerate() {
out.devices.push(parse_device(d, pos)?);
}
for n in notes {
out.notes.push(parse_note(n)?);
}
Ok(out)
}
fn field<'a>(v: &'a Value, key: &str) -> Option<&'a Value> {
v.get(key).filter(|v| !v.is_null())
}
fn kind(v: &Value) -> &'static str {
match v {
Value::Null => "null",
Value::Bool(_) => "boolean",
Value::Number(_) => "number",
Value::String(_) => "string",
Value::Array(_) => "array",
Value::Object(_) => "object",
}
}
const DEVICE_KEYS: &[&str] = &["index", "vendor", "name", "arch", "sm", "vram_mb"];
const NOTE_KEYS: &[&str] = &["vendor", "kind", "message"];
fn reject_unknown_keys(v: &Value, allowed: &[&str], what: &str) -> Result<(), String> {
if let Value::Object(map) = v {
for k in map.keys() {
if !allowed.contains(&k.as_str()) {
return Err(format!(
"unknown {what} key {k:?} (allowed: {})",
allowed.join(", ")
));
}
}
}
Ok(())
}
fn parse_device(v: &Value, pos: usize) -> Result<GpuInfo, String> {
if !v.is_object() {
return Err(format!("each device must be an object, found {}", kind(v)));
}
reject_unknown_keys(v, DEVICE_KEYS, "device")?;
let vendor = match field(v, "vendor") {
None => GpuVendor::Nvidia,
Some(j) => {
let s = j.as_str().ok_or_else(|| {
format!("device {pos}: `vendor` must be a string, found {}", kind(j))
})?;
GpuVendor::parse(s).ok_or_else(|| format!("device {pos}: unknown vendor {s:?}"))?
}
};
let raw_arch = field(v, "arch")
.or_else(|| field(v, "sm"))
.ok_or_else(|| format!("device {pos}: missing required `arch`"))?;
let token = raw_arch.as_str().ok_or_else(|| {
format!(
"device {pos}: `arch` must be a string, found {}",
kind(raw_arch)
)
})?;
let arch = GpuArch::parse(vendor, token)
.ok_or_else(|| format!("device {pos}: {token:?} is not a valid {vendor} arch"))?;
let index = match field(v, "index") {
None => u8::try_from(pos).map_err(|_| {
format!("device {pos}: array position exceeds the device-index range (0..255)")
})?,
Some(j) => {
let n = j.as_u64().ok_or_else(|| {
format!(
"device {pos}: `index` must be a non-negative integer, found {}",
kind(j)
)
})?;
u8::try_from(n).map_err(|_| {
format!("device {pos}: `index` {n} exceeds the device-index range (0..255)")
})?
}
};
let total_memory_mb = match field(v, "vram_mb") {
None => 0,
Some(j) => j.as_u64().ok_or_else(|| {
format!(
"device {pos}: `vram_mb` must be a non-negative integer, found {}",
kind(j)
)
})?,
};
let name = match field(v, "name") {
None => format!("{vendor} spoofed {arch}"),
Some(j) => j
.as_str()
.ok_or_else(|| format!("device {pos}: `name` must be a string, found {}", kind(j)))?
.to_string(),
};
Ok(GpuInfo {
index,
vendor,
name,
arch,
total_memory_mb,
})
}
fn parse_note(v: &Value) -> Result<SurveyNote, String> {
if !v.is_object() {
return Err(format!("each note must be an object, found {}", kind(v)));
}
reject_unknown_keys(v, NOTE_KEYS, "note")?;
let vendor = match field(v, "vendor") {
None => GpuVendor::Nvidia,
Some(j) => {
let s = j
.as_str()
.ok_or_else(|| format!("note: `vendor` must be a string, found {}", kind(j)))?;
GpuVendor::parse(s).ok_or_else(|| format!("note: unknown vendor {s:?}"))?
}
};
let note_kind = match field(v, "kind") {
None => NoteKind::HardwareUnusable,
Some(j) => {
let s = j
.as_str()
.ok_or_else(|| format!("note: `kind` must be a string, found {}", kind(j)))?;
NoteKind::parse(s).ok_or_else(|| {
format!(
"note: unknown kind {s:?} (allowed: {})",
NoteKind::ALL_NAMES
)
})?
}
};
let message = match field(v, "message") {
None => return Err("note: missing required `message`".to_string()),
Some(j) => j
.as_str()
.ok_or_else(|| format!("note: `message` must be a string, found {}", kind(j)))?
.to_string(),
};
Ok(SurveyNote {
vendor,
kind: note_kind,
message,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn bare_array_spoofs_devices() {
let s = parse_survey(
r#"[{"vendor":"amd","arch":"gfx1030","vram_mb":16384,"name":"AMD Radeon RX 6800"},
{"vendor":"nvidia","arch":"sm_120","vram_mb":16311,"index":3}]"#,
)
.unwrap();
assert_eq!(s.devices.len(), 2);
assert_eq!(s.devices[0].vendor, GpuVendor::Amd);
assert_eq!(s.devices[0].arch, GpuArch::Gfx("gfx1030".into()));
assert_eq!(s.devices[0].index, 0, "index defaults to array position");
assert_eq!(s.devices[0].short_name(), "Radeon RX 6800");
assert_eq!(s.devices[1].index, 3, "explicit index wins");
assert_eq!(
s.devices[1].arch,
GpuArch::Sm {
major: 12,
minor: 0
}
);
assert!(s.notes.is_empty());
}
#[test]
fn replays_a_captured_probe_gpus_array() {
let s = parse_survey(
r#"[{"index":0,"name":"NVIDIA GeForce RTX 5060 Ti","vendor":"nvidia",
"arch":"sm_120","sm":"sm_120","vram_mb":16311}]"#,
)
.unwrap();
assert_eq!(s.devices.len(), 1);
assert_eq!(
s.devices[0].arch,
GpuArch::Sm {
major: 12,
minor: 0
}
);
assert_eq!(s.devices[0].total_memory_mb, 16311);
}
#[test]
fn sm_alone_works_for_a_legacy_capture() {
let s = parse_survey(r#"[{"sm":"sm_86","vram_mb":24564}]"#).unwrap();
assert_eq!(s.devices[0].arch, GpuArch::Sm { major: 8, minor: 6 });
assert_eq!(s.devices[0].vendor, GpuVendor::Nvidia);
}
#[test]
fn envelope_carries_notes_with_no_devices() {
let s = parse_survey(
r#"{"gpus":[],"notes":[{"vendor":"amd","kind":"hardware_unusable",
"message":"an AMD GPU is present but ROCm is not installed"}]}"#,
)
.unwrap();
assert!(s.devices.is_empty());
assert_eq!(s.notes.len(), 1);
assert_eq!(s.notes[0].vendor, GpuVendor::Amd);
assert_eq!(s.notes[0].kind, NoteKind::HardwareUnusable);
let err = s.require_devices().unwrap_err();
assert!(err.contains("ROCm is not installed"), "got: {err}");
}
#[test]
fn envelope_fields_are_optional() {
assert!(parse_survey("{}").unwrap().devices.is_empty());
assert!(parse_survey(r#"{"gpus":[]}"#).unwrap().notes.is_empty());
assert!(parse_survey("[]").unwrap().devices.is_empty());
}
#[test]
fn explicit_null_reads_as_absent() {
let s = parse_survey(r#"[{"arch":"sm_86","name":null,"index":null}]"#).unwrap();
assert_eq!(s.devices[0].index, 0);
assert_eq!(s.devices[0].name, "NVIDIA spoofed sm_86");
}
#[test]
fn a_name_is_generated_when_omitted() {
let s = parse_survey(r#"[{"vendor":"amd","arch":"gfx942"}]"#).unwrap();
assert_eq!(s.devices[0].name, "AMD spoofed gfx942");
}
#[test]
fn missing_arch_is_an_error_not_a_default() {
let e = parse_survey(r#"[{"vendor":"amd","vram_mb":8192}]"#).unwrap_err();
assert!(e.contains("missing required `arch`"), "got: {e}");
}
#[test]
fn arch_must_match_its_vendors_shape() {
let e = parse_survey(r#"[{"vendor":"amd","arch":"sm_120"}]"#).unwrap_err();
assert!(e.contains("not a valid AMD arch"), "got: {e}");
let e = parse_survey(r#"[{"vendor":"nvidia","arch":"gfx1030"}]"#).unwrap_err();
assert!(e.contains("not a valid NVIDIA arch"), "got: {e}");
}
#[test]
fn a_typo_in_a_key_fails_loudly() {
let e = parse_survey(r#"[{"arch":"sm_86","vram":8192}]"#).unwrap_err();
assert!(e.contains("unknown device key \"vram\""), "got: {e}");
}
#[test]
fn rejects_wrong_types_and_out_of_range_values() {
for (input, want) in [
(
r#"[{"arch":"sm_86","index":"0"}]"#,
"must be a non-negative integer",
),
(
r#"[{"arch":"sm_86","index":300}]"#,
"exceeds the device-index range",
),
(
r#"[{"arch":"sm_86","index":-1}]"#,
"must be a non-negative integer",
),
(
r#"[{"arch":"sm_86","vram_mb":-5}]"#,
"must be a non-negative integer",
),
(
r#"[{"arch":"sm_86","vram_mb":1.5}]"#,
"must be a non-negative integer",
),
(r#"[{"arch":"sm_86","name":7}]"#, "`name` must be a string"),
(r#"[{"arch":7}]"#, "`arch` must be a string"),
(r#"[{"vendor":"intel","arch":"sm_86"}]"#, "unknown vendor"),
(r#"["sm_86"]"#, "must be an object"),
(r#"{"gpus":7}"#, "`gpus` must be an array"),
(r#"{"notes":7}"#, "`notes` must be an array"),
(
r#"{"notes":[{"vendor":"amd"}]}"#,
"missing required `message`",
),
(
r#"{"notes":[{"kind":"nope","message":"m"}]}"#,
"unknown kind",
),
("7", "expected an array of devices"),
("\"hi\"", "expected an array of devices"),
] {
let e = parse_survey(input).unwrap_err();
assert!(
e.contains(want),
"input {input}\n got: {e}\n want: {want}"
);
}
}
#[test]
fn malformed_json_is_rejected_by_serde() {
for bad in ["", "[", "[1,", "{\"a\" 1}", "[1] trailing", "{\"a\":1,}"] {
assert!(parse_survey(bad).is_err(), "should reject {bad:?}");
}
let deep = format!("{}{}", "[".repeat(300), "]".repeat(300));
assert!(
parse_survey(&deep).is_err(),
"recursion limit must reject, not overflow"
);
}
#[test]
fn note_kind_round_trips_through_its_name() {
for k in [
NoteKind::HardwareUnusable,
NoteKind::ToolFailed,
NoteKind::Unparsable,
NoteKind::MaskApplied,
] {
assert_eq!(NoteKind::parse(k.as_str()), Some(k));
assert!(NoteKind::ALL_NAMES.contains(k.as_str()), "{k:?} listed");
}
}
}