#![allow(dead_code)]
use std::sync::OnceLock;
use regex::Regex;
#[derive(Debug, Default, Clone)]
pub struct PreprocessedSource {
pub text: String,
pub launches: Vec<LaunchSite>,
pub marker_offsets: Vec<MarkerOffset>,
}
#[derive(Debug, Clone)]
pub struct LaunchSite {
pub start: usize,
pub end: usize,
pub line: usize,
pub kernel: String,
pub grid: String,
pub block: String,
pub smem: Option<String>,
pub stream: Option<String>,
pub args: Vec<String>,
}
#[derive(Debug, Clone, Copy)]
pub struct MarkerOffset {
pub start: usize,
pub end: usize,
pub original_start: usize,
pub original_end: usize,
}
pub fn preprocess(source: &str) -> PreprocessedSource {
let re = match launch_re() {
Some(r) => r,
None => {
return PreprocessedSource {
text: source.to_string(),
..Default::default()
};
}
};
let skip = skip_spans_for(source);
let in_skip = |pos: usize| -> bool {
skip.iter().any(|(s, e)| pos >= *s && pos < *e)
};
let bytes = source.as_bytes();
let mut out = String::with_capacity(source.len());
let mut last = 0usize;
let mut launches = Vec::new();
let mut marker_offsets = Vec::new();
for caps in re.captures_iter(source) {
let m = caps.get(0).unwrap();
if in_skip(m.start()) {
continue;
}
let kernel = caps.name("k").unwrap().as_str().to_string();
let grid = caps.name("grid").unwrap().as_str().to_string();
let block = caps.name("block").unwrap().as_str().to_string();
let smem = caps.name("smem").map(|m| m.as_str().to_string());
let stream = caps.name("stream").map(|m| m.as_str().to_string());
let after_triple = m.end();
let mut p = after_triple;
while p < bytes.len() && bytes[p].is_ascii_whitespace() {
p += 1;
}
if p >= bytes.len() || bytes[p] != b'(' {
out.push_str(&source[last..m.end()]);
last = m.end();
continue;
}
let (arg_inner_end, args_inner) = match balanced_parens(source, p) {
Some(v) => v,
None => {
out.push_str(&source[last..m.end()]);
last = m.end();
continue;
}
};
let launch_end = arg_inner_end + 1;
let args = split_call_args_pub(args_inner)
.unwrap_or_else(|| vec![args_inner.to_string()]);
out.push_str(&source[last..m.start()]);
let line = 1 + source[..m.start()].bytes().filter(|b| *b == b'\n').count();
let marker = format!("__decuda_launch({kernel}, {grid}, {block}");
out.push_str(&marker);
if let Some(s) = &smem {
out.push_str(&format!(", smem={s}"));
}
if let Some(s) = &stream {
out.push_str(&format!(", stream={s}"));
}
out.push(')');
let new_end = out.len();
marker_offsets.push(MarkerOffset {
start: m.start(),
end: new_end,
original_start: m.start(),
original_end: launch_end,
});
launches.push(LaunchSite {
start: m.start(),
end: launch_end,
line,
kernel,
grid,
block,
smem,
stream,
args,
});
last = launch_end;
}
out.push_str(&source[last..]);
PreprocessedSource {
text: out,
launches,
marker_offsets,
}
}
fn skip_spans_for(source: &str) -> Vec<(usize, usize)> {
let mut out = Vec::new();
let bytes = source.as_bytes();
let mut i = 0usize;
while i < bytes.len() {
let b = bytes[i];
if b == b'/' && i + 1 < bytes.len() && bytes[i + 1] == b'/' {
let start = i;
while i < bytes.len() && bytes[i] != b'\n' {
i += 1;
}
out.push((start, i));
continue;
}
if b == b'/' && i + 1 < bytes.len() && bytes[i + 1] == b'*' {
let start = i;
i += 2;
while i + 1 < bytes.len() && !(bytes[i] == b'*' && bytes[i + 1] == b'/') {
i += 1;
}
i = (i + 2).min(bytes.len());
out.push((start, i));
continue;
}
if b == b'"' {
let start = i;
i += 1;
while i < bytes.len() {
if bytes[i] == b'\\' {
i += 2;
continue;
}
if bytes[i] == b'"' {
i += 1;
break;
}
i += 1;
}
out.push((start, i));
continue;
}
if b == b'\'' {
let start = i;
i += 1;
while i < bytes.len() {
if bytes[i] == b'\\' {
i += 2;
continue;
}
if bytes[i] == b'\'' {
i += 1;
break;
}
i += 1;
}
out.push((start, i));
continue;
}
i += 1;
}
out
}
pub fn balanced_parens(source: &str, i: usize) -> Option<(usize, &str)> {
let bytes = source.as_bytes();
if i >= bytes.len() || bytes[i] != b'(' {
return None;
}
let mut depth = 0i32;
let mut k = i;
while k < bytes.len() {
match bytes[k] {
b'(' => depth += 1,
b')' => {
depth -= 1;
if depth == 0 {
return Some((k, &source[i + 1..k]));
}
}
_ => {}
}
k += 1;
}
None
}
pub fn split_call_args_pub(s: &str) -> Option<Vec<String>> {
let s = s.trim();
if s.is_empty() {
return Some(Vec::new());
}
let mut out = Vec::new();
let mut depth = 0i32;
let mut angle = 0i32;
let mut bracket = 0i32;
let mut in_str = false;
let mut in_chr = false;
let mut esc = false;
let mut start = 0usize;
let bytes = s.as_bytes();
let mut i = 0usize;
while i < bytes.len() {
let c = bytes[i] as char;
if esc {
esc = false;
i += 1;
continue;
}
if in_str {
if c == '\\' {
esc = true;
} else if c == '"' {
in_str = false;
}
i += 1;
continue;
}
if in_chr {
if c == '\\' {
esc = true;
} else if c == '\'' {
in_chr = false;
}
i += 1;
continue;
}
match c {
'"' => in_str = true,
'\'' => in_chr = true,
'(' => depth += 1,
')' => {
depth -= 1;
if depth < 0 {
return None;
}
}
'<' => angle += 1,
'>' => angle -= 1,
'[' => bracket += 1,
']' => bracket -= 1,
',' if depth == 0 && angle == 0 && bracket == 0 => {
out.push(s[start..i].trim().to_string());
start = i + 1;
}
_ => {}
}
i += 1;
}
if depth != 0 || angle != 0 || bracket != 0 {
return None;
}
out.push(s[start..].trim().to_string());
Some(out)
}
fn launch_re() -> Option<&'static Regex> {
static R: OnceLock<Regex> = OnceLock::new();
Some(R.get_or_init(|| {
Regex::new(
r"(?x)
(?P<k>[A-Za-z_][A-Za-z0-9_:\s*&]*?)
\s*<<<\s*
(?P<grid>[^,]+?)
\s*,\s*
(?P<block>[^,]+?)
(?:\s*,\s*(?P<smem>[^,]+?))?
(?:\s*,\s*(?P<stream>[^,]+?))?
\s*>>>
",
)
.expect("kernel launch regex")
}))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn extracts_simple_launch() {
let src = "foo<<<grid, block>>>(a, b);";
let p = preprocess(src);
assert_eq!(p.launches.len(), 1);
let l = &p.launches[0];
assert_eq!(l.kernel, "foo");
assert_eq!(l.grid, "grid");
assert_eq!(l.block, "block");
assert!(l.smem.is_none());
assert_eq!(l.args, vec!["a".to_string(), "b".to_string()]);
assert!(p.text.contains("__decuda_launch(foo, grid, block)"));
assert!(!p.text.contains("<<<"));
}
#[test]
fn extracts_full_launch() {
let src = "k<<<g, b, 32, stream>>>(x, y);";
let p = preprocess(src);
assert_eq!(p.launches.len(), 1);
let l = &p.launches[0];
assert_eq!(l.smem.as_deref(), Some("32"));
assert_eq!(l.stream.as_deref(), Some("stream"));
}
#[test]
fn handles_nested_parens_in_args() {
let src = "k<<<1, 1>>>(foo(a, b), c);";
let p = preprocess(src);
let l = &p.launches[0];
assert_eq!(l.args, vec!["foo(a, b)".to_string(), "c".to_string()]);
}
#[test]
fn handles_no_args() {
let src = "k<<<1, 1>>>();";
let p = preprocess(src);
let l = &p.launches[0];
assert!(l.args.is_empty());
}
#[test]
fn ignores_comparisons() {
let src = "if (a < b && b < c) { x = 1; }";
let p = preprocess(src);
assert!(p.launches.is_empty());
assert!(p.text.contains("a < b && b < c"));
}
#[test]
fn multiple_launches_in_one_file() {
let src = "k1<<<1, 1>>>(a);\nk2<<<2, 2>>>(b, c);\n";
let p = preprocess(src);
assert_eq!(p.launches.len(), 2);
assert_eq!(p.launches[0].kernel, "k1");
assert_eq!(p.launches[1].kernel, "k2");
}
}