use crate::{AssetId, PayloadLocator};
use alloc::collections::BTreeMap;
use alloc::string::String;
use alloc::vec::Vec;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, serde::Serialize, serde::Deserialize)]
#[serde(rename_all = "lowercase")]
#[derive(Default)]
pub enum ShaderKind {
#[default]
Vertex,
Fragment,
#[serde(rename = "vertex_instanced", alias = "vertexinstanced")]
VertexInstanced,
}
impl ShaderKind {
pub fn compile_kind(&self) -> &'static str {
match self {
ShaderKind::Vertex | ShaderKind::VertexInstanced => "vertex",
ShaderKind::Fragment => "fragment",
}
}
}
#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
pub struct StageSource {
#[serde(default)]
pub source: String,
#[serde(default)]
pub sources: Option<BTreeMap<String, String>>,
}
impl StageSource {}
#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
pub struct Shader {
#[serde(skip)]
pub asset_id: AssetId,
pub vertex: StageSource,
pub fragment: StageSource,
#[serde(default)]
pub vertex_instanced: Option<StageSource>,
#[serde(skip)]
pub locator: Option<PayloadLocator>,
}
impl Shader {
pub fn stage(&self, kind: ShaderKind) -> Option<&StageSource> {
match kind {
ShaderKind::Vertex => Some(&self.vertex),
ShaderKind::Fragment => Some(&self.fragment),
ShaderKind::VertexInstanced => self.vertex_instanced.as_ref(),
}
}
}
#[derive(Debug, Clone, Default, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct ShaderPayload {
pub stages: Vec<(ShaderKind, Vec<u8>)>,
}
impl ShaderPayload {
pub fn encode(&self) -> Result<Vec<u8>, postcard::Error> {
postcard::to_allocvec(self)
}
pub fn decode(bytes: &[u8]) -> Result<Self, postcard::Error> {
postcard::from_bytes(bytes)
}
pub fn stage(&self, kind: ShaderKind) -> Option<&[u8]> {
self.stages
.iter()
.find(|(k, _)| *k == kind)
.map(|(_, b)| b.as_slice())
}
}
#[cfg(test)]
mod tests {
use super::*;
use alloc::vec;
#[test]
fn payload_round_trips_and_indexes_by_kind() {
let payload = ShaderPayload {
stages: vec![
(ShaderKind::Vertex, vec![1, 2, 3]),
(ShaderKind::Fragment, vec![4, 5]),
],
};
let bytes = payload.encode().expect("encode");
let decoded = ShaderPayload::decode(&bytes).expect("decode");
assert_eq!(decoded, payload);
assert_eq!(decoded.stage(ShaderKind::Vertex), Some(&[1u8, 2, 3][..]));
assert_eq!(decoded.stage(ShaderKind::Fragment), Some(&[4u8, 5][..]));
assert_eq!(decoded.stage(ShaderKind::VertexInstanced), None);
}
#[test]
fn stage_lookup_covers_every_kind() {
let s = Shader::default();
assert!(s.stage(ShaderKind::Vertex).is_some());
assert!(s.stage(ShaderKind::Fragment).is_some());
assert!(s.stage(ShaderKind::VertexInstanced).is_none());
}
#[test]
fn the_instanced_vertex_stage_compiles_as_a_vertex_stage() {
assert_eq!(ShaderKind::Vertex.compile_kind(), "vertex");
assert_eq!(ShaderKind::VertexInstanced.compile_kind(), "vertex");
assert_eq!(ShaderKind::Fragment.compile_kind(), "fragment");
assert_eq!(ShaderKind::default(), ShaderKind::Vertex);
}
#[test]
fn stage_kinds_parse_from_their_authored_spellings() {
let kind = |s: &str| serde_json::from_str::<ShaderKind>(s).unwrap();
assert_eq!(kind(r#""vertex""#), ShaderKind::Vertex);
assert_eq!(kind(r#""fragment""#), ShaderKind::Fragment);
assert_eq!(kind(r#""vertex_instanced""#), ShaderKind::VertexInstanced);
assert_eq!(kind(r#""vertexinstanced""#), ShaderKind::VertexInstanced);
assert_eq!(
serde_json::to_string(&ShaderKind::VertexInstanced).unwrap(),
r#""vertex_instanced""#
);
}
#[test]
fn a_shader_parses_from_authored_args() {
let s: Shader = serde_json::from_str(
r#"{"vertex":{"sources":{"metal":"my.metal"}},"fragment":{"source":"my.metal"}}"#,
)
.unwrap();
assert_eq!(s.fragment.source, "my.metal");
assert_eq!(
s.vertex.sources.as_ref().expect("per-platform")["metal"],
"my.metal"
);
assert!(s.vertex_instanced.is_none());
assert_eq!(s.asset_id, AssetId::default());
assert!(s.locator.is_none());
let bytes = postcard::to_allocvec(&s).unwrap();
let back: Shader = postcard::from_bytes(&bytes).unwrap();
assert_eq!(back.fragment.source, "my.metal");
}
#[test]
fn an_empty_payload_has_no_stages() {
let payload = ShaderPayload::default();
assert!(payload.stages.is_empty());
assert_eq!(payload.stage(ShaderKind::Vertex), None);
assert_eq!(
ShaderPayload::decode(&payload.encode().unwrap()),
Ok(payload)
);
}
#[test]
fn decoding_garbage_is_an_error_not_a_panic() {
assert!(ShaderPayload::decode(&[0xff, 0xff, 0xff]).is_err());
}
}