From 43623ec9bf0eaaf7113285c46e8a09018f181b18 Mon Sep 17 00:00:00 2001 From: Max Brunsfeld Date: Thu, 20 Aug 2026 14:02:11 -0700 Subject: [PATCH] Upgrade to latest wasmtime, improve robustness and perf of wasm-based parsing (#5847) * build(deps): upgrade wasmtime C API to 48.0.0 Upgrade the Rust and Zig Wasmtime dependencies and enable reference values with the null GC collector. Wasmtime 48 requires a newer Rust toolchain, whose Clippy version identifies three item helpers that can be const. Mark them const so the workspace continues to pass Clippy with warnings denied. * feat(benchmark): support Wasm grammars * perf(wasm): cache language function handles * fix(wasm): improve validation and failure cleanup Bounds-check dylink metadata parsing, require exact import and export names, and restore memory and function-table allocation offsets when language loading fails. * fix(wasm): copy the complete supertype map --- Cargo.lock | Bin 69613 -> 71934 bytes build.zig.zon | 64 +++---- crates/cli/benches/benchmark.rs | 40 ++++- crates/generate/src/build_tables/item.rs | 6 +- crates/xtask/src/benchmark.rs | 16 ++ crates/xtask/src/main.rs | 3 + lib/Cargo.toml | 4 +- lib/src/wasm_store.c | 211 ++++++++++++++--------- 8 files changed, 225 insertions(+), 119 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 0acd8cec26c9b27c0d4252defd2d5f2389b2b149..24fa7ee44a04507b2ef4d1023cb260e6dda4bad0 100644 GIT binary patch delta 5745 zcmb_gYiylm9naaN+qyBv)}FSzxeXWtYc88VNzvNX$Nvpl;w`L@qK|_~L>RfA49}TnB#hL)UXoci!ju z-+q_>`#)Ip$4Q}$@(a<|J;fANSwC_(7lfz@9o3?M-Li%W*Ewp2Bz_6@b zN{$)jP4Xf}E;x16C+ek*%n}z8AzAtyO=6Bp>jQH{CKjktx*UG`=0?}@{srB=emdtg zap5_nkTlfKi8etiX1!vOXCkSp#vU2*@UFC(%Ys<_d zx1d%jO2bV`-ZSMG{!P+k6{VLP6*1N&8;SAGQ<6~4NC{6$;M#~_dc8>Cbu3oTG4;aNa}s7;y&%@nd3j8gbi$bjThX^^ia z-Q~W^+RHaD{o33E2aGd;D;gqWIg%VvPDNc-RbiFG86+r{wLv%0;1GySB3X^IvLZ64 z@_ee7`!Bn3?s@&*sKA9w)bNm*KvdLZgK7$s6HRCccCsHDB+CTet-p#rBymqR&r zF#^XG|9$n6vTEC+@}*O2%0p#C`CMOn`O)R|)-i(0jDBQF8KZ2XXcq33avHONGKFNk zPbx`d)$a-(@lA;n;1&cTVy zMZ^FgcqNfk<01;Cty7dsCh&@@EIICxvr}AhUxp6V%bi!e-l6;Gme03u+OluhkIcRp z+g)~?ysG4FSCj`&Ub1kHW98*{E-%lXTvVRjwzABWjm?M;Ze3E|-`3f>5TfOR5J@eH zmTcnDSTw5*nMy$gwSq9`wPcdG?2~hxG7~+zj!FDU&J}AJ8lbcM;mXHb+0;l?Mqdk} zHL(;ep9KZjgn;2^72+rnT@v0ns-l(XY?*lu(Q$6FJaD*H>u%n6)z#;c=LC}`iWJdB zsIINbpHmQcoToY9b}9onbnr5TYW97M*>e+|;nDI`=v(!2?CR@Vv5QQo@+upA-3jG* zmN5f%42uTEMhSdoNPN@=J)!Z(AR@jXV?uRUJQw?sN0*oBlk4Vn*7Dx&RxUiny*w2b z0juOV`pc&T=oJEaSC|Tnh^A<^tWj%5olk-b;4(GVr}CTBg?U(3wqLd6yJr@!+T$ko z_=&N5Mgd3^eNayH-qi{z6V$}y%p=!XvJ^F^F$35cfK(}}B_`Pzp`wo%D3uZ;xQW#~ zQkIe~$NF^Z+UQQ2yiUYaFC;({K#U7YT+0*ABnr$?2*&Y)hX@83N`!i9obVD*dfv6$ z%I~fjY~7M2;*|!1Amfxs2E-u&`Kjg{b;tx+=7doT=yKMHKy1)D1>78PN;~dAMPs9p1B#yj zB4~xENX2M9^YFhpTvPXsr_t$=GIRO5#t-Rt7Ip1Tlb<3w(x(euKG>fk5Z=(Nkugk; zRsbTQfXj+X2SCA`ha52v=po#LIA@q9)Nxz{SFc8V!@jYwx5abSn1}NtZBo+J>Q~Df zN3VaNZS@5V#QS)mX>h<|R3k%Ed*?86E;EixU?eONI5HUu3=epX@F_ttG)2cm5C9TJ z94~I?#@qI|(b2T8nS(q%H9e6gk?n~z?}Op7d2c6&NB7SAV9HI-|6qD*cwafXbA9>n zhPIwCcf@?vFc(BW=DNZpRSA;>ga?eI2wICoQ>Rt7j&kp4w3Oh#Na>+goKv34go20S z{#h^kclOjK%lO!aa{7j)W%Uzl7Pg*l8J=u0%nYRfpd?NR_g173a4v|DJBfyt&gmed zt{6s17YRiYMFxF-tGjw;2Qu(k*M}MjBs)W;#;|cbI_DgnqXEJO0-<{A(Rc)Gs60to z2tk#!lnJgm0xi#XbX5n`n>qcgBL)Y6&_Xbft43ZjjTE>+j5ZS3%6ZCBL$`4i1%Os+ zZ#YIP<$wgmWD;$iv9Z}RAqtAS0t^tO%oXuOAXfr90{k?jg=Q%ml@S7<98(G4V#b4> zAv1F8VN1&|Z>(1j>?e>_4yr~$;CLrvR0#|yBcPcGfRlrW5s3!GpoJp}IcAvJkf`w5 zV$eUX-+r#~{;$nAl19RwiALFZX>H?qa_@pdQh}kqyM5PnMk-;0lL%5XrlI9F(;Al40}>2Go$Mo&waVGR^yc* z>nEsXC0wC#q_hU&0*ba0u7cw*>p@8b1B$$+BveR3Qw36#9*04`R5|O;qx+lZ%;O5x zM;kLYEnYAu{*6GlpoktUkuA_6^D6figpFQB0}6+Bf`2K@{IBp5DSv;syj57NY}5fXh1 zx|JRP^o4Y44Q5IFHciX1D^2_B^h zlwvWukoriWVw$6qDp9%pp|uM^n2oPLbZEr_Van0%xAa5IN^fy98B{PiSQ#BebdFvU zVo8KrO(0Dj5{OlTT|tywlL-0k5J(eYDs_3}QHWt^tVNHloqP0v1*Qa0ZY9wTzTcE^ zuw2#jIYu%3*`u<^Hq4hJEJ4UeNM)p*(LnTpGb?MiHkLkq+v1umcRxHd05Z{_=}J#a zTn$4F#Ud3KhHEJ$Hy%($Ii?LCAWZ@y+!`=7sT$G9rU2?2T3yy1=&2G>^nnfK;cXp_ zI}WI2>;8wMp?xq9?t`K0!08G)6db|~Y!;r069b%Kfbnn`UKTw8<3Ob;0EZ{m1MHH3 zOeCpU&@y;c_#Vrud!f#o*Nso-+iNqk1nzImx!P;f$7wa3{&pgmwe)i(la{RYX zl@~fX%aNCtH(q?HXX&{aJs)e408);Aa`Sc2;>lyc91Z_ zBVZ9&1^JO1fQQ|T35pD+6t4oXvOKJV2~4=obM2Xp)2_qp*)e zH8lWqJ*0x8mc~lrff03vQi3tiP}C9T&q7d4*2~`G>lR7TM+!Z?u{?FW-uUy0wp#hg z@y_NUw_FFM1L0${wlGXEp#FcXRLUQht4uzFXbud|fcDfd2T3dRibNeafCZ+; zCk_MHRSv$swRQO|<=xjWYK=b-lTtWn$cF?w5a2!%EK0Io!smvRVO9WGusl*T4S$S* z&6V6uNSlT0J6P{=F;1|3!=tmZ=)6@(`T5^2UouBu<&JN3m9;0IoV~@0v%L$Uqsr#9 zx?J6D^ZC}2(W@+50QamSy2j|02i(}z|~`Kf|Nlx za32ID7hIKu5U;ce*tnU<_7!3Lzx$6ebHx|R_*rk>L|&~y2o{6bNMQp7ff6X8U=#;&0^lbBOBn7T)GIIeaeYAuk zl(lC!mi{x{<(2o=%uUj4%sby`|E!ZMoYU}Tpj(d)0WR54DFKb{0b+vVAMZG zZXquWVuIg@2GhE}eDPFg6BrI}qRo=6`fy{b#RKZVD9{+K9T=8h0X2Y$%^rjavWJ0z zkbv~Bs25959zw=$%`vPQ0Fv|92=#LD+0G(PUx|cW){3C|BKB*r02Gvt)g*&H0g_ve u*h7Hnhd#&3MMWO5g0Wa^q6@&)_z>jBW(CWuADviOR{d*nkj$vQ`BrU35aQTA21SIg@>j z%va@Iq)o35wa;n&wWmGUKRTy9dg9gv^tY{R+J6{&tFL|4>aJFokn1_EoV#gES68@X zG8DmubDpQDz41I)m5b8K<}8?t$!X4msKVKlf{0$mD3cYA!jA`O4@FU*`QZOte+OC#`Z&U8@{y$Q@)rIHT9Yo9X57nI>PD6Nx9iQ8N&FR{?x zTbhHH&;=`o>G0Yf+I808%;5s1+(g^IYUS z1$?Z65x;zp4zw519fx}8xnB&?p|t~a%h`*jFL>mp@zD@eklJP@S<&RqJ8yB7aOIByl`Qn<-~6uq-*}!M;CwonLfQ07raOloW=Dem>Pw&O6pQH^SZ6+ty4$))_ZK+&#j}G!!-wDjjjpL@))XMI#p(JMQ4Jrq7<*ZF*X&&(M)f3$s#eA zl`(($fu#YAM7o?YiNXmO{y{MK6 z5v^lJYLRWA*1kdd`-byp4lIWG>|;>e#02nxQB^7l)C8WJ;WDI^6fvBHp9Z#rRK;|ffloiv{ zn6#FhHa_0;xR=(;S%6Q$h!8m>M7b0&7(me}89kRc&YL82&BaR{s?Jt3hy*$#F-cq| z42@qhP1B~@Icoe^wPS(H2$e9bY83=HpQUYPrT{-(vjh4|s2MPiF;;G*6)D#e6Z&F= z3$*`^A^KRo)PCB0udBWJoZs}FxMRme+)>s|OqJbXYH~N-usWP8LIS1&XO)Yb1Ox%M zai^Wns`6??sf_a|GD}H8xFU}+IG`~Hu4w$RW$nI;_IA_bn}!-u97PRyqq2}y1{tsk zFJ-}MA+?IG009`QO|j<06!#@Le6i$^oh~dV+P`?19{chF8n|?@BeLKa2hyFxhZe8LGumOo{AMmEa?{5Z%y$@6DS|rWd#Lw2y2#wbh=#bxDg}zUb0X$SDX9 zELNb*g<3#75i_gKOJ-6@97-*-XBK7wVT!3yKxeC{z>1|u<}Mwm;lF0oe(tOD=FokY zJx{+pIMjal@--`I#n#QE*(f-Ih?1ek81Wv=gf$HT1Q|3*#*&@ol|@?|ug!$fYvHZY zIjRn;U%GHU?Yd;B{oU*N!j`4)Uq3!pt2KCOWs+y-j7-vn0Di$EQczXP0H_77Yoh@X z)U)s&G`Pwr{%iw*GEu*0X8992LgkmywX4=vl@zv-iPhXd44D z6(HbUwH)YCg8AYw92Alu?V;#tC2EA@iM#Dlg>EgQW%n$beMgKw*bG?^hC7TCO2I&% z7-WG88<%{957h$gQl&blY=uKX2Y9s1MjC$=mc0K|Vo$DYulUi|=Ft@oZ64JE-_D(c z>I<%XZ9Kie;_&mZE!Ya8g-cH29PJu1VH`lVg;!{AWvAH(1_o&T{mbXVa_HjwSG4!s zKi1Q}=BGFJw?BIH`r+O%HMM&petW9yp`B03Q)d3Us`sk0d(Xt=u90&`R-u#YC57`g zjY3r6^#Nn?AigphI8w}L9+I(`w`Dx}qziCvlLj=jtxz=wPex;N2gTFxTt)Z&xts2N z>LF5>FIh0TE0rr!*+CocTTYYD_O-V^ee#0##upx#*M9hwvF?_nf%7+yS)nTA3j!gN zP(WCt77_qx_~Mgs_~5(^fF5#*Gw@{z`Uqt?WVTq^xO$N2;PNh@1G0)Se)HWfI&g6C z==lL6(V(^&npO!G!VoJ&D#dtVngEazf(k?fq6Fkf!(3oUKgI@_kcG*wqI+Je?SCBX zZqfWVK0R0DDujkIH0D&fj>;e{L8}5YrdBWD+#nGY3puXZM#&2P#H|{@kG!EMz43M* z9eCrE*|jfNa@?3^Vctvd2!aBqt-1=NMdPwV_h4%b@_^K)Yz#S}N6l48UJ^+XmWgEK zIqg3%NPm3u#Cf~EvwiYvdgaa4-KlO{H&JQdYiH6A_72eLhki0$$rT?D)0f^qX>QoF zlU}_0G`j9zJ#^~f>!z8q_x+)+Jv%}qd_vcL)I$#(mh}4F1M{w)*p;@^xIc*=yLkns z&_jK14b39d7{aZU*7(E$>0(i`f~MmM9=4W46B=qs1;mmF-2tSs0^h}ogt$S9kG1Si zZ(ViVxs9)CiBY2OkV~Z+$cmV}w-6!NO!YQa5i0m=kZcV^z$3Q!W=6&xJr@ptMAGk{ zT}sy;8Kx(H-2}%C^x=`^OB!f(!rj&A$d$V%c1@ML-aesq3f=zkK&MF~_nb6S(D->7 zdMc544Ahe&1_U1Px)PyC(PJ(qv8swHuoxf*x{-23<|#oZq$S$>!4TbX%>uga^dYoy z$(?Wa)0%sZlRt-_Kv;tobFK_5O+}401WJW84m)!URtD(}Oh!W^NuIg`ijo19x~2%uce5u&;DFg0Ap}JC(~IXOT>87M!}Q<>OX&9Zdyb;`RJsqSp8GR~mM*4Ehq26j zIJ>7~a3G(PM0>B*~DT5BM=Ne3p1%+BBENvCB$_MR2pv3FhsNrEyW7<&WG2| zJZaMyj5Hc+y@Clr^gaBwaD+x55s`2kJX<1)V6$NwI}>ES3Ih({2KXjo;CKPOGu+ zanLa)iI;g3J7KCB#kC*J#@~6KVIxhUWXQ8d+Z1^11@=Cq2#7R?=;5DHaWfoa;h;OKQ0A0{(xU?7`IDOWDMneI99F*Qk3^ygrRAV;YHk&p=gU!jZ zh)*O|WCz4SL!7hTy8x2NW;6aubMA0Q{Qgg`Ml~9OEj58K*qgF8a8jgutVcK(+%=7p zLISZd!#})3F17Hg#E{#ejr8(w2b)LHyL{=QZ4*0pOmxI(GJdP1DcT?g>Cr=rX7g)l zqu2|ZZM0L?x-^*!-T3d3_Rl|ge(oHI^o?g9*K>y@)1xQ_%jk$L7BUAqvdV1oSV%!R zhS-D_7=$YPKZt~k#hMh-Ou2Jkr{%QejG> = static REPETITION_COUNT: LazyLock = LazyLock::new(|| { env::var("TREE_SITTER_BENCHMARK_REPETITION_COUNT").map_or(5, |s| s.parse::().unwrap()) }); +static WASM: LazyLock = LazyLock::new(|| env::var_os("TREE_SITTER_BENCHMARK_WASM").is_some()); static TEST_LOADER: LazyLock = LazyLock::new(|| Loader::with_parser_lib_path(SCRATCH_DIR.clone())); +#[cfg(feature = "wasm")] +static WASM_ENGINE: LazyLock = LazyLock::new(Default::default); #[expect( clippy::type_complexity, @@ -94,7 +99,11 @@ fn main() { } info!("\nLanguage: {language_name}"); - let language = get_language(language_path); + let language = if *WASM { + get_wasm_language(language_name, &mut parser) + } else { + get_language(language_path) + }; parser.set_language(&language).unwrap(); info!(" Constructing Queries"); @@ -214,3 +223,32 @@ fn get_language(path: &Path) -> Language { .with_context(|| format!("Failed to load language at path {}", src_path.display())) .unwrap() } + +#[cfg(feature = "wasm")] +fn get_wasm_language(language_name: &str, parser: &mut Parser) -> Language { + let wasm_language_name = language_name.replace('-', "_"); + let wasm_path = ROOT_DIR + .join("target") + .join("release") + .join(format!("tree-sitter-{language_name}.wasm")); + let wasm = fs::read(&wasm_path) + .with_context(|| { + format!( + "Failed to read {}. Generate Wasm fixtures with `cargo xtask generate-fixtures --wasm`", + wasm_path.display() + ) + }) + .unwrap(); + let mut store = WasmStore::new(&WASM_ENGINE).expect("Failed to create Wasm store"); + let language = store + .load_language(&wasm_language_name, &wasm) + .with_context(|| format!("Failed to load Wasm language at {}", wasm_path.display())) + .unwrap(); + parser.set_wasm_store(store).unwrap(); + language +} + +#[cfg(not(feature = "wasm"))] +fn get_wasm_language(_language_name: &str, _parser: &mut Parser) -> Language { + panic!("Wasm benchmarking requires the `wasm` feature"); +} diff --git a/crates/generate/src/build_tables/item.rs b/crates/generate/src/build_tables/item.rs index d6657c268..f8e3494b9 100644 --- a/crates/generate/src/build_tables/item.rs +++ b/crates/generate/src/build_tables/item.rs @@ -139,7 +139,7 @@ struct ItemContent<'a> { impl Eq for ItemContent<'_> {} impl ItemContent<'_> { - fn prec(&self) -> Precedence { + const fn prec(&self) -> Precedence { if self.dot > 0 { self.production.steps[self.dot - 1].precedence() } else { @@ -147,7 +147,7 @@ impl ItemContent<'_> { } } - fn assoc(&self) -> Option { + const fn assoc(&self) -> Option { if self.dot > 0 { self.production.steps[self.dot - 1].associativity() } else { @@ -486,7 +486,7 @@ impl<'a> ParseItem<'a> { /// This item's identity keys at the current dot. #[must_use] - fn dot_keys(&self) -> DotKeys { + const fn dot_keys(&self) -> DotKeys { self.keys[self.step_index as usize] } } diff --git a/crates/xtask/src/benchmark.rs b/crates/xtask/src/benchmark.rs index 8da80d52c..2b5262f19 100644 --- a/crates/xtask/src/benchmark.rs +++ b/crates/xtask/src/benchmark.rs @@ -3,6 +3,10 @@ use anyhow::Result; use crate::{Benchmark, bail_on_err}; pub fn run(args: &Benchmark) -> Result<()> { + if args.wasm { + unsafe { std::env::set_var("TREE_SITTER_BENCHMARK_WASM", "1") }; + } + if let Some(ref example) = args.example_file_name { unsafe { std::env::set_var("TREE_SITTER_BENCHMARK_EXAMPLE_FILTER", example) }; } @@ -26,6 +30,12 @@ pub fn run(args: &Benchmark) -> Result<()> { .arg("benchmark") .arg("-p") .arg("tree-sitter-cli") + .args( + args.wasm + .then_some(["--features", "wasm"]) + .into_iter() + .flatten(), + ) .arg("--no-run") .arg("--message-format=json") .spawn()? @@ -66,6 +76,12 @@ pub fn run(args: &Benchmark) -> Result<()> { .arg("benchmark") .arg("-p") .arg("tree-sitter-cli") + .args( + args.wasm + .then_some(["--features", "wasm"]) + .into_iter() + .flatten(), + ) .status()?; if !status.success() { diff --git a/crates/xtask/src/main.rs b/crates/xtask/src/main.rs index f7eab1b56..05e15d5c0 100644 --- a/crates/xtask/src/main.rs +++ b/crates/xtask/src/main.rs @@ -71,6 +71,9 @@ struct Benchmark { /// Whether to run the benchmarks in debug mode. #[arg(long, short = 'g')] debug: bool, + /// Benchmark Wasm grammars instead of native grammars. + #[arg(long)] + wasm: bool, } #[derive(Args)] diff --git a/lib/Cargo.toml b/lib/Cargo.toml index b304b876f..7c12f00ce 100644 --- a/lib/Cargo.toml +++ b/lib/Cargo.toml @@ -47,9 +47,9 @@ streaming-iterator = "0.1.9" tree-sitter-language.workspace = true [dependencies.wasmtime-c-api] -version = "36.0.13" +version = "48.0.0" default-features = false -features = [ "cranelift", "gc-drc" ] +features = [ "cranelift", "gc", "gc-null" ] optional = true package = "wasmtime-c-api-impl" diff --git a/lib/src/wasm_store.c b/lib/src/wasm_store.c index f7ef81652..525b8c78b 100644 --- a/lib/src/wasm_store.c +++ b/lib/src/wasm_store.c @@ -75,13 +75,13 @@ typedef struct { WasmLanguageId *language_id; wasmtime_instance_t instance; int32_t external_states_address; - int32_t lex_main_fn_index; - int32_t lex_keyword_fn_index; - int32_t scanner_create_fn_index; - int32_t scanner_destroy_fn_index; - int32_t scanner_serialize_fn_index; - int32_t scanner_deserialize_fn_index; - int32_t scanner_scan_fn_index; + wasmtime_func_t lex_main_fn; + wasmtime_func_t lex_keyword_fn; + wasmtime_func_t scanner_create_fn; + wasmtime_func_t scanner_destroy_fn; + wasmtime_func_t scanner_serialize_fn; + wasmtime_func_t scanner_deserialize_fn; + wasmtime_func_t scanner_scan_fn; } LanguageWasmInstance; typedef struct { @@ -188,19 +188,31 @@ typedef struct { * WasmDylinkMemoryInfo ***********************/ -static uint8_t read_u8(const uint8_t **p) { - return *(*p)++; +typedef struct { + const uint8_t *data; + size_t offset; + size_t size; +} WasmReader; + +static bool wasm_reader__read_u8(WasmReader *reader, uint8_t *result) { + if (reader->offset >= reader->size) return false; + *result = reader->data[reader->offset++]; + return true; } -static inline uint64_t read_uleb128(const uint8_t **p, const uint8_t *end) { - uint64_t value = 0; - unsigned shift = 0; - do { - if (*p == end) return UINT64_MAX; - value += (uint64_t)(**p & 0x7f) << shift; - shift += 7; - } while (*((*p)++) >= 128); - return value; +static bool wasm_reader__read_uleb128(WasmReader *reader, uint32_t *result) { + uint32_t value = 0; + for (unsigned shift = 0; shift < 32; shift += 7) { + uint8_t byte; + if (!wasm_reader__read_u8(reader, &byte)) return false; + if (shift == 28 && (byte & 0xf0) != 0) return false; + value |= (uint32_t)(byte & 0x7f) << shift; + if ((byte & 0x80) == 0) { + *result = value; + return true; + } + } + return false; } static bool wasm_dylink_info__parse( @@ -213,45 +225,64 @@ static bool wasm_dylink_info__parse( const uint8_t WASM_CUSTOM_SECTION = 0x0; const uint8_t WASM_DYLINK_MEM_INFO = 0x1; - const uint8_t *p = bytes; - const uint8_t *end = bytes + length; - if (length < 8) return false; - if (memcmp(p, WASM_MAGIC_NUMBER, 4) != 0) return false; - p += 4; - if (memcmp(p, WASM_VERSION, 4) != 0) return false; - p += 4; + if (memcmp(bytes, WASM_MAGIC_NUMBER, 4) != 0) return false; + if (memcmp(bytes + 4, WASM_VERSION, 4) != 0) return false; - while (p < end) { - uint8_t section_id = read_u8(&p); - uint32_t section_length = read_uleb128(&p, end); - const uint8_t *section_end = p + section_length; - if (section_end > end) return false; + WasmReader reader = { + .data = bytes, + .offset = 8, + .size = length, + }; + + while (reader.offset < reader.size) { + uint8_t section_id; + uint32_t section_length; + if ( + !wasm_reader__read_u8(&reader, §ion_id) || + !wasm_reader__read_uleb128(&reader, §ion_length) || + section_length > reader.size - reader.offset + ) return false; + size_t section_end = reader.offset + section_length; if (section_id == WASM_CUSTOM_SECTION) { - uint32_t name_length = read_uleb128(&p, section_end); - const uint8_t *name_end = p + name_length; - if (name_end > section_end) return false; + size_t previous_size = reader.size; + reader.size = section_end; + uint32_t name_length; + if ( + !wasm_reader__read_uleb128(&reader, &name_length) || + name_length > reader.size - reader.offset + ) return false; + size_t name_end = reader.offset + name_length; - if (name_length == 8 && memcmp(p, "dylink.0", 8) == 0) { - p = name_end; - while (p < section_end) { - uint8_t subsection_type = read_u8(&p); - uint32_t subsection_size = read_uleb128(&p, section_end); - const uint8_t *subsection_end = p + subsection_size; - if (subsection_end > section_end) return false; + if (name_length == 8 && memcmp(&reader.data[reader.offset], "dylink.0", 8) == 0) { + reader.offset = name_end; + while (reader.offset < section_end) { + uint8_t subsection_type; + uint32_t subsection_size; + if ( + !wasm_reader__read_u8(&reader, &subsection_type) || + !wasm_reader__read_uleb128(&reader, &subsection_size) || + subsection_size > section_end - reader.offset + ) return false; + size_t subsection_end = reader.offset + subsection_size; if (subsection_type == WASM_DYLINK_MEM_INFO) { - info->memory_size = read_uleb128(&p, subsection_end); - info->memory_align = read_uleb128(&p, subsection_end); - info->table_size = read_uleb128(&p, subsection_end); - info->table_align = read_uleb128(&p, subsection_end); + reader.size = subsection_end; + if ( + !wasm_reader__read_uleb128(&reader, &info->memory_size) || + !wasm_reader__read_uleb128(&reader, &info->memory_align) || + !wasm_reader__read_uleb128(&reader, &info->table_size) || + !wasm_reader__read_uleb128(&reader, &info->table_align) || + reader.offset != subsection_end + ) return false; return true; } - p = subsection_end; + reader.offset = subsection_end; } } + reader.size = previous_size; } - p = section_end; + reader.offset = section_end; } return false; } @@ -540,7 +571,8 @@ static void delete_partially_loaded_language( } static bool name_eq(const wasm_name_t *name, const char *string) { - return strncmp(string, name->data, name->size) == 0; + size_t length = strlen(string); + return name->size == length && memcmp(name->data, string, length) == 0; } static inline wasm_functype_t* wasm_functype_new_4_0( @@ -1075,6 +1107,18 @@ static uint32_t ts_wasm_store__serialization_buffer_address(TSWasmStore *self) { return self->current_memory_offset; } +static wasmtime_func_t ts_wasm_store__get_function( + TSWasmStore *self, + int32_t function_index +) { + wasmtime_context_t *context = wasmtime_store_context(self->store); + wasmtime_val_t value; + bool succeeded = wasmtime_table_get(context, &self->function_table, function_index, &value); + ts_assert(succeeded); + ts_assert(value.kind == WASMTIME_FUNCREF); + return value.of.funcref; +} + static bool ts_wasm_store__instantiate( TSWasmStore *self, wasmtime_module_t *module, @@ -1090,6 +1134,8 @@ static bool ts_wasm_store__instantiate( char *language_function_name = NULL; wasmtime_extern_t *imports = NULL; wasmtime_context_t *context = wasmtime_store_context(self->store); + uint32_t initial_memory_offset = self->current_memory_offset; + uint32_t initial_function_table_offset = self->current_function_table_offset; // Grow the function table to make room for the new functions. wasmtime_val_t initializer = {.kind = WASMTIME_FUNCREF}; @@ -1256,6 +1302,8 @@ static bool ts_wasm_store__instantiate( return true; error: + self->current_memory_offset = initial_memory_offset; + self->current_function_table_offset = initial_function_table_offset; if (language_function_name) ts_free(language_function_name); if (message.size) wasm_byte_vec_delete(&message); if (error) wasmtime_error_delete(error); @@ -1282,6 +1330,8 @@ const TSLanguage *ts_wasm_store_load_language( TSLanguage *language = NULL; StringData symbol_name_buffer = array_new(); StringData field_name_buffer = array_new(); + uint32_t initial_memory_offset = self->current_memory_offset; + uint32_t initial_function_table_offset = self->current_function_table_offset; wasm_error->kind = TSWasmErrorKindNone; if (!wasm_dylink_info__parse((const unsigned char *)wasm, wasm_len, &dylink_info)) { @@ -1508,14 +1558,13 @@ const TSLanguage *ts_wasm_store_load_language( ); if (!valid_wasm_memory) goto invalid_language_memory; - TSSymbol last_supertype = language->supertype_symbols[language->supertype_count - 1]; - TSMapSlice last_slice = language->supertype_map_slices[last_supertype]; + TSMapSlice last_slice = language->supertype_map_slices[largest_supertype]; uint32_t supertype_map_entry_count = last_slice.index + last_slice.length; language->supertype_map_entries = copy( &wasm_memory, wasm_language.supertype_map_entries, - supertype_map_entry_count * sizeof(char *), + supertype_map_entry_count * sizeof(TSSymbol), &valid_wasm_memory ); if (!valid_wasm_memory) goto invalid_language_memory; @@ -1651,13 +1700,13 @@ const TSLanguage *ts_wasm_store_load_language( .language_id = language_id_clone(result->language_id), .instance = instance, .external_states_address = wasm_language.external_scanner.states, - .lex_main_fn_index = wasm_language.lex_fn, - .lex_keyword_fn_index = wasm_language.keyword_lex_fn, - .scanner_create_fn_index = wasm_language.external_scanner.create, - .scanner_destroy_fn_index = wasm_language.external_scanner.destroy, - .scanner_serialize_fn_index = wasm_language.external_scanner.serialize, - .scanner_deserialize_fn_index = wasm_language.external_scanner.deserialize, - .scanner_scan_fn_index = wasm_language.external_scanner.scan, + .lex_main_fn = ts_wasm_store__get_function(self, wasm_language.lex_fn), + .lex_keyword_fn = ts_wasm_store__get_function(self, wasm_language.keyword_lex_fn), + .scanner_create_fn = ts_wasm_store__get_function(self, wasm_language.external_scanner.create), + .scanner_destroy_fn = ts_wasm_store__get_function(self, wasm_language.external_scanner.destroy), + .scanner_serialize_fn = ts_wasm_store__get_function(self, wasm_language.external_scanner.serialize), + .scanner_deserialize_fn = ts_wasm_store__get_function(self, wasm_language.external_scanner.deserialize), + .scanner_scan_fn = ts_wasm_store__get_function(self, wasm_language.external_scanner.scan), })); return language; @@ -1668,6 +1717,8 @@ invalid_language_memory: goto error; error: + self->current_memory_offset = initial_memory_offset; + self->current_function_table_offset = initial_function_table_offset; delete_partially_loaded_language(result, &symbol_name_buffer, &field_name_buffer); if (module) wasmtime_module_delete(module); return NULL; @@ -1699,6 +1750,8 @@ bool ts_wasm_store_add_language( // If the language module has not been instantiated in this store, then add // it to this store. if (!exists) { + uint32_t initial_memory_offset = self->current_memory_offset; + uint32_t initial_function_table_offset = self->current_function_table_offset; *index = self->language_instances.size; char *message; wasmtime_instance_t instance; @@ -1723,19 +1776,21 @@ bool ts_wasm_store_add_language( .size = wasmtime_memory_data_size(context, &self->memory), }; if (!wasm_memory__read(&wasm_memory, language_address, &wasm_language, sizeof(LanguageInWasmMemory))) { + self->current_memory_offset = initial_memory_offset; + self->current_function_table_offset = initial_function_table_offset; return false; } array_push(&self->language_instances, ((LanguageWasmInstance) { .language_id = language_id_clone(language_data->language_id), .instance = instance, .external_states_address = wasm_language.external_scanner.states, - .lex_main_fn_index = wasm_language.lex_fn, - .lex_keyword_fn_index = wasm_language.keyword_lex_fn, - .scanner_create_fn_index = wasm_language.external_scanner.create, - .scanner_destroy_fn_index = wasm_language.external_scanner.destroy, - .scanner_serialize_fn_index = wasm_language.external_scanner.serialize, - .scanner_deserialize_fn_index = wasm_language.external_scanner.deserialize, - .scanner_scan_fn_index = wasm_language.external_scanner.scan, + .lex_main_fn = ts_wasm_store__get_function(self, wasm_language.lex_fn), + .lex_keyword_fn = ts_wasm_store__get_function(self, wasm_language.keyword_lex_fn), + .scanner_create_fn = ts_wasm_store__get_function(self, wasm_language.external_scanner.create), + .scanner_destroy_fn = ts_wasm_store__get_function(self, wasm_language.external_scanner.destroy), + .scanner_serialize_fn = ts_wasm_store__get_function(self, wasm_language.external_scanner.serialize), + .scanner_deserialize_fn = ts_wasm_store__get_function(self, wasm_language.external_scanner.deserialize), + .scanner_scan_fn = ts_wasm_store__get_function(self, wasm_language.external_scanner.scan), })); } @@ -1774,19 +1829,13 @@ void ts_wasm_store_reset(TSWasmStore *self) { static void ts_wasm_store__call( TSWasmStore *self, - int32_t function_index, + wasmtime_func_t *func, wasmtime_val_raw_t *args_and_results, size_t args_and_results_len ) { wasmtime_context_t *context = wasmtime_store_context(self->store); - wasmtime_val_t value; - bool succeeded = wasmtime_table_get(context, &self->function_table, function_index, &value); - ts_assert(succeeded); - ts_assert(value.kind == WASMTIME_FUNCREF); - wasmtime_func_t func = value.of.funcref; - wasm_trap_t *trap = NULL; - wasmtime_error_t *error = wasmtime_func_call_unchecked(context, &func, args_and_results, args_and_results_len, &trap); + wasmtime_error_t *error = wasmtime_func_call_unchecked(context, func, args_and_results, args_and_results_len, &trap); if (error) { // wasm_message_t message; // wasmtime_error_message(error, &message); @@ -1819,7 +1868,7 @@ typedef struct { TSSymbol result_symbol; } TSLexerDataPrefix; -static bool ts_wasm_store__call_lex_function(TSWasmStore *self, unsigned function_index, TSStateId state) { +static bool ts_wasm_store__call_lex_function(TSWasmStore *self, wasmtime_func_t *func, TSStateId state) { wasmtime_context_t *context = wasmtime_store_context(self->store); uint8_t *memory_data = wasmtime_memory_data(context, &self->memory); memcpy( @@ -1832,7 +1881,7 @@ static bool ts_wasm_store__call_lex_function(TSWasmStore *self, unsigned functio {.i32 = self->lexer_address}, {.i32 = state}, }; - ts_wasm_store__call(self, function_index, args, 2); + ts_wasm_store__call(self, func, args, 2); if (self->has_error) return false; bool result = args[0].i32; @@ -1847,7 +1896,7 @@ static bool ts_wasm_store__call_lex_function(TSWasmStore *self, unsigned functio bool ts_wasm_store_call_lex_main(TSWasmStore *self, TSStateId state) { return ts_wasm_store__call_lex_function( self, - self->current_instance->lex_main_fn_index, + &self->current_instance->lex_main_fn, state ); } @@ -1855,14 +1904,14 @@ bool ts_wasm_store_call_lex_main(TSWasmStore *self, TSStateId state) { bool ts_wasm_store_call_lex_keyword(TSWasmStore *self, TSStateId state) { return ts_wasm_store__call_lex_function( self, - self->current_instance->lex_keyword_fn_index, + &self->current_instance->lex_keyword_fn, state ); } uint32_t ts_wasm_store_call_scanner_create(TSWasmStore *self) { wasmtime_val_raw_t args[1] = {{.i32 = 0}}; - ts_wasm_store__call(self, self->current_instance->scanner_create_fn_index, args, 1); + ts_wasm_store__call(self, &self->current_instance->scanner_create_fn, args, 1); if (self->has_error) return 0; return args[0].i32; } @@ -1870,7 +1919,7 @@ uint32_t ts_wasm_store_call_scanner_create(TSWasmStore *self) { void ts_wasm_store_call_scanner_destroy(TSWasmStore *self, uint32_t scanner_address) { if (self->current_instance) { wasmtime_val_raw_t args[1] = {{.i32 = scanner_address}}; - ts_wasm_store__call(self, self->current_instance->scanner_destroy_fn_index, args, 1); + ts_wasm_store__call(self, &self->current_instance->scanner_destroy_fn, args, 1); } } @@ -1896,7 +1945,7 @@ bool ts_wasm_store_call_scanner_scan( {.i32 = self->lexer_address}, {.i32 = valid_tokens_address} }; - ts_wasm_store__call(self, self->current_instance->scanner_scan_fn_index, args, 3); + ts_wasm_store__call(self, &self->current_instance->scanner_scan_fn, args, 3); if (self->has_error) return false; memcpy( @@ -1920,7 +1969,7 @@ uint32_t ts_wasm_store_call_scanner_serialize( {.i32 = scanner_address}, {.i32 = serialization_buffer_address}, }; - ts_wasm_store__call(self, self->current_instance->scanner_serialize_fn_index, args, 2); + ts_wasm_store__call(self, &self->current_instance->scanner_serialize_fn, args, 2); if (self->has_error) return 0; uint32_t length = args[0].i32; @@ -1962,7 +2011,7 @@ void ts_wasm_store_call_scanner_deserialize( {.i32 = serialization_buffer_address}, {.i32 = length}, }; - ts_wasm_store__call(self, self->current_instance->scanner_deserialize_fn_index, args, 3); + ts_wasm_store__call(self, &self->current_instance->scanner_deserialize_fn, args, 3); } bool ts_wasm_store_has_error(const TSWasmStore *self) {