use std::collections::{BTreeSet, HashMap, HashSet};
use crate::dep_graph::{DepGraph, DepNode};
use crate::graph_utils::tarjan_scc::{SccEvent, SccId, TarjanSccResult};
use crate::read::{ExportId, ImportId, InputFuncId, InputModule};
use eyre::{anyhow, bail, Result};
use lazy_static::lazy_static;
use regex::Regex;
use tracing::{trace, warn};
use wasmparser::TypeRef;
#[derive(Debug, PartialEq, Eq, Clone)]
pub struct SplitPoint {
pub module_name: String,
pub import: ImportId,
pub import_func: InputFuncId,
pub export: ExportId,
pub export_func: InputFuncId,
}
pub fn get_split_points(module: &InputModule) -> Result<Vec<SplitPoint>> {
macro_rules! process_imports_or_exports {
($pattern:expr, $map:ident, $member:ident, $id_ty:ty) => {
let mut $map = HashMap::<(String, String), $id_ty>::new();
{
lazy_static! {
static ref PATTERN: Regex = Regex::new($pattern).unwrap();
}
for (id, item) in module.$member.iter().enumerate() {
let Some(captures) = PATTERN.captures(&item.name) else {
continue;
};
let (_, [module_name, unique_id]) = captures.extract();
$map.insert((module_name.into(), unique_id.into()), id);
}
}
};
}
process_imports_or_exports!(
"__wasm_split_00(.*)00_import_([0-9a-f]{32})",
import_map,
imports,
ImportId
);
process_imports_or_exports!(
"__wasm_split_00(.*)00_export_([0-9a-f]{32})",
export_map,
exports,
ExportId
);
let split_points = import_map
.drain()
.map(|(key, import_id)| -> Result<SplitPoint> {
let export_id = export_map
.remove(&key)
.ok_or_else(|| anyhow!("No corresponding export for split import {key:?}"))?;
let export = module.exports[export_id];
let wasmparser::Export {
kind: wasmparser::ExternalKind::Func,
index,
..
} = export
else {
bail!("Expected exported function but received: {export:?}");
};
let &import_func = module.imported_func_map.get(&import_id).ok_or_else(|| {
anyhow!(
"Expected imported function but received: {:?}",
&module.imports[import_id]
)
})?;
Ok(SplitPoint {
module_name: key.0,
import: import_id,
import_func,
export: export_id,
export_func: index as InputFuncId,
})
})
.collect::<Result<Vec<SplitPoint>>>()?;
if !export_map.is_empty() {
warn!(
"No corresponding imports for split export(s) {:?}",
export_map.keys().collect::<Vec<_>>()
);
}
Ok(split_points)
}
#[derive(Debug, Default)]
pub struct ReachabilityGraph {
pub reachable: HashSet<DepNode>,
}
#[derive(Debug, Default)]
pub struct OutputModuleInfo {
pub included_symbols: HashSet<DepNode>,
pub used_shared_deps: HashSet<DepNode>,
pub is_empty: bool,
}
impl OutputModuleInfo {
pub fn print(&self, module_name: &str, module: &InputModule) {
print_deps(module_name, module, &self.included_symbols);
}
}
impl From<ReachabilityGraph> for OutputModuleInfo {
fn from(reachability: ReachabilityGraph) -> Self {
Self {
included_symbols: reachability.reachable,
..Default::default()
}
}
}
pub fn trace_enabled(verbose: bool) -> bool {
verbose && tracing::event_enabled!(tracing::Level::TRACE)
}
fn print_deps(module_name: &str, module: &InputModule, reachable: &HashSet<DepNode>) {
let format_dep = |dep: &DepNode| match dep {
DepNode::Function(index) => {
let name = module.names.functions.get(index);
format!("func[{index}] <{name:?}>")
}
DepNode::DataSymbol(index) => {
let symbol = module.reloc_info.symbols[*index];
format!("{symbol:?}")
}
DepNode::Global(index) => {
format!("global[{index}]")
}
DepNode::Table(index) => {
format!("table[{index}]")
}
DepNode::Tag(index) => {
format!("tag[{index}]")
}
DepNode::Memory(index) => {
format!("memory[{index}]")
}
};
trace!("SPLIT: ============== {module_name}");
let mut total_size: usize = 0;
for dep in reachable.iter() {
if let DepNode::Function(index) = dep {
let size = index
.checked_sub(module.imported_funcs.len())
.map(|defined_index| module.defined_funcs[defined_index].body.range().len())
.unwrap_or_default();
total_size += size;
trace!(" {} size={size:?}", format_dep(dep));
} else {
trace!(" {}", format_dep(dep));
}
}
trace!("SPLIT: ============== {module_name} : total size: {total_size}");
}
fn get_main_module_roots(module: &InputModule, split_points: &[SplitPoint]) -> HashSet<DepNode> {
let mut roots: HashSet<DepNode> = HashSet::new();
if let Some(id) = module.start {
roots.insert(DepNode::Function(id));
}
for func_id in 0..module.imported_funcs.len() {
roots.insert(DepNode::Function(func_id));
}
for global_id in 0..module.imported_globals_num {
roots.insert(DepNode::Global(global_id));
}
for table_id in 0..module.imported_tables_num {
roots.insert(DepNode::Table(table_id));
}
for tag_id in 0..module.imported_tags_num {
roots.insert(DepNode::Tag(tag_id));
}
for tag_id in 0..module.imported_memories_num {
roots.insert(DepNode::Memory(tag_id));
}
roots.insert(module.main_memory());
for wasmparser::Export { index, kind, .. } in module.exports.iter() {
roots.insert(match kind {
wasmparser::ExternalKind::Func | wasmparser::ExternalKind::FuncExact => {
DepNode::Function(*index as usize)
}
wasmparser::ExternalKind::Table => DepNode::Table(*index as usize),
wasmparser::ExternalKind::Global => DepNode::Global(*index as usize),
wasmparser::ExternalKind::Tag => DepNode::Tag(*index as usize),
wasmparser::ExternalKind::Memory => DepNode::Memory(*index as usize),
});
}
for &func_id in &module.reloc_info.visible_indirects {
roots.insert(DepNode::Function(func_id));
}
for split_point in split_points.iter() {
roots.remove(&DepNode::Function(split_point.export_func));
roots.remove(&DepNode::Function(split_point.import_func));
}
roots
}
fn wbg_rooting_funs(_dep_graph: &DepGraph, module: &InputModule) -> HashSet<DepNode> {
let mut users_must_be_in_main = HashSet::new();
let mut _wbg_describe_cast = None;
for (import_id, import) in module.imports.iter().enumerate() {
if import.module != "__wbindgen_placeholder__" || !matches!(import.ty, TypeRef::Func(_)) {
continue;
}
if import.name == "__wbindgen_describe_cast" {
let func_id = module.imported_func_map.get(&import_id).cloned().unwrap();
_wbg_describe_cast = Some(func_id);
users_must_be_in_main.insert(DepNode::Function(func_id));
}
}
users_must_be_in_main
}
fn get_split_roots(splits_in_module: &[&SplitPoint]) -> HashSet<DepNode> {
let mut roots = HashSet::<DepNode>::new();
for entry_point in splits_in_module {
roots.insert(DepNode::Function(entry_point.export_func));
}
roots
}
pub fn get_split_points_by_module(
split_points: &[SplitPoint],
) -> HashMap<String, Vec<&SplitPoint>> {
split_points
.iter()
.fold(HashMap::new(), |mut map, split_point| {
map.entry(split_point.module_name.clone())
.or_default()
.push(split_point);
map
})
}
#[derive(Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Clone)]
pub enum SplitModuleIdentifier {
Main,
Split(String),
Chunk(BTreeSet<String>),
}
impl SplitModuleIdentifier {
pub fn filename(&self, module_index: usize) -> String {
match self {
Self::Main => unreachable!("main wasm filepath is handled separately"),
Self::Split(name) => format!("split_{name}"),
Self::Chunk(_) => format!("chunk_{module_index}"),
}
}
pub fn loader_name(&self) -> String {
match self {
Self::Split(name) => format!("__wasm_split_load_{name}"),
_ => unreachable!("only whole modules have a loader"),
}
}
fn also_in(&mut self, other: &SplitModuleIdentifier) {
let mut needed_by = BTreeSet::new();
match self {
SplitModuleIdentifier::Main => return,
SplitModuleIdentifier::Split(split) => {
needed_by.insert(split.clone());
}
SplitModuleIdentifier::Chunk(chunk) => std::mem::swap(chunk, &mut needed_by),
};
match other {
SplitModuleIdentifier::Main => {
*self = SplitModuleIdentifier::Main;
return;
}
SplitModuleIdentifier::Split(split) => {
needed_by.insert(split.clone());
}
SplitModuleIdentifier::Chunk(chunk) => {
needed_by.extend(chunk.iter().cloned());
}
};
if needed_by.len() == 1 {
*self = SplitModuleIdentifier::Split(needed_by.pop_first().unwrap());
} else {
debug_assert!(needed_by.len() > 1);
match self {
SplitModuleIdentifier::Chunk(chunk) => std::mem::swap(chunk, &mut needed_by),
_ => *self = SplitModuleIdentifier::Chunk(needed_by),
}
}
}
}
#[derive(Debug, Default)]
pub struct SplitProgramInfo {
pub output_modules: Vec<(SplitModuleIdentifier, OutputModuleInfo)>,
pub split_points: Vec<SplitPoint>,
pub split_point_exports: HashSet<InputFuncId>,
pub shared_deps: HashSet<DepNode>,
pub symbol_output_module: HashMap<DepNode, usize>,
}
struct DepGraphAnalysis<'g> {
dep_graph: &'g DepGraph,
wbg_rooting_deps: &'g HashSet<DepNode>,
scc_searcher: TarjanSccResult,
scc_root_colors: HashMap<SccId, SplitModuleIdentifier>,
}
struct DepGraphPainter<It> {
scc_events: It,
scc_colors: HashMap<SccId, SplitModuleIdentifier>,
color: SplitModuleIdentifier,
}
fn color_all<It, T>(
color_map: &mut HashMap<T, SplitModuleIdentifier>,
roots: It,
color: SplitModuleIdentifier,
) where
It: IntoIterator<Item = T>,
T: Eq + std::hash::Hash,
{
for root in roots {
color_map
.entry(root)
.and_modify(|module| module.also_in(&color))
.or_insert(color.clone());
}
}
impl<'g> DepGraphAnalysis<'g> {
fn new(dep_graph: &'g DepGraph, wbg_rooting_deps: &'g HashSet<DepNode>) -> Self {
Self {
dep_graph,
wbg_rooting_deps,
scc_searcher: TarjanSccResult::new(),
scc_root_colors: HashMap::new(),
}
}
fn explore(&mut self, roots: HashSet<DepNode>, color: SplitModuleIdentifier) {
let scc_roots = self
.scc_searcher
.explore(&roots, self.dep_graph, self.wbg_rooting_deps);
color_all(&mut self.scc_root_colors, scc_roots, color);
}
fn into_painter(self) -> DepGraphPainter<impl Iterator<Item = SccEvent>> {
let topsort = self.scc_searcher.into_topsort();
DepGraphPainter {
scc_events: topsort.into_iter().rev(),
scc_colors: self.scc_root_colors,
color: SplitModuleIdentifier::Main,
}
}
}
impl<It: Iterator<Item = SccEvent>> DepGraphPainter<It> {
fn next(&mut self) -> Option<(DepNode, &SplitModuleIdentifier)> {
loop {
match self.scc_events.next()? {
SccEvent::Next {
this,
out_edges,
marked_for_main,
} => {
self.color = self.scc_colors.remove(&this).unwrap();
if marked_for_main {
self.color = SplitModuleIdentifier::Main;
}
color_all(&mut self.scc_colors, out_edges, self.color.clone());
}
SccEvent::Member { node } => {
return Some((node, &self.color));
}
}
}
}
}
pub fn compute_split_modules(
module: &InputModule,
dep_graph: &DepGraph,
split_points: Vec<SplitPoint>,
) -> Result<SplitProgramInfo> {
let split_points_by_module = get_split_points_by_module(&split_points);
trace!("split_points={split_points:?}");
let split_func_map: HashMap<_, _> = split_points
.iter()
.map(|split_point| (split_point.import_func, split_point.export_func))
.collect();
let all_imports: HashSet<_> = split_func_map
.keys()
.map(|&import| DepNode::Function(import))
.collect();
let wbg_rooting_deps = wbg_rooting_funs(dep_graph, module);
let mut graph_analysis = DepGraphAnalysis::new(dep_graph, &wbg_rooting_deps);
let main_roots = get_main_module_roots(module, &split_points);
graph_analysis.explore(main_roots, SplitModuleIdentifier::Main);
for (module_name, entry_points) in &split_points_by_module {
let roots = get_split_roots(entry_points);
graph_analysis.explore(roots, SplitModuleIdentifier::Split(module_name.clone()));
}
let mut painter = graph_analysis.into_painter();
let mut split_module_contents = HashMap::<SplitModuleIdentifier, OutputModuleInfo>::new();
while let Some((node, color)) = painter.next() {
if all_imports.contains(&node) {
continue;
}
split_module_contents
.entry(color.clone())
.or_default()
.included_symbols
.insert(node);
}
let mut program_info = SplitProgramInfo::default();
for out_module in split_module_contents.values_mut() {
let needed_symbols = out_module
.included_symbols
.iter()
.filter_map(|dep| dep_graph.get(dep))
.flatten()
.cloned()
.chain([module.main_memory(), module.indirect_function_table()]);
for mut dep_to_check in needed_symbols {
if let DepNode::Function(called_func_id) = &mut dep_to_check {
if let Some(mapped_func_id) = split_func_map.get(called_func_id) {
*called_func_id = *mapped_func_id;
}
}
let in_other_module = !out_module.included_symbols.contains(&dep_to_check);
if !in_other_module {
continue;
}
if let DepNode::DataSymbol(_) = dep_to_check {
continue;
}
out_module.used_shared_deps.insert(dep_to_check);
program_info.shared_deps.insert(dep_to_check);
}
}
for out_module in split_points_by_module.keys() {
split_module_contents
.entry(SplitModuleIdentifier::Split(out_module.clone()))
.or_insert_with(|| {
OutputModuleInfo {
is_empty: true,
..OutputModuleInfo::default()
}
});
}
for split_point in &split_points {
program_info
.shared_deps
.insert(DepNode::Function(split_point.export_func));
program_info
.split_point_exports
.insert(split_point.export_func);
}
program_info.split_points = split_points;
program_info.output_modules = split_module_contents.into_iter().collect();
program_info
.output_modules
.sort_by_key(|(identifier, _)| (*identifier).clone());
for (output_index, (_, info)) in program_info.output_modules.iter().enumerate() {
for &symbol in info.included_symbols.iter() {
program_info
.symbol_output_module
.insert(symbol, output_index);
}
}
Ok(program_info)
}