use std::sync::{Arc, Mutex, OnceLock, PoisonError, Weak};
use arrayvec::ArrayVec;
use crate::{LayerShape, Rect};
const RUNTIME_SHADER_INLINE_UNIFORMS: usize = 16;
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Hash)]
pub enum TileMode {
#[default]
Clamp,
Repeated,
Mirror,
Decal,
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct BlurredEdgeTreatment {
shape: Option<LayerShape>,
}
impl BlurredEdgeTreatment {
pub const RECTANGLE: Self = Self {
shape: Some(LayerShape::Rectangle),
};
pub const UNBOUNDED: Self = Self { shape: None };
pub const fn with_shape(shape: LayerShape) -> Self {
Self { shape: Some(shape) }
}
pub fn shape(self) -> Option<LayerShape> {
self.shape
}
pub fn clip(self) -> bool {
self.shape.is_some()
}
pub fn tile_mode(self) -> TileMode {
if self.clip() {
TileMode::Clamp
} else {
TileMode::Decal
}
}
}
impl Default for BlurredEdgeTreatment {
fn default() -> Self {
Self::RECTANGLE
}
}
pub const RUNTIME_SHADER_PRELUDE_WGSL: &str = concat!(
include_str!("../shaders/fullscreen_quad_vs.wgsl"),
include_str!("../shaders/runtime_shader_bindings.wgsl"),
);
#[derive(Clone, Debug)]
pub struct RuntimeShader {
source: Arc<str>,
source_hash: u64,
uniforms: RuntimeShaderUniforms,
specialization: Option<Arc<ShaderSpecialization>>,
input_padding: f32,
output_padding: f32,
batched_source: bool,
preserves_transparency: bool,
domains: Option<Box<ShaderDomains>>,
}
#[derive(Clone, Debug, Default)]
struct ShaderSpecialization {
overrides: Vec<(&'static str, f64)>,
overrides_hash: OnceLock<u64>,
substrates: ArrayVec<SubstrateSpec, MAX_SUBSTRATES>,
draw_split: Option<&'static str>,
exact: bool,
}
pub(crate) struct ShaderSpecializationCache<K, const N: usize> {
entries: ArrayVec<CachedShaderSpecialization<K>, N>,
}
struct CachedShaderSpecialization<K> {
source: Option<Arc<ShaderSpecialization>>,
key: K,
result: Option<Arc<ShaderSpecialization>>,
}
impl<K: PartialEq, const N: usize> ShaderSpecializationCache<K, N> {
pub(crate) const fn new() -> Self {
assert!(N > 0);
Self {
entries: ArrayVec::new_const(),
}
}
pub(crate) fn apply(
&mut self,
shader: &mut RuntimeShader,
key: K,
specialize: impl FnOnce(&mut RuntimeShader, &K),
) {
let hit = self.entries.iter().rposition(|entry| {
entry.key == key
&& match (&entry.source, &shader.specialization) {
(Some(source), Some(current)) => Arc::ptr_eq(source, current),
(None, None) => true,
_ => false,
}
});
if let Some(index) = hit {
let entry = self.entries.remove(index);
shader.specialization.clone_from(&entry.result);
self.entries.push(entry);
return;
}
if shader
.specialization
.as_ref()
.is_some_and(|source| Arc::strong_count(source) == 1)
{
specialize(shader, &key);
return;
}
let source = shader.specialization.clone();
specialize(shader, &key);
if self.entries.is_full() {
self.entries.remove(0);
}
self.entries.push(CachedShaderSpecialization {
source,
key,
result: shader.specialization.clone(),
});
}
}
static DEFAULT_SHADER_SPECIALIZATION: ShaderSpecialization = ShaderSpecialization {
overrides: Vec::new(),
overrides_hash: OnceLock::new(),
substrates: ArrayVec::new_const(),
draw_split: None,
exact: false,
};
#[derive(Clone, Copy, Debug, Default, PartialEq)]
struct ShaderDomains {
output_support: Option<Rect>,
sample_domain: Option<Rect>,
}
fn finite_rect(rect: Option<Rect>) -> Option<Rect> {
rect.filter(|rect| {
rect.x.is_finite()
&& rect.y.is_finite()
&& rect.width.is_finite()
&& rect.height.is_finite()
})
}
pub const MAX_SUBSTRATES: usize = 3;
#[derive(Clone, Copy, Debug, PartialEq)]
pub enum SubstrateSpec {
Mean,
Average { block: u32 },
Blur { radius_px: f32 },
}
impl SubstrateSpec {
fn same_bits(&self, other: &Self) -> bool {
match (self, other) {
(Self::Mean, Self::Mean) => true,
(Self::Average { block: a }, Self::Average { block: b }) => a == b,
(Self::Blur { radius_px: a }, Self::Blur { radius_px: b }) => {
a.to_bits() == b.to_bits()
}
_ => false,
}
}
fn hash_bits<H: std::hash::Hasher>(&self, state: &mut H) {
use std::hash::Hash;
match self {
Self::Mean => 2u8.hash(state),
Self::Average { block } => {
0u8.hash(state);
block.hash(state);
}
Self::Blur { radius_px } => {
1u8.hash(state);
radius_px.to_bits().hash(state);
}
}
}
}
#[derive(Clone, Debug, PartialEq)]
struct RuntimeShaderUniforms {
len: usize,
inline: [f32; RUNTIME_SHADER_INLINE_UNIFORMS],
heap: Option<Vec<f32>>,
}
impl RuntimeShaderUniforms {
fn new() -> Self {
Self {
len: 0,
inline: [0.0; RUNTIME_SHADER_INLINE_UNIFORMS],
heap: None,
}
}
fn as_slice(&self) -> &[f32] {
if let Some(heap) = &self.heap {
heap.as_slice()
} else {
&self.inline[..self.len]
}
}
fn len(&self) -> usize {
self.as_slice().len()
}
fn ensure_len(&mut self, min_len: usize) {
if let Some(heap) = &mut self.heap {
if heap.len() < min_len {
heap.resize(min_len, 0.0);
}
return;
}
if min_len <= RUNTIME_SHADER_INLINE_UNIFORMS {
self.len = self.len.max(min_len);
return;
}
let mut heap = Vec::with_capacity(min_len);
heap.extend_from_slice(&self.inline[..self.len]);
heap.resize(min_len, 0.0);
self.heap = Some(heap);
}
fn set(&mut self, index: usize, value: f32) {
if let Some(heap) = &mut self.heap {
heap[index] = value;
} else {
self.inline[index] = value;
}
}
#[cfg(test)]
fn is_inline(&self) -> bool {
self.heap.is_none()
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq, thiserror::Error)]
pub enum RuntimeShaderUniformError {
#[error(
"uniform range starting at {index} with width {width} exceeds user uniform range 0..{max_user_uniforms}; slots {reserved_start}..{max_uniforms} are reserved for renderer data"
)]
OutOfUserRange {
index: usize,
width: usize,
max_user_uniforms: usize,
reserved_start: usize,
max_uniforms: usize,
},
}
impl RuntimeShader {
pub const MAX_UNIFORMS: usize = 256;
pub const RESERVED_UNIFORM_START: usize = 224;
pub const SUBSTRATE_REGION_UNIFORMS: [usize; MAX_SUBSTRATES] = [232, 228, 224];
pub const SOURCE_REGION_UNIFORM: usize = 236;
pub const MASK_RECT_UNIFORM: usize = 240;
pub const MASK_RADII_UNIFORM: usize = 244;
pub const EFFECT_RECT_UNIFORM: usize = 248;
pub const LOGICAL_SIZE_UNIFORM: usize = 252;
pub const ALPHA_UNIFORM: usize = 254;
pub const MAX_USER_UNIFORMS: usize = Self::RESERVED_UNIFORM_START;
#[track_caller]
pub fn new(wgsl_source: &str) -> Self {
let (source, source_hash) =
cached_shader_source(std::panic::Location::caller(), wgsl_source);
Self::with_source(source, source_hash)
}
pub fn from_shared_source(source: Arc<str>) -> Self {
let source_hash = cached_shared_shader_source_hash(&source);
Self::with_source(source, source_hash)
}
fn with_source(source: Arc<str>, source_hash: u64) -> Self {
Self {
source,
source_hash,
uniforms: RuntimeShaderUniforms::new(),
specialization: None,
input_padding: 0.0,
output_padding: 0.0,
batched_source: false,
preserves_transparency: false,
domains: None,
}
}
fn specialization(&self) -> &ShaderSpecialization {
self.specialization
.as_deref()
.unwrap_or(&DEFAULT_SHADER_SPECIALIZATION)
}
fn specialization_mut(&mut self) -> &mut ShaderSpecialization {
Arc::make_mut(self.specialization.get_or_insert_with(Arc::default))
}
pub fn set_override(&mut self, name: &'static str, value: f64) {
let position = self
.overrides()
.binary_search_by(|(existing, _)| existing.cmp(&name));
if position.is_ok_and(|index| self.overrides()[index].1.to_bits() == value.to_bits()) {
return;
}
let specialization = self.specialization_mut();
specialization.overrides_hash.take();
let overrides = &mut specialization.overrides;
match position {
Ok(index) => overrides[index].1 = value,
Err(index) => overrides.insert(index, (name, value)),
}
}
pub fn clear_override(&mut self, name: &str) -> bool {
let Ok(index) = self
.overrides()
.binary_search_by(|(existing, _)| (*existing).cmp(name))
else {
return false;
};
let specialization = self.specialization_mut();
specialization.overrides_hash.take();
specialization.overrides.remove(index);
true
}
pub fn overrides(&self) -> &[(&'static str, f64)] {
&self.specialization().overrides
}
pub fn overrides_hash(&self) -> u64 {
let specialization = self.specialization();
if specialization.overrides.is_empty() {
return 0;
}
*specialization.overrides_hash.get_or_init(|| {
#[cfg(test)]
OVERRIDE_HASH_COMPUTATIONS.with(|count| count.set(count.get() + 1));
hash_shader_bytes(specialization.overrides.iter().flat_map(|(name, value)| {
name.bytes().chain([0]).chain(value.to_bits().to_le_bytes())
}))
})
}
pub fn set_input_padding(&mut self, padding: f32) {
self.input_padding = if padding.is_finite() {
padding.max(0.0)
} else {
0.0
};
}
pub fn input_padding(&self) -> f32 {
self.input_padding
}
pub fn set_output_padding(&mut self, padding: f32) {
self.output_padding = if padding.is_finite() {
padding.max(0.0)
} else {
0.0
};
}
pub fn output_padding(&self) -> f32 {
self.output_padding
}
pub fn set_output_support(&mut self, support: Option<Rect>) {
self.set_domains(ShaderDomains {
output_support: finite_rect(support),
sample_domain: self.sample_domain(),
});
}
pub fn output_support(&self) -> Option<Rect> {
self.domains
.as_ref()
.and_then(|domains| domains.output_support)
}
fn set_domains(&mut self, domains: ShaderDomains) {
self.domains = (domains != ShaderDomains::default()).then(|| Box::new(domains));
}
pub fn set_sample_domain(&mut self, domain: Option<Rect>) {
self.set_domains(ShaderDomains {
output_support: self.output_support(),
sample_domain: finite_rect(domain),
});
}
pub fn sample_domain(&self) -> Option<Rect> {
self.domains
.as_ref()
.and_then(|domains| domains.sample_domain)
}
pub fn set_float(&mut self, index: usize, value: f32) {
let _ = self.try_set_float(index, value);
}
pub fn try_set_float(
&mut self,
index: usize,
value: f32,
) -> Result<(), RuntimeShaderUniformError> {
self.try_ensure_capacity(index, 1)?;
self.uniforms.set(index, value);
Ok(())
}
pub fn set_float2(&mut self, index: usize, x: f32, y: f32) {
let _ = self.try_set_float2(index, x, y);
}
pub fn try_set_float2(
&mut self,
index: usize,
x: f32,
y: f32,
) -> Result<(), RuntimeShaderUniformError> {
self.try_ensure_capacity(index, 2)?;
self.uniforms.set(index, x);
self.uniforms.set(index + 1, y);
Ok(())
}
pub fn set_float4(&mut self, index: usize, x: f32, y: f32, z: f32, w: f32) {
let _ = self.try_set_float4(index, x, y, z, w);
}
pub fn try_set_float4(
&mut self,
index: usize,
x: f32,
y: f32,
z: f32,
w: f32,
) -> Result<(), RuntimeShaderUniformError> {
self.try_ensure_capacity(index, 4)?;
self.uniforms.set(index, x);
self.uniforms.set(index + 1, y);
self.uniforms.set(index + 2, z);
self.uniforms.set(index + 3, w);
Ok(())
}
pub fn set_batched_source(&mut self, batched: bool) {
self.batched_source = batched;
}
pub fn batched_source(&self) -> bool {
self.batched_source
}
pub fn set_preserves_transparency(&mut self, preserves: bool) {
self.preserves_transparency = preserves;
}
pub fn preserves_transparency(&self) -> bool {
self.preserves_transparency
}
pub fn set_substrates(&mut self, substrates: &[SubstrateSpec]) {
assert!(
substrates.len() <= MAX_SUBSTRATES,
"a runtime shader declares at most {MAX_SUBSTRATES} substrates"
);
if self.substrates().len() == substrates.len()
&& self
.substrates()
.iter()
.zip(substrates)
.all(|(existing, incoming)| existing.same_bits(incoming))
{
return;
}
self.specialization_mut().substrates = substrates.iter().copied().collect();
}
pub fn substrates(&self) -> &[SubstrateSpec] {
&self.specialization().substrates
}
pub fn hash_substrates<H: std::hash::Hasher>(&self, state: &mut H) {
use std::hash::Hash;
self.substrates().len().hash(state);
for substrate in self.substrates() {
substrate.hash_bits(state);
}
self.draw_split().hash(state);
}
pub fn set_draw_split(&mut self, override_name: Option<&'static str>) {
if self.draw_split() == override_name {
return;
}
self.specialization_mut().draw_split = override_name;
}
pub fn draw_split(&self) -> Option<&'static str> {
self.specialization().draw_split
}
pub fn set_specialization_exact(&mut self, exact: bool) {
if self.specialization_exact() == exact {
return;
}
self.specialization_mut().exact = exact;
}
pub fn specialization_exact(&self) -> bool {
self.specialization().exact
}
pub fn source(&self) -> &str {
&self.source
}
pub fn uniforms(&self) -> &[f32] {
self.uniforms.as_slice()
}
pub fn uniforms_padded(&self) -> [f32; Self::MAX_UNIFORMS] {
let mut padded = [0.0f32; Self::MAX_UNIFORMS];
let len = self.uniforms.len().min(Self::MAX_UNIFORMS);
padded[..len].copy_from_slice(&self.uniforms.as_slice()[..len]);
padded
}
pub fn source_hash(&self) -> u64 {
self.source_hash
}
fn try_ensure_capacity(
&mut self,
index: usize,
width: usize,
) -> Result<(), RuntimeShaderUniformError> {
let min_len = index
.checked_add(width)
.ok_or_else(|| Self::uniform_range_error(index, width))?;
if min_len > Self::MAX_USER_UNIFORMS {
return Err(Self::uniform_range_error(index, width));
}
self.uniforms.ensure_len(min_len);
Ok(())
}
fn uniform_range_error(index: usize, width: usize) -> RuntimeShaderUniformError {
RuntimeShaderUniformError::OutOfUserRange {
index,
width,
max_user_uniforms: Self::MAX_USER_UNIFORMS,
reserved_start: Self::RESERVED_UNIFORM_START,
max_uniforms: Self::MAX_UNIFORMS,
}
}
}
#[cfg(test)]
thread_local! {
static OVERRIDE_HASH_COMPUTATIONS: std::cell::Cell<usize> = const { std::cell::Cell::new(0) };
}
impl PartialEq for RuntimeShader {
fn eq(&self, other: &Self) -> bool {
self.source_hash == other.source_hash
&& (Arc::ptr_eq(&self.source, &other.source)
|| self.source.as_ref() == other.source.as_ref())
&& self.uniforms == other.uniforms
&& self.overrides().len() == other.overrides().len()
&& self
.overrides()
.iter()
.zip(other.overrides())
.all(|(a, b)| a.0 == b.0 && a.1.to_bits() == b.1.to_bits())
&& self.input_padding.to_bits() == other.input_padding.to_bits()
&& self.output_padding.to_bits() == other.output_padding.to_bits()
&& self.batched_source == other.batched_source
&& self.preserves_transparency == other.preserves_transparency
&& self.substrates() == other.substrates()
&& self.draw_split() == other.draw_split()
&& self.domains == other.domains
}
}
fn hash_shader_source(source: &str) -> u64 {
hash_shader_bytes(source.bytes())
}
fn hash_shader_bytes(bytes: impl IntoIterator<Item = u8>) -> u64 {
const FNV_OFFSET_BASIS: u64 = 0xcbf2_9ce4_8422_2325;
const FNV_PRIME: u64 = 0x0000_0100_0000_01b3;
bytes.into_iter().fold(FNV_OFFSET_BASIS, |hash, byte| {
(hash ^ u64::from(byte)).wrapping_mul(FNV_PRIME)
})
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct ShaderSourceCallsite {
file: &'static str,
line: u32,
column: u32,
}
struct CachedShaderSource {
callsite: ShaderSourceCallsite,
source_hash: u64,
source: Arc<str>,
}
struct CachedSharedShaderSourceHash {
byte_ptr: usize,
len: usize,
source_hash: u64,
source: Weak<str>,
}
fn cached_shared_shader_source_hash(source: &Arc<str>) -> u64 {
static CACHE: OnceLock<Mutex<Vec<CachedSharedShaderSourceHash>>> = OnceLock::new();
let byte_ptr = source.as_ptr() as usize;
let len = source.len();
let mut cache = CACHE
.get_or_init(|| Mutex::new(Vec::new()))
.lock()
.unwrap_or_else(PoisonError::into_inner);
cache.retain(|entry| entry.source.strong_count() > 0);
if let Some(entry) = cache.iter().find(|entry| {
entry.byte_ptr == byte_ptr
&& entry.len == len
&& entry
.source
.upgrade()
.is_some_and(|cached| Arc::ptr_eq(&cached, source))
}) {
return entry.source_hash;
}
let source_hash = hash_shader_source(source);
cache.push(CachedSharedShaderSourceHash {
byte_ptr,
len,
source_hash,
source: Arc::downgrade(source),
});
source_hash
}
fn cached_shader_source(
location: &'static std::panic::Location<'static>,
source: &str,
) -> (Arc<str>, u64) {
static CACHE: OnceLock<Mutex<Vec<CachedShaderSource>>> = OnceLock::new();
let callsite = ShaderSourceCallsite {
file: location.file(),
line: location.line(),
column: location.column(),
};
let mut cache = CACHE
.get_or_init(|| Mutex::new(Vec::new()))
.lock()
.unwrap_or_else(PoisonError::into_inner);
if let Some(entry) = cache.iter_mut().find(|entry| entry.callsite == callsite) {
if entry.source.as_ref() == source {
return (entry.source.clone(), entry.source_hash);
}
let source_hash = hash_shader_source(source);
entry.source_hash = source_hash;
entry.source = Arc::<str>::from(source);
return (entry.source.clone(), entry.source_hash);
}
let source_hash = hash_shader_source(source);
let shared = Arc::<str>::from(source);
cache.push(CachedShaderSource {
callsite,
source_hash,
source: shared.clone(),
});
(shared, source_hash)
}
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
pub enum ShaderTarget {
Page,
Layer,
}
#[derive(Clone, Debug, PartialEq)]
pub struct ShaderWarmUp {
pub shader: RuntimeShader,
pub target: ShaderTarget,
}
#[derive(Clone, Debug, PartialEq)]
pub enum RenderEffect {
Blur {
radius_x: f32,
radius_y: f32,
edge_treatment: TileMode,
},
Offset { offset_x: f32, offset_y: f32 },
Shader {
shader: Arc<RuntimeShader>,
},
Chain {
first: Arc<RenderEffect>,
second: Arc<RenderEffect>,
},
}
impl RenderEffect {
pub fn blur(radius: f32) -> Self {
Self::blur_with_edge_treatment(radius, TileMode::default())
}
pub fn blur_with_edge_treatment(radius: f32, edge_treatment: TileMode) -> Self {
Self::Blur {
radius_x: radius,
radius_y: radius,
edge_treatment,
}
}
pub fn blur_xy(radius_x: f32, radius_y: f32, edge_treatment: TileMode) -> Self {
Self::Blur {
radius_x,
radius_y,
edge_treatment,
}
}
pub fn offset(offset_x: f32, offset_y: f32) -> Self {
Self::Offset { offset_x, offset_y }
}
pub fn runtime_shader(shader: RuntimeShader) -> Self {
Self::Shader {
shader: Arc::new(shader),
}
}
pub fn then(self, other: RenderEffect) -> Self {
Self::Chain {
first: Arc::new(self),
second: Arc::new(other),
}
}
pub fn contains_runtime_shader(&self) -> bool {
match self {
RenderEffect::Shader { .. } => true,
RenderEffect::Chain { first, second } => {
first.contains_runtime_shader() || second.contains_runtime_shader()
}
_ => false,
}
}
pub fn preserves_transparency(&self) -> bool {
match self {
RenderEffect::Blur { .. } | RenderEffect::Offset { .. } => true,
RenderEffect::Shader { shader } => shader.preserves_transparency(),
RenderEffect::Chain { first, second } => {
first.preserves_transparency() && second.preserves_transparency()
}
}
}
pub fn input_padding(&self) -> f32 {
match self {
RenderEffect::Blur {
radius_x, radius_y, ..
} => radius_x.abs().max(radius_y.abs()),
RenderEffect::Offset { offset_x, offset_y } => offset_x.abs().max(offset_y.abs()),
RenderEffect::Shader { shader } => shader.input_padding(),
RenderEffect::Chain { first, second } => first.input_padding() + second.input_padding(),
}
}
pub fn output_padding(&self) -> f32 {
match self {
RenderEffect::Blur { .. } | RenderEffect::Offset { .. } => 0.0,
RenderEffect::Shader { shader } => shader.output_padding(),
RenderEffect::Chain { first, second } => {
first.output_padding() + second.output_padding()
}
}
}
pub fn output_support(&self) -> Option<Rect> {
match self {
RenderEffect::Blur { .. } | RenderEffect::Offset { .. } => None,
RenderEffect::Shader { shader } => shader.output_support(),
RenderEffect::Chain { second, .. } => second.output_support(),
}
}
pub fn sample_domain(&self) -> Option<Rect> {
match self {
RenderEffect::Blur { .. } | RenderEffect::Offset { .. } => None,
RenderEffect::Shader { shader } => shader.sample_domain(),
RenderEffect::Chain { second, .. } => second.sample_domain(),
}
}
}
#[cfg(test)]
#[path = "tests/render_effect_tests.rs"]
mod tests;