diff --git a/docs/CHECKED_PROGRAM.md b/docs/CHECKED_PROGRAM.md new file mode 100644 index 00000000..9c53f8ea --- /dev/null +++ b/docs/CHECKED_PROGRAM.md @@ -0,0 +1,297 @@ +# Checked-program contract + +The compiler records lexical identity, solved types and source locations once in +checked bodies. Specialization supplies concrete function/global targets, and +consumers retain their own analyses, storage and execution rules. The compiler +keeps checked templates alongside optional specialized output; editor recovery +uses a separate `SourceAnalysis` snapshot. + +## Identity and ownership + +| Handle | Owner and meaning | +| --- | --- | +| `ExprID` | A node coordinate in one `CheckedBody`. Shared coordinates do not imply a value snapshot or single evaluation. | +| `LocalId` | A binding in that body, introduced by a parameter, size binder, declaration, loop or lambda parameter. | +| `RequirementId` | A coordinate in that body's interface-requirement inventory; member references also record a `DefId`. | +| `DefId` | A definition in one declaration table, including separately indexed interface members, independent of sorted storage coordinates. | +| `InstanceId` | An entry in one specialized inventory, keyed by definition plus type and size arguments, independent of declaration sorting. | + +Handles have no historical meaning across reparsing, checking or independent +owners. They contain no generation tag. Equal numeric local IDs in different +bodies are expected; callers must carry the owning body/program. Validation +checks coordinates against that owner, but cannot detect a caller substituting +an equal numeric handle from another owner. Names support diagnostics, entry +selection, host layout lookup and emitted symbols; they do not establish local +or concrete target identity. + +A `CheckedBody` is the source expression tree plus side tables indexed by +`ExprID`: the result type, the resolved `Reference` of each identifier or type +application, and the `LocalId`s each declaration or lambda introduces. The tree +keeps source spellings for diagnostics; only the tables establish identity. Each +`Local` owns its binding type and mutability. Checked `let`/`var` nodes have a +`void` result and no source annotation, even when the binding stores an array. +Reading a reference parameter can have type `T` while its local record has type +`&T`; result, storage and target-signature types need not be identical. + +## Publication and validation + +`CheckedProgram::try_new` and `SpecializedProgram::try_from_instances` perform +fallible structural validation. Their `new`/`from_instances` convenience forms +panic on invalid input. Public data remains editable: constructing or mutating +it directly does not produce an immutable proof object or rerun type inference. + +Phase-specific `validate()` checks all retained nodes, including nodes outside +runtime roots: roots/edges are in range and acyclic, referenced locals have +body-owned binders, binders and binding occurrences are unique, and references +and types obey their phase. Declaration tables validate interface-member +inventories; body requirements must belong to their recorded interfaces. +Preconditions and global assumption roots must be boolean. Calls need function +types and matching arity; direct function-instance calls also match the actual +target's parameter count. Validation does not repeat lexical resolution, choose +overloads, prove every typing rule or certify safety or transformation legality. + +The compiler validates templates before specialization, concrete output and its +origins before safety/analysis, and the result again after hoisting. Direct +consumers of mutable public program data must validate before passing it to +another phase and reestablish any semantic or safety guarantees their edits affect. + +## Checked templates + +`Compiler::check()` publishes `CheckedProgram` only when all accumulated inputs +parse and all source bodies and prototypes pass type/resolution and structural +checking. A neighboring invalid function prevents whole-program publication. +Safety-only failures retain checked type facts, but `check()` returns false and +execution remains blocked. An independently constructed `CheckedProgram` only +has the structural guarantees of its constructor. + +Template references identify locals, size parameters, global definitions, ordered +overload sets or members of recorded interface requirements. They contain no +instance references. Named generic types and symbolic array sizes are allowed; +anonymous inference variables are not. An unused generic template may retain an +interface obligation with no current implementation, including no candidates. + +`Reference::Functions` is an overload set. Checking fixes candidate order from +source declaration order; specialization consumes it with the existing inference, +coercion and size-generic precedence rules, without repeating lexical lookup. +Candidates may include function-valued globals; borrowed-call checks use the +solved callee signature when the set contains no function declarations. +Interface selection has a separate policy: first exact signature match in the +recorded order. Generic implementations and reference/array-to-slice coercions do +not qualify merely because they unify. Its top-level `Var`/`Anon` deferral rule +does not permit anonymous inference variables in published templates. + +## Specialized programs + +Every retained function/global declaration has exactly one instance record, every +record owns one such declaration, and keys are unique. Function signatures, +locals, node types, embedded type arguments/annotations, instance type arguments +and global types contain no type variables or symbolic array sizes. Functions +and globals have no generic binders. Interface requirements are discharged and +removed; all retained references are `Local` or `Instance`. + +An instance's inventory target is authoritative. Callable locals and +function-valued globals can still be indirect calls. Operand types and recorded +callee types may differ from the selected function's signature through permitted +coercions; backends still materialize implicit conversions. + +The declaration list also retains generic struct layout definitions, enums, +constants and concrete global assumptions. These are not function/global +instances; struct fields need not be concrete. `ArraySize::Known(0)` retains its +existing unknown/unspecified-size convention. Interfaces and macro declarations +do not survive specialization. Non-generic globals remain present even when +unreachable, and host storage follows the existing sorted declaration order. + +`validate_origins(templates)` checks each instance's source definition kind and +type/size argument counts. `Compiler` retains the owning template inventory, so +origin `DefId`s remain resolvable through `checked_program()`. Standalone concrete +validation cannot resolve origins; clients keeping concrete output independently +must also retain its templates if they need origin queries or validation. + +## Compiler lifecycle and diagnostics + +| Operation | Retained results and execution gate | +| --- | --- | +| `parse(contents, path)` | Appends an input and invalidates templates, concrete output and editor facts, even on failure. Clears derived diagnostics and retains parse errors from all inputs; its return value describes only the new input. | +| `check()` | Revalidates all inputs with current options, replaces templates and clears concrete/editor output. Parse/type failures publish no templates; safety-only failures retain type facts but block specialization. | +| `analyze()` | Uses the same checking pipeline with recovery and all-body diagnostic collection, replacing the editor snapshot and clearing concrete output. Only complete successful type checking can also publish templates. | +| Change effective entry points | Retains templates/editor facts; clears concrete output and specialization diagnostics. Empty roots and explicit `main` are equivalent; order remains significant. Missing entries retain the existing skip policy. | +| `specialize()` | Borrows templates; publishes only after specialization, concrete safety, hoisting and structural validation succeed. Failure publishes no output and leaves templates usable for retry or different roots without rechecking. | +| Compile a backend | Borrows concrete output for current roots and validation options. Requires successful source validation and specialization; consumes neither artifact. | + +`check()` stops at accumulated parse errors before constructing checked bodies, +even with `check_all`. `analyze()` visits recovered syntax for editor facts but +retains the same execution gate. Parsing is additive even for repeated paths; +replacing or removing an input requires rebuilding the compiler. + +Source-validation and specialization diagnostics have separate owners. Public +`last_errors` and `last_safety_errors` combine their current messages; +`last_parse_errors` covers all inputs and `last_type_errors` the last checking +pass. These mutable lists are views, not validation authority: clearing them +cannot authorize execution. Each failed specialization attempt replaces its own +diagnostics; root changes clear only that phase's diagnostics. + +Validation options are snapshotted during checking. Currently `no_recursion` is +the validation-affecting option; future such options must join `ValidationOptions`. +`specialize()` and `specialized_program()` reject reuse while its current value +differs from the checked snapshot and require revalidation. Restoring the checked +value permits reuse; writes alone do not create a generation. Rechecking clears +old output even on failure. `quiet`, `check_all` and `print_ir` control reporting, +diagnostic collection or code generation, not validation policy. Already returned +compiled programs are independently owned and are not revoked. + +Use `checked_program()` for source-semantic queries and `specialized_program()` +for concrete consumers. The compatibility `decls()` view prefers current concrete +declarations, otherwise templates, and panics before type publication. +Executable layout users must call `globals_info_with_offset()` after successful +specialization, with the backend's reserved offset. FFI compilation continues to +check again. Retained templates extend memory lifetime through code generation; +no automatic FFI reuse or additional caching is provided. + +## Partial editor facts + +The LSP calls `analyze()` and queries immutable `SourceAnalysis` for hover and +definition requests. Ordinary `check()` does not retain the extra source inventory +and fact vectors. Parsing and either checking entry point replace old snapshots; +root changes preserve them. The LSP rebuilds from all open documents on edits and +publishes diagnostics for every document, including empty lists after repair. + +A snapshot owns the normalized, macro-expanded source declaration inventory and +its `DefId -> BodyAnalysis` map. Early macro failure retains only the unexpanded +inventory. Expression coordinates belong to the snapshot's source function arena, +not a checked/concrete arena. Local records retain names, optional types and +binder locations; requirements identify interface declarations and can have gaps +after errors. Cloning a snapshot clones its owners together. No IDs survive edits. + +`ExpressionFacts` separates optional type, recorded reference and declaration +binder facts. Missing/unvisited facts are `None`. An empty unresolved overload +set supplies neither a target nor inferred type; a nonempty set still represents +candidates. Arithmetic without a named overload supplies no target reference. +Type admission is deliberately bounded: + +- Successful, unrecovered bodies expose solved types without anonymous variables + and with valid named annotations, including declared generic/size parameters. +- Failed bodies expose no types justified only by solver substitutions. A concrete + type after failed unification, a bad call/cast or incomplete struct literal is + insufficient; inferred local types are withheld too. +- Independently established facts can survive body failure: range-checked suffixed + integers, fixed-type real/string/character/boolean literals and reads of valid + explicitly typed parameters/locals or fixed loop/size bindings. These use + original annotations, not failed substitutions. Unsuffixed integer inference, + unresolved/malformed annotations and void value bindings supply no type. +- All body types from a parse-damaged file are withheld, including expansions with + that provenance. Clean files retain useful facts, but dependencies on recovered + signatures/type declarations invalidate inferred caller facts. Independently + established facts can still survive in those callers. + +Established references can support navigation even when a type is unavailable. +Local navigation follows the recorded binder; overload navigation retains its +exact-type/first-candidate heuristic, not semantic target selection. Field +navigation requires a known base type. Parameter locations currently identify +the enclosing function/lambda. Missing facts never trigger a substitute resolver. + +Partial facts cannot authorize specialization. Complete safety-error-only input +retains type/reference facts, but incomplete programs receive no safety analysis. +Parser synchronization, lost/unvisited declarations, precise token spans and +richer overload presentation remain editor limitations. There is no cross-edit +identity/cache scheme or completion handler. + +## Safety and identified body operations + +Source safety checks definitions, including unreachable ones. Its template call +policy uses the first recorded function with matching arity and exact signature; +explicit applications, unmatched generic/interface candidates and size-generic +bodies defer to concrete checking. Source safety failure blocks specialization. + +After specialization, the existing safety traversal checks every retained concrete +function body before field hoisting or backend lowering, including ordinary +callers/wrappers and type/size instances. Lambda bodies use the enclosing +traversal's existing rules. A direct function `InstanceId` supplies the exact +callee and its `require` clauses, including externs and selected interface +implementations. Safety retains an arity guard but performs no signature-equality +selection or assertion. Explicit applications become instance reads as well. + +For example, specialization rejects `bounded(-1, true)` when +`bounded(x: i32, value: T) require x >= 0 {}`. An ordinary wrapper must establish +that requirement through its own contract or guard; a valid value at the wrapper's +caller does not specialize the wrapper on that value. + +Concrete checking uses the existing proof language: call-site proofs support +`true`, conjunctions and the existing less-than/greater-or-equal interval and +array-length/size patterns. Unsupported clauses fail conservatively; even a valid +`require x == 0` call can be unprovable. Local function values and function-valued +globals remain indirect and carry no direct-call contract. Structural validation +is not a safety proof, and this boundary supplies no general higher-order analysis. + +Equivalent failed requirements are deduplicated by source call location, +requirement location and concrete clause text; the first diagnostic retains its +concrete callee name. Distinct sites, requirement origins or substituted bounds +remain distinct. Other safety diagnostics use location/message deduplication. + +Safety storage roots use identities and field paths. Call substitution maps +callee parameter IDs to caller arguments and analyzes callee-only expressions +in a fresh context. Only nonlocal assumption facts cross body boundaries; equal +numeric locals cannot transfer proofs. Normalization, checking, specialization +and safety take assumption bodies/roots directly. Function signature, +borrowed-return and escape restrictions stay at function boundaries; lvalue and +borrowed-call checks use the owning body. No synthetic assumption function or +name-based body adapter is needed. + +Borrowed-argument assignability and no-alias checking share one classification. +Recorded overloads conservatively union borrowed positions; explicit type +arguments are substituted before classification. Indirect/function-valued calls +use solved signatures, so a local callee follows its own signature under +shadowing. References require assignability; references and slices participate +in the existing syntactic base/field-path alias check. These operations do not +select another overload or establish general memory disjointness. + +## Captures, transformations and consumer responsibilities + +Shared free-local discovery uses `BindingFacts` over checked bodies or in-progress +checker records, without requiring solved types or whole-program publication. +It excludes internal binders by identity, visits nested lambdas for construction +dependencies, and retains first-use order without duplicate captures. Size +parameters are compile-time values. Incomplete results are explicit: an empty +list with missing facts cannot prove noncapture; known free locals still establish +capture. Escape taint uses `LocalId` and capturing/noncapturing/unknown states and +continues to diagnose established capture after type errors. + +Escape checking retains its existing scope: explicit returns, initializer taint, +block/if combinations and conservative propagation from callees/arguments (which +can reject scalar results). Implicit returns, assignment propagation, returns +inside lambda bodies, aggregate escapes and broader escape soundness are separate +work. Backends share capture discovery/extraction and own addressable storage, +closure ABI and linking. Extracted lambdas are `CheckedFunction`s with cloned +enclosing arenas and backend-managed queues/patches. They are backend extraction +views, not independently published program bodies, new source instances or a +program-wide lowered lambda inventory. + +Whole-body clones preserve expression/local/requirement coordinates and source +locations under a new owner. `CheckedBody::duplicate` freshens nodes and internal +binders, remaps their uses and preserves enclosing references and source locations. +Shared reads are allowed; shared subtrees declaring bindings need freshening. +Macro occurrence normalization precedes checking so expansions resolve in their +actual lexical scopes. `replace` keeps a coordinate/location and takes the new +result type, retaining recorded references/binders only when they still apply +to the new expression. + +Derived analyses must be recomputed after input changes. Validation and preserved +coordinates/locations neither refresh analyses nor prove behavior preservation. +Hoisting refreshes capture/alias/read/write facts and preserves the global-write/ +call information in its precomputed may-write summary; that summary is not a +purity proof. Consumers still own value snapshots, effect order, alias/dependence +proofs, code-motion legality and cancellation behavior. + +| Consumer | Facts consumed; responsibility retained | +| --- | --- | +| Specialization | Definition/argument keys and recorded candidates; reachability, selection, recursion limits and unique emitted symbols. | +| Safety | Identified storage and concrete callees; proof grammar, call obligations and conservative failure. | +| Field hoisting | Instance-keyed global-write summaries; borrowed/global aliases, captured storage, opaque calls and movement legality. | +| Copy elision | Local identities and binding types; live ranges and escape conditions. | +| Cranelift, LLVM, register VM, Stack | Local/instance maps and ordered captures; representation recovery, implicit conversions, aggregate copies, storage, ABI and runtime checks. | +| VM expression inlining | Complete callee `BodyContext`, restored after emission; eligibility, allocation and code emission. Current eligibility excludes declarations, loops, calls, returns and lambdas. | +| LSP and host APIs | Source facts/concrete inventories; presentation, entry policy, layout offsets, buffer validity and compiled-program lifetime. | + +Runtime cancellation remains a backend responsibility; LLVM AOT omits callback +cancellation. Fully selected ordinary template callees, explicit conversion nodes, +shared closure lowering, broader safety proofs and new optimization/value IRs are +separate work. diff --git a/docs/MONOMORPHIZATION.md b/docs/MONOMORPHIZATION.md index e0c92121..04d5c308 100644 --- a/docs/MONOMORPHIZATION.md +++ b/docs/MONOMORPHIZATION.md @@ -1,271 +1,125 @@ -# Monomorphization Pass - -This document describes the monomorphization implementation for the Lyte compiler. - -## Overview - -Monomorphization is the process of generating specialized versions of generic functions for each unique set of concrete type arguments they're called with. This transforms generic code into concrete, type-specific code. - -For example: -```lyte -id(x: T) → T { x } - -main { - let a = id(42) // Generates id$i32 - let b = id(true) // Generates id$bool -} -``` - -## Architecture - -### Module: `src/monomorph_pass.rs` - -The monomorphization pass is implemented in a separate module that can be invoked between type checking and JIT compilation. - -### Key Components - -#### `MonomorphPass` - -The main struct that manages the monomorphization process: - -```rust -pub struct MonomorphPass { - instantiations: HashMap, // Tracks what's been generated - recursion_detector: RecursionDetector, // Detects infinite recursion - specialized_decls: Vec, // Newly generated declarations - worklist: VecDeque, // Functions to process - processed: HashSet, // Functions already processed -} -``` - -#### Algorithm - -The pass uses a **demand-driven** approach: - -1. **Start from entry point** (e.g., `main`) -2. **Walk the call graph** by processing each function's body -3. **For each generic call**: - - Compute concrete type arguments - - Check if already instantiated - - Check for infinite recursion - - Generate specialized version - - Add to worklist -4. **Repeat** until worklist is empty - -### Type Argument Inference - -The pass infers concrete type arguments from call sites by examining the resolved types after type checking. This is simpler than full Hindley-Milner inference since types are already known. - -### Name Mangling - -Specialized functions are given mangled names using the existing `mangle_name` function: -- `id` → `id$i32` -- `map` → `map$i32$bool` -- `process<[f32;10]>` → `process$[f32;10]` - -### Type Substitution - -For each specialized version: -1. Create an `Instance` (type substitution map) from type parameters to type arguments -2. Clone the generic function declaration -3. Substitute all type variables in: - - Return type - - Parameter types - - Body expression types -4. Clear the `typevars` field (no longer generic) - -### Infinite Recursion Detection - -The pass uses the existing `RecursionDetector` from `src/monomorph.rs` to detect: -- Direct recursion: `foo>()` calling `foo>>()` -- Mutually recursive: `f()` calling `g>()` calling `f>>()` - -## Integration Points - -### Where to Hook In - -The monomorphization pass should be called in `src/compiler.rs`: - -```rust -pub fn check(&mut self) -> bool { - // Parse and collect declarations - let mut decls = self.collect_decls(); - - // Type check all declarations - self.typecheck_decls(&mut decls)?; - - // NEW: Monomorphize generics - let specialized = self.monomorphize(&mut decls)?; - decls.extend(specialized); - - // Freeze into immutable DeclTable for JIT - self.decls = DeclTable::new(decls); - - true -} -``` - -### Calling Convention - -```rust -fn monomorphize(&mut self, decls: &DeclTable) -> Result, String> { - let mut pass = MonomorphPass::new(); - let entry_point = Name::str("main"); // or configurable - pass.monomorphize(decls, entry_point) -} -``` - -## Testing - -The module includes 26 unit tests covering: - -- **Basic instantiation**: Single generic parameter, multiple parameters -- **Deduplication**: Same type args don't create duplicates -- **Type substitution**: Nested types, arrays, tuples -- **Constraints**: Interface constraints are preserved -- **Expression traversal**: All expression types handled -- **Recursion detection**: Uses existing infrastructure -- **Edge cases**: Non-generic functions, empty declarations - -### Running Tests - -```bash -cargo test --lib monomorph_pass::tests +# Checked programs and specialization + +Specialization consumes checked templates and produces concrete function/global +targets for safety analysis and code generation. The authoritative ownership, +lifecycle, editor and mutation rules are in the +[checked-program contract](CHECKED_PROGRAM.md). + +```text +source syntax → CheckedProgram → SpecializedProgram → backend IR + checking specialization, + and source concrete safety, + safety field hoisting, + final validation ``` -## Current Limitations - -The current implementation has some simplifications that will need to be enhanced: - -1. **Type Argument Inference**: Currently uses a simplified heuristic. May need to be enhanced to handle complex cases. - -2. **Call Site Resolution**: Currently only handles direct `Expr::Id` calls. Doesn't yet handle: - - Function values assigned to variables - - Higher-order function calls - - Method calls - -3. **Generic Structs**: The pass focuses on functions. Struct monomorphization is partially handled by the existing `Instance` mechanism in the JIT. - -4. **Separate Compilation**: All monomorphization happens at link time from a single entry point. - -## Future Enhancements - -### 1. Enhanced Type Inference - -Improve `infer_type_arguments` to: -- Match function types against generic signatures -- Handle partial application -- Support higher-rank types - -### 2. Incremental Monomorphization - -Support multiple entry points for library compilation: -```rust -pub fn monomorphize_multiple(&mut self, entry_points: &[Name]) -> Result, String> -``` - -### 3. Optimization Opportunities - -After monomorphization: -- Dead code elimination (remove unused specializations) -- Cross-function inlining -- Specialization-specific optimizations - -### 4. Better Diagnostics - -- Report which generic function caused infinite recursion -- Show the type argument chain -- Suggest fixes (e.g., add runtime bounds) - -### 5. Struct Monomorphization - -Explicitly generate specialized struct declarations: -```rust -struct Vec { ... } -// Generate: -// struct Vec$i32 { ... } -// struct Vec$bool { ... } -``` - -## Examples - -### Example 1: Simple Generic Function - -**Input:** -```lyte -id(x: T) → T { x } - -main { - let a = id(42) -} -``` - -**Generated:** -```lyte -id$i32(x: i32) → i32 { x } - -main { - let a = id$i32(42) -} -``` - -### Example 2: Multiple Type Parameters - -**Input:** -```lyte -map(a: [T0], f: T0 → T1) → [T1] { ... } - -main { - let result = map([1, 2, 3], |x| x > 0) -} -``` - -**Generated:** -```lyte -map$i32$bool(a: [i32], f: i32 → bool) → [bool] { ... } - -main { - let result = map$i32$bool([1, 2, 3], |x| x > 0) -} +`Compiler` retains checked templates alongside optional concrete output. Changing +effective entry points invalidates only concrete output and its diagnostics; +different roots can specialize without rechecking. Parsing invalidates all derived +results. Checking replaces templates; changes to validation policy such as +`no_recursion` require revalidation before execution. Code generation borrows a +successfully specialized program and consumes neither artifact. + +`analyze()` runs the shared checker with partial editor publication. Its separate +`SourceAnalysis` owns recovered source facts for hover/navigation; incomplete facts +never serve as executable input. Parsed syntax remains available for diagnostics +and editing, but backends consume checked bodies. + +## Bodies and identities + +`CheckedBody` owns the expression tree, its per-expression type, reference and +binder tables, local records and interface requirements. Binding types belong to +`LocalId` records; a checked `let`/`var` statement is `void` and has no remaining +source annotation. Source and checked bodies share the same `Expr` type; the +checked tables, not the spelled names, carry resolution and binding identity. + +`DefId` identifies a checked definition, including an interface member. +`InstanceId` identifies a concrete function/global inventory entry. `ExprID`, +`LocalId` and `RequirementId` belong to one body. Equal numeric coordinates in +different owners are unrelated; spelling and symbol mangling are not semantic +identity. A size binder links its `LocalId` to the type-level `ArraySize::Var` +symbol, so diagnostic renaming does not affect size substitution. + +`DeclTable` owns definition indexing over a `DeclarationList`. Concrete programs +own a declaration list and instance inventory. The common list provides nominal +layout and host symbol lookup without source identity APIs. Declaration sorting +preserves definition/instance identities by remapping storage coordinates. + +## Specialization and publication + +`MonomorphPass::monomorphize_multi` validates its checked input and starts from +entry definitions. It interns a `MonomorphKey` of definition, concrete type +arguments and size arguments, reserving an instance before visiting its body. +Recursive calls and repeated reachability share that instance. The recursion +guard still rejects increasingly complex recursive type specializations. + +Templates can retain named generics, symbolic sizes and deferred interface +obligations, but no anonymous inference variables. `Reference::Functions` stores +ordered candidates, not a selected callee. Ordinary specialization retains its +inference/coercion and size-generic precedence rules over those recorded IDs. +Interface selection separately uses the first exact signature match, without +generic overload unification or coercions. Interface selections remain in the +owning body's specialization frame while recursive callees are instantiated. + +Specialization substitutes node/local types, resolves size parameters, selects +candidates/requirements and rewrites nonlocal references to `InstanceId`s. +Generic globals use the same interning mechanism: repeated use of one concrete +global shares storage. Non-generic globals remain present even when unreachable. +Final sorting preserves the existing host layout order. + +Fulfilled requirements, interfaces and macros are removed. Function/global +signatures and bodies are concrete, and all retained references are `Local` or +`Instance`, including nodes outside runtime roots. Generic struct definitions +remain for layout; concrete global assumptions remain body/condition records. +No specialized body returns to source syntax for another type check. + +Fallible constructors validate structure and complete instance inventories; +origin validation checks definition kind and type/size argument counts against +the retained templates. All retained concrete function bodies then run through +the existing safety traversal, including ordinary callers of generics. Direct +function instances supply exact contracts without signature filtering. This +precedes field hoisting; structural validation runs again before compiler +publication. Failure retains templates and publishes no concrete program. + +## Analyses and lowering + +Safety owns its existing interval/constraint proofs. Identified storage roots +and fresh callee contexts prevent equal numeric locals in different bodies from +sharing facts. Indirect function values carry no direct-call contract, and +unsupported precondition syntax can still fail conservatively. + +Field hoisting uses instance-keyed may-write summaries and accounts for borrowed/ +global aliases, captured storage, transitive writes and opaque calls. Binding +identity alone proves neither disjoint memory nor safe code motion. Copy elision +retains its own liveness/escape rules. Mutating public checked data requires +validation and recomputation of affected analyses; coordinates and source +locations do not keep old results valid. + +Backends use local/instance maps and shared ordered free-local discovery. They +retain representation recovery, implicit conversions, captured storage, closure +ABI and generated-lambda queues/patches. Extracted lambdas clone enclosing arenas; +they are not additional source instances. VM expression inlining switches the +complete callee body/storage context and restores its caller. Runtime cancellation +belongs to backend lowering; LLVM AOT retains its omission of callback cancellation. + +## Maintaining the contract + +Boundary tests live in `src/checked.rs` and `src/checked/validate.rs`; compiler +lifecycle/safety/assumption tests and LSP recovery tests exercise phase consumers. +The CLI golden corpus covers backend execution and emitted-code checks, with +per-case check-only flags and backend exclusions. Build the CLI before each +workspace golden run; the runner uses Cargo's `CARGO_BIN_EXE_lyte`, including with +a custom target directory: + +```sh +cargo build --workspace +cargo test --workspace +cargo build --workspace --features llvm +cargo test --workspace --features llvm +cargo check --lib --no-default-features ``` -### Example 3: Nested Generics - -**Input:** -```lyte -process(arr: [T; 10]) → T { ... } - -main { - let arr: [i32; 10] - let x = process(arr) -} -``` - -**Generated:** -```lyte -process$i32(arr: [i32; 10]) → i32 { ... } - -main { - let arr: [i32; 10] - let x = process$i32(arr) -} -``` - -## Implementation Status - -- ✅ Core monomorphization infrastructure -- ✅ Type substitution -- ✅ Name mangling -- ✅ Infinite recursion detection -- ✅ Comprehensive unit tests -- ✅ Integration with compiler (not yet added) -- ⏳ Enhanced type inference -- ⏳ Struct monomorphization -- ⏳ Integration tests with real code - -## References - -- **Name Mangling**: `src/monomorph.rs` - `mangle_name()` -- **Recursion Detection**: `src/monomorph.rs` - `RecursionDetector` -- **Type Substitution**: `src/types.rs` - `TypeID::subst()` -- **Declaration Table**: `src/decl_table.rs` +LLVM requires the configured LLVM toolchain. Production C Stack runtime/golden +suites require the supported Clang-built interpreter; Rust StackVM tests are +separate. The assembly VM requires AArch64. AOT emission tests establish +object/header generation, not execution on the target device. diff --git a/lsp/src/analysis.rs b/lsp/src/analysis.rs index 9f7c7464..fc54c0b0 100644 --- a/lsp/src/analysis.rs +++ b/lsp/src/analysis.rs @@ -49,7 +49,7 @@ impl AnalysisState { compiler.parse(text, &path); } - compiler.check(); + compiler.analyze(); self.compiler = Some(compiler); } } diff --git a/lsp/src/goto_def.rs b/lsp/src/goto_def.rs index f85bd7c4..19748c34 100644 --- a/lsp/src/goto_def.rs +++ b/lsp/src/goto_def.rs @@ -1,7 +1,7 @@ use crate::analysis::{self, AnalysisState}; use crate::hover::find_expr_at; use lsp_types::*; -use lyte::{Decl, Expr, Loc, Name, Type}; +use lyte::{BodyAnalysis, Decl, DeclTable, Expr, Loc, Name, Reference, Type, TypeID}; pub fn handle_goto_definition( state: &AnalysisState, @@ -16,68 +16,76 @@ pub fn handle_goto_definition( let line = pos.line + 1; let col = pos.character + 1; - let decls = compiler.decls(); + let analysis = compiler.source_analysis()?; + let decls = analysis.declarations(); - for decl in &decls.decls { - if let Decl::Func(func) = decl { - if func.loc.file != file_name { - continue; - } - if let Some((id, _)) = find_expr_at(func, file_name, line, col) { - match &func.arena.exprs[id] { - Expr::Id(name) => { - if let Some(loc) = find_decl_loc(decls, *name) { - return Some(loc_to_response(&loc)); - } - } - Expr::Call(callee, _) => { - if let Expr::Id(name) = &func.arena.exprs[*callee] { - if let Some(loc) = find_decl_loc(decls, *name) { - return Some(loc_to_response(&loc)); - } - } - } - Expr::Field(base_id, field_name) => { - // Try to resolve the base type and find the field declaration. - if let Some(&base_ty) = func.types.get(*base_id) { - if let Type::Name(struct_name, _) = &*base_ty { - let found = decls.find(*struct_name); - for d in found { - if let Some(field) = d.find_field(field_name) { - return Some(loc_to_response(&field.loc)); - } - } - } + for (index, decl) in decls.decls.iter().enumerate() { + let Decl::Func(func) = decl else { continue }; + if func.loc.file != file_name { + continue; + } + let Some(body) = analysis.body(decls.id_at(index)) else { + continue; + }; + let Some((id, _)) = find_expr_at(func, file_name, line, col) else { + continue; + }; + let target = match &func.arena[id] { + Expr::Call(callee, _) => *callee, + Expr::Field(base, field_name) => { + if let Some(Type::Name(struct_name, _)) = + body.expression(*base).and_then(|facts| facts.ty).as_deref() + { + for decl in decls.find(*struct_name) { + if let Some(field) = decl.find_field(field_name) { + return Some(loc_to_response(&field.loc)); } } - _ => {} } + continue; + } + _ => id, + }; + let Some(facts) = body.expression(target) else { + continue; + }; + if let Some(reference) = &facts.reference { + if let Some(loc) = find_decl_loc(decls, body, reference, facts.ty) { + return Some(loc_to_response(&loc)); } } } - None } -/// Find the source location of the first declaration with the given name. -/// Skips stdlib declarations (those in files starting with '<'). -fn find_decl_loc(decls: &lyte::DeclTable, name: Name) -> Option { - let found = decls.find(name); - for decl in found { - match decl { - Decl::Func(f) => { - // Skip stdlib functions. - if f.loc.file.starts_with('<') { - continue; - } - return Some(f.loc); - } - _ => { - // Other decl types don't have loc yet; skip for now. - } +/// Use recorded identities even when the expression has no established type. +/// Exact-type/first-candidate navigation is a presentation heuristic, not call +/// selection. Never fall back to spelling for an unresolved or shadowed use. +fn find_decl_loc( + decls: &DeclTable, + body: &BodyAnalysis, + reference: &Reference, + ty: Option, +) -> Option { + match reference { + Reference::Local(local) | Reference::SizeParameter(local) => { + body.local(*local).map(|local| local.loc) + } + Reference::Functions(candidates) => { + let functions: Vec<_> = candidates + .iter() + .filter_map(|id| decls.function(*id)) + .filter(|f| !f.loc.file.starts_with('<')) + .collect(); + functions + .iter() + .find(|f| ty.is_some() && f.annotated_ty() == ty) + .or_else(|| functions.first()) + .map(|f| f.loc) } + Reference::InterfaceMember { member, .. } => decls.function(*member).map(|f| f.loc), + _ => None, } - None } fn loc_to_response(loc: &Loc) -> GotoDefinitionResponse { diff --git a/lsp/src/hover.rs b/lsp/src/hover.rs index 7d3f1a95..1a58b3b6 100644 --- a/lsp/src/hover.rs +++ b/lsp/src/hover.rs @@ -1,6 +1,6 @@ use crate::analysis::{self, AnalysisState}; use lsp_types::*; -use lyte::{Decl, Expr, ExprID, FuncDecl, Name}; +use lyte::{BodyAnalysis, Decl, Expr, ExprID, FuncDecl, Name}; pub fn handle_hover(state: &AnalysisState, params: &HoverParams) -> Option { let compiler = state.compiler()?; @@ -13,14 +13,18 @@ pub fn handle_hover(state: &AnalysisState, params: &HoverParams) -> Option Option Option { +fn hover_in_func( + func: &FuncDecl, + body: &BodyAnalysis, + file: Name, + line: u32, + col: u32, +) -> Option { let (id, _) = find_expr_at(func, file, line, col)?; - - if let Some(&ty) = func.types.get(id) { - let type_str = ty.pretty_print(); - let label = match &func.arena.exprs[id] { - Expr::Id(name) => format!("{}", name), - Expr::Field(_, name) => format!(".{}", name), - _ => String::new(), - }; - - if label.is_empty() { - Some(format!("```lyte\n{}\n```", type_str)) - } else { - Some(format!("```lyte\n{}: {}\n```", label, type_str)) + let facts = body.expression(id)?; + let (ty, label) = match &func.arena[id] { + Expr::Id(name) | Expr::TypeApp(name, _) => (facts.ty?, name.to_string()), + Expr::Let(..) | Expr::Var(..) => { + let local = body.local(facts.binding?)?; + (local.ty?, local.name.to_string()) } + Expr::Field(_, name) => (facts.ty?, format!(".{}", name)), + _ => (facts.ty?, String::new()), + }; + Some(if label.is_empty() { + format!("```lyte\n{}\n```", ty.pretty_print()) } else { - None - } + format!("```lyte\n{}: {}\n```", label, ty.pretty_print()) + }) } /// Find the expression in a function whose location best matches the cursor. @@ -67,7 +75,7 @@ pub fn find_expr_at( let mut best_id: Option = None; let mut best_col: u32 = 0; - for (id, loc) in func.arena.locs.iter().enumerate() { + for (id, &loc) in func.arena.locs.iter().enumerate() { if loc.file == file && loc.line == line && loc.col <= col && loc.col >= best_col { best_col = loc.col; best_id = Some(id); diff --git a/lsp/src/main.rs b/lsp/src/main.rs index 925580ae..b73a229d 100644 --- a/lsp/src/main.rs +++ b/lsp/src/main.rs @@ -5,6 +5,8 @@ mod analysis; mod diagnostics; mod goto_def; mod hover; +#[cfg(test)] +mod recovery_tests; fn main() { let (connection, io_threads) = Connection::stdio(); @@ -41,7 +43,7 @@ fn main_loop(connection: &Connection, state: &mut analysis::AnalysisState) { } } -fn handle_request(connection: &Connection, state: &mut analysis::AnalysisState, req: Request) { +fn handle_request(connection: &Connection, state: &analysis::AnalysisState, req: Request) { if let Some(params) = cast_request::(&req) { let result = hover::handle_hover(state, ¶ms); let resp = Response::new_ok(req.id, serde_json::to_value(result).unwrap()); diff --git a/lsp/src/recovery_tests.rs b/lsp/src/recovery_tests.rs new file mode 100644 index 00000000..84200342 --- /dev/null +++ b/lsp/src/recovery_tests.rs @@ -0,0 +1,445 @@ +//! Exercise the real notification, request dispatch and diagnostic publication +//! paths. Every edit rebuilds the compiler, just as it does in a live session. +use super::*; +use serde_json::{json, Value}; +use std::collections::HashMap; + +struct Editor { + state: analysis::AnalysisState, + server: Connection, + client: Connection, + documents: HashMap, + diagnostics: HashMap>, +} + +impl Editor { + fn new() -> Self { + let (server, client) = Connection::memory(); + Self { + state: analysis::AnalysisState::new(), + server, + client, + documents: HashMap::new(), + diagnostics: HashMap::new(), + } + } + + fn edit(&mut self, file: &str, text: &str) { + let uri = format!("file:///{}", file); + let (method, params) = if self.documents.insert(file.into(), text.into()).is_some() { + ( + "textDocument/didChange", + json!({"textDocument": {"uri": uri, "version": 2}, "contentChanges": [{"text": text}]}), + ) + } else { + ( + "textDocument/didOpen", + json!({"textDocument": {"uri": uri, "version": 1, "languageId": "lyte", "text": text}}), + ) + }; + // Clear the observed messages to require a fresh publication for every + // open file, including an empty list when old errors have disappeared. + self.diagnostics.clear(); + handle_notification( + &self.server, + &mut self.state, + Notification::new(method.into(), params), + ); + while let Ok(message) = self.client.receiver.try_recv() { + let Message::Notification(notification) = message else { + panic!("expected diagnostics") + }; + assert_eq!(notification.method, "textDocument/publishDiagnostics"); + let params: PublishDiagnosticsParams = + serde_json::from_value(notification.params).unwrap(); + self.diagnostics + .insert(params.uri.as_str().into(), params.diagnostics); + } + assert_eq!(self.diagnostics.len(), self.documents.len()); + } + + fn request(&self, file: &str, needle: &str, method: &str) -> Value { + let source = &self.documents[file]; + let offset = source.find(needle).expect("query text in document"); + let prefix = &source[..offset]; + let line = prefix.bytes().filter(|&b| b == b'\n').count(); + let character = prefix.rsplit('\n').next().unwrap().encode_utf16().count(); + let request = Request::new( + 1.into(), + method.into(), + json!({ + "textDocument": {"uri": format!("file:///{}", file)}, + "position": {"line": line, "character": character} + }), + ); + handle_request(&self.server, &self.state, request); + let Message::Response(response) = self.client.receiver.try_recv().unwrap() else { + panic!("expected response") + }; + assert!(response.error.is_none(), "{:?}", response.error); + response.result.unwrap() + } + + fn hover(&self, file: &str, needle: &str) -> Option { + let result = self.request(file, needle, "textDocument/hover"); + if result.is_null() { + return None; + } + Some(result["contents"]["value"].as_str().unwrap().into()) + } + + fn definition(&self, file: &str, needle: &str) -> Option { + serde_json::from_value(self.request(file, needle, "textDocument/definition")).unwrap() + } + + fn assert_hover(&self, file: &str, needle: &str, text: &str) { + assert_eq!( + self.hover(file, needle).as_deref(), + Some(format!("```lyte\n{}\n```", text).as_str()) + ); + } + + fn assert_definition(&self, file: &str, needle: &str, target_file: &str, line: u32) { + let location = self + .definition(file, needle) + .expect("established reference"); + assert_eq!(location.uri.as_str(), format!("file:///{}", target_file)); + assert_eq!(location.range.start.line, line); + } + + fn diagnostics(&self, file: &str) -> &[Diagnostic] { + &self.diagnostics[&format!("file:///{}", file)] + } + + fn assert_blocked(&self) { + let compiler = self.state.compiler().unwrap(); + assert!(compiler.specialized_program().is_err()); + assert!(compiler.compile_vm().is_err()); + assert!(compiler.compile_stack().is_err()); + } +} + +#[test] +fn valid_function_beside_type_invalid_function_in_same_or_other_file() { + for separate in [false, true] { + let mut editor = Editor::new(); + let good = "target(x: i32) -> i32 { x }\ngood() -> i32 {\n let x = 1\n target(x)\n}"; + let bad = "bad() -> i32 { true }"; + editor.edit( + "good.lyte", + &if separate { + good.into() + } else { + format!("{}\n{}", good, bad) + }, + ); + if separate { + editor.edit("bad.lyte", bad); + } + assert!(editor.state.compiler().unwrap().checked_program().is_none()); + editor.assert_hover("good.lyte", "x)", "x: i32"); + editor.assert_definition("good.lyte", "target(x)", "good.lyte", 0); + assert!(!editor + .diagnostics(if separate { "bad.lyte" } else { "good.lyte" }) + .is_empty()); + editor.assert_blocked(); + } +} + +#[test] +fn generic_types_and_local_references_survive_invalid_neighbors() { + let mut editor = Editor::new(); + editor.edit( + "generic.lyte", + "identity(x: T) -> T {\n let value = x\n value\n}", + ); + editor.edit("bad.lyte", "bad() -> i32 { missing }"); + editor.assert_hover("generic.lyte", "x\n", "x: T"); + editor.assert_hover("generic.lyte", "value\n", "value: T"); + editor.assert_definition("generic.lyte", "value\n", "generic.lyte", 1); + editor.assert_definition("generic.lyte", "x\n", "generic.lyte", 0); +} + +#[test] +fn failing_body_keeps_identity_separate_from_failed_inference() { + let mut editor = Editor::new(); + editor.edit("failing.lyte", "target(x: i32) -> i32 { x }\nbroken(p: T) -> i32 {\n var number: i32\n let inferred = target(true)\n number\n p\n inferred\n}"); + editor.assert_hover("failing.lyte", "number\n", "number: i32"); + editor.assert_hover("failing.lyte", "p\n", "p: T"); + // The solver can leave concrete types for these despite rejecting the call. + assert!(editor.hover("failing.lyte", "target(true)").is_none()); + assert!(editor.hover("failing.lyte", "inferred\n").is_none()); + editor.assert_definition("failing.lyte", "target(true)", "failing.lyte", 0); + editor.assert_definition("failing.lyte", "inferred\n", "failing.lyte", 3); + editor.assert_definition("failing.lyte", "number\n", "failing.lyte", 2); + editor.assert_hover("failing.lyte", "true", "bool"); + editor.assert_blocked(); +} + +#[test] +fn unavailable_types_and_unresolved_names_do_not_invent_answers() { + let mut editor = Editor::new(); + editor.edit("unknown.lyte", "target() -> i32 { 1 }\nbroken(p: Unknown) {\n let cast = true as i32\n let target = missing\n p\n cast\n target\n 42\n 999999999999i32\n}"); + for needle in [ + "p\n", + "cast\n", + "target\n", + "missing", + "42\n", + "999999999999i32", + ] { + assert!(editor.hover("unknown.lyte", needle).is_none(), "{}", needle); + } + assert!(editor.definition("unknown.lyte", "missing").is_none()); + // The local shadows the top-level function, even with no established type. + editor.assert_definition("unknown.lyte", "target\n", "unknown.lyte", 3); + editor.assert_definition("unknown.lyte", "p\n", "unknown.lyte", 1); +} + +#[test] +fn unresolved_explicit_calls_withhold_contextual_types_but_keep_bindings() { + for call in [ + "missing⟨i32⟩(1)", + "identity⟨i32, bool⟩(1)", + "unused⟨i32⟩(1)", + ] { + let mut editor = Editor::new(); + editor.edit( + "explicit.lyte", + &format!( + "bad() -> i32 {{ unknown }}\nprobe() -> i32 {{\n let value = {}\n true\n value\n}}\nidentity(x: T) -> T {{ x }}\nunused(x: T) -> T {{ x }}", + call + ), + ); + // The surrounding return constraint can solve `value` to i32 even + // though the explicit application has no candidate. Neither the use + // nor the binding may publish that inferred type. + for needle in [call, "let value", "value\n"] { + assert!( + editor.hover("explicit.lyte", needle).is_none(), + "{}: {}", + call, + needle + ); + } + assert!( + editor.definition("explicit.lyte", call).is_none(), + "{}", + call + ); + editor.assert_definition("explicit.lyte", "value\n", "explicit.lyte", 2); + editor.assert_hover("explicit.lyte", "true\n", "bool"); + editor.assert_blocked(); + } +} + +#[test] +fn builtin_arithmetic_keeps_types_without_named_overload_candidates() { + let mut editor = Editor::new(); + editor.edit( + "arithmetic.lyte", + "probe() -> i32 {\n let value = 1 + 2\n value\n}\nbad() -> i32 { unknown }", + ); + editor.assert_hover("arithmetic.lyte", "value\n", "value: i32"); + editor.assert_definition("arithmetic.lyte", "value\n", "arithmetic.lyte", 1); + editor.assert_blocked(); +} + +#[test] +fn parse_recovery_preserves_other_files_without_publishing_placeholders() { + for malformed in [ + "bad() -> i32 { 1wat }", + "bad() -> [i32; 1] { [1wat] }", + "bad() -> i32 { true + ) }", + "bad() -> i32 { abs(1wat) }", + "bad() { let x = 1; x. }", + "fn broken( {", + ] { + let mut editor = Editor::new(); + editor.edit("good.lyte", "good() -> i32 {\n let x = 1\n x\n}"); + editor.edit("parse.lyte", malformed); + assert!( + !editor + .state + .compiler() + .unwrap() + .last_parse_errors + .is_empty(), + "{}", + malformed + ); + assert!(editor.state.compiler().unwrap().checked_program().is_none()); + editor.assert_hover("good.lyte", "x\n", "x: i32"); + editor.assert_definition("good.lyte", "x\n", "good.lyte", 1); + assert!(!editor.diagnostics("parse.lyte").is_empty()); + editor.assert_blocked(); + } +} + +#[test] +fn recovered_declaration_cannot_supply_a_confident_type_to_clean_caller() { + let mut editor = Editor::new(); + // Recovery substitutes void for the malformed return annotation. Even a + // fully solved clean caller must not expose that placeholder as a fact. + editor.edit("damaged.lyte", "recovered() -> ?\n"); + editor.edit( + "caller.lyte", + "caller() {\n recovered()\n}\nunrelated() -> i32 {\n let value = 42\n value\n}", + ); + assert!(editor.state.compiler().unwrap().last_type_errors.is_empty()); + assert!(editor.hover("caller.lyte", "recovered()").is_none()); + editor.assert_definition("caller.lyte", "recovered()", "damaged.lyte", 0); + editor.assert_hover("caller.lyte", "value\n", "value: i32"); + editor.assert_blocked(); +} + +#[test] +fn safety_diagnostics_preserve_types_and_references_but_block_execution() { + let mut editor = Editor::new(); + editor.edit("safety.lyte", "broken() -> i32 {\n let x = 1\n x / 0\n}"); + assert!(editor.state.compiler().unwrap().checked_program().is_some()); + editor.assert_hover("safety.lyte", "x /", "x: i32"); + editor.assert_definition("safety.lyte", "x /", "safety.lyte", 1); + assert!(editor + .diagnostics("safety.lyte") + .iter() + .any(|d| d.severity == Some(DiagnosticSeverity::WARNING))); + editor.assert_blocked(); +} + +#[test] +fn edits_replace_facts_and_publish_empty_diagnostics_after_repair() { + let mut editor = Editor::new(); + editor.edit("edit.lyte", "main() -> i32 {\n let old = 1\n old\n}"); + editor.assert_hover("edit.lyte", "old\n", "old: i32"); + assert!(editor.diagnostics("edit.lyte").is_empty()); + editor.edit( + "edit.lyte", + "main() -> i32 {\n let new = missing\n old\n}", + ); + assert!(editor.hover("edit.lyte", "old\n").is_none()); + assert!(editor.definition("edit.lyte", "old\n").is_none()); + assert!(!editor.diagnostics("edit.lyte").is_empty()); + editor.edit("edit.lyte", "main() -> bool {\n let new = true\n new\n}"); + editor.assert_hover("edit.lyte", "new\n", "new: bool"); + editor.assert_definition("edit.lyte", "new\n", "edit.lyte", 1); + assert!(editor.diagnostics("edit.lyte").is_empty()); + assert!(editor.state.compiler().unwrap().checked_program().is_some()); +} + +#[test] +fn field_navigation_in_unaffected_body_survives_errors() { + let mut editor = Editor::new(); + editor.edit( + "field.lyte", + "struct S {\n value: i32\n}\nget(s: S) -> i32 {\n s.value\n}\nbad() -> i32 { missing }", + ); + editor.assert_hover("field.lyte", ".value", ".value: i32"); + editor.assert_definition("field.lyte", ".value", "field.lyte", 1); +} + +#[test] +fn incomplete_parameter_and_generic_field_queries_do_not_panic() { + for source in [ + "bad(x) -> i32 { x }\ncaller() -> i32 { bad(1) }", + "struct Box { value: T }\ncaller(x: Box) -> i32 { x.value }", + "bad(x) -> i32 { 1 }\ncaller() -> i32 { bad⟨i32⟩(1) }", + "__add(x) -> i32 { 1 }\ncaller() -> i32 { 1 + 2 }", + "interface I { apply(x) -> T }\ncaller(x: T) -> T where I { apply(x) }", + "interface I { apply(x: T) -> T }\napply(x) -> i32 { 1 }\ncaller(x: T) -> T where I { apply(x) }\nmain() -> i32 { caller(1) }", + "macro incomplete(x: i32)\ncaller() -> i32 { @incomplete(1) }", + "macro incomplete(x: i32\ncaller() -> i32 { @incomplete(1) }", + ] { + let mut editor = Editor::new(); + editor.edit("incomplete.lyte", source); + assert!(!editor.diagnostics("incomplete.lyte").is_empty()); + editor.assert_blocked(); + } +} + +#[test] +fn recovered_body_keeps_bindings_and_never_hovers_error_nodes() { + let mut editor = Editor::new(); + editor.edit( + "recover.lyte", + "broken() -> i32 {\n let value = 1wat\n value\n}", + ); + assert!(editor.hover("recover.lyte", "1wat").is_none()); + assert!(editor.hover("recover.lyte", "value\n").is_none()); + editor.assert_definition("recover.lyte", "value\n", "recover.lyte", 1); +} + +#[test] +fn recovered_struct_cannot_supply_field_types_to_clean_body() { + let mut editor = Editor::new(); + editor.edit("damaged.lyte", "struct Damaged { field: ? }"); + editor.edit("caller.lyte", "caller(x: Damaged) {\n x.field\n}"); + assert!(editor.hover("caller.lyte", "x.field").is_none()); + assert!(editor.hover("caller.lyte", ".field").is_none()); + assert!(editor.definition("caller.lyte", ".field").is_none()); + editor.assert_definition("caller.lyte", "x.field", "caller.lyte", 0); +} + +#[test] +fn generic_interface_facts_keep_their_requirement_and_declaration_owners() { + let mut editor = Editor::new(); + editor.edit("interface.lyte", "interface I { apply(x: T) -> T }\ngood(x: T) -> T where I {\n let result = apply(x)\n result\n}\nbad(x: T) -> T where Missing I {\n apply(x)\n}"); + editor.assert_hover("interface.lyte", "result\n", "result: T"); + editor.assert_definition("interface.lyte", "apply(x)\n", "interface.lyte", 0); + // Inspect ownership as well as sending requests: the failed first where + // clause leaves a gap before the recorded requirement used by `apply`. + let analysis = editor.state.compiler().unwrap().source_analysis().unwrap(); + let decls = analysis.declarations(); + for record in decls.records() { + let Some(body) = analysis.body(record.definition) else { + continue; + }; + let Some(function) = decls.function(record.definition) else { + continue; + }; + for id in 0..function.arena.exprs.len() { + if let Some(lyte::Reference::InterfaceMember { + requirement, + member, + }) = &body.expression(id).unwrap().reference + { + let owner = body + .requirement_interface(*requirement) + .expect("requirement owner"); + assert!(decls.interface_members(owner).contains(member)); + assert!(decls.function(*member).is_some()); + assert!(decls.definition(owner).is_some()); + } + } + } +} + +#[test] +fn successful_literal_types_and_invalid_annotations_are_distinguished() { + let mut editor = Editor::new(); + editor.edit("annotations.lyte", "struct Box { value: T }\nbroken(a: (i32, Unknown), b: Box, c: Box) {\n a\n b\n c\n 1.0f64\n 7i32\n missing\n}"); + for needle in ["a\n", "b\n"] { + assert!(editor.hover("annotations.lyte", needle).is_none()); + editor.assert_definition("annotations.lyte", needle, "annotations.lyte", 1); + } + editor.assert_hover("annotations.lyte", "c\n", "c: Box"); + editor.assert_hover("annotations.lyte", "1.0f64", "f64"); + editor.assert_hover("annotations.lyte", "7i32", "i32"); +} + +#[test] +fn invalid_void_binding_keeps_identity_without_claiming_a_value_type() { + let mut editor = Editor::new(); + editor.edit("void.lyte", "bad() {\n var ghost: void\n ghost\n}"); + assert!(editor.hover("void.lyte", "var ghost").is_none()); + assert!(editor.hover("void.lyte", "ghost\n").is_none()); + editor.assert_definition("void.lyte", "ghost\n", "void.lyte", 1); +} + +#[test] +fn calls_with_deferred_size_parameters_preserve_successful_body_types() { + let mut editor = Editor::new(); + editor.edit("sizes.lyte", "count(values: [i32; N]) -> i32 { N }\ngood() -> i32 {\n let values = [1, 2]\n let total = count(values)\n total\n}\nbad() -> i32 { missing }"); + editor.assert_hover("sizes.lyte", "total\n", "total: i32"); + editor.assert_definition("sizes.lyte", "count(values)", "sizes.lyte", 0); +} diff --git a/src/README.md b/src/README.md index fee5497a..923014bb 100644 --- a/src/README.md +++ b/src/README.md @@ -1,6 +1,10 @@ # Notes +The [checked-program boundary](../docs/CHECKED_PROGRAM.md) specifies the distinct +template and specialized-program guarantees, validation, lifecycle/editor behavior, +and consumer responsibilities. + ✅ = implemented feature ❌ = not planned diff --git a/src/checked.rs b/src/checked.rs new file mode 100644 index 00000000..8ceea171 --- /dev/null +++ b/src/checked.rs @@ -0,0 +1,1023 @@ +//! The program after lexical and type checking. +//! +//! A body keeps the source expression tree and records solved types, resolved +//! references and binding identities in side tables indexed by `ExprID`. +//! Expression and local IDs are handles in one body, not historical identities: +//! cloning a whole body preserves the handles, while duplication inside that +//! body freshens its declarations. Analyses must be recomputed after mutation. +//! +//! See `docs/CHECKED_PROGRAM.md` for the separate template and concrete contracts. +use crate::*; +use std::collections::{HashMap, HashSet}; +use std::convert::TryInto; +use std::ops::{Deref, Index}; + +mod validate; + +macro_rules! identity { + ($name:ident) => { + #[derive(Clone, Copy, Debug, Eq, PartialEq, Hash, Ord, PartialOrd)] + pub struct $name(pub u32); + impl $name { + pub fn index(self) -> usize { + self.0 as usize + } + } + }; +} +identity!(DefId); +identity!(InstanceId); +identity!(LocalId); +identity!(RequirementId); + +#[derive(Clone, Debug, Eq, PartialEq, Hash)] +pub enum Reference { + Local(LocalId), + SizeParameter(LocalId), + Global(DefId), + /// An ordered overload set, not a selected callee. Specialization consumes + /// these recorded candidates without repeating source-name lookup. + Functions(Vec), + InterfaceMember { + requirement: RequirementId, + member: DefId, + }, + Instance(InstanceId), +} + +#[derive(Clone, Debug, Eq, PartialEq, Hash)] +pub struct CheckedParam { + pub local: LocalId, +} + +/// Links a checked size binder to the existing type-level array-size symbol. +/// Local diagnostic spelling is independent of this semantic correspondence. +#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)] +pub struct SizeParameter { + pub symbol: Name, + pub local: LocalId, +} + +#[derive(Clone, Debug, Eq, PartialEq, Hash)] +pub struct Local { + /// Diagnostic spelling; references never resolve this name again. + pub name: Name, + pub ty: TypeID, + pub mutable: bool, +} + +pub type CheckedDecl = Decl; +pub type CheckedDeclTable = DeclTable; +pub type CheckedDeclarations = DeclarationList; + +/// A checked body: the source expression tree plus side tables indexed by +/// `ExprID`. The tree keeps its source spellings for diagnostics; solved +/// types, resolved references and binding identities live in the tables. +#[derive(Clone, Debug, Default, Eq, PartialEq, Hash)] +pub struct CheckedBody { + syntax: ExprArena, + types: Vec, + /// The recorded resolution of each `Id`/`TypeApp` node; `None` elsewhere. + references: Vec>, + /// The locals introduced by each `Let`/`Var`/`For` node (one) or `Lambda` + /// node (one per parameter); empty elsewhere. + binders: Vec>, + pub locals: Vec, + pub requirements: Vec, +} + +impl CheckedBody { + pub fn new() -> Self { + Self::default() + } + pub fn from_parts( + syntax: ExprArena, + types: Vec, + references: Vec>, + binders: Vec>, + locals: Vec, + requirements: Vec, + ) -> Self { + let n = syntax.exprs.len(); + assert!( + syntax.locs.len() == n + && types.len() == n + && references.len() == n + && binders.len() == n, + "checked body tables must cover every expression" + ); + Self { + syntax, + types, + references, + binders, + locals, + requirements, + } + } + pub fn len(&self) -> usize { + self.syntax.exprs.len() + } + pub fn is_empty(&self) -> bool { + self.syntax.exprs.is_empty() + } + /// Every expression handle in this body. + pub fn ids(&self) -> std::ops::Range { + 0..self.len() + } + /// The expression tree with its source locations. + pub fn syntax(&self) -> &ExprArena { + &self.syntax + } + pub fn exprs(&self) -> &[Expr] { + &self.syntax.exprs + } + pub fn ty(&self, id: ExprID) -> TypeID { + self.types[id] + } + pub fn loc(&self, id: ExprID) -> Loc { + self.syntax.locs[id] + } + /// The recorded resolution of an `Id` or `TypeApp` node. Other nodes + /// record nothing. + pub fn reference(&self, id: ExprID) -> Option<&Reference> { + self.references[id].as_ref() + } + /// The locals a node introduces: one for `let`/`var`/`for`, one per + /// lambda parameter, none otherwise. + pub fn binders(&self, id: ExprID) -> &[LocalId] { + &self.binders[id] + } + /// The local introduced by a `let`, `var` or `for` node. + pub fn binder(&self, id: ExprID) -> LocalId { + match &self.syntax.exprs[id] { + Expr::Let(..) | Expr::Var(..) | Expr::For { .. } => *self.binders[id] + .first() + .unwrap_or_else(|| panic!("expression {} records no binding", id)), + _ => panic!("expression {} is not a declaration", id), + } + } + pub fn local(&self, id: LocalId) -> &Local { + &self.locals[id.index()] + } + pub fn add_local(&mut self, name: Name, ty: TypeID, mutable: bool) -> LocalId { + let id = LocalId(self.locals.len().try_into().expect("too many locals")); + self.locals.push(Local { name, ty, mutable }); + id + } + /// Append an expression with no recorded reference or binders. + pub fn add(&mut self, expr: Expr, ty: TypeID, loc: Loc) -> ExprID { + let id = self.syntax.add(expr, loc); + self.types.push(ty); + self.references.push(None); + self.binders.push(vec![]); + id + } + /// Append an identifier with its resolution. + pub fn add_id(&mut self, name: Name, reference: Reference, ty: TypeID, loc: Loc) -> ExprID { + let id = self.add(Expr::Id(name), ty, loc); + self.references[id] = Some(reference); + id + } + /// Append a read of a body-local binding, spelled with the local's name. + pub fn add_local_read(&mut self, local: LocalId, ty: TypeID, loc: Loc) -> ExprID { + self.add_id(self.local(local).name, Reference::Local(local), ty, loc) + } + /// Append a declaration or lambda together with the locals it introduces. + pub fn add_binding( + &mut self, + expr: Expr, + binders: Vec, + ty: TypeID, + loc: Loc, + ) -> ExprID { + let id = self.add(expr, ty, loc); + self.set_binders(id, binders); + id + } + /// Append `let local = init`, spelled with the local's name. + pub fn add_let(&mut self, local: LocalId, init: ExprID, loc: Loc) -> ExprID { + let name = self.local(local).name; + self.add_binding( + Expr::Let(name, init, None), + vec![local], + mk_type(Type::Void), + loc, + ) + } + /// Append `var local = init`, spelled with the local's name. + pub fn add_var(&mut self, local: LocalId, init: Option, loc: Loc) -> ExprID { + let name = self.local(local).name; + self.add_binding( + Expr::Var(name, init, None), + vec![local], + mk_type(Type::Void), + loc, + ) + } + pub fn set_ty(&mut self, id: ExprID, ty: TypeID) { + self.types[id] = ty; + } + /// Record the resolution of an `Id` or `TypeApp` node. + pub fn set_reference(&mut self, id: ExprID, reference: Reference) { + assert!( + matches!(self.syntax.exprs[id], Expr::Id(_) | Expr::TypeApp(..)), + "expression {} is not a reference", + id + ); + self.references[id] = Some(reference); + } + /// Record the locals a declaration or lambda introduces. + pub fn set_binders(&mut self, id: ExprID, binders: Vec) { + let expected = match &self.syntax.exprs[id] { + Expr::Let(..) | Expr::Var(..) | Expr::For { .. } => 1, + Expr::Lambda { params, .. } => params.len(), + _ => panic!("expression {} declares no bindings", id), + }; + assert_eq!( + binders.len(), + expected, + "expression {} declares {} bindings", + id, + expected + ); + self.binders[id] = binders; + } + /// Replace an expression, keeping its handle and location. Recorded facts + /// that still apply to the new expression are kept: a reference for an + /// `Id`/`TypeApp`, binders for a declaration or a lambda with the same + /// parameter count. Any other recorded fact is cleared. + pub fn replace(&mut self, id: ExprID, expr: Expr, ty: TypeID) { + let binders = match &expr { + Expr::Let(..) | Expr::Var(..) | Expr::For { .. } => 1, + Expr::Lambda { params, .. } => params.len(), + _ => 0, + }; + if !matches!(expr, Expr::Id(_) | Expr::TypeApp(..)) { + self.references[id] = None; + } + if self.binders[id].len() != binders { + self.binders[id].clear(); + } + self.syntax.exprs[id] = expr; + self.types[id] = ty; + } + pub fn substitute(&mut self, instance: &Instance) { + for ty in &mut self.types { + *ty = ty.subst(instance); + } + for expr in &mut self.syntax.exprs { + match expr { + Expr::AsTy(_, ty) => *ty = ty.subst(instance), + Expr::TypeApp(_, args) => { + for ty in args { + *ty = ty.subst(instance); + } + } + Expr::Let(_, _, ty) | Expr::Var(_, _, ty) => { + if let Some(ty) = ty { + *ty = ty.subst(instance); + } + } + _ => {} + } + } + for local in &mut self.locals { + local.ty = local.ty.subst(instance); + } + for requirement in &mut self.requirements { + *requirement = requirement.subst(instance); + } + } + + /// Source-oriented diagnostics over checked expressions. + pub fn pretty_print(&self, id: ExprID, indent: usize) -> String { + self.syntax.exprs[id].pretty_print(&self.syntax, indent) + } + + /// Duplicate evaluations. Bindings declared in the copied subtree receive + /// fresh local IDs; references to enclosing bindings keep their identities. + /// Whole-body cloning instead preserves every body-local index. + /// The source subtree must have unique binding occurrences, as required by + /// program validation; shared reads can occur more than once. + pub fn duplicate(&mut self, root: ExprID) -> ExprID { + fn declarations(body: &CheckedBody, id: ExprID, found: &mut HashSet) { + found.extend(body.binders(id).iter().copied()); + for child in body[id].subexprs() { + declarations(body, child, found); + } + } + let mut declared = HashSet::new(); + declarations(self, root, &mut declared); + let mut declared: Vec<_> = declared.into_iter().collect(); + declared.sort(); + let mut locals = HashMap::new(); + for old in declared { + let local = self.local(old).clone(); + locals.insert(old, self.add_local(local.name, local.ty, local.mutable)); + } + fn copy(body: &mut CheckedBody, id: ExprID, locals: &HashMap) -> ExprID { + let remap = |local: LocalId| locals.get(&local).copied().unwrap_or(local); + let mut expr = body.syntax.exprs[id].clone(); + let ty = body.types[id]; + let loc = body.syntax.locs[id]; + let reference = body.references[id].clone().map(|reference| match reference { + Reference::Local(local) => Reference::Local(remap(local)), + Reference::SizeParameter(local) => Reference::SizeParameter(remap(local)), + other => other, + }); + let binders: Vec<_> = body.binders[id].iter().map(|&local| remap(local)).collect(); + expr.map_children(|child| copy(body, child, locals)); + let new = body.add(expr, ty, loc); + body.references[new] = reference; + body.binders[new] = binders; + new + } + copy(self, root, &locals) + } + + /// Free local references of a lambda, in source evaluation order. Descending + /// into nested lambdas includes the captures needed to construct them. + pub fn captures(&self, root: ExprID, bound: impl IntoIterator) -> Vec { + crate::free_locals::free_locals(self, root, bound).locals + } + pub fn captured_locals(&self) -> HashSet { + let mut captured = HashSet::new(); + for id in self.ids() { + if let Expr::Lambda { body, .. } = &self[id] { + captured.extend(self.captures(*body, self.binders(id).iter().copied())); + } + } + captured + } +} +impl Index for CheckedBody { + type Output = Expr; + fn index(&self, id: ExprID) -> &Self::Output { + &self.syntax.exprs[id] + } +} + +#[derive(Clone, Debug, Eq, PartialEq, Hash)] +pub struct CheckedFunction { + pub name: Name, + pub typevars: Vec, + pub size_vars: Vec, + pub params: Vec, + pub body: Option, + pub ret: TypeID, + pub requires: Vec, + pub loc: Loc, + pub arena: CheckedBody, + pub closure_vars: Vec, + pub is_extern: bool, +} +impl CheckedFunction { + pub fn param_types(&self) -> Vec { + self.params + .iter() + .map(|param| self.arena.local(param.local).ty) + .collect() + } + pub fn domain(&self) -> TypeID { + mk_type(Type::Tuple(self.param_types())) + } + pub fn ty(&self) -> TypeID { + func(self.domain(), self.ret) + } + pub fn captured_locals(&self) -> HashSet { + self.arena.captured_locals() + } + pub fn extract_lambda(&self, expression: ExprID, name: Name) -> Self { + let Expr::Lambda { body, .. } = &self.arena[expression] else { + panic!("expected lambda"); + }; + let Type::Func(_, ret) = &*self.arena.ty(expression) else { + panic!("checked lambda type"); + }; + let locals = self.arena.binders(expression); + Self { + name, + typevars: Vec::new(), + size_vars: Vec::new(), + params: locals.iter().map(|&local| CheckedParam { local }).collect(), + body: Some(*body), + ret: *ret, + requires: Vec::new(), + loc: self.arena.loc(expression), + arena: self.arena.clone(), + closure_vars: self.arena.captures(*body, locals.iter().copied()), + is_extern: false, + } + } +} +impl FunctionInfo for CheckedFunction { + type Arena = CheckedBody; + fn name(&self) -> Name { + self.name + } + fn ty(&self) -> TypeID { + self.ty() + } +} + +#[derive(Clone, Debug)] +pub struct CheckedProgram { + pub decls: CheckedDeclTable, +} +impl CheckedProgram { + pub fn new(decls: CheckedDeclTable) -> Self { + Self::try_new(decls).expect("invalid checked program") + } + pub fn try_new(decls: CheckedDeclTable) -> Result { + let program = Self { decls }; + program.validate()?; + Ok(program) + } + pub fn function(&self, definition: DefId) -> Option<&CheckedFunction> { + self.decls.function(definition) + } +} +impl Deref for CheckedProgram { + type Target = CheckedDeclarations; + fn deref(&self) -> &Self::Target { + &self.decls + } +} + +#[derive(Clone, Debug, Eq, PartialEq, Hash)] +pub struct InstanceRecord { + pub definition: DefId, + pub type_args: Vec, + pub size_args: Vec, + /// Storage coordinate only; semantic references use InstanceId. + pub declaration: usize, +} + +/// Concrete bodies and instance targets. Source declaration IDs remain origins; +/// they are never used to choose among several concrete specializations. +#[derive(Clone, Debug)] +pub struct SpecializedProgram { + pub decls: CheckedDeclarations, + pub instances: Vec, +} +impl SpecializedProgram { + pub fn from_instances( + declarations: Vec, + instances: Vec, + ) -> Self { + Self::try_from_instances(declarations, instances).expect("invalid specialized program") + } + pub fn try_from_instances( + declarations: Vec, + mut instances: Vec, + ) -> Result { + let mut sorted: Vec<_> = declarations.into_iter().enumerate().collect(); + sorted.sort_by_key(|(_, declaration)| declaration.name()); + let mut remap = vec![0; sorted.len()]; + for (new, (old, _)) in sorted.iter().enumerate() { + remap[*old] = new; + } + for instance in &mut instances { + instance.declaration = *remap + .get(instance.declaration) + .ok_or("instance declaration is outside the program")?; + } + let decls = DeclarationList::from_sorted( + sorted + .into_iter() + .map(|(_, declaration)| declaration) + .collect(), + ); + let program = Self { decls, instances }; + program.validate()?; + Ok(program) + } + pub fn instance(&self, id: InstanceId) -> &CheckedDecl { + &self.decls.decls[self.instances[id.index()].declaration] + } + pub fn function_instance(&self, id: InstanceId) -> Option<&CheckedFunction> { + match self.instance(id) { + Decl::Func(function) => Some(function), + _ => None, + } + } + pub fn instance_name(&self, id: InstanceId) -> Name { + self.instance(id).name() + } + pub fn functions(&self) -> impl Iterator { + self.instances + .iter() + .enumerate() + .filter_map(move |(index, _)| { + let id = InstanceId(index as u32); + self.function_instance(id).map(|function| (id, function)) + }) + } + pub fn globals(&self) -> impl Iterator { + self.instances + .iter() + .enumerate() + .filter_map(move |(index, _)| { + let id = InstanceId(index as u32); + let decl = self.instance(id); + matches!(decl, Decl::Global { .. }).then_some((id, decl)) + }) + } + /// Instance inventory in storage order, preserving host global/extern layout. + pub fn storage_instances(&self) -> impl Iterator { + let mut ids: Vec<_> = self + .instances + .iter() + .enumerate() + .map(|(index, record)| (record.declaration, InstanceId(index as u32))) + .collect(); + ids.sort_by_key(|(coordinate, _)| *coordinate); + ids.into_iter().map(move |(_, id)| (id, self.instance(id))) + } + pub fn instance_for_entry(&self, name: Name) -> Option { + self.functions() + .find_map(|(id, function)| (function.name == name).then_some(id)) + } +} +impl Deref for SpecializedProgram { + type Target = CheckedDeclarations; + fn deref(&self) -> &Self::Target { + &self.decls + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn duplication_freshens_internal_bindings_and_keeps_enclosing_references() { + let mut body = CheckedBody::new(); + let ty = mk_type(Type::Int32); + let outer = body.add_local(Name::str("x"), ty, false); + let inner = body.add_local(Name::str("x"), ty, false); + let outer_read = body.add_local_read(outer, ty, test_loc()); + let binding = body.add_let(inner, outer_read, test_loc()); + let inner_read = body.add_local_read(inner, ty, test_loc()); + let root = body.add(Expr::Block(vec![binding, inner_read]), ty, test_loc()); + let copy = body.duplicate(root); + let Expr::Block(children) = &body[copy] else { + panic!() + }; + let Expr::Let(_, initializer, _) = body[children[0]] else { + panic!() + }; + let fresh = body.binder(children[0]); + assert_ne!(fresh, inner); + assert_eq!(body.reference(initializer), Some(&Reference::Local(outer))); + assert_eq!(body.reference(children[1]), Some(&Reference::Local(fresh))); + assert_eq!(body.local(fresh), body.local(inner)); + assert_eq!(body.ty(copy), body.ty(root)); + assert_eq!(body.loc(copy), body.loc(root)); + assert_eq!(body.reference(inner_read), Some(&Reference::Local(inner))); + } + + #[test] + fn captures_follow_bindings_through_shadowing_and_nested_closures() { + let mut body = CheckedBody::new(); + let ty = mk_type(Type::Int32); + let callable = func(tuple(vec![]), ty); + let outer = body.add_local(Name::str("x"), callable, true); + let inner = body.add_local(Name::str("x"), ty, false); + let read_outer = body.add_local_read(outer, callable, test_loc()); + let read_inner = body.add_local_read(inner, ty, test_loc()); + let applied_outer = body.add(Expr::TypeApp(Name::str("x"), vec![]), callable, test_loc()); + body.set_reference(applied_outer, Reference::Local(outer)); + let nested_root = body.add( + Expr::Block(vec![read_inner, applied_outer, read_outer, read_inner]), + ty, + test_loc(), + ); + let nested = body.add_binding( + Expr::Lambda { + params: vec![], + body: nested_root, + }, + vec![], + func(mk_type(Type::Tuple(vec![])), ty), + test_loc(), + ); + // Preserve first use, including explicit applications and repeated reads, + // rather than sorting captures by ID or diagnostic spelling. + assert_eq!(body.captures(nested_root, []), vec![inner, outer]); + assert_eq!(body.captures(nested, [inner]), vec![outer]); + } + + #[test] + fn whole_body_substitution_retains_local_and_requirement_handles() { + let mut body = CheckedBody::new(); + let generic = typevar("T"); + let local = body.add_local(Name::str("x"), generic, false); + let read = body.add_local_read(local, generic, test_loc()); + body.requirements.push(InterfaceRequirement { + id: RequirementId(0), + interface: DefId(1), + type_args: vec![generic], + members: vec![InterfaceMember { + definition: DefId(2), + signature: generic, + candidates: vec![DefId(3)], + }], + }); + let mut copy = body.clone(); + let instance: Instance = [(generic, mk_type(Type::Int32))].iter().copied().collect(); + copy.substitute(&instance); + assert_eq!(copy[read], body[read]); + assert_eq!(copy.ty(read), mk_type(Type::Int32)); + assert_eq!(copy.local(local).ty, mk_type(Type::Int32)); + assert_eq!(body.ty(read), generic); + assert_eq!(copy.requirements[0].id, body.requirements[0].id); + assert_eq!(copy.requirements[0].members[0].definition, DefId(2)); + assert_eq!(copy.requirements[0].members[0].candidates, vec![DefId(3)]); + assert_eq!(copy.requirements[0].type_args, vec![mk_type(Type::Int32)]); + assert_eq!( + copy.requirements[0].members[0].signature, + mk_type(Type::Int32) + ); + } + + #[cfg(any(feature = "cranelift", feature = "llvm"))] + #[test] + fn checked_captures_preserve_native_parameter_storage() { + for source in [ + "capture(value: i32, create: bool) -> i32 { if create { let read = || { value }; }; value } main() -> i32 { capture(42, false) }", + "capture(values: [i32; 2]) -> i32 { var total = 0; for i in 0 .. 2 { let read = || { values[0] }; total = total + read() }; total } main() -> i32 { capture([21, 0]) }", + "capture(value: &i32) -> i32 { for i in 0 .. 2 { let increment = || { value = value + 1 }; increment() }; value } main() -> i32 { var value = 40; capture(value) }", + ] { + let mut compiler = Compiler::new(); + compiler.quiet = true; + assert!(compiler.parse(source, ".")); + assert!(compiler.check()); + compiler.specialize().unwrap(); + let program = compiler.specialized_program().unwrap(); + #[cfg(feature = "cranelift")] + { + let mut jit = crate::JIT::default(); + let (entry, size) = jit.compile(program).unwrap(); + let mut globals = vec![0u8; size]; + unsafe { + crate::cancel::set_cancel_callback(globals.as_mut_ptr(), None, std::ptr::null_mut()); + let entry: extern "C" fn(*mut u8, *mut u8) -> i32 = std::mem::transmute(entry); + assert_eq!(entry(globals.as_mut_ptr(), std::ptr::null_mut()), 42, "Cranelift: {}", source); + } + jit.free_memory(); + } + #[cfg(feature = "llvm")] + { + let jit = crate::LLVMJIT::new(); + let compiled = jit.compile_only(program, &[Name::str("main")]).unwrap(); + let mut globals = vec![0u8; compiled.globals_size]; + unsafe { + crate::cancel::set_cancel_callback(globals.as_mut_ptr(), None, std::ptr::null_mut()); + let entry: extern "C" fn(*mut u8, *mut u8) -> i32 = std::mem::transmute(compiled.entry_points[&Name::str("main")]); + assert_eq!(entry(globals.as_mut_ptr(), std::ptr::null_mut()), 42, "LLVM: {}", source); + } + } + } + } + + /// Host buffers are rebound for every invocation while the DSP state stays + /// in the same globals allocation. Guards detect writes outside each slice. + fn exercise_audio_buffers( + compiler: &Compiler, + globals_size: usize, + globals_base: usize, + backend: &str, + mut process: impl FnMut(*mut u8), + ) { + let metadata = compiler.globals_info_with_offset(globals_base); + let offset = |name: &str| metadata.iter().find(|global| global.0 == name).unwrap().1; + let mut globals = vec![0u8; globals_size]; + let mut expected_state = 0.0f32; + let mut sample_count = 0usize; + for frames in [0usize, 1, 17, 239, 240, 255, 256, 257] { + let mut input = vec![-8192.0f32; frames + 2]; + for (index, sample) in input[1..frames + 1].iter_mut().enumerate() { + *sample = ((sample_count + index) % 7) as f32 * 0.125 - 0.25; + } + let original_input = input.clone(); + let mut output = vec![-16384.0f32; frames + 2]; + let expected: Vec<_> = input[1..frames + 1] + .iter() + .map(|sample| { + expected_state += sample; + expected_state * 0.5 + }) + .collect(); + unsafe { + crate::ffi::lyte_globals_bind_slice( + globals.as_mut_ptr(), + offset("input"), + input.as_ptr().add(1).cast(), + frames as i32, + ); + crate::ffi::lyte_globals_bind_slice( + globals.as_mut_ptr(), + offset("output"), + output.as_mut_ptr().add(1).cast(), + frames as i32, + ); + std::ptr::write_unaligned( + globals.as_mut_ptr().add(offset("frames")).cast::(), + frames as i32, + ); + } + process(globals.as_mut_ptr()); + assert_eq!( + &output[1..frames + 1], + expected.as_slice(), + "{} at {} frames", + backend, + frames + ); + assert_eq!(output[0], -16384.0, "{} prefix guard", backend); + assert_eq!(output[frames + 1], -16384.0, "{} suffix guard", backend); + assert_eq!(input, original_input, "{} changed input", backend); + let actual_state = unsafe { + std::ptr::read_unaligned(globals.as_ptr().add(offset("state")).cast::()) + }; + assert_eq!( + actual_state, expected_state, + "{} persistent state at {} frames", + backend, frames + ); + sample_count += frames; + } + } + + #[test] + fn checked_dsp_preserves_arbitrary_buffer_lengths_and_persistent_state() { + let mut compiler = Compiler::new(); + compiler.quiet = true; + compiler.no_recursion = true; + compiler.set_entry_points(&["process"]); + assert!(compiler.parse( + r#" + var frames: i32 + var input: [f32] + var output: [f32] + var state: f32 + assume frames >= 0 && frames <= input.len && frames <= output.len + "#, + "" + )); + assert!(compiler.parse( + r#" + step(input: f32) -> f32 { + state = state + input + state * 0.5 + } + process { + for i in 0 .. frames { output[i] = step(input[i]) } + } + "#, + "audio.lyte" + )); + assert!(compiler.check(), "{:?}", compiler.last_errors); + compiler.specialize().unwrap(); + #[cfg(any(feature = "cranelift", feature = "llvm"))] + let program = compiler.specialized_program().unwrap(); + let entry_name = Name::str("process"); + + let vm_program = compiler.compile_vm().unwrap(); + let linked = crate::vm::LinkedProgram::from_program(&vm_program); + let mut vm = crate::vm::VM::new(); + exercise_audio_buffers( + &compiler, + vm_program.globals_size, + 0, + "register VM", + |globals| unsafe { + vm.call_with_external_globals( + &linked, + &vm_program, + vm_program.entry_points[&entry_name], + globals, + vm_program.globals_size, + ); + assert!(!vm.cancelled); + }, + ); + #[cfg(has_stack_interp)] + { + let stack_program = compiler.compile_stack().unwrap(); + let mut stack = crate::stack_interp_bridge::StackBackend::new(&stack_program); + exercise_audio_buffers( + &compiler, + stack_program.globals_size, + crate::cancel::CANCEL_FLAG_RESERVED as usize, + "C Stack", + |globals| { + stack.call_entry(stack_program.entry_points[&entry_name], globals); + assert_eq!(stack.trap_reason(), crate::cancel::TRAP_NONE); + }, + ); + } + #[cfg(feature = "cranelift")] + { + let mut jit = crate::JIT::default(); + jit.no_recursion = true; + let (entries, size) = jit.compile_multi(program, &[entry_name]).unwrap(); + let entry: unsafe extern "C" fn(*mut u8, *mut u8) = + unsafe { std::mem::transmute(entries[&entry_name]) }; + exercise_audio_buffers( + &compiler, + size, + crate::cancel::CANCEL_FLAG_RESERVED as usize, + "Cranelift", + |globals| unsafe { + crate::cancel::set_cancel_callback(globals, None, std::ptr::null_mut()); + entry(globals, std::ptr::null_mut()); + }, + ); + jit.free_memory(); + } + #[cfg(feature = "llvm")] + { + let mut jit = crate::LLVMJIT::new(); + jit.no_recursion = true; + let compiled = jit.compile_only(program, &[entry_name]).unwrap(); + let entry: unsafe extern "C" fn(*mut u8, *mut u8) = + unsafe { std::mem::transmute(compiled.entry_points[&entry_name]) }; + exercise_audio_buffers( + &compiler, + compiled.globals_size, + crate::cancel::CANCEL_FLAG_RESERVED as usize, + "LLVM", + |globals| unsafe { + crate::cancel::set_cancel_callback(globals, None, std::ptr::null_mut()); + entry(globals, std::ptr::null_mut()); + }, + ); + } + } + + #[cfg(has_stack_interp)] + #[test] + fn checked_stack_loop_cancels_and_reenters_through_c_interpreter() { + unsafe extern "C" fn cancel(user_data: *mut u8) -> bool { + *user_data.cast::() += 1; + true + } + + let mut compiler = Compiler::new(); + compiler.quiet = true; + assert!(compiler.parse( + r#" + var progress: i32 + var completed: i32 + main { + for i in 0 .. 8192 { progress = progress + 1 } + completed = 1 + } + "#, + "cancel.lyte" + )); + assert!(compiler.check(), "{:?}", compiler.last_errors); + compiler.specialize().unwrap(); + let metadata = + compiler.globals_info_with_offset(crate::cancel::CANCEL_FLAG_RESERVED as usize); + let offset = |name: &str| metadata.iter().find(|global| global.0 == name).unwrap().1; + let program = compiler.compile_stack().unwrap(); + let entry = program.entry_points[&Name::str("main")]; + let mut stack = crate::stack_interp_bridge::StackBackend::new(&program); + let mut globals = vec![0u8; program.globals_size]; + let globals_ptr = globals.as_mut_ptr(); + let progress = unsafe { globals_ptr.add(offset("progress")).cast::() }; + let completed = unsafe { globals_ptr.add(offset("completed")).cast::() }; + drop(compiler); + + let mut callbacks = 0u32; + stack.set_cancel_callback(Some(cancel), (&mut callbacks as *mut u32).cast()); + stack.call_entry(entry, globals_ptr); + assert_eq!(callbacks, 1); + assert!(stack.cancelled()); + assert_eq!(stack.trap_reason(), crate::cancel::TRAP_CANCELLED); + let interrupted = unsafe { std::ptr::read_unaligned(progress) }; + assert!(interrupted > 0 && interrupted < 8192); + assert_eq!(unsafe { std::ptr::read_unaligned(completed) }, 0); + + stack.set_cancel_callback(None, std::ptr::null_mut()); + stack.call_entry(entry, globals_ptr); + assert!(!stack.cancelled()); + assert_eq!(stack.trap_reason(), crate::cancel::TRAP_NONE); + assert_eq!( + unsafe { std::ptr::read_unaligned(progress) }, + interrupted + 8192 + ); + assert_eq!(unsafe { std::ptr::read_unaligned(completed) }, 1); + assert_eq!(callbacks, 1); + } + + #[cfg(feature = "llvm")] + #[test] + fn checked_llvm_loop_cancels_and_reenters_through_host_api() { + use crate::ffi::*; + use std::ffi::{CStr, CString}; + unsafe extern "C" fn cancel(user_data: *mut u8) -> bool { + *user_data.cast::() += 1; + true + } + unsafe { + let compiler = lyte_compiler_new(std::ptr::null(), 0); + let source = CString::new("var progress: i32\nvar completed: i32\nmain { for i in 0 .. 8192 { progress = progress + 1 }; completed = 1 }").unwrap(); + let filename = CString::new("cancel.lyte").unwrap(); + assert!(lyte_compiler_add_source( + compiler, + source.as_ptr(), + filename.as_ptr() + )); + let program = lyte_compiler_compile(compiler); + assert!(!program.is_null(), "LLVM FFI compilation failed"); + lyte_compiler_free(compiler); + let size = lyte_program_get_globals_size(program); + let globals = lyte_globals_alloc(program); + assert!(!globals.is_null()); + let offset = |name: &str| { + let index = (0..lyte_program_get_globals_count(program)) + .find(|index| { + CStr::from_ptr(lyte_program_get_global_name(program, *index)) + .to_str() + .unwrap() + == name + }) + .unwrap(); + lyte_program_get_global_offset(program, index) + }; + let progress = globals.add(offset("progress")).cast::(); + let completed = globals.add(offset("completed")).cast::(); + let mut callbacks = 0u32; + lyte_program_set_cancel_callback( + program, + Some(cancel), + (&mut callbacks as *mut u32).cast(), + ); + assert!(!lyte_entry_point_call(program, 0, globals)); + assert_eq!(callbacks, 1); + assert_eq!( + crate::cancel::read_trap_reason(globals), + crate::cancel::TRAP_CANCELLED + ); + let interrupted = std::ptr::read_unaligned(progress); + assert!(interrupted > 0 && interrupted < 8192); + assert_eq!(std::ptr::read_unaligned(completed), 0); + + lyte_program_set_cancel_callback(program, None, std::ptr::null_mut()); + assert!(lyte_entry_point_call(program, 0, globals)); + assert_eq!( + crate::cancel::read_trap_reason(globals), + crate::cancel::TRAP_NONE + ); + assert_eq!(std::ptr::read_unaligned(progress), interrupted + 8192); + assert_eq!(std::ptr::read_unaligned(completed), 1); + lyte_globals_free(globals, size); + lyte_program_free(program); + } + } + + #[test] + fn concrete_instance_targets_survive_symbol_sorting() { + let declarations = vec![ + Decl::Global { + name: Name::str("z"), + typevars: vec![], + ty: mk_type(Type::Int32), + }, + Decl::Global { + name: Name::str("a"), + typevars: vec![], + ty: mk_type(Type::Bool), + }, + ]; + let records = vec![ + InstanceRecord { + definition: DefId(0), + type_args: vec![], + size_args: vec![], + declaration: 0, + }, + InstanceRecord { + definition: DefId(0), + type_args: vec![mk_type(Type::Bool)], + size_args: vec![], + declaration: 1, + }, + ]; + let program = SpecializedProgram::from_instances(declarations, records); + assert_eq!(program.instance_name(InstanceId(0)), Name::str("z")); + assert_eq!(program.instance_name(InstanceId(1)), Name::str("a")); + assert_eq!( + program + .storage_instances() + .map(|(id, _)| id) + .collect::>(), + vec![InstanceId(1), InstanceId(0)] + ); + } +} diff --git a/src/checked/validate.rs b/src/checked/validate.rs new file mode 100644 index 00000000..7f09a893 --- /dev/null +++ b/src/checked/validate.rs @@ -0,0 +1,1052 @@ +//! Structural validation at publication. This does not repeat type inference, +//! overload selection, safety proofs, or backend representation decisions. +use super::*; + +enum Phase<'a> { + Template(&'a CheckedDeclTable), + Concrete(&'a SpecializedProgram), +} + +#[derive(Clone, Copy, PartialEq, Eq)] +enum BinderKind { + Value, + Size, +} + +impl CheckedProgram { + /// Validate body-owned handles and recorded definition/requirement targets. + /// Generic types, symbolic sizes and deferred overloads are valid templates. + /// Call again after editing public program data, before consuming it. + pub fn validate(&self) -> Result<(), String> { + let phase = Phase::Template(&self.decls); + for declaration in &self.decls.decls { + phase.declaration(declaration)?; + } + Ok(()) + } +} + +impl SpecializedProgram { + /// Validate concrete bodies and the complete function/global inventory. + /// Generic struct definitions remain layout templates in this phase too. + /// Origin DefIds require the source inventory; see `validate_origins`. + pub fn validate(&self) -> Result<(), String> { + let mut owned = HashSet::new(); + let mut keys = HashSet::new(); + for (id, instance) in self.instances.iter().enumerate() { + let declaration = self + .decls + .decls + .get(instance.declaration) + .ok_or_else(|| format!("instance {} declaration is outside the program", id))?; + if !matches!(declaration, Decl::Func(_) | Decl::Global { .. }) { + return Err(format!("instance {} is neither a function nor global", id)); + } + if !owned.insert(instance.declaration) { + return Err("multiple instances own the same declaration".into()); + } + if !keys.insert(( + instance.definition, + &instance.type_args, + &instance.size_args, + )) { + return Err("duplicate concrete instance key".into()); + } + for &ty in &instance.type_args { + validate_type(ty, true)?; + } + } + let phase = Phase::Concrete(self); + for (index, declaration) in self.decls.decls.iter().enumerate() { + if matches!(declaration, Decl::Func(_) | Decl::Global { .. }) && !owned.contains(&index) + { + return Err(format!("declaration {} has no concrete instance", index)); + } + phase.declaration(declaration)?; + } + Ok(()) + } + + /// Check provenance against the owning checked definition inventory. + /// Compiler retains that inventory alongside its specialized output; + /// standalone clients must retain it to resolve or validate origins. + pub fn validate_origins(&self, templates: &CheckedProgram) -> Result<(), String> { + for instance in &self.instances { + let target = self + .decls + .decls + .get(instance.declaration) + .ok_or("instance declaration is outside the program")?; + let (types, sizes) = match (templates.decls.definition(instance.definition), target) { + (Some(Decl::Func(source)), Decl::Func(_)) => { + (source.typevars.len(), source.size_vars.len()) + } + (Some(Decl::Global { typevars, .. }), Decl::Global { .. }) => (typevars.len(), 0), + _ => { + return Err( + "instance origin is missing or has the wrong declaration kind".into(), + ) + } + }; + if instance.type_args.len() != types || instance.size_args.len() != sizes { + return Err("instance arguments do not match its definition's parameters".into()); + } + } + Ok(()) + } +} + +fn validate_type(ty: TypeID, concrete: bool) -> Result<(), String> { + match &*ty { + Type::Anon(_) => return Err(format!("unsolved type {}", ty.pretty_print())), + Type::Var(_) if concrete => return Err(format!("non-concrete type {}", ty.pretty_print())), + Type::Array(element, size) => { + if concrete && matches!(size, ArraySize::Var(_)) { + return Err(format!("non-concrete array size in {}", ty.pretty_print())); + } + validate_type(*element, concrete)?; + } + Type::Slice(element) | Type::Reference(element) => validate_type(*element, concrete)?, + Type::Tuple(elements) | Type::Name(_, elements) => { + for &element in elements { + validate_type(element, concrete)?; + } + } + Type::Func(domain, result) => { + validate_type(*domain, concrete)?; + validate_type(*result, concrete)?; + } + _ => {} + } + Ok(()) +} + +impl Phase<'_> { + fn concrete(&self) -> bool { + matches!(self, Self::Concrete(_)) + } + + fn declaration(&self, declaration: &CheckedDecl) -> Result<(), String> { + let result = match declaration { + Decl::Func(function) => self.function(function), + Decl::Assume { arena, cond } => { + self.body(arena, &[], &[], &[], &[*cond])?; + if arena.ty(*cond) != mk_type(Type::Bool) { + return Err("global assumption is not boolean".into()); + } + Ok(()) + } + Decl::Interface(interface) if !self.concrete() => { + for function in &interface.funcs { + self.function(function)?; + } + Ok(()) + } + Decl::Interface(_) | Decl::Macro(_) => Err("source-only declaration in program".into()), + Decl::Global { typevars, ty, .. } => { + if self.concrete() && !typevars.is_empty() { + return Err("generic global in concrete program".into()); + } + validate_type(*ty, self.concrete()) + } + Decl::Struct(structure) => { + for field in &structure.fields { + validate_type(field.ty, false)?; + } + Ok(()) + } + Decl::Enum { .. } | Decl::Const { .. } => Ok(()), + }; + result.map_err(|error| format!("{}: {}", declaration.name(), error)) + } + + fn function(&self, function: &CheckedFunction) -> Result<(), String> { + if self.concrete() && (!function.typevars.is_empty() || !function.size_vars.is_empty()) { + return Err("generic function in concrete program".into()); + } + validate_type(function.ret, self.concrete())?; + let roots: Vec<_> = function + .requires + .iter() + .copied() + .chain(function.body) + .collect(); + self.body( + &function.arena, + &function.params, + &function.size_vars, + &function.closure_vars, + &roots, + )?; + for &root in &function.requires { + if function.arena.ty(root) != mk_type(Type::Bool) { + return Err("precondition is not boolean".into()); + } + } + Ok(()) + } + + fn body( + &self, + body: &CheckedBody, + params: &[CheckedParam], + sizes: &[SizeParameter], + captures: &[LocalId], + roots: &[ExprID], + ) -> Result<(), String> { + validate_edges(body, roots)?; + self.requirements(body)?; + let bindings = collect_bindings(body, params, sizes, captures)?; + // Unused records, e.g. substituted size binders, can remain after + // transformation; their types must still obey the phase. + for local in &body.locals { + validate_type(local.ty, self.concrete())?; + } + for id in body.ids() { + self.node(id, body, &bindings) + .map_err(|error| format!("expression {}: {}", id, error))?; + } + Ok(()) + } + + fn node( + &self, + id: ExprID, + body: &CheckedBody, + bindings: &[Option], + ) -> Result<(), String> { + validate_type(body.ty(id), self.concrete())?; + let expression = &body[id]; + let reference = body.reference(id); + let declares = match expression { + Expr::Let(..) | Expr::Var(..) | Expr::For { .. } => 1, + Expr::Lambda { params, .. } => params.len(), + _ => 0, + }; + if body.binders(id).len() != declares { + return Err(format!( + "expression declares {} bindings but records {}", + declares, + body.binders(id).len() + )); + } + match expression { + Expr::Id(_) | Expr::TypeApp(..) => { + let Some(reference) = reference else { + return Err("reference has no recorded resolution".into()); + }; + self.reference(reference, body, bindings)?; + if let Expr::TypeApp(_, args) = expression { + for &ty in args { + validate_type(ty, self.concrete())?; + } + } + } + _ if reference.is_some() => { + return Err("non-reference expression records a resolution".into()) + } + Expr::AsTy(_, ty) => validate_type(*ty, self.concrete())?, + Expr::Let(_, _, annotation) | Expr::Var(_, _, annotation) => { + if annotation.is_some() { + return Err("checked declaration retains a source annotation".into()); + } + if body.ty(id) != mk_type(Type::Void) { + return Err("checked declaration result type is not void".into()); + } + } + Expr::Lambda { params, .. } => { + if params.iter().any(|param| param.ty.is_some()) { + return Err("checked lambda parameter retains a source annotation".into()); + } + } + Expr::Macro(..) | Expr::Error => return Err("unexpanded or invalid expression".into()), + Expr::Call(callee, args) => { + let Type::Func(domain, _) = &*body.ty(*callee) else { + return Err("call target has no function type".into()); + }; + let Type::Tuple(parameters) = &**domain else { + return Err("call target has no parameter tuple".into()); + }; + if parameters.len() != args.len() { + return Err("call arity does not match checked signature".into()); + } + if let Self::Concrete(program) = self { + if matches!(body[*callee], Expr::Id(_) | Expr::TypeApp(..)) { + if let Some(reference @ Reference::Instance(instance)) = + body.reference(*callee) + { + // The callee node can occur later in the arena. + self.reference(reference, body, bindings)?; + if let Some(target) = program.function_instance(*instance) { + if target.params.len() != args.len() { + return Err( + "call arity does not match function instance".into() + ); + } + } + } + } + } + } + _ => {} + } + Ok(()) + } + + fn requirements(&self, body: &CheckedBody) -> Result<(), String> { + let Self::Template(decls) = self else { + return if body.requirements.is_empty() { + Ok(()) + } else { + Err("unresolved interface requirements in concrete body".into()) + }; + }; + for (index, requirement) in body.requirements.iter().enumerate() { + if requirement.id.index() != index { + return Err("requirement identity does not match its body coordinate".into()); + } + let Some(Decl::Interface(interface)) = decls.definition(requirement.interface) else { + return Err("requirement interface is outside the definition inventory".into()); + }; + if requirement.type_args.len() != interface.typevars.len() { + return Err("requirement arguments do not match its interface".into()); + } + let members = decls.interface_members(requirement.interface); + if !requirement + .members + .iter() + .map(|member| member.definition) + .eq(members.iter().copied()) + { + return Err("requirement members do not belong to its interface".into()); + } + for &ty in &requirement.type_args { + validate_type(ty, false)?; + } + for member in &requirement.members { + validate_type(member.signature, false)?; + // An empty candidate list can be a deferred template obligation. + validate_candidates(&member.candidates, decls)?; + } + } + Ok(()) + } + + fn reference( + &self, + reference: &Reference, + body: &CheckedBody, + bindings: &[Option], + ) -> Result<(), String> { + match (self, reference) { + (_, Reference::Local(local)) => { + if bindings.get(local.index()) != Some(&Some(BinderKind::Value)) { + return Err("local reference has no value binder in its body".into()); + } + } + (Self::Template(_), Reference::SizeParameter(local)) => { + if bindings.get(local.index()) != Some(&Some(BinderKind::Size)) { + return Err("size reference has no size binder in its body".into()); + } + } + (Self::Template(decls), Reference::Global(id)) => { + if !matches!(decls.definition(*id), Some(Decl::Global { .. })) { + return Err("global reference has no global definition".into()); + } + } + (Self::Template(decls), Reference::Functions(candidates)) => { + if candidates.is_empty() { + return Err("empty checked overload set".into()); + } + validate_candidates(candidates, decls)?; + } + ( + Self::Template(_), + Reference::InterfaceMember { + requirement, + member, + }, + ) => { + if !body + .requirements + .get(requirement.index()) + .map_or(false, |requirement| { + requirement + .members + .iter() + .any(|candidate| candidate.definition == *member) + }) + { + return Err("interface member reference has no owning requirement".into()); + } + } + (Self::Concrete(program), Reference::Instance(id)) => { + if id.index() >= program.instances.len() { + return Err("instance reference is outside its program".into()); + } + } + (Self::Template(_), Reference::Instance(_)) => { + return Err("instance reference in checked template".into()) + } + (Self::Concrete(_), _) => { + return Err("Unresolved checked reference in concrete body".into()) + } + } + Ok(()) + } +} + +/// A referenced local must have one binder in this body, including formal +/// parameters and captures. Collect all binders before validating references. +fn collect_bindings( + body: &CheckedBody, + params: &[CheckedParam], + sizes: &[SizeParameter], + captures: &[LocalId], +) -> Result>, String> { + let mut bindings = vec![None; body.locals.len()]; + let mut bind = |local: LocalId, kind: BinderKind| -> Result<(), String> { + let slot = bindings + .get_mut(local.index()) + .ok_or("local binder is outside its body")?; + if slot.replace(kind).is_some() { + return Err(format!("local {} has multiple binders", local.0)); + } + Ok(()) + }; + for local in params + .iter() + .map(|param| param.local) + .chain(captures.iter().copied()) + { + bind(local, BinderKind::Value)?; + } + for size in sizes { + bind(size.local, BinderKind::Size)?; + } + for id in body.ids() { + for &local in body.binders(id) { + bind(local, BinderKind::Value)?; + } + } + Ok(bindings) +} + +fn validate_candidates(candidates: &[DefId], decls: &CheckedDeclTable) -> Result<(), String> { + let mut seen = HashSet::new(); + for &candidate in candidates { + if !matches!( + decls.definition(candidate), + Some(Decl::Func(_) | Decl::Global { .. }) + ) { + return Err("candidate is not a function/global definition".into()); + } + if !seen.insert(candidate) { + return Err("duplicate overload candidate".into()); + } + } + Ok(()) +} + +/// Validate every retained edge, including nodes outside the current roots. +/// Forward edges are valid (operator publication appends callees); cycles are +/// not. Use an explicit stack so malformed data cannot recurse indefinitely. +fn validate_edges(body: &CheckedBody, roots: &[ExprID]) -> Result<(), String> { + if roots.iter().any(|&root| root >= body.len()) { + return Err("expression root is outside its body".into()); + } + let mut uses = vec![0usize; body.len()]; + for &root in roots { + uses[root] += 1; + } + for id in body.ids() { + for child in body[id].subexprs() { + *uses + .get_mut(child) + .ok_or("expression edge is outside its body")? += 1; + } + } + let mut state = vec![0u8; body.len()]; + let mut contains_binder = vec![false; body.len()]; + for start in 0..body.len() { + if state[start] == 2 { + continue; + } + let mut pending = vec![(start, false)]; + while let Some((id, leaving)) = pending.pop() { + let visit = state + .get_mut(id) + .ok_or("expression edge is outside its body")?; + if leaving { + contains_binder[id] = !body.binders(id).is_empty() + || body[id] + .subexprs() + .iter() + .any(|&child| contains_binder[child]); + // Sharing reads is allowed. Sharing a declaration-containing + // subtree would give distinct lexical occurrences one binder, + // undoing the normalization required before checking/duplication. + if uses[id] > 1 && contains_binder[id] { + return Err("shared binding occurrence in checked body".into()); + } + *visit = 2; + continue; + } + match *visit { + 2 => continue, + 1 => return Err("cyclic checked expression graph".into()), + _ => *visit = 1, + } + pending.push((id, true)); + pending.extend(body[id].subexprs().into_iter().map(|child| (child, false))); + } + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn sample_function() -> CheckedFunction { + let mut arena = CheckedBody::new(); + let ty = mk_type(Type::Int32); + let local = arena.add_local(Name::str("x"), ty, false); + let body = arena.add_local_read(local, ty, test_loc()); + CheckedFunction { + name: Name::str("main"), + typevars: vec![], + size_vars: vec![], + params: vec![CheckedParam { local }], + body: Some(body), + ret: ty, + requires: vec![], + loc: test_loc(), + arena, + closure_vars: vec![], + is_extern: false, + } + } + + fn record(declaration: usize) -> InstanceRecord { + InstanceRecord { + definition: DefId(declaration as u32), + type_args: vec![], + size_args: vec![], + declaration, + } + } + + fn concrete(function: CheckedFunction) -> Result { + SpecializedProgram::try_from_instances(vec![Decl::Func(function)], vec![record(0)]) + } + + fn checked(source: &str) -> CheckedProgram { + let mut compiler = Compiler::new(); + compiler.quiet = true; + assert!( + compiler.parse(source, "boundary.lyte"), + "{:?}", + compiler.last_errors + ); + assert!(compiler.check(), "{:?}", compiler.last_errors); + compiler.checked_program().unwrap().clone() + } + + #[test] + fn publication_rejects_invalid_body_handles_in_both_phases() { + let corruptions: &[(fn(&mut CheckedFunction), &str)] = &[ + (|f| f.body = Some(f.arena.len()), "root is outside"), + (|f| f.requires.push(f.arena.len()), "root is outside"), + ( + |f| { + f.arena.add(Expr::Return(999), f.ret, test_loc()); + }, + "edge is outside", + ), + ( + |f| { + let id = f.arena.len(); + f.arena.add(Expr::Return(id), f.ret, test_loc()); + }, + "cyclic", + ), + (|f| f.params[0].local = LocalId(999), "binder is outside"), + (|f| f.params.push(f.params[0].clone()), "multiple binders"), + ( + |f| { + f.arena.add_id( + Name::str("x"), + Reference::Local(LocalId(999)), + f.ret, + test_loc(), + ); + }, + "no value binder", + ), + ( + |f| { + let local = f.arena.add_local(Name::str("orphan"), f.ret, false); + f.arena.add_local_read(local, f.ret, test_loc()); + }, + "no value binder", + ), + ( + |f| { + f.arena.add_binding( + Expr::Lambda { + params: vec![Param { + name: Name::str("p"), + ty: None, + }], + body: 0, + }, + vec![LocalId(999)], + f.ty(), + test_loc(), + ); + }, + "binder is outside", + ), + ( + |f| { + f.arena.add(Expr::Error, f.ret, test_loc()); + }, + "invalid expression", + ), + ( + |f| { + f.arena.add( + Expr::Macro(Name::str("unexpanded"), vec![]), + f.ret, + test_loc(), + ); + }, + "unexpanded", + ), + ( + |f| f.arena.locals[0].ty = mk_type(Type::Anon(100)), + "unsolved type", + ), + ]; + for &(corrupt, expected) in corruptions { + let mut function = sample_function(); + corrupt(&mut function); + let error = + CheckedProgram::try_new(CheckedDeclTable::new(vec![Decl::Func(function.clone())])) + .unwrap_err(); + assert!(error.contains(expected), "{}: {}", expected, error); + let error = concrete(function).unwrap_err(); + assert!(error.contains(expected), "{}: {}", expected, error); + } + } + + #[test] + fn declaration_nodes_require_consumed_annotations_and_void_results_in_both_phases() { + for declaration in ["let stored = x", "var stored = x", "var stored: i32"] { + let templates = checked(&format!("main(x: i32) -> i32 {{ {}; x }}", declaration)); + let definition = templates.decls.named_ids(Name::str("main"))[0]; + let function = templates.function(definition).unwrap(); + let id = function + .arena + .ids() + .find(|&id| matches!(function.arena[id], Expr::Let(..) | Expr::Var(..))) + .unwrap(); + concrete(function.clone()).unwrap(); + for (annotation, result, expected) in [ + ( + Some(function.ret), + mk_type(Type::Void), + "retains a source annotation", + ), + (None, function.ret, "result type is not void"), + ] { + let mut invalid = function.clone(); + let mut kind = invalid.arena[id].clone(); + match &mut kind { + Expr::Let(_, _, ty) | Expr::Var(_, _, ty) => *ty = annotation, + _ => unreachable!(), + } + invalid.arena.replace(id, kind, result); + let errors = [ + CheckedProgram::try_new(CheckedDeclTable::new(vec![Decl::Func( + invalid.clone(), + )])) + .unwrap_err(), + concrete(invalid).unwrap_err(), + ]; + for error in errors { + assert!(error.contains(expected), "{}: {}", declaration, error); + } + } + } + } + + #[test] + fn direct_calls_validate_selected_instance_arity_even_with_stale_node_types() { + for type_application in [false, true] { + for (target_id, target_arity, expected_error) in [ + (InstanceId(1), 1, None), + ( + InstanceId(1), + 2, + Some("call arity does not match function instance"), + ), + ( + InstanceId(2), + 1, + Some("instance reference is outside its program"), + ), + ] { + let mut caller = sample_function(); + // Equal diagnostic names cannot substitute for the selected ID. + let mut target = caller.clone(); + if target_arity == 2 { + let local = target + .arena + .add_local(Name::str("second"), target.ret, false); + target.params.push(CheckedParam { local }); + } + let recorded_type = caller.ty(); // Still says one parameter. + let argument = caller.body.unwrap(); + let callee = caller.arena.len() + 1; + caller.body = Some(caller.arena.add( + Expr::Call(callee, vec![argument]), + caller.ret, + test_loc(), + )); + let kind = if type_application { + Expr::TypeApp(Name::str("target"), vec![]) + } else { + Expr::Id(Name::str("target")) + }; + // Validate the forward reference before looking up its target. + assert_eq!(caller.arena.add(kind, recorded_type, test_loc()), callee); + caller + .arena + .set_reference(callee, Reference::Instance(target_id)); + let result = SpecializedProgram::try_from_instances( + vec![Decl::Func(caller), Decl::Func(target)], + vec![record(0), record(1)], + ); + if let Some(expected) = expected_error { + let error = result.unwrap_err(); + assert!(error.contains(expected), "{}", error); + } else { + result.unwrap(); + } + } + } + } + + #[test] + fn instance_arity_validation_preserves_indirect_calls_and_deferred_overloads() { + let templates = checked("var callback: i32 -> i32 + choose() -> i32 { 0 } + choose(value: i32) -> i32 { value } + main { callback = choose; let local = |value: i32| { value }; callback(1); local(1); choose() }"); + let main = templates + .function(templates.decls.named_ids(Name::str("main"))[0]) + .unwrap(); + assert!(main.arena.ids().any(|id| { + matches!(main.arena.reference(id), Some(Reference::Functions(candidates)) if candidates.len() == 2) + })); + MonomorphPass::new() + .monomorphize(&templates, Name::str("main")) + .unwrap() + .validate() + .unwrap(); + } + + #[test] + fn shared_reads_are_valid_but_shared_binding_occurrences_need_duplication() { + let mut function = sample_function(); + let read = function.body.unwrap(); + let root = function + .arena + .add(Expr::Block(vec![read, read]), function.ret, test_loc()); + function.body = Some(root); + concrete(function.clone()).unwrap(); + let local = function + .arena + .add_local(Name::str("y"), function.ret, false); + let binding = function.arena.add_let(local, read, test_loc()); + let block = function + .arena + .add(Expr::Block(vec![binding]), mk_type(Type::Void), test_loc()); + function + .arena + .replace(root, Expr::Block(vec![block, block, read]), function.ret); + assert!(concrete(function.clone()) + .unwrap_err() + .contains("shared binding occurrence")); + let copy = function.arena.duplicate(block); + function + .arena + .replace(root, Expr::Block(vec![block, copy, read]), function.ret); + concrete(function).unwrap(); + } + + #[test] + #[should_panic(expected = "invalid member inventory")] + fn declaration_table_rejects_missing_interface_member_handles() { + CheckedDeclTable::from_records(vec![DeclRecord { + definition: DefId(0), + declaration: Decl::Interface(Interface { + name: Name::str("Value"), + typevars: vec![], + funcs: vec![sample_function()], + loc: test_loc(), + }), + members: vec![], + }]); + } + + #[test] + fn templates_validate_reference_domains_without_selecting_overloads() { + let references = [ + Reference::Instance(InstanceId(0)), + Reference::Global(DefId(0)), // This definition is a function. + Reference::Functions(vec![]), + Reference::Functions(vec![DefId(999)]), + Reference::Functions(vec![DefId(0), DefId(0)]), + Reference::SizeParameter(LocalId(0)), // This binder is a value parameter. + Reference::InterfaceMember { + requirement: RequirementId(0), + member: DefId(0), + }, + ]; + for reference in references { + let mut function = sample_function(); + function + .arena + .add_id(Name::str("x"), reference, function.ty(), test_loc()); + assert!( + CheckedProgram::try_new(CheckedDeclTable::new(vec![Decl::Func(function)])).is_err() + ); + } + let mut function = sample_function(); + function.arena.add_id( + Name::str("x"), + Reference::Functions(vec![DefId(0)]), + function.ty(), + test_loc(), + ); + let template = + CheckedProgram::new(CheckedDeclTable::new(vec![Decl::Func(function.clone())])); + assert!(template.validate().is_ok()); + // Even an orphan node must satisfy the concrete phase contract. + assert!(concrete(function) + .unwrap_err() + .contains("Unresolved checked reference")); + } + + #[test] + fn template_generics_sizes_and_unfulfilled_requirements_remain_valid() { + let templates = checked("struct Box { value: T } interface Missing { missing(x: T) -> T } deferred(x: T) -> T where Missing { let copy = x; var value: T; value = copy; missing(value) } sized(xs: [i32; N]) -> i32 { N } main {} "); + let deferred = templates.decls.named_ids(Name::str("deferred"))[0]; + assert!( + templates.function(deferred).unwrap().arena.requirements[0].members[0] + .candidates + .is_empty() + ); + let program = MonomorphPass::new() + .monomorphize(&templates, Name::str("main")) + .unwrap(); + assert!(program.validate().is_ok()); + assert!(program + .decls + .decls + .iter() + .any(|decl| matches!(decl, Decl::Struct(structure) if !structure.typevars.is_empty()))); + assert!(program + .functions() + .all(|(_, function)| function.arena.requirements.is_empty())); + } + + #[test] + fn requirement_handles_and_candidate_kinds_are_owned_by_templates() { + let templates = checked("interface Value { value(x: T) -> T } value(x: i32) -> i32 { x } use(x: T) -> T where Value { value(x) } main {} "); + let corruptions: &[fn(&mut InterfaceRequirement)] = &[ + |r| r.id = RequirementId(1), + |r| r.interface = DefId(99999), + |r| r.members[0].definition = r.members[0].candidates[0], + |r| r.members[0].candidates = vec![r.interface], + |r| r.type_args.clear(), + ]; + for corrupt in corruptions { + let mut records: Vec<_> = templates.decls.records().collect(); + for record in &mut records { + if let Decl::Func(function) = &mut record.declaration { + if function.name == Name::str("use") { + corrupt(&mut function.arena.requirements[0]); + } + } + } + assert!(CheckedProgram::try_new(CheckedDeclTable::from_records(records)).is_err()); + } + } + + #[test] + fn concrete_types_include_locals_annotations_and_instance_arguments() { + let symbolic = mk_type(Type::Array( + mk_type(Type::Int32), + ArraySize::Var(Name::str("N")), + )); + for ty in [typevar("T"), symbolic] { + let mut function = sample_function(); + function.arena.locals[0].ty = ty; + assert!(concrete(function).unwrap_err().contains("non-concrete")); + let mut function = sample_function(); + function + .arena + .add(Expr::AsTy(0, ty), function.ret, test_loc()); + assert!(concrete(function).unwrap_err().contains("non-concrete")); + let mut instance = record(0); + instance.type_args.push(ty); + assert!(SpecializedProgram::try_from_instances( + vec![Decl::Func(sample_function())], + vec![instance] + ) + .unwrap_err() + .contains("non-concrete")); + } + } + + #[test] + fn concrete_inventory_is_complete_and_one_to_one() { + let declarations = vec![Decl::Func(sample_function())]; + for records in [vec![], vec![record(1)], vec![record(0), record(0)]] { + assert!(SpecializedProgram::try_from_instances(declarations.clone(), records).is_err()); + } + let mut other = record(1); + other.definition = DefId(0); + let error = SpecializedProgram::try_from_instances( + vec![Decl::Func(sample_function()), Decl::Func(sample_function())], + vec![record(0), other], + ) + .unwrap_err(); + assert!(error.contains("duplicate concrete instance key")); + assert!(SpecializedProgram::try_from_instances( + vec![Decl::Const { + name: Name::str("constant"), + value: 1 + }], + vec![record(0)] + ) + .is_err()); + let mut function = sample_function(); + function.arena.add_id( + Name::str("x"), + Reference::Instance(InstanceId(1)), + function.ty(), + test_loc(), + ); + assert!(concrete(function) + .unwrap_err() + .contains("instance reference is outside")); + } + + #[test] + fn instance_origins_are_validated_while_templates_are_available() { + let templates = + CheckedProgram::new(CheckedDeclTable::new(vec![Decl::Func(sample_function())])); + let mut program = concrete(sample_function()).unwrap(); + program.validate_origins(&templates).unwrap(); + program.instances[0].definition = DefId(100); + assert!(program.validate_origins(&templates).is_err()); + program.instances[0].definition = DefId(0); + program.instances[0].type_args.push(mk_type(Type::Int32)); + assert!(program.validate_origins(&templates).is_err()); + } + + #[test] + fn concrete_validation_preserves_coercing_calls_and_reference_storage_types() { + let templates = checked("read(xs: [i32]) -> i32 { 42 } update(x: &i32) { x = 1 } main { var x = 0; update(x); read([1, 2]) }"); + let program = MonomorphPass::new() + .monomorphize(&templates, Name::str("main")) + .unwrap(); + program.validate().unwrap(); + let update = program + .functions() + .find(|(_, function)| function.name == Name::str("update")) + .unwrap() + .1; + assert!(matches!( + *update.arena.local(update.params[0].local).ty, + Type::Reference(_) + )); + assert!(update.arena.ids().any(|id| { + matches!(update.arena.reference(id), Some(Reference::Local(_))) + && update.arena.ty(id) == mk_type(Type::Int32) + })); + } + + #[test] + fn unsuccessful_concretization_publishes_no_program() { + let mut compiler = Compiler::new(); + compiler.quiet = true; + assert!(compiler.parse("generic(x: T) -> T { x }", "generic-root.lyte")); + assert!(compiler.check()); + compiler.set_entry_points(&["generic"]); + assert!(compiler.specialize().unwrap_err().contains("non-concrete")); + assert!(compiler.specialized_program().is_err()); + assert!(compiler.checked_program().is_some()); + } + + #[test] + fn duplicated_loop_and_lambda_binders_keep_only_their_outer_captures() { + let templates = checked("main { let outer = 2; for i in 0 .. 2 { let apply = |value: i32| { outer + i + value }; apply(i) } }"); + let mut records: Vec<_> = templates.decls.records().collect(); + let function = records + .iter_mut() + .find_map(|record| match &mut record.declaration { + Decl::Func(function) if function.name == Name::str("main") => Some(function), + _ => None, + }) + .unwrap(); + let original = function + .arena + .ids() + .find(|&id| matches!(function.arena[id], Expr::For { .. })) + .unwrap(); + let captures = function.arena.captures(original, []); + assert_eq!(captures.len(), 1); + let previous_locals = function.arena.locals.len(); + let copy = function.arena.duplicate(original); + assert_eq!(function.arena.captures(copy, []), captures); + // The loop variable, local function value and lambda parameter all freshen. + assert_eq!(function.arena.locals.len(), previous_locals + 3); + assert_eq!(function.arena.ty(copy), function.arena.ty(original)); + assert_eq!(function.arena.loc(copy), function.arena.loc(original)); + let root = function.body.unwrap(); + let Expr::Block(mut statements) = function.arena[root].clone() else { + panic!() + }; + statements.push(copy); + function + .arena + .replace(root, Expr::Block(statements), mk_type(Type::Void)); + let duplicated = CheckedProgram::try_new(CheckedDeclTable::from_records(records)).unwrap(); + MonomorphPass::new() + .monomorphize(&duplicated, Name::str("main")) + .unwrap() + .validate() + .unwrap(); + } + + #[test] + fn deferred_size_safety_failure_blocks_publication_and_retains_templates() { + let mut compiler = Compiler::new(); + compiler.quiet = true; + assert!(compiler.parse("struct P { x: i32 } probe(xs: [i32; N]) { var p: P; for i in 0 .. 2 { let value = p.x; xs[5] } } main { probe([1, 2]) }", "safety-before-motion.lyte")); + assert!(compiler.check(), "{:?}", compiler.last_errors); + assert!(compiler + .specialize() + .unwrap_err() + .contains("safety check failed")); + assert!(!compiler.last_safety_errors.is_empty()); + assert!(compiler.specialized_program().is_err()); + assert!(compiler.checked_program().is_some()); + } +} diff --git a/src/checker.rs b/src/checker.rs index 4e1984d8..8682fd74 100644 --- a/src/checker.rs +++ b/src/checker.rs @@ -1,5 +1,26 @@ +use crate::checked::{ + CheckedBody, CheckedFunction, CheckedParam, Local, LocalId, Reference, RequirementId, +}; +use crate::free_locals::{free_locals, BindingFacts, BindingNode}; use crate::*; -use std::collections::HashSet; +use std::collections::{HashMap, HashSet}; + +#[derive(Clone, Copy, PartialEq, Eq)] +enum BindingKind { + Parameter, + LambdaParameter, + Let, + Var, + For, + SizeParameter, +} + +/// One borrowed position across the recorded candidates. A reference in any +/// candidate requires assignability; slices and references both require no-alias. +struct BorrowedParameter { + position: usize, + reference: bool, +} /// Type checking errors generated by /// both the type checker and constraint @@ -11,11 +32,12 @@ pub struct TypeError { } /// Local variable declaration. -#[derive(Copy, Clone, Debug)] +#[derive(Clone, Debug)] struct Var { name: Name, ty: TypeID, mutable: bool, + resolution: Reference, } /// Type-checks ASTs. @@ -53,10 +75,14 @@ pub struct Checker { pub types: Vec, /// Which expressions `check_expr` actually visited in the current function. - /// Unvisited entries in `types` keep the fill value from `check_fn_decl`, + /// Unvisited entries in `types` keep the fill value from `begin_body`, /// so post-check passes must skip them rather than trust their type. visited: Vec, + /// Types established without solving: checked literals and reads of bindings + /// with an explicit/fixed type. Never filled from a failed substitution. + independent_types: Vec>, + /// Is an expression's value used, rather than discarded? See /// `mark_value_positions`. value_pos: Vec, @@ -73,6 +99,14 @@ pub struct Checker { /// Return types of the enclosing functions and lambdas, innermost last. /// A `return` expression is constrained against the last entry. ret_types: Vec, + + // Construction state exists only while checking. A successful check freezes + // it into canonical checked nodes; source syntax never owns solved meaning. + references: Vec>, + locals: Vec, + local_kinds: Vec, + binders: Vec>, + requirements: Vec, } /// Returns true if the type is or contains a borrowed type (`[T]` or `&T`). @@ -200,6 +234,7 @@ impl Checker { Self { types: vec![], visited: vec![], + independent_types: vec![], value_pos: vec![], lvalue: vec![], inst: Instance::new(), @@ -214,9 +249,45 @@ impl Checker { errors: vec![], loop_depth: 0, ret_types: vec![], + references: vec![], + locals: vec![], + local_kinds: vec![], + binders: vec![], + requirements: vec![], } } + fn record_use(&mut self, expr: ExprID, resolution: Reference) { + self.references[expr] = Some(resolution); + } + + fn bind( + &mut self, + name: Name, + ty: TypeID, + mutable: bool, + kind: BindingKind, + binder: Option, + ) { + let id = LocalId(self.locals.len() as u32); + self.locals.push(Local { name, ty, mutable }); + self.local_kinds.push(kind); + if let Some(expr) = binder { + self.binders[expr].push(id); + } + let resolution = if kind == BindingKind::SizeParameter { + Reference::SizeParameter(id) + } else { + Reference::Local(id) + }; + self.vars.push(Var { + name, + ty, + mutable, + resolution, + }); + } + fn eq(&mut self, lhs: TypeID, rhs: TypeID, loc: Loc, hint: &str) { let hint = if hint.is_empty() { None @@ -292,7 +363,12 @@ impl Checker { // suffix range checks, so `-2147483648i32` is valid even // though `2147483648i32` is not. if let Expr::Int(n, Some(suffix @ IntLiteralSuffix::I32)) = &arena[arg] { + let errors_before = self.errors.len(); let ty = self.check_int_literal_suffix(arg, *n, *suffix, arena, true); + self.visited[arg] = true; + if self.errors.len() == errors_before { + self.independent_types[arg] = Some(ty); + } self.types[arg] = ty; return ty; } @@ -377,9 +453,22 @@ impl Checker { } let overload_name = Name::new(op.overload_name().into()); + let candidates: Vec<_> = decls + .named_ids(overload_name) + .into_iter() + .filter(|id| matches!(decls.definition(*id), Some(Decl::Func(_)))) + .collect(); + // Built-in arithmetic needs no named target. An empty optional + // overload inventory is not an unresolved identifier reference. + if !candidates.is_empty() { + self.record_use(id, Reference::Functions(candidates)); + } for d in decls.find(overload_name) { - if let Decl::Func(_) = d { - let dt = d.ty().fresh(&mut self.next_anon); + if let Decl::Func(function) = d { + let Some(signature) = function.annotated_ty() else { + continue; + }; + let dt = signature.fresh(&mut self.next_anon); alts.push(Alt { ty: dt, interfaces: vec![], @@ -454,6 +543,7 @@ impl Checker { fn check_expr(&mut self, id: ExprID, arena: &ExprArena, decls: &DeclTable) -> TypeID { self.visited[id] = true; + let errors_before = self.errors.len(); let ty = match &arena[id] { Expr::True | Expr::False => mk_type(Type::Bool), Expr::Int(_, None) => { @@ -485,6 +575,7 @@ impl Checker { Expr::TypeApp(name, type_args) => { // Explicit type application: name⟨i32⟩. if let Some(Decl::Global { typevars, ty, .. }) = decls.find(*name).first() { + self.record_use(id, Reference::Global(decls.named_ids(*name)[0])); self.lvalue[id] = true; let mut inst = Instance::new(); for (tv, ta) in typevars.iter().zip(type_args.iter()) { @@ -496,16 +587,28 @@ impl Checker { // whose typevar count matches, substitute the explicit type // args, and let the solver pick the right arity. let fn_decls = decls.find(*name); + let candidates = decls + .named_ids(*name) + .into_iter() + .filter(|id| { + matches!(decls.definition(*id), Some(Decl::Func(f)) + if f.typevars.len() == type_args.len()) + }) + .collect(); + self.record_use(id, Reference::Functions(candidates)); let mut alts: Vec = Vec::new(); for d in fn_decls { if let Decl::Func(fdecl) = d { - if fdecl.typevars.len() == type_args.len() { + if let Some(signature) = fdecl + .annotated_ty() + .filter(|_| fdecl.typevars.len() == type_args.len()) + { let mut inst = Instance::new(); for (tv, ta) in fdecl.typevars.iter().zip(type_args.iter()) { inst.insert(mk_type(Type::Var(*tv)), *ta); } alts.push(Alt { - ty: fdecl.ty().subst(&inst), + ty: signature.subst(&inst), interfaces: vec![], }); } @@ -526,9 +629,11 @@ impl Checker { // Local variables will override all declarations. // Is this what we want? if let Some(v) = self.find(*name) { + self.record_use(id, v.resolution); self.lvalue[id] = v.mutable; v.ty } else if let Some(Decl::Global { typevars, ty, .. }) = decls.find(*name).first() { + self.record_use(id, Reference::Global(decls.named_ids(*name)[0])); self.lvalue[id] = true; if typevars.is_empty() { *ty @@ -544,6 +649,17 @@ impl Checker { } } else { let t = self.fresh(); + let candidates = decls + .named_ids(*name) + .into_iter() + .filter(|id| { + matches!( + decls.definition(*id), + Some(Decl::Func(_) | Decl::Global { .. }) + ) + }) + .collect(); + self.record_use(id, Reference::Functions(candidates)); let alts: Vec = decls .alts(*name) .iter() @@ -553,7 +669,15 @@ impl Checker { if alts.is_empty() { self.errors.push(TypeError { location: arena.locs[id], - message: format!("undeclared identifier: {}", *name), + message: if decls + .find(*name) + .iter() + .any(|decl| matches!(decl, Decl::Func(_))) + { + format!("function '{}' has an incomplete parameter signature", name) + } else { + format!("undeclared identifier: {}", *name) + }, }); } @@ -713,11 +837,7 @@ impl Checker { ); } - self.vars.push(Var { - name: *name, - ty, - mutable: true, - }); + self.bind(*name, ty, true, BindingKind::Var, Some(id)); // A var declaration is a statement, not an expression: it does // not produce a value. Record the variable's type for the @@ -741,11 +861,7 @@ impl Checker { "variable initializer type must match", ); - self.vars.push(Var { - name: *name, - ty, - mutable: false, - }); + self.bind(*name, ty, false, BindingKind::Let, Some(id)); // A let declaration is a statement too — see the comment on // Expr::Var above. @@ -999,11 +1115,7 @@ impl Checker { "end of for loop must be an integer", ); - self.vars.push(Var { - name: *var, - ty, - mutable: false, - }); + self.bind(*var, ty, false, BindingKind::For, Some(id)); self.loop_depth += 1; self.check_expr(*body, arena, decls); @@ -1025,11 +1137,13 @@ impl Checker { let mut param_types = vec![]; for param in params { let ty = param.ty.unwrap_or_else(|| self.fresh()); - self.vars.push(Var { - name: param.name, - mutable: false, + self.bind( + param.name, ty, - }); + false, + BindingKind::LambdaParameter, + Some(id), + ); param_types.push(ty); } @@ -1055,6 +1169,27 @@ impl Checker { } Expr::Error => self.fresh(), }; + if self.errors.len() == errors_before + && (matches!( + arena[id], + Expr::True + | Expr::False + | Expr::Int(..) + | Expr::Real(..) + | Expr::Char(_) + | Expr::String(_) + ) || matches!( + (&arena[id], &self.references[id]), + ( + Expr::Id(_), + Some(Reference::Local(_) | Reference::SizeParameter(_)) + ) + )) + { + if ty != mk_type(Type::Void) { + self.independent_types[id] = Some(ty); + } + } self.types[id] = ty; ty } @@ -1062,13 +1197,69 @@ impl Checker { fn find(&self, name: Name) -> Option { for v in self.vars.iter().rev() { if v.name == name { - return Some(*v); + return Some(v.clone()); } } None } + /// Start an independent source body without manufacturing a function scope. + fn begin_body(&mut self, arena: &ExprArena) { + let n = arena.exprs.len(); + self.references = vec![None; n]; + self.locals.clear(); + self.local_kinds.clear(); + self.binders = vec![vec![]; n]; + self.requirements.clear(); + self.types = vec![mk_type(Type::Void); n]; + self.visited = vec![false; n]; + self.independent_types = vec![None; n]; + self.value_pos = vec![false; n]; + self.lvalue = vec![false; n]; + self.vars.clear(); + self.inst.clear(); + self.constraints.clear(); + self.ret_types.clear(); + self.loop_depth = 0; + } + + fn check_expected_expr( + &mut self, + root: ExprID, + arena: &ExprArena, + decls: &DeclTable, + expected: TypeID, + message: &str, + ) { + let ty = self.check_expr(root, arena, decls); + self.eq(ty, expected, arena.locs[root], message); + } + + fn check_assumption(&mut self, arena: &ExprArena, cond: ExprID, decls: &DeclTable) { + let errors_before = self.errors.len(); + self.begin_body(arena); + mark_value_positions(cond, arena, true, &mut self.value_pos); + // An assumption is a boolean-valued body. Preserve that expectation for + // explicit returns in it as well, without constructing a function record. + self.ret_types.push(mk_type(Type::Bool)); + self.check_expected_expr( + cond, + arena, + decls, + mk_type(Type::Bool), + "assume condition must be a boolean expression", + ); + self.ret_types.pop(); + self.solve_body(arena, decls, errors_before); + let solved = self.solved_types(); + self.check_void_declarations(arena, &solved); + if self.errors.len() == errors_before { + self.check_unsolved_types(arena, &solved, false); + } + } + fn check_fn_decl(&mut self, func_decl: &FuncDecl, decls: &DeclTable) { + self.begin_body(&func_decl.arena); // `self.errors` accumulates across every decl in the table, so the // post-check passes below gate on errors from *this* function rather // than on the accumulator — otherwise one bad function silently @@ -1087,108 +1278,98 @@ impl Checker { }); } - let n = func_decl.arena.exprs.len(); - self.types.resize(n, mk_type(Type::Void)); - self.visited.clear(); - self.visited.resize(n, false); - self.value_pos.clear(); - self.value_pos.resize(n, false); - self.lvalue.resize(n, false); - if let Some(body) = func_decl.body { - // println!("🟧 checking function {:?} 🟧", *func_decl.name); - - // The body's value is the return value, unless the function - // returns void — see the `self.eq(ty, func_decl.ret, ..)` below. let body_used = func_decl.ret != mk_type(Type::Void); mark_value_positions(body, &func_decl.arena, body_used, &mut self.value_pos); - for &req in &func_decl.requires { - mark_value_positions(req, &func_decl.arena, true, &mut self.value_pos); - } - - self.inst.clear(); - self.constraints.clear(); - - for param in &func_decl.params { - let ty = match param.ty { - Some(ty) => ty, - None => { - self.errors.push(TypeError { - location: func_decl.loc, - message: format!( - "parameter '{}' is missing a type annotation", - *param.name - ), - }); - return; - } - }; - let (var_ty, mutable) = match &*ty { - Type::Reference(inner) => (*inner, true), - _ => (ty, false), - }; - self.vars.push(Var { - name: param.name, - ty: var_ty, - mutable, - }); - } + } + for &req in &func_decl.requires { + mark_value_positions(req, &func_decl.arena, true, &mut self.value_pos); + } - // Size variables are available as i32 values in the body. - for &sv in &func_decl.size_vars { - self.vars.push(Var { - name: sv, - ty: mk_type(Type::Int32), - mutable: false, - }); - } + for param in &func_decl.params { + let ty = match param.ty { + Some(ty) => ty, + None => { + self.errors.push(TypeError { + location: func_decl.loc, + message: format!( + "parameter '{}' is missing a type annotation", + *param.name + ), + }); + return; + } + }; + let (var_ty, mutable) = match &*ty { + Type::Reference(inner) => (*inner, true), + _ => (ty, false), + }; + self.bind(param.name, var_ty, mutable, BindingKind::Parameter, None); + } - // Add interface functions to available functions. - for constraint in &func_decl.constraints { - if let Some(Decl::Interface(interface)) = - decls.find(constraint.interface_name).first() - { - let mut inst = Instance::new(); + for (local, param) in self.locals.iter_mut().zip(&func_decl.params) { + local.ty = param.ty.expect("checked parameter annotation"); + } - interface - .typevars - .iter() - .zip(&constraint.typevars) - .for_each(|pair| { - let t0 = mk_type(Type::Var(*pair.0)); - let t1 = mk_type(Type::Var(*pair.1)); - if t0 != t1 { - inst.insert(t0, t1); - } - }); + // Size variables are available as i32 values in the body. + for &sv in &func_decl.size_vars { + self.bind( + sv, + mk_type(Type::Int32), + false, + BindingKind::SizeParameter, + None, + ); + } - for func in &interface.funcs { - self.vars.push(Var { - name: func.name, - ty: func.ty().subst(&inst), - mutable: false, - }); - } - } else { - self.errors.push(TypeError { - location: func_decl.loc, - message: format!("unknown interface: {}", constraint.interface_name), - }) + // Interface members retain their requirement and declaration IDs; + // specialization selects only from these checked candidates. + for (ordinal, constraint) in func_decl.constraints.iter().enumerate() { + let requirement_id = RequirementId(ordinal as u32); + let args = constraint + .typevars + .iter() + .map(|name| typevar(name)) + .collect(); + if let Some(requirement) = + decls.interface_requirement(requirement_id, constraint.interface_name, args) + { + for member in &requirement.members { + let function = decls + .function(member.definition) + .expect("checked interface member"); + self.vars.push(Var { + name: function.name, + ty: member.signature, + mutable: false, + resolution: Reference::InterfaceMember { + requirement: requirement_id, + member: member.definition, + }, + }); } + self.requirements.push(requirement); + } else { + self.errors.push(TypeError { + location: func_decl.loc, + message: format!("unknown interface: {}", constraint.interface_name), + }); } + } - // Type-check require clauses: each must have type bool. - for &req in &func_decl.requires { - let req_ty = self.check_expr(req, &func_decl.arena, decls); - self.eq( - req_ty, - mk_type(Type::Bool), - func_decl.arena.locs[req], - "require clause must be a boolean expression", - ); - } + // Type-check require clauses: each must have type bool. + for &req in &func_decl.requires { + self.check_expected_expr( + req, + &func_decl.arena, + decls, + mk_type(Type::Bool), + "require clause must be a boolean expression", + ); + } - // Check the body of the function. + // Prototypes still own checked parameters and preconditions. + if let Some(body) = func_decl.body { self.ret_types.push(func_decl.ret); let ty = self.check_expr(body, &func_decl.arena, decls); self.ret_types.pop(); @@ -1201,48 +1382,55 @@ impl Checker { "return type must match function return type", ); } + } + self.solve_body(&func_decl.arena, decls, errors_before); + } - self.vars.clear(); + fn solve_body(&mut self, arena: &ExprArena, decls: &DeclTable, errors_before: usize) { + self.vars.clear(); - if self.errors.len() == errors_before { - solve_constraints( - &mut self.constraints, - &mut self.inst, - decls, - &mut self.errors, - ); - } + if self.errors.len() == errors_before { + solve_constraints( + &mut self.constraints, + &mut self.inst, + decls, + &mut self.errors, + ); + } - // Check lvalue validity now that types are solved. - if self.errors.len() == errors_before { - self.check_lvalues(func_decl, decls); - } + // Check lvalue validity now that types are solved. + if self.errors.len() == errors_before { + self.check_lvalues(arena, decls); + } - // Check that no two borrowed parameters alias (Fortran-style no-alias rule). - if self.errors.len() == errors_before { - self.check_slice_aliasing(func_decl, decls); - } + // Check that no two borrowed parameters alias (Fortran-style no-alias rule). + if self.errors.len() == errors_before { + self.check_slice_aliasing(arena, decls); } } /// Check that all assignments have valid lvalue targets, using solved types /// to determine slice mutability. - fn check_lvalues(&mut self, func_decl: &FuncDecl, decls: &DeclTable) { - for id in 0..func_decl.arena.exprs.len() { - if let Expr::Binop(Binop::Assign, lhs, _) = &func_decl.arena[id] { - if !self.is_lvalue(*lhs, func_decl) { + fn check_lvalues(&mut self, arena: &ExprArena, decls: &DeclTable) { + for id in 0..arena.exprs.len() { + if let Expr::Binop(Binop::Assign, lhs, _) = &arena[id] { + if !self.is_lvalue(*lhs, arena) { self.errors.push(TypeError { - location: func_decl.arena.locs[id], + location: arena.locs[id], message: "left-hand side of assignment isn't assignable".to_string(), }); } } - if let Expr::Call(f, args) = &func_decl.arena[id] { - for pos in self.reference_arg_positions(*f, decls, func_decl) { - if pos < args.len() && !self.is_lvalue(args[pos], func_decl) { + if let Expr::Call(f, args) = &arena[id] { + for parameter in self.borrowed_parameters(*f, arena, decls) { + if !parameter.reference { + continue; + } + let pos = parameter.position; + if pos < args.len() && !self.is_lvalue(args[pos], arena) { self.errors.push(TypeError { - location: func_decl.arena.locs[args[pos]], + location: arena.locs[args[pos]], message: "reference argument must be assignable".to_string(), }); } @@ -1251,55 +1439,63 @@ impl Checker { } } - fn reference_arg_positions( + /// Classify borrowing once for both post-solve consumers. Recorded overloads + /// (including explicit TypeApp candidates) retain the conservative union; + /// local and other function-valued callees use their solved signature. + fn borrowed_parameters( &self, callee: ExprID, + arena: &ExprArena, decls: &DeclTable, - func_decl: &FuncDecl, - ) -> Vec { - let mut positions = Vec::new(); - let mut found_decl = false; - - if let Expr::Id(callee_name) = &func_decl.arena[callee] { - for d in decls.find(*callee_name) { - if let Decl::Func(fd) = d { - found_decl = true; - for (i, param) in fd.params.iter().enumerate() { - if let Some(ty) = param.ty { - if matches!(*ty, Type::Reference(_)) && !positions.contains(&i) { - positions.push(i); - } + ) -> Vec { + let mut borrowed: Vec = Vec::new(); + let mut add = |position, ty: TypeID| { + if matches!(*ty, Type::Slice(_) | Type::Reference(_)) { + let reference = matches!(*ty, Type::Reference(_)); + if let Some(parameter) = borrowed.iter_mut().find(|p| p.position == position) { + parameter.reference |= reference; + } else { + borrowed.push(BorrowedParameter { + position, + reference, + }); + } + } + }; + let mut found_function = false; + if let Some(Reference::Functions(candidates)) = &self.references[callee] { + for &candidate in candidates { + if let Some(function) = decls.function(candidate) { + found_function = true; + let explicit: Instance = match &arena[callee] { + Expr::TypeApp(_, args) => function + .typevars + .iter() + .zip(args) + .map(|(&var, &ty)| (mk_type(Type::Var(var)), ty)) + .collect(), + _ => Instance::new(), + }; + for (position, parameter) in function.params.iter().enumerate() { + if let Some(ty) = parameter.ty { + add(position, ty.subst(&explicit)); } } } } } - - if found_decl { - return positions; - } - - // Indirect call through a fat pointer. There is no declaration to - // consult, so take the parameter types from the callee expression's - // solved type — the same source the backends use to decide which - // arguments to pass by address. - self.callee_param_types(callee) - .iter() - .enumerate() - .filter(|(_, ty)| matches!(***ty, Type::Reference(_))) - .map(|(i, _)| i) - .collect() - } - - /// Parameter types of a callee expression, from its solved type. Used for - /// indirect calls, where there is no declaration to consult. - fn callee_param_types(&self, callee: ExprID) -> Vec { - if let Type::Func(from, _) = &*self.types[callee].subst(&self.inst) { - if let Type::Tuple(params) = &**from { - return params.clone(); + // A recorded set can contain only function-valued globals. Without an + // actual function declaration, borrowing comes from the solved signature. + if !found_function { + if let Type::Func(from, _) = &*self.types[callee].subst(&self.inst) { + if let Type::Tuple(params) = &**from { + for (position, &ty) in params.iter().enumerate() { + add(position, ty); + } + } } } - Vec::new() + borrowed } /// Enforce the no-alias rule: two borrowed parameters in the same call must not @@ -1309,59 +1505,24 @@ impl Checker { /// This enables Fortran-style `restrict` semantics: the compiler (and LLVM /// backend) can assume slice parameters never overlap, enabling load hoisting, /// store reordering, and vectorization. - fn check_slice_aliasing(&mut self, func_decl: &FuncDecl, decls: &DeclTable) { - for id in 0..func_decl.arena.exprs.len() { - let Expr::Call(f, args) = &func_decl.arena[id] else { + fn check_slice_aliasing(&mut self, arena: &ExprArena, decls: &DeclTable) { + for id in 0..arena.exprs.len() { + let Expr::Call(f, args) = &arena[id] else { continue; }; - // Find which parameter positions are borrowed in the callee's declaration. - // For overloaded functions, union all borrowed positions (conservative). - let mut borrow_positions: Vec<(usize, bool)> = Vec::new(); - let mut found_decl = false; - if let Expr::Id(callee_name) = &func_decl.arena[*f] { - for d in decls.find(*callee_name) { - if let Decl::Func(fd) = d { - found_decl = true; - for (i, param) in fd.params.iter().enumerate() { - if let Some(ty) = param.ty { - let is_slice = matches!(*ty, Type::Slice(_)); - if matches!(*ty, Type::Slice(_) | Type::Reference(_)) - && !borrow_positions.iter().any(|(pos, _)| *pos == i) - { - borrow_positions.push((i, is_slice)); - } - } - } - } - } - } - - // Indirect call through a fat pointer: fall back to the callee - // expression's solved type. Slices can't appear in function types, - // but references can, so this catches aliased borrowed arguments - // that would otherwise slip through the closure call path. - if !found_decl { - for (i, ty) in self.callee_param_types(*f).iter().enumerate() { - let is_slice = matches!(**ty, Type::Slice(_)); - if matches!(**ty, Type::Slice(_) | Type::Reference(_)) - && !borrow_positions.iter().any(|(pos, _)| *pos == i) - { - borrow_positions.push((i, is_slice)); - } - } - } - - if borrow_positions.len() < 2 { + let borrowed = self.borrowed_parameters(*f, arena, decls); + if borrowed.len() < 2 { continue; } - // Collect (position, base_path) for arguments at borrowed positions. - let mut borrow_args: Vec<(usize, Vec)> = Vec::new(); - for &(pos, _) in &borrow_positions { + // Collect base paths for arguments at borrowed positions. + let mut borrow_args: Vec> = Vec::new(); + for parameter in &borrowed { + let pos = parameter.position; if pos < args.len() { - if let Some(path) = expr_base_path(args[pos], &func_decl.arena) { - borrow_args.push((pos, path)); + if let Some(path) = expr_base_path(args[pos], arena) { + borrow_args.push(path); } } } @@ -1369,16 +1530,15 @@ impl Checker { // Check all pairs for aliasing. for i in 0..borrow_args.len() { for j in (i + 1)..borrow_args.len() { - if borrow_args[i].1 == borrow_args[j].1 { + if borrow_args[i] == borrow_args[j] { let name = borrow_args[i] - .1 .iter() .map(|n| n.as_str()) .collect::>() .join("."); self.errors.push(TypeError { - location: func_decl.arena.locs[id], - message: if borrow_positions.iter().all(|(_, is_slice)| *is_slice) { + location: arena.locs[id], + message: if borrowed.iter().all(|p| !p.reference) { format!( "cannot pass '{}' to multiple slice parameters \ (slices must not alias)", @@ -1399,10 +1559,10 @@ impl Checker { } /// Determine if an expression is an lvalue, using solved types. - fn is_lvalue(&self, id: ExprID, func_decl: &FuncDecl) -> bool { - match &func_decl.arena[id] { + fn is_lvalue(&self, id: ExprID, arena: &ExprArena) -> bool { + match &arena[id] { Expr::Id(_) | Expr::TypeApp(_, _) => self.lvalue[id], - Expr::Field(lhs, _) => self.is_lvalue(*lhs, func_decl), + Expr::Field(lhs, _) => self.is_lvalue(*lhs, arena), Expr::ArrayIndex(array_expr, _) => { // Slice indexing is always an lvalue (slices are mutable references). let solved_ty = self.types[*array_expr].subst(&self.inst); @@ -1410,7 +1570,7 @@ impl Checker { return true; } // Array indexing is an lvalue only if the array itself is. - self.is_lvalue(*array_expr, func_decl) + self.is_lvalue(*array_expr, arena) } _ => false, } @@ -1482,6 +1642,7 @@ impl Checker { match decl { Decl::Func(func_decl) => self.check_fn_decl(func_decl, decls), Decl::Macro(func_decl) => self.check_fn_decl(func_decl, decls), + Decl::Assume { arena, cond } => self.check_assumption(arena, *cond, decls), Decl::Interface(Interface { name, funcs, .. }) => self.check_interface(*name, funcs), Decl::Struct(struct_decl) => self.check_struct_decl(struct_decl, decls), _ => (), @@ -1498,13 +1659,21 @@ impl Checker { let errors_before = self.errors.len(); self._check_decl(decl, decls); if let Decl::Func(fd) | Decl::Macro(fd) = decl { - check_escape_in_func(fd, &mut self.errors); + if let Some(body) = fd.body { + let facts = CheckingBindings { + arena: &fd.arena, + references: &self.references, + binders: &self.binders, + visited: &self.visited, + }; + escape_walk(body, &facts, &mut HashMap::new(), &mut self.errors); + } // Both post-solve passes read the same substituted types, so // compute them once. let solved_types = self.solved_types(); - self.check_void_declarations(fd, &solved_types); + self.check_void_declarations(&fd.arena, &solved_types); if self.errors.len() == errors_before { - self.check_unsolved_types(fd, &solved_types); + self.check_unsolved_types(&fd.arena, &solved_types, !fd.typevars.is_empty()); } } } @@ -1525,8 +1694,8 @@ impl Checker { /// A function that reported its own error still won't produce this /// diagnostic, since `check_fn_decl` skips constraint solving there and /// the declaration's type is left an unsolved type variable. - fn check_void_declarations(&mut self, func_decl: &FuncDecl, solved_types: &[TypeID]) { - for (i, expr) in func_decl.arena.exprs.iter().enumerate() { + fn check_void_declarations(&mut self, arena: &ExprArena, solved_types: &[TypeID]) { + for (i, expr) in arena.exprs.iter().enumerate() { let name = match expr { Expr::Let(name, _, _) | Expr::Var(name, _, _) => name, _ => continue, @@ -1536,7 +1705,7 @@ impl Checker { } if i < solved_types.len() && matches!(&*solved_types[i], Type::Void) { self.errors.push(TypeError { - location: func_decl.arena.locs[i], + location: arena.locs[i], message: format!("variable '{}' cannot have type void", name), }); } @@ -1548,8 +1717,12 @@ impl Checker { /// unsolved vars are usually a symptom of an earlier type error, and a /// function whose constraints never got solved has nothing but unsolved /// vars. - fn check_unsolved_types(&mut self, func_decl: &FuncDecl, solved_types: &[TypeID]) { - let is_generic = !func_decl.typevars.is_empty(); + fn check_unsolved_types( + &mut self, + arena: &ExprArena, + solved_types: &[TypeID], + is_generic: bool, + ) { for (i, solved) in solved_types.iter().enumerate() { // For generic functions, Type::Var is expected (they have named // type parameters). Only flag anonymous type variables (Anon), @@ -1560,9 +1733,9 @@ impl Checker { solved.contains_var() }; if has_problem { - if i < func_decl.arena.locs.len() { + if i < arena.locs.len() { self.errors.push(TypeError { - location: func_decl.arena.locs[i], + location: arena.locs[i], message: format!( "could not fully infer type (resolved to {})", solved.pretty_print() @@ -1574,6 +1747,225 @@ impl Checker { } } + /// Retain a bounded view of this check, without constructing checked syntax + /// for erroneous expressions. A concrete substitution left by failed solving + /// is not evidence: only a successful body may publish inferred types. + pub(crate) fn body_analysis( + &self, + source: &FuncDecl, + analysis: &SourceAnalysis, + ) -> BodyAnalysis { + let source_trusted = analysis.source_is_trusted(source); + let solved = self.solved_types(); + let complete = source_trusted + && self.errors.is_empty() + && self + .references + .iter() + .flatten() + .all(|r| analysis.reference_is_trusted(r)) + && self + .types + .iter() + .chain(&solved) + .all(|ty| analysis.type_is_valid(*ty, source)); + let available = + |ty: TypeID| (source_trusted && analysis.type_is_available(ty, source)).then_some(ty); + let mut locals: Vec<_> = self + .locals + .iter() + .map(|local| AnalyzedLocal { + name: local.name, + ty: available(if complete { + local.ty.subst(&self.inst) + } else { + local.ty + }) + .filter(|ty| *ty != mk_type(Type::Void)), + loc: source.loc, + }) + .collect(); + for (expr, binders) in self.binders.iter().enumerate() { + for binder in binders { + locals[binder.index()].loc = source.arena.locs[expr]; + } + } + let expressions = source + .arena + .exprs + .iter() + .enumerate() + .map(|(id, expr)| { + if !self.visited[id] { + return ExpressionFacts::default(); + } + let reference = self.references[id] + .clone() + .filter(|reference| match reference { + Reference::Functions(candidates) => !candidates.is_empty(), + Reference::Local(local) | Reference::SizeParameter(local) => { + !locals[local.index()].name.is_empty() + } + _ => true, + }); + let ty = if complete { + available(if matches!(expr, Expr::Let(..) | Expr::Var(..)) { + mk_type(Type::Void) + } else { + solved[id] + }) + } else { + self.independent_types[id].and_then(available) + }; + ExpressionFacts { + ty, + reference, + binding: if matches!(expr, Expr::Let(..) | Expr::Var(..) | Expr::For { .. }) { + self.binders[id].first().copied() + } else { + None + }, + } + }) + .collect(); + BodyAnalysis { + expressions, + locals, + requirements: self + .requirements + .iter() + .map(|req| (req.id, req.interface)) + .collect(), + } + } + + /// Freeze a successfully checked source body. Expression references and + /// solved types become one node; declaration types belong to local records. + pub(crate) fn checked_body(&self, source: &ExprArena) -> CheckedBody { + assert!( + self.errors.is_empty(), + "cannot publish an invalid checked body" + ); + let locals: Vec = self + .locals + .iter() + .cloned() + .map(|mut local| { + local.ty = local.ty.subst(&self.inst); + local + }) + .collect(); + let solved = self.solved_types(); + let mut syntax = source.clone(); + let mut types = Vec::with_capacity(syntax.exprs.len()); + let mut references = Vec::with_capacity(syntax.exprs.len()); + for (id, expression) in syntax.exprs.iter_mut().enumerate() { + let mut reference = None; + match expression { + Expr::Id(_) => { + reference = Some(self.references[id].clone().expect("checked reference")); + } + Expr::TypeApp(_, args) => { + reference = Some(self.references[id].clone().expect("checked reference")); + for ty in args { + *ty = ty.subst(&self.inst); + } + } + // Declarations consume their annotations: local records own the types. + Expr::Let(_, _, annotation) | Expr::Var(_, _, annotation) => *annotation = None, + Expr::Lambda { params, .. } => { + for param in params { + param.ty = None; + } + } + Expr::AsTy(_, ty) => *ty = ty.subst(&self.inst), + Expr::Macro(..) | Expr::Error => { + panic!("unexpanded or invalid source in checked body") + } + _ => {} + } + references.push(reference); + types.push(if matches!(expression, Expr::Let(..) | Expr::Var(..)) { + mk_type(Type::Void) + } else { + solved[id] + }); + } + let mut body = CheckedBody::from_parts( + syntax, + types, + references, + self.binders.clone(), + locals, + self.requirements.clone(), + ); + // Operator overloading is lowered while publishing checked meaning, + // rather than rewriting source syntax after its types have been solved. + for id in body.ids() { + if let Expr::Binop(op, lhs, rhs) = body[id].clone() { + if op.arithmetic() + && (matches!(*body.ty(lhs), Type::Name(_, _)) + || (op == Binop::Mod + && matches!(*body.ty(lhs), Type::Float32 | Type::Float64))) + { + let Some(reference) = self.references[id].clone() else { + continue; + }; + let ty = body.ty(id); + let callee = body.add_id( + Name::new(op.overload_name().into()), + reference, + func(tuple(vec![body.ty(lhs), body.ty(rhs)]), ty), + body.loc(id), + ); + body.replace(id, Expr::Call(callee, vec![lhs, rhs]), ty); + } + } + } + body + } + + /// Publish function metadata only at the function boundary. + pub fn checked_function(&self, source: &FuncDecl) -> CheckedFunction { + let arena = self.checked_body(&source.arena); + let params = self + .local_kinds + .iter() + .enumerate() + .filter_map(|(i, kind)| { + (*kind == BindingKind::Parameter).then(|| CheckedParam { + local: LocalId(i as u32), + }) + }) + .collect(); + CheckedFunction { + name: source.name, + typevars: source.typevars.clone(), + size_vars: source + .size_vars + .iter() + .copied() + .zip( + self.local_kinds + .iter() + .enumerate() + .filter_map(|(index, kind)| { + (*kind == BindingKind::SizeParameter).then_some(LocalId(index as u32)) + }), + ) + .map(|(symbol, local)| crate::checked::SizeParameter { symbol, local }) + .collect(), + params, + body: source.body, + ret: source.ret.subst(&self.inst), + requires: source.requires.clone(), + loc: source.loc, + arena, + closure_vars: vec![], + is_extern: source.is_extern, + } + } + /// Returns types we've solved for. pub fn solved_types(&self) -> Vec { self.types.iter().map(|t| t.subst(&self.inst)).collect() @@ -1655,240 +2047,142 @@ fn expr_base_path(id: ExprID, arena: &ExprArena) -> Option> { } } -/// Verify that no capturing lambda (one with free variables from the enclosing -/// function) escapes the function, either directly returned or via a variable. -/// Such a closure would capture stack addresses that dangle after the frame exits. -fn check_escape_in_func(func_decl: &FuncDecl, errors: &mut Vec) { - let Some(body) = func_decl.body else { - return; - }; - let mut scope: HashSet = func_decl - .params - .iter() - .map(|p| p.name.to_string()) - .collect(); - let mut tainted: HashSet = HashSet::new(); - escape_walk(body, &func_decl.arena, &mut scope, &mut tainted, errors); +/// Binding-only view of this check. It deliberately has no solved types and +/// remains usable when source errors prevent checked-program publication. +struct CheckingBindings<'a> { + arena: &'a ExprArena, + references: &'a [Option], + binders: &'a [Vec], + visited: &'a [bool], } -/// Walk `expr` tracking: -/// - `scope`: all names declared in the enclosing function (used by lambda capture detection) -/// - `tainted`: names that are bound to a capturing lambda and must not be returned +impl BindingFacts for CheckingBindings<'_> { + fn binding_node(&self, id: ExprID) -> BindingNode { + let mut node = BindingNode { + children: self.arena[id].subexprs(), + used: None, + declared: self.binders[id].clone(), + complete: self.visited[id], + }; + match &self.arena[id] { + Expr::Id(_) | Expr::TypeApp(_, _) => match &self.references[id] { + Some(Reference::Local(local)) => node.used = Some(*local), + None => node.complete = false, + Some(Reference::Functions(candidates)) if candidates.is_empty() => { + node.complete = false; + } + _ => {} + }, + Expr::Error | Expr::Macro(..) => node.complete = false, + _ => {} + } + node + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum CaptureState { + Noncapturing, + Unknown, + Capturing, +} + +impl CaptureState { + fn union(self, other: Self) -> Self { + match (self, other) { + (Self::Capturing, _) | (_, Self::Capturing) => Self::Capturing, + (Self::Unknown, _) | (_, Self::Unknown) => Self::Unknown, + _ => Self::Noncapturing, + } + } +} + +/// Preserve the existing escape traversal and conservative call-result taint. +/// Binding IDs remove scope reconstruction: taint for a departed scope cannot +/// affect a sibling binding. Implicit returns, assignment propagation and returns +/// inside lambda bodies remain outside this analysis. fn escape_walk( expr: ExprID, - arena: &ExprArena, - scope: &mut HashSet, - tainted: &mut HashSet, + facts: &CheckingBindings<'_>, + tainted: &mut HashMap, errors: &mut Vec, ) { - match &arena[expr] { - Expr::Return(e) => { - // Walk first so any declarations inside `e` update `tainted`. - escape_walk(*e, arena, scope, tainted, errors); - if expr_is_tainted(*e, arena, scope, tainted) { + match &facts.arena[expr] { + Expr::Return(value) => { + escape_walk(*value, facts, tainted, errors); + if expr_capture_state(*value, facts, tainted) == CaptureState::Capturing { errors.push(TypeError { - location: arena.locs[*e], + location: facts.arena.locs[*value], message: "closure with captured variables cannot be returned \ (captured addresses would dangle after the frame exits)" .to_string(), }); } } - Expr::Lambda { .. } => { - // Do not recurse into the lambda body — it is a separate scope. - } - Expr::Block(exprs) => { - let saved = scope.clone(); - for e in exprs { - escape_walk(*e, arena, scope, tainted, errors); - } - *scope = saved; - } - Expr::Let(name, init, _) => { - escape_walk(*init, arena, scope, tainted, errors); - if expr_is_tainted(*init, arena, scope, tainted) { - tainted.insert(name.to_string()); - } - scope.insert(name.to_string()); - } - Expr::Var(name, init, _) => { - if let Some(init_id) = init { - escape_walk(*init_id, arena, scope, tainted, errors); - if expr_is_tainted(*init_id, arena, scope, tainted) { - tainted.insert(name.to_string()); - } - } - scope.insert(name.to_string()); - } - Expr::If(cond, then, else_) => { - escape_walk(*cond, arena, scope, tainted, errors); - escape_walk(*then, arena, scope, tainted, errors); - if let Some(e) = else_ { - escape_walk(*e, arena, scope, tainted, errors); - } - } - Expr::While(cond, body) => { - escape_walk(*cond, arena, scope, tainted, errors); - escape_walk(*body, arena, scope, tainted, errors); - } - Expr::For { - var, - start, - end, - body, - } => { - escape_walk(*start, arena, scope, tainted, errors); - escape_walk(*end, arena, scope, tainted, errors); - scope.insert(var.to_string()); - escape_walk(*body, arena, scope, tainted, errors); - } - Expr::Call(f, args) => { - escape_walk(*f, arena, scope, tainted, errors); - for a in args { - escape_walk(*a, arena, scope, tainted, errors); - } - } - Expr::Binop(_, lhs, rhs) => { - escape_walk(*lhs, arena, scope, tainted, errors); - escape_walk(*rhs, arena, scope, tainted, errors); - } - Expr::Unop(_, arg) => escape_walk(*arg, arena, scope, tainted, errors), - Expr::Assume(e) | Expr::Field(e, _) | Expr::AsTy(e, _) | Expr::Arena(e) => { - escape_walk(*e, arena, scope, tainted, errors); - } - Expr::ArrayIndex(arr, idx) | Expr::Array(arr, idx) => { - escape_walk(*arr, arena, scope, tainted, errors); - escape_walk(*idx, arena, scope, tainted, errors); - } - Expr::ArrayLiteral(elems) | Expr::Tuple(elems) => { - for e in elems { - escape_walk(*e, arena, scope, tainted, errors); - } - } - Expr::StructLit(_, fields) => { - for (_, fval) in fields { - escape_walk(*fval, arena, scope, tainted, errors); + Expr::Lambda { .. } => {} + Expr::Let(_, init, _) | Expr::Var(_, Some(init), _) => { + escape_walk(*init, facts, tainted, errors); + let state = expr_capture_state(*init, facts, tainted); + if let Some(&local) = facts.binders[expr].first() { + tainted.insert(local, state); } } - Expr::Macro(_, args) => { - for a in args { - escape_walk(*a, arena, scope, tainted, errors); + expression => { + for child in expression.subexprs() { + escape_walk(child, facts, tainted, errors); } } - _ => {} } } -/// Returns true if `expr` might produce a capturing lambda value: -/// - a lambda literal whose body references names from `scope` beyond its own params -/// - an identifier in `tainted` (bound to a capturing lambda earlier) -/// - a block or if/else whose value branch is tainted -fn expr_is_tainted( +/// Unknown facts cannot prove noncapture or create an escape error. A resolved +/// free-local use is sufficient to diagnose capture even with other facts missing. +fn expr_capture_state( expr: ExprID, - arena: &ExprArena, - scope: &HashSet, - tainted: &HashSet, -) -> bool { - match &arena[expr] { - Expr::Lambda { params, body } => { - let lambda_params: HashSet = - params.iter().map(|p| p.name.to_string()).collect(); - lambda_has_captures(*body, arena, &lambda_params, scope) - } - Expr::Id(name) => tainted.contains(&name.to_string()), - Expr::Block(exprs) => exprs - .last() - .map_or(false, |e| expr_is_tainted(*e, arena, scope, tainted)), - Expr::If(_, then, else_) => { - expr_is_tainted(*then, arena, scope, tainted) - || else_.map_or(false, |e| expr_is_tainted(e, arena, scope, tainted)) - } - // Conservatively: if a tainted value flows into a call (as the function - // expression or as any argument), the result is also considered tainted. - // This catches laundering patterns like `return launder(f)` where `f` is - // a capturing lambda. - Expr::Call(fn_expr, args) => { - expr_is_tainted(*fn_expr, arena, scope, tainted) - || args - .iter() - .any(|a| expr_is_tainted(*a, arena, scope, tainted)) - } - _ => false, - } -} - -/// Returns true if `body` (the lambda body) references any name that is in -/// `outer_scope` but not in `lambda_params`. -fn lambda_has_captures( - body: ExprID, - arena: &ExprArena, - lambda_params: &HashSet, - outer_scope: &HashSet, -) -> bool { - match &arena[body] { - Expr::Id(name) => { - let s = name.to_string(); - outer_scope.contains(&s) && !lambda_params.contains(&s) - } - Expr::Binop(_, lhs, rhs) => { - lambda_has_captures(*lhs, arena, lambda_params, outer_scope) - || lambda_has_captures(*rhs, arena, lambda_params, outer_scope) - } - Expr::Unop(_, arg) => lambda_has_captures(*arg, arena, lambda_params, outer_scope), - Expr::Call(f, args) => { - lambda_has_captures(*f, arena, lambda_params, outer_scope) - || args - .iter() - .any(|a| lambda_has_captures(*a, arena, lambda_params, outer_scope)) - } - Expr::Block(exprs) => exprs - .iter() - .any(|e| lambda_has_captures(*e, arena, lambda_params, outer_scope)), - Expr::If(cond, then, else_) => { - lambda_has_captures(*cond, arena, lambda_params, outer_scope) - || lambda_has_captures(*then, arena, lambda_params, outer_scope) - || else_ - .map(|e| lambda_has_captures(e, arena, lambda_params, outer_scope)) - .unwrap_or(false) - } - Expr::Return(e) - | Expr::Assume(e) - | Expr::Field(e, _) - | Expr::AsTy(e, _) - | Expr::Arena(e) => lambda_has_captures(*e, arena, lambda_params, outer_scope), - Expr::While(cond, body) | Expr::ArrayIndex(cond, body) | Expr::Array(cond, body) => { - lambda_has_captures(*cond, arena, lambda_params, outer_scope) - || lambda_has_captures(*body, arena, lambda_params, outer_scope) - } - Expr::For { - start, end, body, .. - } => { - lambda_has_captures(*start, arena, lambda_params, outer_scope) - || lambda_has_captures(*end, arena, lambda_params, outer_scope) - || lambda_has_captures(*body, arena, lambda_params, outer_scope) - } - Expr::Let(_, init, _) => lambda_has_captures(*init, arena, lambda_params, outer_scope), - Expr::Var(_, init, _) => init - .map(|e| lambda_has_captures(e, arena, lambda_params, outer_scope)) - .unwrap_or(false), - Expr::ArrayLiteral(elems) | Expr::Tuple(elems) => elems - .iter() - .any(|e| lambda_has_captures(*e, arena, lambda_params, outer_scope)), - Expr::StructLit(_, fields) => fields - .iter() - .any(|(_, fval)| lambda_has_captures(*fval, arena, lambda_params, outer_scope)), - Expr::Macro(_, args) => args + facts: &CheckingBindings<'_>, + tainted: &HashMap, +) -> CaptureState { + match &facts.arena[expr] { + Expr::Lambda { .. } => { + let free = free_locals(facts, expr, []); + if !free.locals.is_empty() { + CaptureState::Capturing + } else if free.complete { + CaptureState::Noncapturing + } else { + CaptureState::Unknown + } + } + Expr::Id(_) => { + let node = facts.binding_node(expr); + if let Some(local) = node.used { + tainted + .get(&local) + .copied() + .unwrap_or(CaptureState::Noncapturing) + } else if node.complete { + CaptureState::Noncapturing + } else { + CaptureState::Unknown + } + } + Expr::Error => CaptureState::Unknown, + Expr::Block(exprs) => exprs.last().map_or(CaptureState::Noncapturing, |e| { + expr_capture_state(*e, facts, tainted) + }), + Expr::If(_, then, else_) => expr_capture_state(*then, facts, tainted).union( + else_.map_or(CaptureState::Noncapturing, |e| { + expr_capture_state(e, facts, tainted) + }), + ), + // Calls conservatively carry capture from either callee or arguments, + // including laundering patterns such as `return identity(capturing)`. + Expr::Call(callee, args) => args .iter() - .any(|a| lambda_has_captures(*a, arena, lambda_params, outer_scope)), - Expr::Lambda { params, body } => { - // Nested lambda: extend lambda_params with nested params. - let mut inner_params = lambda_params.clone(); - for p in params { - inner_params.insert(p.name.to_string()); - } - lambda_has_captures(*body, arena, &inner_params, outer_scope) - } - _ => false, + .fold(expr_capture_state(*callee, facts, tainted), |state, arg| { + state.union(expr_capture_state(*arg, facts, tainted)) + }), + _ => CaptureState::Noncapturing, } } @@ -2036,8 +2330,275 @@ mod tests { assert!(errors[0].message.contains("slice type")); } + #[test] + fn borrowed_parameters_respect_local_callee_shadowing() { + for source in [ + "call(a: &i32, b: &i32) {} main { let call = |a: i32, b: i32| { a + b }; let x = 1; call(x, x) }", + "call(a: &i32) {} main { let call = |a: i32| { a }; call(1) }", + "call(a: &i32) {} run(call: i32 -> i32) { call(1) }", + ] { + assert!(check(source).is_empty(), "{}", source); + } + } + + #[test] + fn borrowed_local_callee_still_requires_assignable_and_distinct_arguments() { + for (body, message) in [ + ( + "let call = set; call(1, 2)", + "reference argument must be assignable", + ), + ( + "let call = set; var x = 1; call(x, x)", + "multiple borrowed parameters", + ), + ] { + let source = format!( + "call(a: i32, b: i32) {{}} set(a: &i32, b: &i32) {{}} main {{ {} }}", + body + ); + let errors = check(&source); + assert!( + errors.iter().any(|e| e.message.contains(message)), + "{:?}", + errors + ); + } + } + + #[test] + fn borrowed_explicit_type_applications_keep_the_recorded_candidate_union() { + for callee in ["call", "call⟨i32⟩"] { + let errors = check(&format!( + "call(a: T) {{}} call(a: &T, flag: bool) {{}} main {{ {}(1) }}", + callee + )); + assert_eq!(errors.len(), 1, "{:?}", errors); + assert!(errors[0] + .message + .contains("reference argument must be assignable")); + + let errors = check(&format!( + "call(a: &T, b: T) {{}} call(a: T, b: &T, flag: bool) {{}} main {{ var x = 1; {}(x, x) }}", callee + )); + assert_eq!(errors.len(), 1, "{:?}", errors); + assert!(errors[0].message.contains("multiple borrowed parameters")); + } + // Explicit application recorded only the one-type-argument overload. + assert!(check("call(a: T) {} call(a: &T) {} main { call⟨i32⟩(1) }").is_empty()); + // Explicit type arguments can themselves introduce a borrowed type. + let errors = check("call(a: T) {} main { call⟨&i32⟩(1) }"); + assert_eq!(errors.len(), 1, "{:?}", errors); + assert!(errors[0] + .message + .contains("reference argument must be assignable")); + let errors = check("call(a: T, b: T) {} main { var x = 1; call⟨&i32⟩(x, x) }"); + assert_eq!(errors.len(), 1, "{:?}", errors); + assert!(errors[0].message.contains("multiple borrowed parameters")); + } + + #[test] + fn borrowed_slice_and_reference_consumers_remain_distinct() { + assert!(check( + "call(a: [T], b: [T]) {} main { let a = [1, 2]; let b = [3, 4]; call⟨i32⟩(a, b) }" + ) + .is_empty()); + let errors = check("call(a: [T], b: [T]) {} main { let a = [1, 2]; call⟨i32⟩(a, a) }"); + assert_eq!(errors.len(), 1, "{:?}", errors); + assert!(errors[0].message.contains("multiple slice parameters")); + assert!(check( + "set(a: &i32, b: &i32) {} main { var a = 1; var b = 2; let f = set; f(a, b) }" + ) + .is_empty()); + } + + #[test] + fn borrowed_function_valued_expressions_and_globals_use_solved_signatures() { + for callee in ["(if true { set } else { set })", "callbacks⟨i32⟩"] { + let declarations = "set(a: &i32, b: &i32) {} var callbacks: (&T, &T) -> void"; + let errors = check(&format!("{} main {{ {}(1, 2) }}", declarations, callee)); + assert_eq!(errors.len(), 2, "{:?}", errors); + assert!(errors + .iter() + .all(|e| e.message.contains("reference argument must be assignable"))); + let errors = check(&format!( + "{} main {{ var x = 1; {}(x, x) }}", + declarations, callee + )); + assert_eq!(errors.len(), 1, "{:?}", errors); + assert!(errors[0].message.contains("multiple borrowed parameters")); + } + } + + #[test] + fn borrowed_function_global_candidates_require_assignable_arguments() { + let errors = check( + "struct callbacks { value: i32 } var callbacks: (&i32, &i32) -> void main { callbacks(1, 2) }", + ); + assert_eq!(errors.len(), 2, "{:?}", errors); + assert!(errors + .iter() + .all(|e| e.message == "reference argument must be assignable")); + } + + #[test] + fn borrowed_function_global_candidates_require_distinct_arguments() { + let errors = check( + "struct callbacks { value: i32 } var callbacks: (&i32, &i32) -> void main { var x = 1; callbacks(x, x) }", + ); + assert_eq!(errors.len(), 1, "{:?}", errors); + assert_eq!( + errors[0].message, + "cannot pass 'x' to multiple borrowed parameters (borrowed arguments must not alias)", + ); + } + // --- Escape analysis tests --- + #[test] + fn escape_lambda_local_shadowing_does_not_capture_outer_binding() { + for body in [ + "let x = 42; x", + "var x = 42; x", + "for x in 0 .. 2 { let use_x = x; }; 42", + "let x = 42; let nested = || { x }; nested()", + ] { + let source = format!("f(x: i32) -> void -> i32 {{ return (|| {{ {} }}) }}", body); + assert!(check(&source).is_empty(), "{}", source); + } + // The initializer still refers to the outer x before the new x is bound. + let errors = check("f(x: i32) -> void -> i32 { return (|| { let x = x; x }) }"); + assert_eq!(errors.len(), 1, "{:?}", errors); + assert!(errors[0].message.contains("captured")); + } + + #[test] + fn escape_taint_is_owned_by_sibling_and_shadowed_bindings() { + for source in [ + "f(x: i32) -> void -> i32 { { let g = || { x }; g() }; { let g = || { 42 }; return g } }", + "f(x: i32, flag: bool) -> void -> i32 { if flag { let g = || { x }; g() } else { let g = || { 42 }; return g }; return (|| { 0 }) }", + "f(x: i32) -> void -> i32 { let g = || { x }; let g = || { 42 }; return g }", + ] { + assert!(check(source).is_empty(), "{}", source); + } + // Leaving a shadowing scope must not erase the original binding's taint. + let errors = check( + "f(x: i32) -> void -> i32 { let g = || { x }; { let g = || { 42 }; g() }; return g }", + ); + assert_eq!(errors.len(), 1, "{:?}", errors); + assert!(errors[0].message.contains("captured")); + } + + #[test] + fn escape_nested_capture_uses_outer_binding_even_after_inner_shadowing() { + let errors = check("f(x: i32) -> void -> i32 { return (|| { let nested = || { x }; let x = 42; nested() + x }) }"); + assert_eq!(errors.len(), 1, "{:?}", errors); + assert!(errors[0].message.contains("captured")); + assert!(check("f(x: i32) -> void -> i32 { return (|| { let x = 42; let nested = || { x }; nested() }) }").is_empty()); + } + + #[test] + fn escape_free_locals_remain_partial_after_source_errors() { + for (body, captures) in [ + ("missing; x", true), + ("let nested = || { missing; x }; nested()", true), + ("missing", false), + ("let x = missing; x", false), + ("Missing(field: x)", false), // the unknown struct's fields are unvisited + ] { + let source = format!("f(x: i32) -> void -> i32 {{ return (|| {{ {} }}) }}", body); + let mut errors = vec![]; + let table = DeclTable::new(parse_program_str(&source, &mut errors)); + assert!(errors.is_empty()); + let declaration = &table.decls[0]; + let Decl::Func(function) = declaration else { + panic!() + }; + let mut checker = Checker::new(); + checker.check_decl(declaration, &table); + assert!(checker + .errors + .iter() + .any(|e| !e.message.contains("captured"))); + assert_eq!( + checker + .errors + .iter() + .any(|e| e.message.contains("captured")), + captures, + "{}", + source + ); + let lambda = function + .arena + .exprs + .iter() + .rposition(|e| matches!(e, Expr::Lambda { .. })) + .unwrap(); + let facts = CheckingBindings { + arena: &function.arena, + references: &checker.references, + binders: &checker.binders, + visited: &checker.visited, + }; + let free = free_locals(&facts, lambda, []); + assert!(!free.complete, "{}", source); + assert_eq!( + free.locals, + if captures { vec![LocalId(0)] } else { vec![] } + ); + assert_eq!( + expr_capture_state(lambda, &facts, &HashMap::new()), + if captures { + CaptureState::Capturing + } else { + CaptureState::Unknown + } + ); + } + } + + #[test] + fn escape_diagnostics_preserve_editor_reference_and_type_trust() { + let mut compiler = Compiler::new(); + compiler.quiet = true; + assert!(compiler.parse("good() -> i32 { 42 } f(x: i32) -> void -> i32 { let failed = missing; return (|| { failed; x }) }", "partial-capture.lyte")); + assert!(!compiler.analyze()); + assert!(compiler + .last_type_errors + .iter() + .any(|e| e.message.contains("captured"))); + assert!(compiler.checked_program().is_none()); + let analysis = compiler.source_analysis().unwrap(); + let definition = analysis.declarations().named_ids(Name::str("f"))[0]; + let function = analysis.declarations().function(definition).unwrap(); + let body = analysis.body(definition).unwrap(); + for (id, expression) in function.arena.exprs.iter().enumerate() { + if let Expr::Id(name) = expression { + let facts = body.expression(id).unwrap(); + match name.as_str() { + "x" => { + assert!(matches!(facts.reference, Some(Reference::Local(_)))); + assert_eq!(facts.ty, Some(mk_type(Type::Int32))); + } + "failed" => { + assert!(matches!(facts.reference, Some(Reference::Local(_)))); + assert_eq!(facts.ty, None); + } + "missing" => { + assert_eq!(facts.reference, None); + assert_eq!(facts.ty, None); + } + _ => unreachable!(), + } + } + } + compiler.last_errors.clear(); + compiler.last_type_errors.clear(); + assert!(compiler.specialize().is_err()); + assert!(compiler.specialized_program().is_err()); + } + #[test] pub fn test_escape_direct_return_capturing_lambda() { let s = " diff --git a/src/compiler.rs b/src/compiler.rs index df003f5f..5f35c5c8 100644 --- a/src/compiler.rs +++ b/src/compiler.rs @@ -1,8 +1,9 @@ use crate::vm::{LinkedProgram, VMProgram, VM}; use crate::vm_codegen::VMCodegen; use crate::*; +#[cfg(feature = "cranelift")] use core::mem; -use std::collections::HashMap; +use std::collections::{HashMap, HashSet, VecDeque}; use std::fs; /// A compiled program ready for execution, abstracting over backends. @@ -68,8 +69,6 @@ fn builtin_decls() -> Vec { requires: vec![], loc: test_loc(), arena: ExprArena::new(), - types: vec![], - closure_vars: vec![], is_extern: false, }), // print(value: i32) → void @@ -87,8 +86,6 @@ fn builtin_decls() -> Vec { requires: vec![], loc: test_loc(), arena: ExprArena::new(), - types: vec![], - closure_vars: vec![], is_extern: false, }), // putc(c: i32) → void @@ -106,8 +103,6 @@ fn builtin_decls() -> Vec { requires: vec![], loc: test_loc(), arena: ExprArena::new(), - types: vec![], - closure_vars: vec![], is_extern: false, }), ]; @@ -136,8 +131,6 @@ fn builtin_decls() -> Vec { requires: vec![], loc: test_loc(), arena: ExprArena::new(), - types: vec![], - closure_vars: vec![], is_extern: false, })); } @@ -162,8 +155,6 @@ fn builtin_decls() -> Vec { requires: vec![], loc: test_loc(), arena: ExprArena::new(), - types: vec![], - closure_vars: vec![], is_extern: false, })); } @@ -196,8 +187,6 @@ fn builtin_decls() -> Vec { requires: vec![], loc: test_loc(), arena: ExprArena::new(), - types: vec![], - closure_vars: vec![], is_extern: false, })); } @@ -234,8 +223,7 @@ fn builtin_decls() -> Vec { requires: vec![], loc: test_loc(), arena: ExprArena::new(), - types: vec![], - closure_vars: vec![], + is_extern: false, })); @@ -254,8 +242,7 @@ fn builtin_decls() -> Vec { requires: vec![], loc: test_loc(), arena: ExprArena::new(), - types: vec![], - closure_vars: vec![], + is_extern: false, })); @@ -298,148 +285,82 @@ fn rewrite_qualified_enums(arena: &mut ExprArena, decls: &DeclTable) { } } -/// After type-checking, rewrite `Binop(op, lhs, rhs)` into -/// `Call(__add/__sub/__mul/__div, [lhs, rhs])` when the operand type is a -/// named (struct) type. The checker resolves the overload through its Or -/// constraint, but the JIT/VM only handle primitive types in binop codegen. -/// -/// Float `%` is rewritten the same way: no backend has a primitive float -/// remainder instruction, so it lowers to the stdlib's `__mod` overloads. -fn rewrite_overloaded_binops(fdecl: &mut FuncDecl) { - let n = fdecl.arena.exprs.len(); - for i in 0..n { - if let Expr::Binop(op, lhs, rhs) = fdecl.arena.exprs[i].clone() { - if !op.arithmetic() { - continue; - } - let lhs_ty = fdecl.types[lhs]; - let is_float_mod = op == Binop::Mod && matches!(*lhs_ty, Type::Float32 | Type::Float64); - if matches!(*lhs_ty, Type::Name(_, _)) || is_float_mod { - let result_ty = fdecl.types[i]; - let rhs_ty = fdecl.types[rhs]; - // Build the function type: (lhs_ty, rhs_ty) -> result_ty - let fn_ty = func(tuple(vec![lhs_ty, rhs_ty]), result_ty); - - let overload = Name::new(op.overload_name().into()); - let fn_id = fdecl.arena.add(Expr::Id(overload), fdecl.arena.locs[i]); - // Extend types array to cover the new expression - while fdecl.types.len() <= fn_id { - fdecl.types.push(mk_type(Type::Void)); - } - fdecl.types[fn_id] = fn_ty; - fdecl.arena.exprs[i] = Expr::Call(fn_id, vec![lhs, rhs]); - } - } - } +/// Give each source occurrence its own expression node before checking. +/// Macro substitution can share an argument under distinct lexical scopes; +/// resolved references must describe each occurrence in its actual scope. +fn normalize_source_body<'a>( + source: &mut ExprArena, + roots: impl IntoIterator, +) { + fn copy(id: ExprID, source: &ExprArena, output: &mut ExprArena) -> ExprID { + let mut expression = source[id].clone(); + expression.map_children(|child| copy(child, source, output)); + output.add(expression, source.locs[id]) + } + let mut arena = ExprArena::new(); + for root in roots { + *root = copy(*root, source, &mut arena); + } + *source = arena; } -/// Rename non-generic overloaded functions (both declarations and call sites) -/// so each overload gets a unique symbol. e.g. two `add` overloads with -/// different param types become `add$i32$i32` and `add$f32$f32`. -fn rename_overloaded_functions(decls: &mut DeclTable) { - // Count non-generic overloads per name. - let mut counts: HashMap = HashMap::new(); - for d in &decls.decls { - if let Decl::Func(f) = d { - if f.typevars.is_empty() { - *counts.entry(f.name).or_default() += 1; - } - } - } - - // Build a mapping from (original_name, param_types) -> mangled_name - // for overloaded functions. - let mut overload_map: HashMap, Name)>> = HashMap::new(); - for d in &decls.decls { - if let Decl::Func(f) = d { - if f.typevars.is_empty() && counts.get(&f.name).copied().unwrap_or(0) > 1 { - let param_types: Vec = f.params.iter().filter_map(|p| p.ty).collect(); - let mangled = mangle::mangle_overload(f.name, ¶m_types); - overload_map - .entry(f.name) - .or_default() - .push((param_types, mangled)); - } - } - } +/// Options that affect source validation, rather than diagnostics or codegen. +/// Add future validation-affecting options here so cached output cannot bypass +/// their checks. Public option fields are compared at each executable boundary. +#[derive(Clone, Copy, PartialEq, Eq)] +struct ValidationOptions { + no_recursion: bool, +} - if overload_map.is_empty() { - return; - } +struct CheckedState { + templates: CheckedProgram, + specialization: Option, + options: ValidationOptions, + source_valid: bool, +} - // Rename function declarations. - for d in &mut decls.decls { - if let Decl::Func(ref mut f) = d { - if let Some(overloads) = overload_map.get(&f.name) { - let param_types: Vec = f.params.iter().filter_map(|p| p.ty).collect(); - for (pts, mangled) in overloads { - if *pts == param_types { - f.name = *mangled; - break; - } - } - } - } - } +/// Parsing/checking replaces the semantic owner. Specialization only borrows +/// its immutable templates and publishes an independent concrete program. +enum ProgramState { + Unchecked, + Checked(CheckedState), +} - // Rename call-site references in all function bodies. - // For each Expr::Id that matches an overloaded name, look at its resolved - // type to determine which overload it refers to. - for d in &mut decls.decls { - if let Decl::Func(ref mut f) = d { - for i in 0..f.arena.exprs.len() { - if let Expr::Id(name) = &f.arena.exprs[i] { - if let Some(overloads) = overload_map.get(name) { - if let Some(&call_ty) = f.types.get(i) { - // The call-site type is the function type. - // Extract param types from it. - if let Type::Func(dom, _) = &*call_ty { - let call_params = match &**dom { - Type::Tuple(ts) => ts.clone(), - _ => vec![*dom], - }; - for (pts, mangled) in overloads { - if pts.len() == call_params.len() - && pts.iter().zip(call_params.iter()).all(|(a, b)| { - let mut inst = crate::Instance::new(); - crate::types::unify(*a, *b, &mut inst) - }) - { - f.arena.exprs[i] = Expr::Id(*mangled); - break; - } - } - } - } - } - } - } - } - } +#[derive(Default)] +struct PhaseDiagnostics { + messages: Vec, + safety_errors: Vec, } pub struct Compiler { ast: Vec, - decls: DeclTable, + program: ProgramState, + analysis: Option, + source_diagnostics: PhaseDiagnostics, + specialization_diagnostics: PhaseDiagnostics, pub print_ir: bool, /// Number of AST trees that belong to the stdlib (parsed in new()). stdlib_trees: usize, /// Entry point function names. If empty, defaults to ["main"]. entry_points: Vec, - /// Formatted error messages from the last parse/check operation. + /// Formatted source-validation errors plus the current specialization errors. + /// Diagnostic lists are views; clearing them does not authorize compilation. pub last_errors: Vec, - /// Structured parse errors from the last parse operation. + /// Structured parse errors from all accumulated input files. pub last_parse_errors: Vec, /// Structured type errors from the last check operation. pub last_type_errors: Vec, - /// Structured safety errors from the last check operation. + /// Structured source and current specialization safety errors. pub last_safety_errors: Vec, /// When true, suppress all stdout output (for LSP usage). pub quiet: bool, - /// When true, continue checking all declarations even after errors (for LSP). + /// When true, continue type checking declarations after type errors (for LSP). + /// Parse errors still prevent checking with either setting. pub check_all: bool, /// When true, reject recursive functions (direct or mutual) in the safety /// checker and tell the native backends to skip call-depth runtime checks. + /// If this differs from the value used by check(), specialization and code + /// generation require another check(). Checked type facts remain available. pub no_recursion: bool, } @@ -447,7 +368,10 @@ impl Compiler { pub fn new() -> Self { let mut c = Self { ast: Vec::new(), - decls: DeclTable::new(vec![]), + program: ProgramState::Unchecked, + analysis: None, + source_diagnostics: PhaseDiagnostics::default(), + specialization_diagnostics: PhaseDiagnostics::default(), print_ir: false, stdlib_trees: 0, entry_points: Vec::new(), @@ -465,9 +389,80 @@ impl Compiler { } /// Set custom entry point functions. If not called (or called with empty slice), - /// defaults to ["main"]. + /// defaults to ["main"]. Changing effective roots clears specialization and + /// its diagnostics, while retaining checked templates and source diagnostics. pub fn set_entry_points(&mut self, names: &[&str]) { + let previous = self.effective_entry_points(); self.entry_points = names.iter().map(|n| Name::new((*n).into())).collect(); + if self.effective_entry_points() != previous { + self.clear_specialization(); + } + } + + fn validation_options(&self) -> ValidationOptions { + ValidationOptions { + no_recursion: self.no_recursion, + } + } + + fn require_current_validation(&self, checked: &CheckedState) -> Result<(), String> { + if checked.options != self.validation_options() { + return Err( + "validation options changed; call check() before specialization or code generation" + .into(), + ); + } + if !checked.source_valid { + return Err("cannot specialize a program with checking errors".into()); + } + Ok(()) + } + + fn refresh_diagnostics(&mut self) { + self.last_errors = self.source_diagnostics.messages.clone(); + self.last_errors + .extend(self.specialization_diagnostics.messages.iter().cloned()); + self.last_safety_errors = self.source_diagnostics.safety_errors.clone(); + self.last_safety_errors.extend( + self.specialization_diagnostics + .safety_errors + .iter() + .cloned(), + ); + } + + fn clear_specialization(&mut self) { + if let ProgramState::Checked(checked) = &mut self.program { + checked.specialization = None; + } + self.specialization_diagnostics = PhaseDiagnostics::default(); + self.refresh_diagnostics(); + } + + /// Rebuild parse diagnostics from their source owners, including errors in + /// earlier inputs. All facts and diagnostics derived by checking are stale. + fn reset_analysis(&mut self) { + self.program = ProgramState::Unchecked; + self.analysis = None; + self.source_diagnostics = PhaseDiagnostics::default(); + self.specialization_diagnostics = PhaseDiagnostics::default(); + self.last_type_errors.clear(); + self.last_parse_errors = self + .ast + .iter() + .flat_map(|tree| tree.errors.clone()) + .collect(); + self.source_diagnostics.messages = self + .last_parse_errors + .iter() + .map(|err| { + format!( + "{}:{}: {}", + err.location.file, err.location.line, err.message + ) + }) + .collect(); + self.refresh_diagnostics(); } /// Returns the effective entry points (defaults to ["main"] if none set). @@ -486,9 +481,7 @@ impl Compiler { /// True if `name` is declared as a top level function in the parsed source. /// /// Answered from the AST rather than the decl table so it is valid as soon - /// as the source is parsed. `check()` copies every tree decl into - /// `self.decls`, so the two agree once it has run; before that `self.decls` - /// is empty and every entry point would look missing. + /// as the source is parsed, including before checked declarations exist. fn entry_point_is_defined(&self, name: Name) -> bool { self.ast.iter().any(|tree| { tree.decls @@ -534,7 +527,10 @@ impl Compiler { } } + /// Append an input file and invalidate prior checking/specialization. The + /// return value describes this input; errors in earlier files remain active. pub fn parse(&mut self, contents: &str, path: &str) -> bool { + self.program = ProgramState::Unchecked; let mut lexer = Lexer::new(&contents, &path); let mut tree = Tree::default(); @@ -542,8 +538,6 @@ impl Compiler { lexer.next(); tree.decls = parse_program(&mut lexer, &mut tree.errors); - self.last_errors.clear(); - self.last_parse_errors.clear(); for err in &tree.errors { let msg = format!( "{}:{}: {}", @@ -552,27 +546,73 @@ impl Compiler { if !self.quiet { println!("{}", msg); } - self.last_errors.push(msg); - self.last_parse_errors.push(err.clone()); } let success = tree.errors.is_empty(); self.ast.push(tree); + self.reset_analysis(); success } + /// Revalidate all accumulated source using the current validation options. + /// Replaces templates and clears any prior specialization and its errors. pub fn check(&mut self) -> bool { - self.last_errors.clear(); - self.last_type_errors.clear(); - self.last_safety_errors.clear(); + self.reset_analysis(); + // Recovered syntax can retain Expr::Error even when contextual typing + // would succeed. Never build checked bodies from those inputs. + if !self.last_parse_errors.is_empty() { + return false; + } + let success = self.check_source(false); + self.refresh_diagnostics(); + success + } + + /// Check for editor queries, retaining partial facts and visiting recovered + /// input. Uses the same checking pass as check(), but collects all body + /// diagnostics. Parse/type failures still cannot publish CheckedProgram. + pub fn analyze(&mut self) -> bool { + self.reset_analysis(); + let success = self.check_source(true); + self.refresh_diagnostics(); + success + } + + pub fn source_analysis(&self) -> Option<&SourceAnalysis> { + self.analysis.as_ref() + } + + /// Expand and normalize source occurrences before assigning checked meaning. + /// Editor recovery keeps an inventory even if macro expansion fails. + fn prepare_source(&mut self, editor: bool) -> Option { let mut decls = builtin_decls(); + let mut parse_clean = vec![true; decls.len()]; for tree in &self.ast { - decls.append(&mut tree.decls.clone()); + parse_clean.extend(std::iter::repeat(tree.errors.is_empty()).take(tree.decls.len())); + decls.extend(tree.decls.iter().cloned()); + } + let recovered_files: HashSet<_> = self + .ast + .iter() + .flat_map(|tree| tree.errors.iter().map(|err| err.location.file)) + .collect(); + let mut recovered: HashSet<_> = parse_clean + .iter() + .enumerate() + .filter_map(|(index, clean)| (!clean).then_some(DefId(index as u32))) + .collect(); + if editor { + // Even an early macro-expansion failure leaves a source inventory. + self.analysis = Some(SourceAnalysis::new( + DeclTable::new(decls.clone()), + recovered.clone(), + recovered_files.clone(), + )); } // Collect macros for expansion, rejecting duplicates. let mut macros: HashMap = HashMap::new(); - let mut has_errors = false; + let mut has_duplicate_macros = false; for d in &decls { if let Decl::Macro(m) = d { if macros.contains_key(&m.name) { @@ -580,169 +620,289 @@ impl Compiler { if !self.quiet { print_error_with_context(m.loc, &msg); } - self.last_errors.push(format_error(m.loc, &msg)); - has_errors = true; + self.source_diagnostics + .messages + .push(format_error(m.loc, &msg)); + has_duplicate_macros = true; } else { macros.insert(m.name, m.clone()); } } } - if has_errors { - return false; + if has_duplicate_macros { + return None; } // Expand macros in all function bodies. for decl in &mut decls { if let Decl::Func(ref mut fdecl) = decl { if let Err((loc, msg)) = fdecl.expand_macros(¯os) { + self.last_type_errors.push(TypeError { + location: loc, + message: msg.clone(), + }); if !self.quiet { print_error_with_context(loc, &msg); } - self.last_errors.push(format_error(loc, &msg)); - return false; + self.source_diagnostics + .messages + .push(format_error(loc, &msg)); + return None; } } } - self.decls = DeclTable::new(decls); - let orig_decls = self.decls.clone(); - - // Rewrite qualified enum accesses (e.g. Direction.Up -> .Up) - // before type-checking, so the checker and downstream passes - // see them as Expr::Enum nodes. - for decl in &mut self.decls.decls { - if let Decl::Func(ref mut fdecl) | Decl::Macro(ref mut fdecl) = decl { - rewrite_qualified_enums(&mut fdecl.arena, &orig_decls); + // Macro declarations are source templates; expanded bodies are their + // only contribution to the checked program. + recovered.clear(); + decls = decls + .into_iter() + .zip(parse_clean) + .filter(|(decl, _)| !matches!(decl, Decl::Macro(_))) + .enumerate() + .map(|(index, (decl, clean))| { + if !clean { + recovered.insert(DefId(index as u32)); + } + decl + }) + .collect(); + let names = DeclTable::new(decls.clone()); + for decl in &mut decls { + match decl { + Decl::Func(function) => { + rewrite_qualified_enums(&mut function.arena, &names); + normalize_source_body( + &mut function.arena, + function.body.iter_mut().chain(function.requires.iter_mut()), + ); + } + Decl::Assume { arena, cond } => { + rewrite_qualified_enums(arena, &names); + normalize_source_body(arena, std::iter::once(cond)); + } + _ => {} } } + let source = DeclTable::new(decls); + self.analysis = + editor.then(|| SourceAnalysis::new(source.clone(), recovered, recovered_files)); + Some(source) + } - for decl in &mut self.decls.decls { - // Skip type-checking macro declarations (they are untyped templates). - if matches!(decl, Decl::Macro(_)) { - continue; - } - + fn check_source(&mut self, editor: bool) -> bool { + // Input trees, not publicly mutable diagnostic views, own this gate. + let parse_failed = self.ast.iter().any(|tree| !tree.errors.is_empty()); + let Some(source) = self.prepare_source(editor) else { + return false; + }; + let mut has_errors = false; + let mut functions = HashMap::new(); + let mut assumptions = VecDeque::new(); + for record in source.records() { let mut checker = Checker::new(); - checker.check_decl(decl, &orig_decls); - + checker.check_decl(&record.declaration, &source); + match &record.declaration { + Decl::Assume { arena, .. } if checker.errors.is_empty() && !parse_failed => { + assumptions.push_back(checker.checked_body(arena)); + } + Decl::Func(function) => { + if let Some(analysis) = &mut self.analysis { + let body = checker.body_analysis(function, analysis); + analysis.bodies.insert(record.definition, body); + } + if checker.errors.is_empty() && !parse_failed { + functions.insert(record.definition, checker.checked_function(function)); + } + } + Decl::Interface(interface) if checker.errors.is_empty() => { + for (function, &id) in interface.funcs.iter().zip(&record.members) { + checker.check_decl(&Decl::Func(function.clone()), &source); + if let Some(analysis) = &mut self.analysis { + let body = checker.body_analysis(function, analysis); + analysis.bodies.insert(id, body); + } + if checker.errors.is_empty() && !parse_failed { + functions.insert(id, checker.checked_function(function)); + } + } + } + _ => {} + } if !self.quiet { checker.print_errors(); } for err in &checker.errors { - self.last_errors + self.source_diagnostics + .messages .push(format_error(err.location, &err.message)); } - self.last_type_errors.extend(checker.errors.iter().cloned()); - if !checker.errors.is_empty() && !self.check_all { + has_errors |= !checker.errors.is_empty(); + self.last_type_errors.extend(checker.errors); + if has_errors && !self.check_all && !editor { return false; } - has_errors = has_errors || !checker.errors.is_empty(); - - if let Decl::Func(ref mut fdecl) = decl { - fdecl.types = checker.solved_types(); - } } - - // Rewrite overloaded binary ops (e.g. Point + Point) to function calls. - for decl in &mut self.decls.decls { - if let Decl::Func(ref mut fdecl) = decl { - rewrite_overloaded_binops(fdecl); - } + if has_errors || parse_failed { + return false; } - - // Static safety checks (array bounds, division by zero). - let mut safety_checker = SafetyChecker::new(); + let checked = CheckedProgram::try_new(source.map_bodies( + |id, _| { + functions + .remove(&id) + .expect("checked function or interface member") + }, + |_, _| assumptions.pop_front().expect("checked global assumption"), + )); + let checked = match checked { + Ok(checked) => checked, + Err(error) => { + let error = format!("invalid checked program: {}", error); + if !self.quiet { + eprintln!("{}", error); + } + self.source_diagnostics.messages.push(error); + return false; + } + }; + let mut safety = SafetyChecker::new(); if self.no_recursion { - safety_checker.check_recursion(&self.decls); + safety.check_recursion(&checked); } - safety_checker.check(&self.decls); + safety.check(&checked); if !self.quiet { - safety_checker.print_errors(); - } - for err in &safety_checker.errors { - self.last_errors - .push(format_error(err.location, &err.message)); - } - self.last_safety_errors - .extend(safety_checker.errors.iter().cloned()); - if !safety_checker.errors.is_empty() { - has_errors = true; + safety.print_errors(); + } + for error in &safety.errors { + self.source_diagnostics + .messages + .push(format_error(error.location, &error.message)); + } + let source_valid = safety.errors.is_empty(); + self.source_diagnostics.safety_errors = safety.errors; + // Safety errors do not invalidate type/resolution facts, but the saved + // validation result (not the public diagnostic lists) gates execution. + self.program = ProgramState::Checked(CheckedState { + templates: checked, + specialization: None, + options: self.validation_options(), + source_valid, + }); + source_valid + } + + /// Source syntax remains available for diagnostics when checking fails. + /// It is never accepted by a code generator. + pub fn parsed_declarations(&self) -> impl Iterator { + self.ast.iter().flat_map(|tree| tree.decls.iter()) + } + + /// Immutable source templates, also available after specialization or a + /// safety failure. These type facts alone do not authorize execution. + pub fn checked_program(&self) -> Option<&CheckedProgram> { + match &self.program { + ProgramState::Checked(checked) => Some(&checked.templates), + _ => None, } + } - !has_errors + /// Concrete output for the current roots and validation options. Every + /// compiler code-generation entry point passes through this guard. + pub fn specialized_program(&self) -> Result<&SpecializedProgram, String> { + match &self.program { + ProgramState::Checked(checked) => { + self.require_current_validation(checked)?; + checked + .specialization + .as_ref() + .ok_or_else(|| "specialization is required before code generation".into()) + } + _ => Err("specialization is required before code generation".into()), + } } - /// Returns a reference to the declaration table (available after check()). - pub fn decls(&self) -> &DeclTable { - &self.decls + /// Compatibility view: current concrete declarations when available, + /// otherwise checked templates. Prefer checked_program() for source queries + /// and specialized_program() for concrete consumers. + pub fn decls(&self) -> &DeclarationList { + match &self.program { + ProgramState::Checked(checked) => self + .specialized_program() + .map(|program| &program.decls) + .unwrap_or(&checked.templates.decls), + ProgramState::Unchecked => panic!("checking has not produced a program"), + } } pub fn specialize(&mut self) -> Result<(), String> { - let mut pass = MonomorphPass::new(); - let entry_points = self.effective_entry_points(); - let all_decls = pass.monomorphize_multi(&self.decls, &entry_points)?; - // Capture the names of newly-generated specialized decls so we can - // run a focused safety-check pass on them (their bodies were skipped - // pre-monomorph because they contained size variables). - let specialized_names: std::collections::HashSet = - pass.instantiated_names().collect(); - self.decls = DeclTable::new(all_decls); - - // Rename non-generic overloaded functions to unique symbols. - // Must happen after monomorphization so specialized generic bodies - // can resolve overloaded calls (e.g. cmp in a generic quicksort). - rename_overloaded_functions(&mut self.decls); - self.decls = DeclTable::new(self.decls.decls.clone()); - - // Re-run the safety checker on specialized declarations only. - // Their bodies were skipped pre-monomorph because they had non-empty - // size_vars; now sizes are concrete (`[T; Known(K)]`) so bounds and - // require clauses can be verified properly. - if !specialized_names.is_empty() { - let mut sc = SafetyChecker::new(); - for decl in &self.decls.decls { - if let Decl::Func(f) = decl { - if specialized_names.contains(&f.name) { - sc.check_decl(decl, &self.decls); - } - } - } - if !self.quiet { - sc.print_errors(); - } - for err in &sc.errors { - self.last_errors - .push(format_error(err.location, &err.message)); + let ProgramState::Checked(checked) = &self.program else { + return Err("checking is required before specialization".into()); + }; + self.require_current_validation(checked)?; + if checked.specialization.is_some() { + return Ok(()); + } + + // Failed attempts are retryable, even with the same roots. Never leave + // an older output or duplicate diagnostics attached to a new attempt. + self.clear_specialization(); + let result = match self.specialize_checked() { + Ok(program) => { + let ProgramState::Checked(checked) = &mut self.program else { + unreachable!("specialization retains checked templates"); + }; + checked.specialization = Some(program); + Ok(()) } - self.last_safety_errors.extend(sc.errors.iter().cloned()); - if !sc.errors.is_empty() { - return Err(format!("safety check failed for {} call(s)", sc.errors.len())); + Err(error) => { + if self.specialization_diagnostics.messages.is_empty() { + self.specialization_diagnostics.messages.push(error.clone()); + } + Err(error) } - } + }; + self.refresh_diagnostics(); + result + } - // Hoist loop-invariant struct field reads (after monomorphization - // so we operate on concrete types, and after safety checking). - { - let effects = crate::hoist::SideEffects::analyze(&self.decls); - for decl in &mut self.decls.decls { - if let Decl::Func(ref mut fdecl) = decl { - hoist_loop_invariant_fields(fdecl, &effects); - } + fn specialize_checked(&mut self) -> Result { + let checked = self.checked_program().expect("checked templates"); + let entries = self.effective_entry_points(); + let mut pass = MonomorphPass::new(); + let mut program = pass.monomorphize_multi(checked, &entries)?; + // Concrete targets and sizes expose obligations in ordinary callers as + // well as generic instances. Check every retained body before hoisting + // or any other code-moving transformation, then publish atomically. + let mut safety = SafetyChecker::new(); + safety.check(&program); + if !self.quiet { + safety.print_errors(); + } + for error in &safety.errors { + self.specialization_diagnostics + .messages + .push(format_error(error.location, &error.message)); + } + if !safety.errors.is_empty() { + let count = safety.errors.len(); + self.specialization_diagnostics.safety_errors = safety.errors; + return Err(format!("safety check failed for {} call(s)", count)); + } + let effects = crate::hoist::SideEffects::analyze(&program)?; + for declaration in &mut program.decls.decls { + if let Decl::Func(function) = declaration { + hoist_loop_invariant_fields(function, &effects)?; } } - - Ok(()) + program.validate()?; + Ok(program) } pub fn has_decls(&self) -> bool { - // Check if there are any user declarations beyond the built-ins and stdlib - let stdlib_decl_count: usize = self - .ast + self.ast .iter() - .take(self.stdlib_trees) - .map(|t| t.decls.len()) - .sum(); - self.decls.decls.len() > builtin_decls().len() + stdlib_decl_count + .skip(self.stdlib_trees) + .any(|tree| !tree.decls.is_empty()) } /// Returns info about each global variable: (name, offset, size, type_string). @@ -755,10 +915,10 @@ impl Compiler { ) -> Vec<(String, usize, usize, String, bool)> { let mut result = Vec::new(); let mut offset: usize = base_offset; - for decl in &self.decls.decls { + for decl in &self.decls().decls { match decl { Decl::Global { name, ty, .. } => { - let size = ty.size(&self.decls) as usize; + let size = ty.size(self.decls()) as usize; let type_str = ty.pretty_print(); result.push((name.to_string(), offset, size, type_str, false)); offset += size; @@ -770,7 +930,7 @@ impl Compiler { "extern fn({})", f.params .iter() - .map(|p| { p.ty.map_or("?".to_string(), |t| t.pretty_print()) }) + .map(|p| f.arena.local(p.local).ty.pretty_print()) .collect::>() .join(", ") ); @@ -794,7 +954,10 @@ impl Compiler { jit.print_ir = self.print_ir; jit.no_recursion = self.no_recursion; let entry_points = self.effective_entry_points(); - match jit.compile_and_run_multi(&self.decls, &entry_points) { + match jit.compile_and_run_multi( + self.specialized_program().expect("specialized program"), + &entry_points, + ) { Ok((trap_reason, compile_time, exec_time)) => { if let Some(msg) = crate::cancel::trap_reason_message(trap_reason) { eprintln!("trap: {}", msg); @@ -823,9 +986,6 @@ impl Compiler { prefix: &str, target: crate::llvm_aot::AotTarget, ) -> Result<(), String> { - if self.decls.decls.is_empty() { - return Err("No declarations to compile".into()); - } if !self.no_recursion { return Err( "--aot requires --no-recursion (call-depth machinery is unavailable at link time)" @@ -834,7 +994,7 @@ impl Compiler { } let entry_points = self.effective_entry_points(); crate::llvm_aot::compile_aot( - &self.decls, + self.specialized_program()?, &entry_points, output_path, prefix, @@ -851,7 +1011,10 @@ impl Compiler { jit.ir_only = true; jit.no_recursion = self.no_recursion; let entry_points = self.effective_entry_points(); - match jit.compile_and_run_multi(&self.decls, &entry_points) { + match jit.compile_and_run_multi( + self.specialized_program().expect("specialized program"), + &entry_points, + ) { Ok(_) => {} Err(e) => { println!("{}", e); @@ -867,11 +1030,8 @@ impl Compiler { let mut jit = JIT::default(); jit.print_ir = self.print_ir; jit.no_recursion = self.no_recursion; - if self.decls.decls.is_empty() { - return Err(String::from("No declarations to compile")); - } let entry_points = self.effective_entry_points(); - let (map, globals_size) = jit.compile_multi(&self.decls, &entry_points)?; + let (map, globals_size) = jit.compile_multi(self.specialized_program()?, &entry_points)?; let code_ptr = entry_points .iter() .find_map(|name| map.get(name).copied()) @@ -888,11 +1048,8 @@ impl Compiler { let mut jit = JIT::default(); jit.print_ir = self.print_ir; jit.no_recursion = self.no_recursion; - if self.decls.decls.is_empty() { - return Err(String::from("No declarations to compile")); - } let entry_points = self.effective_entry_points(); - let (map, globals_size) = jit.compile_multi(&self.decls, &entry_points)?; + let (map, globals_size) = jit.compile_multi(self.specialized_program()?, &entry_points)?; Ok((map, globals_size, jit)) } @@ -933,22 +1090,16 @@ impl Compiler { /// Compile the declarations to a VM program. pub fn compile_vm(&self) -> Result { - if self.decls.decls.is_empty() { - return Err(String::from("No declarations to compile")); - } let mut codegen = VMCodegen::new(); let entry_points = self.effective_entry_points(); - codegen.compile_multi(&self.decls, &entry_points) + codegen.compile_multi(self.specialized_program()?, &entry_points) } /// Compile to stack-based IR (for Silverfir-nano-style interpreters). pub fn compile_stack(&self) -> Result { - if self.decls.decls.is_empty() { - return Err(String::from("No declarations to compile")); - } let mut codegen = crate::stack_codegen::StackCodegen::new(); let entry_points = self.effective_entry_points(); - let mut program = codegen.compile_multi(&self.decls, &entry_points)?; + let mut program = codegen.compile_multi(self.specialized_program()?, &entry_points)?; // Inline trivial leaf functions (like cmp(a, b) -> a - b) so their // call sites become the raw ops and can participate in fusion. crate::stack_inline::inline_trivial(&mut program); @@ -972,12 +1123,9 @@ impl Compiler { /// Compile to stack IR WITHOUT the fusion optimizer (for profiling). pub fn compile_stack_unfused(&self) -> Result { - if self.decls.decls.is_empty() { - return Err(String::from("No declarations to compile")); - } let mut codegen = crate::stack_codegen::StackCodegen::new(); let entry_points = self.effective_entry_points(); - codegen.compile_multi(&self.decls, &entry_points) + codegen.compile_multi(self.specialized_program()?, &entry_points) } /// Run the code using the stack VM interpreter. @@ -1011,9 +1159,6 @@ impl Compiler { /// Compile to a backend-agnostic CompiledProgram. /// Auto-selects LLVM JIT (when available) or VM. pub fn compile_program(&self) -> Result { - if self.decls.decls.is_empty() { - return Err(String::from("No declarations to compile")); - } let entry_points = self.effective_entry_points(); #[cfg(feature = "llvm")] @@ -1021,14 +1166,14 @@ impl Compiler { let mut jit = crate::llvm_jit::LLVMJIT::new(); jit.print_ir = self.print_ir; jit.no_recursion = self.no_recursion; - let llvm_prog = jit.compile_only(&self.decls, &entry_points)?; + let llvm_prog = jit.compile_only(self.specialized_program()?, &entry_points)?; return Ok(CompiledProgram::Llvm(llvm_prog)); } #[cfg(not(feature = "llvm"))] { let mut codegen = VMCodegen::new(); - let program = codegen.compile_multi(&self.decls, &entry_points)?; + let program = codegen.compile_multi(self.specialized_program()?, &entry_points)?; let linked = LinkedProgram::from_program(&program); let vm = VM::new(); Ok(CompiledProgram::Vm { @@ -1040,10 +1185,66 @@ impl Compiler { } } +#[cfg(test)] +mod lifecycle_tests; + +#[cfg(test)] +mod safety_tests; + +#[cfg(test)] +mod assumption_tests; + #[cfg(test)] mod tests { use super::*; + #[test] + fn macro_arguments_resolve_at_each_expanded_occurrence() { + let mut compiler = Compiler::new(); + assert!(compiler.parse( + r#" + macro sum_at_scopes(value) { + let first = value + { let x = 41; first + value } + } + main() -> i32 { + let x = 1 + @sum_at_scopes(x) + } + "#, + "macro_scopes.lyte", + )); + assert!(compiler.check(), "{:?}", compiler.last_errors); + compiler.specialize().unwrap(); + let program = compiler.compile_vm().unwrap(); + let result = crate::vm::VM::new() + .call(&program, Name::str("main"), &[]) + .unwrap(); + assert_eq!(result, 42); + } + + #[test] + fn extern_preconditions_are_checked_before_publication() { + let mut compiler = Compiler::new(); + assert!(compiler.parse( + "extern fn consume(x: i32) require x >= 1\nmain { consume(1) }", + "extern.lyte" + )); + assert!(compiler.check(), "{:?}", compiler.last_errors); + compiler.specialize().unwrap(); + + for source in [ + "extern fn consume(x: i32) require missing > 0\nmain {}", + "extern fn consume(x)\nmain {}", + "interface Use { use(x) -> T }\nmain {}", + ] { + let mut compiler = Compiler::new(); + assert!(compiler.parse(source, "invalid_prototype.lyte")); + assert!(!compiler.check(), "{}", source); + assert!(compiler.checked_program().is_none()); + } + } + #[cfg(feature = "cranelift")] fn jit(code: &str) { let mut compiler = Compiler::new(); @@ -1052,7 +1253,7 @@ mod tests { compiler.parse(code.into(), &paths[0]); assert!(compiler.check()); compiler.specialize().unwrap(); - assert!(compiler.decls.decls.len() > 0); + assert!(compiler.decls().decls.len() > 0); compiler.run(); } @@ -1112,7 +1313,7 @@ mod tests { compiler.parse(code.into(), &paths[0]); assert!(compiler.check()); compiler.specialize().unwrap(); - assert!(compiler.decls.decls.len() > 0); + assert!(compiler.decls().decls.len() > 0); compiler.run_vm().expect("VM execution failed"); } @@ -1127,7 +1328,7 @@ mod tests { compiler.last_errors ); compiler.specialize().expect("specialize failed"); - assert!(compiler.decls.decls.len() > 0); + assert!(compiler.decls().decls.len() > 0); let program = compiler.compile_vm().expect("VM compile failed"); let mut vm = VM::new(); @@ -1728,18 +1929,32 @@ mod tests { assert!(compiler.check()); compiler.specialize().unwrap(); - assert!(compiler.decls.find(Name::str("sum")).is_empty()); + assert!(compiler.decls().find(Name::str("sum")).is_empty()); let main = compiler - .decls + .decls() .find(Name::str("main")) .into_iter() .find(|decl| matches!(decl, Decl::Func(_))) .expect("main should be present after specialization"); - let printed = main.pretty_print(); - assert!(printed.contains("sum$[i32]")); - assert!(printed.contains("sum$[f32]")); + let Decl::Func(main) = main else { + unreachable!() + }; + let program = compiler.specialized_program().unwrap(); + let targets: Vec<_> = main + .arena + .ids() + .filter_map(|id| { + if let Some(Reference::Instance(target)) = main.arena.reference(id) { + Some(program.instance_name(*target).to_string()) + } else { + None + } + }) + .collect(); + assert!(targets.iter().any(|name| name.contains("sum$[i32]"))); + assert!(targets.iter().any(|name| name.contains("sum$[f32]"))); } #[test] @@ -2094,8 +2309,6 @@ mod tests { // Two intermediate f32 `let`s in a block. Without the fix, each // one emits `LocalTeeF(slot), Drop` directly from translate_void. let code = r#" - var sink: [f32] - assume sink.len == 1 main { for i in 0 .. 4 { let x = i as f32 @@ -2106,7 +2319,8 @@ mod tests { "#; let mut compiler = Compiler::new(); - compiler.parse(code, "test.lyte"); + assert!(compiler.parse("var sink: [f32]\nassume sink.len == 1", "")); + assert!(compiler.parse(code, "test.lyte")); assert!(compiler.check(), "type check failed"); compiler.specialize().expect("specialize failed"); let program = compiler.compile_stack().expect("stack compile failed"); @@ -2116,12 +2330,14 @@ mod tests { .find(|f| f.name == "main") .expect("missing main"); - let bad_pair = main.ops.windows(2).enumerate().find_map(|(i, w)| { - match (&w[0], &w[1]) { + let bad_pair = main + .ops + .windows(2) + .enumerate() + .find_map(|(i, w)| match (&w[0], &w[1]) { (StackOp::LocalTeeF(_), StackOp::Drop) => Some(i), _ => None, - } - }); + }); assert!( bad_pair.is_none(), "f32 `let` compiled to LocalTeeF + Drop (int-window drop) at op {}; \ @@ -2304,52 +2520,71 @@ mod tests { // a value in "statement position" after an f32 expression. let programs: &[(&str, &str)] = &[ // Baseline: f32 let as intermediate statement. - ("f32 let intermediate", r#" + ( + "f32 let intermediate", + r#" fn make_f32() -> f32 { 1.5 } main { let x = make_f32() let y = x + x let z = y + y } - "#), + "#, + ), // Bare f32 expression as block statement (should be dropped). - ("f32 call as statement", r#" + ( + "f32 call as statement", + r#" fn make_f32() -> f32 { 1.5 } main { make_f32() } - "#), + "#, + ), // f32 if-expression in void context. - ("f32 if in void ctx", r#" + ( + "f32 if in void ctx", + r#" main { var x = 0.0f32 if true { x = 1.0 } else { x = 2.0 } } - "#), + "#, + ), // f32 assignment RHS. - ("f32 assign chain", r#" + ( + "f32 assign chain", + r#" main { var x = 0.0f32 var y = 0.0f32 x = 1.5 y = x + x } - "#), + "#, + ), ]; for (label, code) in programs { let mut compiler = Compiler::new(); - compiler.parse(code, "test.lyte"); - if !compiler.check() { - continue; - } - if compiler.specialize().is_err() { - continue; - } - let program = match compiler.compile_stack() { - Ok(p) => p, - Err(_) => continue, - }; + assert!( + compiler.parse(code, "test.lyte"), + "[{}] parse failed: {:?}", + label, + compiler.last_errors + ); + assert!( + compiler.check(), + "[{}] check failed: {:?}", + label, + compiler.last_errors + ); + compiler + .specialize() + .unwrap_or_else(|error| panic!("[{}] specialization failed: {}", label, error)); + let program = compiler + .compile_stack() + .unwrap_or_else(|error| panic!("[{}] Stack compilation failed: {}", label, error)); assert_f_window_balanced(&program, label); } } @@ -2410,7 +2645,11 @@ mod tests { assert_f_window_balanced(&program, &label); compiled += 1; } - assert!(compiled > 50, "expected to sweep >50 corpus files, got {}", compiled); + assert!( + compiled > 50, + "expected to sweep >50 corpus files, got {}", + compiled + ); } #[test] @@ -2581,7 +2820,10 @@ mod tests { vm_program.extern_funcs[0].param_types, vec![crate::vm::ExternType::Ptr, crate::vm::ExternType::I32,] ); - assert_eq!(vm_program.extern_funcs[0].ret_type, crate::vm::ExternType::Bool); + assert_eq!( + vm_program.extern_funcs[0].ret_type, + crate::vm::ExternType::Bool + ); let linked = crate::vm::LinkedProgram::from_program(&vm_program); let mut vm = crate::vm::VM::new(); @@ -2643,9 +2885,7 @@ mod tests { assert!(compiler.check(), "type check failed"); compiler.specialize().expect("specialize failed"); - let stack_program = compiler - .compile_stack() - .expect("stack VM compile failed"); + let stack_program = compiler.compile_stack().expect("stack VM compile failed"); let globals_size = stack_program.globals_size; let globals_info = @@ -2680,10 +2920,7 @@ mod tests { } STACK_SEND_CALLED.store(false, Ordering::SeqCst); - let func_idx = *stack_program - .entry_points - .get(&Name::str("main")) - .unwrap(); + let func_idx = *stack_program.entry_points.get(&Name::str("main")).unwrap(); let result = backend.call_entry(func_idx, globals.as_mut_ptr()); assert!( STACK_SEND_CALLED.load(Ordering::SeqCst), @@ -2714,7 +2951,7 @@ mod tests { let llvm_prog = { let jit = crate::llvm_jit::LLVMJIT::new(); let entry_points = vec![Name::str("main")]; - jit.compile_only(&compiler.decls, &entry_points) + jit.compile_only(compiler.specialized_program().unwrap(), &entry_points) .expect("LLVM compile failed") }; let globals_size = llvm_prog.globals_size; diff --git a/src/compiler/assumption_tests.rs b/src/compiler/assumption_tests.rs new file mode 100644 index 00000000..ccde9487 --- /dev/null +++ b/src/compiler/assumption_tests.rs @@ -0,0 +1,266 @@ +use super::*; + +#[test] +fn normalization_remaps_assumption_roots_and_shared_binding_occurrences() { + let mut errors = vec![]; + let mut lexer = Lexer::new("assume { let value = 1; value >= 0 }", ""); + lexer.next(); + let mut declarations = parse_program(&mut lexer, &mut errors); + assert!(errors.is_empty()); + let Decl::Assume { arena, cond } = &mut declarations[0] else { + panic!() + }; + let loc = arena.locs[*cond]; + // The same binding subtree appears twice. Normalization must make each + // occurrence independently checkable and discard unreachable source nodes. + *cond = arena.add(Expr::Binop(Binop::And, *cond, *cond), loc); + arena.add(Expr::Error, test_loc()); + let old_root = *cond; + let source_locations = arena.locs.clone(); + normalize_source_body(arena, std::iter::once(&mut *cond)); + assert_ne!(*cond, old_root); + assert_eq!(arena.locs[*cond], loc); + assert!(arena.locs.iter().all(|loc| source_locations.contains(loc))); + assert!(!arena.exprs.iter().any(|e| matches!(e, Expr::Error))); + + let table = DeclTable::new(declarations); + let mut checker = Checker::new(); + checker.check_decl(&table.decls[0], &table); + assert!(checker.errors.is_empty(), "{:?}", checker.errors); + let Decl::Assume { arena, cond } = &table.decls[0] else { + panic!() + }; + let body = checker.checked_body(arena); + assert_eq!(body.ty(*cond), mk_type(Type::Bool)); + assert_eq!(body.locals.len(), 2); + let Expr::Binop(Binop::And, left, right) = body[*cond] else { + panic!() + }; + assert_ne!(left, right); + for (root, expected) in [(left, LocalId(0)), (right, LocalId(1))] { + let Expr::Block(statements) = &body[root] else { + panic!() + }; + assert!(matches!(body[statements[0]], Expr::Let(..))); + assert_eq!(body.binder(statements[0]), expected); + let Expr::Binop(Binop::Geq, value, _) = body[statements[1]] else { + panic!() + }; + assert_eq!(body.reference(value), Some(&Reference::Local(expected))); + } + CheckedProgram::try_new(DeclTable::new(vec![Decl::Assume { + arena: body, + cond: *cond, + }])) + .unwrap(); +} + +#[test] +fn assumptions_check_boolean_values_and_body_declarations_at_source_locations() { + for (source, message) in [ + ( + "assume 42i32", + "assume condition must be a boolean expression", + ), + ("assume { return 42i32; true }", "return type must match"), + ("assume { let value = {}; true }", "cannot have type void"), + ( + "assume { let value = missing; true }", + "undeclared identifier", + ), + ] { + let mut compiler = Compiler::new(); + compiler.quiet = true; + assert!( + compiler.parse(source, ""), + "{}: {:?}", + source, + compiler.last_errors + ); + assert!(!compiler.check(), "{}", source); + assert!( + compiler + .last_type_errors + .iter() + .any(|e| e.message.contains(message)), + "{:?}", + compiler.last_errors + ); + assert!(compiler + .last_type_errors + .iter() + .all(|e| e.location.file == Name::str("") + && e.location.line == 1 + && e.location.col > 0)); + assert!(compiler.checked_program().is_none()); + assert!(compiler.specialize().is_err()); + } +} + +#[test] +fn assumption_specialization_preserves_local_ownership_and_concrete_targets() { + let mut compiler = Compiler::new(); + compiler.quiet = true; + assert!(compiler.parse( + r#" + interface Positive { positive(value: T) -> bool } + positive(value: i32) -> bool { value >= 0 } + positive(value: bool) -> bool { value } + nested(value: T) -> bool where Positive { positive(value) } + check_value(value: T) -> bool where Positive { + let other = nested(true) + positive(value) && other + } + sized(values: [i32; N]) -> bool { N >= 1 } + var cache: [T; 2] + assume { + let values = [1, 2] + check_value⟨i32⟩(1) && sized(values) && cache⟨i32⟩.len >= 2 + } + main {} + "#, + "" + )); + assert!(compiler.check(), "{:?}", compiler.last_errors); + let templates = compiler.checked_program().unwrap().clone(); + let (source_body, source_root) = templates + .decls + .decls + .iter() + .find_map(|d| match d { + Decl::Assume { arena, cond } => Some((arena.clone(), *cond)), + _ => None, + }) + .unwrap(); + compiler.specialize().unwrap(); + let output = compiler.specialized_program().unwrap(); + output.validate_origins(&templates).unwrap(); + let (body, root) = output + .decls + .decls + .iter() + .find_map(|d| match d { + Decl::Assume { arena, cond } => Some((arena, *cond)), + _ => None, + }) + .unwrap(); + assert_eq!(root, source_root); + assert_eq!(body.loc(root), source_body.loc(source_root)); + assert_eq!(body.ty(root), mk_type(Type::Bool)); + assert_eq!(body.locals, source_body.locals); + let mut targets = vec![]; + for id in body.ids() { + if let Some(Reference::Instance(target)) = body.reference(id) { + targets.push(output.instances[target.index()].definition); + } + } + for name in ["check_value", "sized", "cache"] { + assert!(targets.contains(&templates.decls.named_ids(Name::str(name))[0])); + } + let sized = output + .instances + .iter() + .find(|record| record.definition == templates.decls.named_ids(Name::str("sized"))[0]) + .unwrap(); + assert_eq!(sized.size_args, vec![2]); + let checked = output + .find_entry_point(Name::str("check_value$i32")) + .unwrap(); + let positive = templates + .decls + .named_ids(Name::str("positive")) + .into_iter() + .find(|id| templates.function(*id).unwrap().param_types() == vec![mk_type(Type::Int32)]) + .unwrap(); + assert!(checked + .arena + .ids() + .any(|id| match checked.arena.reference(id) { + Some(Reference::Instance(target)) => + output.instances[target.index()].definition == positive, + _ => false, + })); + assert!(body + .ids() + .any(|id| body.reference(id) == Some(&Reference::Local(LocalId(0))))); + let retained = compiler + .checked_program() + .unwrap() + .decls + .decls + .iter() + .find_map(|d| match d { + Decl::Assume { arena, .. } => Some(arena), + _ => None, + }) + .unwrap(); + assert_eq!(retained, &source_body); +} + +#[test] +fn assumption_local_facts_do_not_cross_into_other_bodies() { + for source in [ + "var limit: i32\nassume limit >= { let proof = 1; proof }\nmain(divisor: i32) -> i32 { 10 / divisor }", + "var limit: i32\nassume limit >= { let proof = 1; proof }\nassume limit >= { let divisor = 0; 10 / divisor }\nmain {}", + ] { + let mut compiler = Compiler::new(); + compiler.quiet = true; + assert!(compiler.parse(source, "")); + assert!(!compiler.check()); + assert!(compiler.last_type_errors.is_empty(), "{:?}", compiler.last_errors); + assert_eq!(compiler.last_safety_errors.len(), 1, "{:?}", compiler.last_errors); + assert!(compiler.last_safety_errors[0].message.contains("zero")); + } + + // Nonlocal facts still compose across assumptions and reach the function. + let mut compiler = Compiler::new(); + compiler.quiet = true; + assert!(compiler.parse("var divisor: i32\nvar limit: i32\nassume divisor >= 1\nassume limit >= { let proof = 0; proof }\nmain() -> i32 { 10 / divisor }", "")); + assert!(compiler.check(), "{:?}", compiler.last_errors); + compiler.specialize().unwrap(); +} + +#[test] +fn safety_errors_in_assumptions_keep_the_expression_location() { + let mut compiler = Compiler::new(); + compiler.quiet = true; + assert!(compiler.parse( + "var bound: i32\nassume bound >= 10 / 0\nmain {}", + "" + )); + let original_location = compiler + .ast + .last() + .unwrap() + .decls + .iter() + .find_map(|decl| match decl { + Decl::Assume { arena, .. } => arena.exprs.iter().enumerate().find_map(|(id, expr)| { + matches!(expr, Expr::Binop(Binop::Div, ..)).then_some(arena.locs[id]) + }), + _ => None, + }) + .unwrap(); + assert!(!compiler.check()); + assert!(compiler.last_type_errors.is_empty()); + let body = compiler + .checked_program() + .unwrap() + .decls + .decls + .iter() + .find_map(|d| match d { + Decl::Assume { arena, .. } => Some(arena), + _ => None, + }) + .unwrap(); + let division = body + .ids() + .find(|&id| matches!(body[id], Expr::Binop(Binop::Div, ..))) + .map(|id| body.loc(id)) + .unwrap(); + assert_eq!(compiler.last_safety_errors.len(), 1); + assert_eq!(compiler.last_safety_errors[0].location, division); + assert_eq!(division.file, Name::str("")); + assert_eq!(division, original_location); +} diff --git a/src/compiler/lifecycle_tests.rs b/src/compiler/lifecycle_tests.rs new file mode 100644 index 00000000..bd011f07 --- /dev/null +++ b/src/compiler/lifecycle_tests.rs @@ -0,0 +1,568 @@ +use super::*; + +fn parsed(source: &str) -> Compiler { + let mut compiler = Compiler::new(); + compiler.quiet = true; + assert!( + compiler.parse(source, "lifecycle.lyte"), + "{:?}", + compiler.last_errors + ); + compiler +} + +fn assert_no_executable(compiler: &Compiler) { + assert!(compiler.specialized_program().is_err()); + assert!(compiler.compile_vm().is_err()); + assert!(compiler.compile_stack().is_err()); + assert!(compiler.compile_stack_unfused().is_err()); + assert!(compiler.compile_program().is_err()); + #[cfg(feature = "cranelift")] + { + assert!(compiler.jit().is_err()); + assert!(compiler.jit_multi().is_err()); + } +} + +#[test] +fn partial_editor_analysis_never_authorizes_execution() { + for source in [ + "broken() -> i32 { true }", + "broken() -> i32 { 1wat }", + "broken() -> i32 { 1 / 0 }", + ] { + let mut compiler = parsed("good() -> i32 { let x = 42; x }"); + compiler.parse(source, "broken.lyte"); + assert!(!compiler.analyze()); + let snapshot = compiler.source_analysis().unwrap().clone(); + let definition = snapshot.declarations().named_ids(Name::str("good"))[0]; + assert!(snapshot.body(definition).is_some()); + compiler.set_entry_points(&["good"]); + compiler.last_errors.clear(); + compiler.last_parse_errors.clear(); + compiler.last_type_errors.clear(); + compiler.last_safety_errors.clear(); + assert!(compiler.specialize().is_err()); + assert_no_executable(&compiler); + assert!(!compiler.analyze()); + assert!(!compiler.last_errors.is_empty()); + assert!(compiler.specialize().is_err()); + assert_no_executable(&compiler); + + compiler.parse("another() {}", "another.lyte"); + assert!(compiler.source_analysis().is_none()); + // An explicitly cloned snapshot retains its own inventory; querying it + // cannot attach it to new source or grant that compiler executable input. + assert_eq!( + snapshot.declarations().function(definition).unwrap().name, + Name::str("good") + ); + assert!(snapshot.body(definition).is_some()); + assert!(compiler.specialize().is_err()); + assert_no_executable(&compiler); + } +} + +#[test] +fn root_changes_and_backends_reuse_immutable_templates() { + let mut compiler = parsed( + "var count: i32 + identity(value: T) -> T { value } + main() -> i32 { identity(1) } + other() -> i32 { identity(42) }", + ); + assert!(compiler.check(), "{:?}", compiler.last_errors); + let templates = compiler.checked_program().unwrap().clone(); + + for (root, unused, expected) in [ + ("main", "other", 1), + ("other", "main", 42), + ("main", "other", 1), + ] { + compiler.set_entry_points(&[root]); + assert_no_executable(&compiler); + compiler.specialize().unwrap(); + let concrete = compiler.specialized_program().unwrap(); + concrete + .validate_origins(compiler.checked_program().unwrap()) + .unwrap(); + assert!(concrete.decls.find(Name::str(unused)).is_empty()); + assert!(concrete.decls.find(Name::str("identity")).is_empty()); + assert!(!compiler + .checked_program() + .unwrap() + .decls + .find(Name::str(unused)) + .is_empty()); + assert!(std::ptr::eq(compiler.decls(), &concrete.decls)); + + let vm = compiler.compile_vm().unwrap(); + assert_eq!(VM::new().call(&vm, Name::str(root), &[]).unwrap(), expected); + assert_eq!( + crate::stack_vm::StackVM::new().run(&compiler.compile_stack().unwrap()), + expected + ); + assert_eq!( + crate::stack_vm::StackVM::new().run(&compiler.compile_stack_unfused().unwrap()), + expected + ); + let compiled = compiler.compile_program().unwrap(); + assert!(compiled.get_entry_point(Name::str(root)).is_some()); + #[cfg(feature = "cranelift")] + { + let (entries, _, jit) = compiler.jit_multi().unwrap(); + assert!(entries.contains_key(&Name::str(root))); + jit.free_memory(); + } + assert_eq!( + compiler.globals_info_with_offset(0), + vec![("count".into(), 0, 4, "i32".into(), false)] + ); + assert_eq!(compiler.checked_program().unwrap().decls, templates.decls); + assert!(compiler.last_errors.is_empty()); + } + + // An explicit check starts a new validation, even for unchanged source. + assert!(compiler.check()); + assert_no_executable(&compiler); + compiler.specialize().unwrap(); +} + +#[test] +fn equivalent_default_roots_preserve_specialization() { + let mut compiler = parsed("main() -> i32 { 42 }"); + assert!(compiler.check()); + compiler.specialize().unwrap(); + let output = compiler.specialized_program().unwrap() as *const SpecializedProgram; + for roots in [&["main"][..], &[][..], &["main"][..]] { + compiler.set_entry_points(roots); + assert_eq!(compiler.specialized_program().unwrap() as *const _, output); + compiler.specialize().unwrap(); + assert_eq!(compiler.specialized_program().unwrap() as *const _, output); + } +} + +#[test] +fn failed_size_specialization_can_retry_the_same_or_different_roots() { + let mut compiler = parsed( + "get(arr: [i32; N], idx: i32) -> i32 { arr[idx] } + bad() -> i32 { get([10, 20, 30, 40], 99) } + good() -> i32 { 42 }", + ); + compiler.set_entry_points(&["bad"]); + assert!(compiler.check(), "{:?}", compiler.last_errors); + let templates = compiler.checked_program().unwrap().clone(); + + for _ in 0..2 { + compiler.set_entry_points(&["bad"]); + let error = compiler.specialize().unwrap_err(); + assert!(error.contains("safety check failed"), "{}", error); + let diagnostics = compiler.last_errors.clone(); + let safety_count = compiler.last_safety_errors.len(); + assert!(safety_count > 0); + assert_no_executable(&compiler); + assert_eq!(compiler.specialize().unwrap_err(), error); + assert_eq!(compiler.last_errors, diagnostics); + assert_eq!(compiler.last_safety_errors.len(), safety_count); + + compiler.set_entry_points(&["good"]); + assert!(compiler.last_errors.is_empty()); + assert!(compiler.last_safety_errors.is_empty()); + assert_no_executable(&compiler); + compiler.specialize().unwrap(); + let program = compiler.compile_vm().unwrap(); + assert_eq!( + VM::new().call(&program, Name::str("good"), &[]).unwrap(), + 42 + ); + assert_eq!(compiler.checked_program().unwrap().decls, templates.decls); + } + + compiler.set_entry_points(&["bad"]); + assert!(compiler.specialize().is_err()); + assert!(compiler.check()); + assert!(compiler.last_errors.is_empty()); + assert!(compiler.last_safety_errors.is_empty()); + assert_no_executable(&compiler); +} + +#[test] +fn failed_concrete_call_safety_retains_analysis_and_can_retry_valid_roots() { + let mut compiler = parsed( + "bounded(x: i32, value: T) require x >= 0 {} + bad { bounded(-1, true) } + good() -> i32 { bounded⟨bool⟩(0, true); 42 }", + ); + assert!(compiler.analyze(), "{:?}", compiler.last_errors); + let templates = compiler.checked_program().unwrap().clone(); + let snapshot = compiler.source_analysis().unwrap() as *const SourceAnalysis; + + for _ in 0..2 { + compiler.set_entry_points(&["bad"]); + let failure = compiler.specialize().unwrap_err(); + assert!(failure.contains("safety check failed")); + assert_eq!(compiler.last_safety_errors.len(), 1); + let diagnostics = compiler.last_errors.clone(); + assert_eq!(diagnostics.len(), 1); + assert_no_executable(&compiler); + assert_eq!(compiler.checked_program().unwrap().decls, templates.decls); + assert_eq!(compiler.source_analysis().unwrap() as *const _, snapshot); + + compiler.last_errors.clear(); + compiler.last_safety_errors.clear(); + assert_no_executable(&compiler); + assert_eq!(compiler.specialize().unwrap_err(), failure); + assert_eq!(compiler.last_errors, diagnostics); + assert_eq!(compiler.last_safety_errors.len(), 1); + + compiler.set_entry_points(&["good"]); + assert!(compiler.last_errors.is_empty()); + assert!(compiler.last_safety_errors.is_empty()); + compiler.specialize().unwrap(); + let program = compiler.compile_vm().unwrap(); + assert_eq!( + VM::new().call(&program, Name::str("good"), &[]).unwrap(), + 42 + ); + assert_eq!(compiler.checked_program().unwrap().decls, templates.decls); + assert_eq!(compiler.source_analysis().unwrap() as *const _, snapshot); + } +} + +#[test] +fn non_safety_specialization_errors_belong_to_the_current_roots() { + let mut compiler = parsed( + "main(x: i32) -> i32 { x } + main(x: f32) -> f32 { x } + good() -> i32 { 42 }", + ); + assert!(compiler.check(), "{:?}", compiler.last_errors); + let error = compiler.specialize().unwrap_err(); + assert!(error.contains("Multiple overloads"), "{}", error); + assert_eq!(compiler.last_errors, vec![error]); + assert!(compiler.last_safety_errors.is_empty()); + let diagnostics = compiler.last_errors.clone(); + compiler.set_entry_points(&["main"]); + assert_eq!(compiler.last_errors, diagnostics); + assert_no_executable(&compiler); + compiler.set_entry_points(&["good"]); + assert!(compiler.last_errors.is_empty()); + compiler.specialize().unwrap(); + assert!(compiler.compile_vm().is_ok()); +} + +#[test] +fn parsing_invalidates_checked_and_specialized_results() { + let mut compiler = parsed("main() -> i32 { 42 }"); + assert_no_executable(&compiler); + for source in ["first() -> i32 { 1 }", "second() -> i32 { 2 }"] { + assert!(compiler.check()); + assert!(compiler.checked_program().is_some()); + if source.starts_with("second") { + compiler.specialize().unwrap(); + assert!(compiler.compile_vm().is_ok()); + } + assert!(compiler.parse(source, "extra.lyte")); + assert!(compiler.checked_program().is_none()); + assert!(compiler.specialize().is_err()); + assert_no_executable(&compiler); + } + assert!(compiler.check()); + compiler.specialize().unwrap(); + assert!(compiler.parse("other() -> i32 { missing }", "other.lyte")); + assert!(compiler.checked_program().is_none()); + assert_no_executable(&compiler); + assert!(!compiler.check()); + assert!(compiler.checked_program().is_none()); + assert!(compiler.specialize().is_err()); + assert_no_executable(&compiler); +} + +#[test] +fn parse_errors_invalidate_output_even_when_declarations_are_retained() { + let mut compiler = parsed("main() -> i32 { 42 }"); + assert!(compiler.check()); + compiler.specialize().unwrap(); + assert!(!compiler.parse("var sink: [f32]\nassume sink.len == 1", "user.lyte")); + assert!(compiler + .last_parse_errors + .iter() + .any(|error| error.message.contains("assume is only allowed"))); + assert!(compiler + .parsed_declarations() + .any(|decl| matches!(decl, Decl::Global { name, .. } if *name == Name::str("sink")))); + assert!(compiler.checked_program().is_none()); + assert_no_executable(&compiler); + assert!(!compiler.check()); + assert!(compiler.specialize().is_err()); + assert_no_executable(&compiler); + assert!(compiler.parse("other() -> i32 { 1 }", "other.lyte")); + assert!(!compiler.check()); + assert_no_executable(&compiler); +} + +#[test] +fn malformed_expressions_stop_at_parse_errors_with_either_check_all_setting() { + for check_all in [false, true] { + let mut compiler = Compiler::new(); + compiler.quiet = true; + compiler.check_all = check_all; + assert!(!compiler.parse("main() -> i32 { 1wat }", "malformed.lyte")); + let parse_errors = compiler.last_parse_errors.clone(); + let messages = compiler.last_errors.clone(); + assert!(!parse_errors.is_empty()); + + for _ in 0..2 { + // Clearing public diagnostics must not hide the invalid syntax. + compiler.last_parse_errors.clear(); + compiler.last_errors.clear(); + assert!(!compiler.check()); + assert_eq!(compiler.last_parse_errors, parse_errors); + assert_eq!(compiler.last_errors, messages); + assert!(compiler.last_type_errors.is_empty()); + assert!(compiler.last_safety_errors.is_empty()); + assert!(compiler.checked_program().is_none()); + assert!(compiler.specialize().is_err()); + assert_no_executable(&compiler); + } + + assert!(compiler.parse("helper() -> i32 { 42 }", "helper.lyte")); + assert!(!compiler.check()); + assert_eq!(compiler.last_parse_errors, parse_errors); + assert_eq!(compiler.last_errors, messages); + assert_no_executable(&compiler); + } +} + +#[test] +fn parse_errors_from_every_input_survive_parsing_and_rechecking() { + for check_all in [false, true] { + for reverse in [false, true] { + let mut compiler = Compiler::new(); + compiler.quiet = true; + compiler.check_all = check_all; + let mut inputs = [ + ("first_bad.lyte", "fn broken( {", false), + ("good.lyte", "main() -> i32 { 42 }", true), + ("second_bad.lyte", "fn unfinished( {", false), + ]; + if reverse { + inputs.reverse(); + } + for (path, source, valid) in inputs { + let previous_errors = compiler.last_parse_errors.clone(); + assert_eq!(compiler.parse(source, path), valid); + assert!(compiler.last_parse_errors.starts_with(&previous_errors)); + assert_no_executable(&compiler); + } + let parse_errors = compiler.last_parse_errors.clone(); + for path in ["first_bad.lyte", "second_bad.lyte"] { + assert!(parse_errors + .iter() + .any(|error| error.location.file == Name::str(path))); + } + for _ in 0..2 { + assert!(!compiler.check()); + assert_eq!(compiler.last_parse_errors, parse_errors); + let messages = compiler.last_errors.clone(); + assert!(messages + .iter() + .any(|error| error.contains("first_bad.lyte"))); + assert!(messages + .iter() + .any(|error| error.contains("second_bad.lyte"))); + compiler.set_entry_points(&["good"]); + assert_eq!(compiler.last_errors, messages); + assert!(compiler.checked_program().is_none()); + assert!(compiler.specialize().is_err()); + assert_no_executable(&compiler); + // Diagnostics are public views; input trees own parse failure. + compiler.last_errors.clear(); + compiler.last_parse_errors.clear(); + } + } + } +} + +#[test] +fn source_failures_survive_root_changes_and_require_rechecking_after_parse() { + for (source, safety_failure) in [ + ("bad() -> i32 { missing }", false), + ("bad() -> i32 { let values = [1, 2]; values[9] }", true), + ] { + let mut compiler = parsed(source); + assert!(compiler.parse("good() -> i32 { 42 }", "good.lyte")); + assert!(!compiler.check()); + assert_eq!(compiler.checked_program().is_some(), safety_failure); + assert_eq!(!compiler.last_safety_errors.is_empty(), safety_failure); + assert_eq!(!compiler.last_type_errors.is_empty(), !safety_failure); + let errors = compiler.last_errors.clone(); + compiler.set_entry_points(&["good"]); + assert_eq!(compiler.last_errors, errors); + assert!(compiler.specialize().is_err()); + compiler.last_errors.clear(); + compiler.last_type_errors.clear(); + compiler.last_safety_errors.clear(); + assert!(compiler.specialize().is_err()); + assert_no_executable(&compiler); + + assert!(compiler.parse("helper() -> i32 { 1 }", "helper.lyte")); + assert!(compiler.checked_program().is_none()); + assert!(compiler.last_errors.is_empty()); + assert!(compiler.last_type_errors.is_empty()); + assert!(compiler.last_safety_errors.is_empty()); + assert_no_executable(&compiler); + assert!(!compiler.check()); + assert_eq!(compiler.last_errors, errors); + assert!(compiler.specialize().is_err()); + } +} + +#[test] +fn parsing_clears_specialization_diagnostics_and_requires_checking() { + let mut compiler = parsed( + "get(arr: [i32; N], idx: i32) -> i32 { arr[idx] } + main() -> i32 { get([1, 2], 9) }", + ); + assert!(compiler.check()); + assert!(compiler.specialize().is_err()); + assert!(!compiler.last_safety_errors.is_empty()); + assert!(compiler.parse("good() -> i32 { 42 }", "good.lyte")); + assert!(compiler.last_errors.is_empty()); + assert!(compiler.last_safety_errors.is_empty()); + assert!(compiler.checked_program().is_none()); + compiler.set_entry_points(&["good"]); + assert!(compiler.specialize().is_err()); + assert_no_executable(&compiler); + assert!(compiler.check()); + compiler.specialize().unwrap(); +} + +#[test] +fn validation_option_changes_require_checking_before_specialization_or_codegen() { + for original in [false, true] { + for specialized in [false, true] { + let mut compiler = parsed("main() -> i32 { 42 }"); + compiler.no_recursion = original; + assert!(compiler.check()); + if specialized { + compiler.specialize().unwrap(); + } + let templates = compiler.checked_program().unwrap().clone(); + compiler.no_recursion = !original; + assert!(compiler.specialize().unwrap_err().contains("call check()")); + assert_no_executable(&compiler); + assert_eq!(compiler.checked_program().unwrap().decls, templates.decls); + + // Reuse is keyed by option values, not by writes to public fields. + compiler.no_recursion = original; + assert_eq!(compiler.specialized_program().is_ok(), specialized); + compiler.no_recursion = !original; + assert!(compiler.check()); + assert_no_executable(&compiler); + compiler.specialize().unwrap(); + assert!(compiler.compile_vm().is_ok()); + } + } +} + +#[test] +fn enabling_no_recursion_rejects_previously_validated_recursion() { + let mut compiler = parsed( + "recurse(x: i32) -> i32 { if x == 0 { 0 } else { recurse(x - 1) } } + main() -> i32 { recurse(2) }", + ); + assert!(compiler.check()); + compiler.specialize().unwrap(); + compiler.no_recursion = true; + assert_no_executable(&compiler); + assert!(!compiler.check()); + assert!(compiler + .last_errors + .iter() + .any(|error| error.contains("--no-recursion"))); + assert!(compiler.checked_program().is_some()); + assert!(compiler.specialize().is_err()); + assert_no_executable(&compiler); + compiler.no_recursion = false; + assert!(compiler.specialize().unwrap_err().contains("call check()")); + assert!(compiler.check()); + assert!(compiler.last_safety_errors.is_empty()); + compiler.specialize().unwrap(); +} + +#[test] +fn ffi_malformed_expressions_report_parse_errors_without_an_ice() { + use crate::ffi::*; + use std::ffi::{CStr, CString}; + for check_first in [false, true] { + unsafe { + let compiler = lyte_compiler_new(std::ptr::null(), 0); + let source = CString::new("main() -> i32 { 1wat }").unwrap(); + let path = CString::new("malformed.lyte").unwrap(); + assert!(!lyte_compiler_add_source( + compiler, + source.as_ptr(), + path.as_ptr() + )); + let parse_error = CStr::from_ptr(lyte_compiler_get_error(compiler)) + .to_str() + .unwrap() + .to_string(); + assert!(parse_error.contains("malformed.lyte"), "{}", parse_error); + + if check_first { + assert!(!lyte_compiler_check(compiler)); + assert!(!lyte_compiler_had_ice(compiler)); + } + assert!(lyte_compiler_compile(compiler).is_null()); + assert!(!lyte_compiler_had_ice(compiler)); + assert_eq!( + CStr::from_ptr(lyte_compiler_get_error(compiler)) + .to_str() + .unwrap(), + parse_error + ); + assert!(!lyte_compiler_check(compiler)); + assert!(!lyte_compiler_had_ice(compiler)); + assert_eq!( + CStr::from_ptr(lyte_compiler_get_error(compiler)) + .to_str() + .unwrap(), + parse_error + ); + lyte_compiler_free(compiler); + } + } +} + +#[test] +fn ffi_compile_rejects_parse_errors_in_an_earlier_input() { + use crate::ffi::*; + use std::ffi::{CStr, CString}; + unsafe { + let compiler = lyte_compiler_new(std::ptr::null(), 0); + let bad = CString::new("fn broken( {").unwrap(); + let good = CString::new("main() -> i32 { 42 }").unwrap(); + let bad_path = CString::new("bad.lyte").unwrap(); + let good_path = CString::new("good.lyte").unwrap(); + assert!(!lyte_compiler_add_source( + compiler, + bad.as_ptr(), + bad_path.as_ptr() + )); + assert!(lyte_compiler_add_source( + compiler, + good.as_ptr(), + good_path.as_ptr() + )); + assert!(!lyte_compiler_check(compiler)); + assert!(lyte_compiler_compile(compiler).is_null()); + assert!(!lyte_compiler_had_ice(compiler)); + let error = CStr::from_ptr(lyte_compiler_get_error(compiler)) + .to_str() + .unwrap(); + assert!(error.contains("bad.lyte"), "{}", error); + lyte_compiler_free(compiler); + } +} diff --git a/src/compiler/safety_tests.rs b/src/compiler/safety_tests.rs new file mode 100644 index 00000000..d3ac0211 --- /dev/null +++ b/src/compiler/safety_tests.rs @@ -0,0 +1,244 @@ +use super::*; + +fn checked(source: &str) -> Compiler { + let mut compiler = Compiler::new(); + compiler.quiet = true; + assert!(compiler.parse(source, "concrete_safety.lyte")); + assert!(compiler.check(), "{}\n{:?}", source, compiler.last_errors); + compiler +} + +fn assert_requirement_failure(compiler: &mut Compiler, clause: &str) { + let error = compiler.specialize().unwrap_err(); + assert!(error.contains("safety check failed"), "{}", error); + assert!(compiler.specialized_program().is_err()); + assert!(compiler.checked_program().is_some()); + assert_eq!( + compiler.last_safety_errors.len(), + 1, + "{:?}", + compiler.last_errors + ); + assert_eq!(compiler.last_errors.len(), 1, "{:?}", compiler.last_errors); + assert!(compiler.last_safety_errors[0].message.contains(clause)); +} + +#[test] +fn generic_requirements_are_checked_in_ordinary_callers_and_type_applications() { + for target in ["bounded", "bounded⟨bool⟩"] { + for value in [-1, 0] { + let mut compiler = checked(&format!( + "bounded(x: i32, value: T) require x >= 0 {{}}\n\ + main {{ {target}({value}, true) }}" + )); + if value < 0 { + assert_requirement_failure(&mut compiler, "`x >= 0`"); + assert_eq!(compiler.last_safety_errors[0].location.line, 2); + } else { + compiler.specialize().unwrap(); + assert!(compiler.last_errors.is_empty()); + compiler.specialized_program().unwrap().validate().unwrap(); + } + } + } +} + +#[test] +fn ordinary_wrappers_must_establish_generic_call_requirements() { + for body in ["bounded(-1, true)", "bounded(x, true)"] { + let mut compiler = checked(&format!( + "bounded(x: i32, value: T) require x >= 0 {{}}\n\ + wrapper(x: i32) {{ {body} }}\n\ + main {{ wrapper(0) }}" + )); + assert_requirement_failure(&mut compiler, "`x >= 0`"); + assert_eq!(compiler.last_safety_errors[0].location.line, 2); + } + for wrapper in [ + "wrapper(x: i32) require x >= 0 { bounded(x, true) }", + "wrapper(x: i32) { if x >= 0 { bounded(x, true) } }", + ] { + let mut compiler = checked(&format!( + "bounded(x: i32, value: T) require x >= 0 {{}}\n\ + {wrapper}\nmain {{ wrapper(0) }}" + )); + compiler.specialize().unwrap(); + } +} + +#[test] +fn concrete_call_requirements_preserve_array_and_reference_coercions() { + for (definition, call) in [ + ( + "bounded(x: i32, values: [T]) require x >= 0 {}", + "bounded(x, [1, 2])", + ), + ( + "bounded(x: &i32, value: T) require x >= 0 {}", + "bounded(x, true)", + ), + ( + "bounded(x: i32, values: [T]) require x >= 0 require x < values.len {}", + "bounded⟨i32⟩(x, [1, 2])", + ), + ] { + for value in [-1, 0] { + let mut compiler = + checked(&format!("{definition}\nmain {{ var x = {value}; {call} }}")); + if value < 0 { + assert_requirement_failure(&mut compiler, "`x >= 0`"); + } else { + compiler.specialize().unwrap(); + } + } + } +} + +#[test] +fn selected_interface_implementations_and_externs_keep_their_requirements() { + for prefix in ["", "extern fn "] { + let body = if prefix.is_empty() { " {}" } else { "" }; + for value in [-1, 0] { + let mut compiler = checked(&format!( + "interface Bounded {{ bounded(x: i32, value: T) }}\n\ + {prefix}bounded(x: i32, value: bool) require x >= 0{body}\n\ + wrapper(x: i32, value: T) where Bounded {{ bounded({value}, value) }}\n\ + main {{ wrapper(0, true) }}" + )); + if value < 0 { + assert_requirement_failure(&mut compiler, "`x >= 0`"); + } else { + compiler.specialize().unwrap(); + } + } + } +} + +#[test] +fn local_and_global_function_values_remain_indirect() { + // This deliberately characterizes the existing proof limitation: a stored + // function value does not transport its target's require clauses. These + // calls must not be confused with the same-named direct functions. + let mut compiler = checked( + "bounded(x: i32, value: T) require x >= 0 {} + var callback: (i32, bool) -> void + callback(x: i32, value: bool) require x >= 0 {} + main { + callback = bounded⟨bool⟩ + callback(-1, true) + let bounded = bounded⟨bool⟩ + bounded(-1, true) + }", + ); + compiler.specialize().unwrap(); + let program = compiler.specialized_program().unwrap(); + let main = program + .function_instance(program.instance_for_entry(Name::str("main")).unwrap()) + .unwrap(); + let mut local_calls = 0; + let mut global_calls = 0; + for id in main.arena.ids() { + if let Expr::Call(callee, _) = &main.arena[id] { + match main.arena.reference(*callee) { + Some(Reference::Local(_)) => local_calls += 1, + Some(Reference::Instance(target)) => { + assert!(matches!( + program.instance(*target), + CheckedDecl::Global { .. } + )); + global_calls += 1; + } + _ => panic!("unexpected call target"), + } + } + } + assert_eq!((local_calls, global_calls), (1, 1)); +} + +#[test] +fn equivalent_requirements_in_multiple_instances_report_once() { + let mut compiler = checked( + "bounded(x: i32, value: U) require x >= 0 {} + wrapper(value: T) { bounded(-1, value) } + main { wrapper(true); wrapper(1) }", + ); + assert_requirement_failure(&mut compiler, "`x >= 0`"); + assert_eq!(compiler.last_safety_errors[0].location.line, 2); + let diagnostics = compiler.last_errors.clone(); + assert_requirement_failure(&mut compiler, "`x >= 0`"); + assert_eq!(compiler.last_errors, diagnostics); +} + +#[test] +fn separate_call_sites_keep_separate_requirement_diagnostics() { + let mut compiler = checked( + "bounded(x: i32, value: T) require x >= 0 {} + main { bounded(-1, true); bounded(-2, true) }", + ); + assert!(compiler.specialize().is_err()); + assert_eq!(compiler.last_safety_errors.len(), 2); + assert_ne!( + compiler.last_safety_errors[0].location, + compiler.last_safety_errors[1].location + ); +} + +#[test] +fn different_concrete_requirements_at_one_call_site_remain_distinct() { + let mut compiler = checked( + "bounded(xs: [i32; N], idx: i32) require idx < N {} + wrapper(xs: [i32; N]) { bounded(xs, 99) } + main { wrapper([1, 2]); wrapper([1, 2, 3]) }", + ); + assert!(compiler.specialize().is_err()); + assert_eq!( + compiler.last_safety_errors.len(), + 2, + "{:?}", + compiler.last_errors + ); + assert_eq!( + compiler.last_safety_errors[0].location, + compiler.last_safety_errors[1].location + ); + for clause in ["`idx < 2`", "`idx < 3`"] { + assert!( + compiler + .last_errors + .iter() + .any(|error| error.contains(clause)), + "{:?}", + compiler.last_errors + ); + } +} + +#[test] +fn unsupported_requirements_still_fail_conservatively_after_specialization() { + // The call is valid, but equality is outside the existing call-site proof + // grammar. Concrete integration must not silently accept an unproved clause + // or expand the prover to make this witness pass. + let mut compiler = checked( + "bounded(x: i32, value: T) require x == 0 {} + main { bounded(0, true) }", + ); + assert_requirement_failure(&mut compiler, "`x == 0`"); +} + +#[test] +fn template_safety_still_checks_unreachable_definitions_once() { + let mut compiler = Compiler::new(); + compiler.quiet = true; + assert!(compiler.parse( + "bounded(x: i32) require x >= 0 {} + unused { bounded(-1) } + main {}", + "unreachable.lyte", + )); + assert!(!compiler.check()); + let diagnostics = compiler.last_errors.clone(); + assert_eq!(diagnostics.len(), 1); + assert!(compiler.specialize().is_err()); + assert_eq!(compiler.last_errors, diagnostics); + assert_eq!(compiler.last_safety_errors.len(), 1); +} diff --git a/src/copy_elision.rs b/src/copy_elision.rs index 3b4597a7..3593890d 100644 --- a/src/copy_elision.rs +++ b/src/copy_elision.rs @@ -18,10 +18,10 @@ //! the source is observationally identical to copying it, and the backend is //! free to skip the copy. `elidable_let_copies` finds those bindings. -use crate::decl::FuncDecl; -use crate::defs::{Binop, ExprID, Name}; -use crate::expr::Expr; +use crate::checked::{CheckedFunction, LocalId, Reference}; +use crate::defs::{Binop, ExprID}; use crate::types::{Type, TypeID}; +use crate::Expr; use std::collections::HashSet; /// Types that `let` binds by value, and so must copy out of the initializer's @@ -39,7 +39,7 @@ pub fn is_value_aggregate(ty: &TypeID) -> bool { /// /// A copy is elidable when nothing that runs while the binding is live can /// observe the difference — see [`live_range_is_read_only`]. -pub fn elidable_let_copies(decl: &FuncDecl) -> HashSet { +pub fn elidable_let_copies(decl: &CheckedFunction) -> HashSet { let mut elidable = HashSet::new(); if let Some(body) = decl.body { scan_blocks(body, decl, &mut elidable); @@ -47,29 +47,26 @@ pub fn elidable_let_copies(decl: &FuncDecl) -> HashSet { elidable } -fn scan_blocks(id: ExprID, decl: &FuncDecl, elidable: &mut HashSet) { - if let Expr::Block(stmts) = &decl.arena.exprs[id] { +fn scan_blocks(id: ExprID, decl: &CheckedFunction, elidable: &mut HashSet) { + if let Expr::Block(stmts) = &decl.arena[id] { for (i, &stmt) in stmts.iter().enumerate() { - if !matches!(&decl.arena.exprs[stmt], Expr::Let(..)) { + if !matches!(decl.arena[stmt], Expr::Let(..)) { continue; } - if !is_value_aggregate(&decl.types[stmt]) { + let local = decl.arena.binder(stmt); + if !is_value_aggregate(&decl.arena.local(local).ty) { continue; } - // A `let` in tail position is the block's value, so the binding - // outlives the block. Never elide those. + // This analysis requires a following sequence and keeps tail + // declarations conservative. if i + 1 >= stmts.len() { continue; } let rest = &stmts[i + 1..]; - let name = match &decl.arena.exprs[stmt] { - Expr::Let(name, _, _) => *name, - _ => unreachable!(), - }; // The block's value escapes it, so the binding must not reach the // final statement. let last = *stmts.last().unwrap(); - if mentions(last, name, decl) { + if mentions(last, local, decl) { continue; } if rest.iter().all(|&s| live_range_is_read_only(s, decl)) { @@ -78,7 +75,7 @@ fn scan_blocks(id: ExprID, decl: &FuncDecl, elidable: &mut HashSet) { } } - for sub in decl.arena.exprs[id].subexprs() { + for sub in decl.arena[id].subexprs() { scan_blocks(sub, decl, elidable); } } @@ -99,16 +96,16 @@ fn scan_blocks(id: ExprID, decl: &FuncDecl, elidable: &mut HashSet) { /// /// What's left — indexing, field reads, arithmetic, control flow, fresh /// `let`/`var` bindings — only reads. -fn live_range_is_read_only(id: ExprID, decl: &FuncDecl) -> bool { - match &decl.arena.exprs[id] { +fn live_range_is_read_only(id: ExprID, decl: &CheckedFunction) -> bool { + match &decl.arena[id] { Expr::Call(_, _) | Expr::Macro(_, _) | Expr::Lambda { .. } | Expr::Return(_) | Expr::Arena(_) => false, Expr::Binop(Binop::Assign, lhs, rhs) => { - matches!(&decl.arena.exprs[*lhs], Expr::Id(_)) - && !decl.types[*lhs].is_ptr() + matches!(&decl.arena[*lhs], Expr::Id(_)) + && !decl.arena.ty(*lhs).is_ptr() && live_range_is_read_only(*rhs, decl) } expr => expr @@ -119,13 +116,11 @@ fn live_range_is_read_only(id: ExprID, decl: &FuncDecl) -> bool { } /// True if `name` is referenced anywhere in this subtree. -fn mentions(id: ExprID, name: Name, decl: &FuncDecl) -> bool { - if let Expr::Id(n) = &decl.arena.exprs[id] { - if *n == name { - return true; - } +fn mentions(id: ExprID, name: LocalId, decl: &CheckedFunction) -> bool { + if decl.arena.reference(id) == Some(&Reference::Local(name)) { + return true; } - decl.arena.exprs[id] + decl.arena[id] .subexprs() .into_iter() .any(|sub| mentions(sub, name, decl)) diff --git a/src/decl.rs b/src/decl.rs index cc04972d..59635751 100644 --- a/src/decl.rs +++ b/src/decl.rs @@ -21,17 +21,6 @@ pub struct InterfaceConstraint { pub typevars: Vec, } -/// A variable captured by a closure. -/// -/// Closures always capture by address: the closure struct stores a pointer to -/// the variable's stack slot. For `var` bindings the slot already exists; for -/// `let` bindings the JIT allocates a fresh slot and copies the value there. -#[derive(Clone, Debug, Eq, PartialEq, Hash)] -pub struct ClosureVar { - pub name: Name, - pub ty: TypeID, -} - /// Function declaration. #[derive(Clone, Debug, Eq, PartialEq, Hash)] pub struct FuncDecl { @@ -67,16 +56,6 @@ pub struct FuncDecl { /// Expression arena for the function body. pub arena: ExprArena, - /// Solved types from the type checker. - pub types: Vec, - - /// Variables captured from the enclosing scope (non-empty for closures). - /// - /// At runtime the JIT passes a `closure_ptr` pointing to a contiguous - /// array of `i64` slots, one per entry here. Each slot holds the address - /// of the captured variable's storage. - pub closure_vars: Vec, - /// True if this is an extern function provided by the host. /// Extern functions have no body and are called indirectly through /// a {fn_ptr, context} pair stored in the globals buffer. @@ -84,37 +63,6 @@ pub struct FuncDecl { } impl FuncDecl { - /// Every name mentioned anywhere inside a lambda body in this function. - /// - /// A `var` that a lambda captures is shared by address, so it has to live - /// in memory. A backend that would otherwise keep a scalar in a register - /// must skip that promotion for these names — de-promoting later, at the - /// point of capture, puts the spill wherever the lambda happens to sit, - /// and inside a loop that spill re-runs every iteration and clobbers the - /// variable with a stale value. - pub fn names_referenced_in_lambdas(&self) -> std::collections::HashSet { - fn collect( - expr: ExprID, - arena: &ExprArena, - result: &mut std::collections::HashSet, - ) { - if let Expr::Id(name) = &arena[expr] { - result.insert(*name); - } - for sub in arena[expr].subexprs() { - collect(sub, arena, result); - } - } - - let mut result = std::collections::HashSet::new(); - for expr in &self.arena.exprs { - if let Expr::Lambda { body, .. } = expr { - collect(*body, &self.arena, &mut result); - } - } - result - } - /// Get the types of the function parameters. pub fn param_types(&self) -> Vec { self.params @@ -154,7 +102,9 @@ impl FuncDecl { .map(|(p, a)| (p.name, *a)) .collect(); - let body = mac.body.expect("macro must have a body"); + let body = mac + .body + .ok_or_else(|| (loc, format!("macro '{}' has no body", name)))?; let new_body = copy_expr(body, &mac.arena, &mut self.arena, &subst); self.arena.exprs[i] = self.arena.exprs[new_body].clone(); @@ -185,7 +135,12 @@ impl StructDecl { None } - pub fn field_offset(&self, name: &Name, decls: &DeclTable, inst: &Instance) -> i32 { + pub fn field_offset( + &self, + name: &Name, + decls: &DeclarationList, + inst: &Instance, + ) -> i32 { let mut off = 0; for field in &self.fields { if field.name == *name { @@ -200,18 +155,18 @@ impl StructDecl { /// Provides a set of functions that some type variables /// must satisfy. #[derive(Clone, Debug, Eq, PartialEq, Hash)] -pub struct Interface { +pub struct Interface { pub name: Name, pub typevars: Vec, - pub funcs: Vec, + pub funcs: Vec, pub loc: Loc, } /// Top-level declaration. #[derive(Clone, Debug, Eq, PartialEq, Hash)] -pub enum Decl { - Func(FuncDecl), - Macro(FuncDecl), +pub enum Decl { + Func(F), + Macro(F), Struct(StructDecl), Enum { name: Name, @@ -222,18 +177,18 @@ pub enum Decl { typevars: Vec, ty: TypeID, }, - Interface(Interface), + Interface(Interface), Const { name: Name, value: i64, }, Assume { - arena: ExprArena, + arena: F::Arena, cond: ExprID, }, } -impl Decl { +impl Decl { pub fn find_field(&self, name: &Name) -> Option { if let Decl::Struct(st) = self { st.find_field(name) @@ -252,11 +207,11 @@ pub fn find_field(fields: &[Field], name: Name) -> Option { None } -impl Decl { +impl Decl { pub fn name(&self) -> Name { match self { - Decl::Func(FuncDecl { name, .. }) => *name, - Decl::Macro(FuncDecl { name, .. }) => *name, + Decl::Func(function) => function.name(), + Decl::Macro(function) => function.name(), Decl::Struct(StructDecl { name, .. }) => *name, Decl::Enum { name, .. } => *name, Decl::Global { name, .. } => *name, @@ -265,7 +220,9 @@ impl Decl { Decl::Assume { .. } => Name::str("__assume"), } } +} +impl Decl { /// Pretty-print a declaration in lyte syntax. /// /// This method formats a declaration as it would appear in lyte source code, @@ -469,8 +426,6 @@ mod tests { requires: vec![], loc: test_loc(), arena: ExprArena::new(), - types: vec![], - closure_vars: vec![], is_extern: false, }; @@ -498,8 +453,6 @@ mod tests { requires: vec![], loc: test_loc(), arena, - types: vec![], - closure_vars: vec![], is_extern: false, }; @@ -557,8 +510,6 @@ mod tests { requires: vec![], loc: test_loc(), arena: ExprArena::new(), - types: vec![], - closure_vars: vec![], is_extern: false, }], loc: test_loc(), @@ -618,8 +569,6 @@ mod tests { requires: vec![], loc: test_loc(), arena, - types: vec![], - closure_vars: vec![], is_extern: false, }; @@ -628,3 +577,29 @@ mod tests { assert_eq!(output, "increment(x: i32) → i32 {\n x + 1\n}"); } } + +/// Information shared by source signatures and checked definitions. Body access +/// deliberately is not part of this interface. +pub trait FunctionInfo: Clone + std::fmt::Debug + Eq + std::hash::Hash { + type Arena: Clone + std::fmt::Debug + Eq + std::hash::Hash; + fn name(&self) -> Name; + fn ty(&self) -> TypeID; + /// Checked functions always have signatures; source declarations may still + /// be missing parameter annotations during recovery. + fn try_ty(&self) -> Option { + Some(self.ty()) + } +} + +impl FunctionInfo for FuncDecl { + type Arena = ExprArena; + fn name(&self) -> Name { + self.name + } + fn ty(&self) -> TypeID { + self.ty() + } + fn try_ty(&self) -> Option { + self.annotated_ty() + } +} diff --git a/src/decl_table.rs b/src/decl_table.rs index 2a885294..8f0f87b7 100644 --- a/src/decl_table.rs +++ b/src/decl_table.rs @@ -1,47 +1,318 @@ use crate::*; -use std::cmp::Ordering; +use std::ops::Deref; use superslice::Ext; -/// Table of top level declarations. -/// -/// Currently we're just using a sorted Vec. -/// In the future we could use a hash table. +/// Declaration storage sorted for source-name candidate collection. Semantic +/// identities are allocated once before sorting, independently of coordinates. #[derive(Clone, Debug, Eq, PartialEq, Hash)] -pub struct DeclTable { - /// All declarations, sorted by name. - pub decls: Vec, +pub struct DeclTable { + list: DeclarationList, + definitions: Vec, + ids: Vec, + member_ids: Vec>, +} - /// For quickly finding enums which contain a case. - /// This is for resolving .enum_case expressions. - pub enum_cases: Vec<(Name, usize)>, +#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)] +enum DefinitionLocation { + Declaration(usize), + Member(usize, usize), + Missing, } -/// Comparison function to sort declarations by name. -fn decl_cmp(a: &Decl, b: &Decl) -> Ordering { - a.name().cmp(&b.name()) +#[derive(Clone, Debug)] +pub struct DeclRecord { + pub definition: DefId, + pub declaration: Decl, + pub members: Vec, +} + +impl DeclTable { + pub fn new(decls: Vec>) -> Self { + let mut next = decls.len() as u32; + let records = decls + .into_iter() + .enumerate() + .map(|(index, declaration)| { + let members = match &declaration { + Decl::Interface(interface) => interface + .funcs + .iter() + .map(|_| { + let id = DefId(next); + next += 1; + id + }) + .collect(), + _ => Vec::new(), + }; + DeclRecord { + definition: DefId(index as u32), + declaration, + members, + } + }) + .collect(); + Self::from_records(records) + } + pub fn from_records(mut records: Vec>) -> Self { + records.sort_by_key(|record| record.declaration.name()); + let count = records + .iter() + .flat_map(|record| { + std::iter::once(record.definition).chain(record.members.iter().copied()) + }) + .map(|id| id.index() + 1) + .max() + .unwrap_or(0); + let mut definitions = vec![DefinitionLocation::Missing; count]; + let mut ids = Vec::new(); + let mut member_ids = Vec::new(); + let mut decls = Vec::new(); + for (index, record) in records.into_iter().enumerate() { + let member_count = match &record.declaration { + Decl::Interface(interface) => interface.funcs.len(), + _ => 0, + }; + assert_eq!( + record.members.len(), + member_count, + "invalid member inventory" + ); + assert_eq!( + definitions[record.definition.index()], + DefinitionLocation::Missing, + "duplicate declaration identity" + ); + definitions[record.definition.index()] = DefinitionLocation::Declaration(index); + for (member, id) in record.members.iter().enumerate() { + assert_eq!( + definitions[id.index()], + DefinitionLocation::Missing, + "duplicate member identity" + ); + definitions[id.index()] = DefinitionLocation::Member(index, member); + } + ids.push(record.definition); + member_ids.push(record.members); + decls.push(record.declaration); + } + Self { + list: DeclarationList::from_sorted(decls), + definitions, + ids, + member_ids, + } + } + pub fn definition_count(&self) -> usize { + self.definitions.len() + } + pub fn id_at(&self, index: usize) -> DefId { + self.ids[index] + } + pub fn named_ids(&self, name: Name) -> Vec { + let range = self.decls.equal_range_by(|decl| decl.name().cmp(&name)); + self.ids[range].to_vec() + } + pub fn definition(&self, id: DefId) -> Option<&Decl> { + match self.definitions.get(id.index())? { + DefinitionLocation::Declaration(index) => Some(&self.decls[*index]), + _ => None, + } + } + pub fn get(&self, id: DefId) -> Option<&Decl> { + self.definition(id) + } + pub fn function(&self, id: DefId) -> Option<&F> { + match *self.definitions.get(id.index())? { + DefinitionLocation::Declaration(index) => match &self.decls[index] { + Decl::Func(function) | Decl::Macro(function) => Some(function), + _ => None, + }, + DefinitionLocation::Member(index, member) => match &self.decls[index] { + Decl::Interface(interface) => interface.funcs.get(member), + _ => None, + }, + DefinitionLocation::Missing => None, + } + } + pub fn signature(&self, id: DefId) -> Option { + if let Some(function) = self.function(id) { + function.try_ty() + } else { + self.definition(id).map(Decl::ty) + } + } + pub fn interface_members(&self, id: DefId) -> &[DefId] { + match self.definitions.get(id.index()) { + Some(DefinitionLocation::Declaration(index)) => &self.member_ids[*index], + _ => &[], + } + } + pub fn records(&self) -> impl Iterator> + '_ { + self.decls + .iter() + .enumerate() + .map(move |(index, declaration)| DeclRecord { + definition: self.ids[index], + declaration: declaration.clone(), + members: self.member_ids[index].clone(), + }) + } + pub fn map_bodies( + &self, + mut map: impl FnMut(DefId, F) -> G, + mut map_assume: impl FnMut(F::Arena, ExprID) -> G::Arena, + ) -> DeclTable { + DeclTable::from_records( + self.records() + .map(|record| { + let declaration = match record.declaration { + Decl::Func(function) => Decl::Func(map(record.definition, function)), + Decl::Macro(function) => Decl::Macro(map(record.definition, function)), + Decl::Interface(interface) => Decl::Interface(Interface { + name: interface.name, + typevars: interface.typevars, + loc: interface.loc, + funcs: interface + .funcs + .into_iter() + .zip(record.members.iter()) + .map(|(function, id)| map(*id, function)) + .collect(), + }), + Decl::Struct(value) => Decl::Struct(value), + Decl::Enum { name, cases } => Decl::Enum { name, cases }, + Decl::Global { name, typevars, ty } => Decl::Global { name, typevars, ty }, + Decl::Const { name, value } => Decl::Const { name, value }, + Decl::Assume { arena, cond } => Decl::Assume { + arena: map_assume(arena, cond), + cond, + }, + }; + DeclRecord { + definition: record.definition, + declaration, + members: record.members, + } + }) + .collect(), + ) + } + pub fn interface_requirement( + &self, + id: RequirementId, + name: Name, + type_args: Vec, + ) -> Option { + let interface_id = self + .named_ids(name) + .into_iter() + .find(|id| matches!(self.definition(*id), Some(Decl::Interface(_))))?; + let Decl::Interface(interface) = self.definition(interface_id)? else { + unreachable!() + }; + let instance: Instance = interface + .typevars + .iter() + .zip(&type_args) + .map(|(name, ty)| (typevar(name), *ty)) + .collect(); + let members = interface + .funcs + .iter() + .zip(self.interface_members(interface_id)) + .map(|(function, definition)| { + Some(InterfaceMember { + definition: *definition, + signature: function.try_ty()?.subst(&instance), + candidates: self + .named_ids(function.name()) + .into_iter() + .filter(|id| { + matches!( + self.definition(*id), + Some(Decl::Func(_) | Decl::Global { .. }) + ) + }) + .collect(), + }) + }) + .collect::>>()?; + Some(InterfaceRequirement { + id, + interface: interface_id, + type_args, + members, + }) + } + + pub fn interface_alternative(&self, name: Name, type_args: Vec) -> AltInterface { + let requirement = self.interface_requirement(RequirementId(0), name, type_args.clone()); + AltInterface { + interface: name, + typevars: type_args, + members: requirement.map(|requirement| requirement.members), + } + } } impl DeclTable { - pub fn new(mut decls: Vec) -> Self { - decls.sort_by(decl_cmp); + /// Returns all alternatives for a declaration name. + pub fn alts(&self, name: Name) -> Vec { + let sl = self.find(name); + let mut alts = vec![]; - let mut enum_cases = vec![]; + for d in sl { + match d { + Decl::Func(function) => { + let Some(ty) = function.annotated_ty() else { + continue; + }; + let mut interfaces = vec![]; + for c in &function.constraints { + interfaces.push(self.interface_alternative( + c.interface_name, + c.typevars.iter().map(|name| typevar(name)).collect(), + )) + } - for (i, decl) in decls.iter().enumerate() { - if let Decl::Enum { cases, .. } = decl { - for case in cases { - enum_cases.push((*case, i)) + alts.push(Alt { ty, interfaces }); + } + Decl::Global { .. } => { + alts.push(Alt { + ty: d.ty(), + interfaces: vec![], + }); } + _ => (), } } - enum_cases.sort(); + alts + } +} +/// Shared declaration storage and type-layout lookup. This list deliberately +/// provides no source-definition or concrete-instance identity APIs; those live +/// on the owning phase's program/table. +#[derive(Clone, Debug, Eq, PartialEq, Hash)] +pub struct DeclarationList { + pub decls: Vec>, + pub enum_cases: Vec<(Name, usize)>, +} +impl DeclarationList { + pub(crate) fn from_sorted(decls: Vec>) -> Self { + let mut enum_cases = Vec::new(); + for (index, declaration) in decls.iter().enumerate() { + if let Decl::Enum { cases, .. } = declaration { + enum_cases.extend(cases.iter().map(|name| (*name, index))); + } + } + enum_cases.sort(); Self { decls, enum_cases } } - /// Returns a slice of all decls which match name. - pub fn find(&self, name: Name) -> &[Decl] { + pub fn find(&self, name: Name) -> &[Decl] { let range = self.decls.equal_range_by(|x| x.name().cmp(&name)); &self.decls[range] } @@ -55,13 +326,13 @@ impl DeclTable { /// Non-function decls sharing the name (a global, struct, etc.) are skipped /// rather than shadowing the function, since decls with equal names are /// ordered by source position. - pub fn find_entry_point(&self, name: Name) -> Option<&FuncDecl> { + pub fn find_entry_point(&self, name: Name) -> Option<&F> { self.entry_point_overloads(name).next() } /// Returns every function declared with the given name, ignoring non-function /// decls that happen to share it. - pub fn entry_point_overloads(&self, name: Name) -> impl Iterator { + pub fn entry_point_overloads(&self, name: Name) -> impl Iterator { self.find(name).iter().filter_map(|d| match d { Decl::Func(d) => Some(d), _ => None, @@ -99,39 +370,11 @@ impl DeclTable { alts } - - /// Returns all alternatives for a declaration name. - pub fn alts(&self, name: Name) -> Vec { - let sl = self.find(name); - let mut alts = vec![]; - - for d in sl { - match d { - Decl::Func(FuncDecl { constraints, .. }) => { - let mut interfaces = vec![]; - for c in constraints { - interfaces.push(AltInterface { - interface: c.interface_name, - typevars: c.typevars.iter().map(|name| typevar(name)).collect(), - }) - } - - alts.push(Alt { - ty: d.ty(), - interfaces, - }); - } - Decl::Global { .. } => { - alts.push(Alt { - ty: d.ty(), - interfaces: vec![], - }); - } - _ => (), - } - } - - alts +} +impl Deref for DeclTable { + type Target = DeclarationList; + fn deref(&self) -> &Self::Target { + &self.list } } @@ -142,7 +385,7 @@ mod tests { #[test] fn test_sorted_decls() { - let decls = vec![ + let decls: Vec = vec![ Decl::Global { name: Name::new("a".into()), typevars: vec![], diff --git a/src/expr.rs b/src/expr.rs index 306bb20e..3ca86009 100644 --- a/src/expr.rs +++ b/src/expr.rs @@ -168,6 +168,70 @@ impl Expr { } } + /// Rewrite immediate edges without interpreting names, types, or bindings. + pub fn map_children(&mut self, mut map: impl FnMut(ExprID) -> ExprID) { + match self { + Self::Call(function, args) => { + *function = map(*function); + for arg in args { + *arg = map(*arg); + } + } + Self::Macro(_, args) + | Self::ArrayLiteral(args) + | Self::Block(args) + | Self::Tuple(args) => { + for arg in args { + *arg = map(*arg); + } + } + Self::Binop(_, lhs, rhs) + | Self::Array(lhs, rhs) + | Self::ArrayIndex(lhs, rhs) + | Self::While(lhs, rhs) => { + *lhs = map(*lhs); + *rhs = map(*rhs); + } + Self::Unop(_, child) + | Self::Field(child, _) + | Self::AsTy(child, _) + | Self::Let(_, child, _) + | Self::Return(child) + | Self::Arena(child) + | Self::Assume(child) + | Self::Lambda { body: child, .. } => { + *child = map(*child); + } + Self::Var(_, child, _) => { + if let Some(child) = child { + *child = map(*child); + } + } + Self::If(cond, yes, no) => { + *cond = map(*cond); + *yes = map(*yes); + if let Some(no) = no { + *no = map(*no); + } + } + Self::For { + start, end, body, .. + } => { + *start = map(*start); + *end = map(*end); + *body = map(*body); + } + Self::StructLit(_, fields) => { + for (_, value) in fields { + *value = map(*value); + } + } + _ => {} + } + } +} + +impl Expr { /// Pretty-print an expression in lyte syntax. /// /// This method formats an expression as it would appear in lyte source code. @@ -575,7 +639,7 @@ pub fn format_binop(op: Binop) -> &'static str { } } -fn format_unop(op: Unop) -> &'static str { +pub fn format_unop(op: Unop) -> &'static str { match op { Unop::Neg => "-", Unop::Not => "!", diff --git a/src/free_locals.rs b/src/free_locals.rs new file mode 100644 index 00000000..c543958c --- /dev/null +++ b/src/free_locals.rs @@ -0,0 +1,69 @@ +//! Capture discovery needs lexical bindings, not solved types or publication. +use crate::{CheckedBody, ExprID, LocalId, Reference}; +use std::collections::HashSet; + +/// The narrow facts needed from either a checked body or a check in progress. +pub(crate) struct BindingNode { + pub children: Vec, + pub used: Option, + pub declared: Vec, + /// False for unavailable resolution/binder facts. Known facts remain useful. + pub complete: bool, +} + +pub(crate) trait BindingFacts { + fn binding_node(&self, id: ExprID) -> BindingNode; +} + +pub(crate) struct FreeLocals { + /// Unique free locals, in first-use order, including uses in nested lambdas. + pub locals: Vec, + /// An empty list establishes noncapture only when discovery is complete. + pub complete: bool, +} + +pub(crate) fn free_locals( + facts: &impl BindingFacts, + root: ExprID, + bound: impl IntoIterator, +) -> FreeLocals { + fn walk( + facts: &impl BindingFacts, + id: ExprID, + declared: &mut HashSet, + result: &mut FreeLocals, + ) { + let node = facts.binding_node(id); + result.complete &= node.complete; + result.locals.extend(node.used); + declared.extend(node.declared); + for child in node.children { + walk(facts, child, declared, result); + } + } + let mut declared = bound.into_iter().collect(); + let mut result = FreeLocals { + locals: vec![], + complete: true, + }; + walk(facts, root, &mut declared, &mut result); + let mut seen = HashSet::new(); + result + .locals + .retain(|local| !declared.contains(local) && seen.insert(*local)); + result +} + +impl BindingFacts for CheckedBody { + fn binding_node(&self, id: ExprID) -> BindingNode { + BindingNode { + children: self[id].subexprs(), + used: match self.reference(id) { + Some(Reference::Local(local)) => Some(*local), + _ => None, + }, + declared: self.binders(id).to_vec(), + complete: true, + } + } +} diff --git a/src/hoist.rs b/src/hoist.rs index 1be90b65..4053c6ba 100644 --- a/src/hoist.rs +++ b/src/hoist.rs @@ -1,703 +1,575 @@ use crate::*; use std::collections::{HashMap, HashSet}; -/// The globals a function may write, directly or through anything it calls. -/// `None` means "may write any global" — the function reaches a callee we -/// can't see through (an indirect call, or an `extern` body). -type GlobalWrites = Option>; +/// Storage identities establish binding identity, not disjoint pointees. +/// Borrowed arguments and closure captures are invalidated separately below. +#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)] +enum Root { + Local(LocalId), + Global(InstanceId), +} + +type GlobalWrites = Option>; -/// Whole-module may-write summary, used to decide whether a call inside a loop -/// can invalidate a hoisted field read. -/// -/// Lyte function parameters aren't assignable, so a callee's only channel for -/// mutating state its caller can observe is a global. That makes a -/// global-granularity summary exact enough to be useful: a loop calling a -/// function that touches no globals keeps all of its hoists. +/// Global may-write summaries over concrete function instances. This is not a +/// purity analysis: a builtin can print while writing no module global. pub struct SideEffects { - globals: HashSet, - per_func: HashMap, + globals: HashSet, + per_function: HashMap, } impl SideEffects { - /// Build the summary by taking each function's direct global writes and - /// closing over the call graph to a fixpoint. - /// - /// Overloads are merged: the graph is keyed by name, so a name resolves to - /// the union of every overload's writes. That over-approximates, which is - /// the safe direction. - pub fn analyze(decls: &DeclTable) -> Self { - let globals: HashSet = decls - .decls - .iter() - .filter_map(|d| match d { - Decl::Global { name, .. } => Some(*name), - _ => None, - }) - .collect(); - - let known_funcs: HashSet = decls - .decls - .iter() - .filter_map(|d| match d { - Decl::Func(f) => Some(f.name), - _ => None, - }) - .collect(); - - let mut direct: HashMap> = HashMap::new(); - let mut callees: HashMap> = HashMap::new(); - let mut opaque: HashSet = HashSet::new(); - - for d in &decls.decls { - let f = match d { - Decl::Func(f) => f, - _ => continue, - }; - let writes = direct.entry(f.name).or_default(); - let calls = callees.entry(f.name).or_default(); - let body = match f.body { - Some(b) => b, - None => { - // No body: either a builtin (print/assert/putc/math), which - // touches no globals, or an `extern` we can't see into. - if f.is_extern { - opaque.insert(f.name); + pub fn analyze(program: &SpecializedProgram) -> Result { + let mut globals = HashSet::new(); + let mut direct = HashMap::>::new(); + let mut callees = HashMap::>::new(); + let mut opaque = HashSet::new(); + for index in 0..program.instances.len() { + let instance = InstanceId(index as u32); + match program.instance(instance) { + Decl::Global { .. } => { + globals.insert(instance); + } + Decl::Func(function) => { + let writes = direct.entry(instance).or_default(); + let calls = callees.entry(instance).or_default(); + let mut is_opaque = function.is_extern; + if let Some(body) = function.body { + scan_effects(body, function, writes, calls, &mut is_opaque); + } + if is_opaque { + opaque.insert(instance); } - continue; } - }; - let mut is_opaque = false; - scan_effects( - body, - f, - &globals, - &known_funcs, - &shadowed_names(f, &globals), - writes, - calls, - &mut is_opaque, - ); - if is_opaque { - opaque.insert(f.name); + _ => return Err("Concrete instance is neither a function nor global".into()), + } + } + for (&function, calls) in &callees { + if calls.iter().any(|callee| !direct.contains_key(callee)) { + opaque.insert(function); } } - - // Fixpoint: propagate writes and opacity along call edges. Cycles just - // stop changing the sets, so recursion terminates without special - // handling. loop { let mut changed = false; - let names: Vec = direct.keys().copied().collect(); - for name in names { - let targets: Vec = callees[&name].iter().copied().collect(); - for t in targets { - if opaque.contains(&t) && !opaque.contains(&name) { - opaque.insert(name); + for (&function, calls) in &callees { + for callee in calls { + if opaque.contains(callee) && opaque.insert(function) { changed = true; } - let t_writes: Vec = direct - .get(&t) - .map(|w| w.iter().copied().collect()) - .unwrap_or_default(); - let w = direct.get_mut(&name).unwrap(); - for g in t_writes { - if w.insert(g) { - changed = true; - } - } + let writes = direct.get(callee).cloned().unwrap_or_default(); + let own = direct.get_mut(&function).unwrap(); + let previous_len = own.len(); + own.extend(writes); + changed |= own.len() != previous_len; } } if !changed { break; } } + Ok(Self { + globals, + per_function: direct + .into_iter() + .map(|(function, writes)| { + (function, (!opaque.contains(&function)).then_some(writes)) + }) + .collect(), + }) + } - let per_func = direct - .into_iter() - .map(|(name, writes)| { - if opaque.contains(&name) { - (name, None) - } else { - (name, Some(writes)) - } - }) - .collect(); - - SideEffects { globals, per_func } + fn writes_of_call(&self, callee: ExprID, function: &CheckedFunction) -> GlobalWrites { + function_target(callee, function) + .and_then(|target| self.per_function.get(&target).cloned().flatten()) } +} - /// Globals possibly written by a call whose callee expression is `func`. - /// An indirect callee (a local or parameter of function type) is opaque, - /// including one whose name shadows a function of the same name: the - /// summary is keyed by name, so consulting it there would describe a - /// callee this call never reaches. - fn writes_of_call( - &self, - func: ExprID, - fdecl: &FuncDecl, - shadowed: &HashSet, - ) -> GlobalWrites { - match &fdecl.arena.exprs[func] { - Expr::Id(name) if !shadowed.contains(name) => self.writes_of(*name), - _ => None, - } +fn function_target(expr: ExprID, function: &CheckedFunction) -> Option { + match function.arena.reference(expr) { + Some(Reference::Instance(instance)) => Some(*instance), + _ => None, } +} - /// Globals possibly written by calling `name`. `None` means "any global". - fn writes_of(&self, name: Name) -> GlobalWrites { - match self.per_func.get(&name) { - Some(w) => w.clone(), - // Not a function we know about — a call through a local or - // parameter of function type. Assume the worst. - None => None, - } +fn storage_root(expr: ExprID, function: &CheckedFunction) -> Option { + match &function.arena[expr] { + Expr::Id(_) => match function.arena.reference(expr)? { + Reference::Local(local) => Some(Root::Local(*local)), + Reference::Instance(instance) => Some(Root::Global(*instance)), + _ => None, + }, + Expr::Field(base, _) | Expr::ArrayIndex(base, _) => storage_root(*base, function), + _ => None, } } -/// Collect one function's direct global writes and its call edges. -/// -/// Sets `is_opaque` when the function calls something that isn't a statically -/// known function name, since we can't summarize such a callee. -#[allow(clippy::too_many_arguments)] fn scan_effects( - expr_id: ExprID, - fdecl: &FuncDecl, - globals: &HashSet, - known_funcs: &HashSet, - shadowed: &HashSet, - writes: &mut HashSet, - calls: &mut HashSet, - is_opaque: &mut bool, + expr: ExprID, + function: &CheckedFunction, + writes: &mut HashSet, + calls: &mut HashSet, + opaque: &mut bool, ) { - match &fdecl.arena.exprs[expr_id] { + match &function.arena[expr] { Expr::Binop(Binop::Assign, lhs, _) => { - if let Some(base) = assigned_root(*lhs, fdecl) { - if globals.contains(&base) { - writes.insert(base); - } + if let Some(Root::Global(global)) = storage_root(*lhs, function) { + writes.insert(global); } } - Expr::Call(func, _) => match &fdecl.arena.exprs[*func] { - // A name that's also bound locally names the binding, not the - // function, so the call edge would point at the wrong callee. - Expr::Id(name) if known_funcs.contains(name) && !shadowed.contains(name) => { - calls.insert(*name); + Expr::Call(callee, args) => { + if let Some(target) = function_target(*callee, function) { + calls.insert(target); + } else { + *opaque = true; } - _ => *is_opaque = true, - }, + // A wrapper can pass module storage to a callee that only sees a + // parameter. Its summary must account for that borrowed write. + for root in written_argument_roots(*callee, args, function) { + if let Root::Global(global) = root { + writes.insert(global); + } + } + } _ => {} } - - for sub in fdecl.arena.exprs[expr_id].subexprs() { - scan_effects( - sub, - fdecl, - globals, - known_funcs, - shadowed, - writes, - calls, - is_opaque, - ); + // Scanning lambda bodies conservatively includes their eventual effects. + for child in function.arena[expr].subexprs() { + scan_effects(child, function, writes, calls, opaque); } } -/// Names that a call site can't be summarized through: every name the function -/// binds itself — parameters, each `let`, `var`, lambda parameter and `for` -/// variable — plus every module-level global. -/// -/// Calls are summarized by callee name, so a call through a name bound to a -/// value has to be treated as indirect: the summary for the *function* of that -/// name describes a callee the call never reaches. Globals count because a -/// global of function type shadows a same-named function everywhere, not just -/// in the function that declares a local. The local half is function-wide -/// rather than scope-precise: shadowing a function name is rare, and the extra -/// conservatism only costs hoists. -fn shadowed_names(fdecl: &FuncDecl, globals: &HashSet) -> HashSet { - let mut names: HashSet = fdecl.params.iter().map(|p| p.name).collect(); - names.extend(globals.iter().copied()); - if let Some(body) = fdecl.body { - collect_bound_names(body, fdecl, &mut names); +/// Scalar field hoisting, using checked roots and concrete callees. +/// Lambda bodies retain their own evaluation boundary and are never rewritten. +pub fn hoist_loop_invariant_fields( + function: &mut CheckedFunction, + effects: &SideEffects, +) -> Result<(), String> { + if let Some(body) = function.body { + let captured = function.arena.captured_locals(); + hoist_in_expr(body, function, effects, &captured); } - names -} - -/// What the hoist walk needs to know about one function's names. -struct LocalNames { - /// Callee names that don't resolve to the function of the same name. - shadowed: HashSet, - - /// Names mentioned inside a lambda body in this function. A `var` a lambda - /// captures is shared by address, so a call the summary can't see through - /// may be that lambda, writing one of these behind the loop's back. - captured: HashSet, + Ok(()) } -fn collect_bound_names(expr_id: ExprID, fdecl: &FuncDecl, names: &mut HashSet) { - match &fdecl.arena.exprs[expr_id] { - Expr::Let(name, _, _) | Expr::Var(name, _, _) => { - names.insert(*name); +fn hoist_in_expr( + expr: ExprID, + function: &mut CheckedFunction, + effects: &SideEffects, + captured: &HashSet, +) { + match function.arena[expr].clone() { + Expr::Block(statements) => { + for statement in statements { + hoist_in_expr(statement, function, effects, captured); + } + hoist_loops_in_block(expr, function, effects, captured); } - Expr::For { var, .. } => { - names.insert(*var); + Expr::For { body, .. } | Expr::While(_, body) => { + hoist_in_expr(body, function, effects, captured) } - Expr::Lambda { params, .. } => { - for p in params { - names.insert(p.name); + Expr::If(_, then_branch, else_branch) => { + hoist_in_expr(then_branch, function, effects, captured); + if let Some(other) = else_branch { + hoist_in_expr(other, function, effects, captured); } } _ => {} } - - for sub in fdecl.arena.exprs[expr_id].subexprs() { - collect_bound_names(sub, fdecl, names); - } } -/// The variable at the root of an assignment target: `g`, `g.f`, `g[i].f`, ... -fn assigned_root(expr_id: ExprID, fdecl: &FuncDecl) -> Option { - match &fdecl.arena.exprs[expr_id] { - Expr::Id(name) => Some(*name), - Expr::Field(base, _) | Expr::ArrayIndex(base, _) => assigned_root(*base, fdecl), - _ => None, - } +type WrittenFields = HashSet<(Root, Option)>; + +struct FieldRead { + root: Root, + field: Name, + expr: ExprID, } -/// Hoist loop-invariant struct field loads out of loops. -/// -/// For each loop (For/While), finds struct field accesses (`expr.field`) where: -/// - The base expression is a simple local variable (`Expr::Id`) -/// - The field is never written to inside the loop body -/// - The field type is a scalar (not a struct/array/slice) -/// -/// Each such access is replaced with a reference to a hoisted `let` binding -/// inserted just before the loop. -pub fn hoist_loop_invariant_fields(fdecl: &mut FuncDecl, effects: &SideEffects) { - if fdecl.body.is_none() { - return; - } - let body = fdecl.body.unwrap(); - let names = LocalNames { - shadowed: shadowed_names(fdecl, &effects.globals), - captured: fdecl.names_referenced_in_lambdas(), - }; - hoist_in_expr(body, fdecl, effects, &names); +/// Reference/slice bindings and aggregate parameters can designate storage +/// owned outside this body. Their identities do not exclude overlap with globals. +fn aliased_local_roots(function: &CheckedFunction) -> HashSet { + let mut roots: HashSet<_> = function + .arena + .locals + .iter() + .enumerate() + .filter(|(_, local)| matches!(&*local.ty, Type::Reference(_) | Type::Slice(_))) + .map(|(index, _)| Root::Local(LocalId(index as u32))) + .collect(); + roots.extend(function.params.iter().filter_map(|parameter| { + let ty = function.arena.local(parameter.local).ty; + (is_ptr_type(ty) || matches!(&*ty, Type::Reference(_) | Type::Float32x4)) + .then_some(Root::Local(parameter.local)) + })); + roots } -/// Recursively walk the AST looking for loops inside blocks. -/// When we find a loop inside a block, we can insert hoisted bindings before it. -fn hoist_in_expr(expr_id: ExprID, fdecl: &mut FuncDecl, effects: &SideEffects, names: &LocalNames) { - match fdecl.arena.exprs[expr_id].clone() { - Expr::Block(stmts) => { - // First, recurse into each statement. - for &s in &stmts { - hoist_in_expr(s, fdecl, effects, names); - } - // Now look for loops in this block and hoist their invariant fields. - hoist_loops_in_block(expr_id, fdecl, effects, names); - } - Expr::For { body, .. } => { - hoist_in_expr(body, fdecl, effects, names); - } - Expr::While(_, body) => { - hoist_in_expr(body, fdecl, effects, names); - } - Expr::If(_, then_branch, else_branch) => { - hoist_in_expr(then_branch, fdecl, effects, names); - if let Some(e) = else_branch { - hoist_in_expr(e, fdecl, effects, names); - } - } - _ => {} +fn invalidate_aliased_writes( + written: &mut WrittenFields, + aliases: &HashSet, + effects: &SideEffects, +) { + let writes_global = written + .iter() + .any(|(root, _)| matches!(root, Root::Global(_))); + let writes_borrowed = written.iter().any(|(root, _)| aliases.contains(root)); + if writes_global || writes_borrowed { + written.extend(aliases.iter().map(|root| (*root, None))); + } + if writes_borrowed { + written.extend( + effects + .globals + .iter() + .map(|global| (Root::Global(*global), None)), + ); } } -/// For each loop statement in a block, hoist invariant struct field reads. fn hoist_loops_in_block( - block_id: ExprID, - fdecl: &mut FuncDecl, + block: ExprID, + function: &mut CheckedFunction, effects: &SideEffects, - names: &LocalNames, + captured: &HashSet, ) { - let stmts = if let Expr::Block(ref stmts) = fdecl.arena.exprs[block_id] { - stmts.clone() - } else { + let Expr::Block(statements) = function.arena[block].clone() else { return; }; - - let mut new_stmts = Vec::with_capacity(stmts.len()); - - for &stmt_id in &stmts { - let loop_body = match &fdecl.arena.exprs[stmt_id] { - Expr::For { body, .. } => Some(*body), - Expr::While(_, body) => Some(*body), - _ => None, - }; - - if let Some(body_id) = loop_body { - // Find all fields written in the loop body. - let mut written_fields: HashSet<(Name, Name)> = HashSet::new(); - collect_written_fields(body_id, fdecl, effects, names, &mut written_fields); - // The hoisted binding is inserted before the whole loop statement, - // so anything the loop's own header evaluates runs after it: a - // `while` condition, and a `for` range, both have to be scanned. - match fdecl.arena.exprs[stmt_id].clone() { - Expr::While(cond, _) => { - collect_written_fields(cond, fdecl, effects, names, &mut written_fields); - } - Expr::For { - var, start, end, .. - } => { - written_fields.insert((var, Name::str("*"))); - collect_written_fields(start, fdecl, effects, names, &mut written_fields); - collect_written_fields(end, fdecl, effects, names, &mut written_fields); - } - _ => {} + let aliases = aliased_local_roots(function); + let mut replacement = Vec::with_capacity(statements.len()); + for statement in statements { + let body = match function.arena[statement] { + Expr::For { body, .. } | Expr::While(_, body) => body, + _ => { + replacement.push(statement); + continue; } - - // Find all field reads that are loop-invariant. - let mut field_reads: Vec = Vec::new(); - collect_invariant_field_reads(body_id, fdecl, &written_fields, &mut field_reads); - if let Expr::While(cond, _) = &fdecl.arena.exprs[stmt_id] { - collect_invariant_field_reads(*cond, fdecl, &written_fields, &mut field_reads); + }; + let mut written = WrittenFields::new(); + collect_written_fields(body, function, effects, captured, &mut written); + // The new initializer precedes both range evaluation and testing the + // condition. Either can invalidate the prospective hoist's source. + match function.arena[statement].clone() { + Expr::While(condition, _) => { + collect_written_fields(condition, function, effects, captured, &mut written) } - - // Deduplicate. - let mut seen = HashSet::new(); - field_reads.retain(|r| seen.insert((r.var, r.field))); - - if !field_reads.is_empty() { - // Create hoisted let bindings and a substitution map. - let mut subst: HashMap<(Name, Name), Name> = HashMap::new(); - let loc = fdecl.arena.locs[stmt_id]; - - for read in &field_reads { - let hoisted_name = - Name::new(format!("__hoisted_{}_{}", &**read.var, &**read.field)); - - // Build: let __hoisted_var_field = var.field - // - // The types come from the nodes we're copying, not from a - // by-name lookup: a same-named binding in a sibling scope - // would otherwise supply the wrong struct type, and the - // hoisted read would use that type's field offset. - let id_expr = fdecl.arena.add(Expr::Id(read.var), loc); - fdecl.types.push(read.base_type); // type for the Id expr - - let field_expr = fdecl.arena.add(Expr::Field(id_expr, read.field), loc); - fdecl.types.push(read.field_type); // type for the Field expr - - let let_expr = fdecl - .arena - .add(Expr::Let(hoisted_name, field_expr, None), loc); - fdecl.types.push(read.field_type); // type for the Let expr (must match init) - - new_stmts.push(let_expr); - subst.insert((read.var, read.field), hoisted_name); - } - - // Replace field accesses in the loop body with hoisted variable references. - replace_field_reads(body_id, fdecl, &subst); - if let Expr::While(cond, _) = fdecl.arena.exprs[stmt_id].clone() { - replace_field_reads(cond, fdecl, &subst); - } + Expr::For { start, end, .. } => { + invalidate_binders(statement, function, &mut written); + collect_written_fields(start, function, effects, captured, &mut written); + collect_written_fields(end, function, effects, captured, &mut written); } - } - - new_stmts.push(stmt_id); + _ => unreachable!(), + } + invalidate_aliased_writes(&mut written, &aliases, effects); + let mut reads = vec![]; + collect_invariant_field_reads(body, function, &written, &mut reads); + if let Expr::While(condition, _) = function.arena[statement] { + collect_invariant_field_reads(condition, function, &written, &mut reads); + } + let mut seen = HashSet::new(); + reads.retain(|read| seen.insert((read.root, read.field))); + let mut substitutions = HashMap::new(); + for read in reads { + let (local, declaration) = create_hoisted_binding(&read, &mut function.arena); + replacement.push(declaration); + substitutions.insert((read.root, read.field), local); + } + replace_field_reads(body, function, &substitutions); + if let Expr::While(condition, _) = function.arena[statement] { + replace_field_reads(condition, function, &substitutions); + } + replacement.push(statement); } + function + .arena + .replace(block, Expr::Block(replacement), function.arena.ty(block)); +} - if new_stmts.len() != stmts.len() { - fdecl.arena.exprs[block_id] = Expr::Block(new_stmts); - } +fn create_hoisted_binding(read: &FieldRead, arena: &mut CheckedBody) -> (LocalId, ExprID) { + let Expr::Field(base, _) = arena[read.expr] else { + unreachable!(); + }; + let Expr::Id(base_name) = arena[base] else { + unreachable!(); + }; + let base_reference = arena.reference(base).cloned().expect("checked reference"); + let (base_ty, base_loc) = (arena.ty(base), arena.loc(base)); + let (field_ty, field_loc) = (arena.ty(read.expr), arena.loc(read.expr)); + // These are fresh evaluations with fresh ExprIDs, while the copied + // outer reference retains its LocalId or global InstanceId. + let base = arena.add_id(base_name, base_reference, base_ty, base_loc); + let initializer = arena.add(Expr::Field(base, read.field), field_ty, field_loc); + let local = arena.add_local( + Name::new(format!("__hoisted_{}", read.field)), + field_ty, + false, + ); + let declaration = arena.add_let(local, initializer, field_loc); + (local, declaration) +} + +fn invalidate_binders(expr: ExprID, function: &CheckedFunction, written: &mut WrittenFields) { + written.extend( + function + .arena + .binders(expr) + .iter() + .map(|&local| (Root::Local(local), None)), + ); } -/// Collect all (variable_name, field_name) pairs that are written to in the -/// expression tree. `(name, "*")` means the whole variable is clobbered. -/// -/// This has to be complete: a write we miss becomes a stale hoisted read. The -/// traversal therefore handles the shapes it cares about and then recurses -/// into every subexpression, so a new `Expr` variant can't silently escape it. fn collect_written_fields( - expr_id: ExprID, - fdecl: &FuncDecl, + expr: ExprID, + function: &CheckedFunction, effects: &SideEffects, - names: &LocalNames, - written: &mut HashSet<(Name, Name)>, + captured: &HashSet, + written: &mut WrittenFields, ) { - match &fdecl.arena.exprs[expr_id] { - Expr::Binop(Binop::Assign, lhs, _) => match &fdecl.arena.exprs[*lhs] { - // `var = ...` replaces the whole variable. - Expr::Id(var_name) => { - written.insert((*var_name, Name::str("*"))); + match &function.arena[expr] { + Expr::Binop(Binop::Assign, lhs, _) => match &function.arena[*lhs] { + Expr::Id(_) => { + if let Some(root) = storage_root(*lhs, function) { + written.insert((root, None)); + } } - // `var.field = ...` clobbers exactly that field. - // - // Deeper targets need nothing: `var.a.b = ...` and `var.arr[i] = ...` - // write through an aggregate field, and `slice[i] = ...` writes - // elements rather than the slice's `len`. Only scalar fields are - // ever hoisted, and none of those writes can reach one. - Expr::Field(base, field_name) => { - if let Expr::Id(var_name) = &fdecl.arena.exprs[*base] { - written.insert((*var_name, *field_name)); + Expr::Field(base, field) if matches!(function.arena[*base], Expr::Id(_)) => { + if let Some(root) = storage_root(*base, function) { + written.insert((root, Some(*field))); } } + // Deeper aggregate writes cannot affect the scalar direct fields + // eligible for this pass's existing read grammar. _ => {} }, - Expr::Call(func, args) => { - // Parameters aren't assignable in Lyte, so a callee we can name - // reaches its caller's state only through globals. - match effects.writes_of_call(*func, fdecl, &names.shadowed) { - Some(gs) => { - for g in gs { - written.insert((g, Name::str("*"))); - } - } + Expr::Call(callee, args) => { + match effects.writes_of_call(*callee, function) { + Some(globals) => written.extend( + globals + .into_iter() + .map(|global| (Root::Global(global), None)), + ), None => { - // A callee we can't name may be a lambda holding the - // address of one of our own locals, so those go too. - for g in &effects.globals { - written.insert((*g, Name::str("*"))); - } - for c in &names.captured { - written.insert((*c, Name::str("*"))); - } - } - } - // Aggregates handed to a call stay tainted as before. - for &arg in args { - if let Expr::Id(var_name) = &fdecl.arena.exprs[arg] { - written.insert((*var_name, Name::str("*"))); + written.extend( + effects + .globals + .iter() + .map(|global| (Root::Global(*global), None)), + ); + written.extend(captured.iter().map(|local| (Root::Local(*local), None))); + written.extend( + aliased_local_roots(function) + .into_iter() + .map(|root| (root, None)), + ); } } + written + .extend(written_argument_roots(*callee, args, function).map(|root| (root, None))); } - Expr::Var(name, _, _) => { - // A binding introduced inside the loop is a fresh variable on every - // iteration, and it doesn't exist before the loop at all — nothing - // about it can be hoisted. - written.insert((*name, Name::str("*"))); - } - Expr::Let(name, _, _) => { - written.insert((*name, Name::str("*"))); - } - Expr::Lambda { params, .. } => { - // Lambda parameters shadow anything of the same name in the - // enclosing scope, so reads through them aren't invariant either. - for p in params { - written.insert((p.name, Name::str("*"))); - } - } - Expr::For { var, .. } => { - written.insert((*var, Name::str("*"))); + Expr::Let(..) | Expr::Var(..) | Expr::For { .. } | Expr::Lambda { .. } => { + invalidate_binders(expr, function, written); } _ => {} } + for child in function.arena[expr].subexprs() { + collect_written_fields(child, function, effects, captured, written); + } +} - for sub in fdecl.arena.exprs[expr_id].subexprs() { - collect_written_fields(sub, fdecl, effects, names, written); +/// Conservatively include direct bindings and borrowed projections. Scalar +/// projections passed by value only read their root. +fn written_argument_roots<'a>( + callee: ExprID, + args: &'a [ExprID], + function: &'a CheckedFunction, +) -> impl Iterator + 'a { + args.iter().enumerate().filter_map(move |(position, &arg)| { + if matches!(function.arena[arg], Expr::Id(_)) + || borrows_argument(callee, position, function) + { + storage_root(arg, function) + } else { + None + } + }) +} + +fn borrows_argument(callee: ExprID, position: usize, function: &CheckedFunction) -> bool { + let Type::Func(domain, _) = &*function.arena.ty(callee) else { + return true; + }; + let Type::Tuple(params) = &**domain else { + return true; + }; + match params.get(position).map(|ty| &**ty) { + Some( + Type::Reference(_) | Type::Slice(_) | Type::Array(..) | Type::Name(..) | Type::Tuple(_), + ) => true, + Some(_) => false, + None => true, } } -/// A field read the hoister decided is loop-invariant, carrying the types of -/// the nodes it was found on so the hoisted copy reproduces them exactly. -struct FieldRead { - var: Name, - field: Name, - base_type: TypeID, - field_type: TypeID, +fn read_subexprs(expr: &Expr) -> Vec { + match expr { + Expr::Lambda { .. } | Expr::Arena(_) | Expr::Array(..) | Expr::Macro(..) => vec![], + Expr::Binop(Binop::Assign, _, rhs) => vec![*rhs], + _ => expr.subexprs(), + } } -/// Collect loop-invariant scalar field reads. fn collect_invariant_field_reads( - expr_id: ExprID, - fdecl: &FuncDecl, - written: &HashSet<(Name, Name)>, + expr: ExprID, + function: &CheckedFunction, + written: &WrittenFields, reads: &mut Vec, ) { - match &fdecl.arena.exprs[expr_id] { - Expr::Field(base, field_name) => { - if let Expr::Id(var_name) = &fdecl.arena.exprs[*base] { - let pair = (*var_name, *field_name); - let wildcard = (*var_name, Name::str("*")); - // Only hoist if the field is never written and the variable isn't wholly reassigned. - if !written.contains(&pair) && !written.contains(&wildcard) { - // Only hoist scalar fields (not sub-structs, arrays, etc.) - let field_type = fdecl.types[expr_id]; - if !is_ptr_type(&field_type) { - reads.push(FieldRead { - var: *var_name, - field: *field_name, - base_type: fdecl.types[*base], - field_type, - }); - } + if let Expr::Field(base, field) = &function.arena[expr] { + if matches!(function.arena[*base], Expr::Id(_)) { + if let Some(root) = storage_root(*base, function) { + if !written.contains(&(root, Some(*field))) + && !written.contains(&(root, None)) + && !is_ptr_type(function.arena.ty(expr)) + { + reads.push(FieldRead { + root, + field: *field, + expr, + }); } } - collect_invariant_field_reads(*base, fdecl, written, reads); - } - Expr::Binop(Binop::Assign, _lhs, rhs) => { - // Don't collect reads from the LHS of assignments. - collect_invariant_field_reads(*rhs, fdecl, written, reads); - } - Expr::Binop(_, lhs, rhs) => { - collect_invariant_field_reads(*lhs, fdecl, written, reads); - collect_invariant_field_reads(*rhs, fdecl, written, reads); - } - Expr::Unop(_, arg) => { - collect_invariant_field_reads(*arg, fdecl, written, reads); - } - Expr::Call(func, args) => { - collect_invariant_field_reads(*func, fdecl, written, reads); - for &arg in args { - collect_invariant_field_reads(arg, fdecl, written, reads); - } - } - Expr::Block(stmts) => { - for &s in stmts { - collect_invariant_field_reads(s, fdecl, written, reads); - } - } - Expr::If(cond, then_b, else_b) => { - collect_invariant_field_reads(*cond, fdecl, written, reads); - collect_invariant_field_reads(*then_b, fdecl, written, reads); - if let Some(e) = else_b { - collect_invariant_field_reads(*e, fdecl, written, reads); - } - } - Expr::While(cond, body) => { - collect_invariant_field_reads(*cond, fdecl, written, reads); - collect_invariant_field_reads(*body, fdecl, written, reads); - } - Expr::For { - start, end, body, .. - } => { - collect_invariant_field_reads(*start, fdecl, written, reads); - collect_invariant_field_reads(*end, fdecl, written, reads); - collect_invariant_field_reads(*body, fdecl, written, reads); - } - Expr::ArrayIndex(base, idx) => { - collect_invariant_field_reads(*base, fdecl, written, reads); - collect_invariant_field_reads(*idx, fdecl, written, reads); - } - Expr::Return(e) | Expr::Assume(e) | Expr::AsTy(e, _) => { - collect_invariant_field_reads(*e, fdecl, written, reads); - } - Expr::Var(_, init, _) => { - if let Some(e) = init { - collect_invariant_field_reads(*e, fdecl, written, reads); - } - } - Expr::Let(_, init, _) => { - collect_invariant_field_reads(*init, fdecl, written, reads); - } - // Deliberately not descending into lambda bodies: capture lists are - // fixed before this pass runs, so a body rewritten to mention a - // hoisted binding would reference something it doesn't capture. - Expr::Lambda { .. } => {} - Expr::Tuple(elems) | Expr::ArrayLiteral(elems) => { - for &e in elems { - collect_invariant_field_reads(e, fdecl, written, reads); - } - } - Expr::StructLit(_, fields) => { - for (_, fval) in fields { - collect_invariant_field_reads(*fval, fdecl, written, reads); - } } - _ => {} + } + for child in read_subexprs(&function.arena[expr]) { + collect_invariant_field_reads(child, function, written, reads); } } -/// Replace field accesses in the expression tree with references to hoisted variables. -fn replace_field_reads(expr_id: ExprID, fdecl: &mut FuncDecl, subst: &HashMap<(Name, Name), Name>) { - match fdecl.arena.exprs[expr_id].clone() { - Expr::Field(base, field_name) => { - if let Expr::Id(var_name) = &fdecl.arena.exprs[base] { - let pair = (*var_name, field_name); - if let Some(hoisted_name) = subst.get(&pair) { - // Replace this Field expression with an Id referencing the hoisted variable. - fdecl.arena.exprs[expr_id] = Expr::Id(*hoisted_name); - return; - } - } - replace_field_reads(base, fdecl, subst); - } - Expr::Binop(Binop::Assign, _lhs, rhs) => { - // Don't replace in LHS of assignments. - replace_field_reads(rhs, fdecl, subst); - } - Expr::Binop(_, lhs, rhs) => { - replace_field_reads(lhs, fdecl, subst); - replace_field_reads(rhs, fdecl, subst); - } - Expr::Unop(_, arg) => { - replace_field_reads(arg, fdecl, subst); - } - Expr::Call(func, args) => { - replace_field_reads(func, fdecl, subst); - for arg in args { - replace_field_reads(arg, fdecl, subst); - } - } - Expr::Block(stmts) => { - for s in stmts { - replace_field_reads(s, fdecl, subst); - } - } - Expr::If(cond, then_b, else_b) => { - replace_field_reads(cond, fdecl, subst); - replace_field_reads(then_b, fdecl, subst); - if let Some(e) = else_b { - replace_field_reads(e, fdecl, subst); - } - } - Expr::While(cond, body) => { - replace_field_reads(cond, fdecl, subst); - replace_field_reads(body, fdecl, subst); - } - Expr::For { - start, end, body, .. - } => { - replace_field_reads(start, fdecl, subst); - replace_field_reads(end, fdecl, subst); - replace_field_reads(body, fdecl, subst); - } - Expr::ArrayIndex(base, idx) => { - replace_field_reads(base, fdecl, subst); - replace_field_reads(idx, fdecl, subst); - } - Expr::Return(e) | Expr::Assume(e) | Expr::AsTy(e, _) => { - replace_field_reads(e, fdecl, subst); - } - Expr::Var(_, init, _) => { - if let Some(e) = init { - replace_field_reads(e, fdecl, subst); - } - } - Expr::Let(_, init, _) => { - replace_field_reads(init, fdecl, subst); - } - // Not descended into, matching collect_invariant_field_reads: a lambda - // body must never be rewritten to mention a hoisted binding. - Expr::Lambda { .. } => {} - Expr::Tuple(elems) | Expr::ArrayLiteral(elems) => { - for e in elems { - replace_field_reads(e, fdecl, subst); - } - } - Expr::StructLit(_, fields) => { - for (_, fval) in fields { - replace_field_reads(fval, fdecl, subst); +fn replace_field_reads( + expr: ExprID, + function: &mut CheckedFunction, + substitutions: &HashMap<(Root, Name), LocalId>, +) { + if let Expr::Field(base, field) = function.arena[expr] { + if matches!(function.arena[base], Expr::Id(_)) { + if let Some(local) = + storage_root(base, function).and_then(|root| substitutions.get(&(root, field))) + { + let (name, ty) = (function.arena.local(*local).name, function.arena.ty(expr)); + function.arena.replace(expr, Expr::Id(name), ty); + function.arena.set_reference(expr, Reference::Local(*local)); + return; } } - _ => {} + } + for child in read_subexprs(&function.arena[expr]) { + replace_field_reads(child, function, substitutions); } } -/// Check if a type is a pointer type (struct, array, slice, tuple). -fn is_ptr_type(ty: &TypeID) -> bool { +fn is_ptr_type(ty: TypeID) -> bool { matches!( - &**ty, + &*ty, Type::Name(_, _) | Type::Tuple(_) | Type::Array(_, _) | Type::Slice(_) ) } + +#[cfg(test)] +mod tests { + use super::*; + + fn specialized(source: &str) -> SpecializedProgram { + let mut compiler = Compiler::new(); + compiler.quiet = true; + assert!(compiler.parse(source, "checked-hoisting.lyte")); + assert!(compiler.check(), "{:?}", compiler.last_errors); + MonomorphPass::new() + .monomorphize(compiler.checked_program().unwrap(), Name::str("main")) + .unwrap() + } + + fn main_mut(program: &mut SpecializedProgram) -> &mut CheckedFunction { + program + .decls + .decls + .iter_mut() + .find_map(|declaration| match declaration { + Decl::Func(function) if function.name == Name::str("main") => Some(function), + _ => None, + }) + .unwrap() + } + + #[test] + fn hoist_binding_identity_is_independent_of_diagnostic_spelling() { + let mut program = specialized("struct P { x: i32 } main { let __hoisted_x = 90; var p: P; p.x = 7; for i in 0 .. 2 { let value = p.x + __hoisted_x } }"); + let effects = SideEffects::analyze(&program).unwrap(); + let function = main_mut(&mut program); + let before = function.arena.locals.len(); + let original = function + .arena + .locals + .iter() + .position(|local| local.name == Name::str("__hoisted_x")) + .unwrap(); + hoist_loop_invariant_fields(function, &effects).unwrap(); + assert_eq!(function.arena.locals.len(), before + 1); + assert_eq!( + function.arena.locals[before].name, + function.arena.locals[original].name + ); + for local in [original, before] { + assert!(function + .arena + .ids() + .any(|id| function.arena.reference(id) + == Some(&Reference::Local(LocalId(local as u32))))); + } + } + + #[test] + fn borrowed_field_writes_prevent_hoisting() { + let mut program = specialized("struct P { x: i32 } bump(x: &i32) { x = x + 1 } main { var p: P; p.x = 1; for i in 0 .. 2 { let value = p.x; bump(p.x) } }"); + let effects = SideEffects::analyze(&program).unwrap(); + let function = main_mut(&mut program); + let before = function.clone(); + hoist_loop_invariant_fields(function, &effects).unwrap(); + assert_eq!(function, &before); + } + + #[test] + fn an_inner_shadow_does_not_invalidate_outer_storage() { + let mut program = specialized("struct P { x: i32 } main { var p: P; p.x = 7; for i in 0 .. 2 { if true { var p: P; p.x = i; let inner = p.x }; let outer = p.x } }"); + let effects = SideEffects::analyze(&program).unwrap(); + let function = main_mut(&mut program); + let before = function.arena.locals.len(); + hoist_loop_invariant_fields(function, &effects).unwrap(); + assert_eq!(function.arena.locals.len(), before + 1); + } + + #[test] + fn borrowed_parameter_and_global_writes_invalidate_aliases() { + // Call-site no-alias checking rejects overlapping arguments, but a + // single borrowed argument may still designate a module global. + for (read, written) in [("p", "g"), ("g", "p")] { + let source = format!( + "struct P {{ x: i32 }} var g: P + sum(p: &P) -> i32 {{ var result = 0 + for i in 0 .. 2 {{ result = result + {read}.x; {written}.x = {written}.x + 1 }} + result }} + main() -> i32 {{ g.x = 1; sum(g) }}" + ); + let mut compiler = Compiler::new(); + compiler.quiet = true; + assert!(compiler.parse(&source, "aliased-hoisting.lyte")); + assert!(compiler.check()); + compiler.specialize().unwrap(); + assert_eq!(crate::vm::VM::new().run(&compiler.compile_vm().unwrap()), 3); + #[cfg(has_stack_interp)] + assert_eq!( + crate::stack_interp_bridge::run(&compiler.compile_stack().unwrap()), + 3 + ); + } + } +} diff --git a/src/interface_resolution.rs b/src/interface_resolution.rs new file mode 100644 index 00000000..3b1c3bcf --- /dev/null +++ b/src/interface_resolution.rs @@ -0,0 +1,194 @@ +//! Checked interface candidates and exact-signature selection. Name lookup is +//! confined to candidate collection; selecting a concrete member uses IDs only. +use crate::*; + +#[derive(Clone, Debug, Eq, PartialEq, Hash)] +pub struct InterfaceMember { + pub definition: DefId, + /// Signature after substituting the enclosing where clause's parameters. + pub signature: TypeID, + pub candidates: Vec, +} + +impl InterfaceMember { + pub fn subst(&self, instance: &Instance) -> Self { + Self { + signature: self.signature.subst(instance), + ..self.clone() + } + } +} + +#[derive(Clone, Debug, Eq, PartialEq, Hash)] +pub struct InterfaceRequirement { + pub id: RequirementId, + pub interface: DefId, + pub type_args: Vec, + pub members: Vec, +} + +impl InterfaceRequirement { + /// Whole-body specialization preserves body-local requirement indices. + pub fn subst(&self, instance: &Instance) -> Self { + Self { + type_args: self.type_args.iter().map(|ty| ty.subst(instance)).collect(), + members: self + .members + .iter() + .map(|member| member.subst(instance)) + .collect(), + ..self.clone() + } + } + + pub fn select( + &self, + instance: &Instance, + decls: &DeclTable, + ) -> Result>, InterfaceSelectionError> { + select_interface_members(&self.type_args, &self.members, instance, decls) + } +} + +#[derive(Clone, Debug, Eq, PartialEq)] +pub struct SelectedInterfaceMember { + pub member: DefId, + pub implementation: DefId, + pub substitution: Instance, +} + +#[derive(Clone, Debug, Eq, PartialEq)] +pub struct InterfaceSelectionError { + pub member: DefId, +} + +impl InterfaceSelectionError { + pub fn message(&self, interface: DefId, decls: &DeclTable) -> String { + let member = decls + .function(self.member) + .expect("interface member identity") + .name(); + let interface = decls + .definition(interface) + .expect("interface identity") + .name(); + format!( + "function {} for interface {} is required", + member, interface + ) + } +} + +/// The signature rule intentionally differs from overload unification: generic +/// implementations and reference/array-to-slice coercions do not match merely +/// because their signatures can unify. +pub fn select_interface_member( + member: &InterfaceMember, + instance: &Instance, + decls: &DeclTable, +) -> Result { + let signature = member.signature.subst(instance); + // Existing interface matching chooses the first exact candidate in source + // order. In particular, programs may redeclare the stdlib's cmp signature. + // Keep that policy separate from ordinary overload ambiguity diagnostics. + let implementation = member + .candidates + .iter() + .copied() + .find(|id| decls.signature(*id) == Some(signature)) + .ok_or(InterfaceSelectionError { + member: member.definition, + })?; + Ok(SelectedInterfaceMember { + member: member.definition, + implementation, + substitution: instance.clone(), + }) +} + +pub fn select_interface_members( + type_args: &[TypeID], + members: &[InterfaceMember], + instance: &Instance, + decls: &DeclTable, +) -> Result>, InterfaceSelectionError> { + // Preserve the existing deferral boundary (top-level Var/Anon only). + if type_args + .iter() + .any(|ty| matches!(&*ty.subst(instance), Type::Var(_) | Type::Anon(_))) + { + return Ok(None); + } + members + .iter() + .map(|member| select_interface_member(member, instance, decls)) + .collect::, _>>() + .map(Some) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn table(source: &str) -> DeclTable { + let mut errors = Vec::new(); + let table = DeclTable::new(crate::parser::parse_program_str(source, &mut errors)); + assert!(errors.is_empty(), "{:?}", errors); + table + } + + #[test] + fn selection_keeps_declaration_identity_and_source_candidate_order() { + let source = table("interface Value { value(x: T) -> i32 }\nvalue(x: i32) -> i32 { 1 }\nvalue(x: i32) -> i32 { 2 }"); + let requirement = source + .interface_requirement( + RequirementId(0), + Name::str("Value"), + vec![mk_type(Type::Int32)], + ) + .unwrap(); + let expected = requirement.members[0].candidates[0]; + let mut records: Vec<_> = source.records().collect(); + records.reverse(); + for record in &mut records { + match &mut record.declaration { + Decl::Func(function) => function.name = Name::str("renamed"), + Decl::Interface(interface) => interface.funcs[0].name = Name::str("renamed_member"), + _ => unreachable!(), + } + } + let renamed = DeclTable::from_records(records); + let selected = requirement + .select(&Instance::new(), &renamed) + .unwrap() + .unwrap(); + assert_eq!(selected[0].implementation, expected); + assert_eq!( + renamed.function(selected[0].member).unwrap().name, + Name::str("renamed_member") + ); + } + + #[test] + fn exact_interface_matching_does_not_adopt_overload_coercions() { + for source in [ + "interface Value { value(x: T) -> i32 }\nvalue(x: U) -> i32 { 1 }", + "interface Value { value(x: [T; 3]) -> i32 }\nvalue(x: [i32]) -> i32 { 1 }", + "interface Value { value(x: &T) -> i32 }\nvalue(x: i32) -> i32 { 1 }", + ] { + let table = table(source); + let requirement = table + .interface_requirement( + RequirementId(0), + Name::str("Value"), + vec![mk_type(Type::Int32)], + ) + .unwrap(); + assert!( + requirement.select(&Instance::new(), &table).is_err(), + "{}", + source + ); + } + } +} diff --git a/src/jit.rs b/src/jit.rs index 2e27325c..eae3159a 100644 --- a/src/jit.rs +++ b/src/jit.rs @@ -2,11 +2,12 @@ // Pulled from https://github.com/bytecodealliance/cranelift-jit-demo use crate::cancel::*; -use crate::decl::*; +use crate::checked::{ + CheckedDecl as Decl, CheckedFunction as FuncDecl, InstanceId, LocalId, Reference, + SpecializedProgram as DeclTable, +}; use crate::defs::*; -use crate::expr::*; -use crate::DeclTable; -use crate::Instance; +use crate::Expr; use crate::TypeID; extern crate cranelift_codegen; use core::panic; @@ -93,10 +94,10 @@ pub struct JIT { /// functions. module: JITModule, - defined_functions: HashSet, + defined_functions: HashSet, /// Global variable offsets from base pointer. - globals: HashMap, + globals: HashMap, /// Total size of global memory needed. globals_size: usize, @@ -187,7 +188,8 @@ impl JIT { let Some(ep_decl) = decls.find_entry_point(ep_name) else { continue; }; - let id = self.compile_function(decls, ep_decl)?; + let instance = decls.instance_for_entry(ep_name).unwrap(); + let id = self.compile_function(decls, ep_decl, Some(instance))?; func_ids.push((ep_name, id)); } @@ -208,6 +210,7 @@ impl JIT { &mut self, decls: &DeclTable, decl: &FuncDecl, + instance: Option, ) -> Result { if decl.body.is_none() { return Err(format!("function '{}' has no body", decl.name)); @@ -254,11 +257,13 @@ impl JIT { // Now that compilation is finished, we can clear out the context state. self.module.clear_context(&mut self.ctx); - self.defined_functions.insert(decl.name); + if let Some(instance) = instance { + self.defined_functions.insert(instance); + } // Compile lambda functions extracted from this function's body. for lambda_decl in pending_lambdas { - self.compile_function(decls, &lambda_decl)?; + self.compile_function(decls, &lambda_decl, None)?; } // Compile any called functions that haven't already been defined. @@ -267,20 +272,16 @@ impl JIT { continue; } - let found = decls.find(name); - assert!(!found.is_empty(), "called function '{}' not found", name); - let decl = if let Decl::Func(d) = &found[0] { - d - } else { - panic!() - }; + let decl = decls + .function_instance(name) + .expect("checked callee instance"); // Skip builtins — they have no body and are called via raw function pointers. if decl.body.is_none() { continue; } - self.compile_function(decls, decl)?; + self.compile_function(decls, decl, Some(name))?; } Ok(id) @@ -289,15 +290,15 @@ impl JIT { fn declare_globals(&mut self, decls: &DeclTable) { // Reserve the first CANCEL_FLAG_RESERVED bytes for the cancel flag at offset 0. let mut offset: i32 = CANCEL_FLAG_RESERVED; - for decl in &decls.decls { + for (instance, decl) in decls.storage_instances() { match decl { - Decl::Global { name, ty, .. } => { - self.globals.insert(*name, offset); + Decl::Global { ty, .. } => { + self.globals.insert(instance, offset); offset += ty.size(decls) as i32; } Decl::Func(f) if f.is_extern => { // Extern functions get 16 bytes: {fn_ptr, context} - self.globals.insert(f.name, offset); + self.globals.insert(instance, offset); offset += 16; } _ => {} @@ -311,7 +312,7 @@ impl JIT { decls: &DeclTable, decl: &FuncDecl, has_globals: bool, - ) -> (HashSet, Vec) { + ) -> (HashSet, Vec) { // Translate into cranelift IR. // Create the builder to build a function. let mut builder = FunctionBuilder::new(&mut self.ctx.func, &mut self.builder_context); @@ -378,32 +379,35 @@ impl JIT { .builder .ins() .load(I64, MemFlags::new(), closure_ptr_val, (i * 8) as i32); - let var = trans.declare_variable(&cv.name.to_string(), I64); + let var = trans.declare_variable(cv, I64); trans.builder.def_var(var, var_ptr); - trans.variable_types.insert(cv.name.to_string(), cv.ty); + let declared = decl.arena.local(*cv).ty; + let storage = match *declared { + crate::Type::Reference(inner) => inner, + _ => declared, + }; + trans.variable_types.insert(*cv, storage); // NOT added to let_bindings → acts as a var binding (accessed through the pointer). } } // Add variables for the function parameters and define them with block param values. for (i, param) in decl.params.iter().enumerate() { - let param_ty = param.ty.expect("expected type"); + let param_ty = decl.arena.local(param.local).ty; let ty = if matches!(&*param_ty, crate::Type::Reference(_)) { I64 } else { param_ty.cranelift_type() }; - let var = trans.declare_variable(¶m.name, ty); + let var = trans.declare_variable(¶m.local, ty); trans.builder.def_var(var, block_params[i + param_idx]); if let crate::Type::Reference(inner) = &*param_ty { - trans.variable_types.insert(param.name.to_string(), *inner); - trans.let_bindings.remove(¶m.name.to_string()); + trans.variable_types.insert(param.local, *inner); + trans.let_bindings.remove(¶m.local); } else { - trans - .variable_types - .insert(param.name.to_string(), param_ty); + trans.variable_types.insert(param.local, param_ty); // Function parameters are like let bindings - they hold values directly. - trans.let_bindings.insert(param.name.to_string()); + trans.let_bindings.insert(param.local); } } @@ -540,27 +544,24 @@ fn fn_sig(module: &JITModule, from: crate::TypeID, to: crate::TypeID) -> Signatu struct FunctionTranslator<'a> { builder: FunctionBuilder<'a>, - variables: HashMap, - variable_types: HashMap, + variables: HashMap, + variable_types: HashMap, module: &'a mut JITModule, /// Next variable index. next_index: usize, - /// For generating code for generics. - current_instance: Instance, - - called_functions: HashSet, + called_functions: HashSet, /// Lambda functions extracted from this function body, to be compiled afterward. pending_lambdas: Vec, /// Tracks which variables are `let` bindings (hold values directly) /// vs `var` bindings (hold pointers to stack slots). - let_bindings: HashSet, + let_bindings: HashSet, /// Global variable offsets from base pointer. - globals: &'a HashMap, + globals: &'a HashMap, /// Base pointer for globals (passed as first param to main). globals_base: Option, @@ -584,14 +585,14 @@ struct FunctionTranslator<'a> { /// Names referenced from inside a lambda body in the function being /// translated. Such a variable is captured by address, so it has to live in /// memory even when its type would otherwise be kept in a register. - lambda_referenced: HashSet, + lambda_referenced: HashSet, } impl<'a> FunctionTranslator<'a> { fn new( builder: FunctionBuilder<'a>, module: &'a mut JITModule, - globals: &'a HashMap, + globals: &'a HashMap, globals_base: Option, output_ptr: Option, lambda_counter: &'a mut usize, @@ -603,7 +604,6 @@ impl<'a> FunctionTranslator<'a> { variable_types: HashMap::new(), module, next_index: 0, - current_instance: Instance::new(), called_functions: HashSet::new(), pending_lambdas: Vec::new(), let_bindings: HashSet::new(), @@ -620,7 +620,7 @@ impl<'a> FunctionTranslator<'a> { fn translate_fn(&mut self, decl: &FuncDecl, decls: &DeclTable) -> Value { self.elidable_lets = crate::copy_elision::elidable_let_copies(decl); - self.lambda_referenced = names_referenced_in_lambdas(decl); + self.lambda_referenced = decl.captured_locals(); // With --no-recursion the safety checker has proved the call graph // is a DAG, so the counter traffic is unnecessary. @@ -750,7 +750,7 @@ impl<'a> FunctionTranslator<'a> { /// Check that a value about to be passed to an if-else merge block matches /// the type the merge block param was declared with. /// - /// The two are derived independently: the param from `decl.types`, the + /// The two are derived independently: the param from the checked result type, the /// value from codegen. They disagree when an expression kind is typed /// non-void by the checker but returns a placeholder here (see issue #22, /// where a `var` declaration was typed f32 but yielded `iconst(I32, 0)`). @@ -777,26 +777,21 @@ impl<'a> FunctionTranslator<'a> { fn translate_lvalue(&mut self, expr: ExprID, decl: &FuncDecl, decls: &DeclTable) -> Value { match &decl.arena[expr] { - Expr::Id(name) => { - if let Some(variable) = self.variables.get(&**name) { - self.builder.use_var(*variable) - } else if let Some(&offset) = self.globals.get(name) { - // Global variable - compute address from base + offset. + Expr::Id(_) => match decl.arena.reference(expr) { + Some(Reference::Local(local)) => self.builder.use_var(self.variables[local]), + Some(Reference::Instance(instance)) => { + let offset = self.globals[instance]; let base = self.globals_base.expect("globals_base not set"); self.builder.ins().iadd_imm(base, offset as i64) - } else { - panic!( - "JIT: unknown lvalue variable {:?} (not local or global)", - name - ); } - } + reference => panic!("unresolved checked reference: {:?}", reference), + }, Expr::Field(lhs, name) => { - let lhs_ty = decl.types[*lhs]; + let lhs_ty = decl.arena.ty(*lhs); let lhs_value = self.translate_lvalue(*lhs, decl, decls); if let crate::Type::Name(struct_name, type_args) = &*lhs_ty { let struct_decl = decls.find(*struct_name); - if let crate::Decl::Struct(s) = &struct_decl[0] { + if let Decl::Struct(s) = &struct_decl[0] { let inst: crate::Instance = s .typevars .iter() @@ -817,7 +812,7 @@ impl<'a> FunctionTranslator<'a> { } } Expr::ArrayIndex(lhs, rhs) => { - let lhs_ty = decl.types[*lhs]; + let lhs_ty = decl.arena.ty(*lhs); let lhs_val = self.translate_lvalue(*lhs, decl, decls); let rhs_val = self.translate_expr(*rhs, decl, decls); let (elem_ty, is_sl) = match &*lhs_ty { @@ -855,46 +850,45 @@ impl<'a> FunctionTranslator<'a> { Expr::Int(imm, _) => self .builder .ins() - .iconst(decl.types[expr].cranelift_type(), *imm), + .iconst(decl.arena.ty(expr).cranelift_type(), *imm), Expr::Real(s, _) => { let val: f64 = s.parse().expect("invalid float literal"); - match &*decl.types[expr] { + match &*decl.arena.ty(expr) { crate::Type::Float32 => self.builder.ins().f32const(val as f32), _ => self.builder.ins().f64const(val), } } Expr::Char(c) => self.builder.ins().iconst(I8, *c as i64), - Expr::Id(name) => { - let ty = &decl.types[expr]; - if let Some(variable) = self.variables.get(&**name) { - let val = self.builder.use_var(*variable); - // Let bindings hold values directly; var bindings hold pointers. - // f32x4 is pointer-typed but lives in a register, so a var - // binding holding one still has to be loaded. - if self.let_bindings.contains(&**name) || is_indirect(*ty) { + Expr::Id(_) => match decl.arena.reference(expr) { + Some(Reference::Local(local)) => { + let ty = decl.arena.ty(expr); + let val = self.builder.use_var(self.variables[local]); + if self.let_bindings.contains(local) || is_indirect(ty) { val } else { self.builder .ins() .load(ty.cranelift_type(), MemFlags::new(), val, 0) } - } else if let Some(&offset) = self.globals.get(name) { - // Global variable - compute address from base + offset. - let base = self.globals_base.expect("globals_base not set"); - let addr = self.builder.ins().iadd_imm(base, offset as i64); - // Composite types (arrays, structs) are pointer-represented: - // return the address, don't load. f32x4 is a value type — load it. - if is_indirect(*ty) { - addr + } + Some(Reference::Instance(instance)) => { + let ty = decl.arena.ty(expr); + if let Some(&offset) = self.globals.get(instance) { + let base = self.globals_base.expect("globals_base not set"); + let addr = self.builder.ins().iadd_imm(base, offset as i64); + if is_indirect(ty) { + addr + } else { + self.builder + .ins() + .load(ty.cranelift_type(), MemFlags::new(), addr, 0) + } } else { - self.builder - .ins() - .load(ty.cranelift_type(), MemFlags::new(), addr, 0) + self.translate_func(*instance, &*ty, decls) } - } else { - self.translate_func(name, &*ty) } - } + reference => panic!("unresolved checked reference: {:?}", reference), + }, Expr::Binop(op, lhs_id, rhs_id) => { self.translate_binop(*op, *lhs_id, *rhs_id, decl, decls) } @@ -902,15 +896,17 @@ impl<'a> FunctionTranslator<'a> { Expr::Call(fn_id, arg_ids) => { // Determine if this is a builtin (assert/print) which has a raw fn_ptr // and no globals/closure parameters, vs a user function with a fat pointer. - let is_builtin = if let Expr::Id(name) = &decl.arena[*fn_id] { - is_builtin_name(name) - } else { - false - }; + let is_builtin = + if let Some(Reference::Instance(instance)) = decl.arena.reference(*fn_id) { + is_builtin_name(&decls.instance_name(*instance)) + } else { + false + }; // f32x4 constructor and splat — emit inline vector construction. - if let Expr::Id(name) = &decl.arena[*fn_id] { - if **name == "f32x4" && arg_ids.len() == 4 { + if let Some(Reference::Instance(instance)) = decl.arena.reference(*fn_id) { + let name = decls.instance_name(*instance); + if *name == "f32x4" && arg_ids.len() == 4 { let x = self.translate_expr(arg_ids[0], decl, decls); let y = self.translate_expr(arg_ids[1], decl, decls); let z = self.translate_expr(arg_ids[2], decl, decls); @@ -921,13 +917,13 @@ impl<'a> FunctionTranslator<'a> { let vec = self.builder.ins().insertlane(vec, w, 3); return vec; } - if **name == "f32x4_splat" && arg_ids.len() == 1 { + if *name == "f32x4_splat" && arg_ids.len() == 1 { let x = self.translate_expr(arg_ids[0], decl, decls); return self.builder.ins().splat(F32X4, x); } } - if let crate::Type::Func(from, to) = *(decl.types[*fn_id]) { + if let crate::Type::Func(from, to) = *(decl.arena.ty(*fn_id)) { // If return type is pointer, allocate stack space and pass as first arg. let output_slot = if returns_via_pointer(to) { let size = to.size(decls) as u32; @@ -947,9 +943,8 @@ impl<'a> FunctionTranslator<'a> { // Use the declaration (not the solved call-site type) because // the solver may retain Array types where the callee expects Slice. let param_types: Vec = - if let Expr::Id(callee_name) = &decl.arena[*fn_id] { - let callee_decls = decls.find(*callee_name); - if let Some(crate::Decl::Func(f)) = callee_decls.first() { + if let Some(Reference::Instance(callee)) = decl.arena.reference(*fn_id) { + if let Some(f) = decls.function_instance(*callee) { f.param_types() } else if let crate::Type::Tuple(pts) = &*from { pts.clone() @@ -963,21 +958,22 @@ impl<'a> FunctionTranslator<'a> { }; // Check if this is a math builtin that can use a direct call. - let math_sym = if let Expr::Id(name) = &decl.arena[*fn_id] { - math_builtin_symbol(name) - } else { - None - }; + let math_sym = + if let Some(Reference::Instance(instance)) = decl.arena.reference(*fn_id) { + math_builtin_symbol(&decls.instance_name(*instance)) + } else { + None + }; // Check if this is an extern function call. - let is_extern_fn = if let Expr::Id(callee_name) = &decl.arena[*fn_id] { - let callee_decls = decls.find(*callee_name); - callee_decls - .first() - .map_or(false, |d| matches!(d, crate::Decl::Func(f) if f.is_extern)) - } else { - false - }; + let is_extern_fn = + if let Some(Reference::Instance(callee)) = decl.arena.reference(*fn_id) { + decls + .function_instance(*callee) + .is_some_and(|function| function.is_extern) + } else { + false + }; let call = if let Some(sym) = math_sym { // Math builtin: use a direct call via declared function. @@ -998,18 +994,16 @@ impl<'a> FunctionTranslator<'a> { self.builder.ins().call(local_callee, &args) } else if is_extern_fn { // Extern function: indirect call through {fn_ptr, context} in globals. - let callee_name = if let Expr::Id(n) = &decl.arena[*fn_id] { - *n - } else { - unreachable!() - }; - let callee_decl = { - let decls_found = decls.find(callee_name); - match decls_found.first() { - Some(crate::Decl::Func(f)) => f.clone(), - _ => unreachable!(), - } - }; + let callee_name = + if let Some(Reference::Instance(n)) = decl.arena.reference(*fn_id) { + *n + } else { + unreachable!() + }; + let callee_decl = decls + .function_instance(callee_name) + .expect("extern instance") + .clone(); let globals_offset = *self .globals .get(&callee_name) @@ -1029,7 +1023,7 @@ impl<'a> FunctionTranslator<'a> { let mut sig = self.module.make_signature(); sig.params.push(AbiParam::new(I64)); // context for param in &callee_decl.params { - let pty = param.ty.unwrap(); + let pty = callee_decl.arena.local(param.local).ty; if matches!(&*pty, crate::Type::Slice(_)) { sig.params.push(AbiParam::new(I64)); // data ptr sig.params.push(AbiParam::new(I32)); // len @@ -1045,7 +1039,7 @@ impl<'a> FunctionTranslator<'a> { let mut args = vec![context]; for (i, arg_id) in arg_ids.iter().enumerate() { - let param_ty = callee_decl.params[i].ty.unwrap(); + let param_ty = callee_decl.arena.local(callee_decl.params[i].local).ty; let arg_val = if matches!(&*param_ty, crate::Type::Reference(_)) { self.translate_lvalue(*arg_id, decl, decls) } else { @@ -1094,8 +1088,8 @@ impl<'a> FunctionTranslator<'a> { // assert needs the globals pointer as the first arg so // it can write trap_reason and longjmp on failure. let is_assert = matches!( - &decl.arena[*fn_id], - Expr::Id(name) if **name == "assert" + decl.arena.reference(*fn_id), + Some(Reference::Instance(instance)) if *decls.instance_name(*instance) == "assert" ); let f = self.translate_expr(*fn_id, decl, decls); let mut args = vec![]; @@ -1168,12 +1162,13 @@ impl<'a> FunctionTranslator<'a> { } else { panic!( "JIT call: expected function type, got {:?}", - decl.types[*fn_id] + decl.arena.ty(*fn_id) ); } } - Expr::Let(name, init, _) => { - let ty = &decl.types[expr]; + Expr::Let(_, init, _) => { + let name = decl.arena.binder(expr); + let ty = &decl.arena.local(name).ty; let init_val = self.translate_expr(*init, decl, decls); let init_val = self.wrap_for_expected_slice(init_val, *ty, *init, decl, decls); @@ -1186,7 +1181,7 @@ impl<'a> FunctionTranslator<'a> { && !self.elidable_lets.contains(&expr) && sz > 0 { - let var = self.declare_variable(name, I64); + let var = self.declare_variable(&name, I64); let slot = self.builder.create_sized_stack_slot(StackSlotData { kind: StackSlotKind::ExplicitSlot, size: sz, @@ -1196,33 +1191,33 @@ impl<'a> FunctionTranslator<'a> { let addr = self.builder.ins().stack_addr(I64, slot, 0); self.builder.def_var(var, addr); self.gen_copy(*ty, addr, init_val, decls); - self.variable_types.insert(name.to_string(), *ty); + self.variable_types.insert(name, *ty); // The binding owns a stack slot now, exactly like a `var`, // so it must not be treated as holding a value directly. - self.let_bindings.remove(&name.to_string()); + self.let_bindings.remove(&name); return addr; } - let var = self.declare_variable(name, ty.cranelift_type()); + let var = self.declare_variable(&name, ty.cranelift_type()); self.builder.def_var(var, init_val); - self.variable_types.insert(name.to_string(), *ty); - self.let_bindings.insert(name.to_string()); + self.variable_types.insert(name, *ty); + self.let_bindings.insert(name); init_val } - Expr::Var(name, init, _) => { - let ty = &decl.types[expr]; - // Remove from let_bindings in case this var shadows a let binding - // (e.g., a for-loop counter with the same name). - self.let_bindings.remove(&name.to_string()); + Expr::Var(_, init, _) => { + let name = decl.arena.binder(expr); + let ty = &decl.arena.local(name).ty; + // This storage is addressed through a pointer. + self.let_bindings.remove(&name); // f32x4: treat as value type (like a let binding) so it lives in // a Cranelift variable (F32X4) rather than a pointer to a stack slot. // A variable a lambda captures is shared by address, so it has to // stay in memory for writes on either side to be visible. if matches!(**ty, crate::types::Type::Float32x4) - && !self.lambda_referenced.contains(&name.to_string()) + && !self.lambda_referenced.contains(&name) { - let var = self.declare_variable(name, F32X4); + let var = self.declare_variable(&name, F32X4); let init_val = if let Some(init_id) = init { self.translate_expr(*init_id, decl, decls) } else { @@ -1230,13 +1225,13 @@ impl<'a> FunctionTranslator<'a> { self.builder.ins().splat(F32X4, zero) }; self.builder.def_var(var, init_val); - self.variable_types.insert(name.to_string(), *ty); - self.let_bindings.insert(name.to_string()); + self.variable_types.insert(name, *ty); + self.let_bindings.insert(name); return init_val; } - let var = self.declare_variable(name, I64); - self.variable_types.insert(name.to_string(), *ty); + let var = self.declare_variable(&name, I64); + self.variable_types.insert(name, *ty); let sz = ty.size(decls) as u32; if sz == 0 { @@ -1270,7 +1265,7 @@ impl<'a> FunctionTranslator<'a> { self.builder.ins().iconst(I32, 0) } Expr::StructLit(struct_name, fields) => { - let ty = &decl.types[expr]; + let ty = &decl.arena.ty(expr); let sz = ty.size(decls) as u32; let slot = self.builder.create_sized_stack_slot(StackSlotData { kind: StackSlotKind::ExplicitSlot, @@ -1283,7 +1278,7 @@ impl<'a> FunctionTranslator<'a> { if let crate::Type::Name(_, type_args) = &**ty { let struct_decl = decls.find(*struct_name); - if let crate::Decl::Struct(s) = &struct_decl[0] { + if let Decl::Struct(s) = &struct_decl[0] { let inst: crate::Instance = s .typevars .iter() @@ -1293,7 +1288,7 @@ impl<'a> FunctionTranslator<'a> { for (fname, fval) in fields { let val = self.translate_expr(*fval, decl, decls); let off = s.field_offset(fname, decls, &inst); - let field_ty = &decl.types[*fval]; + let field_ty = &decl.arena.ty(*fval); let field_addr = self.builder.ins().iadd_imm(addr, off as i64); if is_indirect(*field_ty) { self.gen_copy(*field_ty, field_addr, val, decls); @@ -1308,7 +1303,7 @@ impl<'a> FunctionTranslator<'a> { addr } Expr::Field(lhs, name) => { - let lhs_ty = decl.types[*lhs]; + let lhs_ty = decl.arena.ty(*lhs); // Handle array.len / slice.len. if **name == "len" { match &*lhs_ty { @@ -1338,7 +1333,7 @@ impl<'a> FunctionTranslator<'a> { let lhs_val = self.translate_expr(*lhs, decl, decls); if let crate::Type::Name(struct_name, type_args) = &*lhs_ty { let struct_decl = decls.find(*struct_name); - if let crate::Decl::Struct(s) = &struct_decl[0] { + if let Decl::Struct(s) = &struct_decl[0] { let inst: crate::Instance = s .typevars .iter() @@ -1346,7 +1341,7 @@ impl<'a> FunctionTranslator<'a> { .map(|(tv, ty)| (crate::types::mk_type(crate::Type::Var(*tv)), *ty)) .collect(); let off = s.field_offset(name, decls, &inst); - let field_ty = &decl.types[expr]; + let field_ty = &decl.arena.ty(expr); // Arrays are stored inline, so return the address of the field if is_indirect(*field_ty) { let off_val = self.builder.ins().iconst(I64, off as i64); @@ -1370,7 +1365,7 @@ impl<'a> FunctionTranslator<'a> { for i in 0..index { off += elem_types[i].size(decls) as i32; } - let field_ty = &decl.types[expr]; + let field_ty = &decl.arena.ty(expr); if is_indirect(*field_ty) { let off_val = self.builder.ins().iconst(I64, off as i64); self.builder.ins().iadd(lhs_val, off_val) @@ -1388,13 +1383,13 @@ impl<'a> FunctionTranslator<'a> { } } Expr::ArrayIndex(lhs, rhs) => { - let lhs_ty = decl.types[*lhs]; + let lhs_ty = decl.arena.ty(*lhs); // f32x4 element extraction if matches!(*lhs_ty, crate::types::Type::Float32x4) { let vec = self.translate_expr(*lhs, decl, decls); // Try constant lane index - if let Expr::Int(n, _) = &decl.arena.exprs[*rhs] { + if let Expr::Int(n, _) = &decl.arena[*rhs] { return self.builder.ins().extractlane(vec, *n as u8); } // Dynamic index: spill to stack and load @@ -1434,7 +1429,7 @@ impl<'a> FunctionTranslator<'a> { .imul_imm(rhs_val, elem_ty.size(decls) as i64); let off = self.builder.ins().uextend(I64, off); let p = self.builder.ins().iadd(data_ptr, off); - let result_ty = decl.types[expr]; + let result_ty = decl.arena.ty(expr); if is_indirect(result_ty) { // Composite types (arrays, structs, tuples) are represented // as pointers — return the address directly. @@ -1450,7 +1445,7 @@ impl<'a> FunctionTranslator<'a> { .map(|e| self.translate_expr(*e, decl, decls)) .collect(); - let ty = decl.types[expr]; + let ty = decl.arena.ty(expr); if let crate::Type::Array(elem_ty, _) = &*ty { let element_size = elem_ty.size(decls) as u32; @@ -1477,14 +1472,14 @@ impl<'a> FunctionTranslator<'a> { } else { panic!( "JIT array literal: expected array type, got {:?}", - decl.types[expr] + decl.arena.ty(expr) ); } } Expr::Array(value_expr, _size_expr) => { // Fill-array expression: [value; size], e.g. [0; 5] let fill_value = self.translate_expr(*value_expr, decl, decls); - let ty = decl.types[expr]; + let ty = decl.arena.ty(expr); if let crate::Type::Array(elem_ty, sz) = &*ty { let count = sz.known(); @@ -1509,7 +1504,7 @@ impl<'a> FunctionTranslator<'a> { } else { panic!( "JIT array fill: expected array type, got {:?}", - decl.types[expr] + decl.arena.ty(expr) ); } } @@ -1517,11 +1512,7 @@ impl<'a> FunctionTranslator<'a> { if exprs.is_empty() { self.builder.ins().iconst(I32, 0) } else { - // Save variable scope — variables declared inside this block - // shadow outer names only for the duration of the block. - let saved_vars = self.variables.clone(); - let saved_types = self.variable_types.clone(); - let saved_lets = self.let_bindings.clone(); + // Binding identities remain unambiguous across block boundaries. let mut result = None; for expr in exprs { result = Some(self.translate_expr(*expr, decl, decls)); @@ -1530,9 +1521,6 @@ impl<'a> FunctionTranslator<'a> { break; } } - self.variables = saved_vars; - self.variable_types = saved_types; - self.let_bindings = saved_lets; result.unwrap() } } @@ -1548,13 +1536,13 @@ impl<'a> FunctionTranslator<'a> { // `merge_ty` is Some exactly when it does, so it doubles as the // "produces a value" flag — keeping a separate boolean around // risks the two drifting apart. - let result_ty = decl.types[expr]; + let result_ty = decl.arena.ty(expr); let produces_value = match else_id { Some(else_expr_id) => { !matches!( &*result_ty, crate::Type::Void | crate::Type::Anon(_) | crate::Type::Var(_) - ) && result_ty == decl.types[*else_expr_id] + ) && result_ty == decl.arena.ty(*else_expr_id) } None => false, }; @@ -1617,10 +1605,7 @@ impl<'a> FunctionTranslator<'a> { } } Expr::For { - var, - start, - end, - body, + start, end, body, .. } => { // Evaluate start and end values. Both are outside the loop // variable's scope, so they still see any outer binding of the @@ -1628,18 +1613,15 @@ impl<'a> FunctionTranslator<'a> { let start_val = self.translate_expr(*start, decl, decls); let end_val = self.translate_expr(*end, decl, decls); - // The loop variable is scoped to the loop: save the name-keyed - // state so an outer binding it shadows comes back at loop exit. - let saved_vars = self.variables.clone(); - let saved_types = self.variable_types.clone(); - let saved_lets = self.let_bindings.clone(); + // The checked loop binding has its own local identity. + let var = decl.arena.binder(expr); // Create a variable for the loop counter. - let loop_var = self.declare_variable(var, I32); + let loop_var = self.declare_variable(&var, I32); self.builder.def_var(loop_var, start_val); self.variable_types - .insert(var.to_string(), crate::types::mk_type(crate::Type::Int32)); - self.let_bindings.insert(var.to_string()); + .insert(var, crate::types::mk_type(crate::Type::Int32)); + self.let_bindings.insert(var); // Create blocks for header, body, latch, and exit. let header_block = self.builder.create_block(); @@ -1696,10 +1678,6 @@ impl<'a> FunctionTranslator<'a> { self.builder.switch_to_block(exit_block); self.builder.seal_block(exit_block); - self.variables = saved_vars; - self.variable_types = saved_types; - self.let_bindings = saved_lets; - self.builder.ins().iconst(I32, 0) } Expr::Assume(_) => { @@ -1708,7 +1686,7 @@ impl<'a> FunctionTranslator<'a> { } Expr::Return(expr_id) => { let result = self.translate_expr(*expr_id, decl, decls); - let ret_ty = decl.types[*expr_id]; + let ret_ty = decl.arena.ty(*expr_id); if !self.no_recursion { self.emit_call_depth_release(); @@ -1811,7 +1789,7 @@ impl<'a> FunctionTranslator<'a> { .map(|e| self.translate_expr(*e, decl, decls)) .collect(); - let ty = decl.types[expr]; + let ty = decl.arena.ty(expr); if let crate::Type::Tuple(elem_types) = &*ty { // Calculate total size. @@ -1839,37 +1817,20 @@ impl<'a> FunctionTranslator<'a> { } else { panic!( "JIT tuple expression: expected tuple type, got {:?}", - decl.types[expr] + decl.arena.ty(expr) ); } } - Expr::Lambda { params, body } => { - let lambda_ty = decl.types[expr]; + Expr::Lambda { .. } => { + let lambda_ty = decl.arena.ty(expr); if let crate::Type::Func(dom, rng) = *lambda_ty { - if let crate::Type::Tuple(param_types) = &*dom { + if let crate::Type::Tuple(_) = &*dom { let id = *self.lambda_counter; *self.lambda_counter += 1; let lambda_name = Name::new(format!("__lambda_{}", id)); - let lambda_params: Vec = params - .iter() - .zip(param_types.iter()) - .map(|(p, ty)| Param { - name: p.name, - ty: Some(*ty), - }) - .collect(); - - // Compute free variables captured from the enclosing scope. - let param_names: std::collections::HashSet = - params.iter().map(|p| p.name.to_string()).collect(); - let free_vars = collect_free_var_names( - *body, - &decl.arena, - ¶m_names, - &self.variables, - &decl.types, - ); + let lambda_decl = decl.extract_lambda(expr, lambda_name); + let free_vars = &lambda_decl.closure_vars; // Allocate a closure struct on the stack; each slot holds the address // of one captured variable's storage. @@ -1882,8 +1843,9 @@ impl<'a> FunctionTranslator<'a> { key: None, }); let closure_addr = self.builder.ins().stack_addr(I64, slot, 0); - for (i, (name, ty)) in free_vars.iter().enumerate() { - let var_ptr = self.get_var_address(name, *ty, decls); + for (i, local) in free_vars.iter().enumerate() { + let var_ptr = + self.get_var_address(local, decl.arena.local(*local).ty, decls); self.builder.ins().store( MemFlags::new(), var_ptr, @@ -1896,30 +1858,6 @@ impl<'a> FunctionTranslator<'a> { self.builder.ins().iconst(I64, 0) }; - let closure_vars: Vec = free_vars - .iter() - .map(|(name, ty)| ClosureVar { - name: Name::new(name.clone()), - ty: *ty, - }) - .collect(); - - let lambda_decl = FuncDecl { - name: lambda_name, - typevars: vec![], - size_vars: vec![], - params: lambda_params, - body: Some(*body), - ret: rng, - constraints: vec![], - requires: vec![], - loc: decl.loc, - arena: decl.arena.clone(), - types: decl.types.clone(), - closure_vars, - is_extern: false, - }; - self.pending_lambdas.push(lambda_decl); let mut sig = fn_sig(&self.module, dom, rng); @@ -1957,7 +1895,7 @@ impl<'a> FunctionTranslator<'a> { } Expr::AsTy(expr_id, target_ty) => { let val = self.translate_expr(*expr_id, decl, decls); - let src_ty = decl.types[*expr_id]; + let src_ty = decl.arena.ty(*expr_id); match (&*src_ty, &**target_ty) { (crate::Type::Int32, crate::Type::Float32) => { self.builder.ins().fcvt_from_sint(F32, val) @@ -1995,11 +1933,10 @@ impl<'a> FunctionTranslator<'a> { } } Expr::Enum(case_name) => { - let index = if let crate::Type::Name(enum_name, _) = &*decl.types[expr] { + let index = if let crate::Type::Name(enum_name, _) = &*decl.arena.ty(expr) { let enum_decls = decls.find(*enum_name); - if let Some(crate::Decl::Enum { cases, .. }) = enum_decls - .iter() - .find(|d| matches!(d, crate::Decl::Enum { .. })) + if let Some(Decl::Enum { cases, .. }) = + enum_decls.iter().find(|d| matches!(d, Decl::Enum { .. })) { cases.iter().position(|c| c == case_name).unwrap_or(0) as i64 } else { @@ -2053,12 +1990,12 @@ impl<'a> FunctionTranslator<'a> { unop: Unop, arg_id: ExprID, decl: &FuncDecl, - decls: &crate::DeclTable, + decls: &DeclTable, ) -> Value { let v = self.translate_expr(arg_id, decl, decls); match unop { Unop::Neg => { - let t = decl.types[arg_id]; + let t = decl.arena.ty(arg_id); match *t { crate::types::Type::Float32 | crate::types::Type::Float64 @@ -2076,7 +2013,7 @@ impl<'a> FunctionTranslator<'a> { /// Compute the natural alignment of a type for use with Cranelift memory /// operations. Cranelift requires that alignment ≤ greatest_divisible_power_of_two(size), /// so we base alignment on the element type rather than the total size. - fn type_align(t: &crate::TypeID, decls: &crate::DeclTable) -> u8 { + fn type_align(t: &crate::TypeID, decls: &DeclTable) -> u8 { let size = t.size(decls) as u64; if size == 0 { return 1; @@ -2096,7 +2033,7 @@ impl<'a> FunctionTranslator<'a> { crate::Type::Name(name, vars) => { let decl = decls.find(*name); if decl.len() == 1 { - if let crate::decl::Decl::Struct(sdecl) = &decl[0] { + if let Decl::Struct(sdecl) = &decl[0] { let inst: crate::Instance = sdecl .typevars .iter() @@ -2133,7 +2070,7 @@ impl<'a> FunctionTranslator<'a> { slot: cranelift::codegen::ir::StackSlot, offset: i32, value: Value, - decls: &crate::DeclTable, + decls: &DeclTable, ) { if is_indirect(elem_ty) { let dst = self.builder.ins().iadd_imm(addr, offset as i64); @@ -2143,7 +2080,7 @@ impl<'a> FunctionTranslator<'a> { } } - fn gen_copy(&mut self, t: crate::TypeID, dst: Value, src: Value, decls: &crate::DeclTable) { + fn gen_copy(&mut self, t: crate::TypeID, dst: Value, src: Value, decls: &DeclTable) { if is_indirect(t) { let size = t.size(decls) as u64; let align = Self::type_align(&t, decls); @@ -2162,7 +2099,7 @@ impl<'a> FunctionTranslator<'a> { } } - fn gen_zero(&mut self, t: crate::TypeID, dst: Value, decls: &crate::DeclTable) { + fn gen_zero(&mut self, t: crate::TypeID, dst: Value, decls: &DeclTable) { let size = t.size(decls) as u32; let zero = self.builder.ins().iconst(I8, 0); // Store zero byte-by-byte for the size of the type @@ -2173,13 +2110,7 @@ impl<'a> FunctionTranslator<'a> { } } - fn gen_eq( - &mut self, - t: crate::TypeID, - dst: Value, - src: Value, - decls: &crate::DeclTable, - ) -> Value { + fn gen_eq(&mut self, t: crate::TypeID, dst: Value, src: Value, decls: &DeclTable) -> Value { if let crate::types::Type::Slice(elem) = &*t { let elem_size = elem.size(decls) as i64; return self.gen_slice_eq(dst, src, elem_size); @@ -2239,13 +2170,13 @@ impl<'a> FunctionTranslator<'a> { lhs_id: ExprID, rhs_id: ExprID, decl: &FuncDecl, - decls: &crate::DeclTable, + decls: &DeclTable, ) -> Value { match binop { Binop::Plus => { let lhs = self.translate_expr(lhs_id, decl, decls); let rhs = self.translate_expr(rhs_id, decl, decls); - let t = decl.types[lhs_id]; + let t = decl.arena.ty(lhs_id); match *t { crate::types::Type::Int32 @@ -2261,7 +2192,7 @@ impl<'a> FunctionTranslator<'a> { Binop::Minus => { let lhs = self.translate_expr(lhs_id, decl, decls); let rhs = self.translate_expr(rhs_id, decl, decls); - let t = decl.types[lhs_id]; + let t = decl.arena.ty(lhs_id); match *t { crate::types::Type::Int32 @@ -2277,7 +2208,7 @@ impl<'a> FunctionTranslator<'a> { Binop::Mult => { let lhs = self.translate_expr(lhs_id, decl, decls); let rhs = self.translate_expr(rhs_id, decl, decls); - let t = decl.types[lhs_id]; + let t = decl.arena.ty(lhs_id); match *t { crate::types::Type::Int32 | crate::types::Type::UInt32 @@ -2292,7 +2223,7 @@ impl<'a> FunctionTranslator<'a> { Binop::Div => { let lhs = self.translate_expr(lhs_id, decl, decls); let rhs = self.translate_expr(rhs_id, decl, decls); - let t = decl.types[lhs_id]; + let t = decl.arena.ty(lhs_id); match *t { crate::types::Type::Int32 | crate::types::Type::UInt32 @@ -2307,7 +2238,7 @@ impl<'a> FunctionTranslator<'a> { Binop::Mod => { let lhs = self.translate_expr(lhs_id, decl, decls); let rhs = self.translate_expr(rhs_id, decl, decls); - let t = decl.types[lhs_id]; + let t = decl.arena.ty(lhs_id); match *t { crate::types::Type::Int32 | crate::types::Type::UInt32 @@ -2321,7 +2252,7 @@ impl<'a> FunctionTranslator<'a> { // f32x4 field assignment: v.x = val → insertlane if let Expr::Field(vec_id, field_name) = &decl.arena[lhs_id] { let (vec_id, field_name) = (*vec_id, *field_name); - let vec_ty = decl.types[vec_id]; + let vec_ty = decl.arena.ty(vec_id); if matches!(*vec_ty, crate::types::Type::Float32x4) { let lane: u8 = match &**field_name { "x" | "r" => 0, @@ -2339,11 +2270,11 @@ impl<'a> FunctionTranslator<'a> { // f32x4 element assignment: v[i] = val → insertlane if let Expr::ArrayIndex(vec_id, idx_id) = &decl.arena[lhs_id] { let (vec_id, idx_id) = (*vec_id, *idx_id); - let vec_ty = decl.types[vec_id]; + let vec_ty = decl.arena.ty(vec_id); if matches!(*vec_ty, crate::types::Type::Float32x4) { // The lane index is in 0..4: the safety checker proves it, // and no backend checks it at runtime. - if let Expr::Int(n, _) = &decl.arena.exprs[idx_id] { + if let Expr::Int(n, _) = &decl.arena[idx_id] { let lane = lane_index(*n); let storage = self.f32x4_storage(vec_id, decl, decls); let rhs = self.translate_expr(rhs_id, decl, decls); @@ -2367,8 +2298,8 @@ impl<'a> FunctionTranslator<'a> { self.builder.ins().store(MemFlags::new(), rhs, ptr, 0); return rhs; } - if let Expr::Id(name) = &decl.arena[lhs_id] { - if let Some(&var) = self.variables.get(&**name) { + if let Some(Reference::Local(name)) = decl.arena.reference(lhs_id) { + if let Some(&var) = self.variables.get(name) { self.builder.def_var(var, rhs); return rhs; } @@ -2387,13 +2318,13 @@ impl<'a> FunctionTranslator<'a> { Binop::Equal => { let lhs = self.translate_expr(lhs_id, decl, decls); let rhs = self.translate_expr(rhs_id, decl, decls); - let t = decl.types[lhs_id]; + let t = decl.arena.ty(lhs_id); self.gen_eq(t, lhs, rhs, decls) } Binop::NotEqual => { let lhs = self.translate_expr(lhs_id, decl, decls); let rhs = self.translate_expr(rhs_id, decl, decls); - let t = decl.types[lhs_id]; + let t = decl.arena.ty(lhs_id); match *t { crate::types::Type::Bool @@ -2416,7 +2347,7 @@ impl<'a> FunctionTranslator<'a> { Binop::Less => { let lhs = self.translate_expr(lhs_id, decl, decls); let rhs = self.translate_expr(rhs_id, decl, decls); - let t = decl.types[lhs_id]; + let t = decl.arena.ty(lhs_id); match *t { crate::types::Type::Int32 | crate::types::Type::Int8 => { @@ -2434,7 +2365,7 @@ impl<'a> FunctionTranslator<'a> { Binop::Greater => { let lhs = self.translate_expr(lhs_id, decl, decls); let rhs = self.translate_expr(rhs_id, decl, decls); - let t = decl.types[lhs_id]; + let t = decl.arena.ty(lhs_id); match *t { crate::types::Type::Int32 | crate::types::Type::Int8 => { @@ -2453,7 +2384,7 @@ impl<'a> FunctionTranslator<'a> { Binop::Leq => { let lhs = self.translate_expr(lhs_id, decl, decls); let rhs = self.translate_expr(rhs_id, decl, decls); - let t = decl.types[lhs_id]; + let t = decl.arena.ty(lhs_id); match *t { crate::types::Type::Int32 | crate::types::Type::Int8 => self @@ -2473,7 +2404,7 @@ impl<'a> FunctionTranslator<'a> { Binop::Geq => { let lhs = self.translate_expr(lhs_id, decl, decls); let rhs = self.translate_expr(rhs_id, decl, decls); - let t = decl.types[lhs_id]; + let t = decl.arena.ty(lhs_id); match *t { crate::types::Type::Int32 | crate::types::Type::Int8 => self @@ -2520,7 +2451,13 @@ impl<'a> FunctionTranslator<'a> { self.builder.ins().call_indirect(sref, f_ptr, &vec![value]); } - fn translate_func(&mut self, name: &Name, ty: &crate::Type) -> Value { + fn translate_func( + &mut self, + instance: InstanceId, + ty: &crate::Type, + decls: &DeclTable, + ) -> Value { + let name = &decls.instance_name(instance); if *name == Name::str("assert") { return self .builder @@ -2557,7 +2494,7 @@ impl<'a> FunctionTranslator<'a> { .expect("problem declaring function"); let local_callee = self.module.declare_func_in_func(callee, self.builder.func); - self.called_functions.insert(*name); + self.called_functions.insert(instance); let fn_ptr = self.builder.ins().func_addr(I64, local_callee); @@ -2592,17 +2529,21 @@ impl<'a> FunctionTranslator<'a> { /// a variable, or when the expression isn't a place at all. fn f32x4_storage(&mut self, expr: ExprID, decl: &FuncDecl, decls: &DeclTable) -> Option { match &decl.arena[expr] { - Expr::Id(name) => { - if let Some(&var) = self.variables.get(&**name) { - if self.let_bindings.contains(&**name) { - return None; + Expr::Id(_) => match decl.arena.reference(expr) { + Some(Reference::Local(local)) => { + if self.let_bindings.contains(local) { + None + } else { + Some(self.builder.use_var(self.variables[local])) } - return Some(self.builder.use_var(var)); } - let offset = *self.globals.get(name)?; - let base = self.globals_base.expect("globals_base not set"); - Some(self.builder.ins().iadd_imm(base, offset as i64)) - } + Some(Reference::Instance(instance)) => { + let offset = *self.globals.get(instance)?; + let base = self.globals_base.expect("globals_base not set"); + Some(self.builder.ins().iadd_imm(base, offset as i64)) + } + _ => None, + }, Expr::Field(_, _) | Expr::ArrayIndex(_, _) => { Some(self.translate_lvalue(expr, decl, decls)) } @@ -2627,8 +2568,8 @@ impl<'a> FunctionTranslator<'a> { self.builder.ins().store(MemFlags::new(), new_vec, ptr, 0); return; } - if let Expr::Id(name) = &decl.arena[vec_id] { - if let Some(&var) = self.variables.get(&**name) { + if let Some(Reference::Local(name)) = decl.arena.reference(vec_id) { + if let Some(&var) = self.variables.get(name) { let vec = self.builder.use_var(var); let new_vec = self.builder.ins().insertlane(vec, value, lane); self.builder.def_var(var, new_vec); @@ -2658,8 +2599,8 @@ impl<'a> FunctionTranslator<'a> { self.builder.ins().store(MemFlags::new(), value, addr, 0); return; } - if let Expr::Id(name) = &decl.arena[vec_id] { - if let Some(&var) = self.variables.get(&**name) { + if let Some(Reference::Local(name)) = decl.arena.reference(vec_id) { + if let Some(&var) = self.variables.get(name) { let vec = self.builder.use_var(var); let slot = self.builder.create_sized_stack_slot(StackSlotData { kind: StackSlotKind::ExplicitSlot, @@ -2685,16 +2626,17 @@ impl<'a> FunctionTranslator<'a> { /// Returns the address of a variable's storage for use in a closure struct. /// For var bindings the variable already holds a pointer; for let bindings /// a fresh stack slot is allocated and the value is copied into it. - fn get_var_address( - &mut self, - name: &str, - ty: crate::TypeID, - decls: &crate::DeclTable, - ) -> Value { + fn get_var_address(&mut self, name: &LocalId, ty: crate::TypeID, decls: &DeclTable) -> Value { if let Some(&var) = self.variables.get(name) { if self.let_bindings.contains(name) { // let binding: allocate a slot, copy the value in, capture the slot address. let val = self.builder.use_var(var); + // An aggregate value is already represented by its storage + // address; capturing it must keep that address, not spill the + // pointer as if it were the aggregate's first element. + if is_indirect(ty) { + return val; + } let sz = ty.size(decls) as u32; let slot = self.builder.create_sized_stack_slot(StackSlotData { kind: StackSlotKind::ExplicitSlot, @@ -2709,22 +2651,14 @@ impl<'a> FunctionTranslator<'a> { // var binding: variable holds a pointer to the stack slot directly. self.builder.use_var(var) } - } else if let Some(&offset) = self.globals.get(&Name::new(name.to_string())) { - let base = self.globals_base.expect("globals_base not set"); - self.builder.ins().iadd_imm(base, offset as i64) } else { - panic!("unknown variable in closure capture: {}", name) + panic!("unknown local in closure capture: {:?}", name) } } /// Wraps a sized array value in a slice fat pointer {data_ptr, len} on the stack. /// If the value is already a slice, returns it as-is. - fn wrap_as_slice( - &mut self, - val: Value, - actual_ty: crate::TypeID, - _decls: &crate::DeclTable, - ) -> Value { + fn wrap_as_slice(&mut self, val: Value, actual_ty: crate::TypeID, _decls: &DeclTable) -> Value { match &*actual_ty { crate::Type::Slice(_) => { // Already a slice fat pointer, pass through. @@ -2759,28 +2693,26 @@ impl<'a> FunctionTranslator<'a> { &self, expr: ExprID, decl: &FuncDecl, - decls: &crate::DeclTable, + decls: &DeclTable, ) -> crate::TypeID { - match &decl.arena.exprs[expr] { - Expr::Id(name) => self - .variable_types - .get(name.as_str()) - .copied() - .or_else(|| { - decls.find(*name).iter().find_map(|decl| { - if let crate::Decl::Global { ty, .. } = decl { - Some(*ty) - } else { - None - } - }) - }) - .unwrap_or(decl.types[expr]), + match &decl.arena[expr] { + Expr::Id(_) => match decl.arena.reference(expr) { + Some(Reference::Local(local)) => self + .variable_types + .get(local) + .copied() + .unwrap_or(decl.arena.local(*local).ty), + Some(Reference::Instance(instance)) => match decls.instance(*instance) { + Decl::Global { ty, .. } => *ty, + _ => decl.arena.ty(expr), + }, + _ => decl.arena.ty(expr), + }, Expr::ArrayIndex(arr_id, _) => match &*self.representation_type(*arr_id, decl, decls) { crate::Type::Array(elem, _) | crate::Type::Slice(elem) => *elem, - _ => decl.types[expr], + _ => decl.arena.ty(expr), }, - _ => decl.types[expr], + _ => decl.arena.ty(expr), } } @@ -2790,7 +2722,7 @@ impl<'a> FunctionTranslator<'a> { expected_ty: crate::TypeID, actual_expr: ExprID, decl: &FuncDecl, - decls: &crate::DeclTable, + decls: &DeclTable, ) -> Value { if matches!(&*expected_ty, crate::Type::Slice(_)) { let actual_ty = self.representation_type(actual_expr, decl, decls); @@ -2800,183 +2732,15 @@ impl<'a> FunctionTranslator<'a> { } } - fn declare_variable(&mut self, name: &String, ty: Type) -> Variable { - // Always create a fresh Cranelift variable — even if a variable with - // the same name exists, this declaration shadows it. + fn declare_variable(&mut self, name: &LocalId, ty: Type) -> Variable { + // Allocate storage for this checked binding. let var = self.builder.declare_var(ty); - self.variables.insert(name.into(), var); + self.variables.insert(*name, var); self.next_index += 1; var } } -/// Every name mentioned inside a lambda body anywhere in `decl`, including -/// nested lambdas. An over-approximation of what the function's lambdas -/// capture: a name shadowed by a lambda parameter is included too, which only -/// costs the enclosing variable its register representation. -fn names_referenced_in_lambdas(decl: &FuncDecl) -> HashSet { - decl.names_referenced_in_lambdas() - .iter() - .map(|n| n.to_string()) - .collect() -} - -/// Collect the names (and their types) of free variables referenced in `body` -/// that are present in `local_vars` but not in `exclude` (the lambda's own params). -/// Returns each name at most once. -fn collect_free_var_names( - body: crate::ExprID, - arena: &crate::ExprArena, - exclude: &std::collections::HashSet, - local_vars: &HashMap, - types: &[crate::TypeID], -) -> Vec<(String, crate::TypeID)> { - let mut result = Vec::new(); - let mut seen = std::collections::HashSet::new(); - collect_free_vars_rec( - body, - arena, - exclude, - local_vars, - types, - &mut result, - &mut seen, - ); - result -} - -fn collect_free_vars_rec( - expr: crate::ExprID, - arena: &crate::ExprArena, - exclude: &std::collections::HashSet, - local_vars: &HashMap, - types: &[crate::TypeID], - result: &mut Vec<(String, crate::TypeID)>, - seen: &mut std::collections::HashSet, -) { - match &arena[expr] { - Expr::TypeApp(_, _) => {} // Rewritten to Id by monomorphizer - Expr::Id(name) => { - let s = name.to_string(); - if local_vars.contains_key(&s) && !exclude.contains(&s) && !seen.contains(&s) { - result.push((s.clone(), types[expr])); - seen.insert(s); - } - } - Expr::Call(fn_id, args) => { - collect_free_vars_rec(*fn_id, arena, exclude, local_vars, types, result, seen); - for a in args { - collect_free_vars_rec(*a, arena, exclude, local_vars, types, result, seen); - } - } - Expr::Binop(_, lhs, rhs) => { - collect_free_vars_rec(*lhs, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec(*rhs, arena, exclude, local_vars, types, result, seen); - } - Expr::Unop(_, arg) => { - collect_free_vars_rec(*arg, arena, exclude, local_vars, types, result, seen); - } - Expr::Let(_, init, _) => { - collect_free_vars_rec(*init, arena, exclude, local_vars, types, result, seen); - } - Expr::Var(_, init, _) => { - if let Some(init_id) = init { - collect_free_vars_rec(*init_id, arena, exclude, local_vars, types, result, seen); - } - } - Expr::If(cond, then, else_) => { - collect_free_vars_rec(*cond, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec(*then, arena, exclude, local_vars, types, result, seen); - if let Some(e) = else_ { - collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen); - } - } - Expr::While(cond, body) => { - collect_free_vars_rec(*cond, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec(*body, arena, exclude, local_vars, types, result, seen); - } - Expr::For { - start, end, body, .. - } => { - collect_free_vars_rec(*start, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec(*end, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec(*body, arena, exclude, local_vars, types, result, seen); - } - Expr::Block(exprs) => { - for e in exprs { - collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen); - } - } - Expr::Return(e) | Expr::Assume(e) => { - collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen); - } - Expr::Field(e, _) => { - collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen); - } - Expr::ArrayIndex(arr, idx) => { - collect_free_vars_rec(*arr, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec(*idx, arena, exclude, local_vars, types, result, seen); - } - Expr::ArrayLiteral(elems) => { - for e in elems { - collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen); - } - } - Expr::Tuple(elems) => { - for e in elems { - collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen); - } - } - Expr::AsTy(e, _) => { - collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen); - } - Expr::Arena(e) => { - collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen); - } - Expr::Array(ty_expr, size_expr) => { - collect_free_vars_rec(*ty_expr, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec(*size_expr, arena, exclude, local_vars, types, result, seen); - } - Expr::Lambda { params, body } => { - // Nested lambda: add its params to the exclusion set. - let mut inner_exclude = exclude.clone(); - for p in params { - inner_exclude.insert(p.name.to_string()); - } - collect_free_vars_rec( - *body, - arena, - &inner_exclude, - local_vars, - types, - result, - seen, - ); - } - Expr::Macro(_, args) => { - for a in args { - collect_free_vars_rec(*a, arena, exclude, local_vars, types, result, seen); - } - } - Expr::StructLit(_, fields) => { - for (_, fval) in fields { - collect_free_vars_rec(*fval, arena, exclude, local_vars, types, result, seen); - } - } - // Terminal expressions — no sub-expressions. - Expr::Int(_, _) - | Expr::Real(_, _) - | Expr::String(_) - | Expr::Char(_) - | Expr::True - | Expr::False - | Expr::Enum(_) - | Expr::Break - | Expr::Continue - | Expr::Error => {} - } -} - extern "C" fn lyte_assert(globals: *mut u8, val: i8) { crate::vm::println_output(&format!("assert({})", val != 0)); if val == 0 { diff --git a/src/lib.rs b/src/lib.rs index c298c6a0..ebd91dd0 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,5 +1,13 @@ #![allow(dead_code)] +mod checked; +pub use checked::*; +mod free_locals; +mod source_analysis; +pub use source_analysis::*; +mod interface_resolution; +pub use interface_resolution::*; + mod defs; pub use defs::*; mod types; diff --git a/src/llvm_aot.rs b/src/llvm_aot.rs index 6429a27b..657fe0a1 100644 --- a/src/llvm_aot.rs +++ b/src/llvm_aot.rs @@ -26,11 +26,9 @@ // is emitted that calls the internal mangled function with the appropriate // globals_ptr + null closure_ptr prefix. -use crate::decl::*; -use crate::llvm_jit::{ - build_module, run_default_passes, AotConfig, LLVMJITState, -}; -use crate::{DeclTable, Name}; +use crate::checked::{CheckedDecl as Decl, SpecializedProgram as DeclTable}; +use crate::llvm_jit::{build_module, run_default_passes, AotConfig, LLVMJITState}; +use crate::Name; use inkwell::context::Context; use inkwell::module::Linkage; @@ -167,7 +165,8 @@ pub fn compile_aot( print_ir, // AOT requires no_recursion: the call-depth check refers to a Rust // trap helper that isn't available at link time. - /* no_recursion */ true, + /* no_recursion */ + true, Some(AotConfig { prefix: prefix.to_string(), }), @@ -261,8 +260,8 @@ fn collect_entries(decls: &DeclTable, entry_points: &[Name]) -> Result String { fn collect_globals(decls: &DeclTable, globals_size: usize) -> AotGlobalsLayout { let mut offset: i32 = crate::cancel::CANCEL_FLAG_RESERVED; let mut entries = Vec::new(); - for decl in &decls.decls { + for (_, decl) in decls.storage_instances() { match decl { Decl::Global { name, ty, .. } => { let size = ty.size(decls); @@ -398,8 +397,12 @@ fn emit_wrapper( if entry.returns_via_ptr { idx += 1; } - let user_param_tys: Vec<_> = - inner_ty.get_param_types().iter().skip(idx).copied().collect(); + let user_param_tys: Vec<_> = inner_ty + .get_param_types() + .iter() + .skip(idx) + .copied() + .collect(); // Wrapper signature: (state, ...user_params, [out_ptr]?). let mut param_tys: Vec> = vec![ptr_ty.into()]; @@ -460,11 +463,7 @@ fn emit_wrapper( // ─── Weak hooks ──────────────────────────────────────────────────────────────── -fn emit_weak_hooks( - state: &mut LLVMJITState<'_>, - prefix: &str, - public: &mut HashSet, -) { +fn emit_weak_hooks(state: &mut LLVMJITState<'_>, prefix: &str, public: &mut HashSet) { let asserts_sym = format!("{}_assert", prefix); let print_sym = format!("{}_print_i32", prefix); let putc_sym = format!("{}_putc", prefix); @@ -480,14 +479,11 @@ fn emit_weak_hooks( // Find a previously-declared abort, or declare it now. let abort_ty = void_ty.fn_type(&[], false); - let abort_fn = state - .module - .get_function("abort") - .unwrap_or_else(|| { - state - .module - .add_function("abort", abort_ty, Some(Linkage::External)) - }); + let abort_fn = state.module.get_function("abort").unwrap_or_else(|| { + state + .module + .add_function("abort", abort_ty, Some(Linkage::External)) + }); // _assert — default aborts on cond == 0. { @@ -609,12 +605,11 @@ fn set_internal_linkage(state: &mut LLVMJITState<'_>, public_names: &HashSet String { fn sanitize_macro_name(s: &str) -> String { s.chars() - .map(|c| if c.is_alphanumeric() || c == '_' { c.to_ascii_uppercase() } else { '_' }) + .map(|c| { + if c.is_alphanumeric() || c == '_' { + c.to_ascii_uppercase() + } else { + '_' + } + }) .collect() } @@ -859,7 +913,9 @@ mod tests { let mut compiler = Compiler::new(); compiler.parse("var counter: i32\ninit { counter = 10 }\n", "."); assert!(compiler.check()); - let decls = compiler.decls(); + compiler.set_entry_points(&["init"]); + compiler.specialize().unwrap(); + let decls = compiler.specialized_program().unwrap(); assert!(collect_entries(decls, &[Name::str("init")]).is_ok()); diff --git a/src/llvm_jit.rs b/src/llvm_jit.rs index a9ea099c..7424a834 100644 --- a/src/llvm_jit.rs +++ b/src/llvm_jit.rs @@ -1,10 +1,12 @@ // LLVM JIT backend for the Lyte compiler, using inkwell. // Mirrors the Cranelift JIT backend in jit.rs. -use crate::decl::*; +use crate::checked::{ + CheckedBody as ExprArena, CheckedDecl as Decl, CheckedFunction as FuncDecl, InstanceId, + LocalId, Reference, SpecializedProgram as DeclTable, +}; use crate::defs::*; -use crate::expr::*; -use crate::DeclTable; +use crate::Expr; use crate::TypeID; use std::convert::TryFrom; @@ -493,7 +495,7 @@ pub(crate) fn build_module<'ctx>( let Some(ep_decl) = decls.find_entry_point(ep_name) else { continue; }; - state.compile_function(decls, ep_decl)?; + state.compile_function(decls, ep_decl, decls.instance_for_entry(ep_name))?; } if print_ir { @@ -753,9 +755,9 @@ pub(crate) struct LLVMJITState<'ctx> { pub(crate) context: &'ctx Context, pub(crate) module: Module<'ctx>, pub(crate) builder: Builder<'ctx>, - pub(crate) globals: HashMap, + pub(crate) globals: HashMap, pub(crate) globals_size: usize, - pub(crate) defined_functions: HashSet, + pub(crate) defined_functions: HashSet, pub(crate) lambda_counter: usize, pub(crate) print_ir: bool, /// When true, skip emission of call-depth prologue/epilogue. @@ -792,15 +794,15 @@ impl<'ctx> LLVMJITState<'ctx> { fn declare_globals(&mut self, decls: &DeclTable) { let mut offset: i32 = CANCEL_FLAG_RESERVED; - for decl in &decls.decls { + for (instance, decl) in decls.storage_instances() { match decl { - Decl::Global { name, ty, .. } => { - self.globals.insert(*name, offset); + Decl::Global { ty, .. } => { + self.globals.insert(instance, offset); offset += ty.size(decls) as i32; } Decl::Func(f) if f.is_extern => { // Extern functions get 16 bytes: {fn_ptr, context} - self.globals.insert(f.name, offset); + self.globals.insert(instance, offset); offset += 16; } _ => {} @@ -964,6 +966,7 @@ impl<'ctx> LLVMJITState<'ctx> { &mut self, decls: &DeclTable, decl: &FuncDecl, + instance: Option, ) -> Result, String> { // Build the LLVM function type (globals + closure + params → ret). let fn_type = build_fn_type(self.context, decl.domain(), decl.ret, true); @@ -980,7 +983,9 @@ impl<'ctx> LLVMJITState<'ctx> { let body_id = match decl.body { Some(id) => id, None => { - self.defined_functions.insert(decl.name); + if let Some(instance) = instance { + self.defined_functions.insert(instance); + } return Ok(function); } }; @@ -1006,10 +1011,10 @@ impl<'ctx> LLVMJITState<'ctx> { None }; - // Set up local variable map: name → alloca pointer. - let mut variables: HashMap> = HashMap::new(); - let mut variable_types: HashMap = HashMap::new(); - let mut let_bindings: HashSet = HashSet::new(); + // Set up local variable map: local identity → alloca pointer. + let mut variables: HashMap> = HashMap::new(); + let mut variable_types: HashMap = HashMap::new(); + let mut let_bindings: HashSet = HashSet::new(); // Load closure variables (captured variables). for (i, cv) in decl.closure_vars.iter().enumerate() { @@ -1033,18 +1038,23 @@ impl<'ctx> LLVMJITState<'ctx> { // Store the pointer in an alloca so we can use_var-style access. let alloca = self .builder - .build_alloca(self.ptr_ty(), &cv.name.to_string()) + .build_alloca(self.ptr_ty(), &decl.arena.local(*cv).name) .unwrap(); self.builder.build_store(alloca, var_ptr).unwrap(); - variables.insert(cv.name.to_string(), alloca); - variable_types.insert(cv.name.to_string(), cv.ty); + variables.insert(*cv, alloca); + let declared = decl.arena.local(*cv).ty; + let storage = match *declared { + crate::Type::Reference(inner) => inner, + _ => declared, + }; + variable_types.insert(*cv, storage); // closure vars are treated like var bindings (pointer-indirected) } // Function parameters → let bindings. for (i, param) in decl.params.iter().enumerate() { let param_val = params[param_idx + i]; - let ty = param.ty.expect("param ty"); + let ty = decl.arena.local(param.local).ty; let storage_ty = if let crate::Type::Reference(inner) = &*ty { *inner } else { @@ -1067,13 +1077,16 @@ impl<'ctx> LLVMJITState<'ctx> { let alloca = self .builder - .build_alloca(ty.llvm_basic_type(self.context), ¶m.name.to_string()) + .build_alloca( + ty.llvm_basic_type(self.context), + &decl.arena.local(param.local).name, + ) .unwrap(); self.builder.build_store(alloca, param_val).unwrap(); - variables.insert(param.name.to_string(), alloca); - variable_types.insert(param.name.to_string(), storage_ty); + variables.insert(param.local, alloca); + variable_types.insert(param.local, storage_ty); if !matches!(&*ty, crate::Type::Reference(_)) { - let_bindings.insert(param.name.to_string()); + let_bindings.insert(param.local); } } @@ -1125,29 +1138,22 @@ impl<'ctx> LLVMJITState<'ctx> { let called = trans.called_functions.clone(); let pending = trans.pending_lambdas.clone(); - self.defined_functions.insert(decl.name); + if let Some(instance) = instance { + self.defined_functions.insert(instance); + } - // Compile pending lambdas. for lambda_decl in pending { - if !self.defined_functions.contains(&lambda_decl.name) { - self.compile_function(decls, &lambda_decl)?; - } + self.compile_function(decls, &lambda_decl, None)?; } - - // Compile called user functions. - for name in called { - if self.defined_functions.contains(&name) { - continue; - } - let found = decls.find(name); - if found.is_empty() { + for instance in called { + if self.defined_functions.contains(&instance) { continue; } - if let Decl::Func(d) = &found[0] { - if d.body.is_none() { - continue; - } - self.compile_function(decls, d)?; + let function = decls + .function_instance(instance) + .expect("checked function instance"); + if function.body.is_some() { + self.compile_function(decls, function, Some(instance))?; } } @@ -1160,15 +1166,15 @@ impl<'ctx> LLVMJITState<'ctx> { struct FunctionTranslator<'a, 'ctx> { state: &'a mut LLVMJITState<'ctx>, function: FunctionValue<'ctx>, - /// name → alloca. For let-bindings the alloca holds the value directly. + /// local identity → alloca. For let-bindings the alloca holds the value directly. /// For var-bindings the alloca holds a pointer to the actual stack slot. - variables: HashMap>, - variable_types: HashMap, - let_bindings: HashSet, + variables: HashMap>, + variable_types: HashMap, + let_bindings: HashSet, globals_base: PointerValue<'ctx>, output_ptr: Option>, decls: &'a DeclTable, - called_functions: HashSet, + called_functions: HashSet, pending_lambdas: Vec, /// Stack of (continue_bb, break_bb) for nested loops. loop_stack: Vec<(BasicBlock<'ctx>, BasicBlock<'ctx>)>, @@ -1452,21 +1458,19 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { .unwrap(); } - /// Declare (alloca) a new variable. Returns the alloca pointer. - fn declare_var(&mut self, name: &str, ty: BasicTypeEnum<'ctx>) -> PointerValue<'ctx> { - if let Some(&ptr) = self.variables.get(name) { - return ptr; - } - let alloca = self.entry_alloca(ty, name); - self.variables.insert(name.to_string(), alloca); - alloca - } - /// Get the address of a variable (for lvalue use or closure capture). - fn get_var_addr(&mut self, name: &str, _ty: crate::TypeID) -> PointerValue<'ctx> { + fn get_var_addr(&mut self, name: &LocalId, ty: crate::TypeID) -> PointerValue<'ctx> { if let Some(&alloca) = self.variables.get(name) { if self.let_bindings.contains(name) { - // let binding: alloca holds the value — alloca itself is the address. + // Aggregate values are addresses. The alloca contains their + // address, whereas a scalar's alloca is its actual storage. + if is_indirect(ty) { + return self + .builder() + .build_load(self.ptr_ty(), alloca, "capture_storage") + .unwrap() + .into_pointer_value(); + } return alloca; } else { // var binding: alloca holds a pointer to the slot. @@ -1477,10 +1481,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { .into_pointer_value(); } } - if let Some(&offset) = self.state.globals.get(&Name::new(name.to_string())) { - return self.ptr_at_offset(self.globals_base, offset as u64); - } - panic!("unknown variable in closure capture: {}", name) + panic!("unknown local in closure capture: {:?}", name) } /// Compute globals_base + byte_offset as a pointer. @@ -1541,28 +1542,26 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } fn representation_type(&self, expr: ExprID, decl: &FuncDecl) -> crate::TypeID { - match &decl.arena.exprs[expr] { - Expr::Id(name) => self - .variable_types - .get(name.as_str()) - .copied() - .or_else(|| { - self.decls.find(*name).iter().find_map(|decl| { - if let crate::Decl::Global { ty, .. } = decl { - Some(*ty) - } else { - None - } - }) - }) - .unwrap_or(decl.types[expr]), + match &decl.arena[expr] { + Expr::Id(_) => match decl.arena.reference(expr) { + Some(Reference::Local(local)) => self + .variable_types + .get(local) + .copied() + .unwrap_or(decl.arena.local(*local).ty), + Some(Reference::Instance(instance)) => match self.decls.instance(*instance) { + Decl::Global { ty, .. } => *ty, + _ => decl.arena.ty(expr), + }, + _ => decl.arena.ty(expr), + }, Expr::ArrayIndex(arr_id, _) => match &*self.representation_type(*arr_id, decl) { crate::Type::Array(elem, _) | crate::Type::Slice(elem) | crate::Type::Reference(elem) => *elem, - _ => decl.types[expr], + _ => decl.arena.ty(expr), }, - _ => decl.types[expr], + _ => decl.arena.ty(expr), } } @@ -1701,11 +1700,11 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } /// Match `{ var = expr }` or `var = expr` — returns (var_name, rhs_id) if matched. - fn match_single_var_assign(arena: &ExprArena, expr_id: ExprID) -> Option<(Name, ExprID)> { - let expr = &arena.exprs[expr_id]; + fn match_single_var_assign(arena: &ExprArena, expr_id: ExprID) -> Option<(LocalId, ExprID)> { + let expr = &arena[expr_id]; let inner = if let Expr::Block(stmts) = expr { if stmts.len() == 1 { - &arena.exprs[stmts[0]] + &arena[stmts[0]] } else { return None; } @@ -1713,7 +1712,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { expr }; if let Expr::Binop(Binop::Assign, lhs, rhs) = inner { - if let Expr::Id(name) = &arena.exprs[*lhs] { + if let Some(Reference::Local(name)) = arena.reference(*lhs) { return Some((*name, *rhs)); } } @@ -1724,53 +1723,61 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { fn translate_lvalue(&mut self, expr: ExprID, decl: &FuncDecl) -> PointerValue<'ctx> { match &decl.arena[expr] { - Expr::Id(name) => { - if let Some(&alloca) = self.variables.get(&**name) { - if self.let_bindings.contains(&**name) { - let ty = decl.types[expr]; - if ty.is_ptr() || matches!(&*ty, crate::Type::Slice(_)) { - // Pointer-type let binding (e.g. slice/array/struct param): - // alloca holds a pointer to the data, load it. - self.builder() - .build_load(self.ptr_ty(), alloca, "let_ptr") - .unwrap() - .into_pointer_value() - } else { - alloca - } - } else { - if let Some(ty) = self.variable_types.get(name.as_str()).copied() { - if ty.is_ptr() { + Expr::Id(_) => match decl.arena.reference(expr) { + Some(Reference::Local(name)) => { + if let Some(&alloca) = self.variables.get(name) { + if self.let_bindings.contains(name) { + let ty = decl.arena.ty(expr); + if ty.is_ptr() || matches!(&*ty, crate::Type::Slice(_)) { + // Pointer-type let binding (e.g. slice/array/struct param): + // alloca holds a pointer to the data, load it. self.builder() - .build_load(ty.llvm_basic_type(self.ctx()), alloca, "var_addr") + .build_load(self.ptr_ty(), alloca, "let_ptr") .unwrap() .into_pointer_value() + } else { + alloca + } + } else { + if let Some(ty) = self.variable_types.get(name).copied() { + if ty.is_ptr() { + self.builder() + .build_load( + ty.llvm_basic_type(self.ctx()), + alloca, + "var_addr", + ) + .unwrap() + .into_pointer_value() + } else { + self.builder() + .build_load(self.ptr_ty(), alloca, "var_addr") + .unwrap() + .into_pointer_value() + } } else { self.builder() .build_load(self.ptr_ty(), alloca, "var_addr") .unwrap() .into_pointer_value() } - } else { - self.builder() - .build_load(self.ptr_ty(), alloca, "var_addr") - .unwrap() - .into_pointer_value() } + } else { + panic!("JIT lvalue: unknown variable {:?}", name) } - } else if let Some(&offset) = self.state.globals.get(name) { - self.ptr_at_offset(self.globals_base, offset as u64) - } else { - panic!("JIT lvalue: unknown variable {:?}", name) } - } + Some(Reference::Instance(instance)) => { + self.ptr_at_offset(self.globals_base, self.state.globals[instance] as u64) + } + reference => panic!("unresolved checked reference: {:?}", reference), + }, Expr::Field(lhs_id, field_name) => { - let lhs_ty = decl.types[*lhs_id]; + let lhs_ty = decl.arena.ty(*lhs_id); let base_ptr = self.translate_lvalue(*lhs_id, decl); self.compute_field_ptr(base_ptr, lhs_ty, field_name, decl) } Expr::ArrayIndex(lhs_id, idx_id) => { - let lhs_ty = decl.types[*lhs_id]; + let lhs_ty = decl.arena.ty(*lhs_id); let lhs_ptr = self.translate_lvalue(*lhs_id, decl); let idx_val = self.translate_expr(*idx_id, decl).into_int_value(); self.compute_array_elem_ptr(lhs_ptr, lhs_ty, idx_val) @@ -1785,24 +1792,30 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { /// isn't a place at all. fn f32x4_storage(&mut self, expr: ExprID, decl: &FuncDecl) -> Option> { match &decl.arena[expr] { - Expr::Id(name) => { - if let Some(&alloca) = self.variables.get(&**name) { - if self.let_bindings.contains(&**name) { - // The alloca holds the vector itself. - return Some(alloca); + Expr::Id(_) => match decl.arena.reference(expr) { + Some(Reference::Local(name)) => { + if let Some(&alloca) = self.variables.get(name) { + if self.let_bindings.contains(name) { + // The alloca holds the vector itself. + return Some(alloca); + } + // A captured variable's alloca holds the address of the + // vector in the enclosing frame. + return Some( + self.builder() + .build_load(self.ptr_ty(), alloca, "vec_ptr") + .unwrap() + .into_pointer_value(), + ); } - // A captured variable's alloca holds the address of the - // vector in the enclosing frame. - return Some( - self.builder() - .build_load(self.ptr_ty(), alloca, "vec_ptr") - .unwrap() - .into_pointer_value(), - ); + None } - let offset = *self.state.globals.get(name)?; - Some(self.ptr_at_offset(self.globals_base, offset as u64)) - } + Some(Reference::Instance(instance)) => { + let offset = *self.state.globals.get(instance)?; + Some(self.ptr_at_offset(self.globals_base, offset as u64)) + } + _ => None, + }, Expr::Field(_, _) | Expr::ArrayIndex(_, _) => Some(self.translate_lvalue(expr, decl)), _ => None, } @@ -1837,7 +1850,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { ) -> PointerValue<'ctx> { if let crate::Type::Name(struct_name, type_args) = &*lhs_ty { let struct_decl = self.decls.find(*struct_name); - if let crate::Decl::Struct(s) = &struct_decl[0] { + if let Decl::Struct(s) = &struct_decl[0] { let inst: crate::Instance = s .typevars .iter() @@ -1898,70 +1911,80 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { match &decl.arena[expr] { Expr::True => self.i8_ty().const_int(1, false).into(), Expr::False => self.i8_ty().const_int(0, false).into(), - Expr::Int(n, _) => match &*decl.types[expr] { + Expr::Int(n, _) => match &*decl.arena.ty(expr) { crate::Type::Int8 => self.i8_ty().const_int(*n as u64, true).into(), crate::Type::UInt32 => self.i32_ty().const_int(*n as u64, false).into(), _ => self.i32_ty().const_int(*n as u64, true).into(), }, Expr::Real(s, _) => { let val: f64 = s.parse().expect("invalid float literal"); - match &*decl.types[expr] { + match &*decl.arena.ty(expr) { crate::Type::Float32 => self.state.f32_ty().const_float(val).into(), _ => self.state.f64_ty().const_float(val).into(), } } Expr::Char(c) => self.i8_ty().const_int(*c as u64, false).into(), - Expr::Id(name) => { - let ty = decl.types[expr]; - if let Some(&alloca) = self.variables.get(&**name) { - if self.let_bindings.contains(&**name) || is_indirect(ty) { - // let binding or pointer type: load the value from the alloca. - self.builder() - .build_load(ty.llvm_basic_type(self.ctx()), alloca, &**name) - .unwrap() - } else { - let stored = self - .builder() - .build_load(self.ptr_ty(), alloca, "var_ptr") - .unwrap(); - if let Some(var_ty) = self.variable_types.get(name.as_str()).copied() { - if is_indirect(var_ty) { - stored + Expr::Id(_) => match decl.arena.reference(expr) { + Some(Reference::Local(name)) => { + let ty = decl.arena.ty(expr); + if let Some(&alloca) = self.variables.get(name) { + if self.let_bindings.contains(name) || is_indirect(ty) { + // let binding or pointer type: load the value from the alloca. + self.builder() + .build_load( + ty.llvm_basic_type(self.ctx()), + alloca, + &decl.arena.local(*name).name, + ) + .unwrap() + } else { + let stored = self + .builder() + .build_load(self.ptr_ty(), alloca, "var_ptr") + .unwrap(); + if let Some(var_ty) = self.variable_types.get(name).copied() { + if is_indirect(var_ty) { + stored + } else { + self.builder() + .build_load( + ty.llvm_basic_type(self.ctx()), + stored.into_pointer_value(), + &decl.arena.local(*name).name, + ) + .unwrap() + } } else { self.builder() .build_load( ty.llvm_basic_type(self.ctx()), stored.into_pointer_value(), - &**name, + &decl.arena.local(*name).name, ) .unwrap() } + } + } else { + panic!("missing local storage: {:?}", name) + } + } + Some(Reference::Instance(instance)) => { + let ty = decl.arena.ty(expr); + if let Some(&offset) = self.state.globals.get(instance) { + let addr = self.ptr_at_offset(self.globals_base, offset as u64); + if is_indirect(ty) { + addr.into() } else { self.builder() - .build_load( - ty.llvm_basic_type(self.ctx()), - stored.into_pointer_value(), - &**name, - ) + .build_load(ty.llvm_basic_type(self.ctx()), addr, "global") .unwrap() } - } - } else if let Some(&offset) = self.state.globals.get(name) { - let addr = self.ptr_at_offset(self.globals_base, offset as u64); - // Composite types (arrays, structs) are pointer-represented: - // return the address, don't load. - if is_indirect(ty) { - addr.into() } else { - self.builder() - .build_load(ty.llvm_basic_type(self.ctx()), addr, &**name) - .unwrap() + self.translate_func_ref(*instance, &*ty) } - } else { - // Must be a function reference. - self.translate_func_ref(name, &*ty) } - } + reference => panic!("unresolved checked reference: {:?}", reference), + }, Expr::Binop(op, lhs_id, rhs_id) => { let (op, lhs, rhs) = (*op, *lhs_id, *rhs_id); self.translate_binop(op, lhs, rhs, decl) @@ -1974,9 +1997,9 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { let (fn_id, arg_ids) = (*fn_id, arg_ids.clone()); self.translate_call(fn_id, &arg_ids, expr, decl) } - Expr::Let(name, init_id, _) => { - let (name, init_id) = (*name, *init_id); - let ty = decl.types[expr]; + Expr::Let(_, init_id, _) => { + let (name, init_id) = (decl.arena.binder(expr), *init_id); + let ty = decl.arena.local(name).ty; let init_val = self.translate_expr(init_id, decl); let init_val = self.wrap_for_expected_slice(init_val, ty, init_id, decl); @@ -1992,36 +2015,39 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { let storage = self.entry_array_alloca( self.i8_ty(), sz as u64, - &format!("{}_storage", name), + &format!("{}_storage", decl.arena.local(name).name), ); - let alloca = self.entry_alloca(self.ptr_ty().into(), &*name); + let alloca = + self.entry_alloca(self.ptr_ty().into(), &decl.arena.local(name).name); self.builder().build_store(alloca, storage).unwrap(); - self.variables.insert(name.to_string(), alloca); - self.variable_types.insert(name.to_string(), ty); + self.variables.insert(name, alloca); + self.variable_types.insert(name, ty); // The binding owns storage now, exactly like a `var`, so it // must not be treated as holding a value directly. - self.let_bindings.remove(&name.to_string()); + self.let_bindings.remove(&name); self.gen_copy(ty, storage, init_val); return storage.into(); } - let alloca = self.entry_alloca(ty.llvm_basic_type(self.ctx()), &*name); + let alloca = + self.entry_alloca(ty.llvm_basic_type(self.ctx()), &decl.arena.local(name).name); self.builder().build_store(alloca, init_val).unwrap(); - self.variables.insert(name.to_string(), alloca); - self.variable_types.insert(name.to_string(), ty); - self.let_bindings.insert(name.to_string()); + self.variables.insert(name, alloca); + self.variable_types.insert(name, ty); + self.let_bindings.insert(name); init_val } - Expr::Var(name, init_id, _) => { - let (name, init_id) = (*name, *init_id); - let ty = decl.types[expr]; + Expr::Var(_, init_id, _) => { + let (name, init_id) = (decl.arena.binder(expr), *init_id); + let ty = decl.arena.local(name).ty; let sz = ty.size(self.decls) as usize; assert!(sz > 0, "var size must be > 0"); if !ty.is_ptr() || is_llvm_value_type(ty) { // Scalar/vector var: use a single typed alloca (same as let bindings). // This avoids double-indirection and lets LLVM promote to SSA. - let alloca = self.entry_alloca(ty.llvm_basic_type(self.ctx()), &*name); + let alloca = self + .entry_alloca(ty.llvm_basic_type(self.ctx()), &decl.arena.local(name).name); if let Some(init_id) = init_id { let init_val = self.translate_expr(init_id, decl); self.builder().build_store(alloca, init_val).unwrap(); @@ -2029,20 +2055,21 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { let zero = ty.llvm_basic_type(self.ctx()).const_zero(); self.builder().build_store(alloca, zero).unwrap(); } - self.variables.insert(name.to_string(), alloca); - self.variable_types.insert(name.to_string(), ty); - self.let_bindings.insert(name.to_string()); + self.variables.insert(name, alloca); + self.variable_types.insert(name, ty); + self.let_bindings.insert(name); } else { // Pointer/struct var: allocate byte storage + ptr alloca. let storage = self.entry_array_alloca( self.i8_ty(), sz as u64, - &format!("{}_storage", name), + &format!("{}_storage", decl.arena.local(name).name), ); - let alloca = self.entry_alloca(self.ptr_ty().into(), &*name); + let alloca = + self.entry_alloca(self.ptr_ty().into(), &decl.arena.local(name).name); self.builder().build_store(alloca, storage).unwrap(); - self.variables.insert(name.to_string(), alloca); - self.variable_types.insert(name.to_string(), ty); + self.variables.insert(name, alloca); + self.variable_types.insert(name, ty); if let Some(init_id) = init_id { let init_val = self.translate_expr(init_id, decl); @@ -2056,7 +2083,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } Expr::Field(lhs_id, field_name) => { let (lhs_id, field_name) = (*lhs_id, *field_name); - let lhs_ty = decl.types[lhs_id]; + let lhs_ty = decl.arena.ty(lhs_id); // Handle .len. if *field_name == "len" { match &*lhs_ty { @@ -2094,7 +2121,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } let lhs_val = self.translate_expr(lhs_id, decl).into_pointer_value(); - let field_ty = decl.types[expr]; + let field_ty = decl.arena.ty(expr); let field_ptr = self.compute_field_ptr(lhs_val, lhs_ty, &field_name, decl); if is_indirect(field_ty) { field_ptr.into() @@ -2106,7 +2133,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } Expr::ArrayIndex(lhs_id, rhs_id) => { let (lhs_id, rhs_id) = (*lhs_id, *rhs_id); - let lhs_ty = decl.types[lhs_id]; + let lhs_ty = decl.arena.ty(lhs_id); // f32x4 element extraction: use extractelement if matches!(*lhs_ty, crate::Type::Float32x4) { @@ -2123,7 +2150,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { let lhs_val = self.translate_expr(lhs_id, decl).into_pointer_value(); let rhs_val = self.translate_expr(rhs_id, decl).into_int_value(); let elem_ptr = self.compute_array_elem_ptr(lhs_val, lhs_ty, rhs_val); - let result_ty = decl.types[expr]; + let result_ty = decl.arena.ty(expr); if is_indirect(result_ty) { elem_ptr.into() } else { @@ -2134,7 +2161,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } Expr::ArrayLiteral(elements) => { let elements = elements.clone(); - let ty = decl.types[expr]; + let ty = decl.arena.ty(expr); let elem_values: Vec> = elements .iter() .map(|e| self.translate_expr(*e, decl)) @@ -2156,11 +2183,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { if exprs.is_empty() { self.zero_i32() } else { - // Save variable scope — declarations inside this block - // shadow outer names only for the duration of the block. - let saved_vars = self.variables.clone(); - let saved_types = self.variable_types.clone(); - let saved_lets = self.let_bindings.clone(); + // Binding identities remain unambiguous across block boundaries. let mut result = self.zero_i32(); for e in &exprs { if self.is_block_terminated() { @@ -2168,9 +2191,6 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } result = self.translate_expr(*e, decl); } - self.variables = saved_vars; - self.variable_types = saved_types; - self.let_bindings = saved_lets; result } } @@ -2182,14 +2202,14 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { if let Some((var_name, new_val_id)) = Self::match_single_var_assign(&decl.arena, then_id) { - let var_ty = decl.types[new_val_id]; + let var_ty = decl.arena.ty(new_val_id); let is_scalar_float = matches!(*var_ty, crate::Type::Float32 | crate::Type::Float64); - if is_scalar_float && self.variables.contains_key(&var_name.to_string()) { + if is_scalar_float && self.variables.contains_key(&var_name) { let cond_raw = self.translate_expr(cond_id, decl).into_int_value(); let cond_val = self.to_i1(cond_raw); - let alloca = self.variables[&var_name.to_string()]; - let is_let = self.let_bindings.contains(&var_name.to_string()); + let alloca = self.variables[&var_name]; + let is_let = self.let_bindings.contains(&var_name); // Get the storage pointer. let slot_ptr = if is_let { alloca @@ -2216,9 +2236,9 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { // Determine if this if-else produces a value (both branches // have the same concrete, non-void type). - let result_ty = decl.types[expr]; + let result_ty = decl.arena.ty(expr); let is_value = if let Some(else_expr_id) = else_id { - let else_ty = decl.types[else_expr_id]; + let else_ty = decl.arena.ty(else_expr_id); !matches!( &*result_ty, crate::Type::Void | crate::Type::Anon(_) | crate::Type::Var(_) @@ -2302,31 +2322,25 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { self.zero_i32() } Expr::For { - var, - start, - end, - body, + start, end, body, .. } => { - let (var, start, end, body) = (*var, *start, *end, *body); + let (var, start, end, body) = (decl.arena.binder(expr), *start, *end, *body); // Both bounds are outside the loop variable's scope, so they // still see any outer binding of the same name. let start_val = self.translate_expr(start, decl).into_int_value(); let end_val = self.translate_expr(end, decl).into_int_value(); - // The loop variable is scoped to the loop: save the name-keyed - // state so an outer binding it shadows comes back at loop exit. - let saved_vars = self.variables.clone(); - let saved_types = self.variable_types.clone(); - let saved_lets = self.let_bindings.clone(); + // The checked loop binding has its own local identity. // Allocate loop counter in entry block. - let loop_alloca = self.entry_alloca(self.i32_ty().into(), &*var); + let loop_alloca = + self.entry_alloca(self.i32_ty().into(), &decl.arena.local(var).name); self.builder().build_store(loop_alloca, start_val).unwrap(); // Treat as let binding (holds value directly). - self.variables.insert(var.to_string(), loop_alloca); + self.variables.insert(var, loop_alloca); self.variable_types - .insert(var.to_string(), crate::types::mk_type(crate::Type::Int32)); - self.let_bindings.insert(var.to_string()); + .insert(var, crate::types::mk_type(crate::Type::Int32)); + self.let_bindings.insert(var); let header_bb = self.append_bb("for_header"); let body_bb = self.append_bb("for_body"); @@ -2382,10 +2396,6 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { self.loop_stack.pop(); self.builder().position_at_end(exit_bb); - self.variables = saved_vars; - self.variable_types = saved_types; - self.let_bindings = saved_lets; - self.zero_i32() } Expr::Assume(_) => { @@ -2395,7 +2405,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { Expr::Return(ret_id) => { let ret_id = *ret_id; let result = self.translate_expr(ret_id, decl); - let ret_ty = decl.types[ret_id]; + let ret_ty = decl.arena.ty(ret_id); if !self.state.no_recursion { self.emit_call_depth_release(); @@ -2420,7 +2430,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } Expr::Tuple(elements) => { let elements = elements.clone(); - let ty = decl.types[expr]; + let ty = decl.arena.ty(expr); let elem_vals: Vec> = elements .iter() .map(|e| self.translate_expr(*e, decl)) @@ -2439,17 +2449,16 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { panic!("tuple expr: expected tuple type"); } } - Expr::Lambda { params, body } => { - let (params, body) = (params.clone(), *body); - self.translate_lambda(¶ms, body, expr, decl) + Expr::Lambda { body, .. } => { + let body = *body; + self.translate_lambda(decl.arena.binders(expr), body, expr, decl) } Expr::Enum(case_name) => { let case_name = *case_name; - let index = if let crate::Type::Name(enum_name, _) = &*decl.types[expr] { + let index = if let crate::Type::Name(enum_name, _) = &*decl.arena.ty(expr) { let enum_decls = self.decls.find(*enum_name); - if let Some(crate::Decl::Enum { cases, .. }) = enum_decls - .iter() - .find(|d| matches!(d, crate::Decl::Enum { .. })) + if let Some(Decl::Enum { cases, .. }) = + enum_decls.iter().find(|d| matches!(d, Decl::Enum { .. })) { cases.iter().position(|c| *c == case_name).unwrap_or(0) as u64 } else { @@ -2487,7 +2496,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { Expr::AsTy(src_id, target_ty) => { let (src_id, target_ty) = (*src_id, *target_ty); let val = self.translate_expr(src_id, decl); - let src_ty = decl.types[src_id]; + let src_ty = decl.arena.ty(src_id); self.translate_cast(val, src_ty, target_ty) } Expr::Break => { @@ -2509,7 +2518,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { Expr::StructLit(struct_name, fields) => { let struct_name = *struct_name; let fields: Vec<(Name, ExprID)> = fields.clone(); - let ty = decl.types[expr]; + let ty = decl.arena.ty(expr); let sz = ty.size(self.decls) as u64; let storage = self.entry_array_alloca(self.i8_ty(), sz, "struct_lit"); // Zero-initialize. @@ -2521,7 +2530,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { if let crate::Type::Name(_, type_args) = &*ty { let struct_decl = self.decls.find(struct_name); - if let crate::Decl::Struct(s) = &struct_decl[0] { + if let Decl::Struct(s) = &struct_decl[0] { let inst: crate::Instance = s .typevars .iter() @@ -2531,7 +2540,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { for (fname, fval) in &fields { let val = self.translate_expr(*fval, decl); let off = s.field_offset(fname, self.decls, &inst); - let field_ty = decl.types[*fval]; + let field_ty = decl.arena.ty(*fval); self.store_element(field_ty, storage, off as u64, val); } } @@ -2542,7 +2551,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { // Fill-array expression: [value; size], e.g. [0; 5] let value_expr = *value_expr; let fill_value = self.translate_expr(value_expr, decl); - let ty = decl.types[expr]; + let ty = decl.arena.ty(expr); if let crate::Type::Array(elem_ty, sz) = &*ty { let count = sz.known(); let elem_size = elem_ty.size(self.decls) as u64; @@ -2572,7 +2581,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { decl: &FuncDecl, ) -> BasicValueEnum<'ctx> { let val = self.translate_expr(arg_id, decl); - let ty = decl.types[arg_id]; + let ty = decl.arena.ty(arg_id); match op { Unop::Neg => match *ty { crate::Type::Float32x4 => self @@ -2610,7 +2619,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { ) -> BasicValueEnum<'ctx> { match op { Binop::Plus => { - let t = decl.types[lhs_id]; + let t = decl.arena.ty(lhs_id); let is_float = matches!(*t, crate::Type::Float32 | crate::Type::Float64); // f32x4: vector fadd @@ -2626,13 +2635,13 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { // FMA: a*b + c if is_float { - if let Expr::Binop(Binop::Mult, ma, mb) = decl.arena.exprs[lhs_id] { + if let Expr::Binop(Binop::Mult, ma, mb) = decl.arena[lhs_id] { let a = self.translate_expr(ma, decl).into_float_value(); let b = self.translate_expr(mb, decl).into_float_value(); let c = self.translate_expr(rhs_id, decl).into_float_value(); return self.build_fma(a, b, c, t).into(); } - if let Expr::Binop(Binop::Mult, ma, mb) = decl.arena.exprs[rhs_id] { + if let Expr::Binop(Binop::Mult, ma, mb) = decl.arena[rhs_id] { let c = self.translate_expr(lhs_id, decl).into_float_value(); let a = self.translate_expr(ma, decl).into_float_value(); let b = self.translate_expr(mb, decl).into_float_value(); @@ -2655,7 +2664,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } } Binop::Minus => { - let t = decl.types[lhs_id]; + let t = decl.arena.ty(lhs_id); let is_float = matches!(*t, crate::Type::Float32 | crate::Type::Float64); if matches!(*t, crate::Type::Float32x4) { @@ -2670,7 +2679,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { // FMA: c - a*b => fma(-a, b, c) if is_float { - if let Expr::Binop(Binop::Mult, ma, mb) = decl.arena.exprs[rhs_id] { + if let Expr::Binop(Binop::Mult, ma, mb) = decl.arena[rhs_id] { let c = self.translate_expr(lhs_id, decl).into_float_value(); let a = self.translate_expr(ma, decl).into_float_value(); let b = self.translate_expr(mb, decl).into_float_value(); @@ -2694,7 +2703,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } } Binop::Mult => { - let t = decl.types[lhs_id]; + let t = decl.arena.ty(lhs_id); if matches!(*t, crate::Type::Float32x4) { let lhs = self.translate_expr(lhs_id, decl).into_vector_value(); @@ -2721,7 +2730,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } } Binop::Div => { - let t = decl.types[lhs_id]; + let t = decl.arena.ty(lhs_id); if matches!(*t, crate::Type::Float32x4) { let lhs = self.translate_expr(lhs_id, decl).into_vector_value(); @@ -2754,7 +2763,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } } Binop::Mod => { - let t = decl.types[lhs_id]; + let t = decl.arena.ty(lhs_id); let lhs = self.translate_expr(lhs_id, decl); let rhs = self.translate_expr(rhs_id, decl); match *t { @@ -2777,9 +2786,9 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } Binop::Assign => { // f32x4 field assignment: v.x = val → insert_element + store - if let Expr::Field(vec_id, field_name) = &decl.arena.exprs[lhs_id] { + if let Expr::Field(vec_id, field_name) = &decl.arena[lhs_id] { let (vec_id, field_name) = (*vec_id, *field_name); - let vec_ty = decl.types[vec_id]; + let vec_ty = decl.arena.ty(vec_id); if matches!(*vec_ty, crate::Type::Float32x4) { let lane: u64 = match &**field_name { "x" | "r" => 0, @@ -2798,9 +2807,9 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } } // f32x4 element assignment: v[i] = val → insert_element + store - if let Expr::ArrayIndex(vec_id, idx_id) = &decl.arena.exprs[lhs_id] { + if let Expr::ArrayIndex(vec_id, idx_id) = &decl.arena[lhs_id] { let (vec_id, idx_id) = (*vec_id, *idx_id); - let vec_ty = decl.types[vec_id]; + let vec_ty = decl.arena.ty(vec_id); if matches!(*vec_ty, crate::Type::Float32x4) { // The lane index is in 0..4: the safety checker proves it, // and no backend checks it at runtime. @@ -2836,19 +2845,19 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { Binop::Equal => { let lhs = self.translate_expr(lhs_id, decl); let rhs = self.translate_expr(rhs_id, decl); - let t = decl.types[lhs_id]; + let t = decl.arena.ty(lhs_id); self.gen_eq(t, lhs, rhs) } Binop::NotEqual => { let lhs = self.translate_expr(lhs_id, decl); let rhs = self.translate_expr(rhs_id, decl); - let t = decl.types[lhs_id]; + let t = decl.arena.ty(lhs_id); let eq = self.gen_eq(t, lhs, rhs).into_int_value(); let one = eq.get_type().const_int(1, false); self.builder().build_xor(eq, one, "ne").unwrap().into() } Binop::Less => { - let t = decl.types[lhs_id]; + let t = decl.arena.ty(lhs_id); let lhs = self.translate_expr(lhs_id, decl); let rhs = self.translate_expr(rhs_id, decl); match *t { @@ -2885,7 +2894,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } } Binop::Greater => { - let t = decl.types[lhs_id]; + let t = decl.arena.ty(lhs_id); let lhs = self.translate_expr(lhs_id, decl); let rhs = self.translate_expr(rhs_id, decl); match *t { @@ -2922,7 +2931,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } } Binop::Leq => { - let t = decl.types[lhs_id]; + let t = decl.arena.ty(lhs_id); let lhs = self.translate_expr(lhs_id, decl); let rhs = self.translate_expr(rhs_id, decl); match *t { @@ -2959,7 +2968,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } } Binop::Geq => { - let t = decl.types[lhs_id]; + let t = decl.arena.ty(lhs_id); let lhs = self.translate_expr(lhs_id, decl); let rhs = self.translate_expr(rhs_id, decl); match *t { @@ -3134,29 +3143,28 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { _call_expr_id: ExprID, decl: &FuncDecl, ) -> BasicValueEnum<'ctx> { - let is_builtin = if let Expr::Id(name) = &decl.arena[fn_id] { - is_builtin_name(name) + let is_builtin = if let Some(Reference::Instance(instance)) = decl.arena.reference(fn_id) { + is_builtin_name(&self.decls.instance_name(*instance)) } else { false }; - let fn_type = decl.types[fn_id]; + let fn_type = decl.arena.ty(fn_id); if let crate::Type::Func(from, to) = *fn_type { - let param_types: Vec = if let Expr::Id(callee_name) = &decl.arena[fn_id] - { - let callee_decls = self.decls.find(*callee_name); - if let Some(crate::Decl::Func(f)) = callee_decls.first() { - f.param_types() + let param_types: Vec = + if let Some(Reference::Instance(callee)) = decl.arena.reference(fn_id) { + if let Some(f) = self.decls.function_instance(*callee) { + f.param_types() + } else if let crate::Type::Tuple(pts) = &*from { + pts.clone() + } else { + vec![] + } } else if let crate::Type::Tuple(pts) = &*from { pts.clone() } else { vec![] - } - } else if let crate::Type::Tuple(pts) = &*from { - pts.clone() - } else { - vec![] - }; + }; // Allocate output slot if returning via pointer. let output_slot: Option> = if returns_via_pointer(to) { @@ -3168,8 +3176,9 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { }; // f32x4 constructor and splat. - if let Expr::Id(name) = &decl.arena[fn_id] { - if **name == "f32x4" && arg_ids.len() == 4 { + if let Some(Reference::Instance(instance)) = decl.arena.reference(fn_id) { + let name = self.decls.instance_name(*instance); + if *name == "f32x4" && arg_ids.len() == 4 { let vec_ty = self.ctx().f32_type().vec_type(4); let mut vec = vec_ty.get_undef(); for (i, &arg_id) in arg_ids.iter().enumerate() { @@ -3182,7 +3191,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } return vec.into(); } - if **name == "f32x4_splat" && arg_ids.len() == 1 { + if *name == "f32x4_splat" && arg_ids.len() == 1 { let vec_ty = self.ctx().f32_type().vec_type(4); let val = self.translate_expr(arg_ids[0], decl).into_float_value(); let mut vec = vec_ty.get_undef(); @@ -3198,14 +3207,15 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } // Check for math builtin by name. - let math_intrinsic = if let Expr::Id(name) = &decl.arena[fn_id] { - self.llvm_intrinsic_name(name) - } else { - None - }; + let math_intrinsic = + if let Some(Reference::Instance(instance)) = decl.arena.reference(fn_id) { + self.llvm_intrinsic_name(&self.decls.instance_name(*instance)) + } else { + None + }; let math_sym = if math_intrinsic.is_none() { - if let Expr::Id(name) = &decl.arena[fn_id] { - self.math_builtin_name(name, from) + if let Some(Reference::Instance(instance)) = decl.arena.reference(fn_id) { + self.math_builtin_name(&self.decls.instance_name(*instance), from) } else { None } @@ -3244,17 +3254,17 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } } else if { // Check for extern function. - if let Expr::Id(callee_name) = &decl.arena[fn_id] { - let callee_decls = self.decls.find(*callee_name); - callee_decls - .first() - .map_or(false, |d| matches!(d, crate::Decl::Func(f) if f.is_extern)) + if let Some(Reference::Instance(callee)) = decl.arena.reference(fn_id) { + self.decls + .function_instance(*callee) + .is_some_and(|function| function.is_extern) } else { false } } { // Extern function: load {fn_ptr, context} from globals buffer. - let callee_name = if let Expr::Id(n) = &decl.arena[fn_id] { + let callee_name = if let Some(Reference::Instance(n)) = decl.arena.reference(fn_id) + { *n } else { unreachable!() @@ -3278,16 +3288,14 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { .into_pointer_value(); // Build function type with slices expanded to (ptr, i32). - let callee_decl = { - let decls_found = self.decls.find(callee_name); - match decls_found.first() { - Some(crate::Decl::Func(f)) => f.clone(), - _ => unreachable!(), - } - }; + let callee_decl = self + .decls + .function_instance(callee_name) + .expect("extern instance") + .clone(); let mut param_tys: Vec> = vec![self.ptr_ty().into()]; // context for param in &callee_decl.params { - let pty = param.ty.unwrap(); + let pty = callee_decl.arena.local(param.local).ty; if matches!(&*pty, crate::Type::Slice(_)) { param_tys.push(self.ptr_ty().into()); // data ptr param_tys.push(self.i32_ty().into()); // len @@ -3303,7 +3311,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { let mut args: Vec> = vec![context.into()]; for (i, arg_id) in arg_ids.iter().enumerate() { - let param_ty = callee_decl.params[i].ty.unwrap(); + let param_ty = callee_decl.arena.local(callee_decl.params[i].local).ty; let arg_val = if matches!(&*param_ty, crate::Type::Reference(_)) { self.translate_lvalue(*arg_id, decl).into() } else { @@ -3359,7 +3367,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } else if is_builtin { // assert / print / putc — load raw fn ptr from fat pointer and indirect call. // assert needs globals as its first arg so it can trap via longjmp. - let is_assert = matches!(&decl.arena[fn_id], Expr::Id(n) if **n == "assert"); + let is_assert = matches!(decl.arena.reference(fn_id), Some(Reference::Instance(instance)) if *self.decls.instance_name(*instance) == "assert"); let fat_ptr = self.translate_expr(fn_id, decl).into_pointer_value(); let fn_ptr_val = self .builder() @@ -3750,7 +3758,12 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } /// Translate a function reference (not a call), returning a fat pointer. - fn translate_func_ref(&mut self, name: &Name, ty: &crate::Type) -> BasicValueEnum<'ctx> { + fn translate_func_ref( + &mut self, + instance: InstanceId, + ty: &crate::Type, + ) -> BasicValueEnum<'ctx> { + let name = &self.decls.instance_name(instance); // Builtin raw function pointers. if *name == Name::str("assert") { // assert takes (globals: ptr, val: i8) so it can write trap_reason @@ -3802,7 +3815,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { self.module() .add_function(&**name, full_fn_ty, Some(Linkage::External)) }; - self.called_functions.insert(*name); + self.called_functions.insert(instance); let ptr = f.as_global_value().as_pointer_value(); self.make_fat_ptr(ptr, None) } else { @@ -3847,44 +3860,27 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { fn translate_lambda( &mut self, - params: &[Param], - body: ExprID, + _params: &[LocalId], + _body: ExprID, expr_id: ExprID, decl: &FuncDecl, ) -> BasicValueEnum<'ctx> { - let lambda_ty = decl.types[expr_id]; + let lambda_ty = decl.arena.ty(expr_id); if let crate::Type::Func(dom, rng) = *lambda_ty { - if let crate::Type::Tuple(param_types) = &*dom { + if let crate::Type::Tuple(_) = &*dom { let id = self.state.lambda_counter; self.state.lambda_counter += 1; let lambda_name = Name::new(format!("__lambda_{}", id)); - let lambda_params: Vec = params - .iter() - .zip(param_types.iter()) - .map(|(p, ty)| Param { - name: p.name, - ty: Some(*ty), - }) - .collect(); - - // Collect free variables. - let param_names: HashSet = - params.iter().map(|p| p.name.to_string()).collect(); - let free_vars = collect_free_var_names_llvm( - body, - &decl.arena, - ¶m_names, - &self.variables, - &decl.types, - ); + let lambda_decl = decl.extract_lambda(expr_id, lambda_name); + let free_vars = &lambda_decl.closure_vars; // Build closure struct on stack. let closure_ptr_val: PointerValue<'ctx> = if !free_vars.is_empty() { let sz = (free_vars.len() * 8) as u64; let clos_storage = self.entry_array_alloca(self.i8_ty(), sz, "closure"); - for (i, (name, ty)) in free_vars.iter().enumerate() { - let var_ptr = self.get_var_addr(name, *ty); + for (i, local) in free_vars.iter().enumerate() { + let var_ptr = self.get_var_addr(local, decl.arena.local(*local).ty); let slot = self.ptr_at_offset(clos_storage, i as u64 * 8); self.builder().build_store(slot, var_ptr).unwrap(); } @@ -3893,30 +3889,6 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { self.ptr_ty().const_null() }; - let closure_vars: Vec = free_vars - .iter() - .map(|(n, ty)| ClosureVar { - name: Name::new(n.clone()), - ty: *ty, - }) - .collect(); - - let lambda_decl = FuncDecl { - name: lambda_name, - typevars: vec![], - size_vars: vec![], - params: lambda_params, - body: Some(body), - ret: rng, - constraints: vec![], - requires: vec![], - loc: decl.loc, - arena: decl.arena.clone(), - types: decl.types.clone(), - closure_vars, - is_extern: false, - }; - self.pending_lambdas.push(lambda_decl); // Declare the lambda function to get a pointer. @@ -3938,133 +3910,3 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } } } - -// ─── Free variable collection ─────────────────────────────────────────────────── - -fn collect_free_var_names_llvm( - body: ExprID, - arena: &ExprArena, - exclude: &HashSet, - local_vars: &HashMap>, - types: &[crate::TypeID], -) -> Vec<(String, crate::TypeID)> { - let mut result = Vec::new(); - let mut seen = HashSet::new(); - collect_free_vars_rec_llvm( - body, - arena, - exclude, - local_vars, - types, - &mut result, - &mut seen, - ); - result -} - -fn collect_free_vars_rec_llvm( - expr: ExprID, - arena: &ExprArena, - exclude: &HashSet, - local_vars: &HashMap>, - types: &[crate::TypeID], - result: &mut Vec<(String, crate::TypeID)>, - seen: &mut HashSet, -) { - match &arena[expr] { - Expr::Id(name) => { - let s = name.to_string(); - if local_vars.contains_key(&s) && !exclude.contains(&s) && !seen.contains(&s) { - result.push((s.clone(), types[expr])); - seen.insert(s); - } - } - Expr::Call(fn_id, args) => { - collect_free_vars_rec_llvm(*fn_id, arena, exclude, local_vars, types, result, seen); - for a in args { - collect_free_vars_rec_llvm(*a, arena, exclude, local_vars, types, result, seen); - } - } - Expr::Binop(_, lhs, rhs) => { - collect_free_vars_rec_llvm(*lhs, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec_llvm(*rhs, arena, exclude, local_vars, types, result, seen); - } - Expr::Unop(_, arg) => { - collect_free_vars_rec_llvm(*arg, arena, exclude, local_vars, types, result, seen) - } - Expr::Let(_, init, _) => { - collect_free_vars_rec_llvm(*init, arena, exclude, local_vars, types, result, seen) - } - Expr::Var(_, init, _) => { - if let Some(i) = init { - collect_free_vars_rec_llvm(*i, arena, exclude, local_vars, types, result, seen); - } - } - Expr::If(c, t, e) => { - collect_free_vars_rec_llvm(*c, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec_llvm(*t, arena, exclude, local_vars, types, result, seen); - if let Some(e) = e { - collect_free_vars_rec_llvm(*e, arena, exclude, local_vars, types, result, seen); - } - } - Expr::While(c, b) => { - collect_free_vars_rec_llvm(*c, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec_llvm(*b, arena, exclude, local_vars, types, result, seen); - } - Expr::For { - start, end, body, .. - } => { - collect_free_vars_rec_llvm(*start, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec_llvm(*end, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec_llvm(*body, arena, exclude, local_vars, types, result, seen); - } - Expr::Block(exprs) => { - for e in exprs { - collect_free_vars_rec_llvm(*e, arena, exclude, local_vars, types, result, seen); - } - } - Expr::Return(e) | Expr::Assume(e) => { - collect_free_vars_rec_llvm(*e, arena, exclude, local_vars, types, result, seen) - } - Expr::Field(e, _) => { - collect_free_vars_rec_llvm(*e, arena, exclude, local_vars, types, result, seen) - } - Expr::ArrayIndex(a, i) => { - collect_free_vars_rec_llvm(*a, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec_llvm(*i, arena, exclude, local_vars, types, result, seen); - } - Expr::ArrayLiteral(elems) => { - for e in elems { - collect_free_vars_rec_llvm(*e, arena, exclude, local_vars, types, result, seen); - } - } - Expr::Tuple(elems) => { - for e in elems { - collect_free_vars_rec_llvm(*e, arena, exclude, local_vars, types, result, seen); - } - } - Expr::AsTy(e, _) => { - collect_free_vars_rec_llvm(*e, arena, exclude, local_vars, types, result, seen) - } - Expr::Arena(e) => { - collect_free_vars_rec_llvm(*e, arena, exclude, local_vars, types, result, seen) - } - Expr::Array(t, s) => { - collect_free_vars_rec_llvm(*t, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec_llvm(*s, arena, exclude, local_vars, types, result, seen); - } - Expr::Lambda { params, body } => { - let mut inner = exclude.clone(); - for p in params { - inner.insert(p.name.to_string()); - } - collect_free_vars_rec_llvm(*body, arena, &inner, local_vars, types, result, seen); - } - Expr::StructLit(_, fields) => { - for (_, fval) in fields { - collect_free_vars_rec_llvm(*fval, arena, exclude, local_vars, types, result, seen); - } - } - _ => {} - } -} diff --git a/src/monomorph.rs b/src/monomorph.rs index d0095f3f..9964f707 100644 --- a/src/monomorph.rs +++ b/src/monomorph.rs @@ -1,198 +1,77 @@ use crate::*; -/// Key that uniquely identifies a monomorphized function or struct. +/// A specialization is a source definition and its concrete arguments. Symbol +/// spelling and overload signatures are deliberately absent from its identity. #[derive(Clone, Hash, Eq, PartialEq, Debug)] pub struct MonomorphKey { - pub name: Name, + pub definition: DefId, pub type_args: Vec, - /// Integer size arguments (from size_vars), in declaration order. pub size_args: Vec, - /// Concrete parameter types, used to disambiguate generic overloads - /// that share the same name and type args but differ in arity/params - /// (e.g. `new()` vs `new(cap: i32)`). - pub param_types: Vec, } impl MonomorphKey { - pub fn new(name: Name, type_args: Vec) -> Self { + pub fn new(definition: DefId, type_args: Vec, size_args: Vec) -> Self { Self { - name, - type_args, - size_args: vec![], - param_types: vec![], - } - } - - pub fn new_with_sizes(name: Name, type_args: Vec, size_args: Vec) -> Self { - Self { - name, + definition, type_args, size_args, - param_types: vec![], - } - } - - pub fn with_param_types(mut self, param_types: Vec) -> Self { - self.param_types = param_types; - self - } - - pub fn mangled_name(&self) -> Name { - let base = crate::mangle::mangle_name(self.name, &self.type_args); - let with_sizes = if self.size_args.is_empty() { - base - } else { - let suffix: String = self.size_args.iter().map(|n| format!("${}", n)).collect(); - Name::new(format!("{}{}", base, suffix)) - }; - if self.param_types.is_empty() { - with_sizes - } else { - let param_suffix = - crate::mangle::mangle_name(Name::new(String::new()), &self.param_types); - Name::new(format!("{}#{}", with_sizes, param_suffix)) } } } -/// Tracks monomorphization in progress to detect infinite recursion. -/// -/// Generic recursion happens when instantiating a generic function requires -/// instantiating it again with different (but related) type arguments, potentially -/// creating an infinite chain. -/// -/// # Types of Generic Recursion -/// -/// ## 1. Direct Infinite Recursion (ALWAYS INFINITE) -/// ```ignore -/// fn foo() { -/// foo>() // Creates foo>, foo>>, ... -/// } -/// ``` -/// -/// ## 2. Mutually Recursive (ALWAYS INFINITE) -/// ```ignore -/// fn f() { g>() } -/// fn g() { f>() } -/// // Creates f, g>, f>>, ... -/// ``` -/// -/// ## 3. Bounded Recursion (FINITE - OK!) -/// ```ignore -/// fn recurse(depth: i32) { -/// if depth > 0 { -/// recurse(depth - 1) // Same type args! -/// } -/// } -/// ``` -/// -/// ## 4. Type-decreasing Recursion (FINITE - OK!) -/// ```ignore -/// fn unwrap(x: Vec>) { -/// unwrap(inner) // Type gets simpler, not more complex -/// } -/// ``` -/// -/// # Detection Strategy -/// -/// We use a "stack" of currently-being-instantiated generics. If we try to -/// instantiate something already on the stack with MORE COMPLEX type args, -/// we've found infinite recursion. +/// Retains the language's existing increasing-type-complexity guard. Ordinary +/// recursion is handled by reserving an instance before walking its body. #[derive(Debug, Default)] pub struct RecursionDetector { - /// Stack of monomorphizations currently in progress. - /// If we see the same (name, type_args) twice, we have recursion. in_progress: Vec, } impl RecursionDetector { pub fn new() -> Self { - Self { - in_progress: Vec::new(), - } + Self::default() } - /// Checks if instantiating this key would cause infinite recursion. - /// - /// Returns Ok(()) if safe, Err with error message if recursion detected. - pub fn check(&self, key: &MonomorphKey) -> Result<(), String> { - // Simple check: is this exact instantiation already in progress? + pub fn check(&self, key: &MonomorphKey, name: Name) -> Result<(), String> { if self.in_progress.contains(key) { return Err(format!( "Infinite generic recursion detected: {} with type args {:?} is already being instantiated", - key.name, key.type_args + name, key.type_args )); } - - // Advanced check: is the same function being instantiated with - // increasingly complex type arguments? - for in_progress_key in &self.in_progress { - if in_progress_key.name == key.name { - // Same function name - check if types are getting more complex - if is_more_complex(&key.type_args, &in_progress_key.type_args) { - return Err(format!( - "Infinite generic recursion detected: {} is being instantiated with increasingly complex types.\n\ - Previous: {:?}\n\ - Current: {:?}", - key.name, in_progress_key.type_args, key.type_args - )); - } + for previous in &self.in_progress { + if previous.definition == key.definition + && is_more_complex(&key.type_args, &previous.type_args) + { + return Err(format!( + "Infinite generic recursion detected: {} is being instantiated with increasingly complex types.\n\ + Previous: {:?}\n\ + Current: {:?}", + name, previous.type_args, key.type_args + )); } } - Ok(()) } - /// Marks that we're starting to instantiate this key. - /// Must be paired with `end_instantiation`. pub fn begin_instantiation(&mut self, key: MonomorphKey) { self.in_progress.push(key); } - /// Marks that we've finished instantiating (successfully or not). pub fn end_instantiation(&mut self) { self.in_progress.pop(); } - - /// Gets the current instantiation depth (useful for debugging). - pub fn depth(&self) -> usize { - self.in_progress.len() - } - - /// Gets the stack of in-progress instantiations (for error reporting). - pub fn stack(&self) -> &[MonomorphKey] { - &self.in_progress - } } -/// Checks if `current` type arguments are "more complex" than `previous`. -/// -/// This is a heuristic to detect infinite generic recursion. If types are -/// getting progressively more nested, we're likely in an infinite loop. -/// -/// Examples: -/// - `Vec` is more complex than `i32` -/// - `Vec>` is more complex than `Vec` -/// - `(i32, bool)` is NOT more complex than `i32` (different structure) fn is_more_complex(current: &[TypeID], previous: &[TypeID]) -> bool { - if current.len() != previous.len() { - return false; - } - - for (curr_ty, prev_ty) in current.iter().zip(previous) { - if type_complexity(*curr_ty) > type_complexity(*prev_ty) { - return true; - } - } - - false + current.len() == previous.len() + && current + .iter() + .zip(previous) + .any(|(current, previous)| type_complexity(*current) > type_complexity(*previous)) } -/// Computes a "complexity score" for a type. -/// Higher scores mean more nested/complex types. fn type_complexity(ty: TypeID) -> usize { match &*ty { - // Primitives have complexity 0 Type::Void | Type::Bool | Type::Int8 @@ -201,33 +80,27 @@ fn type_complexity(ty: TypeID) -> usize { | Type::UInt32 | Type::Float32 | Type::Float64 - | Type::Float32x4 => 0, - - // Type variables have complexity 0 (they're placeholders) - Type::Var(_) | Type::Anon(_) => 0, - - // Arrays/slices add 1 + complexity of element + | Type::Float32x4 + | Type::Var(_) + | Type::Anon(_) => 0, Type::Array(elem, _) | Type::Slice(elem) | Type::Reference(elem) => { 1 + type_complexity(*elem) } - - // Tuples: 1 + max complexity of elements - Type::Tuple(types) => 1 + types.iter().map(|t| type_complexity(*t)).max().unwrap_or(0), - - // Functions: 1 + max of domain and range - Type::Func(dom, rng) => 1 + type_complexity(*dom).max(type_complexity(*rng)), - - // Named types (including generics): 1 + max complexity of parameters + Type::Tuple(types) => { + 1 + types + .iter() + .map(|ty| type_complexity(*ty)) + .max() + .unwrap_or(0) + } + Type::Func(domain, result) => 1 + type_complexity(*domain).max(type_complexity(*result)), + Type::Name(_, params) if params.is_empty() => 0, Type::Name(_, params) => { - if params.is_empty() { - 0 // Simple named type like "MyStruct" - } else { - 1 + params - .iter() - .map(|t| type_complexity(*t)) - .max() - .unwrap_or(0) - } + 1 + params + .iter() + .map(|ty| type_complexity(*ty)) + .max() + .unwrap_or(0) } } } @@ -236,211 +109,32 @@ fn type_complexity(ty: TypeID) -> usize { mod tests { use super::*; - // Tests for type complexity calculation - - #[test] - fn test_complexity_primitives() { - assert_eq!(type_complexity(mk_type(Type::Int32)), 0); - assert_eq!(type_complexity(mk_type(Type::Bool)), 0); - assert_eq!(type_complexity(mk_type(Type::Float32)), 0); - } - - #[test] - fn test_complexity_array() { - let i32_array = mk_type(Type::Array(mk_type(Type::Int32), ArraySize::Known(10))); - assert_eq!(type_complexity(i32_array), 1); - - let nested_array = mk_type(Type::Array(i32_array, ArraySize::Known(5))); - assert_eq!(type_complexity(nested_array), 2); - } - - #[test] - fn test_complexity_named_types() { - // Simple named type - let simple = mk_type(Type::Name(Name::new("MyStruct".into()), vec![])); - assert_eq!(type_complexity(simple), 0); - - // Generic with one param - let generic = mk_type(Type::Name( - Name::new("Vec".into()), - vec![mk_type(Type::Int32)], - )); - assert_eq!(type_complexity(generic), 1); - - // Nested generic Vec> - let nested = mk_type(Type::Name(Name::new("Vec".into()), vec![generic])); - assert_eq!(type_complexity(nested), 2); - } - #[test] - fn test_is_more_complex_simple() { - let i32_ty = mk_type(Type::Int32); - let vec_i32 = mk_type(Type::Name(Name::new("Vec".into()), vec![i32_ty])); - - // Vec is more complex than i32 - assert!(is_more_complex(&[vec_i32], &[i32_ty])); - - // i32 is not more complex than Vec - assert!(!is_more_complex(&[i32_ty], &[vec_i32])); - - // Same complexity - assert!(!is_more_complex(&[i32_ty], &[i32_ty])); - } - - #[test] - fn test_is_more_complex_nested() { - let i32_ty = mk_type(Type::Int32); - let vec_i32 = mk_type(Type::Name(Name::new("Vec".into()), vec![i32_ty])); - let vec_vec_i32 = mk_type(Type::Name(Name::new("Vec".into()), vec![vec_i32])); - - // Vec> is more complex than Vec - assert!(is_more_complex(&[vec_vec_i32], &[vec_i32])); - - // Vec> is more complex than i32 - assert!(is_more_complex(&[vec_vec_i32], &[i32_ty])); - } - - // Tests for RecursionDetector - - #[test] - fn test_recursion_detector_no_recursion() { + fn recursion_checks_definitions_rather_than_shared_spelling() { let mut detector = RecursionDetector::new(); - let key1 = MonomorphKey::new(Name::new("foo".into()), vec![mk_type(Type::Int32)]); - - assert!(detector.check(&key1).is_ok()); - detector.begin_instantiation(key1.clone()); - assert_eq!(detector.depth(), 1); - - // Different function is OK - let key2 = MonomorphKey::new(Name::new("bar".into()), vec![mk_type(Type::Int32)]); - assert!(detector.check(&key2).is_ok()); - + let first = MonomorphKey::new(DefId(0), vec![mk_type(Type::Int32)], vec![]); + detector.begin_instantiation(first.clone()); + let nested = mk_type(Type::Array(mk_type(Type::Int32), ArraySize::Known(3))); + let same_name_other_overload = MonomorphKey::new(DefId(1), vec![nested], vec![]); + assert!(detector + .check(&same_name_other_overload, Name::str("f")) + .is_ok()); + let growing = MonomorphKey::new(DefId(0), vec![nested], vec![]); + assert!(detector.check(&growing, Name::str("f")).is_err()); detector.end_instantiation(); - assert_eq!(detector.depth(), 0); - } - - #[test] - fn test_recursion_detector_same_instantiation() { - let mut detector = RecursionDetector::new(); - let key = MonomorphKey::new(Name::new("foo".into()), vec![mk_type(Type::Int32)]); - - detector.begin_instantiation(key.clone()); - - // Trying to instantiate the same thing again - recursion! - let result = detector.check(&key); - assert!(result.is_err()); - assert!(result.unwrap_err().contains("already being instantiated")); - } - - #[test] - fn test_recursion_detector_increasing_complexity() { - let mut detector = RecursionDetector::new(); - let name = Name::new("foo".into()); - - let i32_ty = mk_type(Type::Int32); - let vec_i32 = mk_type(Type::Name(Name::new("Vec".into()), vec![i32_ty])); - let vec_vec_i32 = mk_type(Type::Name(Name::new("Vec".into()), vec![vec_i32])); - - let key1 = MonomorphKey::new(name, vec![i32_ty]); - detector.begin_instantiation(key1); - - // Trying to instantiate foo> while foo is in progress - let key2 = MonomorphKey::new(name, vec![vec_i32]); - let result = detector.check(&key2); - assert!(result.is_err()); - assert!(result.unwrap_err().contains("increasingly complex")); - - // Even more complex should also fail - let key3 = MonomorphKey::new(name, vec![vec_vec_i32]); - let result = detector.check(&key3); - assert!(result.is_err()); - } - - #[test] - fn test_recursion_detector_decreasing_complexity_ok() { - let mut detector = RecursionDetector::new(); - let name = Name::new("unwrap".into()); - - let i32_ty = mk_type(Type::Int32); - let vec_i32 = mk_type(Type::Name(Name::new("Vec".into()), vec![i32_ty])); - - let key1 = MonomorphKey::new(name, vec![vec_i32]); - detector.begin_instantiation(key1); - - // Trying to instantiate unwrap while unwrap> is in progress - // This is OK - complexity is decreasing (type-decreasing recursion) - let key2 = MonomorphKey::new(name, vec![i32_ty]); - let result = detector.check(&key2); - assert!(result.is_ok()); + assert!(detector.check(&growing, Name::str("f")).is_ok()); } #[test] - fn test_recursion_detector_multiple_params() { + fn decreasing_types_remain_accepted() { let mut detector = RecursionDetector::new(); - let name = Name::new("foo".into()); - - let i32_ty = mk_type(Type::Int32); - let bool_ty = mk_type(Type::Bool); - let vec_i32 = mk_type(Type::Name(Name::new("Vec".into()), vec![i32_ty])); - - let key1 = MonomorphKey::new(name, vec![i32_ty, bool_ty]); - detector.begin_instantiation(key1); - - // Increasing complexity in first param - let key2 = MonomorphKey::new(name, vec![vec_i32, bool_ty]); - let result = detector.check(&key2); - assert!(result.is_err()); - } - - #[test] - fn test_recursion_detector_stack() { - let mut detector = RecursionDetector::new(); - - let key1 = MonomorphKey::new(Name::new("f".into()), vec![mk_type(Type::Int32)]); - let key2 = MonomorphKey::new(Name::new("g".into()), vec![mk_type(Type::Bool)]); - - detector.begin_instantiation(key1.clone()); - detector.begin_instantiation(key2.clone()); - - assert_eq!(detector.depth(), 2); - - let stack = detector.stack(); - assert_eq!(stack.len(), 2); - assert_eq!(stack[0], key1); - assert_eq!(stack[1], key2); - - detector.end_instantiation(); - assert_eq!(detector.depth(), 1); - - detector.end_instantiation(); - assert_eq!(detector.depth(), 0); - } - - #[test] - fn test_monomorph_key() { - let name = Name::new("foo".into()); - let type_args = vec![mk_type(Type::Int32), mk_type(Type::Bool)]; - let key = MonomorphKey::new(name, type_args); - - assert_eq!(key.mangled_name(), Name::new("foo$i32$bool".into())); - } - - #[test] - fn test_monomorph_key_equality() { - let name = Name::new("foo".into()); - let type_args = vec![mk_type(Type::Int32)]; - - let key1 = MonomorphKey::new(name, type_args.clone()); - let key2 = MonomorphKey::new(name, type_args); - - assert_eq!(key1, key2); - - // Different name - let key3 = MonomorphKey::new(Name::new("bar".into()), vec![mk_type(Type::Int32)]); - assert_ne!(key1, key3); - - // Different type args - let key4 = MonomorphKey::new(name, vec![mk_type(Type::Bool)]); - assert_ne!(key1, key4); + let nested = mk_type(Type::Array(mk_type(Type::Int32), ArraySize::Known(3))); + detector.begin_instantiation(MonomorphKey::new(DefId(0), vec![nested], vec![])); + assert!(detector + .check( + &MonomorphKey::new(DefId(0), vec![mk_type(Type::Int32)], vec![]), + Name::str("f") + ) + .is_ok()); } } diff --git a/src/monomorph_pass.rs b/src/monomorph_pass.rs index ca5c9c73..fa0c3913 100644 --- a/src/monomorph_pass.rs +++ b/src/monomorph_pass.rs @@ -1,1282 +1,902 @@ use crate::*; -use std::collections::{HashMap, HashSet}; +use std::collections::HashMap; -/// Manages the monomorphization process, generating specialized versions -/// of generic functions and structs. +/// Converts checked definitions into concrete instances. The cache is reserved +/// before visiting a body, so recursive calls and generic globals always name +/// the same instance, regardless of the path by which they are reached. pub struct MonomorphPass { - /// Maps (function_name, type_args) -> mangled_name of the specialized version - instantiations: HashMap, - - /// Detects infinite generic recursion - recursion_detector: RecursionDetector, - - /// Newly generated specialized declarations - out_decls: Vec, + cache: HashMap, + records: Vec, + declarations: Vec>, + recursion: RecursionDetector, +} - /// Non-generic functions whose bodies have been processed (to prevent reprocessing). - processed_non_generic: HashSet, +impl Default for MonomorphPass { + fn default() -> Self { + Self { + cache: HashMap::new(), + records: vec![], + declarations: vec![], + recursion: RecursionDetector::new(), + } + } } impl MonomorphPass { pub fn new() -> Self { - Self { - instantiations: HashMap::new(), - recursion_detector: RecursionDetector::new(), - out_decls: Vec::new(), - processed_non_generic: HashSet::new(), - } + Self::default() } - /// Main entry point: monomorphize all functions starting from the entry point. - /// - /// This performs a demand-driven monomorphization, only specializing - /// generic functions that are actually called with concrete types. - /// - /// Returns all reached function declarations and other declarations. pub fn monomorphize( &mut self, - decls: &DeclTable, + program: &CheckedProgram, entry_point: Name, - ) -> Result, String> { - self.monomorphize_multi(decls, &[entry_point]) + ) -> Result { + self.monomorphize_multi(program, &[entry_point]) } - /// Monomorphize starting from multiple entry points. - /// - /// Each entry point is processed as a root. The `processed_non_generic` - /// set prevents reprocessing shared functions reached from multiple roots. - /// - /// Entry points that aren't defined are skipped — whether a missing entry - /// point is an error is up to the client. pub fn monomorphize_multi( &mut self, - decls: &DeclTable, + program: &CheckedProgram, entry_points: &[Name], - ) -> Result, String> { + ) -> Result { + program.validate()?; + let decls = &program.decls; for &entry_point in entry_points { - if decls.entry_point_overloads(entry_point).count() > 1 { + let roots: Vec<_> = decls + .named_ids(entry_point) + .into_iter() + .filter(|id| decls.function(*id).is_some()) + .collect(); + if roots.len() > 1 { return Err(format!( "Multiple overloads found for entry point function '{}'", entry_point )); } - - let Some(fdecl) = decls.find_entry_point(entry_point) else { - continue; - }; - - if !self.processed_non_generic.contains(&fdecl.name) { - self.processed_non_generic.insert(fdecl.name); - let mut fdecl = fdecl.clone(); - self.process_function(&mut fdecl, decls)?; - self.out_decls.push(Decl::Func(fdecl)); + if let Some(&definition) = roots.first() { + let source = decls.function(definition).unwrap(); + self.instantiate_function(definition, vec![], vec![], source, decls)?; } } - - for decl in decls.decls.iter() { - match decl { - Decl::Func(_) => { - // Functions are processed on demand + // Non-generic globals exist even when no reached function names them: + // their storage is part of the host-visible module layout. + for (index, declaration) in decls.decls.iter().enumerate() { + match declaration { + Decl::Global { typevars, ty, .. } if typevars.is_empty() => { + self.instantiate_global(decls.id_at(index), None, *ty, decls)?; } - Decl::Global { typevars, .. } if !typevars.is_empty() => { - // Generic globals — concrete instances emitted during process_expr - } - _ => { - self.out_decls.push(decl.clone()); + Decl::Func(_) | Decl::Global { .. } | Decl::Interface(_) | Decl::Macro(_) => {} + Decl::Assume { arena, cond } => { + let mut body = arena.clone(); + self.process_body(&mut body, std::iter::once(*cond), &[], decls)?; + self.declarations.push(Some(Decl::Assume { + arena: body, + cond: *cond, + })); } + _ => self.declarations.push(Some(declaration.clone())), } } + let declarations = self + .declarations + .iter() + .cloned() + .map(|decl| decl.ok_or_else(|| "Unfinished concrete instance".to_string())) + .collect::, _>>()?; + let specialized = SpecializedProgram::try_from_instances( + declarations, + self.records.clone(), + )?; + specialized.validate_origins(program)?; + Ok(specialized) + } + + fn reserve(&mut self, key: MonomorphKey) -> InstanceId { + let id = InstanceId(self.records.len() as u32); + self.records.push(InstanceRecord { + definition: key.definition, + type_args: key.type_args.clone(), + size_args: key.size_args.clone(), + declaration: self.declarations.len(), + }); + self.declarations.push(None); + self.cache.insert(key, id); + id + } - // Collect all declarations: original + specialized - Ok(self.out_decls.clone()) + fn finish(&mut self, id: InstanceId, declaration: CheckedDecl) { + let slot = self.records[id.0 as usize].declaration; + self.declarations[slot] = Some(declaration); } - /// Process a single function, finding all generic calls within it - fn process_function(&mut self, fdecl: &mut FuncDecl, decls: &DeclTable) -> Result<(), String> { - if let Some(body) = fdecl.body { - self.process_expr(body, fdecl, decls)?; + fn process_body( + &mut self, + body: &mut CheckedBody, + roots: impl IntoIterator, + size_vars: &[SizeParameter], + decls: &DeclTable, + ) -> Result<(), String> { + // Requirements are body-local. Keep their concrete selection in this + // stack frame while recursively specializing any selected callees. + let mut selections = HashMap::new(); + for requirement in &body.requirements { + let selected = requirement + .select(&Instance::new(), decls) + .map_err(|error| error.message(requirement.interface, decls))? + .ok_or("Unresolved body interface requirement")?; + for member in selected { + selections.insert((requirement.id, member.member), member.implementation); + } } + for root in roots { + self.process_expr(root, body, size_vars, decls, &selections)?; + } + // Publication validates every retained node, including nodes outside + // these roots, without instantiating otherwise unreachable code. + // Selected implementations are now explicit Instance references. Source + // requirement/member IDs have no owner in the concrete program. + body.requirements.clear(); Ok(()) } - /// Recursively process an expression, looking for function calls fn process_expr( &mut self, - expr_id: ExprID, - fdecl: &mut FuncDecl, - decls: &DeclTable, + id: ExprID, + body: &mut CheckedBody, + size_vars: &[SizeParameter], + decls: &DeclTable, + selections: &HashMap<(RequirementId, DefId), DefId>, ) -> Result<(), String> { - // Clone the expression to avoid borrow checker issues - let expr = fdecl.arena[expr_id].clone(); - - match &expr { - Expr::TypeApp(name, type_args) => { - // Explicit type application: name⟨i32⟩. - // The type args are known directly — no inference needed. - let fn_decls = decls.find(*name); - - // Check generic functions. When multiple overloads have the - // same number of type params, use the solved type to pick the - // right one (e.g. different arities like new⟨T⟩() vs new⟨T⟩(cap)). - let solved_type = fdecl.types[expr_id]; - for decl in fn_decls.iter() { - if let Decl::Func(target_fdecl) = decl { - if target_fdecl.typevars.len() == type_args.len() { - // Substitute explicit type args into the generic signature - // and check if it unifies with the solved call-site type. - let mut inst = Instance::new(); - for (tv, ta) in target_fdecl.typevars.iter().zip(type_args.iter()) { - inst.insert(mk_type(Type::Var(*tv)), *ta); - } - let candidate_ty = target_fdecl.ty().subst(&inst); - let mut unify_inst = Instance::new(); - if !unify(candidate_ty, solved_type, &mut unify_inst) { - continue; - } - let mangled = self.instantiate_function( - *name, - type_args.clone(), - target_fdecl, - decls, - )?; - fdecl.arena.exprs[expr_id] = Expr::Id(mangled); - return Ok(()); - } - } + match body[id].clone() { + Expr::Id(_) | Expr::TypeApp(_, _) => { + self.process_reference(id, &[], body, size_vars, decls, selections) + } + Expr::Call(callee, arguments) => { + // Preserve the prior specialization walk's argument order. + for &arg in &arguments { + self.process_expr(arg, body, size_vars, decls, selections)?; } - - // Check generic globals. - for decl in fn_decls.iter() { - if let Decl::Global { - name: gname, - typevars, - ty, - } = decl - { - if typevars.len() == type_args.len() { - let mut inst = Instance::new(); - for (tv, ta) in typevars.iter().zip(type_args.iter()) { - inst.insert(mk_type(Type::Var(*tv)), *ta); - } - let mangled = crate::mangle::mangle_name(*gname, type_args); - if !self.processed_non_generic.contains(&mangled) { - self.processed_non_generic.insert(mangled); - let concrete_ty = ty.subst(&inst); - self.out_decls.push(Decl::Global { - name: mangled, - typevars: vec![], - ty: concrete_ty, - }); - } - fdecl.arena.exprs[expr_id] = Expr::Id(mangled); - return Ok(()); - } - } + if matches!(body[callee], Expr::Id(_) | Expr::TypeApp(_, _)) { + self.process_reference(callee, &arguments, body, size_vars, decls, selections) + } else { + self.process_expr(callee, body, size_vars, decls, selections) } } - Expr::Id(name) => { - // Check if this identifier refers to a function - // Get the solved type for this expression - let solved_type = fdecl.types[expr_id]; - - // Look up the declaration - let fn_decls = decls.find(*name); - - // Only attempt function monomorphization if the solved type is a - // function type. If it's not, this identifier was resolved to a - // local variable/parameter by the type checker (which shadows - // global function names), so skip the function loop. - let is_func_type = matches!(*solved_type, Type::Func(_, _)); - if is_func_type { - for decl in fn_decls { - if let Decl::Func(target_fdecl) = decl { - if !target_fdecl.typevars.is_empty() { - // This is a generic function - compute type arguments from solved type - let type_args = - self.infer_type_arguments(target_fdecl, solved_type, fdecl)?; - - if !type_args.is_empty() { - // Create a specialized version - let mangled_name = self.instantiate_function( - *name, - type_args, - target_fdecl, - decls, - )?; - - // Rewrite the identifier to use the mangled name - fdecl.arena.exprs[expr_id] = Expr::Id(mangled_name); - } - } else { - // Non-generic function - include it and recursively process its body. - let non_generic_overload_count = fn_decls - .iter() - .filter(|d| matches!(d, Decl::Func(f) if f.typevars.is_empty())) - .count(); - - if non_generic_overload_count > 1 { - let func_ty = TypeID::new(Type::Func( - target_fdecl.domain(), - target_fdecl.ret, - )); - let mut inst = Instance::new(); - if unify(func_ty, solved_type, &mut inst) { - let param_types = target_fdecl.param_types(); - let mangled = - crate::mangle::mangle_overload(*name, ¶m_types); - - if !self.processed_non_generic.contains(&mangled) { - self.processed_non_generic.insert(mangled); - let mut func = target_fdecl.clone(); - func.name = mangled; - self.process_function(&mut func, decls)?; - self.out_decls.push(Decl::Func(func)); - } - - fdecl.arena.exprs[expr_id] = Expr::Id(mangled); - } - } else { - if !self.processed_non_generic.contains(&target_fdecl.name) { - self.processed_non_generic.insert(target_fdecl.name); - let mut func = target_fdecl.clone(); - self.process_function(&mut func, decls)?; - self.out_decls.push(Decl::Func(func)); - } - } - } - } - } - } // is_func_type - - // Check for generic globals. - for decl in fn_decls { - if let Decl::Global { - name: gname, - typevars, - ty, - } = decl - { - if !typevars.is_empty() { - // Infer type arguments by unifying the generic type with - // the solved type at this expression. - let mut inst = Instance::new(); - if unify_with_vars(*ty, solved_type, &mut inst) { - let type_args: Vec = typevars - .iter() - .map(|tv| { - let var_ty = mk_type(Type::Var(*tv)); - inst.get(&var_ty).copied().unwrap_or(var_ty) - }) - .collect(); - - let mangled = crate::mangle::mangle_name(*gname, &type_args); - - // Emit the concrete global if not already done. - if !self.processed_non_generic.contains(&mangled) { - self.processed_non_generic.insert(mangled); - let concrete_ty = ty.subst(&inst); - self.out_decls.push(Decl::Global { - name: mangled, - typevars: vec![], - ty: concrete_ty, - }); - } - - fdecl.arena.exprs[expr_id] = Expr::Id(mangled); - } - } - } + expression => { + for child in expression.subexprs() { + self.process_expr(child, body, size_vars, decls, selections)?; } + Ok(()) } - Expr::Call(fn_id, arg_ids) => { - // Process arguments first - let arg_ids = arg_ids.clone(); - for arg_id in &arg_ids { - self.process_expr(*arg_id, fdecl, decls)?; - } - - let fn_id = *fn_id; - - // Check if this is a call to a size-var generic function. - // We handle it here because we need the solved argument types. - if let Expr::Id(fn_name) = fdecl.arena[fn_id].clone() { - let fn_decls = decls.find(fn_name); - let size_var_func = fn_decls.iter().find_map(|d| { - if let Decl::Func(f) = d { - if !f.size_vars.is_empty() { - Some(f.clone()) - } else { - None - } - } else { - None - } - }); - if let Some(target_fdecl) = size_var_func { - // Infer size bindings from argument types vs generic param types. - let size_bindings = infer_size_bindings(&target_fdecl, &arg_ids, fdecl); - if !size_bindings.is_empty() { - let size_args: Vec = target_fdecl - .size_vars - .iter() - .map(|sv| size_bindings.get(sv).copied().unwrap_or(0)) - .collect(); - // Also infer type arguments if the function has type vars. - let type_args = if !target_fdecl.typevars.is_empty() { - let solved_type = fdecl.types[fn_id]; - self.infer_type_arguments(&target_fdecl, solved_type, fdecl)? - } else { - vec![] - }; - let mangled = self.instantiate_function_with_sizes( - fn_name, - type_args, - size_args, - &target_fdecl, - decls, - )?; - fdecl.arena.exprs[fn_id] = Expr::Id(mangled); - // Update ALL caller types to substitute the resolved size vars. - for ty in fdecl.types.iter_mut() { - *ty = subst_size_vars(*ty, &size_bindings); - } - return Ok(()); - } - } - } + } + } - // Process the function expression (which handles Id rewriting for type-var generics) - self.process_expr(fn_id, fdecl, decls)?; - } - Expr::Binop(_, lhs, rhs) => { - self.process_expr(*lhs, fdecl, decls)?; - self.process_expr(*rhs, fdecl, decls)?; - } - Expr::Unop(_, arg) => { - self.process_expr(*arg, fdecl, decls)?; + fn process_reference( + &mut self, + id: ExprID, + arguments: &[ExprID], + body: &mut CheckedBody, + size_vars: &[SizeParameter], + decls: &DeclTable, + selections: &HashMap<(RequirementId, DefId), DefId>, + ) -> Result<(), String> { + let explicit = match body[id].clone() { + Expr::Id(_) => None, + Expr::TypeApp(_, arguments) => Some(arguments), + _ => return Err("Expected a checked reference".into()), + }; + let reference = body + .reference(id) + .cloned() + .ok_or("Expected a checked reference")?; + let solved = body.ty(id); + let candidates = match reference { + Reference::Local(_) | Reference::Instance(_) => return Ok(()), + Reference::SizeParameter(_) => { + return Err(format_error(body.loc(id), "Unsubstituted size parameter")); } - Expr::Let(_, init, _) => { - self.process_expr(*init, fdecl, decls)?; + Reference::Global(definition) => { + let instance = self.instantiate_global(definition, explicit, solved, decls)?; + select_instance(body, id, instance); + return Ok(()); } - Expr::Var(_, init, _) => { - if let Some(init_id) = init { - self.process_expr(*init_id, fdecl, decls)?; + Reference::Functions(candidates) => candidates, + Reference::InterfaceMember { + requirement, + member, + } => vec![*selections + .get(&(requirement, member)) + .ok_or("Missing checked interface selection")?], + }; + + // Size-generic calls historically precede ordinary overload inference. + // Candidate order is the checker's order, never another name lookup. + if explicit.is_none() { + if let Some((definition, target)) = candidates.iter().find_map(|definition| { + decls + .function(*definition) + .filter(|target| !target.size_vars.is_empty()) + .map(|target| (*definition, target)) + }) { + let bindings = infer_size_bindings(target, arguments, body); + if !bindings.is_empty() { + let sizes = target + .size_vars + .iter() + .map(|parameter| bindings.get(¶meter.symbol).copied().unwrap_or(0)) + .collect(); + let types = infer_type_arguments(target, solved)?; + let instance = + self.instantiate_function(definition, types, sizes, target, decls)?; + select_instance(body, id, instance); + substitute_body_sizes(body, &bindings, size_vars); + return Ok(()); } } - Expr::Block(exprs) => { - for expr in exprs { - self.process_expr(*expr, fdecl, decls)?; + } + + // Preserve the old ordinary-overload diagnostic policy: every generic + // candidate must infer successfully for an implicit function reference. + let mut inferred = HashMap::new(); + if explicit.is_none() && matches!(*solved, Type::Func(_, _)) { + for &definition in &candidates { + if let Some(target) = decls.function(definition) { + if !target.typevars.is_empty() { + inferred.insert(definition, infer_type_arguments(target, solved)?); + } } } - Expr::Field(lhs, _) => { - self.process_expr(*lhs, fdecl, decls)?; - } - Expr::ArrayIndex(lhs, rhs) => { - self.process_expr(*lhs, fdecl, decls)?; - self.process_expr(*rhs, fdecl, decls)?; - } - Expr::ArrayLiteral(elements) => { - for elem in elements { - self.process_expr(*elem, fdecl, decls)?; + } + for definition in candidates { + if let Some(Decl::Global { ty, .. }) = decls.definition(definition) { + if !unify_with_vars(*ty, solved, &mut Instance::new()) { + continue; } + let instance = + self.instantiate_global(definition, explicit.clone(), solved, decls)?; + select_instance(body, id, instance); + return Ok(()); } - Expr::If(cond, then_expr, else_expr) => { - self.process_expr(*cond, fdecl, decls)?; - self.process_expr(*then_expr, fdecl, decls)?; - if let Some(else_id) = else_expr { - self.process_expr(*else_id, fdecl, decls)?; + let target = decls + .function(definition) + .ok_or("Resolved function declaration is missing")?; + let types = if let Some(types) = &explicit { + if target.typevars.len() != types.len() { + continue; } - } - Expr::While(cond, body) => { - self.process_expr(*cond, fdecl, decls)?; - self.process_expr(*body, fdecl, decls)?; - } - Expr::Lambda { body, .. } => { - self.process_expr(*body, fdecl, decls)?; - } - Expr::Return(inner) | Expr::Assume(inner) => { - self.process_expr(*inner, fdecl, decls)?; - } - Expr::For { - start, end, body, .. - } => { - let (start, end, body) = (*start, *end, *body); - self.process_expr(start, fdecl, decls)?; - self.process_expr(end, fdecl, decls)?; - self.process_expr(body, fdecl, decls)?; - } - Expr::AsTy(inner, _) => { - self.process_expr(*inner, fdecl, decls)?; - } - Expr::Tuple(elems) => { - for e in elems.clone() { - self.process_expr(e, fdecl, decls)?; + let substitution: Instance = target + .typevars + .iter() + .zip(types) + .map(|(variable, ty)| (mk_type(Type::Var(*variable)), *ty)) + .collect(); + if !unify( + target.ty().subst(&substitution), + solved, + &mut Instance::new(), + ) { + continue; } - } - Expr::Arena(inner) => { - self.process_expr(*inner, fdecl, decls)?; - } - Expr::StructLit(_, fields) => { - let fields = fields.clone(); - for (_, fval) in &fields { - self.process_expr(*fval, fdecl, decls)?; + types.clone() + } else if target.typevars.is_empty() { + if !unify(target.ty(), solved, &mut Instance::new()) { + continue; } - } - Expr::Array(val, sz) => { - let (val, sz) = (*val, *sz); - self.process_expr(val, fdecl, decls)?; - self.process_expr(sz, fdecl, decls)?; - } - // True leaves - Expr::Int(_, _) - | Expr::Real(_, _) - | Expr::String(_) - | Expr::Char(_) - | Expr::True - | Expr::False - | Expr::Enum(_) - | Expr::Error - | Expr::Macro(_, _) - | Expr::Break - | Expr::Continue => {} - } - Ok(()) - } - - /// Infer concrete type arguments for a generic function call. - fn infer_type_arguments( - &self, - generic_fdecl: &FuncDecl, - call_site_type: TypeID, - _caller_fdecl: &FuncDecl, - ) -> Result, String> { - if generic_fdecl.typevars.is_empty() { - return Ok(Vec::new()); - } - // Build the generic function type using the same convention as the - // type checker: domain is always a tuple of parameter types. - let generic_func_type = generic_fdecl.ty(); - // Convert named type variables (Var) to anonymous (Anon) for unification. - let mut fresh_index = 1000; - let mut fresh_inst = Instance::new(); - let fresh_func_type = generic_func_type.fresh_aux(&mut fresh_index, &mut fresh_inst); - // Build map from typevar name to its Anon type. - let var_to_anon: Vec<(Name, TypeID)> = generic_fdecl - .typevars - .iter() - .map(|tv_name| { - let tv = typevar(&tv_name.to_string()); - let anon_ty = fresh_inst.get(&tv).copied().unwrap_or(tv); - (*tv_name, anon_ty) - }) - .collect(); - // Unify the fresh generic type with the call-site type. - let mut unify_inst = Instance::new(); - if !unify(fresh_func_type, call_site_type, &mut unify_inst) { - return Err(format!( - "Cannot infer type arguments for {}", - generic_fdecl.name - )); - } - // Extract type arguments: look up what each Anon resolved to. - let mut type_args = Vec::new(); - for (_tv_name, anon_ty) in &var_to_anon { - let resolved = find(*anon_ty, &unify_inst); - type_args.push(resolved); + vec![] + } else { + let Some(types) = inferred.get(&definition) else { + continue; + }; + types.clone() + }; + let bindings = infer_size_bindings(target, arguments, body); + let sizes = target + .size_vars + .iter() + .map(|parameter| bindings.get(¶meter.symbol).copied().unwrap_or(0)) + .collect(); + let instance = self.instantiate_function(definition, types, sizes, target, decls)?; + select_instance(body, id, instance); + substitute_body_sizes(body, &bindings, size_vars); + return Ok(()); } - Ok(type_args) + Err(format_error( + body.loc(id), + &format!("No checked function candidate matches expression {}", id), + )) } - /// Create a specialized version of a generic function (type vars only). - fn instantiate_function( + fn instantiate_global( &mut self, - name: Name, - type_args: Vec, - generic_fdecl: &FuncDecl, - decls: &DeclTable, - ) -> Result { - self.instantiate_function_with_sizes(name, type_args, vec![], generic_fdecl, decls) - } - - /// Create a specialized version of a generic function with both type and size args. - fn instantiate_function_with_sizes( - &mut self, - name: Name, - type_args: Vec, - size_args: Vec, - generic_fdecl: &FuncDecl, - decls: &DeclTable, - ) -> Result { - // Check if this generic name has multiple generic overloads with the - // same typevar count. If so, include concrete param types in the key - // to disambiguate (e.g. new() vs new(cap: i32)). - let same_typevar_count = decls - .find(name) - .iter() - .filter(|d| { - matches!(d, Decl::Func(f) if !f.typevars.is_empty() - && f.typevars.len() == generic_fdecl.typevars.len()) - }) - .count(); - let key = if same_typevar_count > 1 { - let mut inst = Instance::new(); - for (tv, ta) in generic_fdecl.typevars.iter().zip(type_args.iter()) { - inst.insert(mk_type(Type::Var(*tv)), *ta); + definition: DefId, + explicit: Option>, + solved: TypeID, + decls: &DeclTable, + ) -> Result { + let Some(Decl::Global { name, typevars, ty }) = decls.definition(definition) else { + return Err("Resolved global declaration is missing".into()); + }; + let mut substitution = Instance::new(); + let types = if let Some(types) = explicit { + if types.len() != typevars.len() { + return Err(format!( + "Wrong number of checked global arguments for '{}'", + name + )); } - let concrete_params: Vec = generic_fdecl - .param_types() - .iter() - .map(|t| t.subst(&inst)) - .collect(); - MonomorphKey::new_with_sizes(name, type_args.clone(), size_args.clone()) - .with_param_types(concrete_params) + for (variable, ty) in typevars.iter().zip(&types) { + substitution.insert(mk_type(Type::Var(*variable)), *ty); + } + types + } else if typevars.is_empty() { + vec![] } else { - MonomorphKey::new_with_sizes(name, type_args.clone(), size_args.clone()) + if !unify_with_vars(*ty, solved, &mut substitution) { + return Err(format!( + "Cannot infer checked global arguments for '{}'", + name + )); + } + typevars + .iter() + .map(|variable| { + let variable = mk_type(Type::Var(*variable)); + substitution.get(&variable).copied().unwrap_or(variable) + }) + .collect() }; - - // Check if already instantiated - if let Some(mangled) = self.instantiations.get(&key) { - return Ok(*mangled); + let key = MonomorphKey::new(definition, types.clone(), vec![]); + if let Some(&id) = self.cache.get(&key) { + return Ok(id); } + let id = self.reserve(key); + self.finish( + id, + Decl::Global { + name: if types.is_empty() { + *name + } else { + crate::mangle::mangle_name(*name, &types) + }, + typevars: vec![], + ty: ty.subst(&substitution), + }, + ); + Ok(id) + } - // Check for infinite recursion - self.recursion_detector.check(&key)?; - self.recursion_detector.begin_instantiation(key.clone()); - - // Generate mangled name - let mangled_name = key.mangled_name(); - - // Create type substitution map - let mut instance = Instance::new(); - for (type_param, type_arg) in generic_fdecl.typevars.iter().zip(type_args.iter()) { - let type_var = typevar(&type_param.to_string()); - instance.insert(type_var, *type_arg); + fn instantiate_function( + &mut self, + definition: DefId, + types: Vec, + sizes: Vec, + source: &CheckedFunction, + decls: &DeclTable, + ) -> Result { + let key = MonomorphKey::new(definition, types.clone(), sizes.clone()); + if let Some(&id) = self.cache.get(&key) { + return Ok(id); } - - // Create size substitution map: Name -> i32 - let size_bindings: HashMap = generic_fdecl + let substitution: Instance = source + .typevars + .iter() + .zip(&types) + .map(|(variable, ty)| (mk_type(Type::Var(*variable)), *ty)) + .collect(); + let size_bindings: HashMap<_, _> = source .size_vars .iter() - .zip(size_args.iter()) - .map(|(sv, &n)| (*sv, n)) + .map(|parameter| parameter.symbol) + .zip(sizes) .collect(); - - // Clone and specialize the function declaration - let mut specialized = generic_fdecl.clone(); - specialized.name = mangled_name; - specialized.typevars = Vec::new(); - specialized.size_vars = Vec::new(); // No longer generic - - // Substitute types in return type (type vars + size vars in array sizes) - specialized.ret = subst_size_vars(specialized.ret.subst(&instance), &size_bindings); - - // Substitute types in parameters - for param in &mut specialized.params { - if let Some(ref mut ty) = param.ty { - *ty = subst_size_vars(ty.subst(&instance), &size_bindings); - } - } - - // Substitute types in the body's type map - for ty in specialized.types.iter_mut() { - *ty = subst_size_vars(ty.subst(&instance), &size_bindings); - } - - // Substitute type args in TypeApp expressions (e.g. pool⟨T⟩ → pool⟨i32⟩) - for expr in specialized.arena.exprs.iter_mut() { - if let Expr::TypeApp(_, ref mut args) = expr { - for arg in args.iter_mut() { - *arg = subst_size_vars(arg.subst(&instance), &size_bindings); - } - } - } - - // Substitute size vars in the body expressions (e.g. Expr::Id("N") → Expr::Int(3, None)) - if !size_bindings.is_empty() { - substitute_size_var_exprs(&mut specialized.arena, &size_bindings); - } - - // Record the instantiation - self.instantiations.insert(key, mangled_name); - - // Recursively process the specialized function's body immediately - self.process_function(&mut specialized, decls)?; - - // Add to specialized decls - self.out_decls.push(Decl::Func(specialized)); - - self.recursion_detector.end_instantiation(); - - Ok(mangled_name) - } - - /// Get all generated specialized declarations - pub fn specialized_declarations(&self) -> &[Decl] { - &self.out_decls + self.recursion.check(&key, source.name)?; + self.recursion.begin_instantiation(key.clone()); + let id = self.reserve(key); + // Whole-body copies retain local indices. Their containing InstanceId + // supplies the distinct owner; no occurrence or binding remap is needed. + let mut function = source.clone(); + function.name = instance_symbol(source, &types, &size_bindings, &substitution, decls); + function.typevars.clear(); + function.size_vars.clear(); + function.ret = subst_size_vars(function.ret.subst(&substitution), &size_bindings); + function.arena.substitute(&substitution); + substitute_body_sizes(&mut function.arena, &size_bindings, &source.size_vars); + let roots = function.requires.iter().copied().chain(function.body); + let result = self + .process_body(&mut function.arena, roots, &function.size_vars, decls) + .map_err(|error| format!("{} in '{}'", error, function.name)); + self.recursion.end_instantiation(); + result?; + self.finish(id, Decl::Func(function)); + Ok(id) } +} - /// The mangled names of all functions newly created by monomorphization - /// (i.e., specialized versions of generic functions). Excludes entry - /// points and any pre-existing non-generic functions. - pub fn instantiated_names(&self) -> impl Iterator + '_ { - self.instantiations.values().copied() +/// Symbols are a backend/diagnostic concern, never a specialization key. +fn instance_symbol( + source: &CheckedFunction, + types: &[TypeID], + sizes: &HashMap, + substitution: &Instance, + decls: &DeclTable, +) -> Name { + if source.typevars.is_empty() && source.size_vars.is_empty() { + let overloads = decls.find(source.name).iter().filter(|declaration| { + matches!(declaration, Decl::Func(function) if function.typevars.is_empty()) + }).count(); + return if overloads > 1 { + crate::mangle::mangle_overload(source.name, &source.param_types()) + } else { + source.name + }; } - - /// Get the mangled name for a specific instantiation, if it exists - pub fn get_instantiation(&self, key: &MonomorphKey) -> Option { - self.instantiations.get(key).copied() + let mut name = crate::mangle::mangle_name(source.name, types).to_string(); + for parameter in &source.size_vars { + name.push_str(&format!( + "${}", + sizes.get(¶meter.symbol).copied().unwrap_or(0) + )); + } + let overloads = decls + .find(source.name) + .iter() + .filter(|declaration| { + matches!(declaration, Decl::Func(function) if !function.typevars.is_empty() + && function.typevars.len() == source.typevars.len()) + }) + .count(); + if overloads > 1 { + let parameters: Vec<_> = source + .param_types() + .iter() + .map(|ty| ty.subst(substitution)) + .collect(); + let suffix = crate::mangle::mangle_name(Name::str(""), ¶meters); + name.push_str(&format!("#{}", suffix)); } + Name::new(name) +} - /// Get the full rewrite map for all instantiations - pub fn get_rewrite_map(&self) -> &HashMap { - &self.instantiations - } +fn infer_type_arguments(function: &CheckedFunction, solved: TypeID) -> Result, String> { + if function.typevars.is_empty() { + return Ok(vec![]); + } + let mut fresh_index = 1000; + let mut fresh_substitution = Instance::new(); + let fresh = function + .ty() + .fresh_aux(&mut fresh_index, &mut fresh_substitution); + let mut substitution = Instance::new(); + if !unify(fresh, solved, &mut substitution) { + return Err(format!("Cannot infer type arguments for {}", function.name)); + } + Ok(function + .typevars + .iter() + .map(|variable| { + let ty = mk_type(Type::Var(*variable)); + find( + fresh_substitution.get(&ty).copied().unwrap_or(ty), + &substitution, + ) + }) + .collect()) } -/// Substitute `ArraySize::Var(N)` → `Known(n)` in a type using the size bindings. fn subst_size_vars(ty: TypeID, bindings: &HashMap) -> TypeID { + if bindings.is_empty() { + return ty; + } match &*ty { - Type::Array(elem, ArraySize::Var(name)) => { - let elem2 = subst_size_vars(*elem, bindings); - let size = bindings - .get(name) - .copied() - .map(ArraySize::Known) - .unwrap_or_else(|| ArraySize::Var(*name)); - mk_type(Type::Array(elem2, size)) - } - Type::Array(elem, sz) => mk_type(Type::Array(subst_size_vars(*elem, bindings), sz.clone())), - Type::Tuple(vs) => mk_type(Type::Tuple( - vs.iter().map(|t| subst_size_vars(*t, bindings)).collect(), + Type::Array(element, size) => mk_type(Type::Array( + subst_size_vars(*element, bindings), + match size { + ArraySize::Var(name) => bindings + .get(name) + .copied() + .map(ArraySize::Known) + .unwrap_or_else(|| size.clone()), + _ => size.clone(), + }, + )), + Type::Slice(element) => mk_type(Type::Slice(subst_size_vars(*element, bindings))), + Type::Reference(element) => mk_type(Type::Reference(subst_size_vars(*element, bindings))), + Type::Tuple(types) => mk_type(Type::Tuple( + types + .iter() + .map(|ty| subst_size_vars(*ty, bindings)) + .collect(), )), - Type::Func(a, b) => mk_type(Type::Func( - subst_size_vars(*a, bindings), - subst_size_vars(*b, bindings), + Type::Func(domain, result) => mk_type(Type::Func( + subst_size_vars(*domain, bindings), + subst_size_vars(*result, bindings), )), - Type::Name(n, ps) => mk_type(Type::Name( - *n, - ps.iter().map(|t| subst_size_vars(*t, bindings)).collect(), + Type::Name(name, types) => mk_type(Type::Name( + *name, + types + .iter() + .map(|ty| subst_size_vars(*ty, bindings)) + .collect(), )), _ => ty, } } -/// Replace `Expr::Id(N)` with `Expr::Int(n, None)` for each size var binding, across all arena slots. -fn substitute_size_var_exprs(arena: &mut ExprArena, bindings: &HashMap) { - for slot in arena.exprs.iter_mut() { - if let Expr::Id(name) = slot { - if let Some(&val) = bindings.get(name) { - *slot = Expr::Int(val as i64, None); +/// Record the selected concrete target. An explicit application collapses to +/// a plain identifier: its type arguments are consumed by the selection. +fn select_instance(body: &mut CheckedBody, id: ExprID, instance: InstanceId) { + if let Expr::TypeApp(name, _) = body[id].clone() { + let ty = body.ty(id); + body.replace(id, Expr::Id(name), ty); + } + body.set_reference(id, Reference::Instance(instance)); +} + +fn substitute_body_sizes( + body: &mut CheckedBody, + bindings: &HashMap, + parameters: &[SizeParameter], +) { + if bindings.is_empty() { + return; + } + let values: HashMap<_, _> = parameters + .iter() + .filter_map(|parameter| { + bindings + .get(¶meter.symbol) + .map(|value| (parameter.local, *value)) + }) + .collect(); + for id in 0..body.len() { + let mut kind = body[id].clone(); + match &mut kind { + Expr::Id(_) => { + if let Some(Reference::SizeParameter(local)) = body.reference(id) { + if let Some(value) = values.get(local) { + kind = Expr::Int(i64::from(*value), None); + } + } + } + Expr::TypeApp(_, types) => { + for ty in types { + *ty = subst_size_vars(*ty, bindings); + } } + Expr::AsTy(_, ty) => *ty = subst_size_vars(*ty, bindings), + Expr::Let(_, _, Some(ty)) | Expr::Var(_, _, Some(ty)) => { + *ty = subst_size_vars(*ty, bindings) + } + _ => {} + } + body.replace(id, kind, subst_size_vars(body.ty(id), bindings)); + } + for local in &mut body.locals { + local.ty = subst_size_vars(local.ty, bindings); + } + for requirement in &mut body.requirements { + for ty in &mut requirement.type_args { + *ty = subst_size_vars(*ty, bindings); + } + for member in &mut requirement.members { + member.signature = subst_size_vars(member.signature, bindings); } } } -/// Walk a type pair (generic param type vs concrete arg type) to extract size var bindings. fn infer_size_bindings_pair(generic: TypeID, concrete: TypeID, out: &mut HashMap) { match (&*generic, &*concrete) { - (Type::Array(ge, ArraySize::Var(name)), Type::Array(ce, ArraySize::Known(n))) - if *n != 0 => - { - out.insert(*name, *n); - infer_size_bindings_pair(*ge, *ce, out); + ( + Type::Array(generic, ArraySize::Var(name)), + Type::Array(concrete, ArraySize::Known(size)), + ) if *size != 0 => { + out.insert(*name, *size); + infer_size_bindings_pair(*generic, *concrete, out); + } + (Type::Array(generic, _), Type::Array(concrete, _)) => { + infer_size_bindings_pair(*generic, *concrete, out) } - (Type::Array(ge, _), Type::Array(ce, _)) => infer_size_bindings_pair(*ge, *ce, out), - (Type::Tuple(gs), Type::Tuple(cs)) => { - for (g, c) in gs.iter().zip(cs.iter()) { - infer_size_bindings_pair(*g, *c, out); + (Type::Tuple(generic), Type::Tuple(concrete)) => { + for (generic, concrete) in generic.iter().zip(concrete) { + infer_size_bindings_pair(*generic, *concrete, out); } } - (Type::Func(ga, gb), Type::Func(ca, cb)) => { - infer_size_bindings_pair(*ga, *ca, out); - infer_size_bindings_pair(*gb, *cb, out); + (Type::Func(gd, gr), Type::Func(cd, cr)) => { + infer_size_bindings_pair(*gd, *cd, out); + infer_size_bindings_pair(*gr, *cr, out); } _ => {} } } -/// Infer size var bindings for a call to `target` given the solved types of the argument expressions. fn infer_size_bindings( - target: &FuncDecl, - arg_ids: &[ExprID], - caller: &FuncDecl, + target: &CheckedFunction, + args: &[ExprID], + caller: &CheckedBody, ) -> HashMap { - let mut out = HashMap::new(); - for (param, &arg_id) in target.params.iter().zip(arg_ids.iter()) { - if let Some(param_ty) = param.ty { - let concrete_ty = caller.types[arg_id]; - infer_size_bindings_pair(param_ty, concrete_ty, &mut out); - } + let mut bindings = HashMap::new(); + for (parameter, &argument) in target.params.iter().zip(args) { + infer_size_bindings_pair( + target.arena.local(parameter.local).ty, + caller.ty(argument), + &mut bindings, + ); } - out + bindings } #[cfg(test)] mod tests { use super::*; - fn mk_simple_func(name: &str, typevars: Vec<&str>) -> FuncDecl { - let mut arena = ExprArena::new(); - let body_expr = arena.add(Expr::Int(0, None), test_loc()); - - FuncDecl { - name: Name::str(name), - typevars: typevars.iter().map(|s| Name::str(s)).collect(), - size_vars: vec![], - params: Vec::new(), - constraints: Vec::new(), - requires: vec![], - ret: mk_type(Type::Void), - body: Some(body_expr), - arena, - types: vec![mk_type(Type::Int32)], - loc: test_loc(), - closure_vars: vec![], - is_extern: false, - } + fn checked(source: &str) -> CheckedProgram { + let mut compiler = Compiler::new(); + compiler.quiet = true; + assert!(compiler.parse(source, "checked-specialization.lyte")); + assert!(compiler.check(), "{:?}", compiler.last_errors); + compiler.checked_program().unwrap().clone() } - #[test] - fn test_monomorph_pass_creation() { - let pass = MonomorphPass::new(); - assert_eq!(pass.instantiations.len(), 0); - assert_eq!(pass.out_decls.len(), 0); + fn specialize(source: &str) -> SpecializedProgram { + MonomorphPass::new() + .monomorphize(&checked(source), Name::str("main")) + .unwrap() } - #[test] - fn test_instantiate_simple_function() { - let mut pass = MonomorphPass::new(); - let generic_func = mk_simple_func("id", vec!["T"]); - let decls = DeclTable::new(vec![]); - - let type_args = vec![mk_type(Type::Int32)]; - let result = pass.instantiate_function(Name::str("id"), type_args, &generic_func, &decls); - - assert!(result.is_ok()); - let mangled = result.unwrap(); - assert_eq!(mangled, Name::str("id$i32")); - assert_eq!(pass.out_decls.len(), 1); + fn targets(function: &CheckedFunction) -> Vec { + function + .arena + .ids() + .filter_map(|id| match function.arena.reference(id) { + Some(Reference::Instance(id)) => Some(*id), + _ => None, + }) + .collect() } #[test] - fn test_instantiate_same_function_twice() { - let mut pass = MonomorphPass::new(); - let generic_func = mk_simple_func("id", vec!["T"]); - let decls = DeclTable::new(vec![]); - - let type_args = vec![mk_type(Type::Int32)]; - - // First instantiation - let result1 = - pass.instantiate_function(Name::str("id"), type_args.clone(), &generic_func, &decls); - assert!(result1.is_ok()); - - // Second instantiation - should reuse - let result2 = pass.instantiate_function(Name::str("id"), type_args, &generic_func, &decls); - assert!(result2.is_ok()); - assert_eq!(result1.unwrap(), result2.unwrap()); - - // Should only have one specialized version - assert_eq!(pass.out_decls.len(), 1); + fn concrete_program_drops_fulfilled_interfaces_and_rejects_orphan_source_references() { + let mut source = checked("interface Identity { identify(x: T) -> T } identify(x: i32) -> i32 { x } forward(x: T) -> T where Identity { identify(x) } main { forward(1) }"); + let output = MonomorphPass::new() + .monomorphize(&source, Name::str("main")) + .unwrap(); + assert!(!output + .decls + .decls + .iter() + .any(|decl| matches!(decl, Decl::Interface(_) | Decl::Macro(_)))); + assert!(output + .functions() + .all(|(_, function)| function.arena.requirements.is_empty())); + let definition = source.decls.named_ids(Name::str("identify"))[0]; + let ty = source.function(definition).unwrap().ty(); + let mut records: Vec<_> = source.decls.records().collect(); + let main = records + .iter_mut() + .find_map(|record| match &mut record.declaration { + Decl::Func(function) if function.name == Name::str("main") => Some(function), + _ => None, + }) + .unwrap(); + main.arena.add_id( + Name::str("target"), + Reference::Functions(vec![definition]), + ty, + test_loc(), + ); + source.decls = DeclTable::from_records(records); + let error = MonomorphPass::new() + .monomorphize(&source, Name::str("main")) + .unwrap_err(); + assert!(error.contains("Unresolved checked reference"), "{}", error); } #[test] - fn test_instantiate_different_type_args() { - let mut pass = MonomorphPass::new(); - let generic_func = mk_simple_func("id", vec!["T"]); - let decls = DeclTable::new(vec![]); - - // id - let result1 = pass.instantiate_function( - Name::str("id"), - vec![mk_type(Type::Int32)], - &generic_func, - &decls, + fn size_parameter_diagnostic_names_do_not_control_substitution() { + let mut source = checked("probe(a: [i32; N]) -> i32 { N } main { probe([1, 2, 3]) }"); + let mut records: Vec<_> = source.decls.records().collect(); + let probe = records + .iter_mut() + .find_map(|record| match &mut record.declaration { + Decl::Func(function) if function.name == Name::str("probe") => Some(function), + _ => None, + }) + .unwrap(); + let parameter = probe.size_vars[0]; + assert_eq!(parameter.symbol, Name::str("N")); + probe.arena.locals[parameter.local.index()].name = Name::str("diagnostic_only"); + source.decls = DeclTable::from_records(records); + let output = MonomorphPass::new() + .monomorphize(&source, Name::str("main")) + .unwrap(); + let specialized = output.find_entry_point(Name::str("probe$3")).unwrap(); + assert!(specialized + .arena + .exprs() + .iter() + .any(|expr| *expr == Expr::Int(3, None))); + assert_eq!( + specialized.param_types(), + vec![mk_type(Type::Array( + mk_type(Type::Int32), + ArraySize::Known(3) + ))] ); - assert!(result1.is_ok()); - assert_eq!(result1.unwrap(), Name::str("id$i32")); - - // id - let result2 = pass.instantiate_function( - Name::str("id"), - vec![mk_type(Type::Bool)], - &generic_func, - &decls, + assert_eq!( + specialized.arena.local(parameter.local).name, + Name::str("diagnostic_only") ); - assert!(result2.is_ok()); - assert_eq!(result2.unwrap(), Name::str("id$bool")); - - // Should have two specialized versions - assert_eq!(pass.out_decls.len(), 2); } #[test] - fn test_instantiate_multiple_type_params() { - let mut pass = MonomorphPass::new(); - let generic_func = mk_simple_func("map", vec!["T0", "T1"]); - let decls = DeclTable::new(vec![]); - - let type_args = vec![mk_type(Type::Int32), mk_type(Type::Bool)]; - let result = pass.instantiate_function(Name::str("map"), type_args, &generic_func, &decls); - - assert!(result.is_ok()); - assert_eq!(result.unwrap(), Name::str("map$i32$bool")); - } - - #[test] - fn test_process_expr_block() { - let mut pass = MonomorphPass::new(); - let decls = DeclTable::new(vec![]); - - let mut arena = ExprArena::new(); - let expr1 = arena.add(Expr::Int(1, None), test_loc()); - let expr2 = arena.add(Expr::Int(2, None), test_loc()); - let block = arena.add(Expr::Block(vec![expr1, expr2]), test_loc()); - - let mut fdecl = FuncDecl { - name: Name::str("test"), - typevars: Vec::new(), - size_vars: Vec::new(), - params: Vec::new(), - constraints: Vec::new(), - requires: vec![], - ret: mk_type(Type::Void), - body: Some(block), - arena, - types: vec![mk_type(Type::Void); 3], - loc: test_loc(), - closure_vars: vec![], - is_extern: false, - }; - - let result = pass.process_expr(block, &mut fdecl, &decls); - assert!(result.is_ok()); - } - - #[test] - fn test_process_expr_binop() { - let mut pass = MonomorphPass::new(); - let decls = DeclTable::new(vec![]); - - let mut arena = ExprArena::new(); - let lhs = arena.add(Expr::Int(1, None), test_loc()); - let rhs = arena.add(Expr::Int(2, None), test_loc()); - let binop = arena.add(Expr::Binop(Binop::Plus, lhs, rhs), test_loc()); - - let mut fdecl = FuncDecl { - name: Name::str("test"), - typevars: Vec::new(), - size_vars: Vec::new(), - params: Vec::new(), - constraints: Vec::new(), - requires: vec![], - ret: mk_type(Type::Void), - body: Some(binop), - arena, - types: vec![mk_type(Type::Int32); 3], - loc: test_loc(), - closure_vars: vec![], - is_extern: false, - }; - - let result = pass.process_expr(binop, &mut fdecl, &decls); - assert!(result.is_ok()); + fn nested_specialization_keeps_interface_selection_body_local() { + let source = checked("interface Printable { to_int(x: T) -> i32 } to_int(x: i32) -> i32 { x } to_int(x: bool) -> i32 { if x { 1 } else { 0 } } nested(x: U) -> i32 where Printable { to_int(x) } show(x: T) -> i32 where Printable { let other = nested(true); to_int(x) } main { show(42) }"); + let integer_implementation = source + .decls + .named_ids(Name::str("to_int")) + .into_iter() + .find(|definition| { + source.function(*definition).unwrap().param_types() == vec![mk_type(Type::Int32)] + }) + .unwrap(); + let output = MonomorphPass::new() + .monomorphize(&source, Name::str("main")) + .unwrap(); + let show = output.find_entry_point(Name::str("show$i32")).unwrap(); + assert!(targets(show) + .iter() + .any(|target| output.instances[target.index()].definition == integer_implementation)); } #[test] - fn test_process_expr_array_literal() { - let mut pass = MonomorphPass::new(); - let decls = DeclTable::new(vec![]); - - let mut arena = ExprArena::new(); - let elem1 = arena.add(Expr::Int(1, None), test_loc()); - let elem2 = arena.add(Expr::Int(2, None), test_loc()); - let array = arena.add(Expr::ArrayLiteral(vec![elem1, elem2]), test_loc()); - - let mut fdecl = FuncDecl { - name: Name::str("test"), - typevars: Vec::new(), - size_vars: Vec::new(), - params: Vec::new(), - constraints: Vec::new(), - requires: vec![], - ret: mk_type(Type::Void), - body: Some(array), - arena, - types: vec![mk_type(Type::Int32); 3], - loc: test_loc(), - closure_vars: vec![], - is_extern: false, + fn global_assumptions_follow_the_concrete_storage_instance() { + let mut source = checked("var limit: i32 main {}"); + let global = source.decls.named_ids(Name::str("limit"))[0]; + let mut arena = CheckedBody::new(); + let reference = arena.add_id( + Name::str("limit"), + Reference::Global(global), + mk_type(Type::Int32), + test_loc(), + ); + let zero = arena.add(Expr::Int(0, None), mk_type(Type::Int32), test_loc()); + let condition = arena.add( + Expr::Binop(Binop::Geq, reference, zero), + mk_type(Type::Bool), + test_loc(), + ); + let mut records: Vec<_> = source.decls.records().collect(); + records.push(DeclRecord { + definition: DefId(source.decls.definition_count() as u32), + declaration: Decl::Assume { + arena, + cond: condition, + }, + members: vec![], + }); + source.decls = DeclTable::from_records(records); + let output = MonomorphPass::new() + .monomorphize(&source, Name::str("main")) + .unwrap(); + let assumption = output + .decls + .decls + .iter() + .find_map(|declaration| match declaration { + Decl::Assume { arena, .. } => Some(arena), + _ => None, + }) + .unwrap(); + let Some(&Reference::Instance(instance)) = assumption.reference(reference) else { + panic!("assumption retained a generic-phase reference"); }; - - let result = pass.process_expr(array, &mut fdecl, &decls); - assert!(result.is_ok()); + assert_eq!(output.instances[instance.index()].definition, global); + assert!(matches!(output.instance(instance), Decl::Global { .. })); } #[test] - fn test_get_instantiation() { - let mut pass = MonomorphPass::new(); - let generic_func = mk_simple_func("id", vec!["T"]); - let decls = DeclTable::new(vec![]); - - let type_args = vec![mk_type(Type::Int32)]; - let key = MonomorphKey::new(Name::str("id"), type_args.clone()); - - // Before instantiation - assert!(pass.get_instantiation(&key).is_none()); - - // After instantiation - pass.instantiate_function(Name::str("id"), type_args, &generic_func, &decls) + fn recursive_calls_and_multiple_roots_share_the_reserved_instance() { + let program = checked("recur(x: T, n: i32) -> T { if n > 0 { recur(x, n - 1) } else { x } } first { recur(1, 2) } second { recur(2, 3) }"); + let output = MonomorphPass::new() + .monomorphize_multi( + &program, + &[ + Name::str("missing"), + Name::str("first"), + Name::str("second"), + ], + ) .unwrap(); - assert_eq!(pass.get_instantiation(&key), Some(Name::str("id$i32"))); - } - - #[test] - fn test_specialized_declarations() { - let mut pass = MonomorphPass::new(); - let generic_func = mk_simple_func("id", vec!["T"]); - let decls = DeclTable::new(vec![]); - - assert_eq!(pass.specialized_declarations().len(), 0); - - pass.instantiate_function( - Name::str("id"), - vec![mk_type(Type::Int32)], - &generic_func, - &decls, - ) - .unwrap(); - - assert_eq!(pass.specialized_declarations().len(), 1); - } - - #[test] - fn test_type_substitution_in_specialized_func() { - let mut pass = MonomorphPass::new(); - let decls = DeclTable::new(vec![]); - - // Create a generic function id(x: T) -> T - let t_var = typevar("T"); - let mut generic_func = mk_simple_func("id", vec!["T"]); - generic_func.params = vec![Param { - name: Name::str("x"), - ty: Some(t_var), - }]; - generic_func.ret = t_var; - - // Instantiate with i32 - pass.instantiate_function( - Name::str("id"), - vec![mk_type(Type::Int32)], - &generic_func, - &decls, - ) - .unwrap(); - - // Check the specialized declaration - let specialized = &pass.out_decls[0]; - if let Decl::Func(fdecl) = specialized { - assert_eq!(fdecl.name, Name::str("id$i32")); - assert_eq!(fdecl.typevars.len(), 0); // No longer generic - - // Check that return type was substituted - assert_eq!(*fdecl.ret, Type::Int32); - - // Check that parameter type was substituted - assert_eq!(fdecl.params.len(), 1); - assert_eq!(*fdecl.params[0].ty.unwrap(), Type::Int32); - } else { - panic!("Expected function declaration"); + let recur = output.instance_for_entry(Name::str("recur$i32")).unwrap(); + assert_eq!(output.find(Name::str("recur$i32")).len(), 1); + assert!(targets(output.function_instance(recur).unwrap()).contains(&recur)); + for name in ["first", "second"] { + let function = output.find_entry_point(Name::str(name)).unwrap(); + assert!(targets(function).contains(&recur)); } } #[test] - fn test_monomorphize_with_empty_decls() { - let mut pass = MonomorphPass::new(); - let decls = DeclTable::new(vec![]); - - // A missing entry point is not an error — it's simply skipped. - let result = pass.monomorphize(&decls, Name::str("main")); - assert_eq!(result.unwrap().len(), 0); - } - - #[test] - fn test_instantiate_with_nested_generic() { - let mut pass = MonomorphPass::new(); - let decls = DeclTable::new(vec![]); - - // Create a generic function that works with nested types - let mut generic_func = mk_simple_func("process", vec!["T"]); - let t_var = typevar("T"); - let array_of_t = mk_type(Type::Array(t_var, ArraySize::Known(10))); - generic_func.params = vec![Param { - name: Name::str("arr"), - ty: Some(array_of_t), - }]; - generic_func.ret = t_var; - - // Instantiate with i32 -> should create process$i32 - let result = pass.instantiate_function( - Name::str("process"), - vec![mk_type(Type::Int32)], - &generic_func, - &decls, + fn generic_global_storage_is_shared_across_function_instances() { + let source = checked("var pool: [T; 4] use(x: T) { let a = pool⟨T⟩ } main { use(1); use(true); let b = pool⟨i32⟩ }"); + let definition = source.decls.named_ids(Name::str("pool"))[0]; + let output = MonomorphPass::new() + .monomorphize(&source, Name::str("main")) + .unwrap(); + let globals: Vec<_> = output + .instances + .iter() + .enumerate() + .filter(|(_, record)| record.definition == definition) + .collect(); + assert_eq!(globals.len(), 2); + let integer_global = globals + .iter() + .find(|(_, record)| record.type_args == vec![mk_type(Type::Int32)]) + .unwrap() + .0; + let integer_global = InstanceId(integer_global as u32); + assert!( + targets(output.find_entry_point(Name::str("main")).unwrap()).contains(&integer_global) ); - - assert!(result.is_ok()); - assert_eq!(result.unwrap(), Name::str("process$i32")); - - // Check that the specialized version has the correct array type - let specialized = &pass.out_decls[0]; - if let Decl::Func(fdecl) = specialized { - assert_eq!(fdecl.params.len(), 1); - if let Type::Array(elem_ty, size) = &*fdecl.params[0].ty.unwrap() { - assert_eq!(**elem_ty, Type::Int32); - assert_eq!(*size, ArraySize::Known(10)); - } else { - panic!("Expected array type"); - } - } - } - - #[test] - fn test_multiple_instantiations_same_function() { - let mut pass = MonomorphPass::new(); - let decls = DeclTable::new(vec![]); - - let generic_func = mk_simple_func("id", vec!["T"]); - - // Create three different instantiations - let types = vec![ - mk_type(Type::Int32), - mk_type(Type::Bool), - mk_type(Type::Float32), - ]; - - for ty in types { - pass.instantiate_function(Name::str("id"), vec![ty], &generic_func, &decls) - .unwrap(); - } - - // Should have 3 specialized versions - assert_eq!(pass.out_decls.len(), 3); - assert_eq!(pass.instantiations.len(), 3); - } - - #[test] - fn test_instantiate_function_with_constraints() { - let mut pass = MonomorphPass::new(); - let decls = DeclTable::new(vec![]); - - // Create a generic function with interface constraints - let mut generic_func = mk_simple_func("add", vec!["T"]); - generic_func.constraints = vec![InterfaceConstraint { - interface_name: Name::str("Addable"), - typevars: vec![Name::str("T")], - }]; - - let result = pass.instantiate_function( - Name::str("add"), - vec![mk_type(Type::Int32)], - &generic_func, - &decls, + assert!( + targets(output.find_entry_point(Name::str("use$i32")).unwrap()) + .contains(&integer_global) ); - - assert!(result.is_ok()); - - // Check that constraints are preserved (they're on the original, not the specialized) - let specialized = &pass.out_decls[0]; - if let Decl::Func(fdecl) = specialized { - // Specialized version should have constraints copied - assert_eq!(fdecl.constraints.len(), 1); - assert_eq!(fdecl.constraints[0].interface_name, Name::str("Addable")); - } } #[test] - fn test_instantiation_deduplication() { - let mut pass = MonomorphPass::new(); - let decls = DeclTable::new(vec![]); - - let generic_func = mk_simple_func("id", vec!["T"]); - let type_args = vec![mk_type(Type::Int32)]; - - // First instantiation - let key1 = MonomorphKey::new(Name::str("id"), type_args.clone()); - pass.instantiate_function(Name::str("id"), type_args.clone(), &generic_func, &decls) - .unwrap(); - - assert_eq!(pass.out_decls.len(), 1); - - // Same instantiation again - should not create duplicate - pass.instantiate_function(Name::str("id"), type_args.clone(), &generic_func, &decls) + fn local_function_references_never_reenter_overload_resolution() { + let output = specialize("bump(x: i32) -> i32 { x } bump(x: f32) -> f32 { x } inc(x: i32) -> i32 { x + 1 } main { let bump = inc; bump(0) }"); + let main = output.find_entry_point(Name::str("main")).unwrap(); + let binding = main + .arena + .locals + .iter() + .position(|local| local.name == Name::str("bump")) .unwrap(); - - assert_eq!(pass.out_decls.len(), 1); - assert!(pass.instantiations.contains_key(&key1)); - } - - #[test] - fn test_expr_traversal_coverage() { - let mut pass = MonomorphPass::new(); - let decls = DeclTable::new(vec![]); - - // Test If expression - let mut arena = ExprArena::new(); - let cond = arena.add(Expr::True, test_loc()); - let then_expr = arena.add(Expr::Int(1, None), test_loc()); - let else_expr = arena.add(Expr::Int(2, None), test_loc()); - let if_expr = arena.add(Expr::If(cond, then_expr, Some(else_expr)), test_loc()); - - let mut fdecl = FuncDecl { - name: Name::str("test"), - typevars: Vec::new(), - size_vars: Vec::new(), - params: Vec::new(), - constraints: Vec::new(), - requires: vec![], - ret: mk_type(Type::Void), - body: Some(if_expr), - arena, - types: vec![ - mk_type(Type::Bool), - mk_type(Type::Int32), - mk_type(Type::Int32), - mk_type(Type::Int32), - ], - loc: test_loc(), - closure_vars: vec![], - is_extern: false, - }; - - let result = pass.process_expr(if_expr, &mut fdecl, &decls); - assert!(result.is_ok()); + assert!(main + .arena + .ids() + .any(|id| main.arena.reference(id) == Some(&Reference::Local(LocalId(binding as u32))))); + assert!(output.find(Name::str("bump$i32")).is_empty()); + assert!(output.find(Name::str("bump$f32")).is_empty()); } #[test] - fn test_monomorphize_single_entry_point() { - let mut pass = MonomorphPass::new(); - - // Create a simple non-generic entry point function - // main() { 42 } - let mut arena = ExprArena::new(); - let body_expr = arena.add(Expr::Int(42, None), test_loc()); - - let entry_func = FuncDecl { - name: Name::str("main"), - typevars: Vec::new(), - size_vars: Vec::new(), - params: Vec::new(), - constraints: Vec::new(), - requires: vec![], - ret: mk_type(Type::Int32), - body: Some(body_expr), - arena, - types: vec![mk_type(Type::Int32)], - loc: test_loc(), - closure_vars: vec![], - is_extern: false, - }; - - let decls = DeclTable::new(vec![Decl::Func(entry_func)]); - - // Monomorphize starting from "main" - let result = pass.monomorphize(&decls, Name::str("main")); - - assert!(result.is_ok()); - let all_decls = result.unwrap(); - - // Should have 1 decl (just main, no specializations) - assert_eq!(all_decls.len(), 1); + fn size_substitution_respects_local_shadowing() { + let output = specialize("probe(a: [i32; N]) { let size = N; if true { let N = 99; let local = N }; let again = N } main { probe([1,2,3]) }"); + let function = output.find_entry_point(Name::str("probe$3")).unwrap(); + assert_eq!( + function + .arena + .exprs() + .iter() + .filter(|expr| **expr == Expr::Int(3, None)) + .count(), + 2 + ); + assert!(function.arena.ids().any(|id| matches!(function.arena.reference(id), + Some(Reference::Local(local)) if function.arena.local(*local).name == Name::str("N")))); + assert!(!function + .arena + .ids() + .any(|id| matches!(function.arena.reference(id), Some(Reference::SizeParameter(_))))); } #[test] - fn test_monomorphize_entry_point_calling_generic() { - let mut pass = MonomorphPass::new(); - - // Create a generic function id(x: T) -> T { x } - let t_var = typevar("T"); - let mut id_arena = ExprArena::new(); - let id_param_expr = id_arena.add(Expr::Id(Name::str("x")), test_loc()); - - let id_func = FuncDecl { - name: Name::str("id"), - typevars: vec![Name::str("T")], - size_vars: vec![], - params: vec![Param { - name: Name::str("x"), - ty: Some(t_var), - }], - constraints: Vec::new(), - requires: vec![], - ret: t_var, - body: Some(id_param_expr), - arena: id_arena, - types: vec![t_var], - loc: test_loc(), - closure_vars: vec![], - is_extern: false, - }; - - // Create entry point function that calls id(42) - // main() { id(42) } - let mut main_arena = ExprArena::new(); - let arg_expr = main_arena.add(Expr::Int(42, None), test_loc()); - let fn_expr = main_arena.add(Expr::Id(Name::str("id")), test_loc()); - let call_expr = main_arena.add(Expr::Call(fn_expr, vec![arg_expr]), test_loc()); - - let i32_type = mk_type(Type::Int32); - let func_type = mk_type(Type::Func(tuple(vec![i32_type]), i32_type)); - - let main_func = FuncDecl { - name: Name::str("main"), - typevars: Vec::new(), - size_vars: Vec::new(), - params: Vec::new(), - constraints: Vec::new(), - requires: vec![], - ret: i32_type, - body: Some(call_expr), - arena: main_arena, - types: vec![i32_type, func_type, i32_type], - loc: test_loc(), - closure_vars: vec![], - is_extern: false, - }; - - let decls = DeclTable::new(vec![Decl::Func(id_func), Decl::Func(main_func)]); - - // Monomorphize starting from "main" - let result = pass.monomorphize(&decls, Name::str("main")); - - assert!(result.is_ok()); - let all_decls = result.unwrap(); - - // Should have 2 decls: main, id$i32 (specialized) - assert_eq!(all_decls.len(), 2); - - // Find the specialized version - let specialized_id = all_decls.iter().find(|d| { - if let Decl::Func(f) = d { - f.name.to_string().starts_with("id$") - } else { - false - } - }); - - assert!(specialized_id.is_some()); - if let Some(Decl::Func(fdecl)) = specialized_id { - // Should be specialized (no type variables) - assert_eq!(fdecl.typevars.len(), 0); - // Return type should be i32 - assert_eq!(*fdecl.ret, Type::Int32); - // Name should start with "id$" - assert!(fdecl.name.to_string().starts_with("id$")); - } else { - panic!("Expected function declaration"); - } - - assert_eq!(all_decls[0].pretty_print(), "id$i32(x: i32) → i32 x"); - assert_eq!(all_decls[1].pretty_print(), "main() → i32 id$i32(42)"); + fn ordinary_generic_overload_diagnostics_are_preserved() { + let source = checked( + "choose(x: T) -> i32 { 1 } choose(x: [T; 2]) -> i32 { 2 } main { choose(1) }", + ); + let error = MonomorphPass::new() + .monomorphize(&source, Name::str("main")) + .unwrap_err(); + assert!( + error.contains("Cannot infer type arguments for choose"), + "{}", + error + ); } } diff --git a/src/parser.rs b/src/parser.rs index 878725bb..ba8dff30 100644 --- a/src/parser.rs +++ b/src/parser.rs @@ -1046,8 +1046,7 @@ fn parse_func_decl(name: Name, cx: &mut ParseContext) -> FuncDecl { requires, loc, arena, - types: vec![], - closure_vars: vec![], + is_extern: false, } } diff --git a/src/safety_checker.rs b/src/safety_checker.rs index e38f9251..59258af6 100644 --- a/src/safety_checker.rs +++ b/src/safety_checker.rs @@ -1,6 +1,132 @@ +use crate::checked::{CheckedBody as ExprArena, CheckedFunction as FuncDecl}; use crate::interval::{enclose, IndexInterval}; use crate::*; +/// The safety analysis sees the phase's authoritative call references. Type +/// layout still comes from the phase's declaration inventory. +pub trait SafetyProgram: std::ops::Deref { + fn call_target<'a>( + &'a self, + reference: &Reference, + body: &CheckedBody, + signature: TypeID, + arity: usize, + ) -> Option<&'a CheckedFunction>; +} +impl SafetyProgram for CheckedProgram { + fn call_target<'a>( + &'a self, + reference: &Reference, + body: &CheckedBody, + signature: TypeID, + arity: usize, + ) -> Option<&'a CheckedFunction> { + let ids: &[DefId] = match reference { + Reference::Functions(ids) => ids, + Reference::InterfaceMember { + requirement, + member, + } => { + let requirement = body.requirements.get(requirement.index())?; + let member = requirement + .members + .iter() + .find(|candidate| candidate.definition == *member)?; + &member.candidates + } + _ => return None, + }; + // Templates retain the existing first exact candidate policy. This is + // source diagnostic coverage, not concrete overload selection. + ids.iter() + .filter_map(|id| self.function(*id)) + .find(|function| function.params.len() == arity && function.ty() == signature) + } +} +impl SafetyProgram for SpecializedProgram { + fn call_target<'a>( + &'a self, + reference: &Reference, + _body: &CheckedBody, + _signature: TypeID, + arity: usize, + ) -> Option<&'a CheckedFunction> { + let Reference::Instance(id) = reference else { + return None; + }; + // The selected function is authoritative. Checking and specialization + // establish coercion-aware compatibility; structural validation checks + // arity before safety runs. A function-valued global is still indirect. + self.function_instance(*id) + .filter(|function| function.params.len() == arity) + } +} + +/// Expression analysis needs its owning body and symbolic size binders only. +/// Function signatures and contracts stay at the function/call boundary. +#[derive(Clone, Copy)] +struct SafetyBody<'a> { + arena: &'a CheckedBody, + size_vars: &'a [SizeParameter], +} + +impl<'a> From<&'a FuncDecl> for SafetyBody<'a> { + fn from(function: &'a FuncDecl) -> Self { + Self { + arena: &function.arena, + size_vars: &function.size_vars, + } + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)] +enum PlaceRoot { + Local(LocalId), + Global(DefId), + Instance(InstanceId), +} + +/// Field paths are structural projections of an identified storage root. +/// Diagnostic binding spellings never participate in equality. +#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)] +struct Place { + root: PlaceRoot, + fields: Option, +} +impl Place { + fn local(local: LocalId) -> Self { + Self { + root: PlaceRoot::Local(local), + fields: None, + } + } + fn field(self, field: Name) -> Self { + let fields = match self.fields { + Some(path) => Name::new(format!("{}.{}", path, field)), + None => field, + }; + Self { + fields: Some(fields), + ..self + } + } +} +fn reference_place(reference: &Reference) -> Option { + let root = match reference { + Reference::Local(local) | Reference::SizeParameter(local) => PlaceRoot::Local(*local), + Reference::Global(definition) => PlaceRoot::Global(*definition), + Reference::Instance(instance) => PlaceRoot::Instance(*instance), + _ => return None, + }; + Some(Place { root, fields: None }) +} +fn id_place(id: ExprID, arena: &ExprArena) -> Option { + match &arena[id] { + Expr::Id(_) => arena.reference(id).and_then(reference_place), + _ => None, + } +} + /// Lanes in an `f32x4`. A lane index has to be provably in `0..4` for the same /// reason an array index has to be in range: no backend checks it at runtime. const F32X4_LANES: i64 = 4; @@ -32,18 +158,11 @@ fn collect_size_subst(param_ty: TypeID, arg_ty: TypeID, out: &mut Vec<(Name, i64 } } -/// Extract a trackable name from an expression for constraint tracking. -/// Returns the variable name for `Expr::Id(name)`, or a synthetic compound -/// name for `Expr::Field(base, field)` (e.g., `h.index` becomes a single -/// interned name). This lets the safety checker track bounds on struct fields. -fn expr_constraint_name(id: ExprID, arena: &ExprArena) -> Option { +/// A trackable storage place, including direct field projections. +fn expr_place(id: ExprID, arena: &ExprArena) -> Option { match &arena[id] { - Expr::Id(name) => Some(*name), - Expr::Field(base, field) => { - let base_name = expr_constraint_name(*base, arena)?; - // Intern a compound name "base.field" so it works as a constraint key. - Some(Name::new(format!("{}.{}", *base_name, **field).into())) - } + Expr::Id(_) => arena.reference(id).and_then(reference_place), + Expr::Field(base, field) => Some(expr_place(*base, arena)?.field(*field)), _ => None, } } @@ -56,25 +175,25 @@ fn expr_constraint_name(id: ExprID, arena: &ExprArena) -> Option { /// the two by element type alone, so the length is not part of the match. /// The base of an index or field expression is never itself the argument, so /// walking down from it recovers the array type the coercion hid. -fn array_type(expr: ExprID, decl: &FuncDecl, decls: &DeclTable) -> Option { - if expr < decl.types.len() { - let ty = decl.types[expr]; +fn array_type(expr: ExprID, context: SafetyBody<'_>, decls: &impl SafetyProgram) -> Option { + if expr < context.arena.len() { + let ty = context.arena.ty(expr); if let Type::Array(_, ArraySize::Known(_)) = *ty { return Some(ty); } } - match &decl.arena[expr] { + match &context.arena[expr] { // An element of `[[T; N]; M]` is a `[T; N]`; anything else has no // length to recover. - Expr::ArrayIndex(base, _) => match &*array_type(*base, decl, decls)? { + Expr::ArrayIndex(base, _) => match &*array_type(*base, context, decls)? { Type::Array(elem, _) if matches!(**elem, Type::Array(_, ArraySize::Known(_))) => { Some(*elem) } _ => None, }, Expr::Field(base, field) => { - let base_ty = if *base < decl.types.len() { - decl.types[*base] + let base_ty = if *base < context.arena.len() { + context.arena.ty(*base) } else { return None; }; @@ -95,8 +214,8 @@ fn array_type(expr: ExprID, decl: &FuncDecl, decls: &DeclTable) -> Option Option { - match &*array_type(expr, decl, decls)? { +fn static_len(expr: ExprID, context: SafetyBody<'_>, decls: &impl SafetyProgram) -> Option { + match &*array_type(expr, context, decls)? { Type::Array(_, ArraySize::Known(n)) => Some(*n as i64), _ => None, } @@ -110,7 +229,7 @@ pub struct SafetyError { #[derive(Clone, Debug)] struct IndexConstraint { - pub name: Name, + pub name: Place, pub min: Option, pub max: Option, pub non_zero: bool, @@ -119,29 +238,22 @@ struct IndexConstraint { /// Records that variable `index` has been proven < `array.len`. #[derive(Clone, Debug)] struct LenBound { - pub index: Name, - pub array: Name, + pub index: Place, + pub array: Place, } /// Records that `array.len >= min_len` (the array has at least `min_len` elements). #[derive(Clone, Debug)] struct MinLenBound { - pub array: Name, + pub array: Place, pub min_len: i64, } /// Records that variable `lo` is proven < variable `hi`. #[derive(Clone, Debug)] struct VarBound { - pub lo: Name, - pub hi: Name, -} - -/// Local variable declaration. -#[derive(Copy, Clone, Debug)] -struct Var { - name: Name, - ty: TypeID, + pub lo: Place, + pub hi: Place, } /// Static safety checker using abstract interpretation. @@ -179,9 +291,6 @@ struct Var { /// body, so the body has to guard it too. A directly-called lambda is /// exempt: its definition site is its call site. pub struct SafetyChecker { - /// Currently declared vars, as we're checking. - vars: Vec, - /// Constraints we know about each var. constraints: Vec, @@ -201,7 +310,12 @@ pub struct SafetyChecker { /// checked. A lambda body can run at any point after its definition, so /// anything in here is unconstrained inside a lambda that isn't called /// immediately. - fn_assigned: Vec, + fn_assigned: Vec, + + /// Whole-body specialization preserves source locations. The same source + /// call/requirement can fail in several instances whose emitted names + /// differ. Keep distinct concretized clauses (e.g. different array sizes). + failed_requirements: std::collections::HashMap<(Loc, Loc, String), String>, pub errors: Vec, } @@ -209,13 +323,13 @@ pub struct SafetyChecker { impl SafetyChecker { pub fn new() -> Self { Self { - vars: vec![], constraints: vec![], len_bounds: vec![], leq_len_bounds: vec![], min_len_bounds: vec![], var_bounds: vec![], fn_assigned: vec![], + failed_requirements: std::collections::HashMap::new(), errors: vec![], } } @@ -236,7 +350,7 @@ impl SafetyChecker { self.errors.push(err); } - fn add(&mut self, name: Name, min: Option, max: Option) { + fn add(&mut self, name: Place, min: Option, max: Option) { self.constraints.push(IndexConstraint { name, min, @@ -245,12 +359,12 @@ impl SafetyChecker { }) } - fn replace(&mut self, name: Name, min: Option, max: Option) { + fn replace(&mut self, name: Place, min: Option, max: Option) { self.constraints.retain(|c| c.name != name); self.add(name, min, max); } - fn add_non_zero(&mut self, name: Name) { + fn add_non_zero(&mut self, name: Place) { // If there's already a constraint, mark it non_zero. // Otherwise, add an unconstrained entry with non_zero set. if let Some(c) = self.constraints.iter_mut().find(|c| c.name == name) { @@ -265,9 +379,8 @@ impl SafetyChecker { } } - /// Drop everything we know about `name`. Used when a lambda parameter - /// shadows a captured variable of the same name. - fn forget(&mut self, name: Name) { + /// Drop facts about a storage place before a new dynamic value is bound. + fn forget(&mut self, name: Place) { self.constraints.retain(|c| c.name != name); self.len_bounds .retain(|b| b.index != name && b.array != name); @@ -277,35 +390,35 @@ impl SafetyChecker { self.var_bounds.retain(|b| b.lo != name && b.hi != name); } - fn find(&self, name: Name) -> Option { + fn find(&self, name: Place) -> Option { self.constraints.iter().find(|c| c.name == name).cloned() } /// Given expr evaluates to true, add constraints accordingly. - fn match_expr(&mut self, expr: ExprID, decl: &FuncDecl, decls: &DeclTable) { + fn match_expr(&mut self, expr: ExprID, context: SafetyBody<'_>, decls: &impl SafetyProgram) { // Track bounds from comparisons. Handles both simple variables (Expr::Id) - // and struct field access (Expr::Field) via expr_constraint_name. - if let Expr::Binop(Binop::Less, lhs, rhs) = &decl.arena[expr] { - if let Some(name) = expr_constraint_name(*lhs, &decl.arena) { - let ival = self.check_expr(*rhs, decl, decls); + // and struct field access (Expr::Field) via expr_place. + if let Expr::Binop(Binop::Less, lhs, rhs) = &context.arena[expr] { + if let Some(name) = expr_place(*lhs, context.arena) { + let ival = self.check_expr(*rhs, context, decls); if ival.max != i64::max_value() { self.add(name, None, Some(ival.max - 1)); } } } - if let Expr::Binop(Binop::Leq, lhs, rhs) = &decl.arena[expr] { - if let Some(name) = expr_constraint_name(*lhs, &decl.arena) { - let ival = self.check_expr(*rhs, decl, decls); + if let Expr::Binop(Binop::Leq, lhs, rhs) = &context.arena[expr] { + if let Some(name) = expr_place(*lhs, context.arena) { + let ival = self.check_expr(*rhs, context, decls); if ival.max != i64::max_value() { self.add(name, None, Some(ival.max)); } } } - if let Expr::Binop(Binop::Geq, lhs, rhs) = &decl.arena[expr] { - if let Some(name) = expr_constraint_name(*lhs, &decl.arena) { - let ival = self.check_expr(*rhs, decl, decls); + if let Expr::Binop(Binop::Geq, lhs, rhs) = &context.arena[expr] { + if let Some(name) = expr_place(*lhs, context.arena) { + let ival = self.check_expr(*rhs, context, decls); if ival.min != i64::MIN { self.add(name, Some(ival.min), None); } @@ -313,11 +426,11 @@ impl SafetyChecker { } // match `array.len >= N` — record min length bound - if let Expr::Binop(Binop::Geq, lhs, rhs) = &decl.arena[expr] { - if let Expr::Field(arr_expr, field_name) = &decl.arena[*lhs] { + if let Expr::Binop(Binop::Geq, lhs, rhs) = &context.arena[expr] { + if let Expr::Field(arr_expr, field_name) = &context.arena[*lhs] { if field_name.as_str() == "len" { - if let Expr::Id(array_name) = &decl.arena[*arr_expr] { - let ival = self.check_expr(*rhs, decl, decls); + if let Some(ref array_name) = id_place(*arr_expr, context.arena) { + let ival = self.check_expr(*rhs, context, decls); if ival.min != i64::MAX { self.min_len_bounds.push(MinLenBound { array: *array_name, @@ -330,18 +443,18 @@ impl SafetyChecker { } // match `n <= array.len` — record leq length bound and min length bound - if let Expr::Binop(Binop::Leq, lhs, rhs) = &decl.arena[expr] { - if let Expr::Field(arr_expr, field_name) = &decl.arena[*rhs] { + if let Expr::Binop(Binop::Leq, lhs, rhs) = &context.arena[expr] { + if let Expr::Field(arr_expr, field_name) = &context.arena[*rhs] { if field_name.as_str() == "len" { - if let Expr::Id(array_name) = &decl.arena[*arr_expr] { + if let Some(ref array_name) = id_place(*arr_expr, context.arena) { // Record n <= array.len for transitive propagation - if let Expr::Id(name) = &decl.arena[*lhs] { + if let Some(ref name) = id_place(*lhs, context.arena) { self.leq_len_bounds.push(LenBound { index: *name, array: *array_name, }); } - let ival = self.check_expr(*lhs, decl, decls); + let ival = self.check_expr(*lhs, context, decls); if ival.min != i64::MAX { self.min_len_bounds.push(MinLenBound { array: *array_name, @@ -354,11 +467,11 @@ impl SafetyChecker { } // match `array.len >= n` — record leq length bound (same as n <= array.len) - if let Expr::Binop(Binop::Geq, lhs, rhs) = &decl.arena[expr] { - if let Expr::Field(arr_expr, field_name) = &decl.arena[*lhs] { + if let Expr::Binop(Binop::Geq, lhs, rhs) = &context.arena[expr] { + if let Expr::Field(arr_expr, field_name) = &context.arena[*lhs] { if field_name.as_str() == "len" { - if let Expr::Id(array_name) = &decl.arena[*arr_expr] { - if let Expr::Id(name) = &decl.arena[*rhs] { + if let Some(ref array_name) = id_place(*arr_expr, context.arena) { + if let Some(ref name) = id_place(*rhs, context.arena) { self.leq_len_bounds.push(LenBound { index: *name, array: *array_name, @@ -371,9 +484,9 @@ impl SafetyChecker { // match expressions of the form i < id, where id is another variable // with a constraint - if let Expr::Binop(Binop::Less, lhs, rhs) = &decl.arena[expr] { - if let Expr::Id(name) = &decl.arena[*lhs] { - if let Expr::Id(max_name) = &decl.arena[*rhs] { + if let Expr::Binop(Binop::Less, lhs, rhs) = &context.arena[expr] { + if let Some(ref name) = id_place(*lhs, context.arena) { + if let Some(ref max_name) = id_place(*rhs, context.arena) { if let Some(c) = self.find(*max_name) { if let Some(max) = c.max { self.add(*name, None, Some(max)); @@ -387,9 +500,9 @@ impl SafetyChecker { }); } // match i < array.len — record symbolic length bound - if let Expr::Field(arr_expr, field_name) = &decl.arena[*rhs] { + if let Expr::Field(arr_expr, field_name) = &context.arena[*rhs] { if field_name.as_str() == "len" { - if let Expr::Id(array_name) = &decl.arena[*arr_expr] { + if let Some(ref array_name) = id_place(*arr_expr, context.arena) { self.len_bounds.push(LenBound { index: *name, array: *array_name, @@ -401,39 +514,39 @@ impl SafetyChecker { } // match `x != 0` — mark x as non-zero - if let Expr::Binop(Binop::NotEqual, lhs, rhs) = &decl.arena[expr] { - if let Some(name) = expr_constraint_name(*lhs, &decl.arena) { - if let Expr::Int(0, _) = &decl.arena[*rhs] { + if let Expr::Binop(Binop::NotEqual, lhs, rhs) = &context.arena[expr] { + if let Some(name) = expr_place(*lhs, context.arena) { + if let Expr::Int(0, _) = &context.arena[*rhs] { self.add_non_zero(name); } } - if let Some(name) = expr_constraint_name(*rhs, &decl.arena) { - if let Expr::Int(0, _) = &decl.arena[*lhs] { + if let Some(name) = expr_place(*rhs, context.arena) { + if let Expr::Int(0, _) = &context.arena[*lhs] { self.add_non_zero(name); } } } // match `x > n` — x.min = n + 1 - if let Expr::Binop(Binop::Greater, lhs, rhs) = &decl.arena[expr] { - if let Some(name) = expr_constraint_name(*lhs, &decl.arena) { - let ival = self.check_expr(*rhs, decl, decls); + if let Expr::Binop(Binop::Greater, lhs, rhs) = &context.arena[expr] { + if let Some(name) = expr_place(*lhs, context.arena) { + let ival = self.check_expr(*rhs, context, decls); if ival.min != i64::MAX { self.add(name, Some(ival.min + 1), None); } } // reversed: `n > i` means i < n - if let Some(name) = expr_constraint_name(*rhs, &decl.arena) { - let ival = self.check_expr(*lhs, decl, decls); + if let Some(name) = expr_place(*rhs, context.arena) { + let ival = self.check_expr(*lhs, context, decls); if ival.max != i64::MAX { self.add(name, None, Some(ival.max - 1)); } } // match `array.len > N` — record min length bound - if let Expr::Field(arr_expr, field_name) = &decl.arena[*lhs] { + if let Expr::Field(arr_expr, field_name) = &context.arena[*lhs] { if field_name.as_str() == "len" { - if let Expr::Id(array_name) = &decl.arena[*arr_expr] { - let ival = self.check_expr(*rhs, decl, decls); + if let Some(ref array_name) = id_place(*arr_expr, context.arena) { + let ival = self.check_expr(*rhs, context, decls); if ival.min != i64::MAX { self.min_len_bounds.push(MinLenBound { array: *array_name, @@ -446,11 +559,11 @@ impl SafetyChecker { } // match `N < array.len` — record min length bound - if let Expr::Binop(Binop::Less, lhs, rhs) = &decl.arena[expr] { - if let Expr::Field(arr_expr, field_name) = &decl.arena[*rhs] { + if let Expr::Binop(Binop::Less, lhs, rhs) = &context.arena[expr] { + if let Expr::Field(arr_expr, field_name) = &context.arena[*rhs] { if field_name.as_str() == "len" { - if let Expr::Id(array_name) = &decl.arena[*arr_expr] { - let ival = self.check_expr(*lhs, decl, decls); + if let Some(ref array_name) = id_place(*arr_expr, context.arena) { + let ival = self.check_expr(*lhs, context, decls); if ival.min != i64::MAX { self.min_len_bounds.push(MinLenBound { array: *array_name, @@ -463,9 +576,9 @@ impl SafetyChecker { } // reversed: `n < i` means i > n - if let Expr::Binop(Binop::Less, lhs, rhs) = &decl.arena[expr] { - if let Some(name) = expr_constraint_name(*rhs, &decl.arena) { - let ival = self.check_expr(*lhs, decl, decls); + if let Expr::Binop(Binop::Less, lhs, rhs) = &context.arena[expr] { + if let Some(name) = expr_place(*rhs, context.arena) { + let ival = self.check_expr(*lhs, context, decls); if ival.min != i64::MIN { self.add(name, Some(ival.min + 1), None); } @@ -473,9 +586,9 @@ impl SafetyChecker { } // reversed: `n >= i` means i <= n - if let Expr::Binop(Binop::Geq, lhs, rhs) = &decl.arena[expr] { - if let Some(name) = expr_constraint_name(*rhs, &decl.arena) { - let ival = self.check_expr(*lhs, decl, decls); + if let Expr::Binop(Binop::Geq, lhs, rhs) = &context.arena[expr] { + if let Some(name) = expr_place(*rhs, context.arena) { + let ival = self.check_expr(*lhs, context, decls); if ival.max != i64::MAX { self.add(name, None, Some(ival.max)); } @@ -483,9 +596,9 @@ impl SafetyChecker { } // reversed: `n <= i` means i >= n - if let Expr::Binop(Binop::Leq, lhs, rhs) = &decl.arena[expr] { - if let Some(name) = expr_constraint_name(*rhs, &decl.arena) { - let ival = self.check_expr(*lhs, decl, decls); + if let Expr::Binop(Binop::Leq, lhs, rhs) = &context.arena[expr] { + if let Some(name) = expr_place(*rhs, context.arena) { + let ival = self.check_expr(*lhs, context, decls); if ival.min != i64::MIN { self.add(name, Some(ival.min), None); } @@ -493,22 +606,22 @@ impl SafetyChecker { } // match `a == b` — treat as both `a <= b` and `a >= b` - if let Expr::Binop(Binop::Equal, lhs, rhs) = &decl.arena[expr] { + if let Expr::Binop(Binop::Equal, lhs, rhs) = &context.arena[expr] { // Constrain lhs from rhs value - if let Some(name) = expr_constraint_name(*lhs, &decl.arena) { - let ival = self.check_expr(*rhs, decl, decls); + if let Some(name) = expr_place(*lhs, context.arena) { + let ival = self.check_expr(*rhs, context, decls); self.add(name, Some(ival.min), Some(ival.max)); } // Constrain rhs from lhs value - if let Some(name) = expr_constraint_name(*rhs, &decl.arena) { - let ival = self.check_expr(*lhs, decl, decls); + if let Some(name) = expr_place(*rhs, context.arena) { + let ival = self.check_expr(*lhs, context, decls); self.add(name, Some(ival.min), Some(ival.max)); } // array.len == N → min_len_bound - if let Expr::Field(arr_expr, field_name) = &decl.arena[*lhs] { + if let Expr::Field(arr_expr, field_name) = &context.arena[*lhs] { if field_name.as_str() == "len" { - if let Expr::Id(array_name) = &decl.arena[*arr_expr] { - let ival = self.check_expr(*rhs, decl, decls); + if let Some(ref array_name) = id_place(*arr_expr, context.arena) { + let ival = self.check_expr(*rhs, context, decls); if ival.min != i64::MAX { self.min_len_bounds.push(MinLenBound { array: *array_name, @@ -516,7 +629,7 @@ impl SafetyChecker { }); } // Also record leq bound: n <= array.len - if let Expr::Id(name) = &decl.arena[*rhs] { + if let Some(ref name) = id_place(*rhs, context.arena) { self.leq_len_bounds.push(LenBound { index: *name, array: *array_name, @@ -526,10 +639,10 @@ impl SafetyChecker { } } // N == array.len → min_len_bound - if let Expr::Field(arr_expr, field_name) = &decl.arena[*rhs] { + if let Expr::Field(arr_expr, field_name) = &context.arena[*rhs] { if field_name.as_str() == "len" { - if let Expr::Id(array_name) = &decl.arena[*arr_expr] { - let ival = self.check_expr(*lhs, decl, decls); + if let Some(ref array_name) = id_place(*arr_expr, context.arena) { + let ival = self.check_expr(*lhs, context, decls); if ival.min != i64::MAX { self.min_len_bounds.push(MinLenBound { array: *array_name, @@ -537,7 +650,7 @@ impl SafetyChecker { }); } // Also record leq bound: n <= array.len - if let Expr::Id(name) = &decl.arena[*lhs] { + if let Some(ref name) = id_place(*lhs, context.arena) { self.leq_len_bounds.push(LenBound { index: *name, array: *array_name, @@ -548,9 +661,9 @@ impl SafetyChecker { } } - if let Expr::Binop(Binop::And, lhs, rhs) = &decl.arena[expr] { - self.match_expr(*lhs, decl, decls); - self.match_expr(*rhs, decl, decls); + if let Expr::Binop(Binop::And, lhs, rhs) = &context.arena[expr] { + self.match_expr(*lhs, context, decls); + self.match_expr(*rhs, context, decls); self.propagate_len_bounds(); } } @@ -605,27 +718,29 @@ impl SafetyChecker { } } - fn check_expr(&mut self, expr: ExprID, decl: &FuncDecl, decls: &DeclTable) -> IndexInterval { - match &decl.arena[expr] { + fn check_expr( + &mut self, + expr: ExprID, + context: SafetyBody<'_>, + decls: &impl SafetyProgram, + ) -> IndexInterval { + match &context.arena[expr] { Expr::Int(x, _) => IndexInterval { min: *x, max: *x, non_zero: *x != 0, }, Expr::Block(exprs) => { - let n = self.vars.len(); for e in exprs { - self.check_expr(*e, decl, decls); - } - while self.vars.len() > n { - self.vars.pop(); + self.check_expr(*e, context, decls); } IndexInterval::default() } - Expr::Let(name, init, _) => { - let init_r = self.check_expr(*init, decl, decls); - let ty = decl.types[expr]; - self.vars.push(Var { name: *name, ty }); + Expr::Let(_, init, _) => { + let local = context.arena.binder(expr); + let name = &Place::local(local); + let init_r = self.check_expr(*init, context, decls); + let ty = context.arena.local(local).ty; // Track the interval from the initializer. let mut min = if init_r.min != i64::MIN { @@ -647,7 +762,7 @@ impl SafetyChecker { } // Propagate LenBounds: let x = y inherits y's LenBounds. - if let Expr::Id(src_name) = &decl.arena[*init] { + if let Some(ref src_name) = id_place(*init, context.arena) { let inherited: Vec<_> = self .len_bounds .iter() @@ -664,13 +779,15 @@ impl SafetyChecker { IndexInterval::default() } - Expr::Var(name, init, _) => { + Expr::Var(_, init, _) => { + let local = context.arena.binder(expr); + let name = &Place::local(local); let init_r = if let Some(init) = init { - self.check_expr(*init, decl, decls) + self.check_expr(*init, context, decls) } else { IndexInterval::default() }; - let ty = decl.types[expr]; + let ty = context.arena.local(local).ty; let mut min = if init_r.min != i64::MIN { Some(init_r.min) @@ -692,7 +809,7 @@ impl SafetyChecker { // Propagate LenBounds: var x = y inherits y's LenBounds. if let Some(init) = init { - if let Expr::Id(src_name) = &decl.arena[*init] { + if let Some(ref src_name) = id_place(*init, context.arena) { let inherited: Vec<_> = self .len_bounds .iter() @@ -710,7 +827,11 @@ impl SafetyChecker { IndexInterval::default() } - Expr::Id(name) => { + Expr::Id(_) => { + let Some(place) = context.arena.reference(expr).and_then(reference_place) else { + return IndexInterval::default(); + }; + let name = &place; let mut min = i64::min_value(); let mut max = i64::max_value(); let mut non_zero = false; @@ -735,10 +856,10 @@ impl SafetyChecker { let initial_min_len_bound_count = self.min_len_bounds.len(); let initial_var_bound_count = self.var_bounds.len(); - self.match_expr(*cond, decl, decls); + self.match_expr(*cond, context, decls); self.propagate_len_bounds(); - let mut r = self.check_expr(*then_expr, decl, decls); + let mut r = self.check_expr(*then_expr, context, decls); // Pop condition constraints before checking else branch — // the else branch executes when the condition is false, @@ -751,28 +872,28 @@ impl SafetyChecker { self.var_bounds.truncate(initial_var_bound_count); if let Some(else_expr) = else_expr { - let else_r = self.check_expr(*else_expr, decl, decls); + let else_r = self.check_expr(*else_expr, context, decls); r = enclose(r, else_r); } r } Expr::ArrayIndex(array_expr, index_expr) => { - if *array_expr >= decl.types.len() { + if *array_expr >= context.arena.len() { print_error_with_context( - decl.arena.locs[expr], + context.arena.loc(expr), "internal compiler error: no type found for array index expression", ); return IndexInterval::default(); } - self.check_expr(*array_expr, decl, decls); - let lhs_ty = decl.types[*array_expr]; - let rhs_r = self.check_expr(*index_expr, decl, decls); + self.check_expr(*array_expr, context, decls); + let lhs_ty = context.arena.ty(*array_expr); + let rhs_r = self.check_expr(*index_expr, context, decls); if rhs_r.min < 0 { self.push_error(SafetyError { - location: decl.arena.locs[expr], + location: context.arena.loc(expr), message: format!("couldn't prove index is >= 0"), }); } @@ -783,16 +904,18 @@ impl SafetyChecker { // Also accept a `len_bound { idx, arr }` from a require // clause or `for`/`while` loop condition: this proves // `idx < arr.len`, and arr.len == n for a Known array. - let array_name = if let Expr::Id(name) = &decl.arena[*array_expr] { - Some(*name) - } else { - None - }; - let index_name = if let Expr::Id(name) = &decl.arena[*index_expr] { - Some(*name) - } else { - None - }; + let array_name = + if let Some(ref name) = id_place(*array_expr, context.arena) { + Some(*name) + } else { + None + }; + let index_name = + if let Some(ref name) = id_place(*index_expr, context.arena) { + Some(*name) + } else { + None + }; let len_bound_ok = match (index_name, array_name) { (Some(idx), Some(arr)) => self .len_bounds @@ -802,7 +925,7 @@ impl SafetyChecker { }; if !interval_ok && !len_bound_ok { self.push_error(SafetyError { - location: decl.arena.locs[expr], + location: context.arena.loc(expr), message: format!("couldn't prove index is less than array length"), }); } @@ -812,16 +935,18 @@ impl SafetyChecker { // or a require clause `idx < arr.len`, or // (b) a `var_bound` from `for i in 0 .. N` or a // require clause `idx < N`. - let array_name = if let Expr::Id(name) = &decl.arena[*array_expr] { - Some(*name) - } else { - None - }; - let index_name = if let Expr::Id(name) = &decl.arena[*index_expr] { - Some(*name) - } else { - None - }; + let array_name = + if let Some(ref name) = id_place(*array_expr, context.arena) { + Some(*name) + } else { + None + }; + let index_name = + if let Some(ref name) = id_place(*index_expr, context.arena) { + Some(*name) + } else { + None + }; let has_len_bound = match (index_name, array_name) { (Some(idx), Some(arr)) => self .len_bounds @@ -830,27 +955,31 @@ impl SafetyChecker { _ => false, }; let has_var_bound = if let Some(idx) = index_name { - self.var_bounds - .iter() - .any(|b| b.lo == idx && b.hi == *size_name) + self.var_bounds.iter().any(|b| { + b.lo == idx + && context.size_vars.iter().any(|parameter| { + parameter.symbol == *size_name + && b.hi == Place::local(parameter.local) + }) + }) } else { false }; if !has_len_bound && !has_var_bound { self.push_error(SafetyError { - location: decl.arena.locs[expr], + location: context.arena.loc(expr), message: format!("couldn't prove index is less than array length"), }); } } } else if let Type::Slice(_) = *lhs_ty { // For slices, check if the index has been proven < slice.len. - let array_name = if let Expr::Id(name) = &decl.arena[*array_expr] { + let array_name = if let Some(ref name) = id_place(*array_expr, context.arena) { Some(*name) } else { None }; - let index_name = if let Expr::Id(name) = &decl.arena[*index_expr] { + let index_name = if let Some(ref name) = id_place(*index_expr, context.arena) { Some(*name) } else { None @@ -874,7 +1003,7 @@ impl SafetyChecker { }; if !has_len_bound && !has_min_len_bound { self.push_error(SafetyError { - location: decl.arena.locs[expr], + location: context.arena.loc(expr), message: format!("couldn't prove index is less than slice length"), }); } @@ -883,7 +1012,7 @@ impl SafetyChecker { // index against, so the interval has to prove it on its own. if rhs_r.max >= F32X4_LANES { self.push_error(SafetyError { - location: decl.arena.locs[expr], + location: context.arena.loc(expr), message: format!("couldn't prove index is less than 4"), }); } @@ -897,9 +1026,9 @@ impl SafetyChecker { let saved_leq_len_bounds = self.leq_len_bounds.clone(); let saved_min_len_bounds = self.min_len_bounds.clone(); let saved_var_bounds = self.var_bounds.clone(); - self.match_expr(*cond, decl, decls); + self.match_expr(*cond, context, decls); - self.check_expr(*body, decl, decls); + self.check_expr(*body, context, decls); self.constraints = saved_constraints; self.len_bounds = saved_len_bounds; self.leq_len_bounds = saved_leq_len_bounds; @@ -909,41 +1038,41 @@ impl SafetyChecker { // Invalidate constraints for variables assigned inside the loop. // The restore gives us pre-loop state, but mutations in the body // mean those constraints may not hold at loop exit. - self.invalidate_assigned(*body, &decl.arena); + self.invalidate_assigned(*body, context.arena); IndexInterval::default() } Expr::Binop(op, lhs, rhs) => { if *op == Binop::Plus { - let lhs_range = self.check_expr(*lhs, decl, decls); - let rhs_range = self.check_expr(*rhs, decl, decls); + let lhs_range = self.check_expr(*lhs, context, decls); + let rhs_range = self.check_expr(*rhs, context, decls); return lhs_range + rhs_range; } if *op == Binop::Minus { - let lhs_range = self.check_expr(*lhs, decl, decls); - let rhs_range = self.check_expr(*rhs, decl, decls); + let lhs_range = self.check_expr(*lhs, context, decls); + let rhs_range = self.check_expr(*rhs, context, decls); return lhs_range - rhs_range; } if *op == Binop::Mult { - let lhs_range = self.check_expr(*lhs, decl, decls); - let rhs_range = self.check_expr(*rhs, decl, decls); + let lhs_range = self.check_expr(*lhs, context, decls); + let rhs_range = self.check_expr(*rhs, context, decls); return lhs_range * rhs_range; } if *op == Binop::Div || *op == Binop::Mod { - let lhs_range = self.check_expr(*lhs, decl, decls); - let rhs_range = self.check_expr(*rhs, decl, decls); + let lhs_range = self.check_expr(*lhs, context, decls); + let rhs_range = self.check_expr(*rhs, context, decls); // Only check integer division — float div-by-zero produces Inf/NaN per IEEE 754. - if *rhs < decl.types.len() { - let ty = decl.types[*rhs]; + if *rhs < context.arena.len() { + let ty = context.arena.ty(*rhs); let is_int = matches!(*ty, Type::Int32 | Type::UInt32 | Type::Int8 | Type::UInt8); if is_int && !rhs_range.excludes_zero() { self.push_error(SafetyError { - location: decl.arena.locs[expr], + location: context.arena.loc(expr), message: format!("couldn't prove divisor is non-zero"), }); } @@ -981,10 +1110,10 @@ impl SafetyChecker { } if *op == Binop::Assign { - self.check_expr(*lhs, decl, decls); - let rhs_range = self.check_expr(*rhs, decl, decls); + self.check_expr(*lhs, context, decls); + let rhs_range = self.check_expr(*rhs, context, decls); - if let Some(name) = expr_constraint_name(*lhs, &decl.arena) { + if let Some(name) = expr_place(*lhs, context.arena) { if rhs_range != IndexInterval::default() { self.replace(name, Some(rhs_range.min), Some(rhs_range.max)); } else { @@ -997,33 +1126,26 @@ impl SafetyChecker { // For other binops (==, !=, <, >, etc.), still recurse // into sub-expressions to check array accesses. - self.check_expr(*lhs, decl, decls); - self.check_expr(*rhs, decl, decls); + self.check_expr(*lhs, context, decls); + self.check_expr(*rhs, context, decls); IndexInterval::default() } Expr::Call(callee_expr, args) => { let arg_ivals: Vec<_> = args .iter() - .map(|arg| self.check_expr(*arg, decl, decls)) + .map(|arg| self.check_expr(*arg, context, decls)) .collect(); // An immediately-invoked lambda has known arguments, so check // its body against them rather than unconstrained. - if let Expr::Lambda { params, body } = &decl.arena[*callee_expr] { - let (params, body) = (params.clone(), *body); - self.check_lambda_body( - *callee_expr, - ¶ms, - body, - Some((args, &arg_ivals)), - decl, - decls, - ); + if let Expr::Lambda { body, .. } = &context.arena[*callee_expr] { + let params = context.arena.binders(*callee_expr); + self.check_lambda_body(params, *body, Some((args, &arg_ivals)), context, decls); } - self.check_call_requires(*callee_expr, args, expr, decl, decls); + self.check_call_requires(*callee_expr, args, expr, context, decls); IndexInterval::default() } Expr::Unop(op, expr) => { - let r = self.check_expr(*expr, decl, decls); + let r = self.check_expr(*expr, context, decls); match op { Unop::Neg => { // -[a, b] = [-b, -a] @@ -1039,13 +1161,13 @@ impl SafetyChecker { } } Expr::Return(expr) => { - self.check_expr(*expr, decl, decls); + self.check_expr(*expr, context, decls); IndexInterval::default() } Expr::Assume(cond) => { // Inject constraints from the condition without scoping — // they persist for the rest of the function. - self.match_expr(*cond, decl, decls); + self.match_expr(*cond, context, decls); self.propagate_len_bounds(); IndexInterval::default() } @@ -1056,7 +1178,7 @@ impl SafetyChecker { // of sized-array type, including an element of a nested array // such as `buffers[outer]`. if field.as_str() == "len" { - if let Some(n) = static_len(*base, decl, decls) { + if let Some(n) = static_len(*base, context, decls) { return IndexInterval { min: n, max: n, @@ -1065,8 +1187,8 @@ impl SafetyChecker { } } - // Look up constraints using the compound name (e.g., "h.index"). - if let Some(name) = expr_constraint_name(expr, &decl.arena) { + // Look up facts by the field's identified root and projection. + if let Some(name) = expr_place(expr, context.arena) { let mut min = i64::min_value(); let mut max = i64::max_value(); let mut non_zero = false; @@ -1089,37 +1211,25 @@ impl SafetyChecker { } } Expr::For { - var, - start, - end, - body, + start, end, body, .. } => { - let start_r = self.check_expr(*start, decl, decls); - let end_r = self.check_expr(*end, decl, decls); - - // Everything below binds the loop variable, which is scoped to - // the loop: snapshot first, so the restore after the body drops - // the loop variable's interval and bounds instead of keeping - // them alive — and brings back those of an outer variable of - // the same name, which the binding shadows. Without this, - // `var i = 100; for i in 0 .. 3 {}; a[i]` proved `i < 3` for - // the *outer* i and accepted an out-of-bounds write. + let var = &Place::local(context.arena.binder(expr)); + let start_r = self.check_expr(*start, context, decls); + let end_r = self.check_expr(*end, context, decls); + + // Keep the existing loop transfer rule: restore entry facts + // after visiting the body, then invalidate its assigned roots. + // The loop binding has its own LocalId throughout. let saved_constraints = self.constraints.clone(); let saved_len_bounds = self.len_bounds.clone(); let saved_leq_len_bounds = self.leq_len_bounds.clone(); let saved_min_len_bounds = self.min_len_bounds.clone(); let saved_var_bounds = self.var_bounds.clone(); - let saved_var_count = self.vars.len(); - - self.vars.push(Var { - name: *var, - ty: mk_type(Type::Int32), - }); self.add(*var, Some(start_r.min), Some(end_r.max.saturating_sub(1))); // for i in 0 .. arr.len — record that i < arr.len - if let Expr::Field(arr_expr, field_name) = &decl.arena[*end] { + if let Expr::Field(arr_expr, field_name) = &context.arena[*end] { if field_name.as_str() == "len" { - if let Expr::Id(array_name) = &decl.arena[*arr_expr] { + if let Some(ref array_name) = id_place(*arr_expr, context.arena) { self.len_bounds.push(LenBound { index: *var, array: *array_name, @@ -1128,7 +1238,7 @@ impl SafetyChecker { } } // for i in lo .. hi where hi has a LenBound — transitive bound - if let Expr::Id(end_name) = &decl.arena[*end] { + if let Some(ref end_name) = id_place(*end, context.arena) { self.var_bounds.push(VarBound { lo: *var, hi: *end_name, @@ -1138,16 +1248,15 @@ impl SafetyChecker { // Restoring the snapshot after the body also undoes mutations // inside the loop (e.g. `i = i + 1`), so they don't clobber the // constraints of outer variables after the loop exits. - self.check_expr(*body, decl, decls); + self.check_expr(*body, context, decls); self.constraints = saved_constraints.clone(); self.len_bounds = saved_len_bounds.clone(); self.leq_len_bounds = saved_leq_len_bounds; self.min_len_bounds = saved_min_len_bounds; self.var_bounds = saved_var_bounds; - self.vars.truncate(saved_var_count); // Invalidate constraints for variables assigned inside the loop. - self.invalidate_assigned(*body, &decl.arena); + self.invalidate_assigned(*body, context.arena); // Recover bounds for monotonically incrementing variables. // If a variable is only modified by `var = var + 1`, then: @@ -1159,15 +1268,15 @@ impl SafetyChecker { // To verify initial <= start, find the var's Var declaration // in the AST and check if its initializer is the same identifier // as the loop start (e.g., `var i = lo` with `for j in lo .. hi`). - let start_name = if let Expr::Id(n) = &decl.arena[*start] { + let start_name = if let Some(ref n) = id_place(*start, context.arena) { Some(*n) } else { None }; let mut assigned = Vec::new(); - Self::collect_assigned_vars(*body, &decl.arena, &mut assigned); + Self::collect_assigned_vars(*body, context.arena, &mut assigned); for name in assigned { - if !Self::is_monotonic_increment(name, *body, &decl.arena) { + if !Self::is_monotonic_increment(name, *body, context.arena) { continue; } // Restore the pre-loop min bound (monotonic increase preserves it). @@ -1180,16 +1289,17 @@ impl SafetyChecker { // Scan the AST for `Var(name, Some(init), _)` where init // is `Expr::Id(start_name)`. let initialized_from_start = start_name.is_some_and(|sn| { - decl.arena.exprs.iter().any(|e| { - if let Expr::Var(vn, Some(init), _) = e { - *vn == name && matches!(&decl.arena[*init], Expr::Id(n) if *n == sn) + context.arena.ids().any(|id| { + if let Expr::Var(_, Some(init), _) = &context.arena[id] { + Place::local(context.arena.binder(id)) == name + && id_place(*init, context.arena) == Some(sn) } else { false } }) }); if initialized_from_start { - if let Expr::Id(end_name) = &decl.arena[*end] { + if let Some(ref end_name) = id_place(*end, context.arena) { for b in &saved_len_bounds { if b.index == *end_name { self.len_bounds.push(LenBound { @@ -1206,73 +1316,56 @@ impl SafetyChecker { } Expr::ArrayLiteral(exprs) => { for e in exprs { - self.check_expr(*e, decl, decls); + self.check_expr(*e, context, decls); } IndexInterval::default() } - Expr::Lambda { params, body } => { + Expr::Lambda { body, .. } => { // Nothing is known about the arguments at the definition site, // so the body is checked with its parameters unconstrained. // (A directly-called lambda is handled by the `Call` arm, which // knows the arguments.) - self.check_lambda_body(expr, ¶ms.clone(), *body, None, decl, decls); + let params = context.arena.binders(expr); + self.check_lambda_body(params, *body, None, context, decls); IndexInterval::default() } Expr::Tuple(exprs) => { for e in exprs { - self.check_expr(*e, decl, decls); + self.check_expr(*e, context, decls); } IndexInterval::default() } Expr::StructLit(_, fields) => { for (_, e) in fields { - self.check_expr(*e, decl, decls); + self.check_expr(*e, context, decls); } IndexInterval::default() } Expr::AsTy(e, _) | Expr::Arena(e) => { - self.check_expr(*e, decl, decls); + self.check_expr(*e, context, decls); IndexInterval::default() } _ => IndexInterval::default(), } } - /// Check a lambda body. Parameters shadow any captured variable of the - /// same name; `call_args` supplies the argument intervals (and the - /// argument expressions, for symbolic length bounds) when the lambda is - /// called directly at a known call site, and is `None` at the definition - /// site, where the arguments are unknown. + /// Check a lambda body with its identified parameters and captures. + /// Direct calls supply argument intervals and symbolic length bounds; + /// at the definition site parameter values are unknown. fn check_lambda_body( &mut self, - lambda_expr: ExprID, - params: &[Param], + params: &[LocalId], body: ExprID, call_args: Option<(&[ExprID], &[IndexInterval])>, - decl: &FuncDecl, - decls: &DeclTable, + context: SafetyBody<'_>, + decls: &impl SafetyProgram, ) { - let saved_vars = self.vars.clone(); let saved_constraints = self.constraints.clone(); let saved_len_bounds = self.len_bounds.clone(); let saved_leq_len_bounds = self.leq_len_bounds.clone(); let saved_min_len_bounds = self.min_len_bounds.clone(); let saved_var_bounds = self.var_bounds.clone(); - // Lambda params are usually unannotated, so recover their types from - // the solved function type of the lambda expression. - let solved_param_tys = if lambda_expr < decl.types.len() { - match &*decl.types[lambda_expr] { - Type::Func(dom, _) => match &**dom { - Type::Tuple(tys) => tys.clone(), - _ => vec![*dom], - }, - _ => vec![], - } - } else { - vec![] - }; - // A lambda that isn't invoked right here runs at some unknown later // point, so any capture the enclosing function assigns to — before or // after the definition, including from inside this body — may hold a @@ -1285,16 +1378,12 @@ impl SafetyChecker { } } - for (i, param) in params.iter().enumerate() { - let ty = param.ty.or_else(|| solved_param_tys.get(i).copied()); - let is_u32 = ty == Some(mk_type(Type::UInt32)); + for (i, ¶m) in params.iter().enumerate() { + let ty = context.arena.local(param).ty; + let is_u32 = ty == mk_type(Type::UInt32); - // The param shadows any captured variable of the same name. - self.forget(param.name); - self.vars.push(Var { - name: param.name, - ty: ty.unwrap_or_else(|| mk_type(Type::Void)), - }); + // A repeated analysis of this lambda starts with fresh parameter facts. + self.forget(Place::local(param)); let arg = call_args.and_then(|(exprs, ivals)| Some((exprs.get(i)?, ivals.get(i)?))); match arg { @@ -1304,16 +1393,14 @@ impl SafetyChecker { if is_u32 { min = Some(min.unwrap_or(0).max(0)); } - self.add(param.name, min, max); + self.add(Place::local(param), min, max); if ival.non_zero { - self.add_non_zero(param.name); + self.add_non_zero(Place::local(param)); } // The param inherits the argument's symbolic length bounds. - // Read these from the live state, not the entry snapshot: - // `forget` above has already stripped bounds naming an - // array that this param shadows, and bounds pushed for an - // earlier param should propagate. - if let Expr::Id(arg_name) = &decl.arena[*arg_expr] { + // Read the live state so bounds from earlier parameters + // can propagate. Outer captured roots retain their IDs. + if let Some(ref arg_name) = id_place(*arg_expr, context.arena) { let inherited: Vec<_> = self .len_bounds .iter() @@ -1322,20 +1409,19 @@ impl SafetyChecker { .collect(); for array in inherited { self.len_bounds.push(LenBound { - index: param.name, + index: Place::local(param), array, }); } } } - None if is_u32 => self.add(param.name, Some(0), None), - None => self.add(param.name, None, None), + None if is_u32 => self.add(Place::local(param), Some(0), None), + None => self.add(Place::local(param), None, None), } } - self.check_expr(body, decl, decls); + self.check_expr(body, context, decls); - self.vars = saved_vars; self.constraints = saved_constraints; self.len_bounds = saved_len_bounds; self.leq_len_bounds = saved_leq_len_bounds; @@ -1345,7 +1431,7 @@ impl SafetyChecker { // Assignments in the body take effect whenever the lambda is called, // which we can't pin down, so conservatively drop what we knew about // the variables it writes to. - self.invalidate_assigned(body, &decl.arena); + self.invalidate_assigned(body, context.arena); } /// Check if every assignment to `var_name` in the expression tree is of the @@ -1356,13 +1442,13 @@ impl SafetyChecker { /// agree about where an assignment can hide, or the loop-exit min bound /// gets handed back on the strength of an increment that isn't the only /// write. - fn is_monotonic_increment(var_name: Name, expr: ExprID, arena: &ExprArena) -> bool { + fn is_monotonic_increment(var_name: Place, expr: ExprID, arena: &ExprArena) -> bool { if let Expr::Binop(Binop::Assign, lhs, rhs) = &arena[expr] { - if let Expr::Id(name) = &arena[*lhs] { + if let Some(ref name) = id_place(*lhs, arena) { if *name == var_name { // Check rhs is `var_name + 1` if let Expr::Binop(Binop::Plus, plus_lhs, plus_rhs) = &arena[*rhs] { - let lhs_is_var = matches!(&arena[*plus_lhs], Expr::Id(n) if *n == var_name); + let lhs_is_var = id_place(*plus_lhs, arena) == Some(var_name); let rhs_is_one = matches!(&arena[*plus_rhs], Expr::Int(1, _)); return lhs_is_var && rhs_is_one; } @@ -1376,15 +1462,15 @@ impl SafetyChecker { .all(|child| Self::is_monotonic_increment(var_name, *child, arena)) } - /// Collect all variable names that are assigned (via `=`) inside an + /// Collect all direct storage roots assigned (via `=`) inside an /// expression tree. /// /// Walks every subexpression, including lambda bodies, initializers and /// call arguments — an assignment nested in any of those still happens, and /// missing one would leave a stale constraint in place. - fn collect_assigned_vars(expr: ExprID, arena: &ExprArena, out: &mut Vec) { + fn collect_assigned_vars(expr: ExprID, arena: &ExprArena, out: &mut Vec) { if let Expr::Binop(Binop::Assign, lhs, _) = &arena[expr] { - if let Expr::Id(name) = &arena[*lhs] { + if let Some(ref name) = id_place(*lhs, arena) { out.push(*name); } } @@ -1406,145 +1492,147 @@ impl SafetyChecker { } } - /// Inject constraints from top-level `assume` declarations. - /// Uses a temporary FuncDecl so `match_expr` can access the assume's arena. - fn inject_global_assumes(&mut self, decls: &DeclTable) { + /// Inject each assumption's facts with its own body-local coordinates. + fn inject_global_assumes(&mut self, decls: &impl SafetyProgram) { for decl in &decls.decls { if let Decl::Assume { arena, cond } = decl { - let tmp = FuncDecl { - name: Name::str("__assume"), - typevars: vec![], - size_vars: vec![], - params: vec![], - body: None, - ret: mk_type(Type::Void), - constraints: vec![], - requires: vec![], - loc: test_loc(), - arena: arena.clone(), - types: vec![], - closure_vars: vec![], - is_extern: false, - }; - self.match_expr(*cond, &tmp, decls); + self.match_expr( + *cond, + SafetyBody { + arena, + size_vars: &[], + }, + decls, + ); self.propagate_len_bounds(); + // Only nonlocal facts cross this body boundary. A local with the + // same numeric ID in another assumption or function is unrelated. + let nonlocal = |place: Place| !matches!(place.root, PlaceRoot::Local(_)); + self.constraints.retain(|c| nonlocal(c.name)); + self.len_bounds + .retain(|b| nonlocal(b.index) && nonlocal(b.array)); + self.leq_len_bounds + .retain(|b| nonlocal(b.index) && nonlocal(b.array)); + self.min_len_bounds.retain(|b| nonlocal(b.array)); + self.var_bounds.retain(|b| nonlocal(b.lo) && nonlocal(b.hi)); } } } - /// Check that all `require` clauses on the callee hold at this call site. - /// Reports a SafetyError for any clause that cannot be proved. - /// - /// Resolves the callee by name; if there are multiple overloads, picks the - /// one whose arity matches and whose param types match the caller's argument - /// types. Conservatively skips when the callee can't be uniquely resolved. + /// Check a direct call using the phase's target policy. Concrete instances + /// never undergo another signature-based selection; indirect calls remain + /// deferred. Explicit type applications become instance reads in specialization. fn check_call_requires( &mut self, callee_expr: ExprID, args: &[ExprID], call_expr: ExprID, - caller: &FuncDecl, - decls: &DeclTable, + caller: SafetyBody<'_>, + decls: &impl SafetyProgram, ) { - let Expr::Id(callee_name) = &caller.arena[callee_expr] else { + if !matches!(caller.arena[callee_expr], Expr::Id(_)) { return; - }; - - // Find the callee FuncDecl. Disambiguate overloads by matching the - // function type recorded by the type checker for the callee expression. - let callee_ty = if callee_expr < caller.types.len() { - Some(caller.types[callee_expr]) - } else { - None - }; - - let candidates = decls.find(*callee_name); - let mut callee: Option<&FuncDecl> = None; - for d in candidates { - if let Decl::Func(f) = d { - if f.params.len() != args.len() { - continue; - } - if let Some(cty) = callee_ty { - if f.ty() != cty { - continue; - } - } - callee = Some(f); - break; - } } - let Some(callee) = callee else { + let Some(reference) = caller.arena.reference(callee_expr) else { + return; + }; + let Some(callee) = decls.call_target( + reference, + caller.arena, + caller.arena.ty(callee_expr), + args.len(), + ) else { return; }; - if callee.requires.is_empty() { return; } - - // Build name -> caller-arg-ExprID substitution. - let subst: Vec<(Name, ExprID)> = callee + let subst: Vec<(LocalId, ExprID)> = callee .params .iter() .zip(args) - .map(|(p, &a)| (p.name, a)) + .map(|(parameter, &argument)| (parameter.local, argument)) .collect(); - - // Build size-var -> concrete-i64 substitution by matching declared - // param types (which may contain `[T; N]` with a size variable N) - // against the caller's resolved types for each arg. - let mut size_subst: Vec<(Name, i64)> = vec![]; - for (param, &arg) in callee.params.iter().zip(args) { - if let Some(pty) = param.ty { - if arg < caller.types.len() { - let aty = caller.types[arg]; - collect_size_subst(pty, aty, &mut size_subst); + let mut size_subst = Vec::new(); + for (parameter, &argument) in callee.params.iter().zip(args) { + collect_size_subst( + callee.arena.local(parameter.local).ty, + caller.arena.ty(argument), + &mut size_subst, + ); + } + for &requirement in &callee.requires { + if !self.prove_at_call(requirement, callee, caller, &subst, &size_subst, decls) { + let clause = callee.arena.pretty_print(requirement, 0); + let location = caller.arena.loc(call_expr); + let key = (location, callee.arena.loc(requirement), clause.clone()); + if let Some(message) = self.failed_requirements.get(&key) { + // The diagnostic list is public. Clearing it must allow a + // reused checker to report the same requirement again. + if self.errors.iter().any(|error| { + error.location == location && error.message == *message + }) { + continue; + } } - } - } - - for &req in &callee.requires { - if !self.prove_at_call(req, callee, caller, &subst, &size_subst, decls) { - let msg = format!( + let message = format!( "couldn't prove require clause `{}` for call to `{}`", - callee.arena.exprs[req].pretty_print(&callee.arena, 0), - *callee.name + clause, callee.name ); + self.failed_requirements.insert(key, message.clone()); self.push_error(SafetyError { - location: caller.arena.locs[call_expr], - message: msg, + location, + message, }); } } } - /// Try to prove that the require expression `req` (in callee's arena) - /// holds in the caller's current bound state, after substituting params - /// for the corresponding caller argument expressions. - /// - /// Handles a small grammar of forms structurally: - /// - `idx < arr.len` where arr is a slice param: consult `len_bounds`. - /// - `idx < arr.len` where arr is a `[T; N]` param: compare idx's - /// interval against the concrete N derived from the caller's arg. - /// - `idx < N` where N is a size variable: same, via size_subst. - /// - `lhs >= rhs` for constants: interval check. - /// - `lhs && rhs`: prove both. - /// - `true`: trivially provable. - /// Anything else is conservatively unprovable. + /// Evaluate a callee-only expression with a fresh analysis state. Body-local + /// IDs have different owners in caller and callee; caller facts cannot be + /// applied to a coincidentally equal callee LocalId. + fn callee_interval( + expr: ExprID, + callee: &FuncDecl, + decls: &impl SafetyProgram, + ) -> IndexInterval { + Self::new().check_expr(expr, callee.into(), decls) + } + + /// Preserve the small existing precondition proof grammar, translating + /// parameter references across the call boundary by LocalId. fn prove_at_call( &mut self, req: ExprID, callee: &FuncDecl, - caller: &FuncDecl, - subst: &[(Name, ExprID)], + caller: SafetyBody<'_>, + subst: &[(LocalId, ExprID)], size_subst: &[(Name, i64)], - decls: &DeclTable, + decls: &impl SafetyProgram, ) -> bool { - let lookup = - |n: &Name| -> Option { subst.iter().find(|(p, _)| p == n).map(|(_, a)| *a) }; - let size_lookup = - |n: &Name| -> Option { size_subst.iter().find(|(s, _)| s == n).map(|(_, v)| *v) }; - + let lookup = |reference: &Reference| -> Option { + let Reference::Local(local) = reference else { + return None; + }; + subst + .iter() + .find(|(parameter, _)| parameter == local) + .map(|(_, argument)| *argument) + }; + let size_lookup = |reference: &Reference| -> Option { + let Reference::SizeParameter(local) = reference else { + return None; + }; + let symbol = callee + .size_vars + .iter() + .find(|parameter| parameter.local == *local)? + .symbol; + size_subst + .iter() + .find(|(parameter, _)| *parameter == symbol) + .map(|(_, value)| *value) + }; match &callee.arena[req] { Expr::True => true, Expr::Binop(Binop::And, lhs, rhs) => { @@ -1552,33 +1640,33 @@ impl SafetyChecker { && self.prove_at_call(*rhs, callee, caller, subst, size_subst, decls) } Expr::Binop(Binop::Less, lhs, rhs) => { - // Pattern: < .len - if let (Expr::Id(idx_param), Expr::Field(arr_e, fld)) = + if let (Expr::Id(_), Expr::Field(array, field)) = (&callee.arena[*lhs], &callee.arena[*rhs]) { - if fld.as_str() == "len" { - if let Expr::Id(arr_param) = &callee.arena[*arr_e] { - if let (Some(idx_arg), Some(arr_arg)) = - (lookup(idx_param), lookup(arr_param)) + if field.as_str() == "len" { + if let (Expr::Id(_), Some(index), Some(array)) = ( + &callee.arena[*array], + callee.arena.reference(*lhs), + callee.arena.reference(*array), + ) { + if let (Some(index_arg), Some(array_arg)) = + (lookup(index), lookup(array)) { - // Slice path: consult len_bounds when both sides - // resolve to plain Ids in the caller. - if let (Expr::Id(idx_name), Expr::Id(arr_name)) = - (&caller.arena[idx_arg], &caller.arena[arr_arg]) - { + if let (Some(index), Some(array)) = ( + id_place(index_arg, caller.arena), + id_place(array_arg, caller.arena), + ) { if self .len_bounds .iter() - .any(|b| b.index == *idx_name && b.array == *arr_name) + .any(|bound| bound.index == index && bound.array == array) { return true; } } - // Sized-array path: if the caller's arg is a - // fixed-size array, prove via interval check. - if let Some(k) = static_len(arr_arg, caller, decls) { - let li = self.check_expr(idx_arg, caller, decls); - if li.max != i64::MAX && li.max < k { + if let Some(length) = static_len(array_arg, caller, decls) { + let interval = self.check_expr(index_arg, caller, decls); + if interval.max != i64::MAX && interval.max < length { return true; } } @@ -1586,59 +1674,50 @@ impl SafetyChecker { } } } - // Pattern: < - if let (Expr::Id(lhs_param), Expr::Id(rhs_id)) = - (&callee.arena[*lhs], &callee.arena[*rhs]) - { - if let (Some(la), Some(n_val)) = (lookup(lhs_param), size_lookup(rhs_id)) { - let li = self.check_expr(la, caller, decls); - if li.max != i64::MAX && li.max < n_val { + if let (Expr::Id(_), Expr::Id(_)) = (&callee.arena[*lhs], &callee.arena[*rhs]) { + if let (Some(argument), Some(size)) = ( + callee.arena.reference(*lhs).and_then(lookup), + callee.arena.reference(*rhs).and_then(size_lookup), + ) { + let interval = self.check_expr(argument, caller, decls); + if interval.max != i64::MAX && interval.max < size { return true; } } } - // Fallback: interval comparison `lhs.max < rhs.min`. Resolve - // each side's interval — params via the caller arg, other - // expressions (constants, arithmetic) via the callee arena. - let li = if let Expr::Id(lhs_param) = &callee.arena[*lhs] { - if let Some(la) = lookup(lhs_param) { - Some(self.check_expr(la, caller, decls)) - } else { - None - } + let left = if let Expr::Id(_) = &callee.arena[*lhs] { + callee + .arena + .reference(*lhs) + .and_then(lookup) + .map(|argument| self.check_expr(argument, caller, decls)) } else { - Some(self.check_expr(*lhs, callee, decls)) + Some(Self::callee_interval(*lhs, callee, decls)) }; - let ri = if let Expr::Id(rhs_param) = &callee.arena[*rhs] { - if let Some(ra) = lookup(rhs_param) { - Some(self.check_expr(ra, caller, decls)) - } else { - None - } + let right = if let Expr::Id(_) = &callee.arena[*rhs] { + callee + .arena + .reference(*rhs) + .and_then(lookup) + .map(|argument| self.check_expr(argument, caller, decls)) } else { - Some(self.check_expr(*rhs, callee, decls)) + Some(Self::callee_interval(*rhs, callee, decls)) }; - if let (Some(li), Some(ri)) = (li, ri) { - if li.max != i64::MAX && ri.min != i64::MIN && li.max < ri.min { - return true; + match (left, right) { + (Some(left), Some(right)) => { + left.max != i64::MAX && right.min != i64::MIN && left.max < right.min } + _ => false, } - false } Expr::Binop(Binop::Geq, lhs, rhs) => { - // Pattern: >= → check caller-arg.min >= rhs.min. - if let Expr::Id(lhs_param) = &callee.arena[*lhs] { - if let Some(arg) = lookup(lhs_param) { - let arg_iv = self.check_expr(arg, caller, decls); - // Evaluate rhs in the callee's arena with no constraints — - // works for constant expressions like `0`. - let rhs_iv = self.check_expr(*rhs, callee, decls); - if arg_iv.min != i64::MIN - && rhs_iv.max != i64::MAX - && arg_iv.min >= rhs_iv.max - { - return true; - } + if let Expr::Id(_) = &callee.arena[*lhs] { + if let Some(argument) = callee.arena.reference(*lhs).and_then(lookup) { + let argument = self.check_expr(argument, caller, decls); + let rhs = Self::callee_interval(*rhs, callee, decls); + return argument.min != i64::MIN + && rhs.max != i64::MAX + && argument.min >= rhs.max; } } false @@ -1647,22 +1726,16 @@ impl SafetyChecker { } } - fn check_fn_decl(&mut self, func_decl: &FuncDecl, decls: &DeclTable) { + fn check_fn_decl(&mut self, func_decl: &FuncDecl, decls: &impl SafetyProgram) { if let Some(body) = func_decl.body { // Inject top-level assume constraints before checking the function. self.inject_global_assumes(decls); for param in &func_decl.params { - if let Some(ty) = param.ty { - self.vars.push(Var { - name: param.name, - ty, - }); - if ty == mk_type(Type::UInt32) { - self.add(param.name, Some(0), None); - } else { - self.add(param.name, None, None); - } + if func_decl.arena.local(param.local).ty == mk_type(Type::UInt32) { + self.add(Place::local(param.local), Some(0), None); + } else { + self.add(Place::local(param.local), None, None); } } @@ -1671,18 +1744,14 @@ impl SafetyChecker { // checker can reason about `for i in 0 .. N` and `arr[idx]` for // `[T; N]` parameters. The body still has to prove `idx < N` via // a for-loop, require clause, or local check. - for &sv in &func_decl.size_vars { - self.vars.push(Var { - name: sv, - ty: mk_type(Type::Int32), - }); - self.add(sv, Some(1), None); + for parameter in &func_decl.size_vars { + self.add(Place::local(parameter.local), Some(1), None); } // Inject require clauses as assumptions inside the function body. - // The clauses live in func_decl.arena and reference parameter names. + // The clauses live in this checked body and reference parameter IDs. for &req in &func_decl.requires { - self.match_expr(req, func_decl, decls); + self.match_expr(req, func_decl.into(), decls); } if !func_decl.requires.is_empty() { self.propagate_len_bounds(); @@ -1693,9 +1762,8 @@ impl SafetyChecker { self.fn_assigned.clear(); Self::collect_assigned_vars(body, &func_decl.arena, &mut self.fn_assigned); - self.check_expr(body, &func_decl, decls); + self.check_expr(body, func_decl.into(), decls); - self.vars.clear(); self.constraints.clear(); self.len_bounds.clear(); self.leq_len_bounds.clear(); @@ -1705,7 +1773,7 @@ impl SafetyChecker { } } - pub fn check_decl(&mut self, decl: &Decl, decls: &DeclTable) { + pub fn check_decl(&mut self, decl: &CheckedDecl, decls: &impl SafetyProgram) { match decl { Decl::Func(func_decl) => { // Skip generic functions with size variables — they'll be @@ -1720,337 +1788,208 @@ impl SafetyChecker { } } - pub fn check(&mut self, decls: &DeclTable) { + pub fn check(&mut self, decls: &impl SafetyProgram) { for decl in &decls.decls { self.check_decl(decl, decls); } } - /// Reject recursion under `--no-recursion` by building a conservative - /// call graph that covers direct calls, inline lambda calls, and - /// indirect calls (through parameters, `var`/`let` bindings, struct - /// fields, etc.) and running Tarjan's SCC over it. Any cycle — whether - /// entirely between top-level functions, between lambdas, or spanning - /// both — is reported. - /// - /// The graph has one node per top-level function with a body and one - /// node per `Expr::Lambda` literal in the program. Edges: - /// * Direct call to a top-level function (unshadowed `Expr::Id`) → - /// edge to that function's node. - /// * Direct call to an inline lambda literal → edge to that - /// lambda's node. - /// * Indirect call (anything else — call through a parameter, a - /// var/let/field, a lambda-valued expression, etc.) → conservative - /// edges to every function whose address has been "taken" - /// anywhere in the program. The address-taken set is: - /// (a) every top-level function whose name appears in a - /// non-callee position (passed as an argument, stored in a - /// binding, returned, etc.), plus - /// (b) every lambda literal in the program (creating a lambda - /// produces a first-class function value). - /// - /// Calls inside a lambda body are attributed to the lambda's node, - /// not the enclosing function — so a lambda that calls through a - /// captured fn-typed var forms a self-loop via the indirect edge, - /// while a lambda with no fn-typed captures simply has no outgoing - /// indirect edges. - pub fn check_recursion(&mut self, decls: &DeclTable) { + /// Reject cycles in checked direct calls, inline lambda calls and the + /// conservative graph of indirect calls to address-taken functions. + /// Lexical shadowing has already been resolved; graph construction never + /// reconstructs bindings from function or local spellings. + pub fn check_recursion(&mut self, program: &CheckedProgram) { + use crate::scc::{scc_is_cycle, strongly_connected_components}; use std::collections::{HashMap, HashSet}; - // --- Node representation --- #[derive(Clone, Copy)] - enum NodeKind { - TopLevel, - Lambda { arena_idx: ExprID }, - } - struct NodeInfo { - decl_idx: usize, // containing top-level decl - kind: NodeKind, - } - - // --- Step 1: create top-level nodes. --- - let mut nodes: Vec = Vec::new(); - let mut top_node_of: HashMap = HashMap::new(); - for (i, decl) in decls.decls.iter().enumerate() { - if let Decl::Func(f) = decl { - if f.body.is_some() { - top_node_of.insert(i, nodes.len()); - nodes.push(NodeInfo { - decl_idx: i, - kind: NodeKind::TopLevel, - }); - } + struct Node { + definition: DefId, + lambda: Option, + } + struct Graph { + nodes: Vec, + top: HashMap, + lambdas: HashMap<(DefId, ExprID), usize>, + calls: Vec>, + address_taken: HashSet, + } + fn definitions(reference: &Reference, body: &CheckedBody) -> Vec { + match reference { + Reference::Functions(ids) => ids.clone(), + Reference::InterfaceMember { + requirement, + member, + } => body + .requirements + .get(requirement.index()) + .and_then(|requirement| { + requirement + .members + .iter() + .find(|candidate| candidate.definition == *member) + }) + .map(|member| member.candidates.clone()) + .unwrap_or_default(), + _ => vec![], } } - - // --- Step 2: walk each top-level function recursively. --- - // Per-node call-site list and per-(decl,arena_idx) lambda node - // lookup for inline-lambda-call resolution. - let mut calls_at: Vec> = vec![Vec::new(); nodes.len()]; - let mut lambda_node_of: HashMap<(usize, ExprID), usize> = HashMap::new(); - let mut address_taken_names: HashSet = HashSet::new(); - - #[allow(clippy::too_many_arguments)] fn walk( - expr: ExprID, - arena: &ExprArena, - decl_idx: usize, + expression: ExprID, + body: &CheckedBody, + definition: DefId, current: usize, - in_callee_pos: bool, - nodes: &mut Vec, - calls_at: &mut Vec>, - lambda_node_of: &mut HashMap<(usize, ExprID), usize>, - address_taken: &mut HashSet, + in_callee: bool, + graph: &mut Graph, ) { - let e = &arena.exprs[expr]; - match e { - Expr::Id(name) => { - if !in_callee_pos { - address_taken.insert(*name); + match &body[expression] { + Expr::Id(_) | Expr::TypeApp(_, _) => { + if !in_callee { + if let Some(reference) = body.reference(expression) { + graph.address_taken.extend(definitions(reference, body)); + } } } - Expr::Lambda { body, .. } => { - let new_node = nodes.len(); - nodes.push(NodeInfo { - decl_idx, - kind: NodeKind::Lambda { arena_idx: expr }, + Expr::Lambda { + body: lambda_body, .. + } => { + let node = graph.nodes.len(); + graph.nodes.push(Node { + definition, + lambda: Some(expression), }); - calls_at.push(Vec::new()); - lambda_node_of.insert((decl_idx, expr), new_node); - walk( - *body, - arena, - decl_idx, - new_node, - false, - nodes, - calls_at, - lambda_node_of, - address_taken, - ); - } - Expr::Call(func_id, args) => { - calls_at[current].push(expr); - // Callee is walked in "callee position" so an `Expr::Id` - // callee is not added to the address-taken set. - walk( - *func_id, - arena, - decl_idx, - current, - true, - nodes, - calls_at, - lambda_node_of, - address_taken, - ); - for a in args { - walk( - *a, - arena, - decl_idx, - current, - false, - nodes, - calls_at, - lambda_node_of, - address_taken, - ); + graph.calls.push(Vec::new()); + graph.lambdas.insert((definition, expression), node); + walk(*lambda_body, body, definition, node, false, graph); + } + Expr::Call(callee, arguments) => { + graph.calls[current].push(expression); + walk(*callee, body, definition, current, true, graph); + for &argument in arguments { + walk(argument, body, definition, current, false, graph); } } - _ => { - for child in e.subexprs() { - walk( - child, - arena, - decl_idx, - current, - false, - nodes, - calls_at, - lambda_node_of, - address_taken, - ); + expression => { + for child in expression.subexprs() { + walk(child, body, definition, current, false, graph); } } } } - - // Snapshot the top-level node list to iterate independently of - // the growing `nodes` vector (we push lambda nodes during walk). - let top_level_starts: Vec<(usize, usize)> = nodes - .iter() - .enumerate() - .filter_map(|(n, info)| match info.kind { - NodeKind::TopLevel => Some((n, info.decl_idx)), - _ => None, - }) - .collect(); - - for &(top_node, decl_idx) in &top_level_starts { - if let Decl::Func(func) = &decls.decls[decl_idx] { - if let Some(body) = func.body { - walk( - body, - &func.arena, - decl_idx, - top_node, - false, - &mut nodes, - &mut calls_at, - &mut lambda_node_of, - &mut address_taken_names, - ); + let mut graph = Graph { + nodes: Vec::new(), + top: HashMap::new(), + lambdas: HashMap::new(), + calls: Vec::new(), + address_taken: HashSet::new(), + }; + for (coordinate, declaration) in program.decls.decls.iter().enumerate() { + if let Decl::Func(function) = declaration { + if function.body.is_some() { + let definition = program.decls.id_at(coordinate); + graph.top.insert(definition, graph.nodes.len()); + graph.nodes.push(Node { + definition, + lambda: None, + }); + graph.calls.push(Vec::new()); } } } - - // --- Step 3: compute the address-taken node set. --- - // Every lambda is implicitly address-taken (it's a first-class - // function value). Top-level functions join it only if their - // name appeared outside callee position. - let mut at_nodes: Vec = Vec::new(); - for (n, info) in nodes.iter().enumerate() { - match info.kind { - NodeKind::Lambda { .. } => at_nodes.push(n), - NodeKind::TopLevel => { - if let Decl::Func(f) = &decls.decls[info.decl_idx] { - if address_taken_names.contains(&f.name) { - at_nodes.push(n); - } - } - } - } + let roots = graph.nodes.clone(); + for (node, root) in roots.iter().enumerate() { + let function = program.function(root.definition).unwrap(); + walk( + function.body.unwrap(), + &function.arena, + root.definition, + node, + false, + &mut graph, + ); } - - // --- Step 4: compute per-top-level locals set (for shadowing - // detection). Lambdas inside a top-level reuse this set. - let mut locals_by_decl: HashMap> = HashMap::new(); - for &(_, decl_idx) in &top_level_starts { - let Decl::Func(func) = &decls.decls[decl_idx] else { - continue; - }; - let mut local: HashSet = HashSet::new(); - for p in &func.params { - local.insert(p.name); - } - for expr in &func.arena.exprs { - match expr { - Expr::Let(name, _, _) => { - local.insert(*name); - } - Expr::Var(name, _, _) => { - local.insert(*name); - } - Expr::For { var, .. } => { - local.insert(*var); - } - Expr::Lambda { params, .. } => { - for p in params { - local.insert(p.name); - } - } - _ => {} - } - } - locals_by_decl.insert(decl_idx, local); - } - - // --- Step 5: build adjacency. For each node, classify its calls. --- - let mut adj: Vec> = vec![Vec::new(); nodes.len()]; - for node_idx in 0..nodes.len() { - let decl_idx = nodes[node_idx].decl_idx; - let Decl::Func(func) = &decls.decls[decl_idx] else { - continue; - }; - let empty = HashSet::new(); - let locals = locals_by_decl.get(&decl_idx).unwrap_or(&empty); - - let mut callees: HashSet = HashSet::new(); - for &call_id in &calls_at[node_idx] { - let Expr::Call(callee_expr, _) = &func.arena.exprs[call_id] else { - continue; + let address_taken: Vec<_> = graph + .nodes + .iter() + .enumerate() + .filter_map(|(index, node)| { + (node.lambda.is_some() || graph.address_taken.contains(&node.definition)) + .then_some(index) + }) + .collect(); + let mut adjacency = vec![Vec::new(); graph.nodes.len()]; + for (index, node) in graph.nodes.iter().enumerate() { + let function = program.function(node.definition).unwrap(); + let mut targets = HashSet::new(); + for &call in &graph.calls[index] { + let Expr::Call(callee, _) = &function.arena[call] else { + unreachable!(); }; - - let callee = &func.arena.exprs[*callee_expr]; - let mut resolved_direct = false; - match callee { - Expr::Id(name) if !locals.contains(name) => { - let mut found_any = false; - for (di, d) in decls.decls.iter().enumerate() { - if let Decl::Func(df) = d { - if df.name == *name { - found_any = true; - if df.body.is_some() { - if let Some(&cn) = top_node_of.get(&di) { - callees.insert(cn); - } - } + let direct = match &function.arena[*callee] { + Expr::Id(_) | Expr::TypeApp(_, _) => { + let definitions = function + .arena + .reference(*callee) + .map(|reference| definitions(reference, &function.arena)) + .unwrap_or_default(); + let mut found = false; + for definition in definitions { + if program.function(definition).is_some() { + found = true; + if let Some(&target) = graph.top.get(&definition) { + targets.insert(target); } } } - if found_any { - resolved_direct = true; - } + found } Expr::Lambda { .. } => { - if let Some(&cn) = lambda_node_of.get(&(decl_idx, *callee_expr)) { - callees.insert(cn); - resolved_direct = true; + if let Some(&target) = graph.lambdas.get(&(node.definition, *callee)) { + targets.insert(target); + true + } else { + false } } - _ => {} - } - - if !resolved_direct { - // Indirect call: add edges to every address-taken node. - for &at in &at_nodes { - callees.insert(at); - } + _ => false, + }; + if !direct { + targets.extend(address_taken.iter().copied()); } } - adj[node_idx] = callees.into_iter().collect(); + adjacency[index] = targets.into_iter().collect(); } - - // 3. Run Tarjan's SCC on the adjacency list. - let sccs = strongly_connected_components(&adj); - - // 4. Report cycles. SCC size > 1 is always a cycle; size 1 is a - // cycle only if the single node has a self-edge. - let describe = |n: usize| -> (Loc, String) { - let info = &nodes[n]; - let Decl::Func(f) = &decls.decls[info.decl_idx] else { - unreachable!("non-function decl appears as a call-graph node"); - }; - match info.kind { - NodeKind::TopLevel => (f.loc, format!("function `{}`", f.name)), - NodeKind::Lambda { arena_idx } => { - (f.arena.locs[arena_idx], format!("lambda in `{}`", f.name)) - } + let describe = |index: usize| { + let node = graph.nodes[index]; + let function = program.function(node.definition).unwrap(); + match node.lambda { + Some(expression) => ( + function.arena.loc(expression), + format!("lambda in `{}`", function.name), + ), + None => (function.loc, format!("function `{}`", function.name)), } }; - - for scc in &sccs { - if !scc_is_cycle(scc, &adj) { + for component in strongly_connected_components(&adjacency) { + if !scc_is_cycle(&component, &adjacency) { continue; } - - if scc.len() == 1 { - let (loc, desc) = describe(scc[0]); + if component.len() == 1 { + let (location, description) = describe(component[0]); self.push_error(SafetyError { - location: loc, - message: format!("--no-recursion: {} is recursive", desc), + location, + message: format!("--no-recursion: {} is recursive", description), }); } else { - let descs: Vec = scc.iter().map(|&n| describe(n).1).collect(); - let cycle_desc = descs.join(", "); - for &n in scc { - let (loc, desc) = describe(n); + let descriptions: Vec<_> = component.iter().map(|&node| describe(node).1).collect(); + let cycle = descriptions.join(", "); + for node in component { + let (location, description) = describe(node); self.push_error(SafetyError { - location: loc, + location, message: format!( "--no-recursion: {} participates in a recursive cycle [{}]", - desc, cycle_desc + description, cycle ), }); } @@ -2069,34 +2008,139 @@ impl SafetyChecker { mod tests { use super::*; - pub fn check(s: &str) -> Vec { + fn checked_source(s: &str) -> CheckedProgram { let mut errors = vec![]; let decls = parse_program_str(&s, &mut errors); assert!(errors.is_empty()); - assert_eq!(decls.len(), 1); - let mut table = DeclTable::new(decls); - let mut types = vec![]; - for decl in &table.decls { - let mut type_checker = Checker::new(); - type_checker.check_decl(decl, &table); - assert!(type_checker.errors.is_empty()); - types.push(type_checker.solved_types()); - } - - for i in 0..table.decls.len() { - if let Decl::Func(ref mut fdecl) = &mut table.decls[i] { - fdecl.types = types[i].clone(); - } - } + let table = DeclTable::new(decls); + let check_function = |function: crate::FuncDecl| { + let mut checker = Checker::new(); + checker.check_decl(&Decl::Func(function.clone()), &table); + assert!(checker.errors.is_empty()); + checker.checked_function(&function) + }; + CheckedProgram::new(table.map_bodies( + |_, function| check_function(function), + |arena, cond| { + let mut checker = Checker::new(); + checker.check_decl( + &Decl::Assume { + arena: arena.clone(), + cond, + }, + &table, + ); + assert!(checker.errors.is_empty(), "{:?}", checker.errors); + checker.checked_body(&arena) + }, + )) + } + pub fn check(s: &str) -> Vec { + let checked = checked_source(s); let mut array_checker = SafetyChecker::new(); - array_checker.check(&table); + array_checker.check(&checked); array_checker.print_errors(); array_checker.errors } + #[test] + fn concrete_requirements_use_the_instance_despite_coercing_node_signatures() { + for (definition, main, recorded_parameter) in [ + ( + "bounded(x: i32, values: [i32]) require x >= 0 {}", + "bounded(-1, [1, 2])", + "array", + ), + ( + "bounded(x: &i32) require x >= 0 {}", + "var x = -1; bounded(x)", + "value", + ), + ] { + let templates = checked_source(&format!("{} main {{ {} }}", definition, main)); + let mut program = MonomorphPass::new() + .monomorphize(&templates, Name::str("main")) + .unwrap(); + let main = program.instance_for_entry(Name::str("main")).unwrap(); + let declaration = program.instances[main.index()].declaration; + let Decl::Func(caller) = &mut program.decls.decls[declaration] else { + unreachable!() + }; + let callee = caller + .arena + .exprs() + .iter() + .find_map(|expr| { + if let Expr::Call(callee, _) = expr { + Some(*callee) + } else { + None + } + }) + .unwrap(); + let Some(&Reference::Instance(target)) = caller.arena.reference(callee) else { + unreachable!() + }; + let integer = mk_type(Type::Int32); + let parameters = if recorded_parameter == "array" { + vec![integer, mk_type(Type::Array(integer, ArraySize::Known(2)))] + } else { + vec![integer] + }; + let recorded = mk_type(Type::Func( + mk_type(Type::Tuple(parameters)), + mk_type(Type::Void), + )); + caller.arena.set_ty(callee, recorded); + let actual = program.function_instance(target).unwrap().ty(); + assert_ne!(actual, recorded); + assert!(unify(actual, recorded, &mut Instance::new())); + program.validate().unwrap(); + + let mut safety = SafetyChecker::new(); + for _ in 0..2 { + safety.errors.clear(); + safety.check(&program); + assert_eq!(safety.errors.len(), 1); + assert!(safety.errors[0].message.contains("`x >= 0`")); + } + } + } + + #[test] + fn shadowed_field_constraints_do_not_escape_to_the_outer_binding() { + let errors = check("struct P { x: i32 } f { var a: [i32; 3]; var p: P; p.x = 100; if true { var p: P; p.x = 0; a[p.x] }; a[p.x] }"); + assert_eq!(errors.len(), 1); + assert!(errors[0].message.contains("less than array length")); + } + + #[test] + fn local_callee_does_not_inherit_a_same_named_function_requirement() { + let errors = + check("bounded(x: i32) require x >= 0 {} f { let bounded = |x: i32| {}; bounded(-1) }"); + assert!(errors.is_empty()); + } + + #[test] + fn callee_local_ids_cannot_observe_caller_constraint_slots() { + let errors = check("bounded(x: i32, y: i32) require x >= (y + 1) {} f { let value = 0; let unrelated = -100; bounded(0, 100) }"); + assert_eq!(errors.len(), 1); + assert!(errors[0].message.contains("couldn't prove require clause")); + } + + #[test] + fn recursion_graph_uses_direct_bindings_even_when_a_sibling_scope_shadows_the_name() { + let checked = checked_source("f { if true { let f = |x: i32| {} }; f() }"); + let mut safety = SafetyChecker::new(); + safety.check_recursion(&checked); + assert_eq!(safety.errors.len(), 1); + assert!(safety.errors[0] + .message + .contains("function `f` is recursive")); + } #[test] pub fn test_array_if() { let s = " diff --git a/src/solver.rs b/src/solver.rs index 316e9025..dd4eec9b 100644 --- a/src/solver.rs +++ b/src/solver.rs @@ -13,6 +13,7 @@ use std::hash::{Hash, Hasher}; pub struct AltInterface { pub interface: Name, pub typevars: Vec, + pub members: Option>, } impl AltInterface { @@ -21,23 +22,18 @@ impl AltInterface { AltInterface { interface: self.interface, typevars: self.typevars.iter().map(|ty| ty.subst(inst)).collect(), + members: self + .members + .as_ref() + .map(|members| members.iter().map(|member| member.subst(inst)).collect()), } } /// Is the constraint satisfied in the current environment? - pub fn satisfied(&self, instance: &Instance, decls: &DeclTable, loc: Loc) -> bool { - if let Some(Decl::Interface(interface)) = decls.find(self.interface).first() { - let mut types = vec![]; - for ty in &self.typevars { - types.push(ty.subst(instance)); - } - - let mut tmp_errors = vec![]; - interface.satisfied(&types, decls, &mut tmp_errors, loc) - } else { - // Unknown interface! - false - } + pub fn satisfied(&self, instance: &Instance, decls: &DeclTable, _loc: Loc) -> bool { + self.members.as_ref().is_some_and(|members| { + select_interface_members(&self.typevars, members, instance, decls).is_ok() + }) } } @@ -359,9 +355,22 @@ pub fn iterate_solver( // We've narrowed it down. Better unify! if let Some(field) = find_field(&st.fields, field_name) { let field_ty = if let Type::Var(name) = *field.ty { - let index = - st.typevars.iter().position(|&n| n == name).unwrap(); - vars[index] + let Some(ty) = st + .typevars + .iter() + .position(|&n| n == name) + .and_then(|index| vars.get(index)) + else { + errors.push(TypeError { + location: loc, + message: format!( + "missing type argument for field '{}' of {}", + field_name, struct_name + ), + }); + continue; + }; + *ty } else { field.ty }; diff --git a/src/source_analysis.rs b/src/source_analysis.rs new file mode 100644 index 00000000..92cec798 --- /dev/null +++ b/src/source_analysis.rs @@ -0,0 +1,174 @@ +//! Editor facts from one checking run, including incomplete bodies. +//! +//! This is deliberately not a checked program or executable input. The source +//! inventory owns definition IDs; each body's source arena and local inventory +//! own expression/local IDs. None of these coordinates survive another analysis. +use crate::*; +use std::collections::{HashMap, HashSet}; + +#[derive(Clone, Debug, Default)] +pub struct ExpressionFacts { + /// An established type, never an inference variable or recovery placeholder. + pub ty: Option, + /// Recorded lexical resolution, independent of type availability. An overload + /// set is an ordered candidate inventory, not a selected call target. + pub reference: Option, + /// The binding introduced by a let/var/for expression, when visited. + pub binding: Option, +} + +#[derive(Clone, Debug)] +pub struct AnalyzedLocal { + pub name: Name, + pub ty: Option, + /// Existing source precision: parameters use their enclosing function/lambda + /// location; declarations use the binding statement's location. + pub loc: Loc, +} + +#[derive(Clone, Debug)] +pub struct BodyAnalysis { + pub(crate) expressions: Vec, + pub(crate) locals: Vec, + pub(crate) requirements: HashMap, +} + +impl BodyAnalysis { + pub fn expression(&self, id: ExprID) -> Option<&ExpressionFacts> { + self.expressions.get(id) + } + + pub fn local(&self, id: LocalId) -> Option<&AnalyzedLocal> { + self.locals.get(id.index()) + } + + /// The declaration owning this requirement's members. Failed requirements + /// can leave gaps in the checker's ordinal IDs. This is identity information; + /// partially solved requirement signatures are not exposed as type facts. + pub fn requirement_interface(&self, id: RequirementId) -> Option { + self.requirements.get(&id).copied() + } +} + +/// An immutable snapshot produced by `Compiler::analyze`. Source declarations +/// are inventory, not certified signatures. Missing bodies/facts mean unavailable, +/// never permission to reconstruct semantic meaning by looking up a spelling. +#[derive(Clone, Debug)] +pub struct SourceAnalysis { + declarations: DeclTable, + pub(crate) bodies: HashMap, + unavailable_signatures: HashSet, + recovered_types: HashSet, + recovered_files: HashSet, +} + +impl SourceAnalysis { + pub(crate) fn new( + declarations: DeclTable, + mut recovered: HashSet, + recovered_files: HashSet, + ) -> Self { + let mut recovered_types = HashSet::new(); + for record in declarations.records() { + if recovered.contains(&record.definition) { + recovered.extend(record.members); + if matches!(record.declaration, Decl::Struct(_) | Decl::Enum { .. }) { + recovered_types.insert(record.declaration.name()); + } + } + } + let mut analysis = Self { + declarations, + bodies: HashMap::new(), + unavailable_signatures: recovered, + recovered_types, + recovered_files, + }; + for record in analysis.declarations.records() { + let functions = std::iter::once(record.definition).chain(record.members); + for id in functions { + if let Some(function) = analysis.declarations.function(id) { + let mut signature_scope = function.clone(); + if let Decl::Interface(interface) = &record.declaration { + signature_scope.typevars.extend(&interface.typevars); + } + if !function + .annotated_ty() + .is_some_and(|ty| analysis.type_is_available(ty, &signature_scope)) + { + analysis.unavailable_signatures.insert(id); + } + } + } + } + analysis + } + + pub fn declarations(&self) -> &DeclTable { + &self.declarations + } + + pub fn body(&self, definition: DefId) -> Option<&BodyAnalysis> { + self.bodies.get(&definition) + } + + pub(crate) fn source_is_trusted(&self, source: &FuncDecl) -> bool { + !self.recovered_files.contains(&source.loc.file) + && source + .arena + .locs + .iter() + .all(|loc| !self.recovered_files.contains(&loc.file)) + } + + pub(crate) fn reference_is_trusted(&self, reference: &Reference) -> bool { + match reference { + Reference::Global(id) | Reference::InterfaceMember { member: id, .. } => { + !self.unavailable_signatures.contains(id) + } + Reference::Functions(ids) => { + !ids.is_empty() + && ids + .iter() + .all(|id| !self.unavailable_signatures.contains(id)) + } + Reference::Local(_) | Reference::SizeParameter(_) => true, + Reference::Instance(_) => false, + } + } + + pub(crate) fn type_is_available(&self, ty: TypeID, source: &FuncDecl) -> bool { + !ty.contains_anon() && self.type_is_valid(ty, source) + } + + /// Check annotation provenance and named-type validity using the same source + /// declaration inventory. This is not overload or lexical name resolution. + /// Anonymous variables are allowed here only to inspect pre-solve types; + /// type_is_available additionally excludes them from published facts. + pub(crate) fn type_is_valid(&self, ty: TypeID, source: &FuncDecl) -> bool { + match &*ty { + Type::Name(name, args) => { + !self.recovered_types.contains(name) + && match self.declarations.find(*name).first() { + Some(Decl::Struct(st)) => st.typevars.len() == args.len(), + Some(Decl::Enum { .. }) => args.is_empty(), + _ => false, + } + && args.iter().all(|ty| self.type_is_valid(*ty, source)) + } + Type::Tuple(args) => args.iter().all(|ty| self.type_is_valid(*ty, source)), + Type::Func(domain, ret) => { + self.type_is_valid(*domain, source) && self.type_is_valid(*ret, source) + } + // The parser already checks size-symbol scope. A callee's size + // parameter can remain in a successfully checked call signature + // until specialization, even though it is not a caller parameter. + Type::Array(ty, _) | Type::Slice(ty) | Type::Reference(ty) => { + self.type_is_valid(*ty, source) + } + // Named generic parameters are established symbolic types. + Type::Var(name) => source.typevars.contains(name), + _ => true, + } + } +} diff --git a/src/stack_codegen.rs b/src/stack_codegen.rs index 438d8810..6c8c6e9d 100644 --- a/src/stack_codegen.rs +++ b/src/stack_codegen.rs @@ -1,15 +1,15 @@ //! Stack-based code generator. //! -//! This module translates a DeclTable into a StackProgram that can be +//! This module translates a SpecializedProgram into a StackProgram that can be //! executed by a stack-based virtual machine. It mirrors the register-based //! VM codegen but emits stack IR instructions instead. -use crate::decl::*; +use crate::checked::{CheckedFunction, InstanceId, LocalId, Reference, SpecializedProgram}; +use crate::decl::Decl; use crate::defs::*; -use crate::expr::*; +use crate::expr::Expr; use crate::stack_ir::*; use crate::types::*; -use crate::DeclTable; use std::collections::{HashMap, HashSet}; /// Loop context for break/continue support. @@ -30,17 +30,8 @@ struct PendingCall { func_idx: u32, /// Index of the instruction within that function. instr_idx: usize, - /// Name of the function being called. - callee: Name, -} - -/// A snapshot of a translator's name-keyed binding state. See -/// `FunctionTranslator::save_bindings`. -struct SavedBindings { - variables: HashMap, - variable_types: HashMap, - captured_vars: HashSet, - captured_slots: HashMap, + /// Concrete function instance being called. + callee: InstanceId, } /// How a local variable is stored. @@ -59,14 +50,14 @@ pub struct StackCodegen { /// The program being built. program: StackProgram, - /// Map from function names to their indices in the program. - func_indices: HashMap, + /// Map from function instances to their indices in the program. + func_indices: HashMap, /// Functions that have been compiled. - compiled_functions: HashSet, + compiled_functions: HashSet, /// Functions that need to be compiled. - pending_functions: Vec, + pending_functions: Vec, /// Calls that need to be patched after all functions are compiled. pending_calls: Vec, @@ -75,7 +66,7 @@ pub struct StackCodegen { pending_func_loads: Vec, /// Global variable offsets. - globals: HashMap, + globals: HashMap, /// Counter for generating unique lambda names. lambda_counter: usize, @@ -109,31 +100,22 @@ impl StackCodegen { /// the JIT/LLVM backends, which lets the FFI layer write the stack /// interp's structural trap reason to `TRAP_REASON_OFFSET` so hosts /// can call `read_trap_reason(globals)` uniformly across backends. - fn declare_globals(&mut self, decls: &DeclTable) { + fn declare_globals(&mut self, decls: &SpecializedProgram) { let mut offset: i32 = crate::cancel::CANCEL_FLAG_RESERVED; - for decl in &decls.decls { - match decl { - Decl::Global { - name, typevars, ty, .. - } => { - if typevars.is_empty() { - self.globals.insert(*name, offset); - offset += ty.size(decls) as i32; - } - } - Decl::Func(f) if f.is_extern => { - // Extern functions get 16 bytes: {fn_ptr, context} - self.globals.insert(f.name, offset); - offset += 16; - } - _ => {} - } + for (instance, decl) in decls.storage_instances() { + let size = match decl { + Decl::Global { ty, .. } => ty.size(decls) as i32, + Decl::Func(f) if f.is_extern => 16, + _ => continue, + }; + self.globals.insert(instance, offset); + offset += size; } self.program.globals_size = offset as usize; } - /// Compile a DeclTable into a StackProgram. - pub fn compile(&mut self, decls: &DeclTable) -> Result { + /// Compile a SpecializedProgram into a StackProgram. + pub fn compile(&mut self, decls: &SpecializedProgram) -> Result { let main_name = Name::str("main"); self.compile_multi(decls, &[main_name]) } @@ -144,31 +126,31 @@ impl StackCodegen { /// found show up in `program.entry_points`. pub fn compile_multi( &mut self, - decls: &DeclTable, + decls: &SpecializedProgram, entry_points: &[Name], ) -> Result { self.declare_globals(decls); for &ep_name in entry_points { - if self.compiled_functions.contains(&ep_name) { - continue; - } - let Some(ep_decl) = decls.find_entry_point(ep_name) else { + let Some(instance) = decls.instance_for_entry(ep_name) else { continue; }; - self.compile_function(ep_decl, decls)?; + if self.compiled_functions.contains(&instance) { + continue; + } + let ep_decl = decls + .function_instance(instance) + .expect("entry is a function"); + self.compile_function(ep_decl, decls, Some(instance))?; while let Some(name) = self.pending_functions.pop() { if self.compiled_functions.contains(&name) { continue; } - let func_decls = decls.find(name); - if func_decls.is_empty() { - continue; - } - if let Decl::Func(func_decl) = &func_decls[0] { - self.compile_function(func_decl, decls)?; - } + let func_decl = decls + .function_instance(name) + .expect("call target is a function"); + self.compile_function(func_decl, decls, Some(name))?; } } @@ -177,7 +159,10 @@ impl StackCodegen { // stays at its default and the map is empty. let mut entry_set = false; for &ep_name in entry_points { - if let Some(&idx) = self.func_indices.get(&ep_name) { + if let Some(&idx) = decls + .instance_for_entry(ep_name) + .and_then(|id| self.func_indices.get(&id)) + { self.program.entry_points.insert(ep_name, idx); if !entry_set { self.program.entry = idx; @@ -213,7 +198,12 @@ impl StackCodegen { } /// Compile a single function. - fn compile_function(&mut self, decl: &FuncDecl, decls: &DeclTable) -> Result { + fn compile_function( + &mut self, + decl: &CheckedFunction, + decls: &SpecializedProgram, + instance: Option, + ) -> Result { let mut func = StackFunction::new(&*decl.name); func.param_count = decl.params.len() as u8; @@ -227,8 +217,10 @@ impl StackCodegen { translator.translate(&mut func); let idx = self.program.add_function(func); - self.func_indices.insert(decl.name, idx); - self.compiled_functions.insert(decl.name); + if let Some(instance) = instance { + self.func_indices.insert(instance, idx); + self.compiled_functions.insert(instance); + } // Collect pending calls. let calls_to_patch = std::mem::take(&mut translator.calls_to_patch); @@ -256,7 +248,7 @@ impl StackCodegen { // Compile lambda functions and patch their indices. for lambda_decl in pending_lambdas { let lambda_name = lambda_decl.name; - let lambda_idx = self.compile_function(&lambda_decl, decls)?; + let lambda_idx = self.compile_function(&lambda_decl, decls, None)?; for &(instr_idx, patch_name) in &lambda_patches { if patch_name == lambda_name { if let StackOp::I64Const(ref mut value) = @@ -275,7 +267,7 @@ impl StackCodegen { /// A call instruction that needs patching. struct CallToPatch { instr_idx: usize, - callee: Name, + callee: InstanceId, } /// Check if a type should be returned via output pointer (sret). @@ -304,16 +296,13 @@ fn stack_extern_ret_type(ty: TypeID) -> StackExternRet { /// Translator for a single function body. struct FunctionTranslator<'a> { /// The function declaration being translated. - decl: &'a FuncDecl, + decl: &'a CheckedFunction, /// Declaration table for looking up types and functions. - decls: &'a DeclTable, - - /// Map from variable names to their local storage kind. - variables: HashMap, + decls: &'a SpecializedProgram, - /// Declared representation type for local bindings. - variable_types: HashMap, + /// Map from local identities to their local storage kind. + variables: HashMap, /// Next available scalar local slot. next_scalar: u16, @@ -322,7 +311,7 @@ struct FunctionTranslator<'a> { next_memory_slot: u16, /// Functions that are called and need to be compiled. - pending_functions: &'a mut Vec, + pending_functions: &'a mut Vec, /// Counter for generating unique lambda names. lambda_counter: &'a mut usize, @@ -330,8 +319,8 @@ struct FunctionTranslator<'a> { /// Calls that need patching. calls_to_patch: Vec, - /// Lambda FuncDecls extracted from this function body, to be compiled afterward. - pending_lambdas: Vec, + /// Lambda CheckedFunctions extracted from this function body, to be compiled afterward. + pending_lambdas: Vec, /// I64Const instructions that need to be patched with lambda function indices. lambda_patches: Vec<(usize, Name)>, @@ -340,7 +329,7 @@ struct FunctionTranslator<'a> { func_load_patches: Vec, /// Global variable offsets. - globals: &'a HashMap, + globals: &'a HashMap, /// Memory slot for the sret output pointer (if returning ptr type). output_ptr_slot: Option, @@ -352,14 +341,14 @@ struct FunctionTranslator<'a> { loop_stack: Vec, /// Variables captured from an enclosing scope (double indirection). - captured_vars: HashSet, + captured_vars: HashSet, /// Names a lambda in this function mentions. Shared with the closure by /// address, so they must be memory-backed from the start. - lambda_referenced: HashSet, + lambda_referenced: HashSet, /// Memory slot indices for captured variables (stores pointer-to-storage). - captured_slots: HashMap, + captured_slots: HashMap, /// True when the current expression's result will be discarded. void_ctx: bool, @@ -371,17 +360,17 @@ struct FunctionTranslator<'a> { impl<'a> FunctionTranslator<'a> { fn new( - decl: &'a FuncDecl, - decls: &'a DeclTable, - pending_functions: &'a mut Vec, + decl: &'a CheckedFunction, + decls: &'a SpecializedProgram, + pending_functions: &'a mut Vec, lambda_counter: &'a mut usize, - globals: &'a HashMap, + globals: &'a HashMap, ) -> Self { Self { decl, decls, variables: HashMap::new(), - variable_types: HashMap::new(), + next_scalar: 0, next_memory_slot: 0, pending_functions, @@ -395,7 +384,7 @@ impl<'a> FunctionTranslator<'a> { has_returned: false, loop_stack: Vec::new(), captured_vars: HashSet::new(), - lambda_referenced: decl.names_referenced_in_lambdas(), + lambda_referenced: decl.captured_locals(), captured_slots: HashMap::new(), void_ctx: false, elidable_lets: crate::copy_elision::elidable_let_copies(decl), @@ -420,7 +409,7 @@ impl<'a> FunctionTranslator<'a> { /// Get the type of an expression. fn expr_type(&self, expr: ExprID) -> TypeID { - self.decl.types[expr] + self.decl.arena.ty(expr) } /// Get the type that determines how an expression is represented at runtime. @@ -428,21 +417,18 @@ impl<'a> FunctionTranslator<'a> { /// A call site can solve an array expression as a slice, while codegen still /// has an array address and must explicitly build the slice fat pointer. fn representation_type(&self, expr: ExprID) -> TypeID { - match &self.decl.arena.exprs[expr] { - Expr::Id(name) => self - .variable_types - .get(name) - .copied() - .or_else(|| { - self.decls.find(*name).iter().find_map(|decl| { - if let Decl::Global { ty, .. } = decl { - Some(*ty) - } else { - None - } - }) - }) - .unwrap_or_else(|| self.expr_type(expr)), + match &self.decl.arena[expr] { + Expr::Id(_) => match self.decl.arena.reference(expr) { + Some(Reference::Local(local)) => { + let ty = self.decl.arena.local(*local).ty; + match &*ty { + Type::Reference(inner) => *inner, + _ => ty, + } + } + Some(Reference::Instance(instance)) => self.decls.instance(*instance).ty(), + _ => self.expr_type(expr), + }, Expr::ArrayIndex(arr_id, _) => match &*self.representation_type(*arr_id) { Type::Array(elem, _) | Type::Slice(elem) | Type::Reference(elem) => *elem, _ => self.expr_type(expr), @@ -490,15 +476,14 @@ impl<'a> FunctionTranslator<'a> { let param_offset = if has_sret { 1u16 } else { 0u16 }; for (i, param) in self.decl.params.iter().enumerate() { let param_slot = param_offset + i as u16; - let ty = param.ty.expect("parameter must have type"); + let ty = self.decl.arena.local(param.local).ty; - if let Type::Reference(inner) = &*ty { + if let Type::Reference(_) = &*ty { while self.next_scalar <= param_slot { self.alloc_scalar(); } self.variables - .insert(param.name, LocalKind::Reference(param_slot)); - self.variable_types.insert(param.name, *inner); + .insert(param.local, LocalKind::Reference(param_slot)); } else if !self.is_ptr_type(&ty) { // Scalar parameter: already in local slot param_slot by calling convention. // Just make sure our allocator accounts for it. @@ -506,8 +491,7 @@ impl<'a> FunctionTranslator<'a> { self.alloc_scalar(); } self.variables - .insert(param.name, LocalKind::Scalar(param_slot)); - self.variable_types.insert(param.name, ty); + .insert(param.local, LocalKind::Scalar(param_slot)); } else { // Pointer-represented parameters are passed as addresses. // Keep the address value directly, matching the JIT/LLVM ABI. @@ -515,8 +499,21 @@ impl<'a> FunctionTranslator<'a> { self.alloc_scalar(); } self.variables - .insert(param.name, LocalKind::Scalar(param_slot)); - self.variable_types.insert(param.name, ty); + .insert(param.local, LocalKind::Scalar(param_slot)); + } + } + + // Capture creation can occur on only one branch. Addressable scalar + // parameters therefore need initialized storage before control flow splits. + for (i, param) in self.decl.params.iter().enumerate() { + let ty = self.decl.arena.local(param.local).ty; + if self.lambda_referenced.contains(¶m.local) && !self.is_ptr_type(&ty) { + let mem_slot = self.alloc_memory(self.vm_type_size(&ty)); + func.emit(StackOp::LocalAddr(mem_slot)); + self.emit_local_get(&ty, param_offset + i as u16, func); + self.emit_store_op(&ty, func); + self.variables + .insert(param.local, LocalKind::Memory(mem_slot)); } } @@ -539,12 +536,10 @@ impl<'a> FunctionTranslator<'a> { let addr_local = self.alloc_scalar(); func.emit(StackOp::LocalSet(addr_local)); // Save for later access. - self.captured_vars.insert(cv.name); - self.captured_slots.insert(cv.name, addr_local); + self.captured_vars.insert(*cv); + self.captured_slots.insert(*cv, addr_local); // Also register in variables so nested closures can find this capture. - self.variables - .insert(cv.name, LocalKind::Scalar(addr_local)); - self.variable_types.insert(cv.name, cv.ty); + self.variables.insert(*cv, LocalKind::Scalar(addr_local)); } } @@ -617,27 +612,11 @@ impl<'a> FunctionTranslator<'a> { /// Translate an expression in void context (result will be discarded). /// Only optimizes specific expression types known to be safe. fn translate_void(&mut self, expr: ExprID, func: &mut StackFunction) { - match &self.decl.arena.exprs[expr].clone() { - // Var: the fusion pass already eliminates the i64.const 0 + drop pattern. - // Just translate normally and let the caller drop. - Expr::Var(..) => { - self.translate_expr(expr, func); - func.emit(StackOp::Drop); - } - // Let: translate then drop the result. The let expression's - // type is the initializer type, so f32 lets leave the value - // on the f-window and need DropF — an int Drop here leaks the - // f-window value and eventually overflows the float spill. - Expr::Let(..) => { - self.translate_expr(expr, func); - let ty = self.expr_type(expr); - if matches!(&*ty, Type::Float32) { - func.emit(StackOp::DropF); - } else if matches!(&*ty, Type::Float64) { - func.emit(StackOp::DropD); - } else { - func.emit(StackOp::Drop); - } + match &self.decl.arena[expr].clone() { + // Declarations have no language value. Lower their storage effects + // directly, without materializing an operand-stack placeholder. + Expr::Let(..) | Expr::Var(..) => { + self.translate_expr_inner(expr, func, true); } // An f32x4 assignment in statement position: void context is // what lets translate_assign send the vector ops straight at @@ -658,13 +637,9 @@ impl<'a> FunctionTranslator<'a> { Expr::Block(exprs) => { let exprs = exprs.clone(); if !exprs.is_empty() { - let saved_vars = self.variables.clone(); - let saved_types = self.variable_types.clone(); for &expr_id in exprs.iter() { self.translate_void(expr_id, func); } - self.variables = saved_vars; - self.variable_types = saved_types; } } // If in void context: no need to produce a value on both branches. @@ -728,14 +703,20 @@ impl<'a> FunctionTranslator<'a> { // restored the caller's TOS window from memory, but with the // no-spill op_call/op_return design any trailing value leaves // the callee's final depth unbalanced and leaks into the caller. + let declaration = matches!(self.decl.arena[expr], Expr::Let(..) | Expr::Var(..)) + && matches!(&*self.expr_type(expr), Type::Void); let old_void_ctx = self.void_ctx; - self.void_ctx = void_ctx; + self.void_ctx = void_ctx || declaration; self.translate_expr_inner_body(expr, func); self.void_ctx = old_void_ctx; + if declaration && !void_ctx { + // Generic expression composition uses one inert value for void. + func.emit(StackOp::I64Const(0)); + } } fn translate_expr_inner_body(&mut self, expr: ExprID, func: &mut StackFunction) { - match &self.decl.arena.exprs[expr].clone() { + match &self.decl.arena[expr].clone() { Expr::Int(n, _) => { func.emit(StackOp::I64Const(*n)); } @@ -788,13 +769,15 @@ impl<'a> FunctionTranslator<'a> { func.emit(StackOp::LocalAddr(mem_slot)); } - Expr::Id(name) => { - self.translate_id(*name, expr, func); + Expr::Id(_) => { + let decl = self.decl; + let reference = decl.arena.reference(expr).expect("checked reference"); + self.translate_id(reference, expr, func); } Expr::Enum(case_name) => { let case_name = *case_name; - let index = if let Type::Name(enum_name, _) = &*self.decl.types[expr] { + let index = if let Type::Name(enum_name, _) = &*self.decl.arena.ty(expr) { let enum_decls = self.decls.find(*enum_name); if let Some(Decl::Enum { cases, .. }) = enum_decls.iter().find(|d| matches!(d, Decl::Enum { .. })) @@ -827,10 +810,10 @@ impl<'a> FunctionTranslator<'a> { self.translate_call(*fn_id, &arg_ids, expr, func); } - Expr::Let(name, init, _) => { - let name = *name; + Expr::Let(_, init, _) => { + let name = self.decl.arena.binder(expr); let init = *init; - let ty = self.expr_type(expr); + let ty = self.decl.arena.local(name).ty; if !self.is_ptr_type(&ty) && self.lambda_referenced.contains(&name) { // Captured by a lambda: memory-backed from the start. @@ -841,9 +824,9 @@ impl<'a> FunctionTranslator<'a> { func.emit(StackOp::LocalAddr(mem_slot)); self.emit_local_get(&ty, tmp, func); self.emit_store_op(&ty, func); - self.shadow_outer_binding(&name); + self.variables.insert(name, LocalKind::Memory(mem_slot)); - self.variable_types.insert(name, ty); + if !self.void_ctx { self.emit_local_get(&ty, tmp, func); } @@ -856,9 +839,8 @@ impl<'a> FunctionTranslator<'a> { } else { self.emit_local_tee(&ty, local, func); } - self.shadow_outer_binding(&name); + self.variables.insert(name, LocalKind::Scalar(local)); - self.variable_types.insert(name, ty); } else if crate::copy_elision::is_value_aggregate(&ty) && !self.elidable_lets.contains(&expr) { @@ -875,9 +857,9 @@ impl<'a> FunctionTranslator<'a> { func.emit(StackOp::LocalAddr(mem_slot)); func.emit(StackOp::LocalGet(tmp)); func.emit(StackOp::MemCopy(size)); - self.shadow_outer_binding(&name); + self.variables.insert(name, LocalKind::Memory(mem_slot)); - self.variable_types.insert(name, ty); + if !self.void_ctx { func.emit(StackOp::LocalAddr(mem_slot)); } @@ -891,16 +873,15 @@ impl<'a> FunctionTranslator<'a> { } else { func.emit(StackOp::LocalTee(local)); } - self.shadow_outer_binding(&name); + self.variables.insert(name, LocalKind::Scalar(local)); - self.variable_types.insert(name, ty); } } - Expr::Var(name, init, _) => { - let name = *name; + Expr::Var(_, init, _) => { + let name = self.decl.arena.binder(expr); let init = *init; - let ty = self.expr_type(expr); + let ty = self.decl.arena.local(name).ty; if !self.is_ptr_type(&ty) && self.lambda_referenced.contains(&name) { // Captured by a lambda: memory-backed from the start. @@ -917,9 +898,8 @@ impl<'a> FunctionTranslator<'a> { func.emit(StackOp::LocalAddr(mem_slot)); func.emit(StackOp::MemZero(size)); } - self.shadow_outer_binding(&name); + self.variables.insert(name, LocalKind::Memory(mem_slot)); - self.variable_types.insert(name, ty); } else if !self.is_ptr_type(&ty) { let local = self.alloc_scalar(); if let Some(init_id) = init { @@ -934,9 +914,8 @@ impl<'a> FunctionTranslator<'a> { func.emit(StackOp::I64Const(0)); func.emit(StackOp::LocalSet(local)); } - self.shadow_outer_binding(&name); + self.variables.insert(name, LocalKind::Scalar(local)); - self.variable_types.insert(name, ty); } else { let size = self.vm_type_size(&ty); let mem_slot = self.alloc_memory(size); @@ -945,9 +924,9 @@ impl<'a> FunctionTranslator<'a> { // own storage — no temp, no 16-byte copy. if matches!(&*ty, Type::Float32x4) && self.f32x4_slot_form(init_id) { self.emit_f32x4_into_slot(init_id, mem_slot, func); - self.shadow_outer_binding(&name); + self.variables.insert(name, LocalKind::Memory(mem_slot)); - self.variable_types.insert(name, ty); + if !self.void_ctx { func.emit(StackOp::I64Const(0)); } @@ -957,9 +936,9 @@ impl<'a> FunctionTranslator<'a> { self.emit_f32x4_operands(init_id, func); func.emit(StackOp::LocalAddr(mem_slot)); func.emit(store_op); - self.shadow_outer_binding(&name); + self.variables.insert(name, LocalKind::Memory(mem_slot)); - self.variable_types.insert(name, ty); + if !self.void_ctx { func.emit(StackOp::I64Const(0)); } @@ -976,9 +955,8 @@ impl<'a> FunctionTranslator<'a> { func.emit(StackOp::LocalAddr(mem_slot)); func.emit(StackOp::MemZero(size)); } - self.shadow_outer_binding(&name); + self.variables.insert(name, LocalKind::Memory(mem_slot)); - self.variable_types.insert(name, ty); } // Var expressions produce void; push 0 only if result is needed. if !self.void_ctx { @@ -993,7 +971,6 @@ impl<'a> FunctionTranslator<'a> { func.emit(StackOp::I64Const(0)); } } else { - let saved = self.save_bindings(); for (i, &expr_id) in exprs.iter().enumerate() { if i < exprs.len() - 1 { // Intermediate expressions: void context. @@ -1006,7 +983,6 @@ impl<'a> FunctionTranslator<'a> { self.translate_expr(expr_id, func); } } - self.restore_bindings(saved); } } @@ -1023,12 +999,10 @@ impl<'a> FunctionTranslator<'a> { } Expr::For { - var, - start, - end, - body, + start, end, body, .. } => { - self.translate_for(*var, *start, *end, *body, func); + let var = self.decl.arena.binder(expr); + self.translate_for(var, *start, *end, *body, func); if !self.void_ctx { func.emit(StackOp::I64Const(0)); } @@ -1158,11 +1132,7 @@ impl<'a> FunctionTranslator<'a> { self.translate_cast(*expr_id, *target_ty, func); } - Expr::Lambda { params, body } => { - let params = params.clone(); - let body = *body; - self.translate_lambda(¶ms, body, expr, func); - } + Expr::Lambda { .. } => self.translate_lambda(expr, func), Expr::Assume(_) => { // No-op: assume is only used by the safety checker. @@ -1174,120 +1144,68 @@ impl<'a> FunctionTranslator<'a> { } Expr::TypeApp(_, _) | Expr::Macro(_, _) | Expr::Error => { - func.emit(StackOp::I64Const(0)); + unreachable!("unresolved expression in specialized body") } } } - /// Forget everything known about an outer binding of `name`, so a new - /// binding that shadows it is a clean rebinding. Reads consult - /// `captured_vars` before `variables`, so a leftover entry sends them - /// through the enclosing scope's indirection instead of to this binding. - /// Block scope saves and restores these, so the outer binding's state - /// comes back at block exit. - fn shadow_outer_binding(&mut self, name: &Name) { - self.captured_vars.remove(name); - self.captured_slots.remove(name); - } - - /// Snapshot every name-keyed binding fact, to be restored when the scope - /// that shadowed it ends. Blocks and `for` loops both need this: a binding - /// made inside one must stop being visible when it ends, and an outer - /// binding of the same name must come back. - fn save_bindings(&self) -> SavedBindings { - SavedBindings { - variables: self.variables.clone(), - variable_types: self.variable_types.clone(), - captured_vars: self.captured_vars.clone(), - captured_slots: self.captured_slots.clone(), - } - } - - fn restore_bindings(&mut self, saved: SavedBindings) { - self.variables = saved.variables; - self.variable_types = saved.variable_types; - self.captured_vars = saved.captured_vars; - self.captured_slots = saved.captured_slots; - } - /// Translate an identifier reference. - fn translate_id(&mut self, name: Name, expr: ExprID, func: &mut StackFunction) { + fn translate_id(&mut self, reference: &Reference, expr: ExprID, func: &mut StackFunction) { let ty = self.expr_type(expr); - - // Captured closure variable (double indirection). - if self.captured_vars.contains(&name) { - let addr_local = *self.captured_slots.get(&name).unwrap(); - // Load the pointer to the captured variable's storage. - func.emit(StackOp::LocalGet(addr_local)); - // Aggregates and slices are represented by their address, and the - // captured pointer already is that address — dereferencing it would - // yield the first word of the value. - if !self.is_ptr_type(&ty) { - // Load the value through the pointer. - self.emit_load(&ty, func); - } - return; - } - - // Local variable. - if let Some(&kind) = self.variables.get(&name) { - match kind { - LocalKind::Scalar(slot) => { - self.emit_local_get(&ty, slot, func); - } - LocalKind::Reference(slot) => { - func.emit(StackOp::LocalGet(slot)); + match reference { + Reference::Local(local) => { + if let Some(&addr_local) = self.captured_slots.get(local) { + func.emit(StackOp::LocalGet(addr_local)); if !self.is_ptr_type(&ty) { self.emit_load(&ty, func); } + return; } - LocalKind::Memory(slot) => { - if self.is_ptr_type(&ty) { - // Pointer types: push address. - func.emit(StackOp::LocalAddr(slot)); - } else { - // Scalar in memory slot: load value. + match self.variables[local] { + LocalKind::Scalar(slot) => self.emit_local_get(&ty, slot, func), + LocalKind::Reference(slot) => { + func.emit(StackOp::LocalGet(slot)); + if !self.is_ptr_type(&ty) { + self.emit_load(&ty, func); + } + } + LocalKind::Memory(slot) => { func.emit(StackOp::LocalAddr(slot)); - self.emit_load(&ty, func); + if !self.is_ptr_type(&ty) { + self.emit_load(&ty, func); + } } } } - return; - } - - // Global variable. - if let Some(&offset) = self.globals.get(&name) { - func.emit(StackOp::GlobalAddr(offset)); - if !self.is_ptr_type(&ty) { - self.emit_load(&ty, func); + Reference::Instance(instance) => { + if let Some(&offset) = self.globals.get(instance) { + func.emit(StackOp::GlobalAddr(offset)); + if !self.is_ptr_type(&ty) { + self.emit_load(&ty, func); + } + return; + } + assert!( + self.decls.function_instance(*instance).is_some(), + "reference must name storage or a function" + ); + let mem_slot = self.alloc_memory(16); + func.emit(StackOp::LocalAddr(mem_slot)); + let instr_idx = func.pos(); + func.emit(StackOp::I64Const(0)); + self.pending_functions.push(*instance); + self.func_load_patches.push(CallToPatch { + instr_idx, + callee: *instance, + }); + func.emit(StackOp::Store64); + func.emit(StackOp::LocalAddr(mem_slot)); + func.emit(StackOp::I64Const(0)); + func.emit(StackOp::Store64Off(8)); + func.emit(StackOp::LocalAddr(mem_slot)); } - return; - } - - // Function reference — build fat pointer {func_idx, 0}. - if let Type::Func(_, _) = &*ty { - let mem_slot = self.alloc_memory(16); - // Store func_idx at offset 0. - func.emit(StackOp::LocalAddr(mem_slot)); - let instr_idx = func.pos(); - func.emit(StackOp::I64Const(0)); // placeholder - self.pending_functions.push(name); - self.func_load_patches.push(CallToPatch { - instr_idx, - callee: name, - }); - func.emit(StackOp::Store64); - // Store closure_ptr = 0 at offset 8. - func.emit(StackOp::LocalAddr(mem_slot)); - func.emit(StackOp::I64Const(0)); - func.emit(StackOp::Store64Off(8)); - // Push fat pointer address. - func.emit(StackOp::LocalAddr(mem_slot)); - return; + _ => unreachable!("non-concrete reference in specialized body"), } - - // Unknown: push 0. - func.emit(StackOp::I64Const(0)); } /// The store-form vector op for an f32x4-producing expression, or @@ -1303,7 +1221,7 @@ impl<'a> FunctionTranslator<'a> { if !matches!(&*self.expr_type(expr), Type::Float32x4) { return None; } - match &self.decl.arena.exprs[expr] { + match &self.decl.arena[expr] { Expr::Binop(op, lhs_id, _) => { if !matches!(&*self.expr_type(*lhs_id), Type::Float32x4) { return None; @@ -1327,9 +1245,10 @@ impl<'a> FunctionTranslator<'a> { if self.holds_fat_pointer(*fn_id) { return None; } - let Expr::Id(name) = &self.decl.arena.exprs[*fn_id] else { + let Some(Reference::Instance(instance)) = self.decl.arena.reference(*fn_id) else { return None; }; + let name = self.decls.instance_name(*instance); match (name.as_str(), arg_ids.len()) { ("f32x4", 4) => Some(StackOp::F32x4BuildStore), ("f32x4_splat", 1) => Some(StackOp::F32x4SplatStore), @@ -1354,7 +1273,7 @@ impl<'a> FunctionTranslator<'a> { if self.get_memory_slot(expr).is_some() { return true; } - match &self.decl.arena.exprs[expr] { + match &self.decl.arena[expr] { Expr::Binop(op, lhs_id, rhs_id) => { matches!(op, Binop::Plus | Binop::Minus | Binop::Mult | Binop::Div) && self.f32x4_slot_form(*lhs_id) @@ -1367,7 +1286,7 @@ impl<'a> FunctionTranslator<'a> { /// The operands of `expr` if it is an f32x4 multiplication. fn f32x4_mul_operands(&self, expr: ExprID) -> Option<(ExprID, ExprID)> { - match &self.decl.arena.exprs[expr] { + match &self.decl.arena[expr] { Expr::Binop(Binop::Mult, lhs_id, rhs_id) if matches!(&*self.expr_type(expr), Type::Float32x4) => { @@ -1402,7 +1321,7 @@ impl<'a> FunctionTranslator<'a> { } return; } - match &self.decl.arena.exprs[expr] { + match &self.decl.arena[expr] { Expr::Binop(op, lhs_id, rhs_id) => { let (op, lhs_id, rhs_id) = (*op, *lhs_id, *rhs_id); // `a * b + c`, `c + a * b` and `a * b - c` each collapse to @@ -1453,7 +1372,7 @@ impl<'a> FunctionTranslator<'a> { /// Push the operands of an expression [`Self::f32x4_store_op`] /// accepted, leaving the destination address to the caller. fn emit_f32x4_operands(&mut self, expr: ExprID, func: &mut StackFunction) { - match &self.decl.arena.exprs[expr] { + match &self.decl.arena[expr] { Expr::Binop(_, lhs_id, rhs_id) => { let (lhs_id, rhs_id) = (*lhs_id, *rhs_id); self.translate_expr(lhs_id, func); @@ -1576,14 +1495,12 @@ impl<'a> FunctionTranslator<'a> { Type::UInt32 | Type::UInt8 => func.emit(StackOp::ULt), _ => func.emit(StackOp::ILt), }, - Binop::Greater => { - match &*ty { - Type::Float32 => func.emit(StackOp::FGtF), - Type::Float64 => func.emit(StackOp::DGtD), - Type::UInt32 | Type::UInt8 => func.emit(StackOp::UGt), - _ => func.emit(StackOp::IGt), - } - } + Binop::Greater => match &*ty { + Type::Float32 => func.emit(StackOp::FGtF), + Type::Float64 => func.emit(StackOp::DGtD), + Type::UInt32 | Type::UInt8 => func.emit(StackOp::UGt), + _ => func.emit(StackOp::IGt), + }, Binop::Leq => match &*ty { Type::Float32 => func.emit(StackOp::FLeF), Type::Float64 => func.emit(StackOp::DLeD), @@ -1608,7 +1525,7 @@ impl<'a> FunctionTranslator<'a> { let lhs_ty = self.representation_type(lhs_id); // Check for captured variable assignment (double indirection). - if let Expr::Id(name) = &self.decl.arena.exprs[lhs_id] { + if let Some(Reference::Local(name)) = self.decl.arena.reference(lhs_id) { let name = *name; if self.captured_vars.contains(&name) { self.translate_expr(rhs_id, func); @@ -1632,7 +1549,7 @@ impl<'a> FunctionTranslator<'a> { } // Direct scalar local assignment. - if let Expr::Id(name) = &self.decl.arena.exprs[lhs_id] { + if let Some(Reference::Local(name)) = self.decl.arena.reference(lhs_id) { let name = *name; if let Some(&LocalKind::Scalar(slot)) = self.variables.get(&name) { // Try to emit a register-form `locals[slot] = a OP b` op @@ -1652,7 +1569,7 @@ impl<'a> FunctionTranslator<'a> { } // Slice store: a[i] = rhs where a is a slice of 32-bit elements. - if let Expr::ArrayIndex(arr_id, idx_id) = &self.decl.arena.exprs[lhs_id] { + if let Expr::ArrayIndex(arr_id, idx_id) = &self.decl.arena[lhs_id] { let arr_id = *arr_id; let idx_id = *idx_id; let arr_ty = self.representation_type(arr_id); @@ -1777,7 +1694,7 @@ impl<'a> FunctionTranslator<'a> { // For Func type field assignment, only copy func_idx (8 bytes). if matches!(&*lhs_ty, Type::Func(_, _)) { - if matches!(&self.decl.arena.exprs[lhs_id], Expr::Field(_, _)) { + if matches!(&self.decl.arena[lhs_id], Expr::Field(_, _)) { // rhs is a fat pointer address; load func_idx and store. func.emit(StackOp::Load64); // load func_idx from value (which is fat ptr addr) func.emit(StackOp::Store64); @@ -1796,7 +1713,7 @@ impl<'a> FunctionTranslator<'a> { /// If this expr is an Id that resolves to a memory-backed local, return the slot index. fn get_memory_slot(&self, expr: ExprID) -> Option { - if let Expr::Id(name) = &self.decl.arena.exprs[expr] { + if let Some(Reference::Local(name)) = self.decl.arena.reference(expr) { if let Some(LocalKind::Memory(slot)) = self.variables.get(name) { return Some(*slot); } @@ -1806,7 +1723,7 @@ impl<'a> FunctionTranslator<'a> { /// If this expr is an Id that resolves to a scalar local, return the local index. fn get_scalar_local(&self, expr: ExprID) -> Option { - if let Expr::Id(name) = &self.decl.arena.exprs[expr] { + if let Some(Reference::Local(name)) = self.decl.arena.reference(expr) { if let Some(LocalKind::Scalar(local)) = self.variables.get(name) { return Some(*local); } @@ -1825,7 +1742,7 @@ impl<'a> FunctionTranslator<'a> { rhs_id: ExprID, func: &mut StackFunction, ) -> bool { - let (op, lhs, rhs) = match &self.decl.arena.exprs[rhs_id] { + let (op, lhs, rhs) = match &self.decl.arena[rhs_id] { Expr::Binop(op, lhs, rhs) => (*op, *lhs, *rhs), _ => return false, }; @@ -1858,32 +1775,36 @@ impl<'a> FunctionTranslator<'a> { /// Translate an lvalue expression. Pushes the address onto the stack. fn translate_lvalue(&mut self, expr: ExprID, func: &mut StackFunction) { - match &self.decl.arena.exprs[expr].clone() { - Expr::Id(name) => { - let name = *name; - if let Some(&kind) = self.variables.get(&name) { - match kind { - LocalKind::Scalar(slot) => { - let ty = self.expr_type(expr); - if self.is_ptr_type(&ty) { + match &self.decl.arena[expr].clone() { + Expr::Id(_) => match self.decl.arena.reference(expr) { + Some(Reference::Local(name)) => { + let name = *name; + if let Some(&kind) = self.variables.get(&name) { + match kind { + LocalKind::Scalar(slot) => { + let ty = self.expr_type(expr); + if self.is_ptr_type(&ty) { + func.emit(StackOp::LocalGet(slot)); + } else { + self.emit_var_address(&name, func); + } + } + LocalKind::Reference(slot) => { func.emit(StackOp::LocalGet(slot)); - } else { - self.emit_var_address(&name, func); + } + LocalKind::Memory(slot) => { + func.emit(StackOp::LocalAddr(slot)); } } - LocalKind::Reference(slot) => { - func.emit(StackOp::LocalGet(slot)); - } - LocalKind::Memory(slot) => { - func.emit(StackOp::LocalAddr(slot)); - } + } else { + unreachable!("checked local must have storage"); } - } else if let Some(&offset) = self.globals.get(&name) { - func.emit(StackOp::GlobalAddr(offset)); - } else { - func.emit(StackOp::I64Const(0)); } - } + Some(Reference::Instance(instance)) => { + func.emit(StackOp::GlobalAddr(self.globals[instance])); + } + reference => unreachable!("unresolved checked reference: {:?}", reference), + }, Expr::Field(lhs_id, name) => { let lhs_id = *lhs_id; @@ -2004,8 +1925,9 @@ impl<'a> FunctionTranslator<'a> { } // Check for builtin functions. - if let Expr::Id(name) = &self.decl.arena.exprs[fn_id] { - let name = *name; + if let Some(Reference::Instance(instance)) = self.decl.arena.reference(fn_id) { + let instance = *instance; + let name = self.decls.instance_name(instance); if *name == "print" { if let Some(&arg_id) = arg_ids.first() { @@ -2207,14 +2129,12 @@ impl<'a> FunctionTranslator<'a> { } // Extern function calls. - if let Expr::Id(callee_name) = &self.decl.arena.exprs[fn_id] { - let callee_name = *callee_name; - let callee_decls = self.decls.find(callee_name); - if let Some(Decl::Func(f)) = callee_decls.first() { + { + if let Some(f) = self.decls.function_instance(instance) { if f.is_extern { let globals_offset = *self .globals - .get(&callee_name) + .get(&instance) .expect("extern function not in globals"); // For extern calls, push C-level args. Slices expand @@ -2222,7 +2142,7 @@ impl<'a> FunctionTranslator<'a> { let mut c_arg_count: u8 = 0; for (i, arg_id) in arg_ids.iter().enumerate() { let arg = *arg_id; - let param_ty = f.params[i].ty.unwrap(); + let param_ty = f.arena.local(f.params[i].local).ty; if matches!(&*param_ty, Type::Slice(_)) { self.translate_expr(arg, func); match &*self.representation_type(arg) { @@ -2297,8 +2217,7 @@ impl<'a> FunctionTranslator<'a> { // Get callee param types for slice coercion. let param_types: Vec = { - let callee_decls = self.decls.find(name); - if let Some(Decl::Func(f)) = callee_decls.first() { + if let Some(f) = self.decls.function_instance(instance) { f.param_types() } else { vec![] @@ -2339,7 +2258,7 @@ impl<'a> FunctionTranslator<'a> { arg_ids.len() as u8 }; - self.pending_functions.push(name); + self.pending_functions.push(instance); let instr_idx = func.pos(); func.emit(StackOp::Call { func: 0, @@ -2348,7 +2267,7 @@ impl<'a> FunctionTranslator<'a> { }); self.calls_to_patch.push(CallToPatch { instr_idx, - callee: name, + callee: instance, }); // If sret, push the output address as the result. @@ -2422,21 +2341,13 @@ impl<'a> FunctionTranslator<'a> { /// rather than naming a function declaration. Such calls go through /// `translate_closure_call` instead of the direct-call path. fn holds_fat_pointer(&self, fn_id: ExprID) -> bool { - let Expr::Id(name) = &self.decl.arena.exprs[fn_id] else { - return false; - }; - if self.variables.contains_key(name) { - return true; + match self.decl.arena.reference(fn_id) { + Some(Reference::Local(_)) => true, + Some(Reference::Instance(id)) => { + matches!(self.decls.instance(*id), Decl::Global { .. }) + } + _ => false, } - // Extern functions live in globals memory too, but they are called - // through the direct-call path. - self.globals.contains_key(name) - && matches!(&*self.expr_type(fn_id), Type::Func(_, _)) - && !self - .decls - .find(*name) - .iter() - .any(|d| matches!(d, Decl::Func(f) if f.is_extern)) } /// Parameter types of a callee reached through a fat pointer, taken from @@ -2603,7 +2514,7 @@ impl<'a> FunctionTranslator<'a> { /// Translate a for loop. fn translate_for( &mut self, - var: Name, + var: LocalId, start_id: ExprID, end_id: ExprID, body_id: ExprID, @@ -2621,13 +2532,10 @@ impl<'a> FunctionTranslator<'a> { let end_local = self.alloc_scalar(); func.emit(StackOp::LocalSet(end_local)); - // The counter is a scalar bound to `var` for the duration of the loop. - // Shadowing an outer binding has to forget the outer binding's - // name-keyed state, and the snapshot brings it back at loop exit, where - // the loop variable is out of scope again. - let saved = self.save_bindings(); + // The checked loop binding has its own local identity. + let int_ty = mk_type(Type::Int32); - self.shadow_outer_binding(&var); + let counter_mem = if self.lambda_referenced.contains(&var) { // A lambda shares the counter by address, so it needs memory of // its own, allocated up front the way `let` and `var` do it. @@ -2642,7 +2550,6 @@ impl<'a> FunctionTranslator<'a> { self.variables.insert(var, LocalKind::Scalar(loop_var)); None }; - self.variable_types.insert(var, int_ty); let loop_start = func.pos(); @@ -2672,7 +2579,6 @@ impl<'a> FunctionTranslator<'a> { // Execute body in void context. self.translate_void(body_id, func); - self.restore_bindings(saved); // Increment position (continue target). let increment_pos = func.pos(); @@ -2989,122 +2895,63 @@ impl<'a> FunctionTranslator<'a> { } /// Translate a lambda expression. - fn translate_lambda( - &mut self, - params: &[Param], - body: ExprID, - expr: ExprID, - func: &mut StackFunction, - ) { - let lambda_ty = self.expr_type(expr); - if let Type::Func(dom, rng) = &*lambda_ty { - if let Type::Tuple(param_types) = &**dom { - let id = *self.lambda_counter; - *self.lambda_counter += 1; - let lambda_name = Name::new(format!("__lambda_{}", id)); - - let lambda_params: Vec = params - .iter() - .zip(param_types.iter()) - .map(|(p, ty)| Param { - name: p.name, - ty: Some(*ty), - }) - .collect(); - - // Compute free variables captured from the enclosing scope. - let param_names: HashSet = - params.iter().map(|p| p.name.to_string()).collect(); - let free_vars = collect_free_var_names( - body, - &self.decl.arena, - ¶m_names, - &self.variables, - &self.decl.types, - ); - - // Build closure struct if there are captures. - let has_captures = !free_vars.is_empty(); - let closure_mem_slot = if has_captures { - let n = free_vars.len(); - let slot = self.alloc_memory((n * 8) as u32); - for (i, (name, _ty)) in free_vars.iter().enumerate() { - let var_name = Name::new(name.clone()); - func.emit(StackOp::LocalAddr(slot)); - self.emit_var_address(&var_name, func); - func.emit(StackOp::Store64Off((i * 8) as i32)); - } - Some(slot) - } else { - None - }; - - let closure_vars: Vec = free_vars - .iter() - .map(|(name, ty)| ClosureVar { - name: Name::new(name.clone()), - ty: *ty, - }) - .collect(); - - let lambda_decl = FuncDecl { - name: lambda_name, - typevars: vec![], - size_vars: vec![], - params: lambda_params, - body: Some(body), - ret: *rng, - constraints: vec![], - requires: vec![], - loc: self.decl.loc, - arena: self.decl.arena.clone(), - types: self.decl.types.clone(), - closure_vars, - is_extern: false, - }; - - self.pending_lambdas.push(lambda_decl); - - // Build fat pointer {func_idx, closure_ptr}. - let fat_slot = self.alloc_memory(16); - // Store func_idx. - func.emit(StackOp::LocalAddr(fat_slot)); - let instr_idx = func.pos(); - func.emit(StackOp::I64Const(0)); // placeholder - self.lambda_patches.push((instr_idx, lambda_name)); - func.emit(StackOp::Store64); - // Store closure_ptr. - func.emit(StackOp::LocalAddr(fat_slot)); - if let Some(closure_slot) = closure_mem_slot { - func.emit(StackOp::LocalAddr(closure_slot)); - } else { - func.emit(StackOp::I64Const(0)); - } - func.emit(StackOp::Store64Off(8)); - // Push fat pointer address. - func.emit(StackOp::LocalAddr(fat_slot)); - } else { - panic!( - "stack codegen lambda: expected tuple domain type, got {:?}", - dom - ); + fn translate_lambda(&mut self, expr: ExprID, func: &mut StackFunction) { + let id = *self.lambda_counter; + *self.lambda_counter += 1; + let lambda_name = Name::new(format!("__lambda_{}", id)); + + let lambda_decl = self.decl.extract_lambda(expr, lambda_name); + let free_vars = &lambda_decl.closure_vars; + + // Build closure struct if there are captures. + let has_captures = !free_vars.is_empty(); + let closure_mem_slot = if has_captures { + let n = free_vars.len(); + let slot = self.alloc_memory((n * 8) as u32); + for (i, var_name) in free_vars.iter().enumerate() { + func.emit(StackOp::LocalAddr(slot)); + self.emit_var_address(var_name, func); + func.emit(StackOp::Store64Off((i * 8) as i32)); } + Some(slot) + } else { + None + }; + + self.pending_lambdas.push(lambda_decl); + + // Build fat pointer {func_idx, closure_ptr}. + let fat_slot = self.alloc_memory(16); + // Store func_idx. + func.emit(StackOp::LocalAddr(fat_slot)); + let instr_idx = func.pos(); + func.emit(StackOp::I64Const(0)); // placeholder + self.lambda_patches.push((instr_idx, lambda_name)); + func.emit(StackOp::Store64); + // Store closure_ptr. + func.emit(StackOp::LocalAddr(fat_slot)); + if let Some(closure_slot) = closure_mem_slot { + func.emit(StackOp::LocalAddr(closure_slot)); } else { - panic!( - "stack codegen lambda: expected function type, got {:?}", - lambda_ty - ); + func.emit(StackOp::I64Const(0)); } + func.emit(StackOp::Store64Off(8)); + // Push fat pointer address. + func.emit(StackOp::LocalAddr(fat_slot)); } /// Get the address of a variable for closure capture. - fn emit_var_address(&mut self, name: &Name, func: &mut StackFunction) { + fn emit_var_address(&mut self, name: &LocalId, func: &mut StackFunction) { if self.captured_vars.contains(name) { // Already captured from an enclosing scope: follow indirection. let addr_local = *self.captured_slots.get(name).unwrap(); func.emit(StackOp::LocalGet(addr_local)); } else if let Some(&kind) = self.variables.get(name) { match kind { + LocalKind::Scalar(slot) if self.is_ptr_type(&self.decl.arena.local(*name).ty) => { + // Aggregate and fat-pointer values already hold their storage address. + func.emit(StackOp::LocalGet(slot)); + } LocalKind::Scalar(slot) => { // Scalar: need to spill to memory so we have a stable address. let mem_slot = self.alloc_memory(8); @@ -3123,7 +2970,7 @@ impl<'a> FunctionTranslator<'a> { } } } else { - func.emit(StackOp::I64Const(0)); + unreachable!("checked capture local {:?} has no storage", name); } } @@ -3290,147 +3137,3 @@ impl<'a> FunctionTranslator<'a> { } } } - -/// Collect free variable names referenced in a lambda body that come from the enclosing scope. -fn collect_free_var_names( - body: ExprID, - arena: &ExprArena, - exclude: &HashSet, - local_vars: &HashMap, - types: &[TypeID], -) -> Vec<(String, TypeID)> { - let mut result = Vec::new(); - let mut seen = HashSet::new(); - collect_free_vars_rec( - body, - arena, - exclude, - local_vars, - types, - &mut result, - &mut seen, - ); - result -} - -fn collect_free_vars_rec( - expr: ExprID, - arena: &ExprArena, - exclude: &HashSet, - local_vars: &HashMap, - types: &[TypeID], - result: &mut Vec<(String, TypeID)>, - seen: &mut HashSet, -) { - match &arena[expr] { - Expr::TypeApp(_, _) => {} - Expr::Id(name) => { - let s = name.to_string(); - if local_vars.contains_key(name) && !exclude.contains(&s) && !seen.contains(&s) { - result.push((s.clone(), types[expr])); - seen.insert(s); - } - } - Expr::Call(fn_id, args) => { - collect_free_vars_rec(*fn_id, arena, exclude, local_vars, types, result, seen); - for a in args { - collect_free_vars_rec(*a, arena, exclude, local_vars, types, result, seen); - } - } - Expr::Binop(_, lhs, rhs) => { - collect_free_vars_rec(*lhs, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec(*rhs, arena, exclude, local_vars, types, result, seen); - } - Expr::Unop(_, arg) => { - collect_free_vars_rec(*arg, arena, exclude, local_vars, types, result, seen); - } - Expr::Let(_, init, _) => { - collect_free_vars_rec(*init, arena, exclude, local_vars, types, result, seen); - } - Expr::Var(_, init, _) => { - if let Some(init_id) = init { - collect_free_vars_rec(*init_id, arena, exclude, local_vars, types, result, seen); - } - } - Expr::If(cond, then, else_) => { - collect_free_vars_rec(*cond, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec(*then, arena, exclude, local_vars, types, result, seen); - if let Some(e) = else_ { - collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen); - } - } - Expr::While(cond, body) => { - collect_free_vars_rec(*cond, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec(*body, arena, exclude, local_vars, types, result, seen); - } - Expr::For { - start, end, body, .. - } => { - collect_free_vars_rec(*start, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec(*end, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec(*body, arena, exclude, local_vars, types, result, seen); - } - Expr::Block(exprs) => { - for e in exprs { - collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen); - } - } - Expr::Return(e) | Expr::Assume(e) => { - collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen); - } - Expr::Field(e, _) => { - collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen); - } - Expr::ArrayIndex(arr, idx) => { - collect_free_vars_rec(*arr, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec(*idx, arena, exclude, local_vars, types, result, seen); - } - Expr::ArrayLiteral(elems) | Expr::Tuple(elems) => { - for e in elems { - collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen); - } - } - Expr::AsTy(e, _) | Expr::Arena(e) => { - collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen); - } - Expr::Array(ty_expr, size_expr) => { - collect_free_vars_rec(*ty_expr, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec(*size_expr, arena, exclude, local_vars, types, result, seen); - } - Expr::Lambda { params, body } => { - let mut inner_exclude = exclude.clone(); - for p in params { - inner_exclude.insert(p.name.to_string()); - } - collect_free_vars_rec( - *body, - arena, - &inner_exclude, - local_vars, - types, - result, - seen, - ); - } - Expr::Macro(_, args) => { - for a in args { - collect_free_vars_rec(*a, arena, exclude, local_vars, types, result, seen); - } - } - Expr::StructLit(_, fields) => { - for (_, fval) in fields { - collect_free_vars_rec(*fval, arena, exclude, local_vars, types, result, seen); - } - } - Expr::Int(_, _) - | Expr::Real(_, _) - | Expr::String(_) - | Expr::Char(_) - | Expr::True - | Expr::False - | Expr::Enum(_) - | Expr::Break - | Expr::Continue - | Expr::Error => {} - } -} diff --git a/src/types.rs b/src/types.rs index b8a36694..d7c30bfd 100644 --- a/src/types.rs +++ b/src/types.rs @@ -232,7 +232,7 @@ impl TypeID { } /// Returns the size of a type in bytes. - pub fn size(self, decls: &DeclTable) -> i32 { + pub fn size(self, decls: &DeclarationList) -> i32 { match &*self { Type::Void => 0, Type::Bool => 1, @@ -550,12 +550,12 @@ pub fn unify_with_vars(lhs: TypeID, rhs: TypeID, inst: &mut Instance) -> bool { } } -impl Decl { +impl Decl { pub fn ty(&self) -> TypeID { match self { Decl::Interface { .. } => mk_type(Type::Void), - Decl::Func(FuncDecl { params, ret, .. }) => func(params_ty(params), *ret), - Decl::Macro(FuncDecl { params, ret, .. }) => func(params_ty(params), *ret), + Decl::Func(function) => function.ty(), + Decl::Macro(function) => function.ty(), Decl::Struct(StructDecl { name, .. }) => mk_type(Type::Name(*name, vec![])), Decl::Enum { name, .. } => mk_type(Type::Name(*name, vec![])), Decl::Global { ty, .. } => *ty, @@ -566,6 +566,17 @@ impl Decl { } impl FuncDecl { + /// Source signatures can be incomplete during editing. Do not substitute a + /// made-up type for a missing parameter annotation. + pub fn annotated_ty(&self) -> Option { + let params = self + .params + .iter() + .map(|param| param.ty) + .collect::>>()?; + Some(func(tuple(params), self.ret)) + } + /// Returns the type for this function declaration. pub fn ty(&self) -> TypeID { func(params_ty(&self.params), self.ret) @@ -573,53 +584,6 @@ impl FuncDecl { } impl Interface { - /// Is an interface satisfied? - pub fn satisfied( - &self, - types: &[TypeID], - decls: &DeclTable, - errors: &mut Vec, - loc: Loc, - ) -> bool { - let mut inst = Instance::new(); - for (v, t) in self.typevars.iter().zip(types) { - inst.insert(typevar(v), *t); - } - - // If any type parameter is still unresolved (type variable or anonymous), - // defer the check — it will be verified when the generic function is - // instantiated with concrete types. - if types - .iter() - .any(|t| matches!(&**t, Type::Var(_) | Type::Anon(_))) - { - return true; - } - - let mut satisfied = true; - - // Find functions among decls that have the same type. - for func in &self.funcs { - let d = decls.find(func.name); - - // Do we want to unify instead? - let found = d.iter().any(|d| d.ty() == func.ty().subst(&inst)); - - if !found { - satisfied = false; - errors.push(TypeError { - location: loc, - message: format!( - "function {} for interface {} is required", - func.name, self.name - ), - }); - } - } - - satisfied - } - /// Replaces any named types with type variables. pub fn subst_typevars(&mut self) { let mut inst = Instance::new(); diff --git a/src/vm.rs b/src/vm.rs index 807ca9f5..2e90caf0 100644 --- a/src/vm.rs +++ b/src/vm.rs @@ -1631,16 +1631,16 @@ impl VM { } // SliceEq/SliceNe — compare slice contents by value - // Fat pointer layout: data_ptr (8 bytes) + len (4 bytes) + // Fat pointer layout: data_ptr (8 bytes) + len (4 bytes), without padding. tags::SLICE_EQ => { let fat_a = r!(op.b()) as *const u8; let fat_b = r!(op.c()) as *const u8; let elem_size = (*ops.add(ip)).0 as usize; ip += 1; - let ptr_a = *(fat_a as *const u64) as *const u8; - let len_a = *(fat_a.add(8) as *const u32) as usize; - let ptr_b = *(fat_b as *const u64) as *const u8; - let len_b = *(fat_b.add(8) as *const u32) as usize; + let ptr_a = (fat_a as *const u64).read_unaligned() as *const u8; + let len_a = (fat_a.add(8) as *const u32).read_unaligned() as usize; + let ptr_b = (fat_b as *const u64).read_unaligned() as *const u8; + let len_b = (fat_b.add(8) as *const u32).read_unaligned() as usize; let eq = len_a == len_b && std::slice::from_raw_parts(ptr_a, len_a * elem_size) == std::slice::from_raw_parts(ptr_b, len_b * elem_size); @@ -1652,10 +1652,10 @@ impl VM { let fat_b = r!(op.c()) as *const u8; let elem_size = (*ops.add(ip)).0 as usize; ip += 1; - let ptr_a = *(fat_a as *const u64) as *const u8; - let len_a = *(fat_a.add(8) as *const u32) as usize; - let ptr_b = *(fat_b as *const u64) as *const u8; - let len_b = *(fat_b.add(8) as *const u32) as usize; + let ptr_a = (fat_a as *const u64).read_unaligned() as *const u8; + let len_a = (fat_a.add(8) as *const u32).read_unaligned() as usize; + let ptr_b = (fat_b as *const u64).read_unaligned() as *const u8; + let len_b = (fat_b.add(8) as *const u32).read_unaligned() as usize; let ne = len_a != len_b || std::slice::from_raw_parts(ptr_a, len_a * elem_size) != std::slice::from_raw_parts(ptr_b, len_b * elem_size); @@ -1666,17 +1666,17 @@ impl VM { tags::SLICE_LOAD32 => { // B = slice fat pointer, C = index let fat_ptr = r!(op.b()) as *const u8; - let data_ptr = *(fat_ptr as *const *const u8); + let data_ptr = (fat_ptr as *const *const u8).read_unaligned(); let idx = r!(op.c()) as usize; - let elem = *(data_ptr.add(idx * 4) as *const i32); + let elem = (data_ptr.add(idx * 4) as *const i32).read_unaligned(); r_set!(op.a(), elem as i64 as u64); } tags::SLICE_STORE32 => { // A = src value, B = slice fat pointer, C = index let fat_ptr = r!(op.b()) as *const u8; - let data_ptr = *(fat_ptr as *const *mut u8); + let data_ptr = (fat_ptr as *const *mut u8).read_unaligned(); let idx = r!(op.c()) as usize; - *(data_ptr.add(idx * 4) as *mut i32) = r!(op.a()) as i32; + (data_ptr.add(idx * 4) as *mut i32).write_unaligned(r!(op.a()) as i32); } // Type conversions — AB @@ -1708,7 +1708,8 @@ impl VM { set_i64!(op.a(), (get_i64!(op.b()) as u32) as i64); } - // Memory operations + // Language storage uses packed byte offsets, including pointer headers. + // Indirect accesses therefore cannot assume native integer alignment. tags::LOAD8 => { let ptr = r!(op.b()); self.check_ptr(ptr, 1); @@ -1717,36 +1718,36 @@ impl VM { tags::LOAD32 => { let ptr = r!(op.b()); self.check_ptr(ptr, 4); - set_i64!(op.a(), *(ptr as *const i32) as i64); + set_i64!(op.a(), (ptr as *const i32).read_unaligned() as i64); } tags::LOAD64 => { let ptr = r!(op.b()); self.check_ptr(ptr, 8); - set_i64!(op.a(), *(ptr as *const i64)); + set_i64!(op.a(), (ptr as *const i64).read_unaligned()); } tags::LOAD32_OFF => { let ptr = r!(op.b()).wrapping_add(op.c() as u64); self.check_ptr(ptr, 4); - set_i64!(op.a(), *(ptr as *const i32) as i64); + set_i64!(op.a(), (ptr as *const i32).read_unaligned() as i64); } tags::LOAD32_OFF_WIDE => { let off = (*ops.add(ip)).0 as i64; ip += 1; let ptr = (r!(op.b()) as i64 + off) as u64; self.check_ptr(ptr, 4); - set_i64!(op.a(), *(ptr as *const i32) as i64); + set_i64!(op.a(), (ptr as *const i32).read_unaligned() as i64); } tags::LOAD64_OFF => { let ptr = r!(op.b()).wrapping_add(op.c() as u64); self.check_ptr(ptr, 8); - set_i64!(op.a(), *(ptr as *const i64)); + set_i64!(op.a(), (ptr as *const i64).read_unaligned()); } tags::LOAD64_OFF_WIDE => { let off = (*ops.add(ip)).0 as i64; ip += 1; let ptr = (r!(op.b()) as i64 + off) as u64; self.check_ptr(ptr, 8); - set_i64!(op.a(), *(ptr as *const i64)); + set_i64!(op.a(), (ptr as *const i64).read_unaligned()); } tags::STORE8 => { let ptr = r!(op.a()); @@ -1756,12 +1757,12 @@ impl VM { tags::STORE32 => { let ptr = r!(op.a()); self.check_ptr(ptr, 4); - *(ptr as *mut i32) = get_i64!(op.b()) as i32; + (ptr as *mut i32).write_unaligned(get_i64!(op.b()) as i32); } tags::STORE64 => { let ptr = r!(op.a()); self.check_ptr(ptr, 8); - *(ptr as *mut i64) = get_i64!(op.b()); + (ptr as *mut i64).write_unaligned(get_i64!(op.b())); } tags::STORE8_OFF => { let ptr = r!(op.a()).wrapping_add(op.c() as u64); @@ -1778,26 +1779,26 @@ impl VM { tags::STORE32_OFF => { let ptr = r!(op.a()).wrapping_add(op.c() as u64); self.check_ptr(ptr, 4); - *(ptr as *mut i32) = get_i64!(op.b()) as i32; + (ptr as *mut i32).write_unaligned(get_i64!(op.b()) as i32); } tags::STORE32_OFF_WIDE => { let off = (*ops.add(ip)).0 as i64; ip += 1; let ptr = (r!(op.a()) as i64 + off) as u64; self.check_ptr(ptr, 4); - *(ptr as *mut i32) = get_i64!(op.b()) as i32; + (ptr as *mut i32).write_unaligned(get_i64!(op.b()) as i32); } tags::STORE64_OFF => { let ptr = r!(op.a()).wrapping_add(op.c() as u64); self.check_ptr(ptr, 8); - *(ptr as *mut i64) = get_i64!(op.b()); + (ptr as *mut i64).write_unaligned(get_i64!(op.b())); } tags::STORE64_OFF_WIDE => { let off = (*ops.add(ip)).0 as i64; ip += 1; let ptr = (r!(op.a()) as i64 + off) as u64; self.check_ptr(ptr, 8); - *(ptr as *mut i64) = get_i64!(op.b()); + (ptr as *mut i64).write_unaligned(get_i64!(op.b())); } tags::LOCAL_ADDR => { @@ -2078,8 +2079,8 @@ impl VM { // fat_ptr points to {func_idx: i64, closure_ptr: i64} tags::CALL_CLOSURE => { let fat_ptr = r!(op.a()) as *const u64; - let func_idx = *fat_ptr as FuncIdx; - let closure_ptr_val = *fat_ptr.add(1); + let func_idx = fat_ptr.read_unaligned() as FuncIdx; + let closure_ptr_val = fat_ptr.add(1).read_unaligned(); let args_start = op.b() as usize; let arg_count = op.c() as usize; @@ -2130,8 +2131,8 @@ impl VM { // Read fn_ptr and context from globals buffer. let slot = self.globals.as_ptr().add(globals_offset as usize) as *const u64; - let fn_ptr = *slot as usize; - let context = *slot.add(1) as *mut u8; + let fn_ptr = slot.read_unaligned() as usize; + let context = slot.add(1).read_unaligned() as *mut u8; if fn_ptr == 0 { panic!( diff --git a/src/vm_codegen.rs b/src/vm_codegen.rs index e552778f..46287345 100644 --- a/src/vm_codegen.rs +++ b/src/vm_codegen.rs @@ -1,14 +1,16 @@ //! VM code generator. //! -//! This module translates a DeclTable into a VMProgram that can be +//! This module translates a SpecializedProgram into a VMProgram that can be //! executed by the register-based virtual machine. -use crate::decl::*; +use crate::checked::{ + CheckedBody, CheckedFunction, InstanceId, LocalId, Reference, SpecializedProgram, +}; +use crate::decl::Decl; use crate::defs::*; -use crate::expr::*; +use crate::expr::Expr; use crate::types::*; use crate::vm::*; -use crate::DeclTable; use std::collections::{HashMap, HashSet}; /// Loop context for break/continue support. @@ -29,8 +31,8 @@ struct PendingCall { func_idx: FuncIdx, /// Index of the Call instruction within that function. instr_idx: usize, - /// Name of the function being called. - callee: Name, + /// Concrete function instance being called. + callee: InstanceId, } /// Code generator for the VM. @@ -38,14 +40,14 @@ pub struct VMCodegen { /// The program being built. program: VMProgram, - /// Map from function names to their indices in the program. - func_indices: HashMap, + /// Map from function instances to their indices in the program. + func_indices: HashMap, /// Functions that have been compiled. - compiled_functions: HashSet, + compiled_functions: HashSet, /// Functions that need to be compiled. - pending_functions: Vec, + pending_functions: Vec, /// Calls that need to be patched after all functions are compiled. pending_calls: Vec, @@ -54,7 +56,7 @@ pub struct VMCodegen { pending_func_loads: Vec, /// Global variable offsets. - globals: HashMap, + globals: HashMap, /// Counter for generating unique lambda names. lambda_counter: usize, @@ -81,34 +83,25 @@ impl VMCodegen { } /// Collect global variables and compute their offsets. - fn declare_globals(&mut self, decls: &DeclTable) { + fn declare_globals(&mut self, decls: &SpecializedProgram) { let mut offset: i32 = 0; - for decl in &decls.decls { - match decl { - Decl::Global { - name, typevars, ty, .. - } => { - if typevars.is_empty() { - self.globals.insert(*name, offset); - offset += ty.size(decls) as i32; - } - } - Decl::Func(f) if f.is_extern => { - // Extern functions get 16 bytes: {fn_ptr, context} - self.globals.insert(f.name, offset); - offset += 16; - } - _ => {} - } + for (instance, decl) in decls.storage_instances() { + let size = match decl { + Decl::Global { ty, .. } => ty.size(decls) as i32, + Decl::Func(f) if f.is_extern => 16, + _ => continue, + }; + self.globals.insert(instance, offset); + offset += size; } self.program.globals_size = offset as usize; } - /// Compile a DeclTable into a VMProgram. + /// Compile a SpecializedProgram into a VMProgram. /// /// This looks for a "main" function and compiles it along with all /// functions it calls. - pub fn compile(&mut self, decls: &DeclTable) -> Result { + pub fn compile(&mut self, decls: &SpecializedProgram) -> Result { let main_name = Name::str("main"); self.compile_multi(decls, &[main_name]) } @@ -119,7 +112,7 @@ impl VMCodegen { /// found show up in `program.entry_points`. pub fn compile_multi( &mut self, - decls: &DeclTable, + decls: &SpecializedProgram, entry_points: &[Name], ) -> Result { // First, collect all global variables. @@ -127,26 +120,26 @@ impl VMCodegen { // Compile each entry point root. for &ep_name in entry_points { - if self.compiled_functions.contains(&ep_name) { - continue; - } - let Some(ep_decl) = decls.find_entry_point(ep_name) else { + let Some(instance) = decls.instance_for_entry(ep_name) else { continue; }; - self.compile_function(ep_decl, decls)?; + if self.compiled_functions.contains(&instance) { + continue; + } + let ep_decl = decls + .function_instance(instance) + .expect("entry is a function"); + self.compile_function(ep_decl, decls, Some(instance))?; // Compile any pending functions (called by this entry point or transitively). while let Some(name) = self.pending_functions.pop() { if self.compiled_functions.contains(&name) { continue; } - let func_decls = decls.find(name); - if func_decls.is_empty() { - continue; - } - if let Decl::Func(func_decl) = &func_decls[0] { - self.compile_function(func_decl, decls)?; - } + let func_decl = decls + .function_instance(name) + .expect("call target is a function"); + self.compile_function(func_decl, decls, Some(name))?; } } @@ -155,7 +148,10 @@ impl VMCodegen { // found, program.entry stays at its default and the map is empty. let mut entry_set = false; for &ep_name in entry_points { - if let Some(&idx) = self.func_indices.get(&ep_name) { + if let Some(&idx) = decls + .instance_for_entry(ep_name) + .and_then(|id| self.func_indices.get(&id)) + { self.program.entry_points.insert(ep_name, idx); if !entry_set { self.program.entry = idx; @@ -196,7 +192,12 @@ impl VMCodegen { } /// Compile a single function. - fn compile_function(&mut self, decl: &FuncDecl, decls: &DeclTable) -> Result { + fn compile_function( + &mut self, + decl: &CheckedFunction, + decls: &SpecializedProgram, + instance: Option, + ) -> Result { let mut func = VMFunction::new(&*decl.name); func.param_count = decl.params.len() as u8; @@ -210,13 +211,15 @@ impl VMCodegen { translator.translate(&mut func); // Extract debug info: register and slot names for disassembly. - for (&name, ®) in &translator.variables { - if translator.reg_promoted.contains(&name) { - func.reg_names.push((reg, format!("{}", name))); + for (&name, ®) in &translator.body.variables { + if translator.body.reg_promoted.contains(&name) { + func.reg_names + .push((reg, format!("{}", decl.arena.local(name).name))); } } - for (&name, &slot) in &translator.local_slots { - func.slot_names.push((slot, format!("{}", name))); + for (&name, &slot) in &translator.body.local_slots { + func.slot_names + .push((slot, format!("{}", decl.arena.local(name).name))); } // Peephole optimize: eliminate redundant instructions + register allocation. @@ -226,23 +229,21 @@ impl VMCodegen { // Register allocation compacted the register numbering. // Update locals_size: slot area stays the same, register save area shrinks. func.locals_size = func.local_slots as u32 * 8 + new_reg_count as u32 * 8; - // Update debug register names with the new physical register numbers. - for entry in &mut func.reg_names { - let idx = entry.0 as usize; - if idx < mapping.len() { - let preg = mapping[idx]; - if preg != Reg::MAX { - entry.0 = preg; - } - } + // Debug names follow only registers the allocator retained. + // An eliminated virtual register must never label a reused physical one. + for (register, _) in &mut func.reg_names { + *register = mapping.get(*register as usize).copied().unwrap_or(Reg::MAX); } - // Remove entries where the register was optimized away. - func.reg_names.retain(|&(reg, _)| reg != Reg::MAX); + func.reg_names.retain(|&(register, _)| register != Reg::MAX); + func.reg_names.sort(); + func.reg_names.dedup(); } let idx = self.program.add_function(func); - self.func_indices.insert(decl.name, idx); - self.compiled_functions.insert(decl.name); + if let Some(instance) = instance { + self.func_indices.insert(instance, idx); + self.compiled_functions.insert(instance); + } // Extract data from translator before it's dropped (it borrows self.pending_functions). let calls_to_patch = std::mem::take(&mut translator.calls_to_patch); @@ -276,7 +277,7 @@ impl VMCodegen { // Compile lambda functions and patch their indices into the parent function. for lambda_decl in pending_lambdas { let lambda_name = lambda_decl.name; - let lambda_idx = self.compile_function(&lambda_decl, decls)?; + let lambda_idx = self.compile_function(&lambda_decl, decls, None)?; // Patch every LoadImm placeholder for this lambda. for &(instr_idx, patch_name) in &lambda_patches { if patch_name == lambda_name { @@ -291,25 +292,12 @@ impl VMCodegen { Ok(idx) } - - /// Get or create a function index for a given name. - fn get_or_create_func_index(&mut self, name: Name) -> FuncIdx { - if let Some(&idx) = self.func_indices.get(&name) { - idx - } else { - // Create a placeholder index - the function will be compiled later. - let idx = self.program.functions.len() as FuncIdx; - self.func_indices.insert(name, idx); - self.pending_functions.push(name); - idx - } - } } /// Check if an expression is simple enough to inline: no blocks, control flow, /// let/var bindings, or nested calls. Only arithmetic, comparisons, literals, /// identifiers, type casts, and field access are allowed. -fn is_inline_expr(id: ExprID, arena: &ExprArena) -> bool { +fn is_inline_expr(id: ExprID, arena: &CheckedBody) -> bool { match &arena[id] { Expr::Int(_, _) | Expr::Real(_, _) | Expr::Id(_) | Expr::True | Expr::False => true, Expr::Binop(op, lhs, rhs) => { @@ -332,57 +320,48 @@ fn is_inline_expr(id: ExprID, arena: &ExprArena) -> bool { } } -/// A snapshot of a translator's name-keyed binding state. See -/// `FunctionTranslator::save_bindings`. -struct SavedBindings { - variables: HashMap, - variable_types: HashMap, - local_slots: HashMap, - reg_promoted: HashSet, - reference_vars: HashSet, - captured_vars: HashSet, -} - /// A call instruction that needs patching. struct CallToPatch { /// Index of the Call instruction. instr_idx: usize, - /// Name of the function being called. - callee: Name, + /// Concrete function instance being called. + callee: InstanceId, +} + +/// Storage and derived facts whose IDs belong to one checked body. Inlining +/// switches this context as a unit; lexical scopes need no storage snapshots. +struct BodyContext<'a> { + decl: &'a CheckedFunction, + variables: HashMap, + reg_promoted: HashSet, + lambda_referenced: HashSet, + reference_vars: HashSet, + local_slots: HashMap, + elidable_lets: HashSet, + captured_vars: HashSet, +} + +impl<'a> BodyContext<'a> { + fn new(decl: &'a CheckedFunction) -> Self { + Self { + decl, + variables: HashMap::new(), + reg_promoted: HashSet::new(), + lambda_referenced: decl.captured_locals(), + reference_vars: HashSet::new(), + local_slots: HashMap::new(), + elidable_lets: crate::copy_elision::elidable_let_copies(decl), + captured_vars: HashSet::new(), + } + } } /// Translator for a single function body. struct FunctionTranslator<'a> { - /// The function declaration being translated. - decl: &'a FuncDecl, + body: BodyContext<'a>, /// Declaration table for looking up types and functions. - decls: &'a DeclTable, - - /// Map from variable names to their register numbers. - /// For pointer types, the register holds the memory address. - /// For register-promoted scalars, the register holds the value directly. - variables: HashMap, - - /// Declared representation type for local bindings. - variable_types: HashMap, - - /// Set of variable names that are register-promoted (value in register, not memory). - reg_promoted: HashSet, - - /// Names a lambda in this function mentions. These are shared with the - /// closure by address, so they must never be register-promoted. - lambda_referenced: HashSet, - - /// Set of variable names whose register stores a reference address. - reference_vars: HashSet, - - /// Map from variable names to their local slot indices (for addressable vars). - local_slots: HashMap, - - /// `Expr::Let` ids whose value-copy is unobservable, so the binding can - /// alias the initializer's storage instead. See `crate::copy_elision`. - elidable_lets: HashSet, + decls: &'a SpecializedProgram, /// Next available register. next_reg: Reg, @@ -394,22 +373,16 @@ struct FunctionTranslator<'a> { locals_size: u32, /// Functions that are called and need to be compiled. - pending_functions: &'a mut Vec, + pending_functions: &'a mut Vec, /// Counter for generating unique lambda names. lambda_counter: &'a mut usize, - /// Map from function name to expected function index. - called_functions: HashMap, - - /// Current next function index for placeholders. - next_func_idx: FuncIdx, - /// Calls that need patching. calls_to_patch: Vec, - /// Lambda FuncDecls extracted from this function body, to be compiled afterward. - pending_lambdas: Vec, + /// Lambda CheckedFunctions extracted from this function body, to be compiled afterward. + pending_lambdas: Vec, /// LoadImm instructions that need to be patched with lambda function indices. lambda_patches: Vec<(usize, Name)>, @@ -418,7 +391,7 @@ struct FunctionTranslator<'a> { func_load_patches: Vec, /// Global variable offsets. - globals: &'a HashMap, + globals: &'a HashMap, /// Output pointer register for functions returning pointer types. output_ptr: Option, @@ -438,9 +411,6 @@ struct FunctionTranslator<'a> { /// Byte offset in locals where registers are saved. save_regs_offset: u32, - /// Variables captured from an enclosing scope (accessed via double indirection). - captured_vars: HashSet, - /// Extern function info collected during translation. extern_funcs: Vec, } @@ -473,29 +443,21 @@ fn returns_via_pointer(ty: TypeID) -> bool { impl<'a> FunctionTranslator<'a> { fn new( - decl: &'a FuncDecl, - decls: &'a DeclTable, - pending_functions: &'a mut Vec, + decl: &'a CheckedFunction, + decls: &'a SpecializedProgram, + pending_functions: &'a mut Vec, lambda_counter: &'a mut usize, - globals: &'a HashMap, + globals: &'a HashMap, ) -> Self { Self { - decl, + body: BodyContext::new(decl), decls, - variables: HashMap::new(), - variable_types: HashMap::new(), - reg_promoted: HashSet::new(), - lambda_referenced: decl.names_referenced_in_lambdas(), - reference_vars: HashSet::new(), - local_slots: HashMap::new(), - elidable_lets: crate::copy_elision::elidable_let_copies(decl), + next_reg: 0, next_slot: 0, locals_size: 0, pending_functions, lambda_counter, - called_functions: HashMap::new(), - next_func_idx: 0, calls_to_patch: Vec::new(), pending_lambdas: Vec::new(), lambda_patches: Vec::new(), @@ -505,7 +467,6 @@ impl<'a> FunctionTranslator<'a> { output_ptr_slot: None, has_returned: false, save_regs_offset: 0, - captured_vars: HashSet::new(), loop_stack: Vec::new(), extern_funcs: Vec::new(), } @@ -514,13 +475,13 @@ impl<'a> FunctionTranslator<'a> { /// Translate the function body. fn translate(&mut self, func: &mut VMFunction) { // If return type is a pointer type, first parameter is output pointer. - if returns_via_pointer(self.decl.ret) { + if returns_via_pointer(self.body.decl.ret) { self.output_ptr = Some(self.alloc_reg()); func.param_count += 1; } // Reserve registers for parameters (these are the incoming argument positions). - let param_count = self.decl.params.len(); + let param_count = self.body.decl.params.len(); for _ in 0..param_count { self.alloc_reg(); } @@ -553,24 +514,28 @@ impl<'a> FunctionTranslator<'a> { // Handle parameters. Scalars stay in registers (SaveRegs preserves them // across calls). Pointer types are stored to local slots. - let param_offset = if returns_via_pointer(self.decl.ret) { + let param_offset = if returns_via_pointer(self.body.decl.ret) { 1u8 } else { 0u8 }; - for (i, param) in self.decl.params.iter().enumerate() { + for (i, param) in self.body.decl.params.iter().enumerate() { let src_reg = i as Reg + param_offset as Reg; - let ty = param.ty.expect("parameter must have type"); + let ty = self.body.decl.arena.local(param.local).ty; - if let Type::Reference(inner) = &*ty { + if let Type::Reference(_) = &*ty { let reg = self.alloc_reg(); func.emit(Opcode::Move { dst: reg, src: src_reg, }); - self.variables.insert(param.name, reg); - self.variable_types.insert(param.name, *inner); - self.reference_vars.insert(param.name); + self.body.variables.insert(param.local, reg); + + self.body.reference_vars.insert(param.local); + } else if !self.is_ptr_type(&ty) && self.body.lambda_referenced.contains(¶m.local) { + // A conditional capture must not decide whether parameter storage exists. + let addr = self.alloc_scalar_slot(param.local, ty, func); + self.emit_store(&ty, addr, src_reg, func); } else if !self.is_ptr_type(&ty) { // Scalar parameter: copy to a dedicated register so the param // register can be reused. The copy will be eliminated by @@ -580,9 +545,9 @@ impl<'a> FunctionTranslator<'a> { dst: reg, src: src_reg, }); - self.variables.insert(param.name, reg); - self.variable_types.insert(param.name, ty); - self.reg_promoted.insert(param.name); + self.body.variables.insert(param.local, reg); + + self.body.reg_promoted.insert(param.local); } else { // Pointer-represented parameters are passed as addresses. // Copy them out of the call argument registers so scalar @@ -592,19 +557,18 @@ impl<'a> FunctionTranslator<'a> { dst: reg, src: src_reg, }); - self.variables.insert(param.name, reg); - self.variable_types.insert(param.name, ty); + self.body.variables.insert(param.local, reg); } } // Set up captured closure variables. // The closure pointer was set by CallClosure before entering this function. - if !self.decl.closure_vars.is_empty() { + if !self.body.decl.closure_vars.is_empty() { let closure_ptr_reg = self.alloc_reg(); func.emit(Opcode::GetClosurePtr { dst: closure_ptr_reg, }); - for (i, cv) in self.decl.closure_vars.iter().enumerate() { + for (i, cv) in self.body.decl.closure_vars.iter().enumerate() { // Load the address of the captured variable from closure_struct[i]. let addr_reg = self.alloc_reg(); func.emit(Opcode::Load64Off { @@ -614,7 +578,7 @@ impl<'a> FunctionTranslator<'a> { }); // Store this address in a local slot so it survives across calls. let slot = self.alloc_local(8); - self.local_slots.insert(cv.name, slot); + self.body.local_slots.insert(*cv, slot); let slot_addr = self.alloc_reg(); func.emit(Opcode::LocalAddr { dst: slot_addr, @@ -626,25 +590,25 @@ impl<'a> FunctionTranslator<'a> { }); // The variable maps to the address of the captured storage (a pointer). // Access goes: load addr from local slot → load/store value through addr. - self.variables.insert(cv.name, addr_reg); - self.variable_types.insert(cv.name, cv.ty); + self.body.variables.insert(*cv, addr_reg); + // Mark as a captured variable (accessed via double indirection). - self.captured_vars.insert(cv.name); + self.body.captured_vars.insert(*cv); } } // Translate the body if present. - if let Some(body) = self.decl.body { + if let Some(body) = self.body.decl.body { let result_reg = self.translate_expr(body, func); // Skip epilogue if we already emitted a return (e.g., explicit return statement). if !self.has_returned { // Return the result. - if returns_via_pointer(self.decl.ret) { + if returns_via_pointer(self.body.decl.ret) { // Reload output pointer from local slot (r0 may have been // clobbered by subcalls). let output = self.reload_output_ptr(func); - let size = self.decl.ret.size(self.decls) as u32; + let size = self.body.decl.ret.size(self.decls) as u32; func.emit(Opcode::MemCopy { dst: output, src: result_reg, @@ -748,66 +712,27 @@ impl<'a> FunctionTranslator<'a> { ptr } - /// Forget everything known about an outer binding of `name`, so a new - /// binding that shadows it is a clean rebinding rather than a mix of the - /// two. Every read path consults these name-keyed sets before falling back - /// to `variables`, so a leftover entry sends reads to the outer binding's - /// storage (or through an indirection the new binding doesn't have). - /// Block scope saves and restores all of them, so the outer binding's - /// state comes back at block exit. - fn shadow_outer_binding(&mut self, name: &Name) { - self.local_slots.remove(name); - self.reg_promoted.remove(name); - self.reference_vars.remove(name); - self.captured_vars.remove(name); - } - - /// Snapshot every name-keyed binding fact, to be restored when the scope - /// that shadowed it ends. Blocks and `for` loops both need this: a binding - /// made inside one must stop being visible when it ends, and an outer - /// binding of the same name must come back. - fn save_bindings(&self) -> SavedBindings { - SavedBindings { - variables: self.variables.clone(), - variable_types: self.variable_types.clone(), - local_slots: self.local_slots.clone(), - reg_promoted: self.reg_promoted.clone(), - reference_vars: self.reference_vars.clone(), - captured_vars: self.captured_vars.clone(), - } - } - - fn restore_bindings(&mut self, saved: SavedBindings) { - self.variables = saved.variables; - self.variable_types = saved.variable_types; - self.local_slots = saved.local_slots; - self.reg_promoted = saved.reg_promoted; - self.reference_vars = saved.reference_vars; - self.captured_vars = saved.captured_vars; - } - /// Give a scalar variable a local slot instead of a register, and return a /// register holding the slot's address. - fn alloc_scalar_slot(&mut self, name: Name, ty: TypeID, func: &mut VMFunction) -> Reg { - self.shadow_outer_binding(&name); + fn alloc_scalar_slot(&mut self, name: LocalId, ty: TypeID, func: &mut VMFunction) -> Reg { let slot = self.alloc_local(ty.size(self.decls) as u32); - self.local_slots.insert(name, slot); + self.body.local_slots.insert(name, slot); let addr = self.alloc_reg(); func.emit(Opcode::LocalAddr { dst: addr, slot }); - self.variables.insert(name, addr); - self.variable_types.insert(name, ty); - self.reg_promoted.remove(&name); + self.body.variables.insert(name, addr); + + self.body.reg_promoted.remove(&name); addr } /// Get the address of a variable's storage (for closure capture). /// For register-promoted vars, spills to a local slot first. - fn get_var_address(&mut self, name: &Name, func: &mut VMFunction) -> Reg { - if self.captured_vars.contains(name) { + fn get_var_address(&mut self, name: &LocalId, func: &mut VMFunction) -> Reg { + if self.body.captured_vars.contains(name) { // This variable was itself captured from an enclosing scope. // Our local slot holds a *pointer* to the actual storage, so we // must follow the indirection to return the real address. - let slot = *self.local_slots.get(name).unwrap(); + let slot = *self.body.local_slots.get(name).unwrap(); let slot_addr = self.alloc_reg(); func.emit(Opcode::LocalAddr { dst: slot_addr, @@ -819,23 +744,26 @@ impl<'a> FunctionTranslator<'a> { addr: slot_addr, }); actual_addr - } else if self.reg_promoted.contains(name) { + } else if self.body.reg_promoted.contains(name) { // Register-promoted scalar: spill to a local slot so we have a stable address. - let val_reg = *self.variables.get(name).unwrap(); + let val_reg = *self.body.variables.get(name).unwrap(); let slot = self.alloc_local(8); - self.local_slots.insert(*name, slot); + self.body.local_slots.insert(*name, slot); let addr = self.alloc_reg(); func.emit(Opcode::LocalAddr { dst: addr, slot }); func.emit(Opcode::Store64 { addr, src: val_reg }); // Change variable from register-promoted to stack-allocated. - self.reg_promoted.remove(name); - self.variables.insert(*name, addr); + self.body.reg_promoted.remove(name); + self.body.variables.insert(*name, addr); addr - } else if let Some(&slot) = self.local_slots.get(name) { + } else if let Some(&slot) = self.body.local_slots.get(name) { // Already stack-allocated: return its address. let addr = self.alloc_reg(); func.emit(Opcode::LocalAddr { dst: addr, slot }); addr + } else if self.is_ptr_type(&self.body.decl.arena.local(*name).ty) { + // Borrowed and aggregate parameters carry the captured storage address. + self.body.variables[name] } else { // Should not happen for captured variables. panic!("get_var_address: variable {:?} has no storage", name); @@ -844,7 +772,7 @@ impl<'a> FunctionTranslator<'a> { /// Get the type of an expression. fn expr_type(&self, expr: ExprID) -> TypeID { - self.decl.types[expr] + self.body.decl.arena.ty(expr) } /// Get the type that determines how an expression is represented at runtime. @@ -853,21 +781,18 @@ impl<'a> FunctionTranslator<'a> { /// codegen still receives an array address and must build the slice fat /// pointer itself. fn representation_type(&self, expr: ExprID) -> TypeID { - match &self.decl.arena.exprs[expr] { - Expr::Id(name) => self - .variable_types - .get(name) - .copied() - .or_else(|| { - self.decls.find(*name).iter().find_map(|decl| { - if let Decl::Global { ty, .. } = decl { - Some(*ty) - } else { - None - } - }) - }) - .unwrap_or_else(|| self.expr_type(expr)), + match &self.body.decl.arena[expr] { + Expr::Id(_) => match self.body.decl.arena.reference(expr) { + Some(Reference::Local(local)) => { + let ty = self.body.decl.arena.local(*local).ty; + match &*ty { + Type::Reference(inner) => *inner, + _ => ty, + } + } + Some(Reference::Instance(instance)) => self.decls.instance(*instance).ty(), + _ => self.expr_type(expr), + }, Expr::ArrayIndex(arr_id, _) => match &*self.representation_type(*arr_id) { Type::Array(elem, _) | Type::Slice(elem) | Type::Reference(elem) => *elem, _ => self.expr_type(expr), @@ -878,7 +803,7 @@ impl<'a> FunctionTranslator<'a> { /// Translate an expression and return the register containing the result. fn translate_expr(&mut self, expr: ExprID, func: &mut VMFunction) -> Reg { - match &self.decl.arena.exprs[expr] { + match &self.body.decl.arena[expr] { Expr::Int(n, _) => { let dst = self.alloc_reg(); func.emit(Opcode::LoadImm { dst, value: *n }); @@ -917,122 +842,129 @@ impl<'a> FunctionTranslator<'a> { dst } - Expr::Id(name) => { - let ty = self.expr_type(expr); - - // Check if it's a captured closure variable (double indirection). - if self.captured_vars.contains(name) { - // Load pointer-to-captured-storage from our local slot. - let slot = *self.local_slots.get(name).unwrap(); - let slot_addr = self.alloc_reg(); - func.emit(Opcode::LocalAddr { - dst: slot_addr, - slot, - }); - let captured_addr = self.alloc_reg(); - func.emit(Opcode::Load64 { - dst: captured_addr, - addr: slot_addr, - }); - // Aggregates and slices are represented by their address, - // and the captured pointer already is that address — - // dereferencing it would yield the first word of the value. - if self.is_ptr_type(&ty) { - return captured_addr; - } - // Now load the value from the captured variable's storage. - let dst = self.alloc_reg(); - self.emit_load(&ty, dst, captured_addr, func); - return dst; - } + Expr::Id(_) => match self.body.decl.arena.reference(expr) { + Some(Reference::Local(name)) => { + let ty = self.expr_type(expr); - // Check if it's a local variable. - if let Some(®) = self.variables.get(name) { - if self.reference_vars.contains(name) { + // Check if it's a captured closure variable (double indirection). + if self.body.captured_vars.contains(name) { + // Load pointer-to-captured-storage from our local slot. + let slot = *self.body.local_slots.get(name).unwrap(); + let slot_addr = self.alloc_reg(); + func.emit(Opcode::LocalAddr { + dst: slot_addr, + slot, + }); + let captured_addr = self.alloc_reg(); + func.emit(Opcode::Load64 { + dst: captured_addr, + addr: slot_addr, + }); + // Aggregates and slices are represented by their address, + // and the captured pointer already is that address — + // dereferencing it would yield the first word of the value. if self.is_ptr_type(&ty) { - return reg; - } - let dst = self.alloc_reg(); - self.emit_load(&ty, dst, reg, func); - dst - } else if self.reg_promoted.contains(name) { - // Register-promoted scalar: value is already in the register. - reg - } else if self.is_ptr_type(&ty) { - // Pointer type: re-emit LocalAddr to ensure the register - // is correct after calls that may have clobbered it. - if let Some(&slot) = self.local_slots.get(name) { - func.emit(Opcode::LocalAddr { dst: reg, slot }); + return captured_addr; } - reg - } else if let Some(&slot) = self.local_slots.get(name) { - // Non-promoted scalar in local slot: load from memory. + // Now load the value from the captured variable's storage. let dst = self.alloc_reg(); - func.emit(Opcode::LocalAddr { dst, slot }); - let load_dst = self.alloc_reg(); - self.emit_load(&ty, load_dst, dst, func); - load_dst - } else { - reg + self.emit_load(&ty, dst, captured_addr, func); + return dst; } - } else if let Some(&offset) = self.globals.get(name) { - // Global variable - load from globals memory. - let addr = self.alloc_reg(); - func.emit(Opcode::GlobalAddr { dst: addr, offset }); - // Composite types (arrays, structs) are pointer-represented: - // return the address, don't load the value. - if self.is_ptr_type(&ty) { - addr + + // Check if it's a local variable. + if let Some(®) = self.body.variables.get(name) { + if self.body.reference_vars.contains(name) { + if self.is_ptr_type(&ty) { + return reg; + } + let dst = self.alloc_reg(); + self.emit_load(&ty, dst, reg, func); + dst + } else if self.body.reg_promoted.contains(name) { + // Register-promoted scalar: value is already in the register. + reg + } else if self.is_ptr_type(&ty) { + // Pointer type: re-emit LocalAddr to ensure the register + // is correct after calls that may have clobbered it. + if let Some(&slot) = self.body.local_slots.get(name) { + func.emit(Opcode::LocalAddr { dst: reg, slot }); + } + reg + } else if let Some(&slot) = self.body.local_slots.get(name) { + // Non-promoted scalar in local slot: load from memory. + let dst = self.alloc_reg(); + func.emit(Opcode::LocalAddr { dst, slot }); + let load_dst = self.alloc_reg(); + self.emit_load(&ty, load_dst, dst, func); + load_dst + } else { + reg + } } else { - let dst = self.alloc_reg(); - self.emit_load(&ty, dst, addr, func); - dst + unreachable!("checked local must have storage") } - } else { - // Check if it's a function. - if let Type::Func(_, _) = &*ty { - // Build a 16-byte fat pointer {func_idx, 0} on the stack. - let fat_slot = self.alloc_local(16); - let fat_addr = self.alloc_reg(); - func.emit(Opcode::LocalAddr { - dst: fat_addr, - slot: fat_slot, - }); - // Store func_idx (patched later). - let func_idx_reg = self.alloc_reg(); - self.pending_functions.push(*name); - let instr_idx = func.emit(Opcode::LoadImm { - dst: func_idx_reg, - value: 0, - }); - self.func_load_patches.push(CallToPatch { - instr_idx, - callee: *name, - }); - func.emit(Opcode::Store64 { - addr: fat_addr, - src: func_idx_reg, - }); - // Store closure_ptr = 0. - let zero_reg = self.alloc_reg(); - func.emit(Opcode::LoadImm { - dst: zero_reg, - value: 0, - }); - func.emit(Opcode::Store64Off { - base: fat_addr, - offset: 8, - src: zero_reg, - }); - fat_addr + } + Some(Reference::Instance(name)) => { + let ty = self.expr_type(expr); + if let Some(&offset) = self.globals.get(name) { + // Global variable - load from globals memory. + let addr = self.alloc_reg(); + func.emit(Opcode::GlobalAddr { dst: addr, offset }); + // Composite types (arrays, structs) are pointer-represented: + // return the address, don't load the value. + if self.is_ptr_type(&ty) { + addr + } else { + let dst = self.alloc_reg(); + self.emit_load(&ty, dst, addr, func); + dst + } } else { - // Unknown identifier - this shouldn't happen after type checking. - let dst = self.alloc_reg(); - func.emit(Opcode::LoadImm { dst, value: 0 }); - dst + // Check if it's a function. + if let Type::Func(_, _) = &*ty { + // Build a 16-byte fat pointer {func_idx, 0} on the stack. + let fat_slot = self.alloc_local(16); + let fat_addr = self.alloc_reg(); + func.emit(Opcode::LocalAddr { + dst: fat_addr, + slot: fat_slot, + }); + // Store func_idx (patched later). + let func_idx_reg = self.alloc_reg(); + self.pending_functions.push(*name); + let instr_idx = func.emit(Opcode::LoadImm { + dst: func_idx_reg, + value: 0, + }); + self.func_load_patches.push(CallToPatch { + instr_idx, + callee: *name, + }); + func.emit(Opcode::Store64 { + addr: fat_addr, + src: func_idx_reg, + }); + // Store closure_ptr = 0. + let zero_reg = self.alloc_reg(); + func.emit(Opcode::LoadImm { + dst: zero_reg, + value: 0, + }); + func.emit(Opcode::Store64Off { + base: fat_addr, + offset: 8, + src: zero_reg, + }); + fat_addr + } else { + unreachable!("instance must name storage or a function") + } } } - } + + _ => unreachable!("non-concrete reference in specialized body"), + }, Expr::Binop(op, lhs_id, rhs_id) => self.translate_binop(*op, *lhs_id, *rhs_id, func), @@ -1040,12 +972,14 @@ impl<'a> FunctionTranslator<'a> { Expr::Call(fn_id, arg_ids) => self.translate_call(*fn_id, arg_ids, expr, func), - Expr::Let(name, init, _) => { - let ty = self.expr_type(expr); + Expr::Let(_, init, _) => { + let name = self.body.decl.arena.binder(expr); + let name = &name; + let ty = self.body.decl.arena.local(*name).ty; let init_reg = self.translate_expr(*init, func); let init_reg = self.wrap_for_expected_slice(init_reg, ty, *init, func); - if !self.is_ptr_type(&ty) && self.lambda_referenced.contains(name) { + if !self.is_ptr_type(&ty) && self.body.lambda_referenced.contains(name) { // Captured by a lambda: must live in memory, not a register. let addr = self.alloc_scalar_slot(*name, ty, func); self.emit_store(&ty, addr, init_reg, func); @@ -1056,12 +990,12 @@ impl<'a> FunctionTranslator<'a> { dst: reg, src: init_reg, }); - self.shadow_outer_binding(name); - self.variables.insert(*name, reg); - self.variable_types.insert(*name, ty); - self.reg_promoted.insert(*name); + + self.body.variables.insert(*name, reg); + + self.body.reg_promoted.insert(*name); } else if crate::copy_elision::is_value_aggregate(&ty) - && !self.elidable_lets.contains(&expr) + && !self.body.elidable_lets.contains(&expr) { // `let` binds aggregates by value, so the initializer's // storage has to be copied — otherwise a slice coerced from @@ -1069,16 +1003,16 @@ impl<'a> FunctionTranslator<'a> { // `var`, which has always copied. let size = self.vm_type_size(&ty); let slot = self.alloc_local(size); - self.shadow_outer_binding(name); - self.local_slots.insert(*name, slot); + + self.body.local_slots.insert(*name, slot); let addr_reg = self.alloc_reg(); func.emit(Opcode::LocalAddr { dst: addr_reg, slot, }); - self.variables.insert(*name, addr_reg); - self.variable_types.insert(*name, ty); + self.body.variables.insert(*name, addr_reg); + self.emit_store(&ty, addr_reg, init_reg, func); return addr_reg; } else { @@ -1088,28 +1022,29 @@ impl<'a> FunctionTranslator<'a> { dst: reg, src: init_reg, }); - // This binding has no slot of its own. If it shadows one - // that does, the outer slot mapping has to go: reads - // re-emit LocalAddr for any slot the name still maps to, - // which would clobber this binding's address register. - self.shadow_outer_binding(name); - self.variables.insert(*name, reg); - self.variable_types.insert(*name, ty); + + self.body.variables.insert(*name, reg); + return reg; } init_reg } - Expr::Var(name, init, _) => { - let ty = self.expr_type(expr); + Expr::Var(_, init, _) => { + let name = self.body.decl.arena.binder(expr); + let name = &name; + let ty = self.body.decl.arena.local(*name).ty; - if !self.is_ptr_type(&ty) && self.lambda_referenced.contains(name) { + if !self.is_ptr_type(&ty) && self.body.lambda_referenced.contains(name) { // Captured by a lambda: must live in memory, not a register. let init_reg = if let Some(init_id) = init { self.translate_expr(*init_id, func) } else { let zero = self.alloc_reg(); - func.emit(Opcode::LoadImm { dst: zero, value: 0 }); + func.emit(Opcode::LoadImm { + dst: zero, + value: 0, + }); zero }; let addr = self.alloc_scalar_slot(*name, ty, func); @@ -1126,24 +1061,23 @@ impl<'a> FunctionTranslator<'a> { } else { func.emit(Opcode::LoadImm { dst: reg, value: 0 }); } - self.shadow_outer_binding(name); - self.variables.insert(*name, reg); - self.variable_types.insert(*name, ty); - self.reg_promoted.insert(*name); + + self.body.variables.insert(*name, reg); + + self.body.reg_promoted.insert(*name); } else { // Pointer type: store to local slot. let size = self.vm_type_size(&ty); let slot = self.alloc_local(size); - self.shadow_outer_binding(name); - self.local_slots.insert(*name, slot); + + self.body.local_slots.insert(*name, slot); let addr_reg = self.alloc_reg(); func.emit(Opcode::LocalAddr { dst: addr_reg, slot, }); - self.variables.insert(*name, addr_reg); - self.variable_types.insert(*name, ty); + self.body.variables.insert(*name, addr_reg); if let Some(init_id) = init { let init_reg = self.translate_expr(*init_id, func); @@ -1219,14 +1153,11 @@ impl<'a> FunctionTranslator<'a> { func.emit(Opcode::LoadImm { dst, value: 0 }); dst } else { - // Save variable scope — declarations inside this block - // shadow outer names only for the duration of the block. - let saved = self.save_bindings(); let mut result = 0; for expr_id in exprs { result = self.translate_expr(*expr_id, func); } - self.restore_bindings(saved); + result } } @@ -1238,11 +1169,11 @@ impl<'a> FunctionTranslator<'a> { Expr::While(cond_id, body_id) => self.translate_while(*cond_id, *body_id, func), Expr::For { - var, - start, - end, - body, - } => self.translate_for(*var, *start, *end, *body, func), + start, end, body, .. + } => { + let var = self.body.decl.arena.binder(expr); + self.translate_for(var, *start, *end, *body, func) + } Expr::Assume(_) => { // No-op: assume is only used by the safety checker. @@ -1398,125 +1329,68 @@ impl<'a> FunctionTranslator<'a> { Expr::AsTy(expr_id, target_ty) => self.translate_cast(*expr_id, *target_ty, func), - Expr::Lambda { params, body } => { - let lambda_ty = self.expr_type(expr); - if let Type::Func(dom, rng) = &*lambda_ty { - if let Type::Tuple(param_types) = &**dom { - let id = *self.lambda_counter; - *self.lambda_counter += 1; - let lambda_name = Name::new(format!("__lambda_{}", id)); - - let lambda_params: Vec = params - .iter() - .zip(param_types.iter()) - .map(|(p, ty)| Param { - name: p.name, - ty: Some(*ty), - }) - .collect(); - - // Compute free variables captured from the enclosing scope. - let param_names: std::collections::HashSet = - params.iter().map(|p| p.name.to_string()).collect(); - let free_vars = collect_free_var_names( - *body, - &self.decl.arena, - ¶m_names, - &self.variables, - &self.decl.types, - ); - - // Build closure struct if there are captures. - let closure_ptr_val = if !free_vars.is_empty() { - let n = free_vars.len(); - let closure_slot = self.alloc_local((n * 8) as u32); - let closure_addr = self.alloc_reg(); - func.emit(Opcode::LocalAddr { - dst: closure_addr, - slot: closure_slot, - }); - for (i, (name, _ty)) in free_vars.iter().enumerate() { - let var_name = Name::new(name.clone()); - let var_addr = self.get_var_address(&var_name, func); - func.emit(Opcode::Store64Off { - base: closure_addr, - offset: (i * 8) as i32, - src: var_addr, - }); - } - closure_addr - } else { - let zero = self.alloc_reg(); - func.emit(Opcode::LoadImm { - dst: zero, - value: 0, - }); - zero - }; - - let closure_vars: Vec = free_vars - .iter() - .map(|(name, ty)| ClosureVar { - name: Name::new(name.clone()), - ty: *ty, - }) - .collect(); - - let lambda_decl = FuncDecl { - name: lambda_name, - typevars: vec![], - size_vars: vec![], - params: lambda_params, - body: Some(*body), - ret: *rng, - constraints: vec![], - requires: vec![], - loc: self.decl.loc, - arena: self.decl.arena.clone(), - types: self.decl.types.clone(), - closure_vars, - is_extern: false, - }; + Expr::Lambda { .. } => { + let id = *self.lambda_counter; + *self.lambda_counter += 1; + let lambda_name = Name::new(format!("__lambda_{}", id)); - self.pending_lambdas.push(lambda_decl); + let lambda_decl = self.body.decl.extract_lambda(expr, lambda_name); + let free_vars = &lambda_decl.closure_vars; - // Build a 16-byte fat pointer {func_idx, closure_ptr}. - let fat_slot = self.alloc_local(16); - let fat_addr = self.alloc_reg(); - func.emit(Opcode::LocalAddr { - dst: fat_addr, - slot: fat_slot, - }); - // Store func_idx (patched later). - let func_idx_reg = self.alloc_reg(); - let instr_idx = func.emit(Opcode::LoadImm { - dst: func_idx_reg, - value: 0, - }); - self.lambda_patches.push((instr_idx, lambda_name)); - func.emit(Opcode::Store64 { - addr: fat_addr, - src: func_idx_reg, - }); - // Store closure_ptr. + // Build closure struct if there are captures. + let closure_ptr_val = if !free_vars.is_empty() { + let n = free_vars.len(); + let closure_slot = self.alloc_local((n * 8) as u32); + let closure_addr = self.alloc_reg(); + func.emit(Opcode::LocalAddr { + dst: closure_addr, + slot: closure_slot, + }); + for (i, var_name) in free_vars.iter().enumerate() { + let var_addr = self.get_var_address(var_name, func); func.emit(Opcode::Store64Off { - base: fat_addr, - offset: 8, - src: closure_ptr_val, + base: closure_addr, + offset: (i * 8) as i32, + src: var_addr, }); - fat_addr - } else { - panic!( - "VM codegen lambda: expected tuple domain type, got {:?}", - dom - ); } + closure_addr } else { - panic!( - "VM codegen lambda: expected function type, got {:?}", - lambda_ty - ); - } + let zero = self.alloc_reg(); + func.emit(Opcode::LoadImm { + dst: zero, + value: 0, + }); + zero + }; + + self.pending_lambdas.push(lambda_decl); + + // Build a 16-byte fat pointer {func_idx, closure_ptr}. + let fat_slot = self.alloc_local(16); + let fat_addr = self.alloc_reg(); + func.emit(Opcode::LocalAddr { + dst: fat_addr, + slot: fat_slot, + }); + // Store func_idx (patched later). + let func_idx_reg = self.alloc_reg(); + let instr_idx = func.emit(Opcode::LoadImm { + dst: func_idx_reg, + value: 0, + }); + self.lambda_patches.push((instr_idx, lambda_name)); + func.emit(Opcode::Store64 { + addr: fat_addr, + src: func_idx_reg, + }); + // Store closure_ptr. + func.emit(Opcode::Store64Off { + base: fat_addr, + offset: 8, + src: closure_ptr_val, + }); + fat_addr } Expr::Char(c) => { @@ -1529,7 +1403,8 @@ impl<'a> FunctionTranslator<'a> { } Expr::Enum(case_name) => { - let index = if let crate::Type::Name(enum_name, _) = &*self.decl.types[expr] { + let index = if let crate::Type::Name(enum_name, _) = &*self.body.decl.arena.ty(expr) + { let enum_decls = self.decls.find(*enum_name); if let Some(crate::Decl::Enum { cases, .. }) = enum_decls .iter() @@ -1558,11 +1433,8 @@ impl<'a> FunctionTranslator<'a> { Expr::Arena(inner) => self.translate_expr(*inner, func), - _ => { - // Unimplemented expression - return 0. - let dst = self.alloc_reg(); - func.emit(Opcode::LoadImm { dst, value: 0 }); - dst + Expr::Macro(..) | Expr::TypeApp(..) | Expr::Error => { + unreachable!("unresolved expression in specialized body") } } } @@ -1577,11 +1449,11 @@ impl<'a> FunctionTranslator<'a> { is_f64: bool, func: &mut VMFunction, ) -> Option { - let arena = &self.decl.arena; + let arena = &self.body.decl.arena; if op == Binop::Plus { // a + b*c → FMulAdd { dst, a: b, b: c, c: a } - if let Expr::Binop(Binop::Mult, ma, mb) = arena.exprs[rhs_id] { + if let Expr::Binop(Binop::Mult, ma, mb) = arena[rhs_id] { let c = self.translate_expr(lhs_id, func); let a = self.translate_expr(ma, func); let b = self.translate_expr(mb, func); @@ -1594,7 +1466,7 @@ impl<'a> FunctionTranslator<'a> { return Some(dst); } // b*c + a → FMulAdd { dst, a: b, b: c, c: a } - if let Expr::Binop(Binop::Mult, ma, mb) = arena.exprs[lhs_id] { + if let Expr::Binop(Binop::Mult, ma, mb) = arena[lhs_id] { let a = self.translate_expr(ma, func); let b = self.translate_expr(mb, func); let c = self.translate_expr(rhs_id, func); @@ -1608,7 +1480,7 @@ impl<'a> FunctionTranslator<'a> { } } else if op == Binop::Minus { // b*c - a → FMulSub { dst, a: b, b: c, c: a } (b*c - a) - if let Expr::Binop(Binop::Mult, ma, mb) = arena.exprs[lhs_id] { + if let Expr::Binop(Binop::Mult, ma, mb) = arena[lhs_id] { let a = self.translate_expr(ma, func); let b = self.translate_expr(mb, func); let c = self.translate_expr(rhs_id, func); @@ -1621,7 +1493,7 @@ impl<'a> FunctionTranslator<'a> { return Some(dst); } // a - b*c → FNMulAdd { dst, a: b, b: c, c: a } (a - b*c) - if let Expr::Binop(Binop::Mult, ma, mb) = arena.exprs[rhs_id] { + if let Expr::Binop(Binop::Mult, ma, mb) = arena[rhs_id] { let c = self.translate_expr(lhs_id, func); let a = self.translate_expr(ma, func); let b = self.translate_expr(mb, func); @@ -2105,12 +1977,12 @@ impl<'a> FunctionTranslator<'a> { /// Translate an assignment expression. fn translate_assign(&mut self, lhs_id: ExprID, rhs_id: ExprID, func: &mut VMFunction) -> Reg { // Check for captured variable assignment (double indirection). - if let Expr::Id(name) = &self.decl.arena.exprs[lhs_id] { - if self.captured_vars.contains(name) { + if let Some(Reference::Local(name)) = self.body.decl.arena.reference(lhs_id) { + if self.body.captured_vars.contains(name) { let rhs = self.translate_expr(rhs_id, func); let ty = self.representation_type(lhs_id); let rhs = self.wrap_for_expected_slice(rhs, ty, rhs_id, func); - let slot = *self.local_slots.get(name).unwrap(); + let slot = *self.body.local_slots.get(name).unwrap(); let slot_addr = self.alloc_reg(); func.emit(Opcode::LocalAddr { dst: slot_addr, @@ -2126,10 +1998,10 @@ impl<'a> FunctionTranslator<'a> { } } // Check for direct register-promoted scalar assignment (e.g., `x = expr`). - if let Expr::Id(name) = &self.decl.arena.exprs[lhs_id] { - if self.reg_promoted.contains(name) { + if let Some(Reference::Local(name)) = self.body.decl.arena.reference(lhs_id) { + if self.body.reg_promoted.contains(name) { let rhs = self.translate_expr(rhs_id, func); - let reg = *self.variables.get(name).unwrap(); + let reg = *self.body.variables.get(name).unwrap(); if reg != rhs { func.emit(Opcode::Move { dst: reg, src: rhs }); } @@ -2137,7 +2009,7 @@ impl<'a> FunctionTranslator<'a> { } } // Slice store superinstruction: a[i] = rhs where a is a slice of 32-bit elements. - if let Expr::ArrayIndex(arr_id, idx_id) = &self.decl.arena.exprs[lhs_id] { + if let Expr::ArrayIndex(arr_id, idx_id) = &self.body.decl.arena[lhs_id] { let arr_ty = self.expr_type(*arr_id); if let Type::Slice(elem_ty) = &*arr_ty { let elem_size = elem_ty.size(self.decls); @@ -2162,7 +2034,7 @@ impl<'a> FunctionTranslator<'a> { // Struct field assignment of Func type: only copy 8 bytes (func_idx). // The struct field stores only the func_idx, not the full 16-byte fat pointer. if matches!(&*ty, Type::Func(_, _)) { - if matches!(&self.decl.arena.exprs[lhs_id], Expr::Field(_, _)) { + if matches!(&self.body.decl.arena[lhs_id], Expr::Field(_, _)) { // rhs is a fat pointer address; load func_idx from it and store to struct field. let func_idx_reg = self.alloc_reg(); func.emit(Opcode::Load64 { @@ -2182,37 +2054,37 @@ impl<'a> FunctionTranslator<'a> { /// Translate an lvalue expression (returns address). fn translate_lvalue(&mut self, expr: ExprID, func: &mut VMFunction) -> Reg { - match &self.decl.arena.exprs[expr] { - Expr::Id(name) => { - if let Some(®) = self.variables.get(name) { - if self.reg_promoted.contains(name) { - let ty = self - .variable_types - .get(name) - .copied() - .unwrap_or(self.expr_type(expr)); - let slot = self.alloc_local(ty.size(self.decls) as u32); - let addr = self.alloc_reg(); - func.emit(Opcode::LocalAddr { dst: addr, slot }); - self.emit_store(&ty, addr, reg, func); - self.variables.insert(*name, addr); - self.local_slots.insert(*name, slot); - self.reg_promoted.remove(name); - addr + match &self.body.decl.arena[expr] { + Expr::Id(_) => match self.body.decl.arena.reference(expr) { + Some(Reference::Local(name)) => { + if let Some(®) = self.body.variables.get(name) { + if self.body.reg_promoted.contains(name) { + let ty = self.body.decl.arena.local(*name).ty; + let slot = self.alloc_local(ty.size(self.decls) as u32); + let addr = self.alloc_reg(); + func.emit(Opcode::LocalAddr { dst: addr, slot }); + self.emit_store(&ty, addr, reg, func); + self.body.variables.insert(*name, addr); + self.body.local_slots.insert(*name, slot); + self.body.reg_promoted.remove(name); + addr + } else { + reg + } } else { - reg + unreachable!("checked local must have storage") } - } else if let Some(&offset) = self.globals.get(name) { - // Global variable - return its address. - let dst = self.alloc_reg(); - func.emit(Opcode::GlobalAddr { dst, offset }); - dst - } else { + } + Some(Reference::Instance(instance)) => { let dst = self.alloc_reg(); - func.emit(Opcode::LoadImm { dst, value: 0 }); + func.emit(Opcode::GlobalAddr { + dst, + offset: self.globals[instance], + }); dst } - } + _ => unreachable!("non-concrete reference in specialized body"), + }, Expr::Field(lhs_id, name) => { let lhs_addr = self.translate_lvalue(*lhs_id, func); @@ -2381,7 +2253,9 @@ impl<'a> FunctionTranslator<'a> { } // Special handling for built-in functions. - if let Expr::Id(name) = &self.decl.arena.exprs[fn_id] { + if let Some(Reference::Instance(instance)) = self.body.decl.arena.reference(fn_id) { + let instance = *instance; + let name = &self.decls.instance_name(instance); if **name == "print" { // Print the first argument. if let Some(&arg_id) = arg_ids.first() { @@ -2564,19 +2438,18 @@ impl<'a> FunctionTranslator<'a> { } // Check for extern function calls — emit CallExtern instead of Call. - if let Expr::Id(callee_name) = &self.decl.arena.exprs[fn_id] { - let callee_decls = self.decls.find(*callee_name); - if let Some(Decl::Func(f)) = callee_decls.first() { + { + if let Some(f) = self.decls.function_instance(instance) { if f.is_extern { let globals_offset = *self .globals - .get(callee_name) + .get(&instance) .expect("extern function not in globals"); // Build the C-level parameter types (slices expand to ptr + i32). let mut c_param_types: Vec = Vec::new(); for p in &f.params { - c_param_types.extend(type_to_extern_types(p.ty.unwrap())); + c_param_types.extend(type_to_extern_types(f.arena.local(p.local).ty)); } let ret_types = type_to_extern_types(f.ret); let ret_type = ret_types[0]; @@ -2589,7 +2462,7 @@ impl<'a> FunctionTranslator<'a> { // Translate arguments, wrapping arrays as slices where needed. let mut arg_values = Vec::new(); for (i, arg_id) in arg_ids.iter().enumerate() { - let param_ty = f.params[i].ty.unwrap(); + let param_ty = f.arena.local(f.params[i].local).ty; let arg_reg = if matches!(&*param_ty, Type::Reference(_)) { self.translate_lvalue(*arg_id, func) } else { @@ -2614,7 +2487,7 @@ impl<'a> FunctionTranslator<'a> { let mut c_arg_count: u8 = 0; for (i, param) in f.params.iter().enumerate() { let arg_reg = arg_values[i]; - let param_ty = param.ty.unwrap(); + let param_ty = f.arena.local(param.local).ty; if matches!(&*param_ty, Type::Slice(_)) { // arg_reg points to fat pointer {data_ptr: i64, len: i32}. let ptr_reg = self.alloc_reg(); @@ -2659,14 +2532,9 @@ impl<'a> FunctionTranslator<'a> { // Try to inline small leaf functions. This avoids Call overhead // for trivial functions like `cmp(lhs, rhs) { lhs - rhs }`. - if let Expr::Id(callee_name) = &self.decl.arena.exprs[fn_id] { - let callee_decls = self.decls.find(*callee_name); - for d in callee_decls { - if let Decl::Func(callee) = d { - if let Some(result) = self.try_inline(callee, arg_ids, func) { - return result; - } - } + if let Some(callee) = self.decls.function_instance(instance) { + if let Some(result) = self.try_inline(callee, arg_ids, func) { + return result; } } @@ -2685,17 +2553,11 @@ impl<'a> FunctionTranslator<'a> { // Get the callee's declared parameter types to detect slice params. // We use the declaration (not the solved call-site type) because // the solver may retain Array types where the callee expects Slice. - let param_types: Vec = - if let Expr::Id(callee_name) = &self.decl.arena.exprs[fn_id] { - let callee_decls = self.decls.find(*callee_name); - if let Some(Decl::Func(f)) = callee_decls.first() { - f.param_types() - } else { - vec![] - } - } else { - self.closure_param_types(fn_id) - }; + let param_types = self + .decls + .function_instance(instance) + .expect("direct call target is a function") + .param_types(); // First, translate all arguments to get their values. // We need to do this before allocating the consecutive arg registers @@ -2751,7 +2613,7 @@ impl<'a> FunctionTranslator<'a> { } // Add function to pending list. - self.pending_functions.push(*name); + self.pending_functions.push(instance); // Calculate arg count (including output pointer if present). let arg_count = if output_slot.is_some() { @@ -2771,7 +2633,7 @@ impl<'a> FunctionTranslator<'a> { // Record this call for patching. self.calls_to_patch.push(CallToPatch { instr_idx, - callee: *name, + callee: instance, }); // If returning pointer type, return the address of the output storage. @@ -2804,21 +2666,13 @@ impl<'a> FunctionTranslator<'a> { /// rather than naming a function declaration. Such calls go through /// `translate_closure_call` instead of the direct-call path. fn holds_fat_pointer(&self, fn_id: ExprID) -> bool { - let Expr::Id(name) = &self.decl.arena.exprs[fn_id] else { - return false; - }; - if self.variables.contains_key(name) { - return true; + match self.body.decl.arena.reference(fn_id) { + Some(Reference::Local(_)) => true, + Some(Reference::Instance(id)) => { + matches!(self.decls.instance(*id), Decl::Global { .. }) + } + _ => false, } - // Extern functions live in globals memory too, but they are called - // through the direct-call path. - self.globals.contains_key(name) - && matches!(&*self.expr_type(fn_id), Type::Func(_, _)) - && !self - .decls - .find(*name) - .iter() - .any(|d| matches!(d, Decl::Func(f) if f.is_extern)) } /// Parameter types of a callee reached through a fat pointer, taken from @@ -3025,7 +2879,7 @@ impl<'a> FunctionTranslator<'a> { /// Translate a for loop. fn translate_for( &mut self, - var: Name, + var: LocalId, start_id: ExprID, end_id: ExprID, body_id: ExprID, @@ -3047,9 +2901,9 @@ impl<'a> FunctionTranslator<'a> { // name-keyed state — otherwise reads of the name go to its slot // instead of the counter — and the snapshot brings that state back at // loop exit, where the loop variable is out of scope again. - let saved = self.save_bindings(); + let int_ty = mk_type(Type::Int32); - let counter_slot = if self.lambda_referenced.contains(&var) { + let counter_slot = if self.body.lambda_referenced.contains(&var) { // A lambda shares the counter by address, so it needs storage of // its own, allocated up front the way `let` and `var` do it. // Leaving it register-promoted and letting `get_var_address` spill @@ -3057,12 +2911,11 @@ impl<'a> FunctionTranslator<'a> { // on a conditionally-executed path — iterations that don't reach // it would then read an unwritten slot. self.alloc_scalar_slot(var, int_ty, func); - self.local_slots.get(&var).copied() + self.body.local_slots.get(&var).copied() } else { - self.shadow_outer_binding(&var); - self.variables.insert(var, loop_var); - self.variable_types.insert(var, int_ty); - self.reg_promoted.insert(var); + self.body.variables.insert(var, loop_var); + + self.body.reg_promoted.insert(var); None }; @@ -3098,7 +2951,6 @@ impl<'a> FunctionTranslator<'a> { // Execute body. self.translate_expr(body_id, func); - self.restore_bindings(saved); // Increment position — this is where continue jumps to. let increment_pos = func.code.len(); @@ -3387,10 +3239,10 @@ impl<'a> FunctionTranslator<'a> { /// - Not be recursive /// /// The inlined body is translated using the callee's arena and types, - /// with callee parameter names temporarily bound to argument registers. + /// with callee parameter identities bound in a fresh body context. fn try_inline( &mut self, - callee: &'a FuncDecl, + callee: &'a CheckedFunction, arg_ids: &[ExprID], func: &mut VMFunction, ) -> Option { @@ -3408,10 +3260,9 @@ impl<'a> FunctionTranslator<'a> { // Only inline scalar parameters (no slices, structs, arrays). for param in &callee.params { - if let Some(ty) = param.ty { - if ty.is_ptr() || matches!(&*ty, Type::Slice(_)) { - return None; - } + let ty = callee.arena.local(param.local).ty; + if ty.is_ptr() || matches!(&*ty, Type::Slice(_)) { + return None; } } @@ -3426,25 +3277,15 @@ impl<'a> FunctionTranslator<'a> { .map(|arg| self.translate_expr(*arg, func)) .collect(); - // Temporarily bind callee parameters to argument registers. + // Local IDs are interpreted only in the body context that owns them. + // Register/slot allocation and emitted instructions stay in the caller. + let caller = std::mem::replace(&mut self.body, BodyContext::new(callee)); for (param, ®) in callee.params.iter().zip(arg_regs.iter()) { - self.variables.insert(param.name, reg); - self.reg_promoted.insert(param.name); + self.body.variables.insert(param.local, reg); + self.body.reg_promoted.insert(param.local); } - - // Save the caller's decl and switch to the callee's. - let saved_decl = self.decl; - self.decl = callee; - - // Translate the callee body in-place. let result = self.translate_expr(body, func); - - // Restore caller context. - self.decl = saved_decl; - for param in &callee.params { - self.variables.remove(¶m.name); - self.reg_promoted.remove(¶m.name); - } + self.body = caller; Some(result) } @@ -3761,357 +3602,277 @@ impl<'a> FunctionTranslator<'a> { } } -/// Collect free variable names referenced in a lambda body that come from the enclosing scope. -fn collect_free_var_names( - body: crate::ExprID, - arena: &crate::ExprArena, - exclude: &std::collections::HashSet, - local_vars: &HashMap, - types: &[crate::TypeID], -) -> Vec<(String, crate::TypeID)> { - let mut result = Vec::new(); - let mut seen = std::collections::HashSet::new(); - collect_free_vars_rec( - body, - arena, - exclude, - local_vars, - types, - &mut result, - &mut seen, - ); - result -} - -fn collect_free_vars_rec( - expr: crate::ExprID, - arena: &crate::ExprArena, - exclude: &std::collections::HashSet, - local_vars: &HashMap, - types: &[crate::TypeID], - result: &mut Vec<(String, crate::TypeID)>, - seen: &mut std::collections::HashSet, -) { - match &arena[expr] { - Expr::TypeApp(_, _) => {} // Rewritten to Id by monomorphizer - Expr::Id(name) => { - let s = name.to_string(); - if local_vars.contains_key(name) && !exclude.contains(&s) && !seen.contains(&s) { - result.push((s.clone(), types[expr])); - seen.insert(s); - } - } - Expr::Call(fn_id, args) => { - collect_free_vars_rec(*fn_id, arena, exclude, local_vars, types, result, seen); - for a in args { - collect_free_vars_rec(*a, arena, exclude, local_vars, types, result, seen); - } - } - Expr::Binop(_, lhs, rhs) => { - collect_free_vars_rec(*lhs, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec(*rhs, arena, exclude, local_vars, types, result, seen); - } - Expr::Unop(_, arg) => { - collect_free_vars_rec(*arg, arena, exclude, local_vars, types, result, seen); - } - Expr::Let(_, init, _) => { - collect_free_vars_rec(*init, arena, exclude, local_vars, types, result, seen); - } - Expr::Var(_, init, _) => { - if let Some(init_id) = init { - collect_free_vars_rec(*init_id, arena, exclude, local_vars, types, result, seen); - } - } - Expr::If(cond, then, else_) => { - collect_free_vars_rec(*cond, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec(*then, arena, exclude, local_vars, types, result, seen); - if let Some(e) = else_ { - collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen); - } - } - Expr::While(cond, body) => { - collect_free_vars_rec(*cond, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec(*body, arena, exclude, local_vars, types, result, seen); - } - Expr::For { - start, end, body, .. - } => { - collect_free_vars_rec(*start, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec(*end, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec(*body, arena, exclude, local_vars, types, result, seen); - } - Expr::Block(exprs) => { - for e in exprs { - collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen); - } - } - Expr::Return(e) | Expr::Assume(e) => { - collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen); - } - Expr::Field(e, _) => { - collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen); - } - Expr::ArrayIndex(arr, idx) => { - collect_free_vars_rec(*arr, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec(*idx, arena, exclude, local_vars, types, result, seen); - } - Expr::ArrayLiteral(elems) => { - for e in elems { - collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen); - } - } - Expr::Tuple(elems) => { - for e in elems { - collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen); - } - } - Expr::AsTy(e, _) => { - collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen); - } - Expr::Arena(e) => { - collect_free_vars_rec(*e, arena, exclude, local_vars, types, result, seen); - } - Expr::Array(ty_expr, size_expr) => { - collect_free_vars_rec(*ty_expr, arena, exclude, local_vars, types, result, seen); - collect_free_vars_rec(*size_expr, arena, exclude, local_vars, types, result, seen); - } - Expr::Lambda { params, body } => { - let mut inner_exclude = exclude.clone(); - for p in params { - inner_exclude.insert(p.name.to_string()); - } - collect_free_vars_rec( - *body, - arena, - &inner_exclude, - local_vars, - types, - result, - seen, - ); - } - Expr::Macro(_, args) => { - for a in args { - collect_free_vars_rec(*a, arena, exclude, local_vars, types, result, seen); - } - } - Expr::StructLit(_, fields) => { - for (_, fval) in fields { - collect_free_vars_rec(*fval, arena, exclude, local_vars, types, result, seen); - } - } - Expr::Int(_, _) - | Expr::Real(_, _) - | Expr::String(_) - | Expr::Char(_) - | Expr::True - | Expr::False - | Expr::Enum(_) - | Expr::Break - | Expr::Continue - | Expr::Error => {} - } -} - #[cfg(test)] mod tests { use super::*; + use crate::checked::{DefId, InstanceRecord}; use crate::vm::VM; - /// Helper to create a simple function for testing. - fn make_simple_decl_table(body: Expr, ret_ty: TypeID) -> DeclTable { - let mut arena = ExprArena::new(); - let body_id = arena.add(body, crate::test_loc()); + fn program(expressions: Vec<(Expr, TypeID)>) -> VMProgram { + let mut arena = CheckedBody::new(); + for (expr, ty) in expressions { + arena.add(expr, ty, crate::test_loc()); + } + compile_body(arena) + } - let func = FuncDecl { + fn compile_body(arena: CheckedBody) -> VMProgram { + let body = arena.len() - 1; + let func = CheckedFunction { name: Name::str("main"), typevars: vec![], size_vars: vec![], params: vec![], - body: Some(body_id), - ret: ret_ty, - constraints: vec![], + body: Some(body), + ret: arena.ty(body), + requires: vec![], loc: crate::test_loc(), arena, - types: vec![ret_ty], // Simplified - just use return type. closure_vars: vec![], is_extern: false, }; - - DeclTable::new(vec![Decl::Func(func)]) + let checked = SpecializedProgram::from_instances( + vec![Decl::Func(func)], + vec![InstanceRecord { + definition: DefId(0), + type_args: vec![], + size_args: vec![], + declaration: 0, + }], + ); + VMCodegen::new().compile(&checked).unwrap() } #[test] fn test_compile_simple_int() { - let mut arena = ExprArena::new(); - let expr = arena.add(Expr::Int(42, None), crate::test_loc()); - - let func = FuncDecl { - name: Name::str("main"), - typevars: vec![], - size_vars: vec![], - params: vec![], - body: Some(expr), - ret: mk_type(Type::Int32), - constraints: vec![], - requires: vec![], - loc: crate::test_loc(), - arena, - types: vec![mk_type(Type::Int32)], - closure_vars: vec![], - is_extern: false, - }; - - let decls = DeclTable::new(vec![Decl::Func(func)]); - - let mut codegen = VMCodegen::new(); - let program = codegen.compile(&decls).unwrap(); - - let mut vm = VM::new(); - let result = vm.run(&program); - assert_eq!(result, 42); + let program = program(vec![(Expr::Int(42, None), mk_type(Type::Int32))]); + assert_eq!(VM::new().run(&program), 42); } #[test] fn test_compile_addition() { - let mut arena = ExprArena::new(); - let lhs = arena.add(Expr::Int(10, None), crate::test_loc()); - let rhs = arena.add(Expr::Int(32, None), crate::test_loc()); - let add = arena.add(Expr::Binop(Binop::Plus, lhs, rhs), crate::test_loc()); - - let int32 = mk_type(Type::Int32); - let func = FuncDecl { - name: Name::str("main"), - typevars: vec![], - size_vars: vec![], - params: vec![], - body: Some(add), - ret: int32, - constraints: vec![], - requires: vec![], - loc: crate::test_loc(), - arena, - types: vec![int32, int32, int32], - closure_vars: vec![], - is_extern: false, - }; - - let decls = DeclTable::new(vec![Decl::Func(func)]); - - let mut codegen = VMCodegen::new(); - let program = codegen.compile(&decls).unwrap(); - - let mut vm = VM::new(); - let result = vm.run(&program); - assert_eq!(result, 42); + let program = program(vec![ + (Expr::Int(10, None), mk_type(Type::Int32)), + (Expr::Int(32, None), mk_type(Type::Int32)), + (Expr::Binop(Binop::Plus, 0, 1), mk_type(Type::Int32)), + ]); + assert_eq!(VM::new().run(&program), 42); } #[test] fn test_compile_float_arithmetic() { - let mut arena = ExprArena::new(); - let lhs = arena.add(Expr::Real("1.5".to_string(), None), crate::test_loc()); - let rhs = arena.add(Expr::Real("2.5".to_string(), None), crate::test_loc()); - let mul = arena.add(Expr::Binop(Binop::Mult, lhs, rhs), crate::test_loc()); - - let f32_ty = mk_type(Type::Float32); - let func = FuncDecl { - name: Name::str("main"), - typevars: vec![], - size_vars: vec![], - params: vec![], - body: Some(mul), - ret: f32_ty, - constraints: vec![], - requires: vec![], - loc: crate::test_loc(), - arena, - types: vec![f32_ty, f32_ty, f32_ty], - closure_vars: vec![], - is_extern: false, - }; - - let decls = DeclTable::new(vec![Decl::Func(func)]); + let ty = mk_type(Type::Float32); + let program = program(vec![ + (Expr::Real("1.5".into(), None), ty), + (Expr::Real("2.5".into(), None), ty), + (Expr::Binop(Binop::Mult, 0, 1), ty), + ]); + assert!((VM::new().run_f32(&program) - 3.75).abs() < 0.0001); + } - let mut codegen = VMCodegen::new(); - let program = codegen.compile(&decls).unwrap(); + #[test] + fn test_compile_if_else() { + let program = program(vec![ + (Expr::True, mk_type(Type::Bool)), + (Expr::Int(100, None), mk_type(Type::Int32)), + (Expr::Int(200, None), mk_type(Type::Int32)), + (Expr::If(0, 1, Some(2)), mk_type(Type::Int32)), + ]); + assert_eq!(VM::new().run(&program), 100); + } - let mut vm = VM::new(); - let result = vm.run_f32(&program); - assert!((result - 3.75).abs() < 0.0001); + #[test] + fn test_compile_while_loop() { + let ty = mk_type(Type::Int32); + let void = mk_type(Type::Void); + let boolean = mk_type(Type::Bool); + let loc = crate::test_loc(); + for iterations in [0, 1, 4] { + let mut arena = CheckedBody::new(); + let counter = arena.add_local(Name::str("counter"), ty, true); + let zero = arena.add(Expr::Int(0, None), ty, loc); + let binding = arena.add_var(counter, Some(zero), loc); + let condition = if iterations == 0 { + arena.add(Expr::False, boolean, loc) + } else { + let read = arena.add_local_read(counter, ty, loc); + let limit = arena.add(Expr::Int(iterations, None), ty, loc); + arena.add(Expr::Binop(Binop::Less, read, limit), boolean, loc) + }; + let read = arena.add_local_read(counter, ty, loc); + let one = arena.add(Expr::Int(1, None), ty, loc); + let next = arena.add(Expr::Binop(Binop::Plus, read, one), ty, loc); + let increment = arena.add(Expr::Binop(Binop::Assign, read, next), ty, loc); + let loop_expr = arena.add(Expr::While(condition, increment), void, loc); + // The returned state exposes omitted, extra, or missing iterations. + let result = arena.add_local_read(counter, ty, loc); + arena.add(Expr::Block(vec![binding, loop_expr, result]), ty, loc); + + assert_eq!( + VM::new().run(&compile_body(arena)), + iterations, + "while loop with {} iterations", + iterations + ); + } } #[test] - fn test_compile_if_else() { - let mut arena = ExprArena::new(); - let cond = arena.add(Expr::True, crate::test_loc()); - let then_val = arena.add(Expr::Int(100, None), crate::test_loc()); - let else_val = arena.add(Expr::Int(200, None), crate::test_loc()); - let if_expr = arena.add(Expr::If(cond, then_val, Some(else_val)), crate::test_loc()); - - let int32 = mk_type(Type::Int32); - let bool_ty = mk_type(Type::Bool); - let func = FuncDecl { - name: Name::str("main"), + fn resolved_targets_and_body_local_ids_survive_lowering() { + use crate::checked::CheckedParam; + let ty = mk_type(Type::Int32); + let loc = crate::test_loc(); + let make_function = |name, params, arena: CheckedBody| CheckedFunction { + name, + params, + body: Some(arena.len() - 1), + ret: ty, typevars: vec![], size_vars: vec![], - params: vec![], - body: Some(if_expr), - ret: int32, - constraints: vec![], + requires: vec![], - loc: crate::test_loc(), + loc, arena, - types: vec![bool_ty, int32, int32, int32], closure_vars: vec![], is_extern: false, }; + // Both callee definitions deliberately have the same diagnostic name. + // The checked instance reference chooses the second definition. + let mut wrong_body = CheckedBody::new(); + wrong_body.add(Expr::Int(900, None), ty, loc); + let wrong = make_function(Name::str("step"), vec![], wrong_body); + let mut callee_body = CheckedBody::new(); + let parameter = callee_body.add_local(Name::str("value"), ty, false); + let value = callee_body.add_local_read(parameter, ty, loc); + let one = callee_body.add(Expr::Int(1, None), ty, loc); + callee_body.add(Expr::Binop(Binop::Plus, value, one), ty, loc); + let callee = make_function( + Name::str("step"), + vec![CheckedParam { local: parameter }], + callee_body, + ); - let decls = DeclTable::new(vec![Decl::Func(func)]); - - let mut codegen = VMCodegen::new(); - let program = codegen.compile(&decls).unwrap(); + let mut main_body = CheckedBody::new(); + let outer = main_body.add_local(Name::str("value"), ty, false); + assert_eq!( + outer, parameter, + "the two bodies intentionally reuse local index zero" + ); + let forty = main_body.add(Expr::Int(40, None), ty, loc); + let binding = main_body.add_let(outer, forty, loc); + let target = main_body.add_id( + Name::str("step"), + Reference::Instance(InstanceId(2)), + callee.ty(), + loc, + ); + let arg = main_body.add(Expr::Int(0, None), ty, loc); + let first_call = main_body.add(Expr::Call(target, vec![arg]), ty, loc); + let target = main_body.add_id( + Name::str("step"), + Reference::Instance(InstanceId(2)), + callee.ty(), + loc, + ); + let arg = main_body.add(Expr::Int(0, None), ty, loc); + let second_call = main_body.add(Expr::Call(target, vec![arg]), ty, loc); + let call = main_body.add(Expr::Binop(Binop::Plus, first_call, second_call), ty, loc); + let outer_read = main_body.add_local_read(outer, ty, loc); + let sum = main_body.add(Expr::Binop(Binop::Plus, call, outer_read), ty, loc); + main_body.add(Expr::Block(vec![binding, sum]), ty, loc); + let main = make_function(Name::str("main"), vec![], main_body); + let checked = SpecializedProgram::from_instances( + vec![Decl::Func(main), Decl::Func(wrong), Decl::Func(callee)], + (0..3) + .map(|i| InstanceRecord { + definition: DefId(i), + type_args: vec![], + size_args: vec![], + declaration: i as usize, + }) + .collect(), + ); + let vm = VMCodegen::new().compile(&checked).unwrap(); + assert_eq!(VM::new().run(&vm), 42); + #[cfg(has_stack_interp)] + { + let mut stack = crate::stack_codegen::StackCodegen::new() + .compile(&checked) + .unwrap(); + for function in &mut stack.functions { + crate::stack_rebase_lm::rebase(function); + crate::stack_rebase_lm::patch_call_preserve(function); + } + assert_eq!(crate::stack_interp_bridge::run(&stack), 42); + } + } - let mut vm = VM::new(); - let result = vm.run(&program); - assert_eq!(result, 100); + fn check_capture_program(source: &str) { + let mut compiler = crate::Compiler::new(); + compiler.quiet = true; + compiler.parse(source, "."); + assert!( + compiler.check(), + "capture regression must check successfully" + ); + compiler.specialize().unwrap(); + #[cfg(has_stack_interp)] + { + let stack = compiler.compile_stack().unwrap(); + assert_eq!( + crate::stack_interp_bridge::run(&stack), + 42, + "Stack capture result" + ); + } + let vm = compiler.compile_vm().unwrap(); + assert_eq!(VM::new().run(&vm), 42, "register VM capture result"); } #[test] - fn test_compile_while_loop() { - // while (false) { 1 } - let mut arena = ExprArena::new(); - let cond = arena.add(Expr::False, crate::test_loc()); - let body = arena.add(Expr::Int(1, None), crate::test_loc()); - let while_expr = arena.add(Expr::While(cond, body), crate::test_loc()); - let result = arena.add(Expr::Int(42, None), crate::test_loc()); - let block = arena.add(Expr::Block(vec![while_expr, result]), crate::test_loc()); - - let int32 = mk_type(Type::Int32); - let bool_ty = mk_type(Type::Bool); - let func = FuncDecl { - name: Name::str("main"), - typevars: vec![], - size_vars: vec![], - params: vec![], - body: Some(block), - ret: int32, - constraints: vec![], - requires: vec![], - loc: crate::test_loc(), - arena, - types: vec![bool_ty, int32, int32, int32, int32], - closure_vars: vec![], - is_extern: false, - }; - - let decls = DeclTable::new(vec![Decl::Func(func)]); + fn scalar_parameter_capture_has_storage_before_conditional() { + check_capture_program( + r#" + capture(value: i32, create: bool) -> i32 { + if create { let read = || { value }; } + value + } + main() -> i32 { capture(42, false) } + "#, + ); + } - let mut codegen = VMCodegen::new(); - let program = codegen.compile(&decls).unwrap(); + #[test] + fn aggregate_parameter_capture_keeps_the_aggregate_address() { + check_capture_program( + r#" + capture(values: [i32; 2]) -> i32 { + var total = 0 + for i in 0 .. 2 { + let read = || { values[0] } + total = total + read() + } + total + } + main() -> i32 { capture([21, 0]) } + "#, + ); + } - let mut vm = VM::new(); - let result = vm.run(&program); - assert_eq!(result, 42); + #[test] + fn borrowed_parameter_capture_keeps_the_borrowed_address() { + check_capture_program( + r#" + capture(value: &i32) -> i32 { + for i in 0 .. 2 { + let increment = || { value = value + 1 } + increment() + } + value + } + main() -> i32 { var value = 40; capture(value) } + "#, + ); } } diff --git a/tests/cases/bytecode/biquad.lyte b/tests/cases/bytecode/biquad.lyte index 62f752fd..2ca15d27 100644 --- a/tests/cases/bytecode/biquad.lyte +++ b/tests/cases/bytecode/biquad.lyte @@ -6,7 +6,7 @@ // expected stdout: // fn main (params: 0, locals: 208 bytes): // 0: SaveRegs { start_reg: 0, count: 16, slot: 80 } -// 1: LocalAddr { dst: 15, slot: 0 } +// 1: LocalAddr { dst: 15, slot: 0 } ; bq // 2: LoadF32 { dst: 1, value: 1000.0 } // 3: LoadF32 { dst: 2, value: 44100.0 } // 4: LoadF32 { dst: 3, value: 0.707 } @@ -16,7 +16,7 @@ // 8: Move { dst: 7, src: 3 } // 9: Call { func: 1, args_start: 4, arg_count: 4 } // 10: LocalAddr { dst: 1, slot: 5 } -// 11: LocalAddr { dst: 15, slot: 0 } +// 11: LocalAddr { dst: 15, slot: 0 } ; bq // 12: MemCopy { dst: 15, src: 1, size: 36 } // 13: LoadImm { dst: 1, value: 10000000 } // 14: LoadF32 { dst: 2, value: 0.0 } @@ -41,7 +41,7 @@ // 33: LoadF32 { dst: 11, value: 1.0 } // 34: FSub { dst: 2, a: 2, b: 11 } // => 35: FMul { dst: 11, a: 3, b: 12 } -// 36: LocalAddr { dst: 15, slot: 0 } +// 36: LocalAddr { dst: 15, slot: 0 } ; bq // 37: Load32Off { dst: 13, base: 15, offset: 20 } // 38: FMulAdd { dst: 14, a: 4, b: 13, c: 11 } // 39: Load32Off { dst: 11, base: 15, offset: 24 } @@ -62,6 +62,20 @@ // 54: Move { dst: 0, src: 1 } // 55: RestoreRegs { start_reg: 1, count: 15, slot: 88 } // 56: Return +// locals: +// r1 n +// r10 i +// r12 x +// r13 y +// r2 phase +// r3 __hoisted_b0 +// r4 __hoisted_b1 +// r5 freq +// r6 two_pi +// r7 __hoisted_b2 +// r8 __hoisted_a1 +// r9 __hoisted_a2 +// slot 0 bq // fn lpf (params: 4, locals: 112 bytes): // 0: SaveRegs { start_reg: 0, count: 8, slot: 48 } // 1: LocalAddr { dst: 4, slot: 0 } @@ -80,7 +94,7 @@ // 14: FAdd { dst: 3, a: 1, b: 2 } // 15: LoadF32 { dst: 1, value: 1.0 } // 16: FDiv { dst: 5, a: 1, b: 3 } -// 17: LocalAddr { dst: 1, slot: 1 } +// 17: LocalAddr { dst: 1, slot: 1 } ; bq // 18: MemZero { dst: 1, size: 36 } // 19: LoadF32 { dst: 3, value: 1.0 } // 20: FSub { dst: 6, a: 3, b: 0 } @@ -114,9 +128,12 @@ // 48: RestoreRegs { start_reg: 0, count: 8, slot: 48 } // 49: Return // locals: -// r5 fc -// r6 fs -// r7 q +// r0 cs +// r1 w0 +// r2 alpha +// r3 a0 +// r5 inv +// slot 1 bq struct Biquad { b0: f32, b1: f32, b2: f32, diff --git a/tests/cases/bytecode/fft.lyte b/tests/cases/bytecode/fft.lyte index 99097c29..793d12f1 100644 --- a/tests/cases/bytecode/fft.lyte +++ b/tests/cases/bytecode/fft.lyte @@ -8,9 +8,9 @@ // 0: SaveRegs { start_reg: 0, count: 16, slot: 8224 } // 1: LoadImm { dst: 10, value: 1024 } // 2: LoadImm { dst: 11, value: 2000 } -// 3: LocalAddr { dst: 12, slot: 0 } +// 3: LocalAddr { dst: 12, slot: 0 } ; re // 4: MemZero { dst: 12, size: 4096 } -// 5: LocalAddr { dst: 13, slot: 512 } +// 5: LocalAddr { dst: 13, slot: 512 } ; im // 6: MemZero { dst: 13, size: 4096 } // 7: LoadF32 { dst: 14, value: 0.0 } // 8: LoadImm { dst: 15, value: 0 } @@ -47,12 +47,12 @@ // 39: Store32 { addr: 5, src: 4 } // 40: IAddImm { dst: 3, src: 3, imm: 1 } // 41: Jump { offset: -29 } -// => 42: LocalAddr { dst: 12, slot: 0 } +// => 42: LocalAddr { dst: 12, slot: 0 } ; re // 43: LocalAddr { dst: 1, slot: 1024 } // 44: Store64 { addr: 1, src: 12 } // 45: LoadImm { dst: 2, value: 1024 } // 46: Store32Off { base: 1, offset: 8, src: 2 } -// 47: LocalAddr { dst: 13, slot: 512 } +// 47: LocalAddr { dst: 13, slot: 512 } ; im // 48: LocalAddr { dst: 2, slot: 1026 } // 49: Store64 { addr: 2, src: 13 } // 50: LoadImm { dst: 3, value: 1024 } @@ -60,7 +60,7 @@ // 52: Move { dst: 3, src: 1 } // 53: Move { dst: 4, src: 2 } // 54: Call { func: 1, args_start: 3, arg_count: 2 } -// 55: LocalAddr { dst: 12, slot: 0 } +// 55: LocalAddr { dst: 12, slot: 0 } ; re // 56: LoadImm { dst: 1, value: 0 } // 57: LoadImm { dst: 2, value: 4 } // 58: IMul { dst: 3, a: 1, b: 2 } @@ -73,6 +73,17 @@ // 65: Move { dst: 0, src: 1 } // 66: RestoreRegs { start_reg: 1, count: 15, slot: 8232 } // 67: Return +// locals: +// r1 freq1 +// r10 n +// r11 iters +// r14 checksum +// r15 iter +// r2 freq2 +// r3 i +// r6 t +// slot 0 re +// slot 512 im // fn fft (params: 2, locals: 184 bytes): // 0: SaveRegs { start_reg: 0, count: 23, slot: 0 } // 1: Move { dst: 19, src: 0 } @@ -144,6 +155,27 @@ // 67: Move { dst: 0, src: 1 } // 68: RestoreRegs { start_reg: 1, count: 22, slot: 8 } // 69: Return +// locals: +// r1 size +// r10 __hoisted_len +// r11 k +// r12 wr +// r13 i +// r13 theta +// r14 wi +// r15 j +// r16 ti +// r17 tr +// r2 __hoisted_len +// r21 n +// r22 pi +// r3 __hoisted_len +// r4 group +// r5 half +// r6 angle +// r7 __hoisted_len +// r8 __hoisted_len +// r9 __hoisted_len // fn bit_reverse_permute (params: 2, locals: 88 bytes): // 0: SaveRegs { start_reg: 0, count: 11, slot: 0 } // 1: Move { dst: 5, src: 0 } @@ -180,6 +212,12 @@ // 32: Move { dst: 0, src: 1 } // 33: RestoreRegs { start_reg: 1, count: 10, slot: 8 } // 34: Return +// locals: +// r1 j +// r2 ti +// r2 tr +// r7 __hoisted_len +// r8 __hoisted_len // fn bit_reverse (params: 2, locals: 64 bytes): // 0: SaveRegs { start_reg: 0, count: 8, slot: 0 } // 1: Move { dst: 2, src: 1 } @@ -200,8 +238,10 @@ // 16: RestoreRegs { start_reg: 1, count: 7, slot: 8 } // 17: Return // locals: +// r1 result // r2 bits -// r2 x +// r3 val +// r4 i bit_reverse(x: i32, bits: i32) -> i32 { var result = 0 diff --git a/tests/cases/bytecode/sort.lyte b/tests/cases/bytecode/sort.lyte index 39d8ebc1..092fcd68 100644 --- a/tests/cases/bytecode/sort.lyte +++ b/tests/cases/bytecode/sort.lyte @@ -33,7 +33,7 @@ main { // 0: SaveRegs { start_reg: 0, count: 11, slot: 40016 } // 1: LoadImm { dst: 6, value: 10000 } // 2: LoadImm { dst: 7, value: 50 } -// 3: LocalAddr { dst: 8, slot: 0 } +// 3: LocalAddr { dst: 8, slot: 0 } ; a // 4: MemZero { dst: 8, size: 40000 } // 5: LoadImm { dst: 9, value: 0 } // 6: LoadImm { dst: 10, value: 0 } @@ -56,7 +56,7 @@ main { // 23: Store32 { addr: 2, src: 3 } // 24: IAddImm { dst: 1, src: 1, imm: 1 } // 25: Jump { offset: -13 } -// => 26: LocalAddr { dst: 8, slot: 0 } +// => 26: LocalAddr { dst: 8, slot: 0 } ; a // 27: LocalAddr { dst: 1, slot: 5000 } // 28: Store64 { addr: 1, src: 8 } // 29: LoadImm { dst: 2, value: 10000 } @@ -65,7 +65,7 @@ main { // 32: Call { func: 1, args_start: 2, arg_count: 1 } // 33: LoadImm { dst: 1, value: 1 } // 34: ILtJump { a: 1, b: 6, offset: 14 } -// 35: LocalAddr { dst: 8, slot: 0 } +// 35: LocalAddr { dst: 8, slot: 0 } ; a // 36: LoadImm { dst: 1, value: 0 } // 37: LoadImm { dst: 2, value: 4 } // 38: IMul { dst: 3, a: 1, b: 2 } @@ -85,6 +85,14 @@ main { // 52: Move { dst: 0, src: 1 } // 53: RestoreRegs { start_reg: 1, count: 10, slot: 40024 } // 54: Return +// locals: +// r1 i +// r10 iter +// r3 seed +// r6 n +// r7 iters +// r9 checksum +// slot 0 a // fn sort$i32 (params: 1, locals: 16 bytes): // 0: SaveRegs { start_reg: 0, count: 2, slot: 0 } // 1: Move { dst: 1, src: 0 } @@ -94,9 +102,9 @@ main { // fn quicksort$i32 (params: 1, locals: 608 bytes): // 0: SaveRegs { start_reg: 0, count: 12, slot: 512 } // 1: Move { dst: 1, src: 0 } -// 2: LocalAddr { dst: 2, slot: 0 } +// 2: LocalAddr { dst: 2, slot: 0 } ; lo_stk // 3: MemZero { dst: 2, size: 256 } -// 4: LocalAddr { dst: 3, slot: 32 } +// 4: LocalAddr { dst: 3, slot: 32 } ; hi_stk // 5: MemZero { dst: 3, size: 256 } // 6: LoadImm { dst: 4, value: 0 } // 7: LoadImm { dst: 5, value: 0 } @@ -117,12 +125,12 @@ main { // 22: ILtJump { a: 5, b: 4, offset: 65 } // 23: LoadImm { dst: 5, value: 1 } // 24: ISub { dst: 4, a: 4, b: 5 } -// 25: LocalAddr { dst: 2, slot: 0 } +// 25: LocalAddr { dst: 2, slot: 0 } ; lo_stk // 26: LoadImm { dst: 5, value: 4 } // 27: IMul { dst: 6, a: 4, b: 5 } // 28: IAdd { dst: 5, a: 2, b: 6 } // 29: Load32 { dst: 6, addr: 5 } -// 30: LocalAddr { dst: 3, slot: 32 } +// 30: LocalAddr { dst: 3, slot: 32 } ; hi_stk // 31: LoadImm { dst: 5, value: 4 } // 32: IMul { dst: 7, a: 4, b: 5 } // 33: IAdd { dst: 5, a: 3, b: 7 } @@ -184,4 +192,19 @@ main { // 89: Move { dst: 0, src: 1 } // 90: RestoreRegs { start_reg: 1, count: 11, slot: 520 } // 91: Return +// locals: +// r10 right_lo +// r10 tmp +// r4 top +// r5 left_size +// r5 pivot_val +// r5 tmp +// r6 cur_lo +// r7 cur_hi +// r8 i +// r8 right_size +// r9 j +// r9 left_hi +// slot 0 lo_stk +// slot 32 hi_stk diff --git a/tests/cases/generics/require_concrete.lyte b/tests/cases/generics/require_concrete.lyte new file mode 100644 index 00000000..2e417ad2 --- /dev/null +++ b/tests/cases/generics/require_concrete.lyte @@ -0,0 +1,34 @@ +// expected stdout: +// compilation successful +// assert(true) + +// Concrete checking covers ordinary callers, explicit type applications, +// wrappers, and implicit array-to-slice and reference conversions. +bounded(x: i32, value: T) require x >= 0 {} + +wrapper(x: i32) require x >= 0 { + bounded(x, true) + bounded⟨bool⟩(x, true) +} + +write(xs: [T], idx: i32, value: T) +require idx >= 0 +require idx < xs.len +{ + xs[idx] = value +} + +bump(x: &i32, value: T) require x >= 0 { + x = x + 1 +} + +main { + bounded(0, true) + bounded⟨bool⟩(0, true) + wrapper(0) + var xs = [1, 2] + var idx = 0 + bump(idx, true) + write⟨i32⟩(xs, idx, 42) + assert(xs[1] == 42) +} diff --git a/tests/cases/generics/require_concrete_unproven.lyte b/tests/cases/generics/require_concrete_unproven.lyte new file mode 100644 index 00000000..a95bc7f2 --- /dev/null +++ b/tests/cases/generics/require_concrete_unproven.lyte @@ -0,0 +1,13 @@ +// args: --check +// expected stdout: +// ❌ ../tests/cases/generics/require_concrete_unproven.lyte:10:8: couldn't prove require clause `x >= 0` for call to `bounded$bool` +// main { bounded(-1, true) } +// ^ + +// The ordinary caller must be checked after the generic target is concrete. +bounded(x: i32, value: T) require x >= 0 {} + +main { bounded(-1, true) } + +// expected stderr: +// safety check failed for 1 call(s) diff --git a/tests/cases/lambdas/identity_captures.lyte b/tests/cases/lambdas/identity_captures.lyte new file mode 100644 index 00000000..06a004d5 --- /dev/null +++ b/tests/cases/lambdas/identity_captures.lyte @@ -0,0 +1,21 @@ +// expected stdout: +// compilation successful +// 42 +// 24 + +make() -> void -> i32 { + let x = 99 + return (|| { let x = 42; x }) +} + +main { + let f = make() + print(f()) + var a = 1 + var b = 2 + let outer = |x: i32| { + let inner = || { b * 10 + a + x } + inner() + } + print(outer(3)) +} diff --git a/tests/cases/references/recorded_callees.lyte b/tests/cases/references/recorded_callees.lyte new file mode 100644 index 00000000..7d675606 --- /dev/null +++ b/tests/cases/references/recorded_callees.lyte @@ -0,0 +1,16 @@ +// expected stdout: +// compilation successful +// 42 +// 42 + +call(a: &i32, b: &i32) { a = b } +set(a: &T, b: &T) { a = b } + +main { + let call = |a: i32, b: i32| { a + b } + print(call(21, 21)) + var a = 0 + var b = 42 + set⟨i32⟩(a, b) + print(a) +} diff --git a/tests/cases/slices/empty_string_bytecode.lyte b/tests/cases/slices/empty_string_bytecode.lyte index 8c0c0855..18d83f26 100644 --- a/tests/cases/slices/empty_string_bytecode.lyte +++ b/tests/cases/slices/empty_string_bytecode.lyte @@ -3,7 +3,7 @@ // expected stdout: // fn main (params: 0, locals: 64 bytes): // 0: SaveRegs { start_reg: 0, count: 5, slot: 24 } -// 1: LocalAddr { dst: 1, slot: 0 } +// 1: LocalAddr { dst: 1, slot: 0 } ; s // 2: LocalAddr { dst: 2, slot: 1 } // 3: LoadImm { dst: 3, value: 0 } // 4: Store8Off { base: 2, offset: 0, src: 3 } @@ -26,6 +26,8 @@ // 21: Move { dst: 0, src: 1 } // 22: RestoreRegs { start_reg: 1, count: 4, slot: 32 } // 23: Return +// locals: +// slot 0 s main { var s = ""