1use super::gpu_directive_parse_shared::{
34 directive_program_from_parse_with_source_layout, push_bounded_byte_scan_until,
35 push_c_identifier_span, push_directive_row_bounds, push_hash_and_keyword_start,
36 push_keyword_end, push_ws_skip_from_expr, safe_source_byte_expr,
37 trailing_ws_flag as is_trailing_ws, DirectiveOutputColumn, DirectiveSourceLayout,
38 DirectiveThreadLayout, MAX_DIRECTIVE_WS_PREFIX as MAX_WS_PREFIX,
39};
40use crate::parsing::c::lex::tokens::TOK_PP_DEFINE;
41use vyre::ir::{Expr, Node, Program};
42
43pub const OP_ID: &str = "vyre-libs::parsing::c::preprocess::gpu_define_parse";
45
46pub const BINDING_TOK_STARTS: u32 = 0;
48pub const BINDING_TOK_LENS: u32 = 1;
50pub const BINDING_DIRECTIVE_KINDS: u32 = 2;
52pub const BINDING_SOURCE: u32 = 3;
54pub const BINDING_NAME_START_OUT: u32 = 4;
56pub const BINDING_NAME_LEN_OUT: u32 = 5;
58pub const BINDING_ARGS_START_OUT: u32 = 6;
60pub const BINDING_ARGS_LEN_OUT: u32 = 7;
62pub const BINDING_BODY_START_OUT: u32 = 8;
64pub const BINDING_BODY_LEN_OUT: u32 = 9;
66pub const BINDING_IS_FUNCTION_LIKE_OUT: u32 = 10;
68
69const OUTPUT_COLUMNS: [DirectiveOutputColumn; 7] = [
70 DirectiveOutputColumn {
71 name: "name_start_out",
72 binding: BINDING_NAME_START_OUT,
73 },
74 DirectiveOutputColumn {
75 name: "name_len_out",
76 binding: BINDING_NAME_LEN_OUT,
77 },
78 DirectiveOutputColumn {
79 name: "args_start_out",
80 binding: BINDING_ARGS_START_OUT,
81 },
82 DirectiveOutputColumn {
83 name: "args_len_out",
84 binding: BINDING_ARGS_LEN_OUT,
85 },
86 DirectiveOutputColumn {
87 name: "body_start_out",
88 binding: BINDING_BODY_START_OUT,
89 },
90 DirectiveOutputColumn {
91 name: "body_len_out",
92 binding: BINDING_BODY_LEN_OUT,
93 },
94 DirectiveOutputColumn {
95 name: "is_function_like_out",
96 binding: BINDING_IS_FUNCTION_LIKE_OUT,
97 },
98];
99
100const DEFINE_KW_LEN: u32 = 6;
102
103#[must_use]
111pub fn gpu_define_parse(num_tokens: u32, source_len: u32) -> Program {
112 gpu_define_parse_with_source_layout(num_tokens, source_len, DirectiveSourceLayout::PackedU32)
113}
114
115#[must_use]
117pub fn gpu_define_parse_u8(num_tokens: u32, source_len: u32) -> Program {
118 gpu_define_parse_with_source_layout(num_tokens, source_len, DirectiveSourceLayout::RawU8)
119}
120
121fn gpu_define_parse_with_source_layout(
122 num_tokens: u32,
123 source_len: u32,
124 source_layout: DirectiveSourceLayout,
125) -> Program {
126 let t = Expr::var("t");
127 let safe_load = |addr: Expr| safe_source_byte_expr(source_layout, addr);
128
129 let mut parse: Vec<Node> = Vec::new();
130 push_directive_row_bounds(&mut parse);
131 push_hash_and_keyword_start(&mut parse, source_layout);
132 push_keyword_end(&mut parse, Expr::u32(DEFINE_KW_LEN));
133 push_ws_skip_from_expr(
134 &mut parse,
135 source_layout,
136 "np",
137 Expr::var("post_kw"),
138 "name_skip",
139 "name_start_val",
140 );
141 push_c_identifier_span(
142 &mut parse,
143 source_layout,
144 "name_start_val",
145 "name_len_val",
146 "name_done",
147 );
148
149 parse.push(Node::let_bind(
151 "after_name_idx",
152 Expr::add(Expr::var("name_start_val"), Expr::var("name_len_val")),
153 ));
154 parse.push(Node::let_bind(
155 "after_name_byte",
156 safe_load(Expr::var("after_name_idx")),
157 ));
158 parse.push(Node::let_bind(
159 "is_func_val",
160 Expr::select(
161 Expr::eq(Expr::var("after_name_byte"), Expr::u32(b'(' as u32)),
162 Expr::u32(1),
163 Expr::u32(0),
164 ),
165 ));
166
167 parse.push(Node::let_bind(
172 "args_start_val_raw",
173 Expr::add(Expr::var("after_name_idx"), Expr::u32(1)),
174 ));
175 push_bounded_byte_scan_until(
176 &mut parse,
177 source_layout,
178 "args_i",
179 "args_start_val_raw",
180 "args_scan_limit",
181 "args_byte",
182 "args_len_val_raw",
183 "args_done",
184 Expr::u32(b')' as u32),
185 Expr::eq(Expr::var("is_func_val"), Expr::u32(1)),
186 );
187
188 parse.push(Node::let_bind(
192 "body_pre_start",
193 Expr::select(
194 Expr::eq(Expr::var("is_func_val"), Expr::u32(1)),
195 Expr::select(
196 Expr::eq(Expr::var("args_done"), Expr::u32(1)),
197 Expr::add(
198 Expr::add(
199 Expr::var("args_start_val_raw"),
200 Expr::var("args_len_val_raw"),
201 ),
202 Expr::u32(1),
203 ),
204 Expr::var("tok_end"),
205 ),
206 Expr::var("after_name_idx"),
207 ),
208 ));
209 push_ws_skip_from_expr(
211 &mut parse,
212 source_layout,
213 "bp",
214 Expr::var("body_pre_start"),
215 "body_skip",
216 "body_start_val",
217 );
218
219 for q in 0..MAX_WS_PREFIX {
224 parse.push(Node::let_bind(
226 format!("tb_{q}"),
227 Expr::select(
228 Expr::lt(
229 Expr::add(Expr::var("body_start_val"), Expr::u32(q + 1)),
230 Expr::add(Expr::var("tok_end"), Expr::u32(1)),
231 ),
232 safe_load(Expr::sub(Expr::var("tok_end"), Expr::u32(q + 1))),
233 Expr::u32(0),
234 ),
235 ));
236 }
237 for q in 0..MAX_WS_PREFIX {
238 parse.push(Node::let_bind(
239 format!("tb_ws_{q}"),
240 is_trailing_ws(Expr::var(format!("tb_{q}"))),
241 ));
242 }
243 let trailing_ws_expr = {
247 let mut acc = Expr::u32(MAX_WS_PREFIX);
248 for q in (0..MAX_WS_PREFIX).rev() {
249 let mut prefix_ws = Expr::u32(1);
250 for r in 0..q {
251 prefix_ws = Expr::bitand(prefix_ws, Expr::var(format!("tb_ws_{r}")));
252 }
253 let tb_q_not_ws = Expr::select(
254 Expr::eq(Expr::var(format!("tb_ws_{q}")), Expr::u32(0)),
255 Expr::u32(1),
256 Expr::u32(0),
257 );
258 let cond_u32 = Expr::bitand(tb_q_not_ws, prefix_ws);
259 acc = Expr::select(Expr::eq(cond_u32, Expr::u32(1)), Expr::u32(q), acc);
260 }
261 acc
262 };
263 parse.push(Node::let_bind("trailing_ws_count", trailing_ws_expr));
264 parse.push(Node::let_bind(
266 "body_end_trimmed",
267 Expr::sub(Expr::var("tok_end"), Expr::var("trailing_ws_count")),
268 ));
269 parse.push(Node::let_bind(
270 "body_len_val",
271 Expr::select(
272 Expr::lt(Expr::var("body_start_val"), Expr::var("body_end_trimmed")),
273 Expr::sub(Expr::var("body_end_trimmed"), Expr::var("body_start_val")),
274 Expr::u32(0),
275 ),
276 ));
277
278 parse.push(Node::if_then(
283 Expr::and(
284 Expr::eq(Expr::var("found_hash"), Expr::u32(1)),
285 Expr::gt(Expr::var("name_len_val"), Expr::u32(0)),
286 ),
287 vec![
288 Node::store("name_start_out", t.clone(), Expr::var("name_start_val")),
289 Node::store("name_len_out", t.clone(), Expr::var("name_len_val")),
290 Node::store("body_start_out", t.clone(), Expr::var("body_start_val")),
291 Node::store("body_len_out", t.clone(), Expr::var("body_len_val")),
292 Node::store("is_function_like_out", t.clone(), Expr::var("is_func_val")),
293 Node::if_then(
294 Expr::and(
295 Expr::eq(Expr::var("is_func_val"), Expr::u32(1)),
296 Expr::eq(Expr::var("args_done"), Expr::u32(1)),
297 ),
298 vec![
299 Node::store("args_start_out", t.clone(), Expr::var("args_start_val_raw")),
300 Node::store("args_len_out", t.clone(), Expr::var("args_len_val_raw")),
301 ],
302 ),
303 ],
304 ));
305
306 directive_program_from_parse_with_source_layout(
307 OP_ID,
308 num_tokens,
309 source_len,
310 source_layout,
311 &OUTPUT_COLUMNS,
312 DirectiveThreadLayout::InvocationId,
313 Expr::eq(Expr::var("kind"), Expr::u32(TOK_PP_DEFINE)),
314 parse,
315 )
316}
317
318#[cfg(test)]
319mod tests {
320 use super::*;
321 use vyre::ir::DataType;
322
323 #[test]
324 fn op_id_is_canonical_and_stable() {
325 assert_eq!(OP_ID, "vyre-libs::parsing::c::preprocess::gpu_define_parse");
326 }
327
328 #[test]
329 fn build_program_returns_well_formed_program() {
330 let p = gpu_define_parse(8, 64);
331 assert_eq!(p.buffers().len(), 11);
332 assert_eq!(p.workgroup_size(), [256, 1, 1]);
333 }
334
335 #[test]
336 fn source_buffer_layouts_preserve_packed_abi_and_raw_u8_variant() {
337 let packed = gpu_define_parse(8, 64);
338 let raw_u8 = gpu_define_parse_u8(8, 64);
339 let packed_source = packed
340 .buffers()
341 .iter()
342 .find(|buffer| buffer.name() == "source")
343 .expect("Fix: packed define parser source buffer must exist");
344 let raw_u8_source = raw_u8
345 .buffers()
346 .iter()
347 .find(|buffer| buffer.name() == "source")
348 .expect("Fix: raw-U8 define parser source buffer must exist");
349
350 assert_eq!(packed_source.element(), DataType::U32);
351 assert_ne!(packed_source.count(), 0);
352 assert_eq!(raw_u8_source.element(), DataType::U8);
353 assert_eq!(raw_u8_source.count(), 0);
354 }
355}