Skip to main content

nmbrs_workload/edit/
locate.rs

1// Copyright 2024-2026 Jonathan Shook
2// SPDX-License-Identifier: Apache-2.0
3
4//! Tree-sitter-based YAML locator.
5//!
6//! Given the source bytes of a workload YAML and a target
7//! key path (e.g. `["scenarios", "test_oracles", "report",
8//! "cli_added"]`), find the byte range of:
9//!
10//! - the **value** of the deepest existing key in the path,
11//!   for replacement-style edits, OR
12//! - the **insertion point** for a missing key, for
13//!   add-style edits, with the indentation level the new
14//!   block must adopt.
15//!
16//! The CST preserves comments and whitespace as first-class
17//! nodes, so the byte ranges we return point at YAML data
18//! exclusively — the splicer can do
19//! `original[..start] + emitted + original[end..]` and
20//! every comment / blank line outside the range survives.
21//!
22//! ## Why not just walk a serde tree?
23//!
24//! `serde_yaml` discards comments + most whitespace and
25//! re-emits content with its own formatter. Round-trips are
26//! lossy. Tree-sitter's CST keeps every token in a node
27//! with start/end byte offsets — which is the exact tool
28//! for "find this subtree's bytes without serializing
29//! anything we didn't ask to change."
30
31use std::ops::Range;
32
33use tree_sitter::{Node, Parser, Tree};
34
35/// Outcome of locating a path in the source.
36#[derive(Debug, Clone)]
37pub enum Located {
38    /// The full path resolved to an existing value. The
39    /// returned range covers the value bytes (block scalar
40    /// body, mapping content, etc.) — splice in to replace.
41    Found { range: Range<usize> },
42    /// The path resolved up to `existing_depth`; the
43    /// remaining segments need to be created. `insert_at`
44    /// is the byte offset where new content should be
45    /// inserted (typically end-of-mapping at that depth).
46    /// `indent` is the column the new key/value must start
47    /// at (zero-based byte column of the parent mapping's
48    /// child keys).
49    Missing {
50        existing_depth: usize,
51        insert_at: usize,
52        indent: usize,
53    },
54}
55
56/// Parse `source` with tree-sitter-yaml. Returns the parsed
57/// tree; `source` must be kept alive alongside the tree
58/// because nodes hold byte offsets into it.
59pub fn parse(source: &str) -> Result<Tree, String> {
60    let mut parser = Parser::new();
61    let language = tree_sitter_yaml::LANGUAGE.into();
62    parser
63        .set_language(&language)
64        .map_err(|e| format!("tree-sitter-yaml language load failed: {e}"))?;
65    parser
66        .parse(source, None)
67        .ok_or_else(|| "tree-sitter-yaml parse returned no tree".to_string())
68}
69
70/// Locate `path` in the parsed tree. `path` is a sequence
71/// of mapping keys (the YAML form of a JSONPath); each
72/// segment names a child of the previous segment's mapping.
73///
74/// On success returns either [`Located::Found`] (the path
75/// fully resolved) or [`Located::Missing`] (path resolved
76/// up to a point; new keys need to be inserted at the
77/// returned offset).
78pub fn locate_path(tree: &Tree, source: &str, path: &[&str]) -> Result<Located, String> {
79    let root = tree.root_node();
80    // tree-sitter-yaml's root is `stream`, with one or more
81    // `document` children. The first document holds the
82    // top-level mapping.
83    let document = first_named_child_kind(root, "document")
84        .ok_or_else(|| "yaml has no document".to_string())?;
85    let top =
86        first_block_node_under(document).ok_or_else(|| "yaml document is empty".to_string())?;
87
88    walk_path(top, source, path, 0)
89}
90
91fn first_named_child_kind<'t>(n: Node<'t>, kind: &str) -> Option<Node<'t>> {
92    let mut cursor = n.walk();
93    n.named_children(&mut cursor).find(|c| c.kind() == kind)
94}
95
96/// Skip non-mapping wrapper nodes (e.g. `block_node`,
97/// `flow_node`) to land on the actual `block_mapping` /
98/// `flow_mapping`. tree-sitter-yaml wraps mapping content
99/// in node-type wrappers; we always want the inner
100/// mapping when traversing keys.
101fn first_block_node_under(n: Node<'_>) -> Option<Node<'_>> {
102    if n.kind() == "block_mapping" || n.kind() == "flow_mapping" {
103        return Some(n);
104    }
105    let mut cursor = n.walk();
106    for child in n.named_children(&mut cursor) {
107        if let Some(found) = first_block_node_under(child) {
108            return Some(found);
109        }
110    }
111    None
112}
113
114fn walk_path(
115    mapping: Node<'_>,
116    source: &str,
117    path: &[&str],
118    depth: usize,
119) -> Result<Located, String> {
120    if path.is_empty() {
121        return Ok(Located::Found {
122            range: mapping.byte_range(),
123        });
124    }
125
126    // Walk every `block_mapping_pair` (or `flow_pair`) child
127    // of this mapping looking for a key matching path[0].
128    let key_to_find = path[0];
129    let mut cursor = mapping.walk();
130    let mut last_pair_end: Option<usize> = None;
131    let mut child_indent: Option<usize> = None;
132
133    for pair in mapping.named_children(&mut cursor) {
134        if pair.kind() != "block_mapping_pair" && pair.kind() != "flow_pair" {
135            continue;
136        }
137        let (key_node, value_node) = pair_key_value(pair)
138            .ok_or_else(|| format!("malformed mapping pair at byte {}", pair.start_byte(),))?;
139        let key_text = node_text(key_node, source)
140            .trim()
141            .trim_matches(|c| c == '"' || c == '\'');
142        last_pair_end = Some(pair.end_byte());
143        if child_indent.is_none() {
144            child_indent = Some(pair.start_position().column);
145        }
146        if key_text == key_to_find {
147            // Recurse into this value if there are more
148            // path segments. Otherwise we found the target.
149            if path.len() == 1 {
150                return Ok(Located::Found {
151                    range: value_byte_range(value_node, source),
152                });
153            }
154            // Need to recurse — the value should itself be
155            // a mapping. Strip wrappers (`block_node` etc.)
156            // until we hit `block_mapping`.
157            let inner = first_block_node_under(value_node);
158            return match inner {
159                Some(m) => walk_path(m, source, &path[1..], depth + 1),
160                None => {
161                    // Value isn't a mapping (could be a
162                    // scalar, sequence, or null). The
163                    // path can't continue — treat this
164                    // segment as missing-from-here.
165                    let column = value_node.start_position().column;
166                    Ok(Located::Missing {
167                        existing_depth: depth + 1,
168                        insert_at: value_node.end_byte(),
169                        indent: column,
170                    })
171                }
172            };
173        }
174    }
175
176    // Key not found at this level. Return a Missing with
177    // the insertion point at end-of-mapping plus the
178    // sibling-key indent so the new key aligns.
179    let insert_at = last_pair_end.unwrap_or_else(|| mapping.end_byte());
180    let indent = child_indent.unwrap_or_else(|| mapping.start_position().column);
181    Ok(Located::Missing {
182        existing_depth: depth,
183        insert_at,
184        indent,
185    })
186}
187
188fn pair_key_value<'t>(pair: Node<'t>) -> Option<(Node<'t>, Node<'t>)> {
189    // tree-sitter-yaml exposes `key:` and `value:` named
190    // fields on block_mapping_pair / flow_pair.
191    let key = pair.child_by_field_name("key")?;
192    let value = pair.child_by_field_name("value")?;
193    Some((key, value))
194}
195
196fn node_text<'a>(n: Node<'_>, source: &'a str) -> &'a str {
197    &source[n.byte_range()]
198}
199
200/// Compute the splice range for a value node.
201///
202/// For block scalars (`|`, `>`, multi-line strings), the
203/// range covers the entire scalar including the indicator
204/// and continuation. For flow scalars and primitives it
205/// covers exactly the scalar's bytes. For mapping / list
206/// values the range covers all child content.
207///
208/// The returned range trims a single trailing newline if
209/// present, so the splicer can append a new newline of its
210/// own without doubling.
211fn value_byte_range(value: Node<'_>, source: &str) -> Range<usize> {
212    let mut r = value.byte_range();
213    // Trim a single trailing newline if present — keeps
214    // splice composition predictable.
215    if r.end > r.start && source.as_bytes().get(r.end - 1) == Some(&b'\n') {
216        r.end -= 1;
217    }
218    r
219}
220
221#[cfg(test)]
222mod tests {
223    use super::*;
224
225    fn loc(yaml: &str, path: &[&str]) -> Located {
226        let tree = parse(yaml).expect("parse");
227        locate_path(&tree, yaml, path).expect("locate")
228    }
229
230    #[test]
231    fn locate_root_key_value_finds_block_scalar_body() {
232        let yaml = "scenarios:\n  default:\n    - phase: setup\n";
233        let r = loc(yaml, &["scenarios"]);
234        match r {
235            Located::Found { range } => {
236                let text = &yaml[range];
237                assert!(
238                    text.contains("default"),
239                    "should cover scenarios value, got: {text:?}"
240                );
241            }
242            other => panic!("expected Found, got {other:?}"),
243        }
244    }
245
246    #[test]
247    fn locate_missing_root_key_returns_insert_point_at_eof_of_mapping() {
248        let yaml = "scenarios:\n  default: [a]\n";
249        let r = loc(yaml, &["report"]);
250        match r {
251            Located::Missing {
252                existing_depth,
253                insert_at,
254                indent,
255            } => {
256                assert_eq!(existing_depth, 0);
257                assert_eq!(indent, 0, "root-level keys insert at column 0");
258                // Insert position should be at end of last
259                // root pair (after `default: [a]\n`-ish).
260                assert!(
261                    insert_at >= yaml.len() - 1,
262                    "insert_at {insert_at} should be near eof {}",
263                    yaml.len()
264                );
265            }
266            other => panic!("expected Missing, got {other:?}"),
267        }
268    }
269
270    #[test]
271    fn locate_nested_key_traverses_mappings() {
272        let yaml = r#"
273report:
274  intro:
275    text: hello
276  recall_block:
277    plot: r1
278"#;
279        let r = loc(yaml, &["report", "recall_block"]);
280        match r {
281            Located::Found { range } => {
282                let text = &yaml[range];
283                assert!(
284                    text.contains("plot: r1"),
285                    "expected recall_block body, got: {text:?}"
286                );
287            }
288            other => panic!("expected Found, got {other:?}"),
289        }
290    }
291
292    #[test]
293    fn locate_missing_nested_key_returns_insert_at_parent_end() {
294        let yaml = r#"
295report:
296  intro:
297    text: hello
298"#;
299        let r = loc(yaml, &["report", "cli_added"]);
300        match r {
301            Located::Missing {
302                existing_depth,
303                insert_at,
304                indent,
305            } => {
306                assert_eq!(existing_depth, 1, "report exists, cli_added doesn't");
307                // Indentation should match the existing
308                // sibling key (`intro:`) — column 2.
309                assert_eq!(indent, 2, "child keys of `report:` are at column 2");
310                // insert_at should land after `text: hello\n`-ish.
311                let prefix = &yaml[..insert_at];
312                assert!(prefix.contains("hello"));
313            }
314            other => panic!("expected Missing, got {other:?}"),
315        }
316    }
317
318    #[test]
319    fn locate_path_through_nonmapping_value_returns_missing() {
320        // `report` is set to a scalar — can't recurse into it.
321        let yaml = "report: not_a_mapping\n";
322        let r = loc(yaml, &["report", "cli_added"]);
323        match r {
324            Located::Missing { existing_depth, .. } => {
325                assert_eq!(existing_depth, 1);
326            }
327            other => panic!("expected Missing, got {other:?}"),
328        }
329    }
330
331    #[test]
332    fn locate_preserves_byte_offsets_for_splice() {
333        let yaml = "a: 1\nb: 2\nc: 3\n";
334        let r = loc(yaml, &["b"]);
335        match r {
336            Located::Found { range } => {
337                let prefix = &yaml[..range.start];
338                let suffix = &yaml[range.end..];
339                let value_text = &yaml[range];
340                // Round-trip: prefix + new_value + suffix
341                // should be a valid yaml with `b` replaced.
342                assert_eq!(value_text, "2");
343                let spliced = format!("{prefix}99{suffix}");
344                assert_eq!(spliced, "a: 1\nb: 99\nc: 3\n");
345            }
346            other => panic!("expected Found, got {other:?}"),
347        }
348    }
349}