//! Tree-sitter syntax highlighting: one style byte per content byte, filled by //! running each grammar's highlights.scm query (slurped at build time into the //! ts_queries options module). Grammar set is tiered: `zig`; `minimal` (c, cpp, //! zig); and `full`, which adds ~24 languages lazily on first use. const std = @import("std"); const config = @import("pardes_config"); const tracy = @import("tracy.zig"); const grammar_manifest = @import("grammar_manifest.zig"); pub const enabled = config.syntax_highlighting; const zig_grammar = config.syntax_zig_grammar; const minimal_grammars = config.syntax_minimal_grammars; const full_grammars = config.syntax_full_grammars; const ts = if (enabled) @import("tree-sitter") else struct { pub const Language = opaque {}; pub const Query = opaque {}; }; const ts_queries = if (enabled) @import("ts_queries") else struct {}; const LanguageFn = *const fn () callconv(.c) *const ts.Language; pub const Syn = enum(u8) { none, keyword, string, number, comment }; const Spec = struct { name: []const u8, exts: []const []const u8, language: *const fn () callconv(.c) *const ts.Language, query_src: []const u8, compiled_query: ?*ts.Query = null, }; fn grammarSelected(comptime g: grammar_manifest.Grammar) bool { return switch (g.tier) { .zig => zig_grammar, .minimal => minimal_grammars, .full => full_grammars, }; } fn specCount() comptime_int { var count = 0; for (grammar_manifest.all) |g| { if (grammarSelected(g)) count += 1; } return count; } fn initSpecs() [specCount()]Spec { var out: [specCount()]Spec = undefined; var i = 0; inline for (grammar_manifest.all) |g| { if (grammarSelected(g)) { out[i] = .{ .name = g.name, .exts = g.exts, .language = @extern(LanguageFn, .{ .name = "tree_sitter_" ++ g.name }), .query_src = @field(ts_queries, g.name ++ "_highlights"), }; i += 1; } } return out; } var specs = initSpecs(); const allocation_header_size = 16; var syntax_allocator: std.mem.Allocator = undefined; var syntax_started = false; extern fn ts_set_allocator( new_malloc: ?*const fn (size: usize) callconv(.c) ?*anyopaque, new_calloc: ?*const fn (nmemb: usize, size: usize) callconv(.c) ?*anyopaque, new_realloc: ?*const fn (ptr: ?*anyopaque, size: usize) callconv(.c) ?*anyopaque, new_free: ?*const fn (ptr: ?*anyopaque) callconv(.c) void, ) void; fn syntaxAlloc(size_arg: usize) callconv(.c) ?*anyopaque { const size = @max(size_arg, 1); const total = std.math.add(usize, allocation_header_size, size) catch return null; const bytes = syntax_allocator.alignedAlloc(u8, .@"16", total) catch return null; const header: *align(16) usize = @ptrCast(bytes.ptr); header.* = total; return @ptrCast(bytes.ptr + allocation_header_size); } fn syntaxCalloc(count: usize, size: usize) callconv(.c) ?*anyopaque { const len = std.math.mul(usize, count, size) catch return null; const pointer = syntaxAlloc(len) orelse return null; const bytes: [*]u8 = @ptrCast(pointer); @memset(bytes[0..len], 0); return pointer; } fn syntaxRealloc(ptr: ?*anyopaque, new_size: usize) callconv(.c) ?*anyopaque { const pointer = ptr orelse return syntaxAlloc(new_size); if (new_size == 0) { syntaxFree(pointer); return null; } const user: [*]u8 = @ptrCast(pointer); const header: *align(16) usize = @ptrCast(@alignCast(user - allocation_header_size)); const old_total = header.*; const old_bytes: []align(16) u8 = @as([*]align(16) u8, @ptrCast(header))[0..old_total]; const new_total = std.math.add(usize, allocation_header_size, new_size) catch return null; const new_bytes = syntax_allocator.realloc(old_bytes, new_total) catch return null; const new_header: *align(16) usize = @ptrCast(new_bytes.ptr); new_header.* = new_total; return @ptrCast(new_bytes.ptr + allocation_header_size); } fn syntaxFree(ptr: ?*anyopaque) callconv(.c) void { const pointer = ptr orelse return; const user: [*]u8 = @ptrCast(pointer); const header: *align(16) usize = @ptrCast(@alignCast(user - allocation_header_size)); const total = header.*; const bytes: []align(16) u8 = @as([*]align(16) u8, @ptrCast(header))[0..total]; syntax_allocator.free(bytes); } pub fn start(gpa: std.mem.Allocator) void { if (comptime enabled) { std.debug.assert(!syntax_started); syntax_allocator = gpa; syntax_started = true; ts_set_allocator(syntaxAlloc, syntaxCalloc, syntaxRealloc, syntaxFree); } } pub fn stop() void { if (comptime enabled) { std.debug.assert(syntax_started); for (&specs) |*spec| { if (spec.compiled_query) |query| query.destroy(); spec.compiled_query = null; } ts_set_allocator(null, null, null, null); syntax_started = false; syntax_allocator = undefined; } } const Selected = struct { name: []const u8, lang: *const ts.Language, query: *ts.Query }; // NOTE: don't lang.destroy() — the tree_sitter_*() languages are static // singletons reused on every open; destroying one use-after-frees the next. fn ensure(spec: *Spec) !Selected { const lang = spec.language(); if (spec.compiled_query) |query| return .{ .name = spec.name, .lang = lang, .query = query }; var error_offset: u32 = 0; const query = try ts.Query.create(lang, spec.query_src, &error_offset); spec.compiled_query = query; return .{ .name = spec.name, .lang = lang, .query = query }; } fn forExt(ext: []const u8) !?Selected { for (&specs) |*spec| { for (spec.exts) |choice| { if (std.ascii.eqlIgnoreCase(ext, choice)) return try ensure(spec); } } return null; } const LangAlias = struct { []const u8, []const u8 }; const lang_aliases = [_]LangAlias{ .{ "js", "javascript" }, .{ "jsx", "javascript" }, .{ "mjs", "javascript" }, .{ "py", "python" }, .{ "sh", "bash" }, .{ "shell", "bash" }, .{ "zsh", "bash" }, .{ "bash", "bash" }, .{ "rs", "rust" }, .{ "c++", "cpp" }, .{ "cxx", "cpp" }, .{ "cc", "cpp" }, .{ "cs", "c_sharp" }, .{ "kt", "kotlin" }, .{ "rb", "ruby" }, .{ "ml", "ocaml" }, .{ "hs", "haskell" }, }; fn forLang(name: []const u8) !?Selected { var canonical = name; for (lang_aliases) |a| { if (std.ascii.eqlIgnoreCase(name, a[0])) { canonical = a[1]; break; } } for (&specs) |*spec| { if (std.ascii.eqlIgnoreCase(canonical, spec.name)) return try ensure(spec); } return null; } fn runQuery(styles: []u8, sel: Selected, tree: *ts.Tree, base: usize) void { const cursor = ts.QueryCursor.create(); defer cursor.destroy(); cursor.exec(sel.query, tree.rootNode()); while (cursor.nextMatch()) |match| { for (match.captures) |cap| { const syn = synFor(sel.query.captureNameForId(cap.index) orelse ""); if (syn == .none) continue; var b: usize = base + cap.node.startByte(); const end = @min(base + @as(usize, cap.node.endByte()), styles.len); while (b < end) : (b += 1) styles[b] = @intFromEnum(syn); } } } // Queries compile on first use (`ensure`), never at startup. Pre-compiling the // compact tier in Pardes.init cost EVERY boot ~120ms of ts_query__perform_analysis // (55% of a Debug startup) to save ~40ms on the first .zig/.c/.cpp open — a pane // of prose or a shell paid for a language it never opened. Grammar availability // is unchanged; only the timing moved. fn synFor(name: []const u8) Syn { for ([_]struct { []const u8, Syn }{ .{ "comment", .comment }, .{ "string", .string }, .{ "character", .string }, .{ "number", .number }, .{ "numeric", .number }, .{ "float", .number }, .{ "boolean", .number }, .{ "keyword", .keyword }, .{ "include", .keyword }, .{ "conditional", .keyword }, .{ "repeat", .keyword }, .{ "title", .keyword }, .{ "uri", .string }, .{ "reference", .number }, }) |m| { if (std.mem.indexOf(u8, name, m[0]) != null) return m[1]; } return .none; } /// One Syn byte per content byte in [start, end). Caller frees. pub fn highlightFileRange(gpa: std.mem.Allocator, path: []const u8, content: []const u8, start_byte_raw: usize, end_byte_raw: usize) ![]u8 { const tz = tracy.zone(@src(), "highlightFileRange"); defer tz.end(); if (!enabled) return &.{}; const ext = std.fs.path.extension(path); const selected = (forExt(ext) catch return &.{}) orelse return &.{}; const start_byte = @min(start_byte_raw, content.len); const end_byte = @max(start_byte, @min(end_byte_raw, content.len)); const source = content[start_byte..end_byte]; const styles = try gpa.alloc(u8, source.len); errdefer gpa.free(styles); @memset(styles, 0); const parser = ts.Parser.create(); defer parser.destroy(); parser.setLanguage(selected.lang) catch return styles; const tree = parser.parseString(source, null) orelse return styles; defer tree.destroy(); runQuery(styles, selected, tree, 0); if (std.mem.eql(u8, selected.name, "markdown")) injectCodeBlocks(styles, source, tree.rootNode()); return styles; } fn injectCodeBlocks(styles: []u8, source: []const u8, node: ts.Node) void { if (std.mem.eql(u8, node.kind(), "fenced_code_block")) { highlightCodeBlock(styles, source, node); return; } var i: u32 = 0; const count = node.childCount(); while (i < count) : (i += 1) { if (node.child(i)) |c| injectCodeBlocks(styles, source, c); } } fn childOfKind(node: ts.Node, kind: []const u8) ?ts.Node { var i: u32 = 0; const count = node.childCount(); while (i < count) : (i += 1) { if (node.child(i)) |c| { if (std.mem.eql(u8, c.kind(), kind)) return c; } } return null; } fn highlightCodeBlock(styles: []u8, source: []const u8, block: ts.Node) void { const info = childOfKind(block, "info_string") orelse return; const lang_node = childOfKind(info, "language") orelse return; const content_node = childOfKind(block, "code_fence_content") orelse return; const langtext = source[lang_node.startByte()..lang_node.endByte()]; const sub_sel = (forLang(langtext) catch return) orelse return; const cs: usize = content_node.startByte(); const ce: usize = content_node.endByte(); if (ce > source.len or cs > ce) return; const parser = ts.Parser.create(); defer parser.destroy(); parser.setLanguage(sub_sel.lang) catch return; const tree = parser.parseString(source[cs..ce], null) orelse return; defer tree.destroy(); runQuery(styles, sub_sel, tree, cs); } /// One Syn byte per byte of content[start, end) for unified diffs/patches. /// Pure byte scan; independent of tree-sitter and the `enabled` flag. Caller frees. pub fn highlightDiff(gpa: std.mem.Allocator, content: []const u8, start_byte_raw: usize, end_byte_raw: usize) ![]u8 { const start_byte = @min(start_byte_raw, content.len); const end_byte = @max(start_byte, @min(end_byte_raw, content.len)); const source = content[start_byte..end_byte]; const styles = try gpa.alloc(u8, source.len); errdefer gpa.free(styles); @memset(styles, 0); var offset: usize = 0; var lines = std.mem.splitScalar(u8, source, '\n'); while (lines.next()) |line| { const syn = diffLineSyn(line); if (syn != .none) @memset(styles[offset .. offset + line.len], @intFromEnum(syn)); offset += line.len + 1; } return styles; } fn diffLineSyn(line: []const u8) Syn { if (std.mem.startsWith(u8, line, "@@")) return .keyword; if (std.mem.startsWith(u8, line, "+++") or std.mem.startsWith(u8, line, "---") or std.mem.startsWith(u8, line, "diff ") or std.mem.startsWith(u8, line, "index ") or std.mem.startsWith(u8, line, "\\ No newline")) return .comment; if (line.len == 0) return .none; if (line[0] == '+') return .string; if (line[0] == '-') return .number; return .none; } test "tree-sitter allocator callbacks preserve and free exact allocations" { syntax_allocator = std.testing.allocator; defer syntax_allocator = undefined; var live: ?*anyopaque = syntaxCalloc(4, 1) orelse return error.OutOfMemory; defer if (live) |pointer| syntaxFree(pointer); const original: [*]u8 = @ptrCast(live.?); try std.testing.expectEqualSlices(u8, &.{ 0, 0, 0, 0 }, original[0..4]); @memcpy(original[0..4], "data"); live = syntaxRealloc(live, 32) orelse return error.OutOfMemory; const grown: [*]u8 = @ptrCast(live.?); try std.testing.expectEqualSlices(u8, "data", grown[0..4]); try std.testing.expect(syntaxCalloc(std.math.maxInt(usize), 2) == null); try std.testing.expect(syntaxRealloc(live, 0) == null); live = null; } test "default full grammar set highlights Typst source" { if (!enabled or !full_grammars) return; start(std.testing.allocator); defer stop(); const source = "// note\n#let answer = 42\n#let text = \"hello\"\n"; const styles = try highlightFileRange(std.testing.allocator, "paper.typst", source, 0, source.len); defer std.testing.allocator.free(styles); const comment_at = std.mem.indexOf(u8, source, "// note").?; const keyword_at = std.mem.indexOf(u8, source, "let").?; const number_at = std.mem.indexOf(u8, source, "42").?; const string_at = std.mem.indexOf(u8, source, "\"hello\"").?; try std.testing.expectEqual(Syn.comment, @as(Syn, @enumFromInt(styles[comment_at]))); try std.testing.expectEqual(Syn.keyword, @as(Syn, @enumFromInt(styles[keyword_at]))); try std.testing.expectEqual(Syn.number, @as(Syn, @enumFromInt(styles[number_at]))); try std.testing.expectEqual(Syn.string, @as(Syn, @enumFromInt(styles[string_at]))); const short_ext = try highlightFileRange(std.testing.allocator, "paper.typ", source, 0, source.len); defer std.testing.allocator.free(short_ext); try std.testing.expectEqual(Syn.keyword, @as(Syn, @enumFromInt(short_ext[keyword_at]))); } test "markdown highlights headings and injects fenced code blocks" { if (!enabled or !full_grammars) return; start(std.testing.allocator); defer stop(); const source = "# Heading\n\n```zig\nconst answer = 42;\n```\n"; const styles = try highlightFileRange(std.testing.allocator, "doc.md", source, 0, source.len); defer std.testing.allocator.free(styles); const title_at = std.mem.indexOf(u8, source, "Heading").?; const num_at = std.mem.indexOf(u8, source, "42").?; // the heading is markdown's own highlight; the number proves the zig // grammar was injected into the fenced block and colored its contents. try std.testing.expectEqual(Syn.keyword, @as(Syn, @enumFromInt(styles[title_at]))); try std.testing.expectEqual(Syn.number, @as(Syn, @enumFromInt(styles[num_at]))); } test "highlightDiff colors unified diff lines by prefix" { const diff = "diff --git a/x b/x\n" ++ "--- a/x\n" ++ "+++ b/x\n" ++ "@@ -1,3 +1,3 @@\n" ++ " context\n" ++ "-old line\n" ++ "+new line\n"; const styles = try highlightDiff(std.testing.allocator, diff, 0, diff.len); defer std.testing.allocator.free(styles); const byteSyn = struct { fn at(s: []const u8, src: []const u8, needle: []const u8) Syn { const i = std.mem.indexOf(u8, src, needle).?; return @enumFromInt(s[i]); } }.at; try std.testing.expectEqual(Syn.comment, byteSyn(styles, diff, "diff --git")); try std.testing.expectEqual(Syn.comment, byteSyn(styles, diff, "--- a/x")); try std.testing.expectEqual(Syn.comment, byteSyn(styles, diff, "+++ b/x")); try std.testing.expectEqual(Syn.keyword, byteSyn(styles, diff, "@@ -1,3")); try std.testing.expectEqual(Syn.none, byteSyn(styles, diff, " context")); try std.testing.expectEqual(Syn.number, byteSyn(styles, diff, "-old line")); try std.testing.expectEqual(Syn.string, byteSyn(styles, diff, "+new line")); }