From 9981355fcfba20304a185da755f3bf514360ac05 Mon Sep 17 00:00:00 2001 From: Izaak Branderhorst Date: Wed, 9 Sep 2026 02:48:46 +0200 Subject: [PATCH 1/3] Introduce a checked-program boundary across the compiler Record lexical bindings, solved types, and source locations in owned checked bodies. Specialization produces explicit concrete function/global instances consumed by safety, hoisting, copy elision, Cranelift, LLVM, register VM, and Stack lowering. Replace repeated name lookup, shadowing bookkeeping, and separate capture walkers with recorded identities and shared capture discovery. Retain templates across entry-point changes and specialization retries, validate publication boundaries, and keep partial editor facts separate from executable programs. Check concrete requirements in every retained function body before code motion. Preserve existing syntax and overload/coercion policies. Share borrowed-call classification and add ownership, lifecycle, safety, editor recovery, DSP state, and cancellation regressions. Document the contract in docs/CHECKED_PROGRAM.md. Validated on macOS ARM64 with LLVM 18: default and LLVM workspace builds/tests, backend golden suites, AOT integration tests, and a library check without default features all passed. --- docs/CHECKED_PROGRAM.md | 293 +++ docs/MONOMORPHIZATION.md | 388 ++-- lsp/src/analysis.rs | 2 +- lsp/src/goto_def.rs | 110 +- lsp/src/hover.rs | 50 +- lsp/src/main.rs | 4 +- lsp/src/recovery_tests.rs | 445 ++++ src/README.md | 4 + src/checked.rs | 1076 ++++++++++ src/checked/validate.rs | 1033 +++++++++ src/checker.rs | 1505 +++++++++---- src/compiler.rs | 916 +++++--- src/compiler/assumption_tests.rs | 263 +++ src/compiler/lifecycle_tests.rs | 568 +++++ src/compiler/safety_tests.rs | 244 +++ src/copy_elision.rs | 43 +- src/decl.rs | 121 +- src/decl_table.rs | 363 +++- src/expr.rs | 82 +- src/free_locals.rs | 79 + src/hoist.rs | 1062 +++++----- src/interface_resolution.rs | 194 ++ src/jit.rs | 704 ++----- src/lib.rs | 8 + src/llvm_aot.rs | 148 +- src/llvm_jit.rs | 649 +++--- src/monomorph.rs | 430 +--- src/monomorph_pass.rs | 1867 +++++++---------- src/parser.rs | 3 +- src/safety_checker.rs | 1584 +++++++------- src/solver.rs | 41 +- src/source_analysis.rs | 174 ++ src/stack_codegen.rs | 844 +++----- src/types.rs | 66 +- src/vm.rs | 61 +- src/vm_codegen.rs | 1377 +++++------- tests/cases/bytecode/biquad.lyte | 31 +- tests/cases/bytecode/fft.lyte | 52 +- tests/cases/bytecode/sort.lyte | 37 +- tests/cases/generics/require_concrete.lyte | 34 + .../generics/require_concrete_unproven.lyte | 13 + tests/cases/lambdas/identity_captures.lyte | 21 + tests/cases/references/recorded_callees.lyte | 16 + tests/cases/slices/empty_string_bytecode.lyte | 4 +- 44 files changed, 10391 insertions(+), 6618 deletions(-) create mode 100644 docs/CHECKED_PROGRAM.md create mode 100644 lsp/src/recovery_tests.rs create mode 100644 src/checked.rs create mode 100644 src/checked/validate.rs create mode 100644 src/compiler/assumption_tests.rs create mode 100644 src/compiler/lifecycle_tests.rs create mode 100644 src/compiler/safety_tests.rs create mode 100644 src/free_locals.rs create mode 100644 src/interface_resolution.rs create mode 100644 src/source_analysis.rs create mode 100644 tests/cases/generics/require_concrete.lyte create mode 100644 tests/cases/generics/require_concrete_unproven.lyte create mode 100644 tests/cases/lambdas/identity_captures.lyte create mode 100644 tests/cases/references/recorded_callees.lyte diff --git a/docs/CHECKED_PROGRAM.md b/docs/CHECKED_PROGRAM.md new file mode 100644 index 00000000..a1e13a94 --- /dev/null +++ b/docs/CHECKED_PROGRAM.md @@ -0,0 +1,293 @@ +# 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. + +Each `CheckedNode` owns its operation, result type and source location. 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; `replace_node` can also change provenance. + +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..2a89851b 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 operation/type/location nodes, 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 +expressions share syntax shape and child traversal with distinct reference, +binder and parameter payloads. + +`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..dd63cfbe --- /dev/null +++ b/src/checked.rs @@ -0,0 +1,1076 @@ +//! The program after lexical and type checking. +//! +//! Nodes own their operation, result type and source provenance. 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 CheckedExpr = Expr; +pub type CheckedDecl = Decl; +pub type CheckedDeclTable = DeclTable; +pub type CheckedDeclarations = DeclarationList; + +#[derive(Clone, Debug, Eq, PartialEq, Hash)] +pub struct CheckedNode { + pub kind: CheckedExpr, + pub ty: TypeID, + pub loc: Loc, +} + +#[derive(Clone, Debug, Default, Eq, PartialEq, Hash)] +pub struct CheckedBody { + nodes: Vec, + pub locals: Vec, + pub requirements: Vec, +} + +impl CheckedBody { + pub fn new() -> Self { + Self::default() + } + pub fn from_parts( + nodes: Vec, + locals: Vec, + requirements: Vec, + ) -> Self { + Self { + nodes, + locals, + requirements, + } + } + pub fn len(&self) -> usize { + self.nodes.len() + } + pub fn is_empty(&self) -> bool { + self.nodes.is_empty() + } + pub fn node(&self, id: ExprID) -> &CheckedNode { + &self.nodes[id] + } + pub fn nodes(&self) -> &[CheckedNode] { + &self.nodes + } + pub fn ty(&self, id: ExprID) -> TypeID { + self.nodes[id].ty + } + pub fn loc(&self, id: ExprID) -> Loc { + self.nodes[id].loc + } + 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 + } + pub fn add(&mut self, kind: CheckedExpr, ty: TypeID, loc: Loc) -> ExprID { + let id = self.nodes.len(); + self.nodes.push(CheckedNode { kind, ty, loc }); + id + } + /// Replacing a node retains its handle and provenance, not any analysis of + /// the previous node. Callers provide the replacement's checked result type. + pub fn replace(&mut self, id: ExprID, kind: CheckedExpr, ty: TypeID) { + self.nodes[id].kind = kind; + self.nodes[id].ty = ty; + } + pub fn replace_node(&mut self, id: ExprID, node: CheckedNode) { + self.nodes[id] = node; + } + pub fn substitute(&mut self, instance: &Instance) { + for node in &mut self.nodes { + node.ty = node.ty.subst(instance); + match &mut node.kind { + 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 nodes; spellings are presentation + /// data and are never converted back into unresolved syntax. + pub fn pretty_print(&self, id: ExprID, indent: usize) -> String { + self.pretty_print_with(id, indent, &|reference| match reference { + Reference::Local(local) | Reference::SizeParameter(local) => self.local(*local).name, + Reference::Global(id) => Name::new(format!("global#{}", id.0)), + Reference::Functions(ids) => Name::new(format!("function#{:?}", ids)), + Reference::InterfaceMember { member, .. } => Name::new(format!("member#{}", member.0)), + Reference::Instance(id) => Name::new(format!("instance#{}", id.0)), + }) + } + pub fn pretty_print_with( + &self, + id: ExprID, + indent: usize, + reference_name: &impl Fn(&Reference) -> Name, + ) -> String { + let child = |id| self.pretty_print_with(id, indent, reference_name); + let list = |ids: &[ExprID]| { + ids.iter() + .map(|id| child(*id)) + .collect::>() + .join(", ") + }; + match &self[id] { + Expr::Id(reference) => reference_name(reference).to_string(), + Expr::TypeApp(reference, args) => format!( + "{}⟨{}⟩", + reference_name(reference), + args.iter() + .map(|ty| ty.pretty_print()) + .collect::>() + .join(", ") + ), + Expr::Int(value, suffix) => format!( + "{}{}", + value, + suffix.map(|suffix| suffix.to_string()).unwrap_or_default() + ), + Expr::Real(value, suffix) => format!( + "{}{}", + value, + suffix.map(|suffix| suffix.to_string()).unwrap_or_default() + ), + Expr::String(value) => format!("\"{}\"", value), + Expr::Char(value) => format!("'{}'", value), + Expr::True => "true".into(), + Expr::False => "false".into(), + Expr::Enum(name) => format!(".{}", name), + Expr::Error => "".into(), + Expr::Call(function, args) => format!("{}({})", child(*function), list(args)), + Expr::Macro(name, args) => format!("@{}({})", name, list(args)), + Expr::Binop(op, lhs, rhs) => { + format!("{} {} {}", child(*lhs), format_binop(*op), child(*rhs)) + } + Expr::Unop(op, value) => format!("{}{}", format_unop(*op), child(*value)), + Expr::Lambda { params, body } => format!( + "|{}| {}", + params + .iter() + .map(|param| format!( + "{}: {}", + self.local(param.local).name, + self.local(param.local).ty.pretty_print() + )) + .collect::>() + .join(", "), + child(*body) + ), + Expr::Field(base, name) => format!("{}.{}", child(*base), name), + Expr::Array(element, size) => format!("[{}; {}]", child(*element), child(*size)), + Expr::ArrayLiteral(elements) => format!("[{}]", list(elements)), + Expr::ArrayIndex(array, index) => format!("{}[{}]", child(*array), child(*index)), + Expr::AsTy(value, ty) => format!("{}:{}", child(*value), ty.pretty_print()), + Expr::Let(local, init, annotation) => { + let annotation = annotation + .map(|ty| format!(": {}", ty.pretty_print())) + .unwrap_or_default(); + format!( + "let {}{} = {}", + self.local(*local).name, + annotation, + child(*init) + ) + } + Expr::Var(local, init, annotation) => { + let annotation = annotation + .map(|ty| format!(": {}", ty.pretty_print())) + .unwrap_or_default(); + let init = init + .map(|init| format!(" = {}", child(init))) + .unwrap_or_default(); + format!("var {}{}{}", self.local(*local).name, annotation, init) + } + Expr::If(cond, yes, no) => { + let no = no + .map(|no| { + format!( + " else {}", + self.pretty_print_with(no, indent + 1, reference_name) + ) + }) + .unwrap_or_default(); + format!( + "if {} {}{}", + child(*cond), + self.pretty_print_with(*yes, indent + 1, reference_name), + no + ) + } + Expr::While(cond, body) => format!( + "while {} {}", + child(*cond), + self.pretty_print_with(*body, indent + 1, reference_name) + ), + Expr::For { + var, + start, + end, + body, + } => format!( + "for {} in {} .. {} {}", + self.local(*var).name, + child(*start), + child(*end), + self.pretty_print_with(*body, indent + 1, reference_name) + ), + Expr::Block(exprs) => { + if exprs.is_empty() { + return "{}".into(); + } + let expressions = exprs + .iter() + .map(|expr| { + format!( + "{}{}", + " ".repeat(indent + 1), + self.pretty_print_with(*expr, indent + 1, reference_name) + ) + }) + .collect::>() + .join("\n"); + format!("{{\n{}\n{}}}", expressions, " ".repeat(indent)) + } + Expr::Return(value) => format!("return {}", child(*value)), + Expr::Assume(value) => format!("assume {}", child(*value)), + Expr::Break => "break".into(), + Expr::Continue => "continue".into(), + Expr::Tuple(values) => format!("({})", list(values)), + Expr::StructLit(name, fields) => format!( + "{}({})", + name, + fields + .iter() + .map(|(name, value)| format!("{}: {}", name, child(*value))) + .collect::>() + .join(", ") + ), + Expr::Arena(value) => format!("arena {}", child(*value)), + } + } + + /// 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) { + match &body[id] { + Expr::Let(local, ..) | Expr::Var(local, ..) | Expr::For { var: local, .. } => { + found.insert(*local); + } + Expr::Lambda { params, .. } => { + found.extend(params.iter().map(|param| param.local)); + } + _ => {} + } + 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 mut node = body.node(id).clone(); + let remap = |local: &mut LocalId| { + if let Some(new) = locals.get(local) { + *local = *new; + } + }; + match &mut node.kind { + Expr::Id(Reference::Local(local)) + | Expr::TypeApp(Reference::Local(local), _) + | Expr::Let(local, ..) + | Expr::Var(local, ..) + | Expr::For { var: local, .. } => remap(local), + Expr::Lambda { params, .. } => { + for param in params { + remap(&mut param.local); + } + } + _ => {} + } + node.kind.map_children(|child| copy(body, child, locals)); + body.add(node.kind, node.ty, node.loc) + } + 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, params: &[CheckedParam]) -> Vec { + crate::free_locals::free_locals(self, root, params.iter().map(|param| param.local)).locals + } + pub fn captured_locals(&self) -> HashSet { + let mut captured = HashSet::new(); + for node in &self.nodes { + if let Expr::Lambda { params, body } = &node.kind { + captured.extend(self.captures(*body, params)); + } + } + captured + } +} +impl Index for CheckedBody { + type Output = CheckedExpr; + fn index(&self, id: ExprID) -> &Self::Output { + &self.nodes[id].kind + } +} + +#[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 { params, body } = &self.arena[expression] else { + panic!("expected lambda"); + }; + let Type::Func(_, ret) = &*self.arena.ty(expression) else { + panic!("checked lambda type"); + }; + Self { + name, + typevars: Vec::new(), + size_vars: Vec::new(), + params: params.clone(), + body: Some(*body), + ret: *ret, + requires: Vec::new(), + loc: self.arena.loc(expression), + arena: self.arena.clone(), + closure_vars: self.arena.captures(*body, params), + 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(Expr::Id(Reference::Local(outer)), ty, test_loc()); + let binding = body.add( + Expr::Let(inner, outer_read, None), + mk_type(Type::Void), + test_loc(), + ); + let inner_read = body.add(Expr::Id(Reference::Local(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(fresh, initializer, _) = body[children[0]] else { + panic!() + }; + assert_ne!(fresh, inner); + assert_eq!(body[initializer], Expr::Id(Reference::Local(outer))); + assert_eq!(body[children[1]], Expr::Id(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[inner_read], Expr::Id(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(Expr::Id(Reference::Local(outer)), callable, test_loc()); + let read_inner = body.add(Expr::Id(Reference::Local(inner)), ty, test_loc()); + let applied_outer = body.add( + Expr::TypeApp(Reference::Local(outer), vec![]), + callable, + test_loc(), + ); + let nested_root = body.add( + Expr::Block(vec![read_inner, applied_outer, read_outer, read_inner]), + ty, + test_loc(), + ); + let nested = body.add( + Expr::Lambda { + params: vec![], + body: nested_root, + }, + 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, &[CheckedParam { local: 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(Expr::Id(Reference::Local(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..4c997f91 --- /dev/null +++ b/src/checked/validate.rs @@ -0,0 +1,1033 @@ +//! 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, node) in body.nodes().iter().enumerate() { + self.node(node, body, &bindings) + .map_err(|error| format!("expression {}: {}", id, error))?; + } + Ok(()) + } + + fn node( + &self, + node: &CheckedNode, + body: &CheckedBody, + bindings: &[Option], + ) -> Result<(), String> { + validate_type(node.ty, self.concrete())?; + match &node.kind { + Expr::Id(reference) => self.reference(reference, body, bindings)?, + Expr::TypeApp(reference, args) => { + self.reference(reference, body, bindings)?; + for &ty in args { + validate_type(ty, self.concrete())?; + } + } + 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 node.ty != mk_type(Type::Void) { + return Err("checked declaration result type is not void".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 let Expr::Id(reference @ Reference::Instance(instance)) + | Expr::TypeApp(reference @ Reference::Instance(instance), _) = + &body[*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 node in body.nodes() { + match &node.kind { + Expr::Let(local, ..) | Expr::Var(local, ..) | Expr::For { var: local, .. } => { + bind(*local, BinderKind::Value)? + } + Expr::Lambda { params, .. } => { + for param in params { + bind(param.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 node in body.nodes() { + for child in node.kind.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] = match &body[id] { + Expr::Let(..) | Expr::Var(..) | Expr::For { .. } => true, + Expr::Lambda { params, .. } if !params.is_empty() => true, + _ => 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(Expr::Id(Reference::Local(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(Expr::Id(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(Expr::Id(Reference::Local(local)), f.ret, test_loc()); + }, + "no value binder", + ), + ( + |f| { + f.arena.add( + Expr::Lambda { + params: vec![CheckedParam { + local: LocalId(999), + }], + body: 0, + }, + 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 + .nodes() + .iter() + .position(|node| matches!(node.kind, 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 reference = Reference::Instance(target_id); + let kind = if type_application { + Expr::TypeApp(reference, vec![]) + } else { + Expr::Id(reference) + }; + // Validate the forward reference before looking up its target. + assert_eq!(caller.arena.add(kind, recorded_type, test_loc()), callee); + 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.nodes().iter().any(|node| { + matches!(&node.kind, Expr::Id(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( + Expr::Let(local, read, None), + mk_type(Type::Void), + 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(Expr::Id(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( + Expr::Id(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( + Expr::Id(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 + .nodes() + .iter() + .any(|node| matches!(node.kind, Expr::Id(Reference::Local(_))) + && node.ty == 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 + .nodes() + .iter() + .position(|node| matches!(node.kind, 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..c44009f7 100644 --- a/src/checker.rs +++ b/src/checker.rs @@ -1,5 +1,27 @@ +use crate::checked::{ + CheckedBody, CheckedExpr, CheckedFunction, CheckedNode, 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 +33,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 +76,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 +100,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 +235,7 @@ impl Checker { Self { types: vec![], visited: vec![], + independent_types: vec![], value_pos: vec![], lvalue: vec![], inst: Instance::new(), @@ -214,9 +250,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 +364,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 +454,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 +544,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 +576,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 +588,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 +630,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 +650,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 +670,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 +838,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 +862,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 +1116,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 +1138,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 +1170,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 +1198,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 +1279,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 +1383,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 +1440,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 +1506,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 +1531,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 +1560,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 +1571,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 +1643,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 +1660,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 +1695,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 +1706,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 +1718,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 +1734,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 +1748,248 @@ 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 nodes = Vec::with_capacity(source.exprs.len()); + for (id, expression) in source.exprs.iter().enumerate() { + let resolution = || self.references[id].clone().expect("checked reference"); + let binder = || *self.binders[id].first().expect("checked binder"); + let kind = match expression { + Expr::Id(_) => CheckedExpr::Id(resolution()), + Expr::TypeApp(_, args) => CheckedExpr::TypeApp( + resolution(), + args.iter().map(|t| t.subst(&self.inst)).collect(), + ), + Expr::Let(_, init, _) => CheckedExpr::Let(binder(), *init, None), + Expr::Var(_, init, _) => CheckedExpr::Var(binder(), *init, None), + Expr::For { + start, end, body, .. + } => CheckedExpr::For { + var: binder(), + start: *start, + end: *end, + body: *body, + }, + Expr::Lambda { body, .. } => CheckedExpr::Lambda { + params: self.binders[id] + .iter() + .map(|&local| CheckedParam { local }) + .collect(), + body: *body, + }, + Expr::Int(v, s) => CheckedExpr::Int(*v, *s), + Expr::Real(v, s) => CheckedExpr::Real(v.clone(), *s), + Expr::Call(f, a) => CheckedExpr::Call(*f, a.clone()), + Expr::Binop(op, a, b) => CheckedExpr::Binop(*op, *a, *b), + Expr::Unop(op, a) => CheckedExpr::Unop(*op, *a), + Expr::String(v) => CheckedExpr::String(v.clone()), + Expr::Char(v) => CheckedExpr::Char(*v), + Expr::Field(a, n) => CheckedExpr::Field(*a, *n), + Expr::Array(a, b) => CheckedExpr::Array(*a, *b), + Expr::ArrayLiteral(a) => CheckedExpr::ArrayLiteral(a.clone()), + Expr::ArrayIndex(a, b) => CheckedExpr::ArrayIndex(*a, *b), + Expr::True => CheckedExpr::True, + Expr::False => CheckedExpr::False, + Expr::AsTy(a, t) => CheckedExpr::AsTy(*a, t.subst(&self.inst)), + Expr::If(a, b, c) => CheckedExpr::If(*a, *b, *c), + Expr::While(a, b) => CheckedExpr::While(*a, *b), + Expr::Block(v) => CheckedExpr::Block(v.clone()), + Expr::Return(a) => CheckedExpr::Return(*a), + Expr::Break => CheckedExpr::Break, + Expr::Continue => CheckedExpr::Continue, + Expr::Enum(n) => CheckedExpr::Enum(*n), + Expr::Tuple(v) => CheckedExpr::Tuple(v.clone()), + Expr::StructLit(n, v) => CheckedExpr::StructLit(*n, v.clone()), + Expr::Arena(a) => CheckedExpr::Arena(*a), + Expr::Assume(a) => CheckedExpr::Assume(*a), + Expr::Macro(..) | Expr::Error => { + panic!("unexpanded or invalid source in checked body") + } + }; + let ty = if matches!(expression, Expr::Let(..) | Expr::Var(..)) { + mk_type(Type::Void) + } else { + solved[id] + }; + nodes.push(CheckedNode { + kind, + ty, + loc: source.locs[id], + }); + } + // Operator overloading is lowered while publishing checked meaning, + // rather than rewriting source syntax after its types have been solved. + for id in 0..nodes.len() { + if let CheckedExpr::Binop(op, lhs, rhs) = nodes[id].kind.clone() { + if op.arithmetic() + && (matches!(*nodes[lhs].ty, Type::Name(_, _)) + || (op == Binop::Mod + && matches!(*nodes[lhs].ty, Type::Float32 | Type::Float64))) + { + let Some(reference) = self.references[id].clone() else { + continue; + }; + let callee = nodes.len(); + nodes.push(CheckedNode { + kind: CheckedExpr::Id(reference), + ty: func(tuple(vec![nodes[lhs].ty, nodes[rhs].ty]), nodes[id].ty), + loc: nodes[id].loc, + }); + nodes[id].kind = CheckedExpr::Call(callee, vec![lhs, rhs]); + } + } + } + CheckedBody::from_parts(nodes, locals, self.requirements.clone()) + } + + /// 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 +2071,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], +} + +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 + } } -/// 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 +#[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 +2354,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..776e98bf 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,33 @@ 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 + .nodes() + .iter() + .filter_map(|node| { + if let CheckedExpr::Id(Reference::Instance(target)) = node.kind { + 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 +2310,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 +2320,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 +2331,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 +2521,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 +2646,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 +2821,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 +2886,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 +2921,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 +2952,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..72467ba9 --- /dev/null +++ b/src/compiler/assumption_tests.rs @@ -0,0 +1,263 @@ +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 CheckedExpr::Block(statements) = &body[root] else { + panic!() + }; + assert!(matches!(body[statements[0]], Expr::Let(local, ..) if local == expected)); + let Expr::Binop(Binop::Geq, value, _) = body[statements[1]] else { + panic!() + }; + assert_eq!(body[value], Expr::Id(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 node in body.nodes() { + if let Expr::Id(Reference::Instance(target)) = node.kind { + 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.nodes().iter().any(|node| match node.kind { + Expr::Id(Reference::Instance(target)) => + output.instances[target.index()].definition == positive, + _ => false, + })); + assert!(body + .nodes() + .iter() + .any(|node| matches!(node.kind, Expr::Id(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 + .nodes() + .iter() + .find(|node| matches!(node.kind, Expr::Binop(Binop::Div, ..))) + .unwrap(); + assert_eq!(compiler.last_safety_errors.len(), 1); + assert_eq!(compiler.last_safety_errors[0].location, division.loc); + assert_eq!(division.loc.file, Name::str("")); + assert_eq!(division.loc, 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..8a66aeee --- /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 node in main.arena.nodes() { + if let CheckedExpr::Call(callee, _) = &node.kind { + match &main.arena[*callee] { + CheckedExpr::Id(Reference::Local(_)) => local_calls += 1, + CheckedExpr::Id(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..69cbdedf 100644 --- a/src/copy_elision.rs +++ b/src/copy_elision.rs @@ -18,9 +18,8 @@ //! 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::{CheckedExpr as Expr, CheckedFunction, LocalId, Reference}; +use crate::defs::{Binop, ExprID}; use crate::types::{Type, TypeID}; use std::collections::HashSet; @@ -39,7 +38,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 +46,25 @@ 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(..)) { + let Expr::Let(local, ..) = &decl.arena[stmt] else { continue; - } - if !is_value_aggregate(&decl.types[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 +73,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 +94,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 +114,13 @@ 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] { +fn mentions(id: ExprID, name: LocalId, decl: &CheckedFunction) -> bool { + if let Expr::Id(Reference::Local(n)) = &decl.arena[id] { if *n == 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..079be787 100644 --- a/src/expr.rs +++ b/src/expr.rs @@ -8,9 +8,9 @@ use crate::*; /// tree. It's also faster. Most hierarchical data /// should be represented this way. #[derive(Clone, Debug, Eq, PartialEq, Hash)] -pub enum Expr { +pub enum Expr { /// Identifier expression. - Id(Name), + Id(R), /// Integer literal, with optional explicit suffix. Int(i64, Option), @@ -31,7 +31,7 @@ pub enum Expr { Unop(Unop, ExprID), /// Lambda expression with parameters and body. - Lambda { params: Vec, body: ExprID }, + Lambda { params: Vec

, body: ExprID }, /// String literal. String(String), @@ -61,13 +61,13 @@ pub enum Expr { AsTy(ExprID, TypeID), /// Explicit type application: `name⟨i32⟩` or `name⟨i32, f32⟩`. - TypeApp(Name, Vec), + TypeApp(R, Vec), /// Immutable variable declaration with initializer and optional type. - Let(Name, ExprID, Option), + Let(B, ExprID, Option), /// Mutable variable declaration with optional initializer and type. - Var(Name, Option, Option), + Var(B, Option, Option), /// If expression with optional else branch. If(ExprID, ExprID, Option), @@ -77,7 +77,7 @@ pub enum Expr { /// For loop expression. For { - var: Name, + var: B, start: ExprID, end: ExprID, body: ExprID, @@ -115,7 +115,7 @@ pub enum Expr { Error, } -impl Expr { +impl Expr { /// The immediate subexpression IDs of this expression. /// /// Lambda yields its body: walks that treat a lambda specially still need @@ -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..30f7ca15 --- /dev/null +++ b/src/free_locals.rs @@ -0,0 +1,79 @@ +//! Capture discovery needs lexical bindings, not solved types or publication. +use crate::{CheckedBody, Expr, 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 { + let mut node = BindingNode { + children: self[id].subexprs(), + used: None, + declared: vec![], + complete: true, + }; + match &self[id] { + Expr::Id(Reference::Local(local)) | Expr::TypeApp(Reference::Local(local), _) => { + node.used = Some(*local); + } + Expr::Let(local, ..) | Expr::Var(local, ..) | Expr::For { var: local, .. } => { + node.declared.push(*local); + } + Expr::Lambda { params, .. } => { + node.declared.extend(params.iter().map(|param| param.local)); + } + _ => {} + } + node + } +} diff --git a/src/hoist.rs b/src/hoist.rs index 1be90b65..dec16cd0 100644 --- a/src/hoist.rs +++ b/src/hoist.rs @@ -1,703 +1,583 @@ 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[expr] { + Expr::Id(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(Reference::Local(local)) => Some(Root::Local(*local)), + Expr::Id(Reference::Instance(instance)) => Some(Root::Global(*instance)), + 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, + 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; + } }; - - 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 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) } - - // 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); + 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); } + _ => 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)); +} - // 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); - } +fn create_hoisted_binding(read: &FieldRead, arena: &mut CheckedBody) -> (LocalId, ExprID) { + let Expr::Field(base, _) = arena[read.expr] else { + unreachable!(); + }; + let source_base = arena.node(base).clone(); + let source_field = arena.node(read.expr).clone(); + // These are fresh evaluations with fresh ExprIDs, while the copied + // outer reference retains its LocalId or global InstanceId. + let base = arena.add(source_base.kind, source_base.ty, source_base.loc); + let initializer = arena.add( + Expr::Field(base, read.field), + source_field.ty, + source_field.loc, + ); + let local = arena.add_local( + Name::new(format!("__hoisted_{}", read.field)), + source_field.ty, + false, + ); + let declaration = arena.add( + Expr::Let(local, initializer, None), + mk_type(Type::Void), + source_field.loc, + ); + (local, declaration) +} - // 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); - } - } +fn invalidate_binders(expr: ExprID, function: &CheckedFunction, written: &mut WrittenFields) { + match &function.arena[expr] { + Expr::Let(local, ..) | Expr::Var(local, ..) | Expr::For { var: local, .. } => { + written.insert((Root::Local(*local), None)); } - - new_stmts.push(stmt_id); - } - - if new_stmts.len() != stmts.len() { - fdecl.arena.exprs[block_id] = Expr::Block(new_stmts); + Expr::Lambda { params, .. } => { + written.extend( + params + .iter() + .map(|parameter| (Root::Local(parameter.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); + } +} + +/// 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 + } + }) +} - for sub in fdecl.arena.exprs[expr_id].subexprs() { - collect_written_fields(sub, fdecl, effects, names, written); +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: &CheckedExpr) -> 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))) + { + function.arena.replace( + expr, + Expr::Id(Reference::Local(*local)), + function.arena.ty(expr), + ); + 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 + .nodes() + .iter() + .any(|node| node.kind == Expr::Id(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); + 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..a707b5d8 100644 --- a/src/jit.rs +++ b/src/jit.rs @@ -2,11 +2,11 @@ // Pulled from https://github.com/bytecodealliance/cranelift-jit-demo use crate::cancel::*; -use crate::decl::*; +use crate::checked::{ + CheckedDecl as Decl, CheckedExpr as Expr, CheckedFunction as FuncDecl, InstanceId, LocalId, + Reference, SpecializedProgram as DeclTable, +}; use crate::defs::*; -use crate::expr::*; -use crate::DeclTable; -use crate::Instance; use crate::TypeID; extern crate cranelift_codegen; use core::panic; @@ -93,10 +93,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 +187,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 +209,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 +256,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 +271,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 +289,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 +311,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 +378,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 +543,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 +584,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 +603,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 +619,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 +749,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 +776,18 @@ 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. - 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 - ); - } + Expr::Id(Reference::Local(local)) => self.builder.use_var(self.variables[local]), + Expr::Id(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) } 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 +808,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,36 +846,32 @@ 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) { - 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. + Expr::Id(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) + } + } + Expr::Id(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); - // Composite types (arrays, structs) are pointer-represented: - // return the address, don't load. f32x4 is a value type — load it. - if is_indirect(*ty) { + if is_indirect(ty) { addr } else { self.builder @@ -892,9 +879,10 @@ impl<'a> FunctionTranslator<'a> { .load(ty.cranelift_type(), MemFlags::new(), addr, 0) } } else { - self.translate_func(name, &*ty) + self.translate_func(*instance, &*ty, decls) } } + Expr::Id(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 +890,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 Expr::Id(Reference::Instance(instance)) = &decl.arena[*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 Expr::Id(Reference::Instance(instance)) = &decl.arena[*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 +911,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 +937,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 Expr::Id(Reference::Instance(callee)) = &decl.arena[*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 +952,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 Expr::Id(Reference::Instance(instance)) = &decl.arena[*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 Expr::Id(Reference::Instance(callee)) = &decl.arena[*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 +988,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 Expr::Id(Reference::Instance(n)) = &decl.arena[*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 +1017,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 +1033,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 { @@ -1095,7 +1083,7 @@ impl<'a> FunctionTranslator<'a> { // it can write trap_reason and longjmp on failure. let is_assert = matches!( &decl.arena[*fn_id], - Expr::Id(name) if **name == "assert" + Expr::Id(Reference::Instance(instance)) if *decls.instance_name(*instance) == "assert" ); let f = self.translate_expr(*fn_id, decl, decls); let mut args = vec![]; @@ -1168,12 +1156,12 @@ 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]; + 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); @@ -1196,31 +1184,30 @@ 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()); 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()); + 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 init_val = if let Some(init_id) = init { @@ -1230,13 +1217,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); + self.variable_types.insert(*name, *ty); let sz = ty.size(decls) as u32; if sz == 0 { @@ -1270,7 +1257,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 +1270,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 +1280,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 +1295,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 +1325,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 +1333,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 +1357,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 +1375,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 +1421,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 +1437,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 +1464,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 +1496,7 @@ impl<'a> FunctionTranslator<'a> { } else { panic!( "JIT array fill: expected array type, got {:?}", - decl.types[expr] + decl.arena.ty(expr) ); } } @@ -1517,11 +1504,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 +1513,6 @@ impl<'a> FunctionTranslator<'a> { break; } } - self.variables = saved_vars; - self.variable_types = saved_types; - self.let_bindings = saved_lets; result.unwrap() } } @@ -1548,13 +1528,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, }; @@ -1628,18 +1608,14 @@ 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. // Create a variable for the loop counter. 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 +1672,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 +1680,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 +1783,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 +1811,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 +1837,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 +1852,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 +1889,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 +1927,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 +1984,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 +2007,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 +2027,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 +2064,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 +2074,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 +2093,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 +2104,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 +2164,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 +2186,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 +2202,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 +2217,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 +2232,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 +2246,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 +2264,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 +2292,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 Expr::Id(Reference::Local(name)) = &decl.arena[lhs_id] { + if let Some(&var) = self.variables.get(name) { self.builder.def_var(var, rhs); return rhs; } @@ -2387,13 +2312,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 +2341,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 +2359,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 +2378,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 +2398,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 +2445,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 +2488,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,14 +2523,15 @@ 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; - } - return Some(self.builder.use_var(var)); + Expr::Id(Reference::Local(local)) => { + if self.let_bindings.contains(local) { + None + } else { + Some(self.builder.use_var(self.variables[local])) } - let offset = *self.globals.get(name)?; + } + Expr::Id(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)) } @@ -2627,8 +2559,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 Expr::Id(Reference::Local(name)) = &decl.arena[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 +2590,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 Expr::Id(Reference::Local(name)) = &decl.arena[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 +2617,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 +2642,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 +2684,23 @@ 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 + match &decl.arena[expr] { + Expr::Id(Reference::Local(local)) => self .variable_types - .get(name.as_str()) + .get(local) .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]), + .unwrap_or(decl.arena.local(*local).ty), + Expr::Id(Reference::Instance(instance)) => match decls.instance(*instance) { + Decl::Global { ty, .. } => *ty, + _ => 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 +2710,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 +2720,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..e62fda9b 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, CheckedExpr as Expr, + CheckedFunction as FuncDecl, CheckedParam as Param, InstanceId, LocalId, Reference, + SpecializedProgram as DeclTable, +}; use crate::defs::*; -use crate::expr::*; -use crate::DeclTable; 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) { + for instance in called { + if self.defined_functions.contains(&instance) { continue; } - let found = decls.find(name); - if found.is_empty() { - 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,23 @@ 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 + match &decl.arena[expr] { + Expr::Id(Reference::Local(local)) => self .variable_types - .get(name.as_str()) + .get(local) .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]), + .unwrap_or(decl.arena.local(*local).ty), + Expr::Id(Reference::Instance(instance)) => match self.decls.instance(*instance) { + Decl::Global { ty, .. } => *ty, + _ => 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 +1697,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 +1709,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 Expr::Id(Reference::Local(name)) = &arena[*lhs] { return Some((*name, *rhs)); } } @@ -1724,10 +1720,10 @@ 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]; + Expr::Id(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. @@ -1739,7 +1735,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { alloca } } else { - if let Some(ty) = self.variable_types.get(name.as_str()).copied() { + 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") @@ -1758,19 +1754,20 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { .into_pointer_value() } } - } 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) } } + Expr::Id(Reference::Instance(instance)) => { + self.ptr_at_offset(self.globals_base, self.state.globals[instance] as u64) + } 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,9 +1782,9 @@ 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) { + Expr::Id(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); } @@ -1800,7 +1797,10 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { .into_pointer_value(), ); } - let offset = *self.state.globals.get(name)?; + None + } + Expr::Id(Reference::Instance(instance)) => { + let offset = *self.state.globals.get(instance)?; Some(self.ptr_at_offset(self.globals_base, offset as u64)) } Expr::Field(_, _) | Expr::ArrayIndex(_, _) => Some(self.translate_lvalue(expr, decl)), @@ -1837,7 +1837,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,33 +1898,37 @@ 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) { + Expr::Id(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, &**name) + .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.as_str()).copied() { + if let Some(var_ty) = self.variable_types.get(name).copied() { if is_indirect(var_ty) { stored } else { @@ -1932,7 +1936,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { .build_load( ty.llvm_basic_type(self.ctx()), stored.into_pointer_value(), - &**name, + &decl.arena.local(*name).name, ) .unwrap() } @@ -1941,27 +1945,31 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { .build_load( ty.llvm_basic_type(self.ctx()), stored.into_pointer_value(), - &**name, + &decl.arena.local(*name).name, ) .unwrap() } } - } else if let Some(&offset) = self.state.globals.get(name) { + } else { + panic!("missing local storage: {:?}", name) + } + } + Expr::Id(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); - // 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) + .build_load(ty.llvm_basic_type(self.ctx()), addr, "global") .unwrap() } } else { - // Must be a function reference. - self.translate_func_ref(name, &*ty) + self.translate_func_ref(*instance, &*ty) } } + Expr::Id(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) @@ -1976,7 +1984,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } Expr::Let(name, init_id, _) => { let (name, init_id) = (*name, *init_id); - let ty = decl.types[expr]; + 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 +2000,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]; + 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 +2040,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 +2068,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 +2106,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 +2118,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 +2135,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 +2146,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 +2168,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 +2176,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 +2187,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 +2221,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(_) @@ -2313,20 +2318,17 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { 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 +2384,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 +2393,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 +2418,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)) @@ -2445,11 +2443,10 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } 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 +2484,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 +2506,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 +2518,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 +2528,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 +2539,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 +2569,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 +2607,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 +2623,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 +2652,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 +2667,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 +2691,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 +2718,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 +2751,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 +2774,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 +2795,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 +2833,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 +2882,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 +2919,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 +2956,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 +3131,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 Expr::Id(Reference::Instance(instance)) = &decl.arena[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 Expr::Id(Reference::Instance(callee)) = &decl.arena[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 +3164,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 Expr::Id(Reference::Instance(instance)) = &decl.arena[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 +3179,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 +3195,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) + let math_intrinsic = if let Expr::Id(Reference::Instance(instance)) = &decl.arena[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 Expr::Id(Reference::Instance(instance)) = &decl.arena[fn_id] { + self.math_builtin_name(&self.decls.instance_name(*instance), from) } else { None } @@ -3244,17 +3242,16 @@ 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 Expr::Id(Reference::Instance(callee)) = &decl.arena[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 Expr::Id(Reference::Instance(n)) = &decl.arena[fn_id] { *n } else { unreachable!() @@ -3278,16 +3275,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 +3298,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 +3354,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[fn_id], Expr::Id(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 +3745,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 +3802,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 +3847,27 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { fn translate_lambda( &mut self, - params: &[Param], - body: ExprID, + _params: &[Param], + _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 +3876,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 +3897,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..50975a65 100644 --- a/src/monomorph_pass.rs +++ b/src/monomorph_pass.rs @@ -1,1282 +1,887 @@ 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 (reference, explicit) = match body[id].clone() { + Expr::Id(reference) => (reference, None), + Expr::TypeApp(reference, arguments) => (reference, Some(arguments)), + _ => return Err("Expected a checked reference".into()), + }; + 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)?; + body.replace(id, Expr::Id(Reference::Instance(instance)), solved); + 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)?; + body.replace(id, Expr::Id(Reference::Instance(instance)), solved); + 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)?; + body.replace(id, Expr::Id(Reference::Instance(instance)), solved); + 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)?; + body.replace(id, Expr::Id(Reference::Instance(instance)), solved); + substitute_body_sizes(body, &bindings, size_vars); + return Ok(()); } - Ok(type_args) - } - - /// Create a specialized version of a generic function (type vars only). - fn instantiate_function( - &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) + Err(format_error( + body.loc(id), + &format!("No checked function candidate matches expression {}", id), + )) } - /// Create a specialized version of a generic function with both type and size args. - fn instantiate_function_with_sizes( + fn instantiate_global( &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); +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(Reference::SizeParameter(local)) => { + 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(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::Array(generic, _), Type::Array(concrete, _)) => { + infer_size_bindings_pair(*generic, *concrete, 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 + .nodes() + .iter() + .filter_map(|node| match node.kind { + Expr::Id(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( + Expr::Id(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 + .nodes() + .iter() + .any(|node| node.kind == 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( + Expr::Id(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 Expr::Id(Reference::Instance(instance)) = assumption[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 + .nodes() + .iter() + .any(|node| node.kind == Expr::Id(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 + .nodes() + .iter() + .filter(|node| node.kind == Expr::Int(3, None)) + .count(), + 2 + ); + assert!(function.arena.nodes().iter().any(|node| matches!(node.kind, + Expr::Id(Reference::Local(local)) if function.arena.local(local).name == Name::str("N")))); + assert!(!function + .arena + .nodes() + .iter() + .any(|node| matches!(node.kind, Expr::Id(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..f5a5672f 100644 --- a/src/safety_checker.rs +++ b/src/safety_checker.rs @@ -1,6 +1,132 @@ +use crate::checked::{CheckedBody as ExprArena, CheckedExpr as Expr, 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(reference) => reference_place(reference), + _ => 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(reference) => reference_place(reference), + 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,28 @@ 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(local, init, _) => { + 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 +761,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 +778,14 @@ impl SafetyChecker { IndexInterval::default() } - Expr::Var(name, init, _) => { + Expr::Var(local, init, _) => { + 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 +807,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 +825,11 @@ impl SafetyChecker { IndexInterval::default() } - Expr::Id(name) => { + Expr::Id(reference) => { + let Some(place) = reference_place(reference) 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 +854,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 +870,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 +902,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 +923,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 +933,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 +953,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 +1001,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 +1010,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 +1024,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 +1036,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 +1108,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 +1124,25 @@ 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 { params, body } = &context.arena[*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 +1158,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 +1175,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 +1184,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; @@ -1094,32 +1213,23 @@ impl SafetyChecker { 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(*var); + 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,22 @@ 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) - } else { - false - } - }) + context + .arena + .nodes() + .iter() + .map(|node| &node.kind) + .any(|e| { + if let Expr::Var(vn, Some(init), _) = e { + Place::local(*vn) == 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,7 +1321,7 @@ impl SafetyChecker { } Expr::ArrayLiteral(exprs) => { for e in exprs { - self.check_expr(*e, decl, decls); + self.check_expr(*e, context, decls); } IndexInterval::default() } @@ -1215,64 +1330,46 @@ impl SafetyChecker { // 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); + 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: &[CheckedParam], 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 @@ -1286,15 +1383,11 @@ 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)); - - // 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)), - }); + let ty = context.arena.local(param.local).ty; + let is_u32 = ty == mk_type(Type::UInt32); + + // A repeated analysis of this lambda starts with fresh parameter facts. + self.forget(Place::local(param.local)); let arg = call_args.and_then(|(exprs, ivals)| Some((exprs.get(i)?, ivals.get(i)?))); match arg { @@ -1304,16 +1397,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.local), min, max); if ival.non_zero { - self.add_non_zero(param.name); + self.add_non_zero(Place::local(param.local)); } // 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 +1413,19 @@ impl SafetyChecker { .collect(); for array in inherited { self.len_bounds.push(LenBound { - index: param.name, + index: Place::local(param.local), 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.local), Some(0), None), + None => self.add(Place::local(param.local), 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 +1435,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 +1446,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 +1466,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 +1496,144 @@ 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 { + let Expr::Id(reference) = &caller.arena[callee_expr] else { 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(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 +1641,29 @@ 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(index), 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(array) = &callee.arena[*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 +1671,39 @@ 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(lhs), Expr::Id(rhs)) = (&callee.arena[*lhs], &callee.arena[*rhs]) { + if let (Some(argument), Some(size)) = (lookup(lhs), size_lookup(rhs)) { + 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(reference) = &callee.arena[*lhs] { + lookup(reference).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(reference) = &callee.arena[*rhs] { + lookup(reference).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(reference) = &callee.arena[*lhs] { + if let Some(argument) = lookup(reference) { + 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 +1712,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 +1730,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 +1748,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 +1759,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 +1774,202 @@ 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(reference) | Expr::TypeApp(reference, _) => { + if !in_callee { + 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, - ); - } - } - } - - // --- 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 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 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); + 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 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(reference) | Expr::TypeApp(reference, _) => { + let definitions = definitions(reference, &function.arena); + 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 +1988,141 @@ 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 + .nodes() + .iter() + .find_map(|node| { + if let Expr::Call(callee, _) = node.kind { + Some(callee) + } else { + None + } + }) + .unwrap(); + let Expr::Id(Reference::Instance(target)) = caller.arena[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 + .replace(callee, Expr::Id(Reference::Instance(target)), 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..d7e249e9 100644 --- a/src/stack_codegen.rs +++ b/src/stack_codegen.rs @@ -1,15 +1,16 @@ //! 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::{ + CheckedExpr as Expr, CheckedFunction, InstanceId, LocalId, Reference, SpecializedProgram, +}; +use crate::decl::Decl; use crate::defs::*; -use crate::expr::*; use crate::stack_ir::*; use crate::types::*; -use crate::DeclTable; use std::collections::{HashMap, HashSet}; /// Loop context for break/continue support. @@ -30,17 +31,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 +51,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 +67,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 +101,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 +127,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 +160,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 +199,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 +218,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 +249,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 +268,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 +297,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, + decls: &'a SpecializedProgram, - /// Map from variable names to their local storage kind. - variables: HashMap, - - /// 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 +312,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 +320,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 +330,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 +342,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 +361,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 +385,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 +410,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 +418,15 @@ 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(Reference::Local(local)) => { + let ty = self.decl.arena.local(*local).ty; + match &*ty { + Type::Reference(inner) => *inner, + _ => ty, + } + } + Expr::Id(Reference::Instance(instance)) => self.decls.instance(*instance).ty(), 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 +474,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 +489,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 +497,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 +534,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 +610,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 +635,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 +701,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 +767,13 @@ impl<'a> FunctionTranslator<'a> { func.emit(StackOp::LocalAddr(mem_slot)); } - Expr::Id(name) => { - self.translate_id(*name, expr, func); + Expr::Id(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 { .. })) @@ -830,7 +809,7 @@ impl<'a> FunctionTranslator<'a> { Expr::Let(name, init, _) => { let name = *name; 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 +820,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 +835,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 +853,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 +869,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; 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 +894,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 +910,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 +920,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 +932,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 +951,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 +967,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 +979,6 @@ impl<'a> FunctionTranslator<'a> { self.translate_expr(expr_id, func); } } - self.restore_bindings(saved); } } @@ -1158,11 +1130,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 +1142,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; + _ => unreachable!("non-concrete reference in specialized body"), } - - // 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; - } - - // Unknown: push 0. - func.emit(StackOp::I64Const(0)); } /// The store-form vector op for an f32x4-producing expression, or @@ -1303,7 +1219,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 +1243,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 Expr::Id(Reference::Instance(instance)) = &self.decl.arena[*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 +1271,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 +1284,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 +1319,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 +1370,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 +1493,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 +1523,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 Expr::Id(Reference::Local(name)) = &self.decl.arena[lhs_id] { let name = *name; if self.captured_vars.contains(&name) { self.translate_expr(rhs_id, func); @@ -1632,7 +1547,7 @@ impl<'a> FunctionTranslator<'a> { } // Direct scalar local assignment. - if let Expr::Id(name) = &self.decl.arena.exprs[lhs_id] { + if let Expr::Id(Reference::Local(name)) = &self.decl.arena[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 +1567,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 +1692,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 +1711,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 Expr::Id(Reference::Local(name)) = &self.decl.arena[expr] { if let Some(LocalKind::Memory(slot)) = self.variables.get(name) { return Some(*slot); } @@ -1806,7 +1721,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 Expr::Id(Reference::Local(name)) = &self.decl.arena[expr] { if let Some(LocalKind::Scalar(local)) = self.variables.get(name) { return Some(*local); } @@ -1825,7 +1740,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,8 +1773,8 @@ 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) => { + match &self.decl.arena[expr].clone() { + Expr::Id(Reference::Local(name)) => { let name = *name; if let Some(&kind) = self.variables.get(&name) { match kind { @@ -1878,12 +1793,13 @@ impl<'a> FunctionTranslator<'a> { func.emit(StackOp::LocalAddr(slot)); } } - } else if let Some(&offset) = self.globals.get(&name) { - func.emit(StackOp::GlobalAddr(offset)); } else { - func.emit(StackOp::I64Const(0)); + unreachable!("checked local must have storage"); } } + Expr::Id(Reference::Instance(instance)) => { + func.emit(StackOp::GlobalAddr(self.globals[instance])); + } Expr::Field(lhs_id, name) => { let lhs_id = *lhs_id; @@ -2004,8 +1920,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 Expr::Id(Reference::Instance(instance)) = &self.decl.arena[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 +2124,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 +2137,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 +2212,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 +2253,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 +2262,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 +2336,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[fn_id] { + Expr::Id(Reference::Local(_)) => true, + Expr::Id(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 +2509,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 +2527,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 +2545,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 +2574,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 +2890,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 { - panic!( - "stack codegen lambda: expected function type, got {:?}", - lambda_ty - ); + 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 { + 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 +2965,7 @@ impl<'a> FunctionTranslator<'a> { } } } else { - func.emit(StackOp::I64Const(0)); + unreachable!("checked capture local {:?} has no storage", name); } } @@ -3290,147 +3132,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..b6bc68b2 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, CheckedExpr as Expr, CheckedFunction, InstanceId, LocalId, Reference, + SpecializedProgram, +}; +use crate::decl::Decl; use crate::defs::*; -use crate::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,15 @@ 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(Reference::Local(local)) => { + let ty = self.body.decl.arena.local(*local).ty; + match &*ty { + Type::Reference(inner) => *inner, + _ => ty, + } + } + Expr::Id(Reference::Instance(instance)) => self.decls.instance(*instance).ty(), 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 +800,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,13 +839,13 @@ impl<'a> FunctionTranslator<'a> { dst } - Expr::Id(name) => { + Expr::Id(Reference::Local(name)) => { let ty = self.expr_type(expr); // Check if it's a captured closure variable (double indirection). - if self.captured_vars.contains(name) { + if self.body.captured_vars.contains(name) { // Load pointer-to-captured-storage from our local slot. - 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, @@ -947,25 +869,25 @@ impl<'a> FunctionTranslator<'a> { } // Check if it's a local variable. - if let Some(®) = self.variables.get(name) { - if self.reference_vars.contains(name) { + 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.reg_promoted.contains(name) { + } 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.local_slots.get(name) { + if let Some(&slot) = self.body.local_slots.get(name) { func.emit(Opcode::LocalAddr { dst: reg, slot }); } reg - } else if let Some(&slot) = self.local_slots.get(name) { + } 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 }); @@ -975,7 +897,13 @@ impl<'a> FunctionTranslator<'a> { } else { reg } - } else if let Some(&offset) = self.globals.get(name) { + } else { + unreachable!("checked local must have storage") + } + } + Expr::Id(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 }); @@ -1026,14 +954,13 @@ impl<'a> FunctionTranslator<'a> { }); fat_addr } else { - // Unknown identifier - this shouldn't happen after type checking. - let dst = self.alloc_reg(); - func.emit(Opcode::LoadImm { dst, value: 0 }); - dst + unreachable!("instance must name storage or a function") } } } + Expr::Id(_) => unreachable!("non-concrete reference in specialized body"), + Expr::Binop(op, lhs_id, rhs_id) => self.translate_binop(*op, *lhs_id, *rhs_id, func), Expr::Unop(op, arg_id) => self.translate_unop(*op, *arg_id, func), @@ -1041,11 +968,11 @@ 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); + 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 +983,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 +996,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 +1015,27 @@ 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); + 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 +1052,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 +1144,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 } } @@ -1398,125 +1320,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 +1394,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 +1424,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 +1440,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 +1457,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 +1471,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 +1484,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 +1968,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 Expr::Id(Reference::Local(name)) = &self.body.decl.arena[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 +1989,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 Expr::Id(Reference::Local(name)) = &self.body.decl.arena[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 +2000,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 +2025,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 +2045,34 @@ 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)); + match &self.body.decl.arena[expr] { + Expr::Id(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.variables.insert(*name, addr); - self.local_slots.insert(*name, slot); - self.reg_promoted.remove(name); + self.body.variables.insert(*name, addr); + self.body.local_slots.insert(*name, slot); + self.body.reg_promoted.remove(name); addr } else { reg } - } 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 { - let dst = self.alloc_reg(); - func.emit(Opcode::LoadImm { dst, value: 0 }); - dst + unreachable!("checked local must have storage") } } + Expr::Id(Reference::Instance(instance)) => { + let dst = self.alloc_reg(); + func.emit(Opcode::GlobalAddr { + dst, + offset: self.globals[instance], + }); + dst + } Expr::Field(lhs_id, name) => { let lhs_addr = self.translate_lvalue(*lhs_id, func); @@ -2381,7 +2241,9 @@ impl<'a> FunctionTranslator<'a> { } // Special handling for built-in functions. - if let Expr::Id(name) = &self.decl.arena.exprs[fn_id] { + if let Expr::Id(Reference::Instance(instance)) = &self.body.decl.arena[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 +2426,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 +2450,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 +2475,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 +2520,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 +2541,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 +2601,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 +2621,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 +2654,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[fn_id] { + Expr::Id(Reference::Local(_)) => true, + Expr::Id(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 +2867,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 +2889,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 +2899,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 +2939,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 +3227,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 +3248,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 +3265,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 +3590,269 @@ 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(Expr::Var(counter, Some(zero), None), void, loc); + let condition = if iterations == 0 { + arena.add(Expr::False, boolean, loc) + } else { + let read = arena.add(Expr::Id(Reference::Local(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(Expr::Id(Reference::Local(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(Expr::Id(Reference::Local(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(Expr::Id(Reference::Local(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(Expr::Let(outer, forty, None), mk_type(Type::Void), loc); + let target = main_body.add( + Expr::Id(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( + Expr::Id(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(Expr::Id(Reference::Local(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); + 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(); + 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 = "" From 96a4d080854cded269736396c569de2b3db19006 Mon Sep 17 00:00:00 2001 From: Izaak Branderhorst Date: Wed, 9 Sep 2026 03:47:33 +0200 Subject: [PATCH 2/3] Guard optional Stack interpreter checks in regression tests The Ubuntu LLVM job builds without the Clang-only C Stack interpreter. Three unguarded references to stack_interp_bridge prevented its library tests from compiling. Gate the C Stack portions of the hoisting and VM codegen regressions with has_stack_interp while keeping register VM coverage active on every build. Validation: reproduced all three original errors with has_stack_interp disabled locally; all 439 LLVM library tests now pass in that configuration. All 13 affected hoisting and VM codegen tests also pass with the C Stack interpreter enabled. git diff --check passes. --- src/hoist.rs | 1 + src/vm_codegen.rs | 34 ++++++++++++++++++++-------------- 2 files changed, 21 insertions(+), 14 deletions(-) diff --git a/src/hoist.rs b/src/hoist.rs index dec16cd0..21b82b2d 100644 --- a/src/hoist.rs +++ b/src/hoist.rs @@ -574,6 +574,7 @@ mod tests { 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/vm_codegen.rs b/src/vm_codegen.rs index b6bc68b2..7749cfa4 100644 --- a/src/vm_codegen.rs +++ b/src/vm_codegen.rs @@ -3781,14 +3781,17 @@ mod tests { ); let vm = VMCodegen::new().compile(&checked).unwrap(); assert_eq!(VM::new().run(&vm), 42); - 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); + #[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); + } } fn check_capture_program(source: &str) { @@ -3800,12 +3803,15 @@ mod tests { "capture regression must check successfully" ); compiler.specialize().unwrap(); - let stack = compiler.compile_stack().unwrap(); - assert_eq!( - crate::stack_interp_bridge::run(&stack), - 42, - "Stack capture result" - ); + #[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"); } From 1c7294ece8ecea868df90deed9e8798dcf2377f7 Mon Sep 17 00:00:00 2001 From: Taylor Holliday Date: Tue, 8 Sep 2026 19:40:45 -0700 Subject: [PATCH 3/3] Keep Expr non-generic; record checked facts in side tables The checked-program boundary made `Expr` generic over its reference, binder and parameter payloads so checked bodies could carry `Reference` and `LocalId` in place of names. The repository already records derived facts in tables indexed by `ExprID`, and the checker itself built those tables before transcribing them into the generic form. `Expr` is a single non-generic type again. `CheckedBody` owns the source expression tree plus parallel tables for the solved type, the recorded `Reference` of each identifier/type application, and the `LocalId`s each declaration or lambda introduces. Consumers read identity through `reference(id)`, `binder(id)` and `binders(id)`; the checked pretty printer delegates to the source printer instead of duplicating it. Specialization, safety checking, hoisting, copy elision, Cranelift, LLVM, register VM and stack lowering read the tables. Validation checks that references and binders are recorded exactly where the expression kind requires them. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01BcyQEVKmy98VkZhAnx13Tk --- docs/CHECKED_PROGRAM.md | 8 +- docs/MONOMORPHIZATION.md | 10 +- src/checked.rs | 467 ++++++++++++++----------------- src/checked/validate.rs | 169 ++++++----- src/checker.rs | 122 ++++---- src/compiler.rs | 9 +- src/compiler/assumption_tests.rs | 41 +-- src/compiler/safety_tests.rs | 10 +- src/copy_elision.rs | 18 +- src/expr.rs | 16 +- src/free_locals.rs | 24 +- src/hoist.rs | 71 ++--- src/jit.rs | 180 ++++++------ src/llvm_jit.rs | 269 +++++++++--------- src/monomorph_pass.rs | 77 +++-- src/safety_checker.rs | 144 +++++----- src/stack_codegen.rs | 105 +++---- src/vm_codegen.rs | 346 ++++++++++++----------- 18 files changed, 1046 insertions(+), 1040 deletions(-) diff --git a/docs/CHECKED_PROGRAM.md b/docs/CHECKED_PROGRAM.md index a1e13a94..9c53f8ea 100644 --- a/docs/CHECKED_PROGRAM.md +++ b/docs/CHECKED_PROGRAM.md @@ -24,7 +24,10 @@ 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. -Each `CheckedNode` owns its operation, result type and source location. Each +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 @@ -268,7 +271,8 @@ binders, remaps their uses and preserves enclosing references and source locatio 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; `replace_node` can also change provenance. +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. diff --git a/docs/MONOMORPHIZATION.md b/docs/MONOMORPHIZATION.md index 2a89851b..04d5c308 100644 --- a/docs/MONOMORPHIZATION.md +++ b/docs/MONOMORPHIZATION.md @@ -27,11 +27,11 @@ and editing, but backends consume checked bodies. ## Bodies and identities -`CheckedBody` owns operation/type/location nodes, 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 -expressions share syntax shape and child traversal with distinct reference, -binder and parameter payloads. +`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`, diff --git a/src/checked.rs b/src/checked.rs index dd63cfbe..8ceea171 100644 --- a/src/checked.rs +++ b/src/checked.rs @@ -1,9 +1,10 @@ //! The program after lexical and type checking. //! -//! Nodes own their operation, result type and source provenance. 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. +//! 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::*; @@ -65,21 +66,22 @@ pub struct Local { pub mutable: bool, } -pub type CheckedExpr = Expr; pub type CheckedDecl = Decl; pub type CheckedDeclTable = DeclTable; pub type CheckedDeclarations = DeclarationList; -#[derive(Clone, Debug, Eq, PartialEq, Hash)] -pub struct CheckedNode { - pub kind: CheckedExpr, - pub ty: TypeID, - pub loc: Loc, -} - +/// 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 { - nodes: Vec, + 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, } @@ -89,33 +91,71 @@ impl CheckedBody { Self::default() } pub fn from_parts( - nodes: Vec, + 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 { - nodes, + syntax, + types, + references, + binders, locals, requirements, } } pub fn len(&self) -> usize { - self.nodes.len() + self.syntax.exprs.len() } pub fn is_empty(&self) -> bool { - self.nodes.is_empty() + self.syntax.exprs.is_empty() + } + /// Every expression handle in this body. + pub fn ids(&self) -> std::ops::Range { + 0..self.len() } - pub fn node(&self, id: ExprID) -> &CheckedNode { - &self.nodes[id] + /// The expression tree with its source locations. + pub fn syntax(&self) -> &ExprArena { + &self.syntax } - pub fn nodes(&self) -> &[CheckedNode] { - &self.nodes + pub fn exprs(&self) -> &[Expr] { + &self.syntax.exprs } pub fn ty(&self, id: ExprID) -> TypeID { - self.nodes[id].ty + self.types[id] } pub fn loc(&self, id: ExprID) -> Loc { - self.nodes[id].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()] @@ -125,24 +165,109 @@ impl CheckedBody { self.locals.push(Local { name, ty, mutable }); id } - pub fn add(&mut self, kind: CheckedExpr, ty: TypeID, loc: Loc) -> ExprID { - let id = self.nodes.len(); - self.nodes.push(CheckedNode { kind, ty, loc }); + /// 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 } - /// Replacing a node retains its handle and provenance, not any analysis of - /// the previous node. Callers provide the replacement's checked result type. - pub fn replace(&mut self, id: ExprID, kind: CheckedExpr, ty: TypeID) { - self.nodes[id].kind = kind; - self.nodes[id].ty = ty; + /// 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 } - pub fn replace_node(&mut self, id: ExprID, node: CheckedNode) { - self.nodes[id] = node; + /// 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 node in &mut self.nodes { - node.ty = node.ty.subst(instance); - match &mut node.kind { + 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 { @@ -165,166 +290,9 @@ impl CheckedBody { } } - /// Source-oriented diagnostics over checked nodes; spellings are presentation - /// data and are never converted back into unresolved syntax. + /// Source-oriented diagnostics over checked expressions. pub fn pretty_print(&self, id: ExprID, indent: usize) -> String { - self.pretty_print_with(id, indent, &|reference| match reference { - Reference::Local(local) | Reference::SizeParameter(local) => self.local(*local).name, - Reference::Global(id) => Name::new(format!("global#{}", id.0)), - Reference::Functions(ids) => Name::new(format!("function#{:?}", ids)), - Reference::InterfaceMember { member, .. } => Name::new(format!("member#{}", member.0)), - Reference::Instance(id) => Name::new(format!("instance#{}", id.0)), - }) - } - pub fn pretty_print_with( - &self, - id: ExprID, - indent: usize, - reference_name: &impl Fn(&Reference) -> Name, - ) -> String { - let child = |id| self.pretty_print_with(id, indent, reference_name); - let list = |ids: &[ExprID]| { - ids.iter() - .map(|id| child(*id)) - .collect::>() - .join(", ") - }; - match &self[id] { - Expr::Id(reference) => reference_name(reference).to_string(), - Expr::TypeApp(reference, args) => format!( - "{}⟨{}⟩", - reference_name(reference), - args.iter() - .map(|ty| ty.pretty_print()) - .collect::>() - .join(", ") - ), - Expr::Int(value, suffix) => format!( - "{}{}", - value, - suffix.map(|suffix| suffix.to_string()).unwrap_or_default() - ), - Expr::Real(value, suffix) => format!( - "{}{}", - value, - suffix.map(|suffix| suffix.to_string()).unwrap_or_default() - ), - Expr::String(value) => format!("\"{}\"", value), - Expr::Char(value) => format!("'{}'", value), - Expr::True => "true".into(), - Expr::False => "false".into(), - Expr::Enum(name) => format!(".{}", name), - Expr::Error => "".into(), - Expr::Call(function, args) => format!("{}({})", child(*function), list(args)), - Expr::Macro(name, args) => format!("@{}({})", name, list(args)), - Expr::Binop(op, lhs, rhs) => { - format!("{} {} {}", child(*lhs), format_binop(*op), child(*rhs)) - } - Expr::Unop(op, value) => format!("{}{}", format_unop(*op), child(*value)), - Expr::Lambda { params, body } => format!( - "|{}| {}", - params - .iter() - .map(|param| format!( - "{}: {}", - self.local(param.local).name, - self.local(param.local).ty.pretty_print() - )) - .collect::>() - .join(", "), - child(*body) - ), - Expr::Field(base, name) => format!("{}.{}", child(*base), name), - Expr::Array(element, size) => format!("[{}; {}]", child(*element), child(*size)), - Expr::ArrayLiteral(elements) => format!("[{}]", list(elements)), - Expr::ArrayIndex(array, index) => format!("{}[{}]", child(*array), child(*index)), - Expr::AsTy(value, ty) => format!("{}:{}", child(*value), ty.pretty_print()), - Expr::Let(local, init, annotation) => { - let annotation = annotation - .map(|ty| format!(": {}", ty.pretty_print())) - .unwrap_or_default(); - format!( - "let {}{} = {}", - self.local(*local).name, - annotation, - child(*init) - ) - } - Expr::Var(local, init, annotation) => { - let annotation = annotation - .map(|ty| format!(": {}", ty.pretty_print())) - .unwrap_or_default(); - let init = init - .map(|init| format!(" = {}", child(init))) - .unwrap_or_default(); - format!("var {}{}{}", self.local(*local).name, annotation, init) - } - Expr::If(cond, yes, no) => { - let no = no - .map(|no| { - format!( - " else {}", - self.pretty_print_with(no, indent + 1, reference_name) - ) - }) - .unwrap_or_default(); - format!( - "if {} {}{}", - child(*cond), - self.pretty_print_with(*yes, indent + 1, reference_name), - no - ) - } - Expr::While(cond, body) => format!( - "while {} {}", - child(*cond), - self.pretty_print_with(*body, indent + 1, reference_name) - ), - Expr::For { - var, - start, - end, - body, - } => format!( - "for {} in {} .. {} {}", - self.local(*var).name, - child(*start), - child(*end), - self.pretty_print_with(*body, indent + 1, reference_name) - ), - Expr::Block(exprs) => { - if exprs.is_empty() { - return "{}".into(); - } - let expressions = exprs - .iter() - .map(|expr| { - format!( - "{}{}", - " ".repeat(indent + 1), - self.pretty_print_with(*expr, indent + 1, reference_name) - ) - }) - .collect::>() - .join("\n"); - format!("{{\n{}\n{}}}", expressions, " ".repeat(indent)) - } - Expr::Return(value) => format!("return {}", child(*value)), - Expr::Assume(value) => format!("assume {}", child(*value)), - Expr::Break => "break".into(), - Expr::Continue => "continue".into(), - Expr::Tuple(values) => format!("({})", list(values)), - Expr::StructLit(name, fields) => format!( - "{}({})", - name, - fields - .iter() - .map(|(name, value)| format!("{}: {}", name, child(*value))) - .collect::>() - .join(", ") - ), - Expr::Arena(value) => format!("arena {}", child(*value)), - } + self.syntax.exprs[id].pretty_print(&self.syntax, indent) } /// Duplicate evaluations. Bindings declared in the copied subtree receive @@ -334,15 +302,7 @@ impl CheckedBody { /// 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) { - match &body[id] { - Expr::Let(local, ..) | Expr::Var(local, ..) | Expr::For { var: local, .. } => { - found.insert(*local); - } - Expr::Lambda { params, .. } => { - found.extend(params.iter().map(|param| param.local)); - } - _ => {} - } + found.extend(body.binders(id).iter().copied()); for child in body[id].subexprs() { declarations(body, child, found); } @@ -357,50 +317,44 @@ impl CheckedBody { locals.insert(old, self.add_local(local.name, local.ty, local.mutable)); } fn copy(body: &mut CheckedBody, id: ExprID, locals: &HashMap) -> ExprID { - let mut node = body.node(id).clone(); - let remap = |local: &mut LocalId| { - if let Some(new) = locals.get(local) { - *local = *new; - } - }; - match &mut node.kind { - Expr::Id(Reference::Local(local)) - | Expr::TypeApp(Reference::Local(local), _) - | Expr::Let(local, ..) - | Expr::Var(local, ..) - | Expr::For { var: local, .. } => remap(local), - Expr::Lambda { params, .. } => { - for param in params { - remap(&mut param.local); - } - } - _ => {} - } - node.kind.map_children(|child| copy(body, child, locals)); - body.add(node.kind, node.ty, node.loc) + 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, params: &[CheckedParam]) -> Vec { - crate::free_locals::free_locals(self, root, params.iter().map(|param| param.local)).locals + 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 node in &self.nodes { - if let Expr::Lambda { params, body } = &node.kind { - captured.extend(self.captures(*body, params)); + 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 = CheckedExpr; + type Output = Expr; fn index(&self, id: ExprID) -> &Self::Output { - &self.nodes[id].kind + &self.syntax.exprs[id] } } @@ -435,23 +389,24 @@ impl CheckedFunction { self.arena.captured_locals() } pub fn extract_lambda(&self, expression: ExprID, name: Name) -> Self { - let Expr::Lambda { params, body } = &self.arena[expression] else { + 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: params.clone(), + 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, params), + closure_vars: self.arena.captures(*body, locals.iter().copied()), is_extern: false, } } @@ -602,28 +557,25 @@ mod tests { 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(Expr::Id(Reference::Local(outer)), ty, test_loc()); - let binding = body.add( - Expr::Let(inner, outer_read, None), - mk_type(Type::Void), - test_loc(), - ); - let inner_read = body.add(Expr::Id(Reference::Local(inner)), ty, test_loc()); + 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(fresh, initializer, _) = body[children[0]] else { + let Expr::Let(_, initializer, _) = body[children[0]] else { panic!() }; + let fresh = body.binder(children[0]); assert_ne!(fresh, inner); - assert_eq!(body[initializer], Expr::Id(Reference::Local(outer))); - assert_eq!(body[children[1]], Expr::Id(Reference::Local(fresh))); + 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[inner_read], Expr::Id(Reference::Local(inner))); + assert_eq!(body.reference(inner_read), Some(&Reference::Local(inner))); } #[test] @@ -633,33 +585,28 @@ mod tests { 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(Expr::Id(Reference::Local(outer)), callable, test_loc()); - let read_inner = body.add(Expr::Id(Reference::Local(inner)), ty, test_loc()); - let applied_outer = body.add( - Expr::TypeApp(Reference::Local(outer), vec![]), - callable, - test_loc(), - ); + 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( + 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, &[CheckedParam { local: inner }]), - vec![outer] - ); + assert_eq!(body.captures(nested_root, []), vec![inner, outer]); + assert_eq!(body.captures(nested, [inner]), vec![outer]); } #[test] @@ -667,7 +614,7 @@ mod tests { let mut body = CheckedBody::new(); let generic = typevar("T"); let local = body.add_local(Name::str("x"), generic, false); - let read = body.add(Expr::Id(Reference::Local(local)), generic, test_loc()); + let read = body.add_local_read(local, generic, test_loc()); body.requirements.push(InterfaceRequirement { id: RequirementId(0), interface: DefId(1), diff --git a/src/checked/validate.rs b/src/checked/validate.rs index 4c997f91..7f09a893 100644 --- a/src/checked/validate.rs +++ b/src/checked/validate.rs @@ -202,8 +202,8 @@ impl Phase<'_> { for local in &body.locals { validate_type(local.ty, self.concrete())?; } - for (id, node) in body.nodes().iter().enumerate() { - self.node(node, body, &bindings) + for id in body.ids() { + self.node(id, body, &bindings) .map_err(|error| format!("expression {}: {}", id, error))?; } Ok(()) @@ -211,28 +211,54 @@ impl Phase<'_> { fn node( &self, - node: &CheckedNode, + id: ExprID, body: &CheckedBody, bindings: &[Option], ) -> Result<(), String> { - validate_type(node.ty, self.concrete())?; - match &node.kind { - Expr::Id(reference) => self.reference(reference, body, bindings)?, - Expr::TypeApp(reference, args) => { + 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)?; - for &ty in args { - validate_type(ty, self.concrete())?; + 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 node.ty != mk_type(Type::Void) { + 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 { @@ -245,15 +271,18 @@ impl Phase<'_> { return Err("call arity does not match checked signature".into()); } if let Self::Concrete(program) = self { - if let Expr::Id(reference @ Reference::Instance(instance)) - | Expr::TypeApp(reference @ Reference::Instance(instance), _) = - &body[*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()); + 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() + ); + } } } } @@ -395,17 +424,9 @@ fn collect_bindings( for size in sizes { bind(size.local, BinderKind::Size)?; } - for node in body.nodes() { - match &node.kind { - Expr::Let(local, ..) | Expr::Var(local, ..) | Expr::For { var: local, .. } => { - bind(*local, BinderKind::Value)? - } - Expr::Lambda { params, .. } => { - for param in params { - bind(param.local, BinderKind::Value)?; - } - } - _ => {} + for id in body.ids() { + for &local in body.binders(id) { + bind(local, BinderKind::Value)?; } } Ok(bindings) @@ -438,8 +459,8 @@ fn validate_edges(body: &CheckedBody, roots: &[ExprID]) -> Result<(), String> { for &root in roots { uses[root] += 1; } - for node in body.nodes() { - for child in node.kind.subexprs() { + for id in body.ids() { + for child in body[id].subexprs() { *uses .get_mut(child) .ok_or("expression edge is outside its body")? += 1; @@ -457,14 +478,11 @@ fn validate_edges(body: &CheckedBody, roots: &[ExprID]) -> Result<(), String> { .get_mut(id) .ok_or("expression edge is outside its body")?; if leaving { - contains_binder[id] = match &body[id] { - Expr::Let(..) | Expr::Var(..) | Expr::For { .. } => true, - Expr::Lambda { params, .. } if !params.is_empty() => true, - _ => body[id] + contains_binder[id] = !body.binders(id).is_empty() + || body[id] .subexprs() .iter() - .any(|&child| contains_binder[child]), - }; + .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. @@ -494,7 +512,7 @@ mod tests { 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(Expr::Id(Reference::Local(local)), ty, test_loc()); + let body = arena.add_local_read(local, ty, test_loc()); CheckedFunction { name: Name::str("main"), typevars: vec![], @@ -557,28 +575,33 @@ mod tests { (|f| f.params.push(f.params[0].clone()), "multiple binders"), ( |f| { - f.arena - .add(Expr::Id(Reference::Local(LocalId(999))), f.ret, test_loc()); + 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(Expr::Id(Reference::Local(local)), f.ret, test_loc()); + f.arena.add_local_read(local, f.ret, test_loc()); }, "no value binder", ), ( |f| { - f.arena.add( + f.arena.add_binding( Expr::Lambda { - params: vec![CheckedParam { - local: LocalId(999), + params: vec![Param { + name: Name::str("p"), + ty: None, }], body: 0, }, + vec![LocalId(999)], f.ty(), test_loc(), ); @@ -626,9 +649,8 @@ mod tests { let function = templates.function(definition).unwrap(); let id = function .arena - .nodes() - .iter() - .position(|node| matches!(node.kind, Expr::Let(..) | Expr::Var(..))) + .ids() + .find(|&id| matches!(function.arena[id], Expr::Let(..) | Expr::Var(..))) .unwrap(); concrete(function.clone()).unwrap(); for (annotation, result, expected) in [ @@ -693,14 +715,16 @@ mod tests { caller.ret, test_loc(), )); - let reference = Reference::Instance(target_id); let kind = if type_application { - Expr::TypeApp(reference, vec![]) + Expr::TypeApp(Name::str("target"), vec![]) } else { - Expr::Id(reference) + 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)], @@ -724,8 +748,8 @@ mod tests { let main = templates .function(templates.decls.named_ids(Name::str("main"))[0]) .unwrap(); - assert!(main.arena.nodes().iter().any(|node| { - matches!(&node.kind, Expr::Id(Reference::Functions(candidates)) if candidates.len() == 2) + 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")) @@ -746,11 +770,7 @@ mod tests { let local = function .arena .add_local(Name::str("y"), function.ret, false); - let binding = function.arena.add( - Expr::Let(local, read, None), - mk_type(Type::Void), - test_loc(), - ); + 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()); @@ -800,14 +820,15 @@ mod tests { let mut function = sample_function(); function .arena - .add(Expr::Id(reference), function.ty(), test_loc()); + .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( - Expr::Id(Reference::Functions(vec![DefId(0)])), + function.arena.add_id( + Name::str("x"), + Reference::Functions(vec![DefId(0)]), function.ty(), test_loc(), ); @@ -915,8 +936,9 @@ mod tests { ) .is_err()); let mut function = sample_function(); - function.arena.add( - Expr::Id(Reference::Instance(InstanceId(1))), + function.arena.add_id( + Name::str("x"), + Reference::Instance(InstanceId(1)), function.ty(), test_loc(), ); @@ -954,12 +976,10 @@ mod tests { *update.arena.local(update.params[0].local).ty, Type::Reference(_) )); - assert!(update - .arena - .nodes() - .iter() - .any(|node| matches!(node.kind, Expr::Id(Reference::Local(_))) - && node.ty == mk_type(Type::Int32))); + assert!(update.arena.ids().any(|id| { + matches!(update.arena.reference(id), Some(Reference::Local(_))) + && update.arena.ty(id) == mk_type(Type::Int32) + })); } #[test] @@ -987,15 +1007,14 @@ mod tests { .unwrap(); let original = function .arena - .nodes() - .iter() - .position(|node| matches!(node.kind, Expr::For { .. })) + .ids() + .find(|&id| matches!(function.arena[id], Expr::For { .. })) .unwrap(); - let captures = function.arena.captures(original, &[]); + 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); + 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)); diff --git a/src/checker.rs b/src/checker.rs index c44009f7..8682fd74 100644 --- a/src/checker.rs +++ b/src/checker.rs @@ -1,6 +1,5 @@ use crate::checked::{ - CheckedBody, CheckedExpr, CheckedFunction, CheckedNode, CheckedParam, Local, LocalId, - Reference, RequirementId, + CheckedBody, CheckedFunction, CheckedParam, Local, LocalId, Reference, RequirementId, }; use crate::free_locals::{free_locals, BindingFacts, BindingNode}; use crate::*; @@ -1857,96 +1856,73 @@ impl Checker { }) .collect(); let solved = self.solved_types(); - let mut nodes = Vec::with_capacity(source.exprs.len()); - for (id, expression) in source.exprs.iter().enumerate() { - let resolution = || self.references[id].clone().expect("checked reference"); - let binder = || *self.binders[id].first().expect("checked binder"); - let kind = match expression { - Expr::Id(_) => CheckedExpr::Id(resolution()), - Expr::TypeApp(_, args) => CheckedExpr::TypeApp( - resolution(), - args.iter().map(|t| t.subst(&self.inst)).collect(), - ), - Expr::Let(_, init, _) => CheckedExpr::Let(binder(), *init, None), - Expr::Var(_, init, _) => CheckedExpr::Var(binder(), *init, None), - Expr::For { - start, end, body, .. - } => CheckedExpr::For { - var: binder(), - start: *start, - end: *end, - body: *body, - }, - Expr::Lambda { body, .. } => CheckedExpr::Lambda { - params: self.binders[id] - .iter() - .map(|&local| CheckedParam { local }) - .collect(), - body: *body, - }, - Expr::Int(v, s) => CheckedExpr::Int(*v, *s), - Expr::Real(v, s) => CheckedExpr::Real(v.clone(), *s), - Expr::Call(f, a) => CheckedExpr::Call(*f, a.clone()), - Expr::Binop(op, a, b) => CheckedExpr::Binop(*op, *a, *b), - Expr::Unop(op, a) => CheckedExpr::Unop(*op, *a), - Expr::String(v) => CheckedExpr::String(v.clone()), - Expr::Char(v) => CheckedExpr::Char(*v), - Expr::Field(a, n) => CheckedExpr::Field(*a, *n), - Expr::Array(a, b) => CheckedExpr::Array(*a, *b), - Expr::ArrayLiteral(a) => CheckedExpr::ArrayLiteral(a.clone()), - Expr::ArrayIndex(a, b) => CheckedExpr::ArrayIndex(*a, *b), - Expr::True => CheckedExpr::True, - Expr::False => CheckedExpr::False, - Expr::AsTy(a, t) => CheckedExpr::AsTy(*a, t.subst(&self.inst)), - Expr::If(a, b, c) => CheckedExpr::If(*a, *b, *c), - Expr::While(a, b) => CheckedExpr::While(*a, *b), - Expr::Block(v) => CheckedExpr::Block(v.clone()), - Expr::Return(a) => CheckedExpr::Return(*a), - Expr::Break => CheckedExpr::Break, - Expr::Continue => CheckedExpr::Continue, - Expr::Enum(n) => CheckedExpr::Enum(*n), - Expr::Tuple(v) => CheckedExpr::Tuple(v.clone()), - Expr::StructLit(n, v) => CheckedExpr::StructLit(*n, v.clone()), - Expr::Arena(a) => CheckedExpr::Arena(*a), - Expr::Assume(a) => CheckedExpr::Assume(*a), + 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") } - }; - let ty = if matches!(expression, Expr::Let(..) | Expr::Var(..)) { + _ => {} + } + references.push(reference); + types.push(if matches!(expression, Expr::Let(..) | Expr::Var(..)) { mk_type(Type::Void) } else { solved[id] - }; - nodes.push(CheckedNode { - kind, - ty, - loc: source.locs[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 0..nodes.len() { - if let CheckedExpr::Binop(op, lhs, rhs) = nodes[id].kind.clone() { + for id in body.ids() { + if let Expr::Binop(op, lhs, rhs) = body[id].clone() { if op.arithmetic() - && (matches!(*nodes[lhs].ty, Type::Name(_, _)) + && (matches!(*body.ty(lhs), Type::Name(_, _)) || (op == Binop::Mod - && matches!(*nodes[lhs].ty, Type::Float32 | Type::Float64))) + && matches!(*body.ty(lhs), Type::Float32 | Type::Float64))) { let Some(reference) = self.references[id].clone() else { continue; }; - let callee = nodes.len(); - nodes.push(CheckedNode { - kind: CheckedExpr::Id(reference), - ty: func(tuple(vec![nodes[lhs].ty, nodes[rhs].ty]), nodes[id].ty), - loc: nodes[id].loc, - }); - nodes[id].kind = CheckedExpr::Call(callee, vec![lhs, rhs]); + 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); } } } - CheckedBody::from_parts(nodes, locals, self.requirements.clone()) + body } /// Publish function metadata only at the function boundary. diff --git a/src/compiler.rs b/src/compiler.rs index 776e98bf..5f35c5c8 100644 --- a/src/compiler.rs +++ b/src/compiler.rs @@ -1944,11 +1944,10 @@ mod tests { let program = compiler.specialized_program().unwrap(); let targets: Vec<_> = main .arena - .nodes() - .iter() - .filter_map(|node| { - if let CheckedExpr::Id(Reference::Instance(target)) = node.kind { - Some(program.instance_name(target).to_string()) + .ids() + .filter_map(|id| { + if let Some(Reference::Instance(target)) = main.arena.reference(id) { + Some(program.instance_name(*target).to_string()) } else { None } diff --git a/src/compiler/assumption_tests.rs b/src/compiler/assumption_tests.rs index 72467ba9..ccde9487 100644 --- a/src/compiler/assumption_tests.rs +++ b/src/compiler/assumption_tests.rs @@ -38,14 +38,15 @@ fn normalization_remaps_assumption_roots_and_shared_binding_occurrences() { }; assert_ne!(left, right); for (root, expected) in [(left, LocalId(0)), (right, LocalId(1))] { - let CheckedExpr::Block(statements) = &body[root] else { + let Expr::Block(statements) = &body[root] else { panic!() }; - assert!(matches!(body[statements[0]], Expr::Let(local, ..) if local == expected)); + 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[value], Expr::Id(Reference::Local(expected))); + assert_eq!(body.reference(value), Some(&Reference::Local(expected))); } CheckedProgram::try_new(DeclTable::new(vec![Decl::Assume { arena: body, @@ -148,8 +149,8 @@ fn assumption_specialization_preserves_local_ownership_and_concrete_targets() { assert_eq!(body.ty(root), mk_type(Type::Bool)); assert_eq!(body.locals, source_body.locals); let mut targets = vec![]; - for node in body.nodes() { - if let Expr::Id(Reference::Instance(target)) = node.kind { + for id in body.ids() { + if let Some(Reference::Instance(target)) = body.reference(id) { targets.push(output.instances[target.index()].definition); } } @@ -171,15 +172,17 @@ fn assumption_specialization_preserves_local_ownership_and_concrete_targets() { .into_iter() .find(|id| templates.function(*id).unwrap().param_types() == vec![mk_type(Type::Int32)]) .unwrap(); - assert!(checked.arena.nodes().iter().any(|node| match node.kind { - Expr::Id(Reference::Instance(target)) => - output.instances[target.index()].definition == positive, - _ => false, - })); + assert!(checked + .arena + .ids() + .any(|id| match checked.arena.reference(id) { + Some(Reference::Instance(target)) => + output.instances[target.index()].definition == positive, + _ => false, + })); assert!(body - .nodes() - .iter() - .any(|node| matches!(node.kind, Expr::Id(Reference::Local(LocalId(0)))))); + .ids() + .any(|id| body.reference(id) == Some(&Reference::Local(LocalId(0))))); let retained = compiler .checked_program() .unwrap() @@ -252,12 +255,12 @@ fn safety_errors_in_assumptions_keep_the_expression_location() { }) .unwrap(); let division = body - .nodes() - .iter() - .find(|node| matches!(node.kind, Expr::Binop(Binop::Div, ..))) + .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.loc); - assert_eq!(division.loc.file, Name::str("")); - assert_eq!(division.loc, original_location); + 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/safety_tests.rs b/src/compiler/safety_tests.rs index 8a66aeee..d3ac0211 100644 --- a/src/compiler/safety_tests.rs +++ b/src/compiler/safety_tests.rs @@ -137,11 +137,11 @@ fn local_and_global_function_values_remain_indirect() { .unwrap(); let mut local_calls = 0; let mut global_calls = 0; - for node in main.arena.nodes() { - if let CheckedExpr::Call(callee, _) = &node.kind { - match &main.arena[*callee] { - CheckedExpr::Id(Reference::Local(_)) => local_calls += 1, - CheckedExpr::Id(Reference::Instance(target)) => { + 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 { .. } diff --git a/src/copy_elision.rs b/src/copy_elision.rs index 69cbdedf..3593890d 100644 --- a/src/copy_elision.rs +++ b/src/copy_elision.rs @@ -18,9 +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::checked::{CheckedExpr as Expr, CheckedFunction, LocalId, Reference}; +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 @@ -49,10 +50,11 @@ pub fn elidable_let_copies(decl: &CheckedFunction) -> HashSet { 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() { - let Expr::Let(local, ..) = &decl.arena[stmt] else { + if !matches!(decl.arena[stmt], Expr::Let(..)) { continue; - }; - if !is_value_aggregate(&decl.arena.local(*local).ty) { + } + let local = decl.arena.binder(stmt); + if !is_value_aggregate(&decl.arena.local(local).ty) { continue; } // This analysis requires a following sequence and keeps tail @@ -64,7 +66,7 @@ fn scan_blocks(id: ExprID, decl: &CheckedFunction, elidable: &mut HashSet bool { /// True if `name` is referenced anywhere in this subtree. fn mentions(id: ExprID, name: LocalId, decl: &CheckedFunction) -> bool { - if let Expr::Id(Reference::Local(n)) = &decl.arena[id] { - if *n == name { - return true; - } + if decl.arena.reference(id) == Some(&Reference::Local(name)) { + return true; } decl.arena[id] .subexprs() diff --git a/src/expr.rs b/src/expr.rs index 079be787..3ca86009 100644 --- a/src/expr.rs +++ b/src/expr.rs @@ -8,9 +8,9 @@ use crate::*; /// tree. It's also faster. Most hierarchical data /// should be represented this way. #[derive(Clone, Debug, Eq, PartialEq, Hash)] -pub enum Expr { +pub enum Expr { /// Identifier expression. - Id(R), + Id(Name), /// Integer literal, with optional explicit suffix. Int(i64, Option), @@ -31,7 +31,7 @@ pub enum Expr { Unop(Unop, ExprID), /// Lambda expression with parameters and body. - Lambda { params: Vec

, body: ExprID }, + Lambda { params: Vec, body: ExprID }, /// String literal. String(String), @@ -61,13 +61,13 @@ pub enum Expr { AsTy(ExprID, TypeID), /// Explicit type application: `name⟨i32⟩` or `name⟨i32, f32⟩`. - TypeApp(R, Vec), + TypeApp(Name, Vec), /// Immutable variable declaration with initializer and optional type. - Let(B, ExprID, Option), + Let(Name, ExprID, Option), /// Mutable variable declaration with optional initializer and type. - Var(B, Option, Option), + Var(Name, Option, Option), /// If expression with optional else branch. If(ExprID, ExprID, Option), @@ -77,7 +77,7 @@ pub enum Expr { /// For loop expression. For { - var: B, + var: Name, start: ExprID, end: ExprID, body: ExprID, @@ -115,7 +115,7 @@ pub enum Expr { Error, } -impl Expr { +impl Expr { /// The immediate subexpression IDs of this expression. /// /// Lambda yields its body: walks that treat a lambda specially still need diff --git a/src/free_locals.rs b/src/free_locals.rs index 30f7ca15..c543958c 100644 --- a/src/free_locals.rs +++ b/src/free_locals.rs @@ -1,5 +1,5 @@ //! Capture discovery needs lexical bindings, not solved types or publication. -use crate::{CheckedBody, Expr, ExprID, LocalId, Reference}; +use crate::{CheckedBody, ExprID, LocalId, Reference}; use std::collections::HashSet; /// The narrow facts needed from either a checked body or a check in progress. @@ -56,24 +56,14 @@ pub(crate) fn free_locals( impl BindingFacts for CheckedBody { fn binding_node(&self, id: ExprID) -> BindingNode { - let mut node = BindingNode { + BindingNode { children: self[id].subexprs(), - used: None, - declared: vec![], + used: match self.reference(id) { + Some(Reference::Local(local)) => Some(*local), + _ => None, + }, + declared: self.binders(id).to_vec(), complete: true, - }; - match &self[id] { - Expr::Id(Reference::Local(local)) | Expr::TypeApp(Reference::Local(local), _) => { - node.used = Some(*local); - } - Expr::Let(local, ..) | Expr::Var(local, ..) | Expr::For { var: local, .. } => { - node.declared.push(*local); - } - Expr::Lambda { params, .. } => { - node.declared.extend(params.iter().map(|param| param.local)); - } - _ => {} } - node } } diff --git a/src/hoist.rs b/src/hoist.rs index 21b82b2d..4053c6ba 100644 --- a/src/hoist.rs +++ b/src/hoist.rs @@ -85,16 +85,19 @@ impl SideEffects { } fn function_target(expr: ExprID, function: &CheckedFunction) -> Option { - match function.arena[expr] { - Expr::Id(Reference::Instance(instance)) => Some(instance), + match function.arena.reference(expr) { + Some(Reference::Instance(instance)) => Some(*instance), _ => None, } } fn storage_root(expr: ExprID, function: &CheckedFunction) -> Option { match &function.arena[expr] { - Expr::Id(Reference::Local(local)) => Some(Root::Local(*local)), - Expr::Id(Reference::Instance(instance)) => Some(Root::Global(*instance)), + 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, } @@ -286,43 +289,33 @@ fn create_hoisted_binding(read: &FieldRead, arena: &mut CheckedBody) -> (LocalId let Expr::Field(base, _) = arena[read.expr] else { unreachable!(); }; - let source_base = arena.node(base).clone(); - let source_field = arena.node(read.expr).clone(); + 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(source_base.kind, source_base.ty, source_base.loc); - let initializer = arena.add( - Expr::Field(base, read.field), - source_field.ty, - source_field.loc, - ); + 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)), - source_field.ty, + field_ty, false, ); - let declaration = arena.add( - Expr::Let(local, initializer, None), - mk_type(Type::Void), - source_field.loc, - ); + let declaration = arena.add_let(local, initializer, field_loc); (local, declaration) } fn invalidate_binders(expr: ExprID, function: &CheckedFunction, written: &mut WrittenFields) { - match &function.arena[expr] { - Expr::Let(local, ..) | Expr::Var(local, ..) | Expr::For { var: local, .. } => { - written.insert((Root::Local(*local), None)); - } - Expr::Lambda { params, .. } => { - written.extend( - params - .iter() - .map(|parameter| (Root::Local(parameter.local), None)), - ); - } - _ => {} - } + written.extend( + function + .arena + .binders(expr) + .iter() + .map(|&local| (Root::Local(local), None)), + ); } fn collect_written_fields( @@ -417,7 +410,7 @@ fn borrows_argument(callee: ExprID, position: usize, function: &CheckedFunction) } } -fn read_subexprs(expr: &CheckedExpr) -> Vec { +fn read_subexprs(expr: &Expr) -> Vec { match expr { Expr::Lambda { .. } | Expr::Arena(_) | Expr::Array(..) | Expr::Macro(..) => vec![], Expr::Binop(Binop::Assign, _, rhs) => vec![*rhs], @@ -462,11 +455,9 @@ fn replace_field_reads( if let Some(local) = storage_root(base, function).and_then(|root| substitutions.get(&(root, field))) { - function.arena.replace( - expr, - Expr::Id(Reference::Local(*local)), - function.arena.ty(expr), - ); + 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; } } @@ -530,9 +521,9 @@ mod tests { for local in [original, before] { assert!(function .arena - .nodes() - .iter() - .any(|node| node.kind == Expr::Id(Reference::Local(LocalId(local as u32))))); + .ids() + .any(|id| function.arena.reference(id) + == Some(&Reference::Local(LocalId(local as u32))))); } } diff --git a/src/jit.rs b/src/jit.rs index a707b5d8..eae3159a 100644 --- a/src/jit.rs +++ b/src/jit.rs @@ -3,10 +3,11 @@ use crate::cancel::*; use crate::checked::{ - CheckedDecl as Decl, CheckedExpr as Expr, CheckedFunction as FuncDecl, InstanceId, LocalId, - Reference, SpecializedProgram as DeclTable, + CheckedDecl as Decl, CheckedFunction as FuncDecl, InstanceId, LocalId, Reference, + SpecializedProgram as DeclTable, }; use crate::defs::*; +use crate::Expr; use crate::TypeID; extern crate cranelift_codegen; use core::panic; @@ -776,12 +777,15 @@ impl<'a> FunctionTranslator<'a> { fn translate_lvalue(&mut self, expr: ExprID, decl: &FuncDecl, decls: &DeclTable) -> Value { match &decl.arena[expr] { - Expr::Id(Reference::Local(local)) => self.builder.use_var(self.variables[local]), - Expr::Id(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) - } + 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) + } + reference => panic!("unresolved checked reference: {:?}", reference), + }, Expr::Field(lhs, name) => { let lhs_ty = decl.arena.ty(*lhs); let lhs_value = self.translate_lvalue(*lhs, decl, decls); @@ -855,34 +859,36 @@ impl<'a> FunctionTranslator<'a> { } } Expr::Char(c) => self.builder.ins().iconst(I8, *c as i64), - Expr::Id(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) - } - } - Expr::Id(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 + 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(), addr, 0) + .load(ty.cranelift_type(), MemFlags::new(), val, 0) } - } else { - self.translate_func(*instance, &*ty, decls) } - } - Expr::Id(reference) => panic!("unresolved checked reference: {:?}", reference), + 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.translate_func(*instance, &*ty, decls) + } + } + reference => panic!("unresolved checked reference: {:?}", reference), + }, Expr::Binop(op, lhs_id, rhs_id) => { self.translate_binop(*op, *lhs_id, *rhs_id, decl, decls) } @@ -891,14 +897,14 @@ impl<'a> FunctionTranslator<'a> { // 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(Reference::Instance(instance)) = &decl.arena[*fn_id] { + 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(Reference::Instance(instance)) = &decl.arena[*fn_id] { + 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); @@ -937,7 +943,7 @@ 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(Reference::Instance(callee)) = &decl.arena[*fn_id] { + 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 { @@ -953,7 +959,7 @@ impl<'a> FunctionTranslator<'a> { // Check if this is a math builtin that can use a direct call. let math_sym = - if let Expr::Id(Reference::Instance(instance)) = &decl.arena[*fn_id] { + if let Some(Reference::Instance(instance)) = decl.arena.reference(*fn_id) { math_builtin_symbol(&decls.instance_name(*instance)) } else { None @@ -961,7 +967,7 @@ impl<'a> FunctionTranslator<'a> { // Check if this is an extern function call. let is_extern_fn = - if let Expr::Id(Reference::Instance(callee)) = &decl.arena[*fn_id] { + if let Some(Reference::Instance(callee)) = decl.arena.reference(*fn_id) { decls .function_instance(*callee) .is_some_and(|function| function.is_extern) @@ -989,7 +995,7 @@ impl<'a> FunctionTranslator<'a> { } else if is_extern_fn { // Extern function: indirect call through {fn_ptr, context} in globals. let callee_name = - if let Expr::Id(Reference::Instance(n)) = &decl.arena[*fn_id] { + if let Some(Reference::Instance(n)) = decl.arena.reference(*fn_id) { *n } else { unreachable!() @@ -1082,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(Reference::Instance(instance)) if *decls.instance_name(*instance) == "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![]; @@ -1160,8 +1166,9 @@ impl<'a> FunctionTranslator<'a> { ); } } - Expr::Let(name, init, _) => { - let ty = &decl.arena.local(*name).ty; + 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); @@ -1174,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, @@ -1184,32 +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, *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); + 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, *ty); - self.let_bindings.insert(*name); + self.variable_types.insert(name, *ty); + self.let_bindings.insert(name); init_val } - Expr::Var(name, init, _) => { - let ty = &decl.arena.local(*name).ty; + 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); + 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) + && !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 { @@ -1217,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, *ty); - self.let_bindings.insert(*name); + 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, *ty); + let var = self.declare_variable(&name, I64); + self.variable_types.insert(name, *ty); let sz = ty.size(decls) as u32; if sz == 0 { @@ -1597,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 @@ -1609,13 +1614,14 @@ impl<'a> FunctionTranslator<'a> { let end_val = self.translate_expr(*end, decl, decls); // 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, crate::types::mk_type(crate::Type::Int32)); - self.let_bindings.insert(*var); + .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(); @@ -2292,7 +2298,7 @@ impl<'a> FunctionTranslator<'a> { self.builder.ins().store(MemFlags::new(), rhs, ptr, 0); return rhs; } - if let Expr::Id(Reference::Local(name)) = &decl.arena[lhs_id] { + 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; @@ -2523,18 +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(Reference::Local(local)) => { - if self.let_bindings.contains(local) { - None - } else { - Some(self.builder.use_var(self.variables[local])) + 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])) + } } - } - Expr::Id(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)) - } + 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)) } @@ -2559,7 +2568,7 @@ impl<'a> FunctionTranslator<'a> { self.builder.ins().store(MemFlags::new(), new_vec, ptr, 0); return; } - if let Expr::Id(Reference::Local(name)) = &decl.arena[vec_id] { + 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); @@ -2590,7 +2599,7 @@ impl<'a> FunctionTranslator<'a> { self.builder.ins().store(MemFlags::new(), value, addr, 0); return; } - if let Expr::Id(Reference::Local(name)) = &decl.arena[vec_id] { + 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 { @@ -2687,13 +2696,16 @@ impl<'a> FunctionTranslator<'a> { decls: &DeclTable, ) -> crate::TypeID { match &decl.arena[expr] { - Expr::Id(Reference::Local(local)) => self - .variable_types - .get(local) - .copied() - .unwrap_or(decl.arena.local(*local).ty), - Expr::Id(Reference::Instance(instance)) => match decls.instance(*instance) { - Decl::Global { ty, .. } => *ty, + 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) { diff --git a/src/llvm_jit.rs b/src/llvm_jit.rs index e62fda9b..7424a834 100644 --- a/src/llvm_jit.rs +++ b/src/llvm_jit.rs @@ -2,11 +2,11 @@ // Mirrors the Cranelift JIT backend in jit.rs. use crate::checked::{ - CheckedBody as ExprArena, CheckedDecl as Decl, CheckedExpr as Expr, - CheckedFunction as FuncDecl, CheckedParam as Param, InstanceId, LocalId, Reference, - SpecializedProgram as DeclTable, + CheckedBody as ExprArena, CheckedDecl as Decl, CheckedFunction as FuncDecl, InstanceId, + LocalId, Reference, SpecializedProgram as DeclTable, }; use crate::defs::*; +use crate::Expr; use crate::TypeID; use std::convert::TryFrom; @@ -1543,13 +1543,16 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { fn representation_type(&self, expr: ExprID, decl: &FuncDecl) -> crate::TypeID { match &decl.arena[expr] { - Expr::Id(Reference::Local(local)) => self - .variable_types - .get(local) - .copied() - .unwrap_or(decl.arena.local(*local).ty), - Expr::Id(Reference::Instance(instance)) => match self.decls.instance(*instance) { - Decl::Global { ty, .. } => *ty, + 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) { @@ -1709,7 +1712,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { expr }; if let Expr::Binop(Binop::Assign, lhs, rhs) = inner { - if let Expr::Id(Reference::Local(name)) = &arena[*lhs] { + if let Some(Reference::Local(name)) = arena.reference(*lhs) { return Some((*name, *rhs)); } } @@ -1720,47 +1723,54 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { fn translate_lvalue(&mut self, expr: ExprID, decl: &FuncDecl) -> PointerValue<'ctx> { match &decl.arena[expr] { - Expr::Id(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(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() { + 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 { - panic!("JIT lvalue: unknown variable {:?}", name) } - } - Expr::Id(Reference::Instance(instance)) => { - self.ptr_at_offset(self.globals_base, self.state.globals[instance] as u64) - } + 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.arena.ty(*lhs_id); let base_ptr = self.translate_lvalue(*lhs_id, decl); @@ -1782,27 +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(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); + 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 } - None - } - Expr::Id(Reference::Instance(instance)) => { - let offset = *self.state.globals.get(instance)?; - 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, } @@ -1911,26 +1924,36 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } } Expr::Char(c) => self.i8_ty().const_int(*c as u64, false).into(), - Expr::Id(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 + 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( @@ -1940,36 +1963,28 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { ) .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(), - &decl.arena.local(*name).name, - ) + .build_load(ty.llvm_basic_type(self.ctx()), addr, "global") .unwrap() } - } - } else { - panic!("missing local storage: {:?}", name) - } - } - Expr::Id(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()), addr, "global") - .unwrap() + self.translate_func_ref(*instance, &*ty) } - } else { - self.translate_func_ref(*instance, &*ty) } - } - Expr::Id(reference) => panic!("unresolved checked reference: {:?}", reference), + 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) @@ -1982,8 +1997,8 @@ 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); + 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); @@ -2022,8 +2037,8 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { self.let_bindings.insert(name); init_val } - Expr::Var(name, init_id, _) => { - let (name, init_id) = (*name, *init_id); + 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"); @@ -2307,12 +2322,9 @@ 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(); @@ -2437,9 +2449,9 @@ 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; @@ -3131,7 +3143,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { _call_expr_id: ExprID, decl: &FuncDecl, ) -> BasicValueEnum<'ctx> { - let is_builtin = if let Expr::Id(Reference::Instance(instance)) = &decl.arena[fn_id] { + let is_builtin = if let Some(Reference::Instance(instance)) = decl.arena.reference(fn_id) { is_builtin_name(&self.decls.instance_name(*instance)) } else { false @@ -3140,7 +3152,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { 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(Reference::Instance(callee)) = &decl.arena[fn_id] { + 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 { @@ -3164,7 +3176,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { }; // f32x4 constructor and splat. - if let Expr::Id(Reference::Instance(instance)) = &decl.arena[fn_id] { + 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); @@ -3195,14 +3207,14 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } // Check for math builtin by name. - let math_intrinsic = if let Expr::Id(Reference::Instance(instance)) = &decl.arena[fn_id] - { - self.llvm_intrinsic_name(&self.decls.instance_name(*instance)) - } 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(Reference::Instance(instance)) = &decl.arena[fn_id] { + if let Some(Reference::Instance(instance)) = decl.arena.reference(fn_id) { self.math_builtin_name(&self.decls.instance_name(*instance), from) } else { None @@ -3242,7 +3254,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } } else if { // Check for extern function. - if let Expr::Id(Reference::Instance(callee)) = &decl.arena[fn_id] { + if let Some(Reference::Instance(callee)) = decl.arena.reference(fn_id) { self.decls .function_instance(*callee) .is_some_and(|function| function.is_extern) @@ -3251,7 +3263,8 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { } } { // Extern function: load {fn_ptr, context} from globals buffer. - let callee_name = if let Expr::Id(Reference::Instance(n)) = &decl.arena[fn_id] { + let callee_name = if let Some(Reference::Instance(n)) = decl.arena.reference(fn_id) + { *n } else { unreachable!() @@ -3354,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(Reference::Instance(instance)) if *self.decls.instance_name(*instance) == "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() @@ -3847,7 +3860,7 @@ impl<'a, 'ctx> FunctionTranslator<'a, 'ctx> { fn translate_lambda( &mut self, - _params: &[Param], + _params: &[LocalId], _body: ExprID, expr_id: ExprID, decl: &FuncDecl, diff --git a/src/monomorph_pass.rs b/src/monomorph_pass.rs index 50975a65..fa0c3913 100644 --- a/src/monomorph_pass.rs +++ b/src/monomorph_pass.rs @@ -181,11 +181,15 @@ impl MonomorphPass { decls: &DeclTable, selections: &HashMap<(RequirementId, DefId), DefId>, ) -> Result<(), String> { - let (reference, explicit) = match body[id].clone() { - Expr::Id(reference) => (reference, None), - Expr::TypeApp(reference, arguments) => (reference, Some(arguments)), + 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(()), @@ -194,7 +198,7 @@ impl MonomorphPass { } Reference::Global(definition) => { let instance = self.instantiate_global(definition, explicit, solved, decls)?; - body.replace(id, Expr::Id(Reference::Instance(instance)), solved); + select_instance(body, id, instance); return Ok(()); } Reference::Functions(candidates) => candidates, @@ -225,7 +229,7 @@ impl MonomorphPass { let types = infer_type_arguments(target, solved)?; let instance = self.instantiate_function(definition, types, sizes, target, decls)?; - body.replace(id, Expr::Id(Reference::Instance(instance)), solved); + select_instance(body, id, instance); substitute_body_sizes(body, &bindings, size_vars); return Ok(()); } @@ -251,7 +255,7 @@ impl MonomorphPass { } let instance = self.instantiate_global(definition, explicit.clone(), solved, decls)?; - body.replace(id, Expr::Id(Reference::Instance(instance)), solved); + select_instance(body, id, instance); return Ok(()); } let target = decls @@ -293,7 +297,7 @@ impl MonomorphPass { .map(|parameter| bindings.get(¶meter.symbol).copied().unwrap_or(0)) .collect(); let instance = self.instantiate_function(definition, types, sizes, target, decls)?; - body.replace(id, Expr::Id(Reference::Instance(instance)), solved); + select_instance(body, id, instance); substitute_body_sizes(body, &bindings, size_vars); return Ok(()); } @@ -519,6 +523,16 @@ fn subst_size_vars(ty: TypeID, bindings: &HashMap) -> TypeID { } } +/// 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, @@ -538,9 +552,11 @@ fn substitute_body_sizes( for id in 0..body.len() { let mut kind = body[id].clone(); match &mut kind { - Expr::Id(Reference::SizeParameter(local)) => { - if let Some(value) = values.get(local) { - kind = Expr::Int(i64::from(*value), None); + 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) => { @@ -631,10 +647,9 @@ mod tests { fn targets(function: &CheckedFunction) -> Vec { function .arena - .nodes() - .iter() - .filter_map(|node| match node.kind { - Expr::Id(Reference::Instance(id)) => Some(id), + .ids() + .filter_map(|id| match function.arena.reference(id) { + Some(Reference::Instance(id)) => Some(*id), _ => None, }) .collect() @@ -664,8 +679,9 @@ mod tests { _ => None, }) .unwrap(); - main.arena.add( - Expr::Id(Reference::Functions(vec![definition])), + main.arena.add_id( + Name::str("target"), + Reference::Functions(vec![definition]), ty, test_loc(), ); @@ -697,9 +713,9 @@ mod tests { let specialized = output.find_entry_point(Name::str("probe$3")).unwrap(); assert!(specialized .arena - .nodes() + .exprs() .iter() - .any(|node| node.kind == Expr::Int(3, None))); + .any(|expr| *expr == Expr::Int(3, None))); assert_eq!( specialized.param_types(), vec![mk_type(Type::Array( @@ -738,8 +754,9 @@ mod tests { 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( - Expr::Id(Reference::Global(global)), + let reference = arena.add_id( + Name::str("limit"), + Reference::Global(global), mk_type(Type::Int32), test_loc(), ); @@ -771,7 +788,7 @@ mod tests { _ => None, }) .unwrap(); - let Expr::Id(Reference::Instance(instance)) = assumption[reference] else { + let Some(&Reference::Instance(instance)) = assumption.reference(reference) else { panic!("assumption retained a generic-phase reference"); }; assert_eq!(output.instances[instance.index()].definition, global); @@ -841,9 +858,8 @@ mod tests { .unwrap(); assert!(main .arena - .nodes() - .iter() - .any(|node| node.kind == Expr::Id(Reference::Local(LocalId(binding as u32))))); + .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()); } @@ -855,19 +871,18 @@ mod tests { assert_eq!( function .arena - .nodes() + .exprs() .iter() - .filter(|node| node.kind == Expr::Int(3, None)) + .filter(|expr| **expr == Expr::Int(3, None)) .count(), 2 ); - assert!(function.arena.nodes().iter().any(|node| matches!(node.kind, - Expr::Id(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::Local(local)) if function.arena.local(*local).name == Name::str("N")))); assert!(!function .arena - .nodes() - .iter() - .any(|node| matches!(node.kind, Expr::Id(Reference::SizeParameter(_))))); + .ids() + .any(|id| matches!(function.arena.reference(id), Some(Reference::SizeParameter(_))))); } #[test] diff --git a/src/safety_checker.rs b/src/safety_checker.rs index f5a5672f..59258af6 100644 --- a/src/safety_checker.rs +++ b/src/safety_checker.rs @@ -1,4 +1,4 @@ -use crate::checked::{CheckedBody as ExprArena, CheckedExpr as Expr, CheckedFunction as FuncDecl}; +use crate::checked::{CheckedBody as ExprArena, CheckedFunction as FuncDecl}; use crate::interval::{enclose, IndexInterval}; use crate::*; @@ -122,7 +122,7 @@ fn reference_place(reference: &Reference) -> Option { } fn id_place(id: ExprID, arena: &ExprArena) -> Option { match &arena[id] { - Expr::Id(reference) => reference_place(reference), + Expr::Id(_) => arena.reference(id).and_then(reference_place), _ => None, } } @@ -161,7 +161,7 @@ fn collect_size_subst(param_ty: TypeID, arg_ty: TypeID, out: &mut Vec<(Name, i64 /// A trackable storage place, including direct field projections. fn expr_place(id: ExprID, arena: &ExprArena) -> Option { match &arena[id] { - Expr::Id(reference) => reference_place(reference), + Expr::Id(_) => arena.reference(id).and_then(reference_place), Expr::Field(base, field) => Some(expr_place(*base, arena)?.field(*field)), _ => None, } @@ -736,10 +736,11 @@ impl SafetyChecker { } IndexInterval::default() } - Expr::Let(local, init, _) => { - let name = &Place::local(*local); + 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; + let ty = context.arena.local(local).ty; // Track the interval from the initializer. let mut min = if init_r.min != i64::MIN { @@ -778,14 +779,15 @@ impl SafetyChecker { IndexInterval::default() } - Expr::Var(local, init, _) => { - let name = &Place::local(*local); + 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, context, decls) } else { IndexInterval::default() }; - let ty = context.arena.local(*local).ty; + let ty = context.arena.local(local).ty; let mut min = if init_r.min != i64::MIN { Some(init_r.min) @@ -825,8 +827,8 @@ impl SafetyChecker { IndexInterval::default() } - Expr::Id(reference) => { - let Some(place) = reference_place(reference) else { + Expr::Id(_) => { + let Some(place) = context.arena.reference(expr).and_then(reference_place) else { return IndexInterval::default(); }; let name = &place; @@ -1135,7 +1137,8 @@ impl SafetyChecker { .collect(); // An immediately-invoked lambda has known arguments, so check // its body against them rather than unconstrained. - if let Expr::Lambda { params, body } = &context.arena[*callee_expr] { + 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, context, decls); @@ -1208,12 +1211,9 @@ impl SafetyChecker { } } Expr::For { - var, - start, - end, - body, + start, end, body, .. } => { - let var = &Place::local(*var); + 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); @@ -1289,19 +1289,14 @@ 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| { - context - .arena - .nodes() - .iter() - .map(|node| &node.kind) - .any(|e| { - if let Expr::Var(vn, Some(init), _) = e { - Place::local(*vn) == name - && id_place(*init, context.arena) == Some(sn) - } else { - false - } - }) + 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 Some(ref end_name) = id_place(*end, context.arena) { @@ -1325,11 +1320,12 @@ impl SafetyChecker { } 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.) + let params = context.arena.binders(expr); self.check_lambda_body(params, *body, None, context, decls); IndexInterval::default() } @@ -1358,7 +1354,7 @@ impl SafetyChecker { /// at the definition site parameter values are unknown. fn check_lambda_body( &mut self, - params: &[CheckedParam], + params: &[LocalId], body: ExprID, call_args: Option<(&[ExprID], &[IndexInterval])>, context: SafetyBody<'_>, @@ -1382,12 +1378,12 @@ impl SafetyChecker { } } - for (i, param) in params.iter().enumerate() { - let ty = context.arena.local(param.local).ty; + for (i, ¶m) in params.iter().enumerate() { + let ty = context.arena.local(param).ty; let is_u32 = ty == mk_type(Type::UInt32); // A repeated analysis of this lambda starts with fresh parameter facts. - self.forget(Place::local(param.local)); + self.forget(Place::local(param)); let arg = call_args.and_then(|(exprs, ivals)| Some((exprs.get(i)?, ivals.get(i)?))); match arg { @@ -1397,9 +1393,9 @@ impl SafetyChecker { if is_u32 { min = Some(min.unwrap_or(0).max(0)); } - self.add(Place::local(param.local), min, max); + self.add(Place::local(param), min, max); if ival.non_zero { - self.add_non_zero(Place::local(param.local)); + self.add_non_zero(Place::local(param)); } // The param inherits the argument's symbolic length bounds. // Read the live state so bounds from earlier parameters @@ -1413,14 +1409,14 @@ impl SafetyChecker { .collect(); for array in inherited { self.len_bounds.push(LenBound { - index: Place::local(param.local), + index: Place::local(param), array, }); } } } - None if is_u32 => self.add(Place::local(param.local), Some(0), None), - None => self.add(Place::local(param.local), None, None), + None if is_u32 => self.add(Place::local(param), Some(0), None), + None => self.add(Place::local(param), None, None), } } @@ -1534,7 +1530,10 @@ impl SafetyChecker { caller: SafetyBody<'_>, decls: &impl SafetyProgram, ) { - let Expr::Id(reference) = &caller.arena[callee_expr] else { + if !matches!(caller.arena[callee_expr], Expr::Id(_)) { + return; + } + let Some(reference) = caller.arena.reference(callee_expr) else { return; }; let Some(callee) = decls.call_target( @@ -1641,11 +1640,15 @@ impl SafetyChecker { && self.prove_at_call(*rhs, callee, caller, subst, size_subst, decls) } Expr::Binop(Binop::Less, lhs, rhs) => { - if let (Expr::Id(index), Expr::Field(array, field)) = + if let (Expr::Id(_), Expr::Field(array, field)) = (&callee.arena[*lhs], &callee.arena[*rhs]) { if field.as_str() == "len" { - if let Expr::Id(array) = &callee.arena[*array] { + 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)) { @@ -1671,21 +1674,32 @@ impl SafetyChecker { } } } - if let (Expr::Id(lhs), Expr::Id(rhs)) = (&callee.arena[*lhs], &callee.arena[*rhs]) { - if let (Some(argument), Some(size)) = (lookup(lhs), size_lookup(rhs)) { + 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; } } } - let left = if let Expr::Id(reference) = &callee.arena[*lhs] { - lookup(reference).map(|argument| self.check_expr(argument, caller, decls)) + 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::callee_interval(*lhs, callee, decls)) }; - let right = if let Expr::Id(reference) = &callee.arena[*rhs] { - lookup(reference).map(|argument| self.check_expr(argument, caller, decls)) + 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::callee_interval(*rhs, callee, decls)) }; @@ -1697,8 +1711,8 @@ impl SafetyChecker { } } Expr::Binop(Binop::Geq, lhs, rhs) => { - if let Expr::Id(reference) = &callee.arena[*lhs] { - if let Some(argument) = lookup(reference) { + 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 @@ -1829,9 +1843,11 @@ impl SafetyChecker { graph: &mut Graph, ) { match &body[expression] { - Expr::Id(reference) | Expr::TypeApp(reference, _) => { + Expr::Id(_) | Expr::TypeApp(_, _) => { if !in_callee { - graph.address_taken.extend(definitions(reference, body)); + if let Some(reference) = body.reference(expression) { + graph.address_taken.extend(definitions(reference, body)); + } } } Expr::Lambda { @@ -1910,8 +1926,12 @@ impl SafetyChecker { unreachable!(); }; let direct = match &function.arena[*callee] { - Expr::Id(reference) | Expr::TypeApp(reference, _) => { - let definitions = definitions(reference, &function.arena); + 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() { @@ -2051,17 +2071,17 @@ mod tests { }; let callee = caller .arena - .nodes() + .exprs() .iter() - .find_map(|node| { - if let Expr::Call(callee, _) = node.kind { - Some(callee) + .find_map(|expr| { + if let Expr::Call(callee, _) = expr { + Some(*callee) } else { None } }) .unwrap(); - let Expr::Id(Reference::Instance(target)) = caller.arena[callee] else { + let Some(&Reference::Instance(target)) = caller.arena.reference(callee) else { unreachable!() }; let integer = mk_type(Type::Int32); @@ -2074,9 +2094,7 @@ mod tests { mk_type(Type::Tuple(parameters)), mk_type(Type::Void), )); - caller - .arena - .replace(callee, Expr::Id(Reference::Instance(target)), recorded); + 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())); diff --git a/src/stack_codegen.rs b/src/stack_codegen.rs index d7e249e9..6c8c6e9d 100644 --- a/src/stack_codegen.rs +++ b/src/stack_codegen.rs @@ -4,11 +4,10 @@ //! executed by a stack-based virtual machine. It mirrors the register-based //! VM codegen but emits stack IR instructions instead. -use crate::checked::{ - CheckedExpr as Expr, CheckedFunction, InstanceId, LocalId, Reference, SpecializedProgram, -}; +use crate::checked::{CheckedFunction, InstanceId, LocalId, Reference, SpecializedProgram}; use crate::decl::Decl; use crate::defs::*; +use crate::expr::Expr; use crate::stack_ir::*; use crate::types::*; use std::collections::{HashMap, HashSet}; @@ -419,14 +418,17 @@ impl<'a> FunctionTranslator<'a> { /// has an array address and must explicitly build the slice fat pointer. fn representation_type(&self, expr: ExprID) -> TypeID { match &self.decl.arena[expr] { - Expr::Id(Reference::Local(local)) => { - let ty = self.decl.arena.local(*local).ty; - match &*ty { - Type::Reference(inner) => *inner, - _ => ty, + 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, + } } - } - Expr::Id(Reference::Instance(instance)) => self.decls.instance(*instance).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), @@ -767,7 +769,9 @@ impl<'a> FunctionTranslator<'a> { func.emit(StackOp::LocalAddr(mem_slot)); } - Expr::Id(reference) => { + Expr::Id(_) => { + let decl = self.decl; + let reference = decl.arena.reference(expr).expect("checked reference"); self.translate_id(reference, expr, func); } @@ -806,8 +810,8 @@ 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.decl.arena.local(name).ty; @@ -874,8 +878,8 @@ impl<'a> FunctionTranslator<'a> { } } - Expr::Var(name, init, _) => { - let name = *name; + Expr::Var(_, init, _) => { + let name = self.decl.arena.binder(expr); let init = *init; let ty = self.decl.arena.local(name).ty; @@ -995,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)); } @@ -1243,7 +1245,7 @@ impl<'a> FunctionTranslator<'a> { if self.holds_fat_pointer(*fn_id) { return None; } - let Expr::Id(Reference::Instance(instance)) = &self.decl.arena[*fn_id] else { + let Some(Reference::Instance(instance)) = self.decl.arena.reference(*fn_id) else { return None; }; let name = self.decls.instance_name(*instance); @@ -1523,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(Reference::Local(name)) = &self.decl.arena[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); @@ -1547,7 +1549,7 @@ impl<'a> FunctionTranslator<'a> { } // Direct scalar local assignment. - if let Expr::Id(Reference::Local(name)) = &self.decl.arena[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 @@ -1711,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(Reference::Local(name)) = &self.decl.arena[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); } @@ -1721,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(Reference::Local(name)) = &self.decl.arena[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); } @@ -1774,32 +1776,35 @@ 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[expr].clone() { - Expr::Id(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) { + 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 { - unreachable!("checked local must have storage"); } - } - Expr::Id(Reference::Instance(instance)) => { - func.emit(StackOp::GlobalAddr(self.globals[instance])); - } + 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; @@ -1920,7 +1925,7 @@ impl<'a> FunctionTranslator<'a> { } // Check for builtin functions. - if let Expr::Id(Reference::Instance(instance)) = &self.decl.arena[fn_id] { + if let Some(Reference::Instance(instance)) = self.decl.arena.reference(fn_id) { let instance = *instance; let name = self.decls.instance_name(instance); @@ -2336,9 +2341,9 @@ 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 { - match &self.decl.arena[fn_id] { - Expr::Id(Reference::Local(_)) => true, - Expr::Id(Reference::Instance(id)) => { + match self.decl.arena.reference(fn_id) { + Some(Reference::Local(_)) => true, + Some(Reference::Instance(id)) => { matches!(self.decls.instance(*id), Decl::Global { .. }) } _ => false, diff --git a/src/vm_codegen.rs b/src/vm_codegen.rs index 7749cfa4..46287345 100644 --- a/src/vm_codegen.rs +++ b/src/vm_codegen.rs @@ -4,11 +4,11 @@ //! executed by the register-based virtual machine. use crate::checked::{ - CheckedBody, CheckedExpr as Expr, CheckedFunction, InstanceId, LocalId, Reference, - SpecializedProgram, + CheckedBody, CheckedFunction, InstanceId, LocalId, Reference, SpecializedProgram, }; use crate::decl::Decl; use crate::defs::*; +use crate::expr::Expr; use crate::types::*; use crate::vm::*; use std::collections::{HashMap, HashSet}; @@ -782,14 +782,17 @@ impl<'a> FunctionTranslator<'a> { /// pointer itself. fn representation_type(&self, expr: ExprID) -> TypeID { match &self.body.decl.arena[expr] { - Expr::Id(Reference::Local(local)) => { - let ty = self.body.decl.arena.local(*local).ty; - match &*ty { - Type::Reference(inner) => *inner, - _ => ty, + 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, + } } - } - Expr::Id(Reference::Instance(instance)) => self.decls.instance(*instance).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), @@ -839,127 +842,129 @@ impl<'a> FunctionTranslator<'a> { dst } - Expr::Id(Reference::Local(name)) => { - let ty = self.expr_type(expr); + Expr::Id(_) => match self.body.decl.arena.reference(expr) { + Some(Reference::Local(name)) => { + let ty = self.expr_type(expr); - // 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 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; - } - - // Check if it's a local variable. - if let Some(®) = self.body.variables.get(name) { - if self.body.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; + return captured_addr; } + // Now load the value from the captured variable's storage. 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 }); + self.emit_load(&ty, dst, captured_addr, func); + return dst; + } + + // 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 } - 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 + unreachable!("checked local must have storage") } - } else { - unreachable!("checked local must have storage") } - } - Expr::Id(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 { - // 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 { - unreachable!("instance must name storage or a function") + // 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") + } } } - } - Expr::Id(_) => unreachable!("non-concrete reference in specialized body"), + _ => unreachable!("non-concrete reference in specialized body"), + }, Expr::Binop(op, lhs_id, rhs_id) => self.translate_binop(*op, *lhs_id, *rhs_id, func), @@ -967,7 +972,9 @@ impl<'a> FunctionTranslator<'a> { Expr::Call(fn_id, arg_ids) => self.translate_call(*fn_id, arg_ids, expr, func), - Expr::Let(name, init, _) => { + 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); @@ -1023,7 +1030,9 @@ impl<'a> FunctionTranslator<'a> { init_reg } - Expr::Var(name, init, _) => { + 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.body.lambda_referenced.contains(name) { @@ -1160,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. @@ -1968,7 +1977,7 @@ 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(Reference::Local(name)) = &self.body.decl.arena[lhs_id] { + 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); @@ -1989,7 +1998,7 @@ impl<'a> FunctionTranslator<'a> { } } // Check for direct register-promoted scalar assignment (e.g., `x = expr`). - if let Expr::Id(Reference::Local(name)) = &self.body.decl.arena[lhs_id] { + 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.body.variables.get(name).unwrap(); @@ -2046,33 +2055,36 @@ impl<'a> FunctionTranslator<'a> { /// Translate an lvalue expression (returns address). fn translate_lvalue(&mut self, expr: ExprID, func: &mut VMFunction) -> Reg { match &self.body.decl.arena[expr] { - Expr::Id(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 + 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 { - unreachable!("checked local must have storage") } - } - Expr::Id(Reference::Instance(instance)) => { - let dst = self.alloc_reg(); - func.emit(Opcode::GlobalAddr { - dst, - offset: self.globals[instance], - }); - dst - } + Some(Reference::Instance(instance)) => { + let dst = self.alloc_reg(); + 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); @@ -2241,7 +2253,7 @@ impl<'a> FunctionTranslator<'a> { } // Special handling for built-in functions. - if let Expr::Id(Reference::Instance(instance)) = &self.body.decl.arena[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" { @@ -2654,9 +2666,9 @@ 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 { - match &self.body.decl.arena[fn_id] { - Expr::Id(Reference::Local(_)) => true, - Expr::Id(Reference::Instance(id)) => { + match self.body.decl.arena.reference(fn_id) { + Some(Reference::Local(_)) => true, + Some(Reference::Instance(id)) => { matches!(self.decls.instance(*id), Decl::Global { .. }) } _ => false, @@ -3680,21 +3692,21 @@ mod tests { 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(Expr::Var(counter, Some(zero), None), void, 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(Expr::Id(Reference::Local(counter)), ty, loc); + 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(Expr::Id(Reference::Local(counter)), ty, 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(Expr::Id(Reference::Local(counter)), ty, loc); + let result = arena.add_local_read(counter, ty, loc); arena.add(Expr::Block(vec![binding, loop_expr, result]), ty, loc); assert_eq!( @@ -3732,7 +3744,7 @@ mod tests { 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(Expr::Id(Reference::Local(parameter)), ty, loc); + 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( @@ -3748,23 +3760,25 @@ mod tests { "the two bodies intentionally reuse local index zero" ); let forty = main_body.add(Expr::Int(40, None), ty, loc); - let binding = main_body.add(Expr::Let(outer, forty, None), mk_type(Type::Void), loc); - let target = main_body.add( - Expr::Id(Reference::Instance(InstanceId(2))), + 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( - Expr::Id(Reference::Instance(InstanceId(2))), + 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(Expr::Id(Reference::Local(outer)), 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);