Skip to main content

gpui_rhai/
component_styles.rs

1use std::collections::BTreeMap;
2
3use rhai::{Dynamic, Engine, Map, Scope};
4use thiserror::Error;
5
6use crate::{ComponentRegistry, ModuleId, Style};
7
8const MAX_STYLED_COMPONENTS: usize = 256;
9const MAX_STYLED_PARTS: usize = 4_096;
10
11/// Validated application-wide Style overrides keyed by formal component and part.
12///
13/// Rules contain the same typed [`Style`] values accepted by per-instance
14/// `style`/`part_styles`. They are merged after component source defaults and
15/// before explicit instance overrides.
16#[derive(Clone, Debug, Default, PartialEq)]
17pub struct ComponentStyleSheet {
18    rules: BTreeMap<ModuleId, BTreeMap<String, Style>>,
19}
20
21impl ComponentStyleSheet {
22    #[must_use]
23    pub fn is_empty(&self) -> bool {
24        self.rules.is_empty()
25    }
26
27    #[must_use]
28    pub fn len(&self) -> usize {
29        self.rules.values().map(BTreeMap::len).sum()
30    }
31
32    #[must_use]
33    pub fn component(&self, id: &ModuleId) -> Option<&BTreeMap<String, Style>> {
34        self.rules.get(id)
35    }
36
37    fn new(
38        rules: BTreeMap<ModuleId, BTreeMap<String, Style>>,
39        components: &ComponentRegistry,
40    ) -> Result<Self, ComponentStyleError> {
41        if rules.len() > MAX_STYLED_COMPONENTS {
42            return Err(ComponentStyleError::TooManyComponents(rules.len()));
43        }
44        let part_count = rules.values().map(BTreeMap::len).sum::<usize>();
45        if part_count > MAX_STYLED_PARTS {
46            return Err(ComponentStyleError::TooManyParts(part_count));
47        }
48        for (id, parts) in &rules {
49            let component = components
50                .get(id)
51                .ok_or_else(|| ComponentStyleError::UnknownComponent(id.clone()))?;
52            for part in parts.keys() {
53                if !component.schema.parts.contains(part) {
54                    return Err(ComponentStyleError::UnknownPart {
55                        component: id.clone(),
56                        part: part.clone(),
57                    });
58                }
59            }
60        }
61        Ok(Self { rules })
62    }
63}
64
65/// Compile and evaluate `component_styles() -> map`, then validate every
66/// component ID and named part against the active component registry.
67///
68/// # Errors
69///
70/// Returns compile/evaluation errors, shape/type errors, resource-limit errors,
71/// or an unknown component/part diagnostic.
72pub fn load_component_styles(
73    engine: &Engine,
74    source_name: &str,
75    source: &str,
76    components: &ComponentRegistry,
77) -> Result<ComponentStyleSheet, ComponentStyleError> {
78    let mut ast = engine
79        .compile(source)
80        .map_err(|error| ComponentStyleError::Script(error.to_string()))?;
81    crate::engine::validate_assignment_targets(&ast)
82        .map_err(|error| ComponentStyleError::Script(error.to_string()))?;
83    ast.set_source(source_name);
84    let raw: Dynamic = engine
85        .call_fn(&mut Scope::new(), &ast, "component_styles", ())
86        .map_err(|error| ComponentStyleError::Script(error.to_string()))?;
87    let root = raw
88        .try_cast::<Map>()
89        .ok_or(ComponentStyleError::RootNotMap)?;
90    let mut rules = BTreeMap::new();
91    for (raw_id, raw_parts) in root {
92        let id = ModuleId::parse(raw_id.as_str()).map_err(|source| {
93            ComponentStyleError::InvalidComponentId {
94                id: raw_id.to_string(),
95                source,
96            }
97        })?;
98        let parts = raw_parts
99            .try_cast::<Map>()
100            .ok_or_else(|| ComponentStyleError::ComponentNotMap(id.clone()))?;
101        let mut decoded = BTreeMap::new();
102        for (part, value) in parts {
103            if !value.is::<Style>() {
104                return Err(ComponentStyleError::PartNotStyle {
105                    component: id,
106                    part: part.to_string(),
107                });
108            }
109            decoded.insert(part.to_string(), value.cast::<Style>());
110        }
111        rules.insert(id, decoded);
112    }
113    ComponentStyleSheet::new(rules, components)
114}
115
116#[derive(Debug, Error)]
117pub enum ComponentStyleError {
118    #[error("component stylesheet script failed: {0}")]
119    Script(String),
120    #[error("component_styles() must return a map")]
121    RootNotMap,
122    #[error("component stylesheet ID `{id}` is invalid: {source}")]
123    InvalidComponentId {
124        id: String,
125        source: crate::ModuleIdError,
126    },
127    #[error("component stylesheet rule `{0}` must be a map of named Style values")]
128    ComponentNotMap(ModuleId),
129    #[error("component stylesheet rule `{component}.{part}` must be a Style")]
130    PartNotStyle { component: ModuleId, part: String },
131    #[error("component stylesheet references unavailable component `{0}`")]
132    UnknownComponent(ModuleId),
133    #[error("component stylesheet references unknown part `{component}.{part}`")]
134    UnknownPart { component: ModuleId, part: String },
135    #[error("component stylesheet has {0} components; the limit is {MAX_STYLED_COMPONENTS}")]
136    TooManyComponents(usize),
137    #[error("component stylesheet has {0} part rules; the limit is {MAX_STYLED_PARTS}")]
138    TooManyParts(usize),
139}
140
141#[cfg(test)]
142mod tests {
143    use super::*;
144    use crate::{
145        ComponentDefinition, ComponentInstancePath, ComponentMetadata, ComponentSchema,
146        ExecutionPhase, RuntimeApiRange, UiContext, UiRuntimeState,
147    };
148    use semver::Version;
149    use std::cell::RefCell;
150    use std::rc::Rc;
151
152    fn registry() -> ComponentRegistry {
153        let mut registry = ComponentRegistry::new();
154        registry
155            .register(
156                ComponentDefinition::new(
157                    ComponentMetadata {
158                        id: ModuleId::parse("components/button").unwrap(),
159                        export: "Button".to_owned(),
160                        version: Version::new(0, 1, 0),
161                        runtime_api: RuntimeApiRange::new(2, 3),
162                        dependencies: std::collections::BTreeSet::default(),
163                        capabilities: BTreeMap::default(),
164                        assets: std::collections::BTreeSet::default(),
165                    },
166                    ComponentSchema {
167                        parts: ["root".to_owned(), "label".to_owned()]
168                            .into_iter()
169                            .collect(),
170                        ..ComponentSchema::default()
171                    },
172                )
173                .unwrap(),
174                crate::RUNTIME_API_VERSION,
175            )
176            .unwrap();
177        registry
178    }
179
180    #[test]
181    fn source_decodes_typed_styles_and_validates_parts() {
182        let engine = crate::RuntimeEngine::new();
183        let sheet = load_component_styles(
184            engine.engine(),
185            "styles.rhai",
186            r#"
187                fn component_styles() {
188                    #{ "components/button": #{
189                        root: style().height(px(34)).radius(theme_radius("md")),
190                        label: style().typography("body").font_weight(600),
191                    } }
192                }
193            "#,
194            &registry(),
195        )
196        .unwrap();
197        assert_eq!(sheet.len(), 2);
198        assert_eq!(
199            sheet
200                .component(&ModuleId::parse("components/button").unwrap())
201                .unwrap()["root"]
202                .base
203                .height,
204            Some(crate::LayoutLength::Definite(crate::Length::Pixels(34.0)))
205        );
206    }
207
208    #[test]
209    fn source_rejects_unknown_components_parts_and_non_styles() {
210        let engine = crate::RuntimeEngine::new();
211        for (source, expected) in [
212            (
213                r#"fn component_styles() { #{ "components/missing": #{ root: style() } } }"#,
214                "unavailable component",
215            ),
216            (
217                r#"fn component_styles() { #{ "components/button": #{ icon: style() } } }"#,
218                "unknown part",
219            ),
220            (
221                r#"fn component_styles() { #{ "components/button": #{ root: 12 } } }"#,
222                "must be a Style",
223            ),
224        ] {
225            let error = load_component_styles(engine.engine(), "styles.rhai", source, &registry())
226                .unwrap_err();
227            assert!(error.to_string().contains(expected), "{error}");
228        }
229    }
230
231    #[test]
232    fn stylesheet_precedes_instance_overrides_during_component_render() {
233        let mut engine = crate::RuntimeEngine::new();
234        let compiled = engine
235            .compile(
236                r#"
237                    define_component(#{
238                        metadata: #{
239                            id: "components/button", "export": "Button", version: "0.1.0",
240                            runtime_api: #{ min_inclusive: 2, max_exclusive: 3 },
241                            dependencies: [], capabilities: #{}
242                        },
243                        schema: #{ props: #{}, state: #{ fields: #{} }, events: #{},
244                            slots: #{}, parts: ["root", "label"] },
245                        render: Fn("render_Button")
246                    });
247                    fn Button(props) { render_component("components/button", props) }
248                    fn render_Button(ctx, props) {
249                        text("button").with_style(ctx.component_style("root",
250                            style().height(px(24)).radius(px(0))))
251                    }
252                    fn view(ctx) {
253                        column([
254                            Button(#{ key: "global" }),
255                            Button(#{ key: "instance", style: style().height(px(40)) }),
256                        ])
257                    }
258                "#,
259            )
260            .unwrap();
261        let sheet = load_component_styles(
262            engine.engine(),
263            "styles.rhai",
264            r#"
265                fn component_styles() {
266                    #{ "components/button": #{
267                        root: style().height(px(34)).radius(px(6)),
268                    } }
269                }
270            "#,
271            &registry(),
272        )
273        .unwrap();
274        let mut state = UiRuntimeState::new();
275        state.replace_component_styles_from_host(sheet);
276        let state = Rc::new(RefCell::new(state));
277        let root_path = ComponentInstancePath::root("App", "styles");
278        let context = UiContext::new(
279            Rc::clone(&state),
280            root_path,
281            None,
282            ExecutionPhase::Render,
283            BTreeMap::new(),
284        )
285        .with_generation(compiled.generation());
286        let root = engine.render_with_context(&compiled, context).unwrap();
287        let crate::UiNodeKind::Box { children } = root.kind() else {
288            panic!("view must return a column");
289        };
290        assert_eq!(
291            children[0].style().base.height,
292            Some(crate::LayoutLength::Definite(crate::Length::Pixels(34.0)))
293        );
294        assert_eq!(
295            children[1].style().base.height,
296            Some(crate::LayoutLength::Definite(crate::Length::Pixels(40.0)))
297        );
298        assert_eq!(
299            children[0].style().base.radii.top_left,
300            Some(crate::Length::Pixels(6.0))
301        );
302    }
303}