use std::path::Path;
use std::sync::Arc;
use shape_ast::error::{Result, ShapeError};
use shape_runtime::Runtime;
use shape_runtime::module_loader::ModuleLoader;
use crate::bytecode::BytecodeProgram;
use crate::compiler::BytecodeCompiler;
fn stdlib_compile_logs_enabled() -> bool {
std::env::var("SHAPE_TRACE_STDLIB_COMPILE")
.map(|v| matches!(v.as_str(), "1" | "true" | "TRUE" | "yes" | "YES"))
.unwrap_or(false)
}
#[cfg(not(test))]
const EMBEDDED_CORE_STDLIB: Option<&[u8]> = Some(include_bytes!("../embedded/core_stdlib.msgpack"));
#[cfg(test)]
const EMBEDDED_CORE_STDLIB: Option<&[u8]> = None;
pub fn compile_core_modules(runtime: &Runtime) -> Result<BytecodeProgram> {
let cached: Arc<Result<BytecodeProgram>> = runtime.get_or_init_core_stdlib_cache(|| {
Arc::new(load_core_modules_best_effort())
});
(*cached).clone()
}
fn load_core_modules_best_effort() -> Result<BytecodeProgram> {
if std::env::var("SHAPE_FORCE_SOURCE_STDLIB").is_ok() {
return compile_core_modules_from_source();
}
if let Some(bytes) = EMBEDDED_CORE_STDLIB {
match load_from_embedded(bytes) {
Ok(program) => return Ok(program),
Err(e) => {
if stdlib_compile_logs_enabled() {
eprintln!(
" Embedded stdlib deserialization failed: {}, falling back to source",
e
);
}
}
}
}
compile_core_modules_from_source()
}
fn load_from_embedded(bytes: &[u8]) -> Result<BytecodeProgram> {
let mut program: BytecodeProgram =
rmp_serde::from_slice(bytes).map_err(|e| ShapeError::RuntimeError {
message: format!("Failed to deserialize embedded stdlib: {}", e),
location: None,
})?;
program.ensure_string_index();
Ok(program)
}
pub fn core_binding_names(runtime: &Runtime) -> Vec<String> {
match compile_core_modules(runtime) {
Ok(program) => {
let mut names: Vec<String> = program.functions.iter().map(|f| f.name.clone()).collect();
for name in &program.module_binding_names {
if !names.contains(name) {
names.push(name.clone());
}
}
names
}
Err(_) => Vec::new(),
}
}
pub fn compile_core_modules_from_source() -> Result<BytecodeProgram> {
let trace = stdlib_compile_logs_enabled();
if trace {
eprintln!(" Compiling core stdlib...");
}
let mut loader = ModuleLoader::new();
let core_modules = loader.list_core_stdlib_module_imports()?;
if core_modules.is_empty() {
return Ok(BytecodeProgram::new());
}
let mut merged = BytecodeProgram::new();
for import_path in core_modules {
let file_name = import_path.strip_prefix("std.").unwrap_or(&import_path);
match loader.load_module(&import_path).and_then(|module| {
BytecodeCompiler::compile_module_ast(&module.ast).map(|(program, _)| program)
}) {
Ok(module_program) => {
if trace {
eprintln!(" Compiled {}", file_name);
}
merged.merge_append(module_program);
}
Err(e) => {
if trace {
eprintln!(" Warning: failed to compile {}: {}", file_name, e);
}
}
}
}
if trace {
eprintln!(" Finished core stdlib compilation");
}
Ok(merged)
}
pub fn compile_directory(dir: &Path) -> Result<BytecodeProgram> {
let mut merged = BytecodeProgram::new();
compile_directory_into(&mut merged, dir)?;
Ok(merged)
}
fn compile_directory_into(program: &mut BytecodeProgram, dir: &Path) -> Result<()> {
let entries = std::fs::read_dir(dir).map_err(|e| ShapeError::ModuleError {
message: format!("Failed to read directory {:?}: {}", dir, e),
module_path: Some(dir.to_path_buf()),
})?;
for entry in entries {
let entry = entry.map_err(|e| ShapeError::ModuleError {
message: format!("Failed to read directory entry: {}", e),
module_path: Some(dir.to_path_buf()),
})?;
let path = entry.path();
if path.is_dir() {
compile_directory_into(program, &path)?;
} else if path.extension().and_then(|s| s.to_str()) == Some("shape") {
let file_name = path
.file_name()
.and_then(|s| s.to_str())
.unwrap_or("unknown");
match compile_file(&path) {
Ok(file_program) => {
eprintln!(" Compiled {}", file_name);
program.merge_append(file_program);
}
Err(e) => {
eprintln!(" Warning: failed to compile {}: {}", file_name, e);
}
}
}
}
Ok(())
}
pub fn compile_source(filename: &str, source: &str) -> Result<BytecodeProgram> {
let program = shape_ast::parser::parse_program(source).map_err(|e| ShapeError::ParseError {
message: format!("Failed to parse {}: {}", filename, e),
location: None,
})?;
let mut compiler = BytecodeCompiler::new();
compiler.set_source_with_file(source, filename);
compiler.compile(&program)
}
pub fn compile_file(path: &Path) -> Result<BytecodeProgram> {
let source = std::fs::read_to_string(path).map_err(|e| ShapeError::ModuleError {
message: format!("Failed to read file {:?}: {}", path, e),
module_path: Some(path.to_path_buf()),
})?;
let program =
shape_ast::parser::parse_program(&source).map_err(|e| ShapeError::ParseError {
message: format!("Failed to parse {:?}: {}", path, e),
location: None,
})?;
let mut compiler = BytecodeCompiler::new();
compiler.set_source_with_file(&source, &path.to_string_lossy());
compiler.compile(&program)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_core_bytecode_has_snapshot_schema() {
let runtime = Runtime::new();
let core = compile_core_modules(&runtime).expect("Core modules should compile");
let snapshot = core.type_schema_registry.get("Snapshot");
assert!(
snapshot.is_some(),
"Core bytecode should contain Snapshot enum schema"
);
let enum_info = snapshot.unwrap().get_enum_info();
assert!(enum_info.is_some(), "Snapshot should be an enum");
let info = enum_info.unwrap();
assert!(
info.variant_by_name("Hash").is_some(),
"Snapshot should have Hash variant"
);
assert!(
info.variant_by_name("Resumed").is_some(),
"Snapshot should have Resumed variant"
);
}
#[test]
fn test_core_bytecode_registers_queryable_trait_dispatch_symbols() {
let runtime = Runtime::new();
let core = compile_core_modules(&runtime).expect("Core modules should compile");
let filter = core.lookup_trait_method_symbol("Queryable", "Table", None, "filter");
let map = core.lookup_trait_method_symbol("Queryable", "Table", None, "map");
let execute = core.lookup_trait_method_symbol("Queryable", "Table", None, "execute");
assert_eq!(filter, Some("Table::filter"));
assert_eq!(map, Some("Table::map"));
assert_eq!(execute, Some("Table::execute"));
}
#[test]
fn test_compile_empty_directory() {
let temp_dir = std::env::temp_dir().join("shape_test_empty");
let _ = std::fs::create_dir_all(&temp_dir);
let result = compile_directory(&temp_dir);
assert!(result.is_ok());
let program = result.unwrap();
assert!(
program.instructions.is_empty()
|| program.instructions.last().map(|i| i.opcode)
== Some(crate::bytecode::OpCode::Halt)
);
let _ = std::fs::remove_dir_all(&temp_dir);
}
#[test]
fn test_compile_source_simple_function() {
let source = r#"
fn double(x) { x * 2 }
"#;
let result = compile_source("test.shape", source);
assert!(
result.is_ok(),
"compile_source should succeed: {:?}",
result.err()
);
let program = result.unwrap();
assert!(
!program.functions.is_empty(),
"Should have at least one function"
);
assert!(
program.functions.iter().any(|f| f.name == "double"),
"Should contain 'double' function"
);
}
#[test]
fn test_compile_source_parse_error() {
let source = "fn broken(( { }";
let result = compile_source("broken.shape", source);
assert!(result.is_err(), "Should fail on invalid syntax");
}
#[test]
fn test_compile_source_enum_definition() {
let source = r#"
enum Direction {
Up,
Down,
Left,
Right
}
"#;
let result = compile_source("enums.shape", source);
assert!(
result.is_ok(),
"compile_source should handle enums: {:?}",
result.err()
);
}
#[test]
fn test_embedded_stdlib_round_trip() {
let source = compile_core_modules_from_source().expect("Source compilation should succeed");
let bytes = rmp_serde::to_vec(&source).expect("Serialization should succeed");
let deserialized = load_from_embedded(&bytes).expect("Deserialization should succeed");
assert_eq!(
source.functions.len(),
deserialized.functions.len(),
"Function count should match after round-trip"
);
assert_eq!(
source.instructions.len(),
deserialized.instructions.len(),
"Instruction count should match after round-trip"
);
assert_eq!(
source.constants.len(),
deserialized.constants.len(),
"Constant count should match after round-trip"
);
assert!(
!deserialized.functions.is_empty(),
"Deserialized should have functions"
);
}
#[test]
fn test_body_length_within_bounds() {
let program = compile_core_modules_from_source().expect("compile");
let total = program.instructions.len();
let mut bad = Vec::new();
for (i, f) in program.functions.iter().enumerate() {
let end = f.entry_point + f.body_length;
if end > total {
bad.push(format!(
"func[{}] '{}' entry={} body_length={} end={} exceeds total={}",
i, f.name, f.entry_point, f.body_length, end, total
));
}
}
assert!(
bad.is_empty(),
"Functions with OOB body_length:\n{}",
bad.join("\n")
);
}
#[test]
fn test_core_binding_names() {
let runtime = Runtime::new();
let names = core_binding_names(&runtime);
assert!(!names.is_empty(), "Should have binding names from stdlib");
}
}