use std::any::Any;
use std::sync::Arc;
use wgpu::util::DeviceExt;
use crate::engine::error::{ECSError, ECSResult, ExecutionError};
use crate::engine::types::ChannelID;
use crate::gpu::{GPUBindingDesc, GPUContext, GPUResource};
use super::error::{EnvironmentError, EnvironmentResult};
use super::store::Environment;
pub unsafe trait EnvPod: Any + Clone + Send + Sync + bytemuck::Pod {}
unsafe impl EnvPod for f32 {}
unsafe impl EnvPod for f64 {}
unsafe impl EnvPod for u32 {}
unsafe impl EnvPod for i32 {}
unsafe impl EnvPod for u64 {}
unsafe impl EnvPod for i64 {}
unsafe impl EnvPod for u16 {}
unsafe impl EnvPod for i16 {}
unsafe impl EnvPod for u8 {}
unsafe impl EnvPod for i8 {}
type PackerFn = dyn Fn(&Environment, &mut Vec<u8>) -> Result<(), String> + Send + Sync;
struct Packer {
key: String,
channel_id: ChannelID,
byte_size: usize,
pack: Box<PackerFn>,
}
impl Packer {
fn new<T: EnvPod>(key: impl Into<String>, channel_id: ChannelID) -> Self {
let key = key.into();
let key2 = key.clone();
Self {
byte_size: std::mem::size_of::<T>(),
channel_id,
pack: Box::new(move |env, buf| {
let v: T = env.get::<T>(&key2).map_err(|e| e.to_string())?;
let bytes: &[u8] = bytemuck::bytes_of(&v);
buf.extend_from_slice(bytes);
Ok(())
}),
key,
}
}
}
pub struct EnvUniformBuffer {
env: Arc<Environment>,
packers: Vec<Packer>,
cpu_buf: Vec<u8>,
gpu_buf: Option<wgpu::Buffer>,
cpu_dirty: bool,
}
impl EnvUniformBuffer {
pub fn builder(env: Arc<Environment>) -> EnvUniformBufferBuilder {
EnvUniformBufferBuilder {
env,
packers: Vec::new(),
}
}
pub fn keys(&self) -> impl Iterator<Item = &str> {
self.packers.iter().map(|p| p.key.as_str())
}
pub fn byte_size(&self) -> usize {
self.packers.iter().map(|p| p.byte_size).sum()
}
pub fn mark_cpu_dirty(&mut self) {
self.cpu_dirty = true;
}
pub fn owns_channel(&self, id: ChannelID) -> bool {
self.packers.iter().any(|p| p.channel_id == id)
}
fn repack(&mut self) -> Result<(), ECSError> {
self.cpu_buf.clear();
for p in &self.packers {
(p.pack)(&self.env, &mut self.cpu_buf).map_err(|e| {
ECSError::from(ExecutionError::GpuDispatchFailed {
message: format!("EnvUniformBuffer pack error: {e}").into(),
})
})?;
}
Ok(())
}
}
impl GPUResource for EnvUniformBuffer {
fn name(&self) -> &str {
"EnvUniformBuffer"
}
fn create_gpu(&mut self, ctx: &GPUContext) -> ECSResult<()> {
self.repack()?;
let buf = ctx
.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("EnvUniformBuffer"),
contents: &self.cpu_buf,
usage: wgpu::BufferUsages::UNIFORM | wgpu::BufferUsages::COPY_DST,
});
self.gpu_buf = Some(buf);
let owned: Vec<ChannelID> = self.packers.iter().map(|p| p.channel_id).collect();
self.env.clear_dirty_for_channels(&owned)?;
self.cpu_dirty = false;
Ok(())
}
fn upload(&mut self, ctx: &GPUContext) -> ECSResult<()> {
if !self.cpu_dirty {
return Ok(());
}
self.repack()?;
let buf = self.gpu_buf.as_ref().ok_or_else(|| {
ECSError::from(ExecutionError::GpuDispatchFailed {
message: "EnvUniformBuffer::upload called before create_gpu".into(),
})
})?;
ctx.queue.write_buffer(buf, 0, &self.cpu_buf);
self.cpu_dirty = false;
Ok(())
}
fn download(&mut self, _ctx: &GPUContext) -> ECSResult<()> {
Ok(())
}
fn bindings(&self) -> &[GPUBindingDesc] {
static B: [GPUBindingDesc; 1] = [GPUBindingDesc { read_only: true }];
&B
}
fn encode_bind_group_entries<'a>(
&'a self,
base: u32,
out: &mut Vec<wgpu::BindGroupEntry<'a>>,
) -> ECSResult<()> {
let buf = self.gpu_buf.as_ref().ok_or_else(|| {
ECSError::from(ExecutionError::GpuDispatchFailed {
message: "EnvUniformBuffer not yet created on GPU".into(),
})
})?;
out.push(wgpu::BindGroupEntry {
binding: base,
resource: buf.as_entire_binding(),
});
Ok(())
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
fn as_any_mut(&mut self) -> &mut dyn std::any::Any {
self
}
}
impl EnvUniformBuffer {
pub fn is_cpu_dirty(&self) -> bool {
self.cpu_dirty
}
}
pub struct EnvUniformBufferBuilder {
env: Arc<Environment>,
packers: Vec<Packer>,
}
impl EnvUniformBufferBuilder {
pub fn include<T: EnvPod>(mut self, key: impl Into<String>) -> EnvironmentResult<Self> {
let key = key.into();
let channel_id = self
.env
.channel_of(&key)
.ok_or_else(|| EnvironmentError::KeyNotFound(key.clone()))?;
self.packers.push(Packer::new::<T>(key, channel_id));
Ok(self)
}
pub fn validate_byte_size(self, expected: usize) -> EnvironmentResult<Self> {
let actual: usize = self.packers.iter().map(|p| p.byte_size).sum();
if actual != expected {
return Err(EnvironmentError::UniformLayoutMismatch { expected, actual });
}
Ok(self)
}
pub fn build(self) -> EnvUniformBuffer {
let byte_size: usize = self.packers.iter().map(|p| p.byte_size).sum();
EnvUniformBuffer {
env: self.env,
packers: self.packers,
cpu_buf: Vec::with_capacity(byte_size),
gpu_buf: None,
cpu_dirty: false,
}
}
}
#[cfg(all(test, feature = "gpu"))]
mod tests {
use super::*;
use crate::environment::builder::EnvironmentBuilder;
fn make_env() -> Arc<Environment> {
EnvironmentBuilder::new()
.register::<f32>("rate", 0.05f32)
.unwrap()
.register::<u32>("size", 100u32)
.unwrap()
.build()
.unwrap()
}
#[test]
fn builder_tracks_correct_keys() {
let env = make_env();
let buf = EnvUniformBuffer::builder(Arc::clone(&env))
.include::<f32>("rate")
.unwrap()
.include::<u32>("size")
.unwrap()
.build();
let keys: Vec<&str> = buf.keys().collect();
assert_eq!(keys, ["rate", "size"]);
}
#[test]
fn byte_size_matches_expected() {
let env = make_env();
let buf = EnvUniformBuffer::builder(Arc::clone(&env))
.include::<f32>("rate") .unwrap()
.include::<u32>("size") .unwrap()
.build();
assert_eq!(buf.byte_size(), 8);
}
#[test]
fn validate_byte_size_passes_on_match() {
let env = make_env();
let buf = EnvUniformBuffer::builder(Arc::clone(&env))
.include::<f32>("rate")
.unwrap()
.include::<u32>("size")
.unwrap()
.validate_byte_size(8)
.unwrap()
.build();
assert_eq!(buf.byte_size(), 8);
}
#[test]
fn validate_byte_size_reports_mismatch() {
let env = make_env();
let result = EnvUniformBuffer::builder(Arc::clone(&env))
.include::<f32>("rate")
.unwrap()
.include::<u32>("size")
.unwrap()
.validate_byte_size(16);
match result {
Err(err) => assert_eq!(
err,
EnvironmentError::UniformLayoutMismatch {
expected: 16,
actual: 8,
}
),
Ok(_) => panic!("expected uniform layout mismatch"),
}
}
#[test]
fn include_reports_missing_key() {
let env = make_env();
let result = EnvUniformBuffer::builder(Arc::clone(&env)).include::<f32>("missing");
match result {
Err(err) => assert_eq!(err, EnvironmentError::KeyNotFound("missing".into())),
Ok(_) => panic!("expected missing key error"),
}
}
#[test]
fn is_cpu_dirty_starts_false() {
let env = make_env();
let buf = EnvUniformBuffer::builder(Arc::clone(&env))
.include::<f32>("rate")
.unwrap()
.build();
assert!(!buf.is_cpu_dirty());
}
#[test]
fn mark_cpu_dirty_sets_flag() {
let env = make_env();
let mut buf = EnvUniformBuffer::builder(Arc::clone(&env))
.include::<f32>("rate")
.unwrap()
.build();
assert!(!buf.is_cpu_dirty());
buf.mark_cpu_dirty();
assert!(buf.is_cpu_dirty());
}
#[test]
fn owns_channel_tracks_included_keys() {
let env = make_env();
let id_rate = env.channel_of("rate").unwrap();
let id_size = env.channel_of("size").unwrap();
let buf = EnvUniformBuffer::builder(Arc::clone(&env))
.include::<f32>("rate")
.unwrap()
.build();
assert!(buf.owns_channel(id_rate));
assert!(!buf.owns_channel(id_size));
assert!(!buf.owns_channel(999));
}
}