Skip to main content

shape_vm/
stdlib.rs

1//! Standard library compilation for Shape VM
2//!
3//! This module handles compiling the core stdlib modules at engine initialization.
4//! Core modules are auto-imported and available without explicit imports.
5//! Domain-specific modules (finance, iot, etc.) require explicit imports.
6
7use std::path::Path;
8use std::sync::Arc;
9
10use shape_ast::error::{Result, ShapeError};
11use shape_runtime::Runtime;
12use shape_runtime::module_loader::ModuleLoader;
13
14use crate::bytecode::BytecodeProgram;
15use crate::compiler::BytecodeCompiler;
16
17fn stdlib_compile_logs_enabled() -> bool {
18    std::env::var("SHAPE_TRACE_STDLIB_COMPILE")
19        .map(|v| matches!(v.as_str(), "1" | "true" | "TRUE" | "yes" | "YES"))
20        .unwrap_or(false)
21}
22
23/// Pre-compiled core stdlib bytecode (MessagePack-serialized BytecodeProgram).
24/// Regenerate with: cargo run -p stdlib-gen
25#[cfg(not(test))]
26const EMBEDDED_CORE_STDLIB: Option<&[u8]> = Some(include_bytes!("../embedded/core_stdlib.msgpack"));
27
28// Tests always recompile from source to validate compiler changes
29#[cfg(test)]
30const EMBEDDED_CORE_STDLIB: Option<&[u8]> = None;
31
32/// Compile all core stdlib modules into a single BytecodeProgram
33///
34/// The core modules are those in `stdlib/core/` which are auto-imported
35/// and available without explicit import statements.
36///
37/// Uses precompiled embedded bytecode when available, falling back to
38/// source compilation. Set `SHAPE_FORCE_SOURCE_STDLIB=1` to force source.
39///
40/// The compiled program is cached on the passed-in `runtime`; repeat
41/// calls on the same `Runtime` return a cheap clone of the cached
42/// result. Different `Runtime` instances each build their own cache,
43/// which keeps per-Runtime `TypeSchemaRegistry` ids from colliding
44/// across tests that share the same process.
45///
46/// # Returns
47///
48/// A merged BytecodeProgram containing all core functions, types, and metas.
49pub fn compile_core_modules(runtime: &Runtime) -> Result<BytecodeProgram> {
50    let cached: Arc<Result<BytecodeProgram>> = runtime.get_or_init_core_stdlib_cache(|| {
51        Arc::new(load_core_modules_best_effort())
52    });
53    (*cached).clone()
54}
55
56fn load_core_modules_best_effort() -> Result<BytecodeProgram> {
57    // Env override: force source compilation (for debugging/development)
58    if std::env::var("SHAPE_FORCE_SOURCE_STDLIB").is_ok() {
59        return compile_core_modules_from_source();
60    }
61
62    // Try embedded precompiled artifact first
63    if let Some(bytes) = EMBEDDED_CORE_STDLIB {
64        match load_from_embedded(bytes) {
65            Ok(program) => return Ok(program),
66            Err(e) => {
67                if stdlib_compile_logs_enabled() {
68                    eprintln!(
69                        "  Embedded stdlib deserialization failed: {}, falling back to source",
70                        e
71                    );
72                }
73            }
74        }
75    }
76
77    // Fallback: compile from source
78    compile_core_modules_from_source()
79}
80
81fn load_from_embedded(bytes: &[u8]) -> Result<BytecodeProgram> {
82    let mut program: BytecodeProgram =
83        rmp_serde::from_slice(bytes).map_err(|e| ShapeError::RuntimeError {
84            message: format!("Failed to deserialize embedded stdlib: {}", e),
85            location: None,
86        })?;
87    program.ensure_string_index();
88    Ok(program)
89}
90
91/// Extract top-level binding names from precompiled core bytecode.
92/// Used to seed the compiler with known names without loading AST into persistent context.
93///
94/// Consumes the same per-Runtime cache as [`compile_core_modules`].
95pub fn core_binding_names(runtime: &Runtime) -> Vec<String> {
96    match compile_core_modules(runtime) {
97        Ok(program) => {
98            let mut names: Vec<String> = program.functions.iter().map(|f| f.name.clone()).collect();
99            for name in &program.module_binding_names {
100                if !names.contains(name) {
101                    names.push(name.clone());
102                }
103            }
104            names
105        }
106        Err(_) => Vec::new(),
107    }
108}
109
110/// Compile core stdlib from source (parse + compile). Used as fallback and for tests.
111///
112/// Each module is compiled independently (preserving its own scope for builtins
113/// and intrinsics), then the bytecodes are merged via `merge_append`.
114pub fn compile_core_modules_from_source() -> Result<BytecodeProgram> {
115    let trace = stdlib_compile_logs_enabled();
116    if trace {
117        eprintln!("  Compiling core stdlib...");
118    }
119    let mut loader = ModuleLoader::new();
120    let core_modules = loader.list_core_stdlib_module_imports()?;
121    if core_modules.is_empty() {
122        return Ok(BytecodeProgram::new());
123    }
124
125    let mut merged = BytecodeProgram::new();
126    for import_path in core_modules {
127        let file_name = import_path.strip_prefix("std.").unwrap_or(&import_path);
128        match loader.load_module(&import_path).and_then(|module| {
129            BytecodeCompiler::compile_module_ast(&module.ast).map(|(program, _)| program)
130        }) {
131            Ok(module_program) => {
132                if trace {
133                    eprintln!("    Compiled {}", file_name);
134                }
135                merged.merge_append(module_program);
136            }
137            Err(e) => {
138                if trace {
139                    eprintln!("    Warning: failed to compile {}: {}", file_name, e);
140                }
141            }
142        }
143    }
144
145    if trace {
146        eprintln!("  Finished core stdlib compilation");
147    }
148    Ok(merged)
149}
150
151/// Compile all Shape files in a directory (recursively) into a single BytecodeProgram.
152/// Each file is compiled independently, then merged via `merge_append`.
153pub fn compile_directory(dir: &Path) -> Result<BytecodeProgram> {
154    let mut merged = BytecodeProgram::new();
155    compile_directory_into(&mut merged, dir)?;
156    Ok(merged)
157}
158
159/// Recursively compile all Shape files in a directory and merge into the given program.
160fn compile_directory_into(program: &mut BytecodeProgram, dir: &Path) -> Result<()> {
161    let entries = std::fs::read_dir(dir).map_err(|e| ShapeError::ModuleError {
162        message: format!("Failed to read directory {:?}: {}", dir, e),
163        module_path: Some(dir.to_path_buf()),
164    })?;
165
166    for entry in entries {
167        let entry = entry.map_err(|e| ShapeError::ModuleError {
168            message: format!("Failed to read directory entry: {}", e),
169            module_path: Some(dir.to_path_buf()),
170        })?;
171
172        let path = entry.path();
173
174        if path.is_dir() {
175            compile_directory_into(program, &path)?;
176        } else if path.extension().and_then(|s| s.to_str()) == Some("shape") {
177            let file_name = path
178                .file_name()
179                .and_then(|s| s.to_str())
180                .unwrap_or("unknown");
181            match compile_file(&path) {
182                Ok(file_program) => {
183                    eprintln!("    Compiled {}", file_name);
184                    program.merge_append(file_program);
185                }
186                Err(e) => {
187                    eprintln!("    Warning: failed to compile {}: {}", file_name, e);
188                }
189            }
190        }
191    }
192
193    Ok(())
194}
195
196/// Compile an in-memory Shape source string into a BytecodeProgram.
197/// Used for extension-bundled Shape code (e.g., `include_str!("duckdb.shape")`).
198pub fn compile_source(filename: &str, source: &str) -> Result<BytecodeProgram> {
199    let program = shape_ast::parser::parse_program(source).map_err(|e| ShapeError::ParseError {
200        message: format!("Failed to parse {}: {}", filename, e),
201        location: None,
202    })?;
203
204    let mut compiler = BytecodeCompiler::new();
205    compiler.set_source_with_file(source, filename);
206    compiler.compile(&program)
207}
208
209/// Compile a single Shape file into a BytecodeProgram
210pub fn compile_file(path: &Path) -> Result<BytecodeProgram> {
211    let source = std::fs::read_to_string(path).map_err(|e| ShapeError::ModuleError {
212        message: format!("Failed to read file {:?}: {}", path, e),
213        module_path: Some(path.to_path_buf()),
214    })?;
215
216    let program =
217        shape_ast::parser::parse_program(&source).map_err(|e| ShapeError::ParseError {
218            message: format!("Failed to parse {:?}: {}", path, e),
219            location: None,
220        })?;
221
222    let mut compiler = BytecodeCompiler::new();
223    compiler.set_source_with_file(&source, &path.to_string_lossy());
224    compiler.compile(&program)
225}
226
227#[cfg(test)]
228mod tests {
229    use super::*;
230
231    #[test]
232    fn test_core_bytecode_has_snapshot_schema() {
233        let runtime = Runtime::new();
234        let core = compile_core_modules(&runtime).expect("Core modules should compile");
235        let snapshot = core.type_schema_registry.get("Snapshot");
236        assert!(
237            snapshot.is_some(),
238            "Core bytecode should contain Snapshot enum schema"
239        );
240        let enum_info = snapshot.unwrap().get_enum_info();
241        assert!(enum_info.is_some(), "Snapshot should be an enum");
242        let info = enum_info.unwrap();
243        assert!(
244            info.variant_by_name("Hash").is_some(),
245            "Snapshot should have Hash variant"
246        );
247        assert!(
248            info.variant_by_name("Resumed").is_some(),
249            "Snapshot should have Resumed variant"
250        );
251    }
252
253    #[test]
254    fn test_core_bytecode_registers_queryable_trait_dispatch_symbols() {
255        let runtime = Runtime::new();
256        let core = compile_core_modules(&runtime).expect("Core modules should compile");
257        let filter = core.lookup_trait_method_symbol("Queryable", "Table", None, "filter");
258        let map = core.lookup_trait_method_symbol("Queryable", "Table", None, "map");
259        let execute = core.lookup_trait_method_symbol("Queryable", "Table", None, "execute");
260
261        assert_eq!(filter, Some("Table::filter"));
262        assert_eq!(map, Some("Table::map"));
263        assert_eq!(execute, Some("Table::execute"));
264    }
265
266    #[test]
267    fn test_compile_empty_directory() {
268        // Create a temp directory and compile it
269        let temp_dir = std::env::temp_dir().join("shape_test_empty");
270        let _ = std::fs::create_dir_all(&temp_dir);
271
272        let result = compile_directory(&temp_dir);
273        assert!(result.is_ok());
274
275        let program = result.unwrap();
276        // Should have a Halt instruction at minimum
277        assert!(
278            program.instructions.is_empty()
279                || program.instructions.last().map(|i| i.opcode)
280                    == Some(crate::bytecode::OpCode::Halt)
281        );
282
283        let _ = std::fs::remove_dir_all(&temp_dir);
284    }
285
286    #[test]
287    fn test_compile_source_simple_function() {
288        let source = r#"
289            fn double(x) { x * 2 }
290        "#;
291        let result = compile_source("test.shape", source);
292        assert!(
293            result.is_ok(),
294            "compile_source should succeed: {:?}",
295            result.err()
296        );
297
298        let program = result.unwrap();
299        assert!(
300            !program.functions.is_empty(),
301            "Should have at least one function"
302        );
303        assert!(
304            program.functions.iter().any(|f| f.name == "double"),
305            "Should contain 'double' function"
306        );
307    }
308
309    #[test]
310    fn test_compile_source_parse_error() {
311        let source = "fn broken(( { }";
312        let result = compile_source("broken.shape", source);
313        assert!(result.is_err(), "Should fail on invalid syntax");
314    }
315
316    #[test]
317    fn test_compile_source_enum_definition() {
318        let source = r#"
319            enum Direction {
320                Up,
321                Down,
322                Left,
323                Right
324            }
325        "#;
326        let result = compile_source("enums.shape", source);
327        assert!(
328            result.is_ok(),
329            "compile_source should handle enums: {:?}",
330            result.err()
331        );
332    }
333
334    #[test]
335    fn test_embedded_stdlib_round_trip() {
336        // Compile from source, serialize, deserialize, and verify key properties match
337        let source = compile_core_modules_from_source().expect("Source compilation should succeed");
338        let bytes = rmp_serde::to_vec(&source).expect("Serialization should succeed");
339        let deserialized = load_from_embedded(&bytes).expect("Deserialization should succeed");
340
341        assert_eq!(
342            source.functions.len(),
343            deserialized.functions.len(),
344            "Function count should match after round-trip"
345        );
346        assert_eq!(
347            source.instructions.len(),
348            deserialized.instructions.len(),
349            "Instruction count should match after round-trip"
350        );
351        assert_eq!(
352            source.constants.len(),
353            deserialized.constants.len(),
354            "Constant count should match after round-trip"
355        );
356        assert!(
357            !deserialized.functions.is_empty(),
358            "Deserialized should have functions"
359        );
360    }
361
362    #[test]
363    fn test_body_length_within_bounds() {
364        let program = compile_core_modules_from_source().expect("compile");
365        let total = program.instructions.len();
366        let mut bad = Vec::new();
367        for (i, f) in program.functions.iter().enumerate() {
368            let end = f.entry_point + f.body_length;
369            if end > total {
370                bad.push(format!(
371                    "func[{}] '{}' entry={} body_length={} end={} exceeds total={}",
372                    i, f.name, f.entry_point, f.body_length, end, total
373                ));
374            }
375        }
376        assert!(
377            bad.is_empty(),
378            "Functions with OOB body_length:\n{}",
379            bad.join("\n")
380        );
381    }
382
383    #[test]
384    fn test_core_binding_names() {
385        let runtime = Runtime::new();
386        let names = core_binding_names(&runtime);
387        assert!(!names.is_empty(), "Should have binding names from stdlib");
388    }
389}