diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index 4f480604e..0e4d37326 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -65,3 +65,32 @@ If `net_added_lines` is positive, add: - [ ] No `#[allow(clippy::...)]` added outside generated code. - [ ] Every commit signed off (`git commit -s`); no `Co-Authored-By` for AI. - [ ] External material (if any) attributed and Apache-2.0 compatible. + +## Authorized Broken-Main Recovery (optional) + + + +- `authorization_ref` (durable validation integration PR body or maintainer-controlled GitHub artifact): +- `authorized_by` (ClimaMind Rumoca repository maintainer): +- `batch_id`: +- Authorized ordered `owner_prs`: +- `target_branch` / baseline `head_sha`: +- RFC 3339 UTC `expires_at`: +- Owner PR / final `head_sha`: +- Independent technical review / reviewed `head_sha`: +- Owner mechanism test / tested `head_sha`: +- Recovery-rule PR / final `head_sha`: +- Integration PR / `head_sha`: +- Hosted CI workflow / `head_sha`: +- [ ] Authorization exists, matches this merge, and is unexpired. +- [ ] `authorized_by` is a ClimaMind Rumoca maintainer; this task's explicit authorization is sufficient, with no additional maintainer or approval. +- [ ] Recovery is inactive and fails closed if `expires_at` passed, every authorized owner PR landed, or target `main` has all required CI green. +- [ ] Evidence is bound to the owner final head and recorded in order; merge only after all required hosted CI is green on the integration head. +- [ ] Evidence order is recorded without skipping: authorization verification; independent technical review on final owner `head_sha`; passing owner mechanism test on that same `head_sha`; exact-head integration; hosted CI green on the integration head; then merge. No later step occurs before its predecessor. +- [ ] Every listed final owner head, including the recovery-rule PR head, is a Git ancestor of the integration head; no cherry-pick, patch-id, squash, or content-equivalent substitute. +- [ ] Recorded target baseline `head_sha` is a Git ancestor of the integration head. +- [ ] Integration history = target baseline + listed exact owner histories + signed merge commits only; no integration-only production, test, spec, workflow, baseline, validator, tolerance, fixture, or content commit. +- [ ] Every integration merge commit has exactly one `Signed-off-by` trailer and no `Co-Authored-By` trailer. +- [ ] CI workflow head = integration head. +- [ ] Any owner, baseline, or integration head change fails closed; rebuild and rerun affected evidence and CI. +- [ ] Draft, validation-only, never merge. diff --git a/.github/scripts/msl-nix-closure.sh b/.github/scripts/msl-nix-closure.sh new file mode 100755 index 000000000..ae199cd11 --- /dev/null +++ b/.github/scripts/msl-nix-closure.sh @@ -0,0 +1,184 @@ +#!/usr/bin/env bash +set -euo pipefail + +readonly archive_name=closure.nar +readonly manifest_name=manifest +readonly -a manifest_fields=(version commit system out_path archive_sha256) +readonly -a required_binaries=( + msl_tests + rumoca-worker + rumoca-sim-worker + rumoca-msl-tools +) +declare -a cleanup_paths=() + +cleanup() { + if ((${#cleanup_paths[@]})); then + rm -f -- "${cleanup_paths[@]}" + fi +} +trap cleanup EXIT + +die() { + echo "msl-nix-closure: $*" >&2 + exit 1 +} + +require_inputs() { + local commit=$1 + local system=$2 + [[ $commit =~ ^[0-9a-fA-F]{40}$ ]] || die "invalid Git commit: $commit" + [[ $system =~ ^[a-zA-Z0-9._+-]+$ ]] || die "invalid Nix system: $system" +} + +manifest_value() { + local manifest=$1 + local key=$2 + local -a values + mapfile -t values < <( + awk -v key="$key" 'index($0, key "=") == 1 { print substr($0, length(key) + 2) }' "$manifest" + ) + [[ ${#values[@]} -eq 1 && -n ${values[0]} ]] || die "invalid manifest field: $key" + printf '%s\n' "${values[0]}" +} + +verify_executables() { + local out_path=$1 + local manifest=${2-} + local binary path actual expected + for binary in "${required_binaries[@]}"; do + path="$out_path/bin/$binary" + [[ -f $path && -x $path ]] || die "missing required executable: $binary" + if [[ -n $manifest ]]; then + expected=$(manifest_value "$manifest" "binary_${binary}_sha256") + [[ $expected =~ ^[0-9a-f]{64}$ ]] || die "invalid executable checksum: $binary" + actual=$(sha256sum "$path" | awk '{ print $1 }') + [[ $actual == "$expected" ]] || die "required executable checksum mismatch: $binary" + fi + done +} + +pack() { + [[ $# -eq 4 ]] || die "usage: $0 pack OUT_LINK ARTIFACT_DIR COMMIT SYSTEM" + local out_link=$1 + local artifact_dir=$2 + local commit=$3 + local system=$4 + local out_path archive manifest paths_file archive_tmp binary + local -a closure_paths + + require_inputs "$commit" "$system" + command -v nix-store >/dev/null || die "nix-store is required" + command -v sha256sum >/dev/null || die "sha256sum is required" + [[ -e $out_link ]] || die "missing realized MSL output: $out_link" + out_path=$(realpath "$out_link") + [[ $out_path == /* && $out_path != *[[:space:]]* ]] || die "invalid Nix output path: $out_path" + verify_executables "$out_path" + + mkdir -p "$artifact_dir" + archive="$artifact_dir/$archive_name" + manifest="$artifact_dir/$manifest_name" + paths_file="$artifact_dir/closure-paths.tmp" + archive_tmp="$archive.tmp" + rm -f "$archive" "$manifest" "$paths_file" "$archive_tmp" + cleanup_paths=("$paths_file" "$archive_tmp" "$archive" "$manifest") + + nix-store --query --requisites "$out_path" > "$paths_file" || + die "failed to query Nix requisites" + grep -Fxq "$out_path" "$paths_file" || printf '%s\n' "$out_path" >> "$paths_file" + LC_ALL=C sort -u -o "$paths_file" "$paths_file" + mapfile -t closure_paths < "$paths_file" + [[ ${#closure_paths[@]} -gt 0 ]] || die "Nix requisites closure is empty" + for out_path in "${closure_paths[@]}"; do + [[ $out_path == /* && $out_path != *[[:space:]]* ]] || + die "invalid path in Nix requisites closure: $out_path" + done + nix-store --export "${closure_paths[@]}" > "$archive_tmp" || + die "failed to export Nix requisites closure" + [[ -s $archive_tmp ]] || die "exported Nix closure archive is empty" + mv "$archive_tmp" "$archive" + + { + echo 'version=1' + echo "commit=$commit" + echo "system=$system" + echo "out_path=$(realpath "$out_link")" + echo "archive_sha256=$(sha256sum "$archive" | awk '{ print $1 }')" + for binary in "${required_binaries[@]}"; do + echo "binary_${binary}_sha256=$(sha256sum "$(realpath "$out_link")/bin/$binary" | awk '{ print $1 }')" + done + } > "$manifest" + rm -f "$paths_file" + cleanup_paths=() +} + +restore() { + [[ $# -eq 4 ]] || die "usage: $0 restore ARTIFACT_DIR OUT_LINK COMMIT SYSTEM" + local artifact_dir=$1 + local out_link=$2 + local expected_commit=$3 + local expected_system=$4 + local archive="$artifact_dir/$archive_name" + local manifest="$artifact_dir/$manifest_name" + local version commit system out_path expected_sha actual_sha paths_file path + local expected_manifest_lines + local -a closure_paths + + require_inputs "$expected_commit" "$expected_system" + command -v nix-store >/dev/null || die "nix-store is required" + command -v sha256sum >/dev/null || die "sha256sum is required" + [[ -f $manifest ]] || die "missing manifest: $manifest" + [[ -f $archive ]] || die "missing closure archive: $archive" + expected_manifest_lines=$((${#manifest_fields[@]} + ${#required_binaries[@]})) + [[ $(wc -l < "$manifest") -eq $expected_manifest_lines ]] || + die "invalid manifest structure" + + version=$(manifest_value "$manifest" version) + commit=$(manifest_value "$manifest" commit) + system=$(manifest_value "$manifest" system) + out_path=$(manifest_value "$manifest" out_path) + expected_sha=$(manifest_value "$manifest" archive_sha256) + [[ $version == 1 ]] || die "unsupported manifest version: $version" + require_inputs "$commit" "$system" + [[ $out_path == /* && $out_path != *[[:space:]]* ]] || die "invalid Nix output path: $out_path" + [[ $expected_sha =~ ^[0-9a-f]{64}$ ]] || die "invalid archive checksum" + [[ $commit == "$expected_commit" ]] || die "commit mismatch: expected $expected_commit, got $commit" + [[ $system == "$expected_system" ]] || die "system mismatch: expected $expected_system, got $system" + actual_sha=$(sha256sum "$archive" | awk '{ print $1 }') + [[ $actual_sha == "$expected_sha" ]] || die "archive checksum mismatch" + + nix-store --import < "$archive" || die "Nix closure import failed" + [[ -e $out_path ]] || die "imported Nix output is missing: $out_path" + paths_file=$(mktemp) + cleanup_paths=("$paths_file") + nix-store --query --requisites "$out_path" > "$paths_file" || + die "failed to query imported Nix requisites" + mapfile -t closure_paths < "$paths_file" + [[ ${#closure_paths[@]} -gt 0 ]] || die "imported Nix requisites closure is empty" + for path in "${closure_paths[@]}"; do + [[ -e $path ]] || die "imported Nix requisite is missing: $path" + done + verify_executables "$out_path" "$manifest" + + if [[ -e $out_link && ! -L $out_link ]]; then + die "refusing to replace non-symlink output: $out_link" + fi + rm -f "$out_link" + ln -s "$out_path" "$out_link" + rm -f "$paths_file" + cleanup_paths=() +} + +case ${1-} in + pack) + shift + pack "$@" + ;; + restore) + shift + restore "$@" + ;; + *) + die "usage: $0 {pack|restore} ..." + ;; +esac diff --git a/.github/scripts/msl-nix-closure.test.mjs b/.github/scripts/msl-nix-closure.test.mjs new file mode 100644 index 000000000..eccbd04cb --- /dev/null +++ b/.github/scripts/msl-nix-closure.test.mjs @@ -0,0 +1,355 @@ +import assert from 'node:assert/strict'; +import { createHash } from 'node:crypto'; +import { + chmodSync, + existsSync, + mkdtempSync, + mkdirSync, + readFileSync, + readdirSync, + realpathSync, + rmSync, + writeFileSync, +} from 'node:fs'; +import { tmpdir } from 'node:os'; +import { join } from 'node:path'; +import { fileURLToPath } from 'node:url'; +import { spawnSync } from 'node:child_process'; +import { afterEach, test } from 'node:test'; + +const repoFile = (path) => new URL(`../../${path}`, import.meta.url); +const helper = fileURLToPath(repoFile('.github/scripts/msl-nix-closure.sh')); +const requiredBinaries = [ + 'msl_tests', + 'rumoca-worker', + 'rumoca-sim-worker', + 'rumoca-msl-tools', +]; +const commit = '0123456789abcdef0123456789abcdef01234567'; +const system = 'x86_64-linux'; +const fixtureRoots = new Set(); + +const cleanupFixtures = () => { + for (const root of fixtureRoots) { + rmSync(root, { force: true, recursive: true }); + } + fixtureRoots.clear(); +}; + +afterEach(cleanupFixtures); + +const sha256 = (value) => createHash('sha256').update(value).digest('hex'); + +const workflowJob = (workflow, job) => { + const start = workflow.indexOf(` ${job}:`); + assert.notEqual(start, -1, `${job} must exist`); + const next = workflow.slice(start + 1).search(/^ [a-zA-Z0-9_-]+:$/m); + return workflow.slice(start, next === -1 ? undefined : start + 1 + next); +}; + +const fixture = ({ exportFails = false, importFails = false } = {}) => { + const root = realpathSync(mkdtempSync(join(tmpdir(), 'rumoca-msl-closure-'))); + fixtureRoots.add(root); + const artifactDir = join(root, 'artifact'); + const outPath = join(root, 'nix-store', 'msl-output'); + const outLink = join(root, 'result-msl-artifacts'); + const fakeBin = join(root, 'fake-bin'); + const nixLog = join(root, 'nix-store.log'); + const tempDir = join(root, 'tmp'); + mkdirSync(artifactDir, { recursive: true }); + mkdirSync(join(outPath, 'bin'), { recursive: true }); + mkdirSync(join(root, 'nix-store', 'dependency')); + mkdirSync(fakeBin); + mkdirSync(tempDir); + for (const binary of requiredBinaries) { + writeFileSync(join(outPath, 'bin', binary), `${binary}\n`); + chmodSync(join(outPath, 'bin', binary), 0o755); + } + const fakeNixStore = `#!/usr/bin/env bash +set -euo pipefail +printf '%s\\n' "$*" >> "${nixLog}" +case "\${1-}" in + --query) + printf '%s\\n' "${join(root, 'nix-store', 'dependency')}" "${outPath}" + ;; + --export) + ${exportFails ? "printf 'partial-archive'; exit 24" : "printf 'complete-closure-archive'"} + ;; + --import) + cat >/dev/null + ${importFails ? 'exit 23' : "printf '%s\\n' '${outPath}'"} + ;; + --verify-path) + ;; + *) + echo "unexpected nix-store invocation: $*" >&2 + exit 97 + ;; +esac +`; + writeFileSync(join(fakeBin, 'nix-store'), fakeNixStore); + chmodSync(join(fakeBin, 'nix-store'), 0o755); + return { + artifactDir, + env: { ...process.env, PATH: `${fakeBin}:${process.env.PATH}`, TMPDIR: tempDir }, + nixLog, + outLink, + outPath, + root, + tempDir, + }; +}; + +const writeArtifact = (fx, overrides = {}) => { + const archive = overrides.archive ?? 'complete-closure-archive'; + writeFileSync(join(fx.artifactDir, 'closure.nar'), archive); + const values = { + version: '1', + commit, + system, + out_path: fx.outPath, + archive_sha256: sha256(archive), + ...Object.fromEntries( + requiredBinaries.map((binary) => [ + `binary_${binary}_sha256`, + sha256(`${binary}\n`), + ]), + ), + ...overrides, + }; + delete values.archive; + writeFileSync( + join(fx.artifactDir, 'manifest'), + `${Object.entries(values).map(([key, value]) => `${key}=${value}`).join('\n')}\n`, + ); +}; + +const restore = (fx, expectedCommit = commit, expectedSystem = system) => + spawnSync( + 'bash', + [helper, 'restore', fx.artifactDir, fx.outLink, expectedCommit, expectedSystem], + { encoding: 'utf8', env: fx.env }, + ); + +test('CI architecture is permanently independent of Cachix', () => { + const workflow = readFileSync(repoFile('.github/workflows/ci.yml'), 'utf8'); + const flake = readFileSync(repoFile('flake.nix'), 'utf8'); + assert.doesNotMatch(workflow, /cachix\/cachix-action/i); + assert.doesNotMatch(workflow, /CACHIX_AUTH_TOKEN/); + assert.doesNotMatch(workflow, /rumoca\.cachix\.org/i); + assert.doesNotMatch(flake, /Cachix/i); +}); + +test('workflow transports one rerun-stable MSL Nix closure artifact', () => { + const workflow = readFileSync(repoFile('.github/workflows/ci.yml'), 'utf8'); + const artifactName = 'msl-nix-closure-${{ github.run_id }}-${{ env.RUMOCA_CI_HEAD_SHA }}'; + assert.match(workflow, /GitHub Actions artifact[^\n]*complete Nix closure/i); + + const producer = workflowJob(workflow, 'nix-build-msl'); + assert.match( + producer, + /nix build --print-build-logs \.#msl-artifacts --out-link result-msl-artifacts/, + ); + assert.match(producer, /\.github\/scripts\/msl-nix-closure\.sh pack/); + assert.ok(producer.includes(`name: ${artifactName}`)); + assert.match(producer, /uses: actions\/upload-artifact@v6/); + assert.match(producer, /path: target\/msl-nix-closure/); + assert.match(producer, /if-no-files-found: error/); + assert.match(producer, /overwrite: true/); + assert.match(producer, /retention-days: [1-3]/); + + for (const job of ['msl-shards', 'msl-merge', 'modelicatest-gate']) { + const body = workflowJob(workflow, job); + assert.match(body, /needs: (?:\[?[^\n]*\b)nix-build-msl\b/); + assert.match(body, /uses: actions\/download-artifact@v8/); + assert.ok(body.includes(`name: ${artifactName}`)); + assert.match(body, /\.github\/scripts\/msl-nix-closure\.sh restore/); + assert.doesNotMatch(body, /nix build[^\n]*\.#msl-artifacts/); + const restoreIndex = body.indexOf('.github/scripts/msl-nix-closure.sh restore'); + const useIndex = body.indexOf('result-msl-artifacts/bin/'); + assert.ok(restoreIndex !== -1 && restoreIndex < useIndex, `${job} must restore before use`); + } + + const artifactNames = workflow + .split('\n') + .filter((line) => line.trimStart().startsWith('name: msl-nix-closure-')); + assert.deepEqual( + artifactNames.map((line) => line.trim()), + Array(4).fill(`name: ${artifactName}`), + 'producer overwrite and partial consumer reruns must share one stable artifact name', + ); +}); + +test('every checkout and MSL closure provenance uses the selected CI head', () => { + const workflow = readFileSync(repoFile('.github/workflows/ci.yml'), 'utf8'); + const lines = workflow.split('\n'); + const checkoutLines = lines + .map((line, index) => [line, index]) + .filter(([line]) => line.includes('uses: actions/checkout@v5')); + assert.ok(checkoutLines.length > 0); + for (const [, index] of checkoutLines) { + const block = lines.slice(index, index + 8).join('\n'); + assert.match(block, /ref: \$\{\{ env\.RUMOCA_CI_HEAD_SHA \}\}/); + } + + for (const job of ['nix-build-msl', 'msl-shards', 'msl-merge', 'modelicatest-gate']) { + const body = workflowJob(workflow, job); + assert.doesNotMatch(body, /\$\{\{ github\.sha \}\}|\$GITHUB_SHA/); + assert.match(body, /RUMOCA_CI_HEAD_SHA/); + } +}); + +test('manifest line count is derived from fixed fields and required binaries', () => { + const source = readFileSync(helper, 'utf8'); + assert.doesNotMatch(source, /wc -l[^\n]*-eq\s+9/); + assert.match( + source, + /expected_manifest_lines=\$\(\(\$\{#manifest_fields\[@\]\} \+ \$\{#required_binaries\[@\]\}\)\)/, + ); +}); + +test('helper cleanup is process-scoped instead of relying on RETURN traps', () => { + const source = readFileSync(helper, 'utf8'); + assert.doesNotMatch(source, /trap[^\n]*RETURN/); + assert.match(source, /trap\s+cleanup\s+EXIT/); +}); + +test('pack exports the complete Nix requisites closure and records provenance', () => { + const fx = fixture(); + const result = spawnSync( + 'bash', + [helper, 'pack', fx.outPath, fx.artifactDir, commit, system], + { encoding: 'utf8', env: fx.env }, + ); + assert.equal(result.status, 0, result.stderr); + const nixLog = readFileSync(fx.nixLog, 'utf8'); + assert.match(nixLog, new RegExp(`--query --requisites ${fx.outPath}`)); + assert.match(nixLog, /--export .*dependency.*msl-output/); + const manifest = readFileSync(join(fx.artifactDir, 'manifest'), 'utf8'); + assert.match(manifest, new RegExp(`commit=${commit}`)); + assert.match(manifest, new RegExp(`system=${system}`)); + assert.match(manifest, new RegExp(`out_path=${fx.outPath}`)); + assert.match(manifest, /archive_sha256=[0-9a-f]{64}/); + for (const binary of requiredBinaries) { + assert.match(manifest, new RegExp(`binary_${binary}_sha256=[0-9a-f]{64}`)); + } +}); + +test('restore rejects a missing manifest', () => { + const fx = fixture(); + writeFileSync(join(fx.artifactDir, 'closure.nar'), 'archive'); + const result = restore(fx); + assert.notEqual(result.status, 0); + assert.match(result.stderr, /missing manifest/i); +}); + +test('restore rejects a missing closure archive', () => { + const fx = fixture(); + writeArtifact(fx); + rmSync(join(fx.artifactDir, 'closure.nar')); + const result = restore(fx); + assert.notEqual(result.status, 0); + assert.match(result.stderr, /missing closure archive/i); +}); + +test('restore rejects an archive checksum mismatch before Nix import', () => { + const fx = fixture(); + writeArtifact(fx, { archive_sha256: '0'.repeat(64) }); + const result = restore(fx); + assert.notEqual(result.status, 0); + assert.match(result.stderr, /archive checksum mismatch/i); + assert.equal(existsSync(fx.nixLog), false); +}); + +test('restore rejects a commit mismatch', () => { + const fx = fixture(); + writeArtifact(fx, { commit: 'f'.repeat(40) }); + const result = restore(fx); + assert.notEqual(result.status, 0); + assert.match(result.stderr, /commit mismatch/i); +}); + +test('restore rejects a Nix system mismatch', () => { + const fx = fixture(); + writeArtifact(fx, { system: 'aarch64-linux' }); + const result = restore(fx); + assert.notEqual(result.status, 0); + assert.match(result.stderr, /system mismatch/i); +}); + +test('restore fails closed when Nix import fails', () => { + const fx = fixture({ importFails: true }); + writeArtifact(fx); + const result = restore(fx); + assert.notEqual(result.status, 0); + assert.match(result.stderr, /Nix closure import failed/i); +}); + +test('pack failure removes all temporary closure files', () => { + const fx = fixture({ exportFails: true }); + const result = spawnSync( + 'bash', + [helper, 'pack', fx.outPath, fx.artifactDir, commit, system], + { encoding: 'utf8', env: fx.env }, + ); + assert.notEqual(result.status, 0); + assert.match(result.stderr, /failed to export Nix requisites closure/i); + assert.deepEqual(readdirSync(fx.artifactDir), []); +}); + +test('restore failure removes its mktemp file', () => { + const fx = fixture(); + writeArtifact(fx); + rmSync(join(fx.root, 'nix-store', 'dependency'), { recursive: true }); + const result = restore(fx); + assert.notEqual(result.status, 0); + assert.match(result.stderr, /imported Nix requisite is missing/i); + assert.deepEqual(readdirSync(fx.tempDir), []); +}); + +test('restore rejects a missing required binary', () => { + const fx = fixture(); + writeArtifact(fx); + rmSync(join(fx.outPath, 'bin', 'rumoca-worker')); + const result = restore(fx); + assert.notEqual(result.status, 0); + assert.match(result.stderr, /missing required executable.*rumoca-worker/i); +}); + +test('restore rejects a non-executable required binary', () => { + const fx = fixture(); + writeArtifact(fx); + chmodSync(join(fx.outPath, 'bin', 'rumoca-sim-worker'), 0o644); + const result = restore(fx); + assert.notEqual(result.status, 0); + assert.match(result.stderr, /missing required executable.*rumoca-sim-worker/i); +}); + +test('restore rejects a required binary checksum mismatch', () => { + const fx = fixture(); + writeArtifact(fx, { binary_msl_tests_sha256: 'f'.repeat(64) }); + const result = restore(fx); + assert.notEqual(result.status, 0); + assert.match(result.stderr, /required executable checksum mismatch.*msl_tests/i); +}); + +test('restore imports, verifies, and recreates the exact output link', () => { + const fx = fixture(); + writeArtifact(fx); + const result = restore(fx); + assert.equal(result.status, 0, result.stderr); + assert.equal(readFileSync(fx.nixLog, 'utf8').includes('--import'), true); + assert.equal(realpathSync(fx.outLink), fx.outPath); +}); + +test('fixture lifecycle cleanup removes every registered temporary root', () => { + const first = fixture(); + const second = fixture(); + assert.equal(existsSync(first.root), true); + assert.equal(existsSync(second.root), true); + + cleanupFixtures(); + + assert.equal(existsSync(first.root), false); + assert.equal(existsSync(second.root), false); +}); diff --git a/.github/scripts/recovery-exact-head-contract.mjs b/.github/scripts/recovery-exact-head-contract.mjs new file mode 100644 index 000000000..0c84ec513 --- /dev/null +++ b/.github/scripts/recovery-exact-head-contract.mjs @@ -0,0 +1,124 @@ +import { readFileSync } from 'node:fs'; +import { fileURLToPath } from 'node:url'; + +const batchId = 'climamind-rumoca-broken-main-2026-07'; +const label = 'recovery-exact-head-ci'; +const shaPattern = /^[0-9a-f]{40}$/; +const utcTimestampPattern = /^(\d{4})-(\d{2})-(\d{2})T(\d{2}):(\d{2}):(\d{2})(?:\.\d+)?Z$/; + +const requireValue = (condition, message) => { + if (!condition) throw new Error(message); +}; + +const bodyFields = (body) => { + const wanted = new Set([ + 'recovery_batch_id', + 'recovery_head_sha', + 'recovery_expires_at', + ]); + const fields = new Map(); + for (const line of String(body ?? '').split(/\r?\n/)) { + const match = line.match(/^\s*(recovery_[a-z_]+):\s*(\S+)\s*$/); + if (!match || !wanted.has(match[1])) continue; + requireValue(!fields.has(match[1]), `duplicate ${match[1]}`); + fields.set(match[1], match[2]); + } + for (const field of wanted) requireValue(fields.has(field), `missing ${field}`); + return fields; +}; + +const utcTimestampMillis = (value) => { + const match = String(value).match(utcTimestampPattern); + requireValue(match, 'recovery_expires_at must be a valid RFC 3339 UTC timestamp'); + const [, year, month, day, hour, minute, second] = match.map(Number); + const calendar = new Date(Date.UTC(year, month - 1, day, hour, minute, second)); + requireValue( + calendar.getUTCFullYear() === year + && calendar.getUTCMonth() === month - 1 + && calendar.getUTCDate() === day + && calendar.getUTCHours() === hour + && calendar.getUTCMinutes() === minute + && calendar.getUTCSeconds() === second, + 'recovery_expires_at must be a valid RFC 3339 UTC timestamp', + ); + const millis = Date.parse(value); + requireValue(!Number.isNaN(millis), 'recovery_expires_at must be a valid RFC 3339 UTC timestamp'); + return millis; +}; + +export const validateRecoveryContract = ({ + event, + eventName, + githubSha, + repository, + selectedSha, + now = new Date(), +}) => { + requireValue(shaPattern.test(selectedSha), 'selected SHA must be a lowercase 40-hex commit'); + requireValue(shaPattern.test(githubSha), 'GITHUB_SHA must be a lowercase 40-hex commit'); + if (eventName !== 'pull_request') { + requireValue(selectedSha === githubSha, 'normal CI must select GITHUB_SHA'); + return { mode: 'normal', sha: selectedSha }; + } + const pr = event?.pull_request; + requireValue(pr, 'pull_request payload data is required'); + const labels = Array.isArray(pr.labels) ? pr.labels : []; + const recoveryRequested = String(pr.head?.ref ?? '').startsWith('integration/recovery-') + || labels.some((item) => item?.name === label) + || /^\s*recovery_[a-z_]+:/m.test(String(pr.body ?? '')); + if (!recoveryRequested) { + requireValue(selectedSha === githubSha, 'normal CI must select GITHUB_SHA'); + return { mode: 'normal', sha: selectedSha }; + } + + requireValue(event?.repository?.full_name === repository, 'event repository mismatch'); + requireValue( + pr.head?.repo?.full_name === repository && pr.base?.repo?.full_name === repository, + 'recovery head and base must use the same repository', + ); + requireValue(pr.draft === true, 'recovery pull request must remain Draft'); + requireValue( + String(pr.head?.ref ?? '').startsWith('integration/recovery-'), + 'recovery branch prefix must be integration/recovery-', + ); + requireValue(event?.action === 'labeled', 'recovery must be authorized by a labeled event'); + requireValue(event?.label?.name === label, 'recovery labeled event must apply recovery-exact-head-ci'); + requireValue( + labels.some((item) => item?.name === label), + `recovery pull request requires the ${label} label`, + ); + const fields = bodyFields(pr.body); + requireValue(fields.get('recovery_batch_id') === batchId, 'recovery batch id mismatch'); + requireValue( + fields.get('recovery_head_sha') === pr.head?.sha, + 'body-recorded head must equal payload head SHA', + ); + const expiresAt = fields.get('recovery_expires_at'); + const expiryMillis = utcTimestampMillis(expiresAt); + requireValue(expiryMillis > now.getTime(), 'recovery authorization is expired'); + requireValue(pr.head?.sha === selectedSha, 'selected SHA must equal payload head SHA'); + return { mode: 'recovery', sha: selectedSha }; +}; + +const main = () => { + const eventPath = process.env.GITHUB_EVENT_PATH; + requireValue(eventPath, 'GITHUB_EVENT_PATH is required'); + const event = JSON.parse(readFileSync(eventPath, 'utf8')); + const result = validateRecoveryContract({ + event, + eventName: process.env.GITHUB_EVENT_NAME, + githubSha: process.env.GITHUB_SHA, + repository: process.env.GITHUB_REPOSITORY, + selectedSha: process.env.RUMOCA_CI_HEAD_SHA, + }); + console.log(`${result.mode} CI head verified: ${result.sha}`); +}; + +if (process.argv[1] === fileURLToPath(import.meta.url)) { + try { + main(); + } catch (error) { + console.error(`recovery-exact-head-contract: ${error.message}`); + process.exitCode = 1; + } +} diff --git a/.github/scripts/recovery-exact-head-contract.test.mjs b/.github/scripts/recovery-exact-head-contract.test.mjs new file mode 100644 index 000000000..79eb5f947 --- /dev/null +++ b/.github/scripts/recovery-exact-head-contract.test.mjs @@ -0,0 +1,192 @@ +import assert from 'node:assert/strict'; +import { mkdtempSync, readFileSync, rmSync, writeFileSync } from 'node:fs'; +import { tmpdir } from 'node:os'; +import { join } from 'node:path'; +import { fileURLToPath } from 'node:url'; +import { spawnSync } from 'node:child_process'; +import { afterEach, test } from 'node:test'; + +import { validateRecoveryContract } from './recovery-exact-head-contract.mjs'; + +const script = fileURLToPath(new URL('./recovery-exact-head-contract.mjs', import.meta.url)); +const headSha = '1'.repeat(40); +const mergeSha = '2'.repeat(40); +const repo = 'climamind/rumoca'; +const roots = []; + +afterEach(() => { + for (const root of roots.splice(0)) rmSync(root, { recursive: true, force: true }); +}); + +const recoveryBody = ({ + batch = 'climamind-rumoca-broken-main-2026-07', + head = headSha, + expiresAt = '2099-07-29T23:59:59Z', +} = {}) => [ + `recovery_batch_id: ${batch}`, + `recovery_head_sha: ${head}`, + `recovery_expires_at: ${expiresAt}`, +].join('\n'); + +const pullRequestEvent = (overrides = {}) => ({ + action: 'labeled', + label: { name: 'recovery-exact-head-ci' }, + pull_request: { + author_association: 'OWNER', + base: { repo: { full_name: repo } }, + body: recoveryBody(), + draft: true, + head: { + ref: 'integration/recovery-clean-v4-20260722', + repo: { full_name: repo }, + sha: headSha, + }, + labels: [{ name: 'recovery-exact-head-ci' }], + ...overrides, + }, + repository: { full_name: repo }, +}); + +const validateRecovery = (event = pullRequestEvent(), selectedSha = headSha) => + validateRecoveryContract({ + event, + eventName: 'pull_request', + githubSha: mergeSha, + repository: repo, + selectedSha, + now: new Date('2026-07-22T12:00:00Z'), + }); + +test('accepts the exact authorized same-repository recovery head', () => { + const event = pullRequestEvent({ author_association: 'NONE' }); + assert.deepEqual(validateRecovery(event), { mode: 'recovery', sha: headSha }); + const fractional = pullRequestEvent({ + body: recoveryBody({ expiresAt: '2099-07-29T23:59:59.123Z' }), + }); + assert.deepEqual(validateRecovery(fractional), { mode: 'recovery', sha: headSha }); +}); + +test('accepts normal CI only on GitHub synthetic merge SHA', () => { + const event = pullRequestEvent({ + body: 'ordinary pull request', + draft: false, + head: { + ref: 'fix/ordinary-change', + repo: { full_name: repo }, + sha: headSha, + }, + labels: [], + }); + assert.deepEqual( + validateRecoveryContract({ + event, + eventName: 'pull_request', + githubSha: mergeSha, + repository: repo, + selectedSha: mergeSha, + now: new Date('2026-07-22T12:00:00Z'), + }), + { mode: 'normal', sha: mergeSha }, + ); + assert.throws( + () => validateRecoveryContract({ + event: {}, eventName: 'push', githubSha: mergeSha, repository: repo, + selectedSha: headSha, now: new Date('2026-07-22T12:00:00Z'), + }), + /normal CI must select GITHUB_SHA/, + ); +}); + +for (const [name, mutate, message] of [ + ['fork head', (pr) => { pr.head.repo.full_name = 'outsider/rumoca'; }, /same repository/], + ['foreign base', (pr) => { pr.base.repo.full_name = 'other/rumoca'; }, /same repository/], + ['non-Draft PR', (pr) => { pr.draft = false; }, /Draft/], + ['absent label', (pr) => { pr.labels = []; }, /label/], + ['wrong branch prefix', (pr) => { pr.head.ref = 'feature/recovery'; }, /branch prefix/], + ['absent batch id', (pr) => { pr.body = 'recovery_head_sha: 111\n'; }, /recovery_batch_id/], + ['wrong batch id', (pr) => { pr.body = recoveryBody({ batch: 'other' }); }, /batch/], + ['absent recorded head', (pr) => { pr.body = 'recovery_batch_id: climamind-rumoca-broken-main-2026-07\n'; }, /recovery_head_sha/], + ['wrong recorded head', (pr) => { pr.body = recoveryBody({ head: '3'.repeat(40) }); }, /recorded head/], + ['absent expiry', (pr) => { pr.body = recoveryBody().split('\n').slice(0, 2).join('\n'); }, /recovery_expires_at/], + ['invalid expiry', (pr) => { pr.body = recoveryBody({ expiresAt: 'next-week' }); }, /RFC 3339 UTC/], + ['invalid calendar date', (pr) => { pr.body = recoveryBody({ expiresAt: '2027-02-29T00:00:00Z' }); }, /RFC 3339 UTC/], + ['expired authorization', (pr) => { pr.body = recoveryBody({ expiresAt: '2026-07-21T23:59:59Z' }); }, /expired/], + ['duplicate body field', (pr) => { pr.body += `\nrecovery_head_sha: ${headSha}`; }, /duplicate recovery_head_sha/], +]) { + test(`rejects recovery with ${name}`, () => { + const event = structuredClone(pullRequestEvent()); + mutate(event.pull_request); + assert.throws(() => validateRecovery(event), message); + }); +} + +for (const action of ['opened', 'edited', 'synchronize']) { + test('rejects a persistent recovery label on ' + action, () => { + const event = structuredClone(pullRequestEvent()); + event.action = action; + assert.throws(() => validateRecovery(event), /labeled event/); + }); +} + +test('rejects a different label action even when the PR retains the recovery label', () => { + const event = pullRequestEvent({ author_association: 'NONE' }); + event.label.name = 'other-label'; + assert.throws(() => validateRecovery(event), /must apply recovery-exact-head-ci/); +}); + +test('rejects recovery when selected SHA differs from payload head', () => { + assert.throws(() => validateRecovery(pullRequestEvent(), '4'.repeat(40)), /selected SHA/); +}); + +test('does not reinterpret a malformed recovery request as normal merge-ref CI', () => { + const event = pullRequestEvent({ labels: [] }); + assert.throws(() => validateRecovery(event, mergeSha), /label/); +}); + +test('CLI independently reads and validates GITHUB_EVENT_PATH', () => { + const root = mkdtempSync(join(tmpdir(), 'rumoca-recovery-contract-')); + roots.push(root); + const eventPath = join(root, 'event.json'); + writeFileSync(eventPath, JSON.stringify(pullRequestEvent())); + const result = spawnSync(process.execPath, [script], { + encoding: 'utf8', + env: { + ...process.env, + GITHUB_EVENT_NAME: 'pull_request', + GITHUB_EVENT_PATH: eventPath, + GITHUB_REPOSITORY: repo, + GITHUB_SHA: mergeSha, + RUMOCA_CI_HEAD_SHA: headSha, + }, + }); + assert.equal(result.status, 0, result.stderr); + assert.match(result.stdout, new RegExp(`recovery.*${headSha}`)); +}); + +test('workflow selects recovery head narrowly and proves checked-out provenance', () => { + const workflow = readFileSync(new URL('../workflows/ci.yml', import.meta.url), 'utf8'); + for (const type of [ + 'opened', 'synchronize', 'reopened', 'ready_for_review', + 'converted_to_draft', 'edited', 'labeled', 'unlabeled', + ]) assert.match(workflow, new RegExp(`\\b${type}\\b`)); + const selection = workflow.slice( + workflow.indexOf(' RUMOCA_CI_HEAD_SHA:'), + workflow.indexOf('\n\n# Prevent duplicate runs'), + ); + assert.match(selection, /github\.event_name == 'pull_request'/); + assert.match(selection, /github\.event\.action == 'labeled'/); + assert.match(selection, /github\.event\.label\.name == 'recovery-exact-head-ci'/); + assert.match(selection, /pull_request\.draft == true/); + assert.match(selection, /integration\/recovery-/); + assert.match(selection, /head\.repo\.full_name == github\.repository/); + assert.match(selection, /base\.repo\.full_name == github\.repository/); + assert.match(selection, /recovery-exact-head-ci/); + assert.doesNotMatch(selection, /author_association/); + assert.match(selection, /climamind-rumoca-broken-main-2026-07/); + assert.match(selection, /recovery_head_sha: \{0\}.*head\.sha/); + assert.match(selection, /recovery_expires_at:/); + assert.match(selection, /head\.sha[\s\S]*\|\| github\.sha/); + assert.match(workflow, /recovery-exact-head-contract\.mjs/); + assert.match(workflow, /node --test \.github\/scripts\/recovery-exact-head-contract\.test\.mjs/); + assert.match(workflow, /git rev-parse HEAD/); +}); diff --git a/.github/scripts/toolchain-pins.test.mjs b/.github/scripts/toolchain-pins.test.mjs new file mode 100644 index 000000000..291be36a9 --- /dev/null +++ b/.github/scripts/toolchain-pins.test.mjs @@ -0,0 +1,233 @@ +import assert from 'node:assert/strict'; +import { createHash } from 'node:crypto'; +import { + chmodSync, + mkdtempSync, + mkdirSync, + readFileSync, + rmSync, + writeFileSync, +} from 'node:fs'; +import { spawnSync } from 'node:child_process'; +import { tmpdir } from 'node:os'; +import { join } from 'node:path'; +import { fileURLToPath } from 'node:url'; +import { test } from 'node:test'; + +const root = fileURLToPath(new URL('../..', import.meta.url)); +const read = (path) => readFileSync(new URL(`../../${path}`, import.meta.url), 'utf8'); +const pin = '1.27.0~1-gd7e2907-1'; +const pool = 'https://build.openmodelica.org/apt/pool/contrib-noble'; +const packages = [ + ['omc', 'amd64'], + ['omc-common', 'all'], + ['libomc', 'amd64'], + ['libomcsimulation', 'amd64'], +]; + +const fixture = ({ metadata = {}, rows = packages } = {}) => { + const directory = mkdtempSync(join(tmpdir(), 'rumoca-omc-fixture-')); + const packagesDirectory = join(directory, 'packages'); + const binDirectory = join(directory, 'bin'); + const manifest = join(directory, 'manifest'); + mkdirSync(packagesDirectory); + mkdirSync(binDirectory); + + const manifestRows = packages.map(([packageName, architecture]) => { + const filename = `${packageName}_${pin}_${architecture}.deb`; + const fields = { + Package: packageName, + Version: pin, + Architecture: architecture, + ...metadata[packageName], + }; + const contents = Object.entries(fields) + .map(([name, value]) => `${name}: ${value}`) + .join('\n'); + writeFileSync(join(packagesDirectory, filename), `${contents}\npayload\n`); + const sha256 = createHash('sha256') + .update(readFileSync(join(packagesDirectory, filename))) + .digest('hex'); + return `${packageName} ${pin} ${architecture} ${filename} ${sha256} ${pool}/${filename}`; + }); + writeFileSync(manifest, `${rows.map((row) => ( + Array.isArray(row) ? manifestRows[packages.indexOf(row)] : row + )).join('\n')}\n`); + const dpkgDeb = join(binDirectory, 'dpkg-deb'); + writeFileSync(dpkgDeb, `#!/usr/bin/env bash +set -euo pipefail +[[ "$1" == "-f" ]] +sed -n "s/^$3: //p" "$2" +`); + chmodSync(dpkgDeb, 0o755); + + return { + directory, + manifest, + packagesDirectory, + run: () => spawnSync( + join(root, 'scripts/ci/verify-openmodelica-packages.sh'), + [manifest, pin, packagesDirectory], + { + encoding: 'utf8', + env: { ...process.env, PATH: `${binDirectory}:${process.env.PATH}` }, + }, + ), + }; +}; + +const withFixture = (options, check) => { + const value = fixture(options); + try { + check(value); + } finally { + rmSync(value.directory, { recursive: true, force: true }); + } +}; + +const workflowJob = (workflow, job) => { + const start = workflow.indexOf(` ${job}:`); + assert.notEqual(start, -1, `${job} must exist`); + const next = workflow.slice(start + 1).search(/^ [a-zA-Z0-9_-]+:$/m); + return workflow.slice(start, next === -1 ? undefined : start + 1 + next); +}; + +test('OpenModelica uses one exact repository pin', () => { + const version = read('toolchains/openmodelica-version').trim(); + assert.match(version, /^\d+\.\d+\.\d+~1-g[0-9a-f]+-\d+$/); + assert.equal(version, pin); +}); + +test('OpenModelica manifest pins the four official package artifacts', () => { + assert.equal( + read('toolchains/openmodelica-packages.txt'), + `# source_commit d7e2907f419d8061ce0a461ac9709dcb84fa70a1 +omc ${pin} amd64 omc_${pin}_amd64.deb a4511acb19f7377275347f9fc27af92c307db0903e7f46e44085410a060d86ab ${pool}/omc_${pin}_amd64.deb +omc-common ${pin} all omc-common_${pin}_all.deb 9781f4efa44274ecff4538227ef8054b10d3b9abe6cf3d36bb9b23237c0b1119 ${pool}/omc-common_${pin}_all.deb +libomc ${pin} amd64 libomc_${pin}_amd64.deb 90b650ecf9da174cd477115a8f61d54f8bf22b6c87ba1707cb5482d933ae3f5f ${pool}/libomc_${pin}_amd64.deb +libomcsimulation ${pin} amd64 libomcsimulation_${pin}_amd64.deb bdc2c5c307aacf516017f8b64fc4878e6a75d7b0587366d2b93ec8ae3590632d ${pool}/libomcsimulation_${pin}_amd64.deb +`, + ); +}); + +test('offline package verifier accepts the exact artifact set', () => { + withFixture({}, ({ run }) => { + const result = run(); + assert.equal(result.status, 0, result.stderr); + }); +}); + +test('offline package verifier rejects a bad checksum', () => { + withFixture({}, ({ manifest, run }) => { + const contents = readFileSync(manifest, 'utf8'); + writeFileSync(manifest, contents.replace(/[0-9a-f]{64}/, '0'.repeat(64))); + const result = run(); + assert.notEqual(result.status, 0); + assert.match(result.stderr, /checksum mismatch/i); + }); +}); + +test('offline package verifier rejects incorrect Debian metadata', () => { + for (const [field, value] of [ + ['Package', 'wrong-package'], + ['Version', '1.27.0~2-g4470062-1'], + ['Architecture', 'arm64'], + ]) { + withFixture({ metadata: { omc: { [field]: value } } }, ({ run }) => { + const result = run(); + assert.notEqual(result.status, 0, `${field}=${value}`); + assert.match(result.stderr, new RegExp(`${field} mismatch`, 'i')); + }); + } +}); + +test('offline package verifier rejects missing and duplicate manifest entries', () => { + for (const rows of [packages.slice(0, -1), [...packages, packages[0]]]) { + withFixture({ rows }, ({ run }) => { + const result = run(); + assert.notEqual(result.status, 0, result.stderr); + assert.match(result.stderr, /missing|duplicate/i); + }); + } +}); + +test('local package tooling defaults to the same Node major as CI', () => { + assert.equal(read('.nvmrc').trim(), '20'); +}); + +test('all MSL jobs use the shared OpenModelica installer', () => { + const workflow = read('.github/workflows/ci.yml'); + const installerCalls = workflow.match(/scripts\/ci\/install-openmodelica\.sh/g) ?? []; + + assert.equal(installerCalls.length, 2); + assert.doesNotMatch(workflow, /omc_channel=/); + assert.doesNotMatch(workflow, /apt-get install[^\n]*\bom[c]?\b/); +}); + +test('the installer uses content-verified pool artifacts without a rolling index', () => { + const installer = read('scripts/ci/install-openmodelica.sh'); + + assert.match(installer, /openmodelica-packages\.txt/); + assert.match(installer, /verify-openmodelica-packages\.sh/); + assert.match(installer, /mktemp -d/); + assert.match(installer, /trap .*EXIT/); + assert.match(installer, /curl .*--proto ['"]?=https/); + assert.doesNotMatch(installer, /apt-cache madison/); + assert.doesNotMatch(installer, /openmodelica\.list|openmodelica-keyring/); + assert.doesNotMatch(installer, /build\.openmodelica\.org\/apt .*stable/); +}); + +test('every MSL consumer installs OpenModelica before verifying omc', () => { + const workflow = read('.github/workflows/ci.yml'); + for (const job of ['msl-shards', 'modelicatest-gate']) { + const jobBody = workflowJob(workflow, job); + const install = jobBody.indexOf('scripts/ci/install-openmodelica.sh'); + + assert.notEqual(install, -1, `${job} must use the shared installer`); + for (const verification of jobBody.matchAll(/omc --version/g)) { + assert.ok( + install < verification.index, + `${job} must install OpenModelica before verification`, + ); + } + } +}); + +test('the installer validates the installed version against the pin', () => { + const installer = read('scripts/ci/install-openmodelica.sh'); + + assert.match(installer, /toolchains\/openmodelica-version/); + assert.match(installer, /omc --version/); + assert.match(installer, /installed_version/); + assert.match(installer, /expected_version/); +}); + +test('the installer accepts the release output implied by the exact package pin', () => { + assert.equal(read('toolchains/openmodelica-version').trim(), pin); + const result = spawnSync( + 'scripts/ci/install-openmodelica.sh', + ['--check-output', 'OpenModelica 1.27.0~1-gd7e2907'], + { encoding: 'utf8' }, + ); + + assert.equal(result.status, 0, result.stderr); + assert.equal(result.stdout, 'OpenModelica 1.27.0~1-gd7e2907\n'); +}); + +test('the installer rejects runtime identities other than the exact pin-derived release', () => { + for (const output of [ + 'OpenModelica 1.27.0', + 'OpenModelica 1.27.0~1-g0000000', + 'OpenModelica 1.27.0~2-g4470062', + 'OpenModelica 1.27.0~dev.beta.1-4-ge5d8071', + ]) { + const result = spawnSync( + 'scripts/ci/install-openmodelica.sh', + ['--check-output', output], + { encoding: 'utf8' }, + ); + + assert.notEqual(result.status, 0, output); + assert.match(result.stderr, /OpenModelica version mismatch/, output); + } +}); diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 1eb1969b1..849e61bd4 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -6,11 +6,26 @@ on: tags: ['v*'] pull_request: branches: ['main'] + types: [opened, synchronize, reopened, ready_for_review, converted_to_draft, edited, labeled, unlabeled] env: CARGO_TERM_COLOR: always CARGO_INCREMENTAL: 0 RUST_BACKTRACE: 1 + RUMOCA_CI_HEAD_SHA: >- + ${{ github.event_name == 'pull_request' + && github.event.action == 'labeled' + && github.event.label.name == 'recovery-exact-head-ci' + && github.event.pull_request.draft == true + && startsWith(github.event.pull_request.head.ref, 'integration/recovery-') + && github.event.pull_request.head.repo.full_name == github.repository + && github.event.pull_request.base.repo.full_name == github.repository + && contains(github.event.pull_request.labels.*.name, 'recovery-exact-head-ci') + && contains(github.event.pull_request.body, 'recovery_batch_id: climamind-rumoca-broken-main-2026-07') + && contains(github.event.pull_request.body, format('recovery_head_sha: {0}', github.event.pull_request.head.sha)) + && contains(github.event.pull_request.body, 'recovery_expires_at: ') + && github.event.pull_request.head.sha + || github.sha }} # Prevent duplicate runs concurrency: @@ -26,6 +41,30 @@ permissions: pull-requests: write jobs: + ci-head-provenance: + name: CI head provenance + runs-on: ubuntu-24.04 + steps: + - uses: actions/checkout@v5 + with: + ref: ${{ env.RUMOCA_CI_HEAD_SHA }} + token: ${{ github.token }} + persist-credentials: false + + - name: Validate recovery selection contract + run: node .github/scripts/recovery-exact-head-contract.mjs + + - name: Verify checked-out commit + shell: bash + run: | + set -euo pipefail + actual=$(git rev-parse HEAD) + test "$actual" = "$RUMOCA_CI_HEAD_SHA" || { + echo "::error::checked out $actual, expected $RUMOCA_CI_HEAD_SHA" + exit 1 + } + echo "Verified CI head: $actual" + detect-release-tag: name: Detect release tag on commit if: github.event_name == 'push' && (github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags/v')) @@ -36,6 +75,7 @@ jobs: steps: - uses: actions/checkout@v5 with: + ref: ${{ env.RUMOCA_CI_HEAD_SHA }} token: ${{ github.token }} persist-credentials: true fetch-depth: 0 @@ -72,6 +112,7 @@ jobs: steps: - uses: actions/checkout@v5 with: + ref: ${{ env.RUMOCA_CI_HEAD_SHA }} token: ${{ github.token }} persist-credentials: true @@ -85,6 +126,12 @@ jobs: - name: Check formatting run: cargo fmt --all -- --check + - name: Verify repository toolchain pins + run: | + node --test .github/scripts/recovery-exact-head-contract.test.mjs + node --test .github/scripts/toolchain-pins.test.mjs + node --test .github/scripts/msl-nix-closure.test.mjs + # Architectural invariant, checked without compiling anything: `xtask` must # carry NO rumoca-* workspace crate. It parses args, moves files, and shells # out; heavy compiler-linked work runs on demand via `cargo run -p ` @@ -107,14 +154,13 @@ jobs: steps: - uses: actions/checkout@v5 with: + ref: ${{ env.RUMOCA_CI_HEAD_SHA }} token: ${{ github.token }} persist-credentials: true fetch-depth: 0 - name: Install system dependencies - run: | - sudo apt-get update - sudo apt-get install -y libudev-dev + run: scripts/ci/apt-install.sh libudev-dev - name: Install Rust run: | @@ -163,18 +209,16 @@ jobs: path: target/crate-dag/workspace-crate-dag.dot if-no-files-found: error - # Reproducible build + clippy + fmt via the root crane+fenix flake, backed by - # the Cachix binary cache (rumoca.cachix.org). First run is cold and pushes - # the built closure; subsequent runs (and local `cachix use rumoca`) restore - # the prebuilt dependency closure + rumoca binary in seconds. Replaces the - # deprecated magic-nix-cache, which rate-limited the GitHub Actions cache API - # on a closure this size (roadmap M4 slice 2). + # Reproducible build via the root crane+fenix flake. This independent check + # uses the normal Nix substituters or a cold build; only the heavy MSL bundle + # has a workflow-owned build-once/share-many transport below. nix-checks: name: Nix flake checks (Linux) runs-on: ubuntu-24.04 steps: - uses: actions/checkout@v5 with: + ref: ${{ env.RUMOCA_CI_HEAD_SHA }} token: ${{ github.token }} persist-credentials: true @@ -183,19 +227,12 @@ jobs: with: determinate: false - - name: Cachix (rumoca binary cache) - uses: cachix/cachix-action@v17 - with: - name: rumoca - # Public cache: pulls need no auth; the token lets CI push new paths. - authToken: ${{ secrets.CACHIX_AUTH_TOKEN }} - - name: Nix build the flake package check # Build ONLY the reproducible package check. Format and clippy are owned # by the standalone Format/Lint jobs, and the heavy release MSL artifact # bundle is built once by the dedicated `nix-build-msl` producer and - # consumed via Cachix; building any of that here too would duplicate - # release+LTO work. + # consumed via a GitHub Actions artifact; building any of that here too + # would duplicate release+LTO work. # # We deliberately do NOT run `nix flake check`: it eval-builds every output # (re-building those packages), and its `--no-build` eval trips over the @@ -218,6 +255,7 @@ jobs: steps: - uses: actions/checkout@v5 with: + ref: ${{ env.RUMOCA_CI_HEAD_SHA }} token: ${{ github.token }} persist-credentials: true @@ -424,6 +462,7 @@ jobs: steps: - uses: actions/checkout@v5 with: + ref: ${{ env.RUMOCA_CI_HEAD_SHA }} token: ${{ github.token }} persist-credentials: true @@ -432,11 +471,6 @@ jobs: with: determinate: false - - name: Cachix (rumoca binary cache) - uses: cachix/cachix-action@v17 - with: - name: rumoca - - name: Rust cache (template-runtime) uses: Swatinem/rust-cache@v2 with: @@ -551,6 +585,7 @@ jobs: steps: - uses: actions/checkout@v5 with: + ref: ${{ env.RUMOCA_CI_HEAD_SHA }} token: ${{ github.token }} persist-credentials: true @@ -580,9 +615,7 @@ jobs: cache-dependency-path: '**/package-lock.json' - name: Install Xvfb and system dependencies - run: | - sudo apt-get update - sudo apt-get install -y xvfb xauth libudev-dev + run: scripts/ci/apt-install.sh xvfb xauth libudev-dev - name: Select headless browser shell: bash @@ -689,6 +722,7 @@ jobs: steps: - uses: actions/checkout@v5 with: + ref: ${{ env.RUMOCA_CI_HEAD_SHA }} token: ${{ github.token }} persist-credentials: true @@ -769,18 +803,17 @@ jobs: run: | ./target/debug/xtask coverage gate --allowed-workspace-line-coverage-drop 3.0 - # Build the heavy MSL artifacts (release + LTO) ONCE via Nix/crane and push to - # Cachix, so the shard / merge / ModelicaTest consumers restore them instead - # of each recompiling + re-LTO'ing the workspace. LTO is a link-time cost no - # per-crate cache can avoid — only build-once does. On an unchanged dep closure - # the crane `cargoArtifacts` layer is a Cachix hit; only the workspace crates - # rebuild here, once, for all consumers. + # Build the heavy MSL artifacts (release + LTO) ONCE via Nix/crane. Export a + # GitHub Actions artifact containing the complete Nix closure so shard, merge, + # and ModelicaTest consumers restore exactly this commit/system output without + # credentials or recompilation. LTO is a link-time cost only build-once avoids. nix-build-msl: name: Build MSL artifacts (Nix, once) runs-on: ubuntu-24.04 steps: - uses: actions/checkout@v5 with: + ref: ${{ env.RUMOCA_CI_HEAD_SHA }} token: ${{ github.token }} persist-credentials: true @@ -789,18 +822,23 @@ jobs: with: determinate: false - - name: Cachix (rumoca binary cache) - uses: cachix/cachix-action@v17 - with: - name: rumoca - # Public cache: consumer pulls need no auth; the token lets this - # producer push the freshly built closure. Absent on fork PRs, in which - # case consumers Cachix-miss and fall back to compiling — still correct. - authToken: ${{ secrets.CACHIX_AUTH_TOKEN }} - - - name: Build + cache MSL artifacts + - name: Build MSL artifacts once + shell: bash run: | - nix build --print-build-logs .#msl-artifacts + set -euo pipefail + nix build --print-build-logs .#msl-artifacts --out-link result-msl-artifacts + system=$(nix eval --impure --raw --expr builtins.currentSystem) + .github/scripts/msl-nix-closure.sh pack \ + result-msl-artifacts target/msl-nix-closure "$RUMOCA_CI_HEAD_SHA" "$system" + + - name: Upload complete MSL Nix closure + uses: actions/upload-artifact@v6 + with: + name: msl-nix-closure-${{ github.run_id }}-${{ env.RUMOCA_CI_HEAD_SHA }} + path: target/msl-nix-closure + if-no-files-found: error + overwrite: true + retention-days: 2 # ============================================================================ # MSL parity gate, sharded. The ~54min root-example parity run (575 models + @@ -808,10 +846,10 @@ jobs: # (so the timeout tail spreads evenly); the msl-merge job then runs the # baseline ratchet ONCE on the merged results. ModelicaTest is its own job. # - # Each consumer restores the prebuilt MSL artifact bundle from Cachix (needs: - # nix-build-msl) and runs it via `--prebuilt-test-binary` / - # `--prebuilt-sim-worker`, so no workspace compile/LTO happens here — only the - # light `cargo xtask` wrapper builds. + # Each consumer downloads and validates the producer's complete Nix closure + # GitHub Actions artifact (needs: nix-build-msl) and runs it via + # `--prebuilt-test-binary` / `--prebuilt-sim-worker`, so no workspace + # compile/LTO happens here — only the light `cargo xtask` wrapper builds. # ============================================================================ msl-shards: name: MSL Quality Gate (shard ${{ matrix.shard }}/4) @@ -825,18 +863,22 @@ jobs: steps: - uses: actions/checkout@v5 with: + ref: ${{ env.RUMOCA_CI_HEAD_SHA }} token: ${{ github.token }} persist-credentials: true - name: Install system dependencies - run: | - sudo apt-get update - sudo apt-get install -y libudev-dev + run: scripts/ci/apt-install.sh libudev-dev - - name: Install Rust + - name: Install OpenModelica + shell: bash + run: scripts/ci/install-openmodelica.sh + + - name: Verify container toolchain run: | - curl --proto '=https' --tlsv1.2 -sSf https://sh.rustup.rs | sh -s -- -y --default-toolchain "$(awk -F'"' '/^channel/ { print $2 }' rust-toolchain.toml)" - echo "$HOME/.cargo/bin" >> $GITHUB_PATH + rustc --version + cargo --version + omc --version - name: Rust cache (msl-gate) uses: Swatinem/rust-cache@v2 @@ -853,6 +895,89 @@ jobs: path: target/msl/ModelicaStandardLibrary-4.1.0 key: msl-v4.1.0-release-zip-layout-v2 + - name: Cache OMC parity checkpoint + id: omc-parity-checkpoint-cache + uses: actions/cache@v5 + continue-on-error: true + with: + path: | + target/msl/results/omc_simulation_reference.json + target/msl/results/omc_sim_work + target/msl/results/sim_traces/omc + target/msl/results/omc_parity_cache + key: msl-omc-parity-${{ runner.os }}-${{ hashFiles('Cargo.lock', '.github/workflows/ci.yml', 'crates/rumoca-test-msl/tests/msl_tests/msl_quality_baseline.json') }}-${{ github.run_id }}-${{ github.run_attempt }} + restore-keys: | + msl-omc-parity-${{ runner.os }}-${{ hashFiles('Cargo.lock', '.github/workflows/ci.yml', 'crates/rumoca-test-msl/tests/msl_tests/msl_quality_baseline.json') }}- + + - name: Download prior OMC parity checkpoint artifact + if: steps.omc-parity-checkpoint-cache.outputs.cache-hit != 'true' + id: omc-parity-checkpoint-artifact + uses: actions/github-script@v8 + with: + script: | + const fs = require('fs'); + const path = require('path'); + const { owner, repo } = context.repo; + const branch = context.ref.replace('refs/heads/', ''); + const runs = await github.paginate(github.rest.actions.listWorkflowRuns, { + owner, + repo, + workflow_id: 'ci.yml', + branch, + status: 'completed', + per_page: 20, + }); + for (const run of runs) { + const artifacts = await github.paginate(github.rest.actions.listWorkflowRunArtifacts, { + owner, + repo, + run_id: run.id, + per_page: 100, + }); + const artifact = artifacts.find((item) => + item.name === 'msl-compatibility-report' && !item.expired + ); + if (!artifact) { + continue; + } + const response = await github.rest.actions.downloadArtifact({ + owner, + repo, + artifact_id: artifact.id, + archive_format: 'zip', + }); + const out = path.join(process.env.RUNNER_TEMP, 'omc-parity-checkpoint.zip'); + fs.writeFileSync(out, Buffer.from(response.data)); + core.setOutput('found', 'true'); + core.setOutput('path', out); + core.info(`Downloaded OMC parity checkpoint artifact ${artifact.id} from run ${run.id}.`); + return; + } + core.setOutput('found', 'false'); + core.info('No prior OMC parity checkpoint artifact found for this commit.'); + + - name: Restore prior OMC parity checkpoint artifact + if: steps.omc-parity-checkpoint-artifact.outputs.found == 'true' + shell: bash + run: | + set -euo pipefail + restore_dir="$RUNNER_TEMP/omc-parity-checkpoint" + rm -rf "$restore_dir" + mkdir -p "$restore_dir" target/msl/results + unzip -q "${{ steps.omc-parity-checkpoint-artifact.outputs.path }}" -d "$restore_dir" + for path in \ + omc_simulation_reference.json \ + omc_sim_work \ + sim_traces/omc \ + omc_parity_cache + do + if [[ -e "$restore_dir/results/$path" ]]; then + mkdir -p "target/msl/results/$(dirname "$path")" + rm -rf "target/msl/results/$path" + cp -a "$restore_dir/results/$path" "target/msl/results/$path" + fi + done + - name: Ensure MSL shell: bash run: | @@ -871,23 +996,6 @@ jobs: test -f "$msl_dir/Complex.mo" test -f "$msl_dir/Modelica 4.1.0/package.mo" - - name: Install OpenModelica - shell: bash - run: | - set -euo pipefail - omc_channel="stable" - sudo apt-get update - sudo apt-get install -y --no-install-recommends \ - build-essential ca-certificates clang cmake curl gnupg \ - libexpat1-dev liblapack-dev lsb-release unzip zip - curl -fsSL https://build.openmodelica.org/apt/openmodelica.asc \ - | sudo gpg --dearmor -o /usr/share/keyrings/openmodelica-keyring.gpg - echo "deb [arch=amd64 signed-by=/usr/share/keyrings/openmodelica-keyring.gpg] https://build.openmodelica.org/apt $(lsb_release -cs) $omc_channel" \ - | sudo tee /etc/apt/sources.list.d/openmodelica.list - sudo apt-get update - sudo apt-get install -y --no-install-recommends omc - omc --version - - name: Clean stale MSL results shell: bash run: rm -rf target/msl/results @@ -897,14 +1005,19 @@ jobs: with: determinate: false - - name: Cachix (rumoca binary cache) - uses: cachix/cachix-action@v17 + - name: Download complete MSL Nix closure + uses: actions/download-artifact@v8 with: - name: rumoca + name: msl-nix-closure-${{ github.run_id }}-${{ env.RUMOCA_CI_HEAD_SHA }} + path: target/msl-nix-closure - - name: Restore prebuilt MSL binaries (Cachix) + - name: Restore and verify prebuilt MSL closure + shell: bash run: | - nix build .#msl-artifacts --out-link result-msl-artifacts + set -euo pipefail + system=$(nix eval --impure --raw --expr builtins.currentSystem) + .github/scripts/msl-nix-closure.sh restore \ + target/msl-nix-closure result-msl-artifacts "$RUMOCA_CI_HEAD_SHA" "$system" - name: Run MSL parity shard ${{ matrix.shard }}/4 run: | @@ -919,6 +1032,14 @@ jobs: --sim-total-memory-mb 6144 \ --monitor-interval-secs 10 + - name: Validate shard results before upload + shell: bash + run: | + set -euo pipefail + test -s target/msl/results/msl_results.json + test -s target/msl/results/omc_simulation_reference.json + test -s target/msl/results/sim_trace_comparison.json + - name: Upload shard results uses: actions/upload-artifact@v6 with: @@ -940,6 +1061,7 @@ jobs: steps: - uses: actions/checkout@v5 with: + ref: ${{ env.RUMOCA_CI_HEAD_SHA }} token: ${{ github.token }} persist-credentials: true @@ -972,14 +1094,19 @@ jobs: with: determinate: false - - name: Cachix (rumoca binary cache) - uses: cachix/cachix-action@v17 + - name: Download complete MSL Nix closure + uses: actions/download-artifact@v8 with: - name: rumoca + name: msl-nix-closure-${{ github.run_id }}-${{ env.RUMOCA_CI_HEAD_SHA }} + path: target/msl-nix-closure - - name: Restore prebuilt MSL binaries (Cachix) + - name: Restore and verify prebuilt MSL closure + shell: bash run: | - nix build .#msl-artifacts --out-link result-msl-artifacts + set -euo pipefail + system=$(nix eval --impure --raw --expr builtins.currentSystem) + .github/scripts/msl-nix-closure.sh restore \ + target/msl-nix-closure result-msl-artifacts "$RUMOCA_CI_HEAD_SHA" "$system" - name: Merge shards + run quality gate run: | @@ -1254,6 +1381,7 @@ jobs: steps: - uses: actions/checkout@v5 with: + ref: ${{ env.RUMOCA_CI_HEAD_SHA }} token: ${{ github.token }} persist-credentials: true @@ -1302,20 +1430,7 @@ jobs: - name: Install OpenModelica shell: bash - run: | - set -euo pipefail - omc_channel="stable" - sudo apt-get update - sudo apt-get install -y --no-install-recommends \ - build-essential ca-certificates clang cmake curl gnupg \ - libexpat1-dev liblapack-dev lsb-release unzip zip - curl -fsSL https://build.openmodelica.org/apt/openmodelica.asc \ - | sudo gpg --dearmor -o /usr/share/keyrings/openmodelica-keyring.gpg - echo "deb [arch=amd64 signed-by=/usr/share/keyrings/openmodelica-keyring.gpg] https://build.openmodelica.org/apt $(lsb_release -cs) $omc_channel" \ - | sudo tee /etc/apt/sources.list.d/openmodelica.list - sudo apt-get update - sudo apt-get install -y --no-install-recommends omc - omc --version + run: scripts/ci/install-openmodelica.sh - name: Prepare ModelicaTest sources shell: bash @@ -1331,14 +1446,19 @@ jobs: with: determinate: false - - name: Cachix (rumoca binary cache) - uses: cachix/cachix-action@v17 + - name: Download complete MSL Nix closure + uses: actions/download-artifact@v8 with: - name: rumoca + name: msl-nix-closure-${{ github.run_id }}-${{ env.RUMOCA_CI_HEAD_SHA }} + path: target/msl-nix-closure - - name: Restore prebuilt MSL binaries (Cachix) + - name: Restore and verify prebuilt MSL closure + shell: bash run: | - nix build .#msl-artifacts --out-link result-msl-artifacts + set -euo pipefail + system=$(nix eval --impure --raw --expr builtins.currentSystem) + .github/scripts/msl-nix-closure.sh restore \ + target/msl-nix-closure result-msl-artifacts "$RUMOCA_CI_HEAD_SHA" "$system" - name: Run ModelicaTest semantic gate shell: bash @@ -1383,13 +1503,12 @@ jobs: steps: - uses: actions/checkout@v5 with: + ref: ${{ env.RUMOCA_CI_HEAD_SHA }} token: ${{ github.token }} persist-credentials: true - name: Install system dependencies - run: | - sudo apt-get update - sudo apt-get install -y libudev-dev + run: scripts/ci/apt-install.sh libudev-dev - name: Install Rust run: | @@ -1424,13 +1543,12 @@ jobs: steps: - uses: actions/checkout@v5 with: + ref: ${{ env.RUMOCA_CI_HEAD_SHA }} token: ${{ github.token }} persist-credentials: true - name: Install system dependencies - run: | - sudo apt-get update - sudo apt-get install -y libudev-dev + run: scripts/ci/apt-install.sh libudev-dev - name: Install Rust run: | @@ -1464,6 +1582,7 @@ jobs: steps: - uses: actions/checkout@v5 with: + ref: ${{ env.RUMOCA_CI_HEAD_SHA }} token: ${{ github.token }} persist-credentials: true @@ -1472,11 +1591,6 @@ jobs: with: determinate: false - - name: Cachix (rumoca binary cache) - uses: cachix/cachix-action@v17 - with: - name: rumoca - - name: Rust cache (wasm) uses: Swatinem/rust-cache@v2 with: @@ -1597,13 +1711,12 @@ jobs: steps: - uses: actions/checkout@v5 with: + ref: ${{ env.RUMOCA_CI_HEAD_SHA }} token: ${{ github.token }} persist-credentials: true - name: Install system dependencies - run: | - sudo apt-get update - sudo apt-get install -y libudev-dev + run: scripts/ci/apt-install.sh libudev-dev - name: Install Rust nightly with rust-src run: | @@ -1666,6 +1779,7 @@ jobs: steps: - uses: actions/checkout@v5 with: + ref: ${{ env.RUMOCA_CI_HEAD_SHA }} token: ${{ github.token }} persist-credentials: true @@ -1680,9 +1794,7 @@ jobs: - name: Install musl tools (Linux musl) if: contains(matrix.target, 'musl') - run: | - sudo apt-get update - sudo apt-get install -y musl-tools + run: scripts/ci/apt-install.sh musl-tools - name: Configure musl build (Linux musl) if: contains(matrix.target, 'musl') @@ -1802,6 +1914,7 @@ jobs: steps: - uses: actions/checkout@v5 with: + ref: ${{ env.RUMOCA_CI_HEAD_SHA }} token: ${{ github.token }} persist-credentials: true @@ -1811,12 +1924,6 @@ jobs: with: determinate: false - - name: Cachix (rumoca binary cache) - if: runner.os != 'Windows' - uses: cachix/cachix-action@v17 - with: - name: rumoca - - uses: actions/setup-python@v6 if: runner.os == 'Windows' with: @@ -1904,6 +2011,7 @@ jobs: steps: - uses: actions/checkout@v5 with: + ref: ${{ env.RUMOCA_CI_HEAD_SHA }} token: ${{ github.token }} persist-credentials: true @@ -1912,11 +2020,6 @@ jobs: with: determinate: false - - name: Cachix (rumoca binary cache) - uses: cachix/cachix-action@v17 - with: - name: rumoca - - name: Build sdist run: | nix develop --command bash -lc ' @@ -1946,6 +2049,7 @@ jobs: steps: - uses: actions/checkout@v5 with: + ref: ${{ env.RUMOCA_CI_HEAD_SHA }} token: ${{ github.token }} persist-credentials: true fetch-depth: 0 diff --git a/.nvmrc b/.nvmrc new file mode 100644 index 000000000..209e3ef4b --- /dev/null +++ b/.nvmrc @@ -0,0 +1 @@ +20 diff --git a/AGENTS.md b/AGENTS.md index 9a075fd5f..23b4648be 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -13,7 +13,7 @@ spec, it's not a rule — propose a spec change first. how specs themselves work; what the statuses mean; how to propose a new spec. - [CONTRIBUTING.md](CONTRIBUTING.md) — local setup and `cargo xtask` CLI usage. -- [spec/SPEC_0032_DEVELOPMENT_PROCESS.md](spec/SPEC_0032_DEVELOPMENT_PROCESS.md) — +- [spec/SPEC_0033_DEVELOPMENT_PROCESS.md](spec/SPEC_0033_DEVELOPMENT_PROCESS.md) — operational workflow, triage proof requirements, upstream-first fix policy, and MSL-backed validation expectations. @@ -29,7 +29,7 @@ spec, it's not a rule — propose a spec change first. | Diagnostics, spans, error codes, tracing | [SPEC_0008](spec/SPEC_0008_PHASE_ERRORS.md) | | Tool config (`rumoca-tool-*`) | [SPEC_0018](spec/SPEC_0018_TOOL_CONFIG.md) | | Function length, nesting, file size, deterministic collections, code-size policy | [SPEC_0021](spec/SPEC_0021_CODE_COMPLEXITY.md) | -| Development workflow, bug triage, root-cause proof, upstream-first fixes | [SPEC_0032](spec/SPEC_0032_DEVELOPMENT_PROCESS.md) | +| Development workflow, bug triage, root-cause proof, upstream-first fixes | [SPEC_0033](spec/SPEC_0033_DEVELOPMENT_PROCESS.md) | | Opening a PR (workflow, metrics, verification commands, MSL gates, done criteria) | [SPEC_0025](spec/SPEC_0025_PR_REVIEW_PROCESS.md) | ## Rules of thumb diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index ddc6566c0..49d8e8938 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -40,6 +40,23 @@ The canonical top-level command groups are: ## Local Prerequisites +Put rustup's proxy directory before package-manager Rust installations so a +plain `cargo` command honors `rust-toolchain.toml`: + +```bash +export PATH="$HOME/.cargo/bin:$PATH" +rustc --version +cargo --version +``` + +MSL parity uses the exact OpenModelica Debian build recorded in +`toolchains/openmodelica-version`. CI downloads the four official package files +listed in `toolchains/openmodelica-packages.txt`, verifies their committed +SHA-256 digests and Debian metadata before installation, and then requires each +installed package version and the full `omc --version` build identity to match +the pin. Older, silently upgraded, pre-release, and different build outputs are +rejected. + Rust-only workflows do not require Node/npm: ```bash @@ -53,6 +70,7 @@ Package, playground, VS Code, and browser-asset workflows do require Node/npm. CI uses Node 20, so local package validation should use Node 20 as well: ```bash +nvm use node --version npm --version ``` diff --git a/Cargo.lock b/Cargo.lock index a9981b515..24b459aff 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -387,18 +387,18 @@ dependencies = [ [[package]] name = "bytemuck" -version = "1.25.0" +version = "1.25.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c8efb64bd706a16a1bdde310ae86b351e4d21550d98d056f22f8a7f7a2183fec" +checksum = "d6aedf8ae72766347502cf3cb4f41cf5e9cc37d28bee90f1fdaaae15f9cf9424" dependencies = [ "bytemuck_derive", ] [[package]] name = "bytemuck_derive" -version = "1.10.2" +version = "1.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f9abbd1bc6865053c427f7198e6af43bfdedc55ab791faed4fbd361d789575ff" +checksum = "f65693059b6b9c588b9f62fed1cedbf0a8b805631457ea162d68f0de186f3de5" dependencies = [ "proc-macro2", "quote", @@ -413,9 +413,9 @@ checksum = "1fd0f2584146f6f2ef48085050886acf353beff7305ebd1ae69500e27c67f64b" [[package]] name = "bytes" -version = "1.12.0" +version = "1.12.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8ae3f5d315924270530207e2a68396c3cc547f6dca3fbdca317cfb1a51edb593" +checksum = "fc652a48c352aef3ea3aed32080501cf3ef6ed5da78602a020c991775b0aff04" [[package]] name = "bzip2" @@ -438,9 +438,9 @@ dependencies = [ [[package]] name = "cc" -version = "1.2.66" +version = "1.2.67" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f5d6cac793997bd970000024b2934968efe83b382de4fdcf4fcb46b6ee4ad996" +checksum = "e17dd265a7d0f31ef544e1b20e03add05d3b45b491b633b10d67145d2acc1a38" dependencies = [ "find-msvc-tools", "jobserver", @@ -888,18 +888,18 @@ dependencies = [ [[package]] name = "crossbeam-channel" -version = "0.5.15" +version = "0.5.16" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "82b8f8f868b36967f9606790d1903570de9ceaf870a7bf9fbbd3016d636a2cb2" +checksum = "d85363c37faeca707aef026efa9f3b34d077bce547e48f770770625c6013679e" dependencies = [ "crossbeam-utils", ] [[package]] name = "crossbeam-deque" -version = "0.8.6" +version = "0.8.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9dd111b7b7f7d55b72c0a6ae361660ee5853c9af73f70c3c2ef6858b950e2e51" +checksum = "5181e0de7b61eb03a81e347d6dd8797bae9da5146707b51077e2d71a54ec0ceb" dependencies = [ "crossbeam-epoch", "crossbeam-utils", @@ -907,27 +907,27 @@ dependencies = [ [[package]] name = "crossbeam-epoch" -version = "0.9.18" +version = "0.9.20" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5b82ac4a3c2ca9c3460964f020e1402edd5753411d7737aa39c3714ad1b5420e" +checksum = "2d6914041f254d6e9176c01941b21115dcfb7089e55135a35411081bd106ef3f" dependencies = [ "crossbeam-utils", ] [[package]] name = "crossbeam-queue" -version = "0.3.12" +version = "0.3.13" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0f58bbc28f91df819d0aa2a2c00cd19754769c2fad90579b3592b1c9ba7a3115" +checksum = "803d13fb3b09d88be9f4dbc29062c66b19bf7170867ceb746d2a8689bf6c7a26" dependencies = [ "crossbeam-utils", ] [[package]] name = "crossbeam-utils" -version = "0.8.21" +version = "0.8.22" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d0a5c400df2834b80a4c3327b3aad3a4c4cd4de0629063962b03235697506a28" +checksum = "61803da095bee82a81bb1a452ecc25d3b2f1416d1897eb86430c6159ef717c17" [[package]] name = "crossterm" @@ -1441,7 +1441,7 @@ dependencies = [ "num-traits", "private-gemm-x86", "pulp", - "rand 0.9.4", + "rand 0.9.5", "rand_distr", "rayon", "reborrow", @@ -2314,9 +2314,9 @@ dependencies = [ [[package]] name = "inotify" -version = "0.11.2" +version = "0.11.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "533e68a5842e734946fe159fb03fc9bbbb254f590dd0d8ad321ae5ff7beca2c1" +checksum = "153be1941a183ec9ccd095ddbe17a8b8d435ef6c76e9e02451b933c3999af2c8" dependencies = [ "bitflags 2.13.0", "inotify-sys", @@ -2325,9 +2325,9 @@ dependencies = [ [[package]] name = "inotify-sys" -version = "0.1.7" +version = "0.1.8" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9ea94e891b3606826e9c998be69ddca42247dad8ad50b1649a5cb7e1c9ae06fd" +checksum = "c033f80b2c113cdf91ab7a33faa9cbc014726dcad99880c8609af2a370edf37d" dependencies = [ "libc", ] @@ -2402,9 +2402,9 @@ checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682" [[package]] name = "jiff" -version = "0.2.31" +version = "0.2.32" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ccfe6121cbe750cf81efa362d85c0bde7ea298ec43092d3a193baca59cdbd634" +checksum = "961d16382652bfdd8c6f68b223b26a8c93e0d475c672f414411db31c6c5c900e" dependencies = [ "defmt", "jiff-static", @@ -2416,9 +2416,9 @@ dependencies = [ [[package]] name = "jiff-static" -version = "0.2.31" +version = "0.2.32" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e165e897f662d428f3cd3828a919dbe067c2d42bb1031eede74ef9d27ecdedd2" +checksum = "d0879bd39df99c4c5e2c6615ccc026391a423dde10532c573e6086eb94a802cc" dependencies = [ "proc-macro2", "quote", @@ -2476,11 +2476,11 @@ dependencies = [ [[package]] name = "jobserver" -version = "0.1.34" +version = "0.1.35" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9afb3de4395d6b3e67a780b6de64b51c978ecf11cb9a462c66be7d4ca9039d33" +checksum = "1c00acbd29eabad4a2392fa0e921c874934dbbf4194312ad20f04a0ed67a3cb3" dependencies = [ - "getrandom 0.3.4", + "getrandom 0.4.3", "libc", ] @@ -2745,9 +2745,9 @@ dependencies = [ [[package]] name = "memchr" -version = "2.8.2" +version = "2.8.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "88904434abc2901f197fe8cc55f0445e7ded921dba5911dad2e2b39b48e663c4" +checksum = "cf8baf1c55e62ffcace7a9f06f4bd9cd3f0c4beb022d3b367256b91b87513d98" [[package]] name = "memo-map" @@ -3101,7 +3101,7 @@ dependencies = [ "num-integer", "num-iter", "num-traits", - "rand 0.8.6", + "rand 0.8.7", "smallvec", "zeroize", ] @@ -3114,7 +3114,7 @@ checksum = "73f88a1307638156682bada9d7604135552957b7818057dcef22705b4d509495" dependencies = [ "bytemuck", "num-traits", - "rand 0.8.6", + "rand 0.8.7", ] [[package]] @@ -3134,11 +3134,10 @@ dependencies = [ [[package]] name = "num-iter" -version = "0.1.45" +version = "0.1.46" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1429034a0490724d0075ebb2bc9e875d6503c3cf69e235a8941aa757d83ef5bf" +checksum = "c92800bd69a1eac91786bcfe9da64a897eb72911b8dc3095decbd07429e8048b" dependencies = [ - "autocfg", "num-integer", "num-traits", ] @@ -3291,7 +3290,7 @@ dependencies = [ "parol-macros", "parol_runtime", "petgraph", - "rand 0.9.4", + "rand 0.9.5", "rand_regex", "rayon", "regex", @@ -3860,9 +3859,9 @@ checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf" [[package]] name = "rand" -version = "0.8.6" +version = "0.8.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5ca0ecfa931c29007047d1bc58e623ab12e5590e8c7cc53200d5202b69266d8a" +checksum = "22f6172bdec972074665ed81ed53b71da00bfc44b65a753cfde883ec4c702a1a" dependencies = [ "libc", "rand_chacha 0.3.1", @@ -3871,9 +3870,9 @@ dependencies = [ [[package]] name = "rand" -version = "0.9.4" +version = "0.9.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "44c5af06bb1b7d3216d91932aed5265164bf384dc89cd6ba05cf59a35f5f76ea" +checksum = "b9ef1d0d795eb7d84685bca4f72f3649f064e6641543d3a8c415898726a57b41" dependencies = [ "rand_chacha 0.9.0", "rand_core 0.9.5", @@ -3941,7 +3940,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6a8615d50dcf34fa31f7ab52692afec947c4dd0ab803cc87cb3b0b4570ff7463" dependencies = [ "num-traits", - "rand 0.9.4", + "rand 0.9.5", ] [[package]] @@ -3959,7 +3958,7 @@ version = "0.18.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "04db2e382d13679a1e42400e90e306cdbb79dc5cd41bb035ba4eae72e78cdf37" dependencies = [ - "rand 0.9.4", + "rand 0.9.5", "regex-syntax", ] @@ -4087,9 +4086,9 @@ dependencies = [ [[package]] name = "regex" -version = "1.12.4" +version = "1.13.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f1292b7759ae1cb9ec195452d1390a074f0cd8541ab7a5a8c31cd6db45d4a6ba" +checksum = "2a0e75113e14dc5acb068cd0786884f214f1312650a3d36d269f5c4f3cdee8a2" dependencies = [ "aho-corasick", "memchr", @@ -4099,9 +4098,9 @@ dependencies = [ [[package]] name = "regex-automata" -version = "0.4.14" +version = "0.4.15" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6e1dd4122fc1595e8162618945476892eefca7b88c52820e74af6262213cae8f" +checksum = "1f388202e4b80542a0921078cc23b6333bcf1409c1e3f86404cae4766a6131db" dependencies = [ "aho-corasick", "memchr", @@ -4626,6 +4625,7 @@ dependencies = [ name = "rumoca-phase-flatten" version = "0.9.19" dependencies = [ + "flate2", "indexmap 2.14.0", "miette", "rumoca-core", @@ -4731,6 +4731,7 @@ name = "rumoca-sim" version = "0.9.19" dependencies = [ "anyhow", + "bincode", "indexmap 2.14.0", "libc", "rumoca-codec", @@ -4955,9 +4956,9 @@ dependencies = [ [[package]] name = "rustc-demangle" -version = "0.1.27" +version = "0.1.28" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b50b8869d9fc858ce7266cce0194bd74df58b9d0e3f6df3a9fc8eb470d95c09d" +checksum = "b74b56ffa8bb2830709a538c2cbcae9aa062db0d2a42563bfb09bdaae44020eb" [[package]] name = "rustc-hash" @@ -5095,9 +5096,9 @@ dependencies = [ [[package]] name = "rustversion" -version = "1.0.22" +version = "1.0.23" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b39cdef0fa800fc44525c84ccb54a029961a8215f9619753635a9c0d2538d46d" +checksum = "cf54715a573b99ac80df0bc206da022bcd442c974952c7b9720069370852e21f" [[package]] name = "ryu" @@ -5393,9 +5394,9 @@ dependencies = [ [[package]] name = "sha1" -version = "0.10.6" +version = "0.10.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba" +checksum = "a978451301f4db1d02937a4ab3ccce137717b81826e79b7d49ffe3244a13c3b8" dependencies = [ "cfg-if", "cpufeatures 0.2.17", @@ -5884,9 +5885,9 @@ dependencies = [ [[package]] name = "thread_local" -version = "1.1.9" +version = "1.1.10" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f60246a4944f24f6e018aa17cdeffb7818b76356965d03b07d6a9886e8962185" +checksum = "1ad99c4c6d32803332c548b1af0540b357b3f5fc0be8f6c6bfe8b2e6ae784070" dependencies = [ "cfg-if", ] @@ -5945,9 +5946,9 @@ dependencies = [ [[package]] name = "tinyvec" -version = "1.11.0" +version = "1.12.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3e61e67053d25a4e82c844e8424039d9745781b3fc4f32b8d55ed50f5f667ef3" +checksum = "bb4ebadaa0af04fab11ae01eb5f9fdb5f9c5b875506e210e71c07873528baa7f" dependencies = [ "tinyvec_macros", ] @@ -6286,7 +6287,7 @@ dependencies = [ "http", "httparse", "log", - "rand 0.8.6", + "rand 0.8.7", "sha1", "thiserror 1.0.69", "utf-8", @@ -6303,7 +6304,7 @@ dependencies = [ "http", "httparse", "log", - "rand 0.9.4", + "rand 0.9.5", "sha1", "thiserror 2.0.18", "utf-8", @@ -6346,7 +6347,7 @@ dependencies = [ "humantime", "lazy_static", "log", - "rand 0.8.6", + "rand 0.8.7", "serde", "spin 0.10.0", ] @@ -7313,7 +7314,7 @@ dependencies = [ "once_cell", "petgraph", "phf", - "rand 0.8.6", + "rand 0.8.7", "rustc_version", "serde", "serde_json", @@ -7417,7 +7418,7 @@ checksum = "44b80a042fc71419fc4952a90c9cbcfb323c0ced048125d8b44fd362f184045f" dependencies = [ "aes", "hmac", - "rand 0.8.6", + "rand 0.8.7", "rand_chacha 0.3.1", "sha3", "zenoh-result", @@ -7432,7 +7433,7 @@ dependencies = [ "getrandom 0.2.17", "hashbrown 0.16.1", "keyed-set", - "rand 0.8.6", + "rand 0.8.7", "schemars 1.2.1", "serde", "token-cell", @@ -7677,7 +7678,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "eeab45020bbecc077f14f06ee8f5aee65ce760af72481e663cac58b5dbfa66dd" dependencies = [ "const_format", - "rand 0.8.6", + "rand 0.8.7", "serde", "uhlc", "zenoh-buffers", @@ -7751,7 +7752,7 @@ dependencies = [ "futures", "lazy_static", "lz4_flex", - "rand 0.8.6", + "rand 0.8.7", "ringbuffer-spsc", "rsa", "serde", @@ -7803,18 +7804,18 @@ dependencies = [ [[package]] name = "zerocopy" -version = "0.8.52" +version = "0.8.54" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ce1022995ff5ff5d841ad7d994facc23098cd40152f2c1d11cd607c6f530653f" +checksum = "b7cbbc0a705a0fd05cc3676525980d2bf5a9bc4adac6d6475209a7887cf59d19" dependencies = [ "zerocopy-derive", ] [[package]] name = "zerocopy-derive" -version = "0.8.52" +version = "0.8.54" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1ae7f38b72ec2a254e2b87ef277cf2cd4fb97cbebf944faa6f33354da0867930" +checksum = "e2e817b7b52d0c7358d3246da9d69935ebb18116b2b102b4230dac079b4862f5" dependencies = [ "proc-macro2", "quote", diff --git a/Cargo.toml b/Cargo.toml index 6290c2b3e..11650183a 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -140,6 +140,9 @@ flate2 = "1.0" indexmap = { version = "2.7", features = ["serde"] } miette = { version = "7.4", features = ["fancy"] } minijinja = { version = "2.5", features = ["debug", "preserve_order"] } +# The generated parser, parol runtime, and scanner exchange concrete scnr2 +# types. Keep these exact so Cargo cannot resolve a newer parol stack that +# pulls an incompatible scnr2 minor version. parol = "=4.2.2" parol_runtime = "=4.2.0" quick-xml = "0.37" diff --git a/NOTICE b/NOTICE new file mode 100644 index 000000000..35dcb9327 --- /dev/null +++ b/NOTICE @@ -0,0 +1,23 @@ +Rumoca +Copyright 2025-2026 CogniPilot Foundation +Copyright 2026 ClimaMind contributors + +This product includes software developed by the Rumoca contributors. + +Rumoca is distributed under the Apache License, Version 2.0. A copy of the +license is included in LICENSE. + +ClimaMind maintains this repository as a public Rumoca line with attribution to +the original CogniPilot project. Downstream distributions should preserve this +NOTICE file and any third-party notices required by bundled dependencies. + +Dependency notes: + +- Rust dependencies are resolved through Cargo.lock and are primarily + permissively licensed. Run `cargo metadata --format-version 1` before a + release to refresh the dependency license inventory. +- The npm workspaces under `editors/` are distributed only with their package + lockfiles and package metadata. Run `npm install` in the relevant workspace + before packaging editor artifacts and retain required npm notices. +- Modelica Standard Library data is used by tests and CI staging. Do not commit + downloaded MSL trees or generated `target/msl` artifacts to this repository. diff --git a/README.md b/README.md index c59b51fcb..0cabe9ff6 100644 --- a/README.md +++ b/README.md @@ -2,13 +2,17 @@ Rumoca Logo -[![CI](https://github.com/cognipilot/rumoca/actions/workflows/ci.yml/badge.svg)](https://github.com/cognipilot/rumoca/actions/workflows/ci.yml) -[![GitHub Pages](https://img.shields.io/badge/GitHub%20Pages-live-2ea44f?logo=github)](https://cognipilot.github.io/rumoca/) +[![CI](https://github.com/climamind/rumoca/actions/workflows/ci.yml/badge.svg)](https://github.com/climamind/rumoca/actions/workflows/ci.yml) +[![GitHub Pages](https://img.shields.io/badge/GitHub%20Pages-live-2ea44f?logo=github)](https://climamind.github.io/rumoca/) [![PyPI](https://img.shields.io/pypi/v/rumoca)](https://pypi.org/project/rumoca/) [![npm](https://img.shields.io/npm/v/@cognipilot/rumoca)](https://www.npmjs.com/package/@cognipilot/rumoca) [![License](https://img.shields.io/badge/license-Apache--2.0-blue.svg)](LICENSE) -**[Try Rumoca in your browser](https://cognipilot.github.io/rumoca/)** (no installation required). +**[Try Rumoca in your browser](https://climamind.github.io/rumoca/)** (no installation required). + +This repository is the ClimaMind-maintained Rumoca line. Rumoca originated in +the CogniPilot community; see [NOTICE](NOTICE) for attribution and third-party +license notes. Rumoca is a modern **Modelica compiler and symbolic interoperability platform** written in Rust. @@ -16,7 +20,7 @@ Rumoca’s goal is not only to compile and simulate Modelica models, but to turn > **Rumoca turns Modelica package trees into modern symbolic systems.** -> **Project status:** Rumoca is in active development. You should expect bugs and rough edges; please file issues at https://github.com/cognipilot/rumoca/issues. +> **Project status:** Rumoca is in active development. You should expect bugs and rough edges; please file issues at https://github.com/climamind/rumoca/issues. ## Why Rumoca Exists @@ -239,19 +243,19 @@ This keeps the compiler’s multi-crate architecture intact (similar to rustc’ #### Binary installer (GitHub Releases) ```bash -curl --proto '=https' --tlsv1.2 -LsSf https://raw.githubusercontent.com/cognipilot/rumoca/main/infra/install/install.sh | bash +curl --proto '=https' --tlsv1.2 -LsSf https://raw.githubusercontent.com/climamind/rumoca/main/infra/install/install.sh | bash ``` Install a specific version (and optionally `rumoca-lsp`): ```bash -curl --proto '=https' --tlsv1.2 -LsSf https://raw.githubusercontent.com/cognipilot/rumoca/main/infra/install/install.sh | bash -s -- --version v0.8.0 --with-lsp +curl --proto '=https' --tlsv1.2 -LsSf https://raw.githubusercontent.com/climamind/rumoca/main/infra/install/install.sh | bash -s -- --version v0.8.0 --with-lsp ``` Windows PowerShell: ```powershell -irm https://raw.githubusercontent.com/cognipilot/rumoca/main/infra/install/install.ps1 | iex +irm https://raw.githubusercontent.com/climamind/rumoca/main/infra/install/install.ps1 | iex ``` The installer defaults to: @@ -299,9 +303,9 @@ After that, the main command groups are: ## Documentation -- Playground: -- User book: (`docs/user-guide/`) -- Developer book: (`docs/dev-guide/`) +- Playground: +- User book: (`docs/user-guide/`) +- Developer book: (`docs/dev-guide/`) - Normative design rules: `spec/` The books explain how to use Rumoca and how the implementation is organized. diff --git a/SECURITY.md b/SECURITY.md new file mode 100644 index 000000000..fb74218e5 --- /dev/null +++ b/SECURITY.md @@ -0,0 +1,27 @@ +# Security Policy + +## Supported Versions + +Security fixes are handled on the default branch first. Release tags and binary +artifacts are supported only when they are published from this repository's +GitHub Releases workflow. + +## Reporting a Vulnerability + +Please report suspected vulnerabilities privately by emailing +security@climamind.com. Include: + +- affected commit, release, or artifact; +- reproduction steps or a minimal input file when possible; +- expected impact and whether the issue affects CLI, LSP, Python bindings, + WASM, VS Code packaging, or generated code. + +Do not open a public issue for vulnerabilities involving arbitrary code +execution, path traversal, malicious Modelica inputs, supply-chain compromise, +or release artifact integrity. + +## Public Disclosure + +ClimaMind will coordinate a fix, credit, and disclosure timeline after the issue +is confirmed. If the issue also affects upstream Rumoca or third-party +dependencies, we will coordinate with the relevant maintainers. diff --git a/crates/rumoca-bind-python/pyproject.toml b/crates/rumoca-bind-python/pyproject.toml index f14fc1023..e99bc4e01 100644 --- a/crates/rumoca-bind-python/pyproject.toml +++ b/crates/rumoca-bind-python/pyproject.toml @@ -40,9 +40,9 @@ notebook = ["ipython"] all = ["numpy", "pandas", "matplotlib", "casadi", "jax", "sympy", "ipython"] [project.urls] -Homepage = "https://github.com/cognipilot/rumoca" -Repository = "https://github.com/cognipilot/rumoca" -Issues = "https://github.com/cognipilot/rumoca/issues" +Homepage = "https://github.com/climamind/rumoca" +Repository = "https://github.com/climamind/rumoca" +Issues = "https://github.com/climamind/rumoca/issues" [tool.maturin] manifest-path = "Cargo.toml" diff --git a/crates/rumoca-bind-wasm/src/gpu_api.rs b/crates/rumoca-bind-wasm/src/gpu_api.rs index c21a476d8..4f421d705 100644 --- a/crates/rumoca-bind-wasm/src/gpu_api.rs +++ b/crates/rumoca-bind-wasm/src/gpu_api.rs @@ -62,6 +62,8 @@ pub fn prepare_gpu_simulation(source: &str, model_name: &str) -> Result Result= 2 then 1 else pre(k) + 1; + when sample(0.02, 0.02) then + if pre(k) == 1 then + prev := 0.0; + else + prev := table[pre(k) - 1]; + end if; + y := u + prev; + k := if pre(k) >= 2 then 1 else pre(k) + 1; + end when; end DiscreteController; "#; @@ -523,13 +525,20 @@ fn test_interactive_session_runs_pure_discrete_model_with_guarded_dynamic_subscr ) .expect("pure discrete session should build"); session.set_input("u", 1.5).expect("set input u"); - session.step(0.02).expect("first discrete tick"); - assert_eq!(session.time(), 0.02); + session + .advance_to(0.01) + .expect("advance before first sample"); + assert_eq!(session.get("y").expect("read y"), Some(0.0)); + assert_eq!(session.get("k").expect("read k"), Some(1.0)); + + session.advance_to(0.02).expect("advance to first sample"); assert_eq!(session.get("y").expect("read y"), Some(1.5)); + assert_eq!(session.get("k").expect("read k"), Some(2.0)); session.set_input("u", 2.0).expect("set input u"); - session.step(0.02).expect("second discrete tick"); + session.advance_to(0.04).expect("advance to second sample"); assert_eq!(session.get("y").expect("read y"), Some(4.0)); + assert_eq!(session.get("k").expect("read k"), Some(1.0)); clear_source_root_cache().expect("clear source-root cache"); } diff --git a/crates/rumoca-bind-wasm/src/tests/simulation_runtime_tests.rs b/crates/rumoca-bind-wasm/src/tests/simulation_runtime_tests.rs index a1daa0fd0..a50aff750 100644 --- a/crates/rumoca-bind-wasm/src/tests/simulation_runtime_tests.rs +++ b/crates/rumoca-bind-wasm/src/tests/simulation_runtime_tests.rs @@ -1,5 +1,7 @@ #[cfg(any(feature = "sim-wasm", feature = "sim-diffsol", feature = "sim-rk45"))] use super::*; +#[cfg(any(feature = "sim-wasm", feature = "sim-diffsol", feature = "sim-rk45"))] +use crate::simulation_api::build_simulation_options; #[cfg(any(feature = "sim-wasm", feature = "sim-diffsol", feature = "sim-rk45"))] #[test] @@ -294,7 +296,203 @@ fn test_prepare_gpu_simulation_settles_wave_initial_equations() { center > 0.9, "GPU preparation must apply Wave2D-style initial equations; u[3,3]={center}" ); + assert!( + y0.iter() + .all(|value| value.as_f64().is_some_and(f64::is_finite)), + "GPU preparation must emit finite settled values" + ); + assert!( + names + .iter() + .zip(y0) + .filter(|(name, _)| name.as_str().is_some_and(|name| name.starts_with("w["))) + .all(|(_, value)| value.as_f64() == Some(0.0)), + "GPU preparation must preserve w initial family at zero" + ); + + clear_source_root_cache().expect("clear source-root cache"); +} +#[cfg(any(feature = "sim-wasm", feature = "sim-diffsol", feature = "sim-rk45"))] +#[test] +fn test_prepare_gpu_simulation_lowers_and_settles_descending_initial_binder() { + let _guard = session_test_guard(); + clear_source_root_cache().expect("clear source-root cache"); + let source = r#" + model DescendingGpuInitial + Real x[3]; + initial equation + for i in 3:-1:1 loop + x[i] = i; + end for; + equation + for i in 3:-1:1 loop + der(x[i]) = 0.0; + end for; + end DescendingGpuInitial; + "#; + + let json = prepare_gpu_simulation(source, "DescendingGpuInitial") + .expect("descending source binder should lower and settle natively"); + let payload: serde_json::Value = serde_json::from_str(&json).expect("valid GPU payload"); + let names = payload["state_names"].as_array().expect("state names"); + let y0 = payload["y0"].as_array().expect("settled y0"); + for (name, expected) in [("x[1]", 1.0), ("x[2]", 2.0), ("x[3]", 3.0)] { + let index = names + .iter() + .position(|candidate| candidate.as_str() == Some(name)) + .expect("state must be present"); + assert_eq!(y0[index].as_f64(), Some(expected)); + } + clear_source_root_cache().expect("clear source-root cache"); +} + +#[cfg(any(feature = "sim-wasm", feature = "sim-diffsol", feature = "sim-rk45"))] +fn assert_n50_compact_initialization(compact: &rumoca_ir_solve::SolveModel) { + let initialization = &compact.problem.initialization; + assert_eq!(compact.problem.layout.y_scalars(), 2 * 50 * 50); + assert_eq!(compact.initial_y.len(), 2 * 50 * 50); + let node_counts = initialization.residual.compute_node_counts(); + assert!(initialization.row_targets.is_empty()); + assert_eq!(initialization.direct_families.len(), 2); + assert_eq!(node_counts.map, 2); + assert_eq!(node_counts.scalar_programs, 0); + assert_eq!( + initialization.residual.nodes.len(), + initialization.direct_families.len() + ); + assert_eq!(initialization.residual.len(), Ok(2 * 50 * 50)); + for family in &initialization.direct_families { + let rumoca_ir_solve::ComputeNode::Map { domain, .. } = + &initialization.residual.nodes[family.node_index] + else { + panic!("direct initialization family must reference a Map") + }; + assert_eq!( + family + .targets + .output_indices(domain) + .map(|indices| indices.len()), + Ok(50 * 50) + ); + } + + let settled = rumoca_sim::settle_gpu_initial_conditions(compact, 0.0) + .expect("N=50 compact initialization should settle through native Map execution"); + assert_eq!(settled.metrics.residual_evaluations, 2); + assert_eq!(settled.metrics.passes, 1); + let max_map_scratch = initialization + .residual + .nodes + .iter() + .filter_map(|node| match node { + rumoca_ir_solve::ComputeNode::Map { + domain, base_ops, .. + } => Some( + domain + .binders + .len() + .saturating_mul(2) + .saturating_add(base_ops.len().max(1)), + ), + _ => None, + }) + .max() + .expect("compact GPU initialization has Map nodes"); + assert!( + settled.metrics.temporary_values + <= initialization + .direct_families + .len() + .saturating_add(max_map_scratch) + ); +} + +#[cfg(any(feature = "sim-wasm", feature = "sim-diffsol", feature = "sim-rk45"))] +#[test] +fn test_prepare_gpu_simulation_settles_wave_initial_equations_n50_in_linear_budget() { + let _guard = session_test_guard(); + clear_source_root_cache().expect("clear source-root cache"); + let source = r#" + model GpuWaveInitialN50 + parameter Integer N = 50; + parameter Real L = 1.0; + parameter Real dx = L / (N - 1); + Real u[N, N]; + Real w[N, N]; + initial equation + for i in 1:N loop + for j in 1:N loop + u[i, j] = exp(-200.0 * (((i - 1) * dx - 0.5 * L) ^ 2 + + ((j - 1) * dx - 0.5 * L) ^ 2)); + w[i, j] = 0.0; + end for; + end for; + equation + for i in 1:N loop + for j in 1:N loop + der(u[i, j]) = w[i, j]; + der(w[i, j]) = 0.0; + end for; + end for; + end GpuWaveInitialN50; + "#; + + let compact = with_singleton_session(|session| { + session.update_document("input.mo", source); + let requested_model = qualify_input_model_name(session, "GpuWaveInitialN50"); + let compiled = compile_requested_model(session, &requested_model)?; + let (opts, _) = build_simulation_options(&compiled, 0.0, 0.0, ""); + rumoca_sim::lower_dae_for_gpu_preparation(&compiled.dae, &opts) + .map_err(|error| JsValue::from_str(&format!("GPU lowering failed: {error}"))) + }) + .expect("N=50 should lower through compact GPU initialization"); + assert_n50_compact_initialization(&compact); + + let mut samples = Vec::new(); + for _ in 0..3 { + let started = std::time::Instant::now(); + let json = prepare_gpu_simulation(source, "GpuWaveInitialN50") + .expect("N=50 GPU preparation should settle explicit initial equations"); + samples.push(started.elapsed()); + let payload: serde_json::Value = + serde_json::from_str(&json).expect("GPU preparation payload should be valid JSON"); + let names = payload["state_names"] + .as_array() + .expect("GPU payload should include state names"); + let y0 = payload["y0"] + .as_array() + .expect("GPU payload should include initial y0"); + assert_eq!(names.len(), 2 * 50 * 50); + assert_eq!(y0.len(), 2 * 50 * 50); + let center = names + .iter() + .position(|name| name.as_str() == Some("u[25,25]")) + .and_then(|index| y0.get(index)) + .and_then(serde_json::Value::as_f64) + .expect("center displacement should be numeric"); + assert!(center > 0.9, "N=50 center must settle, got {center}"); + assert!( + y0.iter() + .all(|value| value.as_f64().is_some_and(f64::is_finite)), + "N=50 settled vector must be finite" + ); + assert!( + names + .iter() + .zip(y0) + .filter(|(name, _)| name.as_str().is_some_and(|name| name.starts_with("w["))) + .all(|(_, value)| value.as_f64() == Some(0.0)), + "N=50 w initial family must remain zero" + ); + } + samples.sort_unstable(); + assert!( + samples + .last() + .is_some_and(|elapsed| elapsed.as_secs_f64() < 10.0), + "N=50 GPU preparation p95-style smoke must retain >=5x debug headroom and stay below 10s: {samples:?}" + ); clear_source_root_cache().expect("clear source-root cache"); } diff --git a/crates/rumoca-codec-flatbuffers/src/bfbs.rs b/crates/rumoca-codec-flatbuffers/src/bfbs.rs index ec2cda59d..cf955403c 100644 --- a/crates/rumoca-codec-flatbuffers/src/bfbs.rs +++ b/crates/rumoca-codec-flatbuffers/src/bfbs.rs @@ -442,17 +442,28 @@ pub fn parse_bfbs(buf: &[u8]) -> anyhow::Result { #[cfg(test)] mod tests { use super::*; + use std::path::PathBuf; + + fn workspace_root() -> PathBuf { + PathBuf::from(env!("CARGO_MANIFEST_DIR")) + .parent() + .and_then(|path| path.parent()) + .expect("workspace root") + .to_path_buf() + } + + fn cerebri_bfbs_path(file_name: &str) -> PathBuf { + workspace_root().join("target/cerebri-bfbs").join(file_name) + } #[test] fn parse_cerebri2_topics() { - let path = Path::new( - "/home/micah/cognipilot/ws/cerebri/build-native_sim/generated/flatbuffers/cerebri2_topics.bfbs", - ); + let path = cerebri_bfbs_path("cerebri2_topics.bfbs"); if !path.exists() { eprintln!("skipping test: bfbs not found"); return; } - let data = std::fs::read(path).unwrap(); + let data = std::fs::read(&path).unwrap(); let schema = parse_bfbs(&data).unwrap(); eprintln!("Objects:"); @@ -487,14 +498,12 @@ mod tests { #[test] fn parse_cerebri2_sil() { - let path = Path::new( - "/home/micah/cognipilot/ws/cerebri/build-native_sim/generated/flatbuffers/cerebri2_sil.bfbs", - ); + let path = cerebri_bfbs_path("cerebri2_sil.bfbs"); if !path.exists() { eprintln!("skipping test: bfbs not found"); return; } - let data = std::fs::read(path).unwrap(); + let data = std::fs::read(&path).unwrap(); let schema = parse_bfbs(&data).unwrap(); eprintln!("file_ident: {:?}", schema.file_ident); diff --git a/crates/rumoca-codec-flatbuffers/src/codec.rs b/crates/rumoca-codec-flatbuffers/src/codec.rs index 88eda51a2..6065411d8 100644 --- a/crates/rumoca-codec-flatbuffers/src/codec.rs +++ b/crates/rumoca-codec-flatbuffers/src/codec.rs @@ -695,21 +695,32 @@ fn split_field_path(field_path: &str) -> Vec<&str> { #[cfg(test)] mod tests { use super::*; - use std::path::Path; + use std::path::PathBuf; + + fn workspace_root() -> PathBuf { + PathBuf::from(env!("CARGO_MANIFEST_DIR")) + .parent() + .and_then(|path| path.parent()) + .expect("workspace root") + .to_path_buf() + } + + fn cerebri_bfbs_paths() -> (PathBuf, PathBuf) { + let dir = workspace_root().join("target/cerebri-bfbs"); + ( + dir.join("cerebri2_topics.bfbs"), + dir.join("cerebri2_sil.bfbs"), + ) + } fn load_test_schema() -> Option { let mut ss = SchemaSet::new(); - let topics = Path::new( - "/home/micah/cognipilot/ws/cerebri/build-native_sim/generated/flatbuffers/cerebri2_topics.bfbs", - ); - let sil = Path::new( - "/home/micah/cognipilot/ws/cerebri/build-native_sim/generated/flatbuffers/cerebri2_sil.bfbs", - ); + let (topics, sil) = cerebri_bfbs_paths(); if !topics.exists() || !sil.exists() { return None; } - ss.load_bfbs(topics).unwrap(); - ss.load_bfbs(sil).unwrap(); + ss.load_bfbs(&topics).unwrap(); + ss.load_bfbs(&sil).unwrap(); Some(ss) } @@ -1120,9 +1131,25 @@ mod integration_tests { use super::*; use crate::bfbs; use std::net::UdpSocket; - use std::path::Path; + use std::path::PathBuf; use std::time::Duration; + fn workspace_root() -> PathBuf { + PathBuf::from(env!("CARGO_MANIFEST_DIR")) + .parent() + .and_then(|path| path.parent()) + .expect("workspace root") + .to_path_buf() + } + + fn cerebri_bfbs_paths() -> (PathBuf, PathBuf) { + let dir = workspace_root().join("target/cerebri-bfbs"); + ( + dir.join("cerebri2_topics.bfbs"), + dir.join("cerebri2_sil.bfbs"), + ) + } + /// Simulate the full cerebri ↔ rumoca loop locally. /// Sends a fake MotorOutput, verifies the sim can unpack it, /// and verifies the packed SimInput is valid. @@ -1130,18 +1157,13 @@ mod integration_tests { #[allow(clippy::too_many_lines)] fn full_loop_simulation() { let mut ss = bfbs::SchemaSet::new(); - let topics = Path::new( - "/home/micah/cognipilot/ws/cerebri/build-native_sim/generated/flatbuffers/cerebri2_topics.bfbs", - ); - let sil = Path::new( - "/home/micah/cognipilot/ws/cerebri/build-native_sim/generated/flatbuffers/cerebri2_sil.bfbs", - ); + let (topics, sil) = cerebri_bfbs_paths(); if !topics.exists() || !sil.exists() { eprintln!("skipping: bfbs files not found"); return; } - ss.load_bfbs(topics).unwrap(); - ss.load_bfbs(sil).unwrap(); + ss.load_bfbs(&topics).unwrap(); + ss.load_bfbs(&sil).unwrap(); // Build a MotorOutput packet the way cerebri does (48 bytes) // Using the pack codec to build a valid MotorOutput diff --git a/crates/rumoca-compile/src/cache.rs b/crates/rumoca-compile/src/cache.rs index 0430786e1..1680131f7 100644 --- a/crates/rumoca-compile/src/cache.rs +++ b/crates/rumoca-compile/src/cache.rs @@ -13,6 +13,7 @@ use crate::source_root_cache::resolve_cache_root_dir; pub const DEFAULT_CACHE_MAX_BYTES: u64 = 10 * 1024 * 1024 * 1024; const CACHE_ACCESS_METADATA_SUFFIX: &str = ".rumoca-access"; const CACHE_PRUNE_LOCK_STALE_AFTER: Duration = Duration::from_secs(30 * 60); +const CACHE_AUTO_PRUNE_CHECK_INTERVAL: Duration = Duration::from_secs(24 * 60 * 60); #[derive(Debug, Clone)] pub struct CacheStatus { @@ -138,6 +139,23 @@ pub(crate) fn maybe_prune_cache_after_write(root: Option<&Path>) { return; } }; + let should_prune = match mark_auto_prune_check_if_due( + &prune_root, + SystemTime::now(), + CACHE_AUTO_PRUNE_CHECK_INTERVAL, + ) { + Ok(should_prune) => should_prune, + Err(err) => { + eprintln!( + "failed to update cache auto-prune marker under {}: {err}", + prune_root.display() + ); + return; + } + }; + if !should_prune { + return; + } let report = match prune_cache_after_write_with_options(Some(&prune_root), &options) { Ok(Some(report)) => report, Ok(None) => return, @@ -495,6 +513,45 @@ fn cache_prune_lock_path(root: &Path) -> PathBuf { .unwrap_or_else(|| PathBuf::from(lock_name)) } +fn mark_auto_prune_check_if_due( + root: &Path, + now: SystemTime, + interval: Duration, +) -> std::io::Result { + let path = cache_auto_prune_marker_path(root); + let due = match fs::read_to_string(&path) { + Ok(content) => content + .trim() + .parse::() + .ok() + .and_then(|seconds| { + now.duration_since(UNIX_EPOCH + Duration::from_secs(seconds)) + .ok() + }) + .map(|age| age > interval) + .unwrap_or(false), + Err(error) if error.kind() == io::ErrorKind::NotFound => false, + Err(_) => false, + }; + if let Some(parent) = path.parent() { + fs::create_dir_all(parent)?; + } + fs::write(&path, system_time_secs_since_unix_epoch(now)?.to_string())?; + Ok(due) +} + +fn cache_auto_prune_marker_path(root: &Path) -> PathBuf { + let marker_name = root + .file_name() + .and_then(|name| name.to_str()) + .filter(|name| !name.is_empty()) + .map(|name| format!(".{name}.rumoca-auto-prune-checked")) + .unwrap_or_else(|| ".rumoca-cache.rumoca-auto-prune-checked".to_string()); + root.parent() + .map(|parent| parent.join(&marker_name)) + .unwrap_or_else(|| PathBuf::from(marker_name)) +} + #[cfg(test)] pub(crate) fn prune_cache_after_write( root: Option<&Path>, @@ -867,6 +924,54 @@ mod tests { assert_eq!(status.subcaches[1].total_bytes, 5); } + #[test] + fn auto_prune_marker_skips_first_foreground_check() { + let temp = tempfile::tempdir().expect("tempdir"); + let root = temp.path().join("cache"); + let now = SystemTime::UNIX_EPOCH + Duration::from_secs(100); + + assert!( + !mark_auto_prune_check_if_due(&root, now, Duration::from_secs(10)) + .expect("mark first check"), + "first foreground write should only create the marker" + ); + assert!(cache_auto_prune_marker_path(&root).exists()); + } + + #[test] + fn auto_prune_marker_rate_limits_foreground_checks() { + let temp = tempfile::tempdir().expect("tempdir"); + let root = temp.path().join("cache"); + let first = SystemTime::UNIX_EPOCH + Duration::from_secs(100); + let recent = first + Duration::from_secs(5); + + assert!( + !mark_auto_prune_check_if_due(&root, first, Duration::from_secs(10)) + .expect("mark first check") + ); + assert!( + !mark_auto_prune_check_if_due(&root, recent, Duration::from_secs(10)) + .expect("mark recent check") + ); + } + + #[test] + fn auto_prune_marker_allows_due_foreground_check() { + let temp = tempfile::tempdir().expect("tempdir"); + let root = temp.path().join("cache"); + let first = SystemTime::UNIX_EPOCH + Duration::from_secs(100); + let due = first + Duration::from_secs(11); + + assert!( + !mark_auto_prune_check_if_due(&root, first, Duration::from_secs(10)) + .expect("mark first check") + ); + assert!( + mark_auto_prune_check_if_due(&root, due, Duration::from_secs(10)) + .expect("mark due check") + ); + } + #[test] fn auto_prune_lock_allows_one_owner_per_cache_root() { let temp = tempfile::tempdir().expect("tempdir"); diff --git a/crates/rumoca-compile/src/codegen_target.rs b/crates/rumoca-compile/src/codegen_target.rs index 767ce0bf6..b1ed0a1c2 100644 --- a/crates/rumoca-compile/src/codegen_target.rs +++ b/crates/rumoca-compile/src/codegen_target.rs @@ -108,6 +108,13 @@ pub enum TensorLayoutCapability { pub struct TargetFile { pub path: String, pub template: String, + /// Whether this artifact may render as an empty file. Required artifacts + /// remain non-empty by default; targets must opt in per file. + #[serde(default)] + pub allow_empty: bool, + /// Optional per-file IR override for mixed-context targets (for example, + /// FMI resource metadata rendered from DAE alongside Solve runtime code). + pub ir: Option, pub render_context: Option, pub mode: Option, /// Stable logical identity of this rendered file within the target @@ -1112,6 +1119,30 @@ render_context = "fmi-model-description" ); } + #[test] + fn target_manifest_file_allow_empty_parses_and_defaults_false() { + let manifest = super::parse_target_manifest( + r#" +version = 1 +ir = "solve" +name = "custom" + +[[files]] +path = "required.txt" +template = "required.txt.jinja" + +[[files]] +path = "optional.txt" +template = "optional.txt.jinja" +allow_empty = true +"#, + ) + .expect("parse target manifest with optional empty output"); + + assert!(!manifest.files[0].allow_empty); + assert!(manifest.files[1].allow_empty); + } + #[test] fn all_builtin_target_manifests_parse() { for target in templates::builtin_targets() { @@ -1459,7 +1490,12 @@ events = false ); let capabilities = manifest.capabilities.as_ref().expect("capabilities"); let mut dae = Dae::new(); - dae.events.scheduled_time_events.push(0.1); + dae.events + .scheduled_time_events + .push(rumoca_ir_dae::DaeScheduledTimeEvent { + time: 0.1, + source_span: None, + }); let err = validate_dae_target_capabilities(&dae, &manifest, capabilities) .expect_err("events should be rejected"); diff --git a/crates/rumoca-compile/src/lib.rs b/crates/rumoca-compile/src/lib.rs index 96da6fbc6..0592ce933 100644 --- a/crates/rumoca-compile/src/lib.rs +++ b/crates/rumoca-compile/src/lib.rs @@ -184,8 +184,11 @@ pub mod galec { /// Read-only DAE analysis helpers exposed through the compile facade. pub mod analysis { - pub use rumoca_phase_dae::balance::BalanceDetail; - pub use rumoca_phase_dae::{balance, balance_detail, equations_unknowns, is_balanced}; + pub use rumoca_phase_dae::balance::{BalanceDetail, InitialClosureBalanceDetail}; + pub use rumoca_phase_dae::{ + balance, balance_detail, equations_unknowns, initial_closure_balance_detail, is_balanced, + is_balanced_for_admission, + }; } /// Structural-analysis primitives (BLT sorting, scalarization). diff --git a/crates/rumoca-compile/src/session/compile_support.rs b/crates/rumoca-compile/src/session/compile_support.rs index 566b5e583..f6a94f8dc 100644 --- a/crates/rumoca-compile/src/session/compile_support.rs +++ b/crates/rumoca-compile/src/session/compile_support.rs @@ -265,11 +265,12 @@ pub(super) fn compile_model_dae_internal_with_options( pub(super) fn compile_model_dae_internal_allow_unbalanced_for_diagnostics( tree: &ast::ClassTree, model_name: &str, + instantiation_options: InstantiateOptions, ) -> DaePhaseResult { let dae_outcome = dae_model_outcome_internal_with_phase_options( tree, model_name, - InstantiateOptions::default(), + instantiation_options, ToDaeOptions { error_on_unbalanced: false, }, diff --git a/crates/rumoca-compile/src/session/session_impl.rs b/crates/rumoca-compile/src/session/session_impl.rs index c28b0852a..691cf57eb 100644 --- a/crates/rumoca-compile/src/session/session_impl.rs +++ b/crates/rumoca-compile/src/session/session_impl.rs @@ -980,7 +980,11 @@ impl Session { 8, )); } - match compile_model_dae_internal_allow_unbalanced_for_diagnostics(tree, model_name) { + match compile_model_dae_internal_allow_unbalanced_for_diagnostics( + tree, + model_name, + self.instantiation_options.clone(), + ) { DaePhaseResult::Success(result) => Ok(result), DaePhaseResult::NeedsInner { missing_inners, .. } => Err(format!( "{model_name} requires inner declarations: {}", diff --git a/crates/rumoca-compile/src/session/tests.rs b/crates/rumoca-compile/src/session/tests.rs index bd1418fb7..62f20f300 100644 --- a/crates/rumoca-compile/src/session/tests.rs +++ b/crates/rumoca-compile/src/session/tests.rs @@ -1137,6 +1137,31 @@ fn test_compile_extracts_experiment_stop_time() { assert_eq!(result.experiment_solver, None); } +#[test] +fn test_compile_extracts_zero_experiment_stop_time_from_nested_model() { + let mut session = Session::default(); + session + .add_document( + "test.mo", + r#" + package P + package Examples + model M + Real x; + equation + x = 1; + annotation(experiment(StopTime=0)); + end M; + end Examples; + end P; + "#, + ) + .unwrap(); + + let result = session.compile_model("P.Examples.M").unwrap(); + assert_eq!(result.experiment_stop_time, Some(0.0)); +} + #[test] fn test_compile_ignores_negative_experiment_stop_time() { let mut session = Session::default(); diff --git a/crates/rumoca-contracts/tests/inst_contracts.rs b/crates/rumoca-contracts/tests/inst_contracts.rs index ca2887f2e..802b0e27e 100644 --- a/crates/rumoca-contracts/tests/inst_contracts.rs +++ b/crates/rumoca-contracts/tests/inst_contracts.rs @@ -213,6 +213,26 @@ fn inst_008_no_cyclic_binding() { ); } +#[test] +fn inst_008_bare_self_default_can_be_overridden_by_parent_modifier() { + expect_success( + r#" + model Child + parameter Real p = p; + Real x; + equation + x = p; + end Child; + + model Test + parameter Real p = 2; + Child child(p = p); + end Test; + "#, + "Test", + ); +} + // ============================================================================= // INST-010: Final immutability // "Element defined as final cannot be modified by modification or redeclaration" diff --git a/crates/rumoca-contracts/tests/sim_contracts.rs b/crates/rumoca-contracts/tests/sim_contracts.rs index 2cd0169d5..5c153ec2f 100644 --- a/crates/rumoca-contracts/tests/sim_contracts.rs +++ b/crates/rumoca-contracts/tests/sim_contracts.rs @@ -401,7 +401,7 @@ fn sim_009_runtime_metadata_consistent_for_hybrid_model() { .events .scheduled_time_events .iter() - .any(|event| (*event - 0.5).abs() <= 1.0e-12), + .any(|event| (event.time - 0.5).abs() <= 1.0e-12), "time-driven discontinuity should be reflected in scheduled_time_events" ); assert!( @@ -410,7 +410,7 @@ fn sim_009_runtime_metadata_consistent_for_hybrid_model() { .events .scheduled_time_events .iter() - .all(|event| event.is_finite()), + .all(|event| event.time.is_finite()), "scheduled_time_events must contain finite values" ); } diff --git a/crates/rumoca-core/src/ir_primitives.rs b/crates/rumoca-core/src/ir_primitives.rs index 35101e7f8..c4c0d6b52 100644 --- a/crates/rumoca-core/src/ir_primitives.rs +++ b/crates/rumoca-core/src/ir_primitives.rs @@ -1203,6 +1203,17 @@ pub struct ExternalFunction { pub output_name: Option, /// Argument names passed to the external function. pub arg_names: Vec, + /// Library names from external function annotation(Library=...). + pub libraries: Vec, + /// Include directory URIs from annotation(IncludeDirectory=...). + pub include_directories: Vec, + /// Library directory URIs from annotation(LibraryDirectory=...). + pub library_directories: Vec, + /// Raw include snippets from annotation(Include=...). + pub includes: Vec, + /// Raw external annotation arguments, retained for diagnostics and future + /// platform-specific native library packaging. + pub annotation: Vec, } /// Function derivative annotation (MLS §12.7.1). diff --git a/crates/rumoca-core/src/lib.rs b/crates/rumoca-core/src/lib.rs index 97b5a4950..8879cf28d 100644 --- a/crates/rumoca-core/src/lib.rs +++ b/crates/rumoca-core/src/lib.rs @@ -498,17 +498,64 @@ pub fn span_to_source_span(span: Span) -> SourceSpan { /// - `StateSelect` - Enumeration for state selection hints (MLS §4.4.4.2) /// - `AssertionLevel` - Enumeration for assertion levels (MLS §8.3.7) pub const BUILTIN_TYPES: &[&str] = &[ - "Real", - "Integer", - "Boolean", - "String", - "ExternalObject", - "Clock", - // Built-in enumerations (MLS §4.4.4.2, §8.3.7) - "StateSelect", - "AssertionLevel", + BuiltinTypeIdentity::Real.name(), + BuiltinTypeIdentity::Integer.name(), + BuiltinTypeIdentity::Boolean.name(), + BuiltinTypeIdentity::String.name(), + BuiltinTypeIdentity::ExternalObject.name(), + BuiltinTypeIdentity::Clock.name(), + BuiltinTypeIdentity::StateSelect.name(), + BuiltinTypeIdentity::AssertionLevel.name(), ]; +/// Compiler-owned identities for MLS builtin types. +/// +/// Resolver registers these definitions before user declarations. The typed +/// identity owns both the canonical name at the source boundary and the +/// semantic `DefId` carried by resolved Flat/DAE references. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[repr(u32)] +pub enum BuiltinTypeIdentity { + Real, + Integer, + Boolean, + String, + ExternalObject, + Clock, + StateSelect, + AssertionLevel, +} + +impl BuiltinTypeIdentity { + pub const ALL: [Self; 8] = [ + Self::Real, + Self::Integer, + Self::Boolean, + Self::String, + Self::ExternalObject, + Self::Clock, + Self::StateSelect, + Self::AssertionLevel, + ]; + + pub const fn name(self) -> &'static str { + match self { + Self::Real => "Real", + Self::Integer => "Integer", + Self::Boolean => "Boolean", + Self::String => "String", + Self::ExternalObject => "ExternalObject", + Self::Clock => "Clock", + Self::StateSelect => "StateSelect", + Self::AssertionLevel => "AssertionLevel", + } + } + + pub const fn def_id(self) -> DefId { + DefId(self as u32 + 1) + } +} + /// Built-in functions (MLS §3.7). /// /// These are predefined operators and functions available in all scopes. @@ -907,6 +954,11 @@ impl PrimaryLabel { } } + /// Source span covered by this label. + pub fn span(&self) -> Span { + self.span + } + /// Attach a human-readable message to the label. pub fn with_message(mut self, msg: impl Into) -> Self { self.message = Some(msg.into()); diff --git a/crates/rumoca-core/src/structured_domain.rs b/crates/rumoca-core/src/structured_domain.rs index 9e404d210..ca6ba1855 100644 --- a/crates/rumoca-core/src/structured_domain.rs +++ b/crates/rumoca-core/src/structured_domain.rs @@ -240,7 +240,8 @@ fn reserve_current_tuple_capacity( } impl StructuredIndexBinder { - fn value_count(&self) -> Result { + /// Number of values enumerated by this binder's inclusive stepped range. + pub fn value_count(&self) -> Result { if self.step == 0 { return Err(StructuredIndexDomainError::ZeroStep { binder_id: self.id, diff --git a/crates/rumoca-eval-ast/src/eval/dimension_inference.rs b/crates/rumoca-eval-ast/src/eval/dimension_inference.rs index 53cf94708..fcd588f08 100644 --- a/crates/rumoca-eval-ast/src/eval/dimension_inference.rs +++ b/crates/rumoca-eval-ast/src/eval/dimension_inference.rs @@ -43,8 +43,12 @@ pub fn infer_dimensions_from_binding_with_scope( Expression::ComponentReference(cr) => { let indexed_path = cr.to_string(); - if let Some(dims) = lookup_structural_with_scope(&indexed_path, scope, &ctx.dimensions) - { + if let Some(dims) = lookup_structural_with_scope( + &indexed_path, + scope, + &ctx.dimensions, + ctx.suffix_index.as_ref(), + ) { return Some(dims.clone()); } @@ -54,9 +58,12 @@ pub fn infer_dimensions_from_binding_with_scope( .map(|p| p.ident.text.as_ref()) .collect::>() .join("."); - let Some(base_dims) = - lookup_structural_with_scope(&unindexed_path, scope, &ctx.dimensions) - else { + let Some(base_dims) = lookup_structural_with_scope( + &unindexed_path, + scope, + &ctx.dimensions, + ctx.suffix_index.as_ref(), + ) else { return scalar_value_known_with_scope(&unindexed_path, ctx, scope).then(Vec::new); }; let base_dims = base_dims.clone(); @@ -65,6 +72,15 @@ pub fn infer_dimensions_from_binding_with_scope( )) } + Expression::ArrayIndex { + base, subscripts, .. + } => { + let base_dims = infer_dimensions_from_binding_with_scope(base, ctx, scope)?; + Some(apply_array_index_subscripts_to_dims( + base_dims, subscripts, ctx, scope, + )) + } + Expression::Parenthesized { inner, .. } => { infer_dimensions_from_binding_with_scope(inner, ctx, scope) } @@ -88,7 +104,13 @@ pub fn infer_dimensions_from_binding_with_scope( Expression::FieldAccess { base, field, .. } => { let base_path = extract_simple_component_path(base)?; let full_path = format!("{base_path}.{field}"); - lookup_structural_with_scope(&full_path, scope, &ctx.dimensions).cloned() + lookup_structural_with_scope( + &full_path, + scope, + &ctx.dimensions, + ctx.suffix_index.as_ref(), + ) + .cloned() } // ArrayComprehension: `{expr for i in range}` -> `[range_len, inner_dims...]`. @@ -103,11 +125,11 @@ pub fn infer_dimensions_from_binding_with_scope( } fn scalar_value_known_with_scope(name: &str, ctx: &TypeCheckEvalContext, scope: &str) -> bool { - lookup_with_scope(name, scope, &ctx.integers).is_some() - || lookup_with_scope(name, scope, &ctx.reals).is_some() - || lookup_with_scope(name, scope, &ctx.booleans).is_some() - || lookup_with_scope(name, scope, &ctx.enums).is_some() - || lookup_with_scope(name, scope, &ctx.enum_ordinals).is_some() + lookup_with_scope(name, scope, &ctx.integers, ctx.suffix_index.as_ref()).is_some() + || lookup_with_scope(name, scope, &ctx.reals, ctx.suffix_index.as_ref()).is_some() + || lookup_with_scope(name, scope, &ctx.booleans, ctx.suffix_index.as_ref()).is_some() + || lookup_with_scope(name, scope, &ctx.enums, ctx.suffix_index.as_ref()).is_some() + || lookup_with_scope(name, scope, &ctx.enum_ordinals, ctx.suffix_index.as_ref()).is_some() } /// Apply component-reference subscripts to a base dimension vector. @@ -138,6 +160,22 @@ fn apply_component_subscripts_to_dims( dims } +fn apply_array_index_subscripts_to_dims( + mut dims: Vec, + subscripts: &[Subscript], + ctx: &TypeCheckEvalContext, + scope: &str, +) -> Vec { + let mut pos = 0usize; + for sub in subscripts { + if pos >= dims.len() { + return dims; + } + apply_subscript_to_dims(sub, &mut dims, &mut pos, ctx, scope); + } + dims +} + fn apply_subscript_to_dims( sub: &Subscript, dims: &mut Vec, @@ -214,6 +252,11 @@ fn infer_range_length( } = range { infer_range_len_numeric(start, step.as_deref(), end, ctx, scope) + } else if let Expression::FunctionCall { comp, args, .. } = range + && comp.to_string() == "linspace" + && args.len() == 3 + { + eval_integer_with_scope(&args[2], ctx, scope).map(|n| n as usize) } else { eval_integer_with_scope(range, ctx, scope).map(|n| n as usize) } diff --git a/crates/rumoca-eval-ast/src/eval/eval_lookup_impl.rs b/crates/rumoca-eval-ast/src/eval/eval_lookup_impl.rs index 43b854914..5637a3dee 100644 --- a/crates/rumoca-eval-ast/src/eval/eval_lookup_impl.rs +++ b/crates/rumoca-eval-ast/src/eval/eval_lookup_impl.rs @@ -5,23 +5,30 @@ use std::borrow::Cow; impl EvalLookup for TypeCheckEvalContext { fn lookup_integer(&self, name: &str, scope: &str) -> Option { - lookup_with_scope(name, scope, &self.integers) + lookup_with_scope(name, scope, &self.integers, self.suffix_index.as_ref()) .copied() - .or_else(|| lookup_with_scope(name, scope, &self.enum_ordinals).copied()) + .or_else(|| { + lookup_with_scope(name, scope, &self.enum_ordinals, self.suffix_index.as_ref()) + .copied() + }) } fn lookup_real(&self, name: &str, scope: &str) -> Option { - lookup_with_scope(name, scope, &self.reals) + lookup_with_scope(name, scope, &self.reals, self.suffix_index.as_ref()) .copied() - .or_else(|| lookup_with_scope(name, scope, &self.integers).map(|value| *value as f64)) + .or_else(|| { + lookup_with_scope(name, scope, &self.integers, self.suffix_index.as_ref()) + .map(|value| *value as f64) + }) } fn lookup_boolean(&self, name: &str, scope: &str) -> Option { - lookup_with_scope(name, scope, &self.booleans).copied() + lookup_with_scope(name, scope, &self.booleans, self.suffix_index.as_ref()).copied() } fn lookup_enum<'a>(&'a self, name: &str, scope: &str) -> Option> { - lookup_with_scope(name, scope, &self.enums).map(|value| Cow::Borrowed(value.as_str())) + lookup_with_scope(name, scope, &self.enums, self.suffix_index.as_ref()) + .map(|value| Cow::Borrowed(value.as_str())) } } @@ -34,10 +41,9 @@ mod tests { let mut ctx = TypeCheckEvalContext::new(); ctx.add_integer("sys.n", 4); ctx.add_real("sys.inner.r", 2.5); - ctx.booleans.insert("sys.flag".to_string(), true); - ctx.enums - .insert("sys.mode".to_string(), "Pkg.Mode.Fast".to_string()); - ctx.enum_ordinals.insert("sys.phase".to_string(), 3); + ctx.add_boolean("sys.flag", true); + ctx.add_enum("sys.mode", "Pkg.Mode.Fast"); + ctx.add_enum_ordinal("sys.phase", 3); assert_eq!(ctx.lookup_integer("n", "sys.inner"), Some(4)); assert_eq!(ctx.lookup_integer("phase", "sys.inner"), Some(3)); diff --git a/crates/rumoca-eval-ast/src/eval/late_inference.rs b/crates/rumoca-eval-ast/src/eval/late_inference.rs index 675058163..8055d174a 100644 --- a/crates/rumoca-eval-ast/src/eval/late_inference.rs +++ b/crates/rumoca-eval-ast/src/eval/late_inference.rs @@ -321,49 +321,55 @@ mod tests { } #[test] - fn test_lookup_with_scope_uses_structured_scope() { + fn test_lookup_with_scope_dotted_name_uses_full_suffix() { let mut map = FxHashMap::default(); map.insert("sys.Medium.nX".to_string(), 4_i64); - assert_eq!(lookup_with_scope("Medium.nX", "sys", &map), Some(&4_i64)); + assert_eq!(lookup_with_scope("Medium.nX", "", &map, None), Some(&4_i64)); } #[test] - fn test_lookup_with_scope_dotted_name_does_not_leaf_fallback() { + fn test_lookup_with_scope_dotted_name_falls_back_to_leaf_when_full_suffix_missing() { let mut map = FxHashMap::default(); map.insert("sys.nX".to_string(), 7_i64); - assert_eq!(lookup_with_scope("Medium.nX", "", &map), None); + assert_eq!(lookup_with_scope("Medium.nX", "", &map, None), Some(&7_i64)); } #[test] - fn test_lookup_with_scope_dotted_name_requires_matching_scoped_name() { + fn test_lookup_with_scope_dotted_name_leaf_fallback_requires_unique_key() { let mut map = FxHashMap::default(); map.insert("a.nX".to_string(), 7_i64); map.insert("b.nX".to_string(), 7_i64); - assert_eq!(lookup_with_scope("Medium.nX", "", &map), None); + assert_eq!(lookup_with_scope("Medium.nX", "", &map, None), None); } #[test] - fn test_lookup_with_scope_dotted_name_ignores_ambiguous_suffixes() { + fn test_lookup_with_scope_dotted_name_does_not_fallback_when_full_suffix_is_ambiguous() { let mut map = FxHashMap::default(); map.insert("a.Medium.nX".to_string(), 1_i64); map.insert("b.Medium.nX".to_string(), 2_i64); - assert_eq!(lookup_with_scope("Medium.nX", "", &map), None); + assert_eq!(lookup_with_scope("Medium.nX", "", &map, None), None); } #[test] - fn test_lookup_with_scope_simple_name_does_not_suffix_fallback() { + fn test_lookup_with_scope_simple_name_still_uses_suffix_fallback() { let mut map = FxHashMap::default(); map.insert("sys.nX".to_string(), 7_i64); - assert_eq!(lookup_with_scope("nX", "", &map), None); + assert_eq!(lookup_with_scope("nX", "", &map, None), Some(&7_i64)); } #[test] - fn test_lookup_with_scope_resolves_indexed_name_with_scope() { + fn test_lookup_with_scope_treats_dot_inside_subscript_as_single_segment() { let mut ctx = TypeCheckEvalContext::new(); ctx.add_integer("sys.arr[data.medium]", 7_i64); + ctx.build_suffix_index(); assert_eq!( - lookup_with_scope("arr[data.medium]", "sys", &ctx.integers), + lookup_with_scope( + "arr[data.medium]", + "", + &ctx.integers, + ctx.suffix_index.as_ref() + ), Some(&7_i64) ); } @@ -372,32 +378,54 @@ mod tests { fn test_lookup_with_scope_does_not_index_fake_suffix_from_subscript_dot() { let mut ctx = TypeCheckEvalContext::new(); ctx.add_integer("sys.arr[data.medium]", 7_i64); + ctx.build_suffix_index(); - assert_eq!(lookup_with_scope("medium]", "", &ctx.integers), None); + assert_eq!( + lookup_with_scope("medium]", "", &ctx.integers, ctx.suffix_index.as_ref()), + None + ); } #[test] - fn test_lookup_with_scope_does_not_scan_suffixes() { + fn test_lookup_with_scope_linear_suffix_match_ignores_subscript_dot_boundary() { let mut map = FxHashMap::default(); map.insert("sys.arr[data.medium].x".to_string(), 7_i64); map.insert("other.scope.x".to_string(), 9_i64); - assert_eq!(lookup_with_scope("medium].x", "", &map), None); + assert_eq!(lookup_with_scope("medium].x", "", &map, None), None); } #[test] - fn test_lookup_with_scope_no_cross_map_suffix_fallback() { + fn test_lookup_with_scope_leaf_fallback_checks_uniqueness_per_target_map() { let mut ctx = TypeCheckEvalContext::new(); ctx.add_integer("a.nX", 7_i64); ctx.add_real("b.nX", 3.0); + ctx.build_suffix_index(); - assert_eq!(lookup_with_scope("Medium.nX", "", &ctx.integers), None); + assert_eq!( + lookup_with_scope("Medium.nX", "", &ctx.integers, ctx.suffix_index.as_ref()), + Some(&7_i64) + ); + } + + #[test] + fn test_lookup_with_scope_suffix_index_sees_keys_added_after_build() { + let mut ctx = TypeCheckEvalContext::new(); + ctx.add_integer("seed.nX", 1_i64); + ctx.build_suffix_index(); + ctx.add_integer("fresh.nXi", 0_i64); + + assert_eq!( + lookup_with_scope("nXi", "", &ctx.integers, ctx.suffix_index.as_ref()), + Some(&0_i64) + ); } #[test] fn test_infer_dims_component_ref_dotted_does_not_leaf_fallback() { let mut ctx = TypeCheckEvalContext::new(); ctx.add_dimensions("sys.arr", vec![7]); + ctx.build_suffix_index(); let expr = make_dotted_comp_ref("Medium.arr"); assert_eq!( diff --git a/crates/rumoca-eval-ast/src/eval/mod.rs b/crates/rumoca-eval-ast/src/eval/mod.rs index dd7c94784..4488cadeb 100644 --- a/crates/rumoca-eval-ast/src/eval/mod.rs +++ b/crates/rumoca-eval-ast/src/eval/mod.rs @@ -1,5 +1,9 @@ //! Compile-time constant evaluation for typecheck-time AST expressions. //! +//! SPEC_0021 file-size exception: AST evaluation still shares lookup, +//! dimension, enum, and expression semantics in one module. split plan: move +//! scoped lookup and suffix indexing into a lookup submodule. +//! //! The typecheck phase needs early evaluation for: //! - structural parameters and dimensions (MLS §10, §18) //! - enum/integer/boolean conditions in guarded expressions @@ -11,7 +15,7 @@ use rumoca_core::{ eval_integer_binary as eval_common_integer_binary, eval_integer_div_builtin, }; use rumoca_ir_ast::{ClassDef, Expression, Statement, StatementBlock, Subscript, TerminalType}; -use rustc_hash::FxHashMap; +use rustc_hash::{FxHashMap, FxHashSet}; use std::borrow::Cow; use std::cell::RefCell; use std::collections::HashSet; @@ -43,31 +47,236 @@ fn lookup_by_scope<'a, T>(name: &str, scope: &str, map: &'a FxHashMap if let Some(val) = try_scope(scope) { return Some(val); } - if let Some(val) = - rumoca_core::find_map_top_level_splits_rev(scope, |base, _suffix| try_scope(base)) - { - return Some(val); + let mut dot_positions = Vec::new(); + collect_top_level_dot_positions(scope, &mut dot_positions); + for dot_idx in dot_positions.into_iter().rev() { + if let Some(val) = try_scope(&scope[..dot_idx]) { + return Some(val); + } } map.get(name) } /// General lookup for constant/scalar evaluation. +/// +/// Unlike `lookup_structural_with_scope`, this permits guarded leaf fallback +/// so package constants referenced through aliases still resolve. fn lookup_with_scope<'a, T: PartialEq>( name: &str, scope: &str, map: &'a FxHashMap, + suffix_index: Option<&SuffixIndex>, ) -> Option<&'a T> { - lookup_by_scope(name, scope, map) + if let Some(val) = lookup_by_scope(name, scope, map) { + return Some(val); + } + if !scope.is_empty() { + return None; + } + + // Suffix fallback for package constants (MLS §7.3). + // For bare names like "nX", look for any entry ending in ".nX". + // For dotted names like "Medium.nX", prefer full dotted suffix + // ".Medium.nX". If no full-dotted match exists, allow a guarded leaf + // fallback only when the leaf suffix resolves to exactly one key. + if has_any_top_level_dot(name) { + match lookup_by_suffix_state(name, map, suffix_index) { + SuffixLookup::Found(val) => return Some(val), + SuffixLookup::Ambiguous => return None, + SuffixLookup::Missing => {} + } + return lookup_by_suffix_unique_key(last_top_level_segment(name), map, suffix_index); + } + + lookup_by_suffix(name, map, suffix_index) } /// Structural lookup used by shape inference and strict dimension resolution. +/// +/// Dotted names intentionally avoid leaf fallback to prevent cross-scope +/// accidental matches (for example `Medium.nX` resolving to unrelated `*.nX`). fn lookup_structural_with_scope<'a, T: PartialEq>( name: &str, scope: &str, map: &'a FxHashMap, + suffix_index: Option<&SuffixIndex>, ) -> Option<&'a T> { - lookup_by_scope(name, scope, map) + if let Some(value) = lookup_by_scope(name, scope, map) { + return Some(value); + } + if !scope.is_empty() { + return None; + } + + if has_any_top_level_dot(name) { + return match lookup_by_suffix_state(name, map, suffix_index) { + SuffixLookup::Found(value) => Some(value), + SuffixLookup::Missing | SuffixLookup::Ambiguous => None, + }; + } + + lookup_by_suffix(name, map, suffix_index) +} + +fn find_unique_value<'a, T: PartialEq>( + candidates: &[usize], + map: &'a FxHashMap, + suffix_index: &SuffixIndex, +) -> SuffixLookup<'a, T> { + let mut found: Option<&T> = None; + for candidate_idx in candidates { + let key = suffix_index.key(*candidate_idx); + if let Some(val) = map.get(key) { + if found.is_some_and(|prev| prev != val) { + return SuffixLookup::Ambiguous; + } + found = Some(val); + } + } + found.map_or(SuffixLookup::Missing, SuffixLookup::Found) +} + +enum SuffixLookup<'a, T> { + Found(&'a T), + Missing, + Ambiguous, +} + +fn lookup_by_suffix_state<'a, T: PartialEq>( + name: &str, + map: &'a FxHashMap, + suffix_index: Option<&SuffixIndex>, +) -> SuffixLookup<'a, T> { + if has_any_top_level_dot(name) { + if let Some(index) = suffix_index { + let Some(candidates) = index.keys_by_dotted_suffix.get(name) else { + return SuffixLookup::Missing; + }; + return find_unique_value(candidates, map, index); + } + + let mut found: Option<&T> = None; + for (key, val) in map { + if !has_top_level_suffix_match(key, name) { + continue; + } + if found.is_some_and(|prev| prev != val) { + return SuffixLookup::Ambiguous; + } + found = Some(val); + } + return found.map_or(SuffixLookup::Missing, SuffixLookup::Found); + } + + if let Some(index) = suffix_index { + let Some(candidates) = index.keys_by_suffix.get(name) else { + return SuffixLookup::Missing; + }; + find_unique_value(candidates, map, index) + } else { + let mut found: Option<&T> = None; + for (key, val) in map { + if !has_top_level_suffix_match(key, name) { + continue; + } + if found.is_some_and(|prev| prev != val) { + return SuffixLookup::Ambiguous; + } + found = Some(val); + } + found.map_or(SuffixLookup::Missing, SuffixLookup::Found) + } +} + +fn lookup_by_suffix<'a, T: PartialEq>( + name: &str, + map: &'a FxHashMap, + suffix_index: Option<&SuffixIndex>, +) -> Option<&'a T> { + match lookup_by_suffix_state(name, map, suffix_index) { + SuffixLookup::Found(val) => Some(val), + SuffixLookup::Missing | SuffixLookup::Ambiguous => None, + } +} + +/// Dotted lookup leaf fallback that requires a unique target-map key. +fn lookup_by_suffix_unique_key<'a, T>( + name: &str, + map: &'a FxHashMap, + suffix_index: Option<&SuffixIndex>, +) -> Option<&'a T> { + if let Some(index) = suffix_index { + let candidates = index.keys_by_suffix.get(name)?; + let mut matched_idx: Option = None; + for candidate_idx in candidates { + let key = index.key(*candidate_idx); + if !map.contains_key(key) { + continue; + } + if matched_idx.replace(*candidate_idx).is_some() { + return None; + } + } + return matched_idx.and_then(|idx| map.get(index.key(idx))); + } + + let mut matched_key: Option<&str> = None; + for key in map.keys() { + if !has_top_level_suffix_match(key, name) { + continue; + } + if matched_key.replace(key.as_str()).is_some() { + return None; + } + } + matched_key.and_then(|key| map.get(key)) +} + +fn has_top_level_suffix_match(key: &str, suffix: &str) -> bool { + if key.len() <= suffix.len() || !key.ends_with(suffix) { + return false; + } + + let boundary_idx = key.len() - suffix.len() - 1; + matches!(key.as_bytes().get(boundary_idx), Some(b'.')) && is_top_level_dot_at(key, boundary_idx) +} + +fn has_any_top_level_dot(path: &str) -> bool { + let mut dots = Vec::new(); + collect_top_level_dot_positions(path, &mut dots); + !dots.is_empty() +} + +fn last_top_level_segment(path: &str) -> &str { + let mut dots = Vec::new(); + collect_top_level_dot_positions(path, &mut dots); + dots.last().map_or(path, |dot_idx| &path[dot_idx + 1..]) +} + +fn collect_top_level_dot_positions(path: &str, out: &mut Vec) { + let mut bracket_depth = 0usize; + for (idx, byte) in path.bytes().enumerate() { + match byte { + b'[' => bracket_depth += 1, + b']' => bracket_depth = bracket_depth.saturating_sub(1), + b'.' if bracket_depth == 0 => out.push(idx), + _ => {} + } + } +} + +fn is_top_level_dot_at(path: &str, dot_index: usize) -> bool { + let mut bracket_depth = 0usize; + for (idx, byte) in path.bytes().enumerate().take(dot_index + 1) { + match byte { + b'[' => bracket_depth += 1, + b']' => bracket_depth = bracket_depth.saturating_sub(1), + b'.' if idx == dot_index => return bracket_depth == 0, + _ => {} + } + } + false } fn component_reference_path(cr: &rumoca_ir_ast::ComponentReference) -> Cow<'_, str> { @@ -97,10 +306,88 @@ pub struct TypeCheckEvalContext { pub func_eval_depth: usize, pub enum_sizes: FxHashMap, pub enum_ordinals: FxHashMap, + suffix_index: Option, + suffix_index_fingerprint: Option, warning_keys: RefCell>, warnings: RefCell>, } +struct SuffixIndex { + keys: Vec, + key_set: FxHashSet, + keys_by_suffix: FxHashMap>, + keys_by_dotted_suffix: FxHashMap>, +} + +#[derive(Clone, Copy, PartialEq, Eq)] +struct SuffixIndexFingerprint { + integers: usize, + reals: usize, + booleans: usize, + enums: usize, + dimensions: usize, + functions: usize, + enum_ordinals: usize, + enum_sizes: usize, +} + +impl SuffixIndex { + fn key(&self, idx: usize) -> &str { + self.keys[idx].as_str() + } + + fn insert_key(&mut self, key: String) { + if !self.key_set.insert(key.clone()) { + return; + } + let key_idx = self.keys.len(); + index_key_suffixes( + &key, + key_idx, + &mut self.keys_by_suffix, + &mut self.keys_by_dotted_suffix, + ); + self.keys.push(key); + } +} + +fn index_key_suffixes( + key: &str, + key_idx: usize, + keys_by_suffix: &mut FxHashMap>, + keys_by_dotted_suffix: &mut FxHashMap>, +) { + let mut bracket_depth = 0usize; + let mut top_level_dots: Vec = Vec::new(); + for (idx, byte) in key.bytes().enumerate() { + match byte { + b'[' => bracket_depth += 1, + b']' => bracket_depth = bracket_depth.saturating_sub(1), + b'.' if bracket_depth == 0 => top_level_dots.push(idx), + _ => {} + } + } + + let Some(&last_dot) = top_level_dots.last() else { + return; + }; + + keys_by_suffix + .entry(key[last_dot + 1..].to_string()) + .or_default() + .push(key_idx); + + for dot_idx in top_level_dots + .iter() + .take(top_level_dots.len().saturating_sub(1)) + { + keys_by_dotted_suffix + .entry(key[*dot_idx + 1..].to_string()) + .or_default() + .push(key_idx); + } +} + impl Default for TypeCheckEvalContext { fn default() -> Self { Self::new() @@ -120,23 +407,119 @@ impl TypeCheckEvalContext { func_eval_depth: 0, enum_sizes: FxHashMap::default(), enum_ordinals: FxHashMap::default(), + suffix_index: None, + suffix_index_fingerprint: None, warning_keys: RefCell::new(HashSet::default()), warnings: RefCell::new(Vec::new()), } } + pub fn build_suffix_index(&mut self) { + let fingerprint = self.suffix_index_fingerprint(); + if self.suffix_index.is_some() && self.suffix_index_fingerprint == Some(fingerprint) { + return; + } + + let mut keys_by_suffix: FxHashMap> = FxHashMap::default(); + let mut keys_by_dotted_suffix: FxHashMap> = FxHashMap::default(); + let mut seen: FxHashSet = FxHashSet::default(); + let mut keys: Vec = Vec::new(); + + for key in self + .integers + .keys() + .chain(self.reals.keys()) + .chain(self.booleans.keys()) + .chain(self.enums.keys()) + .chain(self.dimensions.keys()) + .chain(self.functions.keys()) + .chain(self.enum_ordinals.keys()) + .chain(self.enum_sizes.keys()) + { + let owned = key.clone(); + if seen.insert(owned.clone()) { + keys.push(owned); + } + } + + for (key_idx, key) in keys.iter().enumerate() { + index_key_suffixes( + key, + key_idx, + &mut keys_by_suffix, + &mut keys_by_dotted_suffix, + ); + } + + self.suffix_index = Some(SuffixIndex { + keys, + key_set: seen, + keys_by_suffix, + keys_by_dotted_suffix, + }); + self.suffix_index_fingerprint = Some(fingerprint); + } + + fn suffix_index_fingerprint(&self) -> SuffixIndexFingerprint { + SuffixIndexFingerprint { + integers: self.integers.len(), + reals: self.reals.len(), + booleans: self.booleans.len(), + enums: self.enums.len(), + dimensions: self.dimensions.len(), + functions: self.functions.len(), + enum_ordinals: self.enum_ordinals.len(), + enum_sizes: self.enum_sizes.len(), + } + } + pub fn add_integer(&mut self, name: impl Into, value: i64) { let name = name.into(); + self.add_suffix_index_key(name.as_str()); self.integers.insert(name.clone(), value); self.scalar_spans.remove(&name); } + pub fn add_integer_if_absent(&mut self, name: impl Into, value: i64) { + let name = name.into(); + if !self.integers.contains_key(&name) { + self.add_integer(name, value); + } + } + pub fn add_real(&mut self, name: impl Into, value: f64) { let name = name.into(); + self.add_suffix_index_key(name.as_str()); self.reals.insert(name.clone(), value); self.scalar_spans.remove(&name); } + pub fn add_real_if_absent(&mut self, name: impl Into, value: f64) { + let name = name.into(); + if !self.reals.contains_key(&name) { + self.add_real(name, value); + } + } + + pub fn add_boolean(&mut self, name: impl Into, value: bool) { + let name = name.into(); + self.add_suffix_index_key(name.as_str()); + self.booleans.insert(name, value); + } + + pub fn add_boolean_if_absent(&mut self, name: impl Into, value: bool) { + let name = name.into(); + if !self.booleans.contains_key(&name) { + self.add_boolean(name, value); + } + } + + pub fn add_enum(&mut self, name: impl Into, value: impl Into) { + let name = name.into(); + self.add_suffix_index_key(name.as_str()); + self.enums.insert(name, value.into()); + } + fn remember_scalar_span(&mut self, name: &str, span: Span) { if !span.is_dummy() { self.scalar_spans.insert(name.to_string(), span); @@ -150,7 +533,41 @@ impl TypeCheckEvalContext { } pub fn add_dimensions(&mut self, name: impl Into, dims: Vec) { - self.dimensions.insert(name.into(), dims); + let name = name.into(); + self.add_suffix_index_key(name.as_str()); + self.dimensions.insert(name, dims); + } + + pub fn add_enum_size(&mut self, name: impl Into, size: usize) { + let name = name.into(); + self.add_suffix_index_key(name.as_str()); + self.enum_sizes.insert(name, size); + } + + pub fn add_enum_size_if_absent(&mut self, name: impl Into, size: usize) { + let name = name.into(); + if !self.enum_sizes.contains_key(&name) { + self.add_enum_size(name, size); + } + } + + pub fn add_enum_ordinal(&mut self, name: impl Into, ordinal: i64) { + let name = name.into(); + self.add_suffix_index_key(name.as_str()); + self.enum_ordinals.insert(name, ordinal); + } + + pub fn add_enum_ordinal_if_absent(&mut self, name: impl Into, ordinal: i64) { + let name = name.into(); + if !self.enum_ordinals.contains_key(&name) { + self.add_enum_ordinal(name, ordinal); + } + } + + fn add_suffix_index_key(&mut self, name: &str) { + if let Some(index) = &mut self.suffix_index { + index.insert_key(name.to_string()); + } } pub fn get_integer(&self, name: &str) -> Option { @@ -741,7 +1158,12 @@ fn lookup_dims_with_scope<'a>( ctx: &'a TypeCheckEvalContext, scope: &str, ) -> Option<&'a Vec> { - lookup_with_scope(array_name, scope, &ctx.dimensions) + lookup_with_scope( + array_name, + scope, + &ctx.dimensions, + ctx.suffix_index.as_ref(), + ) } /// Flatten a matrix row element into integer values. @@ -924,7 +1346,7 @@ fn lookup_boolean_with_scope( ctx: &TypeCheckEvalContext, scope: &str, ) -> Option { - lookup_with_scope(ref_path, scope, &ctx.booleans).copied() + lookup_with_scope(ref_path, scope, &ctx.booleans, ctx.suffix_index.as_ref()).copied() } /// Evaluate a numeric comparison (integer then real) with scope-aware lookup. @@ -1080,7 +1502,7 @@ fn lookup_enum_with_scope<'a>( ctx: &'a TypeCheckEvalContext, scope: &str, ) -> Option<&'a str> { - lookup_with_scope(ref_path, scope, &ctx.enums).map(|s| s.as_str()) + lookup_with_scope(ref_path, scope, &ctx.enums, ctx.suffix_index.as_ref()).map(|s| s.as_str()) } /// Compare two enumeration values using suffix matching. @@ -1201,7 +1623,7 @@ fn eval_enum_dimension_with_scope( ) -> Option { if let Expression::ComponentReference(cr) = expr { let ref_path = component_reference_path(cr); - lookup_with_scope(&ref_path, scope, &ctx.enum_sizes).copied() + lookup_with_scope(&ref_path, scope, &ctx.enum_sizes, ctx.suffix_index.as_ref()).copied() } else { None } @@ -1236,9 +1658,17 @@ pub fn eval_integer_with_scope( Expression::ComponentReference(cr) if !cr.parts.is_empty() => { let ref_path = component_reference_path(cr); - lookup_with_scope(&ref_path, scope, &ctx.integers) + lookup_with_scope(&ref_path, scope, &ctx.integers, ctx.suffix_index.as_ref()) .copied() - .or_else(|| lookup_with_scope(&ref_path, scope, &ctx.enum_ordinals).copied()) + .or_else(|| { + lookup_with_scope( + &ref_path, + scope, + &ctx.enum_ordinals, + ctx.suffix_index.as_ref(), + ) + .copied() + }) } Expression::ComponentReference(_) => None, diff --git a/crates/rumoca-eval-ast/src/eval_instantiate/component_params.rs b/crates/rumoca-eval-ast/src/eval_instantiate/component_params.rs index f01f5564d..674a2864c 100644 --- a/crates/rumoca-eval-ast/src/eval_instantiate/component_params.rs +++ b/crates/rumoca-eval-ast/src/eval_instantiate/component_params.rs @@ -1,7 +1,7 @@ use super::{ ConditionEvalEnv, InstantiateEvalCtx, ast, eval_scoped_string_condition_with_depth, get_enum_value_with_depth, resolve_component_ref_expr, try_eval_bool_literal, - try_eval_integer_expr_with_depth, + try_eval_integer_expr_with_depth_and_locals, }; use rumoca_ir_ast::AstIndexMap as IndexMap; use rustc_hash::FxHashMap; @@ -247,28 +247,62 @@ pub fn extract_bool_params_with_mods( /// Returns a map of component names to their integer values. /// Takes the modification environment to check for parameter overrides. pub fn extract_int_params_with_mods(ctx: &InstantiateEvalCtx) -> FxHashMap { + extract_int_params_with_mods_and_known(ctx, &FxHashMap::default()) +} + +pub fn extract_int_params_with_mods_and_known( + ctx: &InstantiateEvalCtx, + known_int_params: &FxHashMap, +) -> FxHashMap { let InstantiateEvalCtx { tree, mod_env, effective_components, resolve_class_components, } = ctx; - let mut int_params = extract_params_with_mods(effective_components, mod_env, |comp, expr| { + let mut int_params = FxHashMap::default(); + let mut eval_locals = known_dotted_integer_params(known_int_params); + + for (name, comp) in *effective_components { if !matches!( comp.variability, rumoca_core::Variability::Parameter(_) | rumoca_core::Variability::Constant(_) ) { - return None; + continue; } - try_eval_integer_expr_with_depth( - expr, - mod_env, - effective_components, - tree, - *resolve_class_components, - 0, - ) - }); + + let mod_path = ast::QualifiedName::from_ident(name); + if let Some(mod_value) = mod_env.get(&mod_path) + && let Some(value) = try_eval_integer_expr_with_depth_and_locals( + &mod_value.value, + mod_env, + effective_components, + tree, + *resolve_class_components, + 0, + Some(&eval_locals), + ) + { + int_params.insert(name.clone(), value); + eval_locals.insert(name.clone(), value); + continue; + } + + if let Some(value_expr) = component_expr_for_structural_eval(comp) + && let Some(value) = try_eval_integer_expr_with_depth_and_locals( + value_expr, + mod_env, + effective_components, + tree, + *resolve_class_components, + 0, + Some(&eval_locals), + ) + { + int_params.insert(name.clone(), value); + eval_locals.insert(name.clone(), value); + } + } // Also add dotted keys from multi-part modifications in mod_env. // This handles record field references like cellData.nRC used in for-loop ranges. @@ -283,15 +317,24 @@ pub fn extract_int_params_with_mods(ctx: &InstantiateEvalCtx) -> FxHashMap>() + .join("."), + value, + ); } } } @@ -304,6 +347,16 @@ pub fn extract_int_params_with_mods(ctx: &InstantiateEvalCtx) -> FxHashMap, +) -> FxHashMap { + known_int_params + .iter() + .filter(|(key, _)| key.contains('.')) + .map(|(key, value)| (key.clone(), *value)) + .collect() +} + fn extract_params_with_mods( effective_components: &IndexMap, mod_env: &ast::ModificationEnvironment, diff --git a/crates/rumoca-eval-ast/src/eval_instantiate/function_eval.rs b/crates/rumoca-eval-ast/src/eval_instantiate/function_eval.rs index cc67316c4..32f04223d 100644 --- a/crates/rumoca-eval-ast/src/eval_instantiate/function_eval.rs +++ b/crates/rumoca-eval-ast/src/eval_instantiate/function_eval.rs @@ -1,10 +1,20 @@ use super::{ - IntegerEvalEnv, MAX_EXPR_EVAL_DEPTH, ast, eval_integer_binary, eval_integer_function_call, + IntegerEvalEnv, MAX_EXPR_EVAL_DEPTH, ResolveClassComponents, ast, evaluate_component_condition_with_depth, try_eval_bool_expr_with_depth_and_locals, try_eval_bool_expr_with_local_values, try_eval_integer_expr_with_depth_and_locals, }; +use rumoca_core::{ + IntegerBinaryOperator, eval_integer_binary as eval_common_integer_binary, + eval_integer_div_builtin, +}; use rustc_hash::FxHashMap; +mod array_indices; +mod function_lookup; + +pub use array_indices::generate_array_indices; +pub(super) use function_lookup::lookup_function_definition; + const MAX_FUNCTION_LOOP_ITERATIONS: usize = 4096; enum LocalValue { @@ -268,6 +278,113 @@ pub(super) fn eval_user_defined_bool_function( locals.bools.get(&output_name).copied() } +pub(super) fn eval_integer_binary(op: &rumoca_core::OpBinary, lhs: i64, rhs: i64) -> Option { + let operator = match op { + rumoca_core::OpBinary::Add => IntegerBinaryOperator::Add, + rumoca_core::OpBinary::Sub => IntegerBinaryOperator::Sub, + rumoca_core::OpBinary::Mul => IntegerBinaryOperator::Mul, + rumoca_core::OpBinary::Div => IntegerBinaryOperator::Div, + _ => return None, + }; + eval_common_integer_binary(operator, lhs, rhs) +} + +pub(super) fn eval_integer_function_call( + comp: &ast::ComponentReference, + args: &[ast::Expression], + env: IntegerEvalEnv<'_>, + depth: usize, + local_ints: Option<&FxHashMap>, +) -> Option { + let func_name = comp + .parts + .iter() + .map(|p| p.ident.text.as_ref()) + .collect::>() + .join("."); + + let qualified_name = comp + .def_id + .and_then(|did| env.tree.def_map.get(&did)) + .cloned(); + + if let Some(value) = eval_integer_builtin(func_name.as_str(), args, env, depth, local_ints) { + return Some(value); + } + + let function_def = lookup_function_definition(&func_name, qualified_name.as_deref(), env.tree)?; + eval_user_defined_integer_function(function_def, args, env, depth, local_ints) +} + +fn eval_integer_builtin( + func_name: &str, + args: &[ast::Expression], + env: IntegerEvalEnv<'_>, + depth: usize, + local_ints: Option<&FxHashMap>, +) -> Option { + let recurse = |e| { + try_eval_integer_expr_with_depth_and_locals( + e, + env.mod_env, + env.effective_components, + env.tree, + env.resolve_class_components, + depth + 1, + local_ints, + ) + }; + + match func_name { + "integer" => recurse(args.first()?), + "mod" => { + let x = recurse(args.first()?)?; + let y = recurse(args.get(1)?)?; + (y != 0).then_some(((x % y) + y) % y) + } + "div" => eval_integer_div_builtin(recurse(args.first()?)?, recurse(args.get(1)?)?), + "abs" => Some(recurse(args.first()?)?.abs()), + "min" => Some(recurse(args.first()?)?.min(recurse(args.get(1)?)?)), + "max" => Some(recurse(args.first()?)?.max(recurse(args.get(1)?)?)), + _ => None, + } +} + +pub(super) fn eval_user_defined_integer_function( + function_def: &ast::ClassDef, + args: &[ast::Expression], + env: IntegerEvalEnv<'_>, + depth: usize, + caller_locals: Option<&FxHashMap>, +) -> Option { + if !function_def.pure || function_def.external.is_some() { + return None; + } + if depth >= MAX_EXPR_EVAL_DEPTH { + return None; + } + + let mut local_values = FxHashMap::default(); + bind_function_inputs( + function_def, + args, + env, + depth + 1, + caller_locals, + &mut local_values, + )?; + initialize_function_locals(function_def, env, depth + 1, &mut local_values); + let output_name = find_scalar_function_output_name(function_def)?; + + for algorithm in &function_def.algorithms { + if interpret_function_statements(algorithm, env, depth + 1, &mut local_values)? { + break; + } + } + + local_values.get(&output_name).copied() +} + fn bind_mixed_function_inputs( function_def: &ast::ClassDef, args: &[ast::Expression], @@ -668,10 +785,7 @@ pub fn evaluate_array_dimensions( mod_env: &ast::ModificationEnvironment, effective_components: &ast::AstIndexMap, tree: &ast::ClassTree, - resolve_class_components: fn( - &ast::ClassTree, - &ast::ClassDef, - ) -> ast::AstIndexMap, + resolve_class_components: ResolveClassComponents, ) -> Option> { // Prefer shape_expr because it reflects active modifications. // Fall back to precomputed shape only if expression evaluation fails. @@ -700,10 +814,7 @@ fn eval_shape_expr( mod_env: &ast::ModificationEnvironment, effective_components: &ast::AstIndexMap, tree: &ast::ClassTree, - resolve_class_components: fn( - &ast::ClassTree, - &ast::ClassDef, - ) -> ast::AstIndexMap, + resolve_class_components: ResolveClassComponents, ) -> Option> { let mut dims = Vec::with_capacity(shape_expr.len()); for sub in shape_expr { @@ -738,10 +849,7 @@ pub fn try_eval_integer_shape_expr( mod_env: &ast::ModificationEnvironment, effective_components: &ast::AstIndexMap, tree: &ast::ClassTree, - resolve_class_components: fn( - &ast::ClassTree, - &ast::ClassDef, - ) -> ast::AstIndexMap, + resolve_class_components: ResolveClassComponents, ) -> Option { try_eval_integer_shape_expr_with_depth( expr, @@ -758,10 +866,7 @@ fn try_eval_integer_shape_expr_with_depth( mod_env: &ast::ModificationEnvironment, effective_components: &ast::AstIndexMap, tree: &ast::ClassTree, - resolve_class_components: fn( - &ast::ClassTree, - &ast::ClassDef, - ) -> ast::AstIndexMap, + resolve_class_components: ResolveClassComponents, depth: usize, ) -> Option { if depth > MAX_EXPR_EVAL_DEPTH { @@ -848,10 +953,7 @@ fn eval_integer_shape_component_ref( mod_env: &ast::ModificationEnvironment, effective_components: &ast::AstIndexMap, tree: &ast::ClassTree, - resolve_class_components: fn( - &ast::ClassTree, - &ast::ClassDef, - ) -> ast::AstIndexMap, + resolve_class_components: ResolveClassComponents, depth: usize, ) -> Option { if depth > MAX_EXPR_EVAL_DEPTH { @@ -991,48 +1093,16 @@ fn shape_component_ref_is_static( true } -/// Generate all array indices for multi-dimensional arrays. -/// For dims = `[2, 3]`, generates: `[[1,1], [1,2], [1,3], [2,1], [2,2], [2,3]]`. -/// Uses 1-based indexing per Modelica semantics (MLS §10.1). -pub fn generate_array_indices(dims: &[i64]) -> Vec> { - if dims.is_empty() { - return vec![]; // Scalar, no indices needed - } - - let total: usize = dims.iter().map(|&d| d as usize).product(); - let mut result = Vec::with_capacity(total); - - // Generate all combinations using iterative approach - let mut indices = vec![1i64; dims.len()]; - loop { - result.push(indices.clone()); - - // Increment indices from right to left (like counting) - let mut i = dims.len(); - while i > 0 { - i -= 1; - indices[i] += 1; - if indices[i] <= dims[i] { - break; - } - // Carry over - if i == 0 { - return result; // All combinations generated - } - indices[i] = 1; - } - } -} - #[cfg(test)] mod tests { use std::sync::Arc; use super::super::{ - InstantiateEvalCtx, enum_values_equal, eval_integer_binary, evaluate_component_condition, - try_eval_integer_expr, + InstantiateEvalCtx, enum_values_equal, evaluate_component_condition, try_eval_integer_expr, }; - use super::evaluate_array_dimensions; + use super::function_lookup::lookup_unique_short_function_name; + use super::{eval_integer_binary, evaluate_array_dimensions}; + use rumoca_core::DefId; use rumoca_ir_ast as ast; use rumoca_ir_ast::AstIndexMap as IndexMap; @@ -1854,4 +1924,77 @@ mod tests { assert_eq!(try_eval_integer_expr(&ctx, &expr), Some(2)); } + + fn function_class(name: &str, def_id: DefId) -> ast::ClassDef { + ast::ClassDef { + def_id: Some(def_id), + name: token(name), + class_type: rumoca_core::ClassType::Function, + pure: true, + ..Default::default() + } + } + + fn package_with_function( + package_name: &str, + package_id: DefId, + function_id: DefId, + ) -> ast::ClassDef { + let mut package = ast::ClassDef { + def_id: Some(package_id), + name: token(package_name), + class_type: rumoca_core::ClassType::Package, + ..Default::default() + }; + package + .classes + .insert("leaf".to_string(), function_class("leaf", function_id)); + package + } + + #[test] + fn short_function_lookup_uses_unique_leaf_index() { + let package_id = DefId::new(1); + let function_id = DefId::new(2); + let mut tree = ast::ClassTree::new(); + tree.definitions.classes.insert( + "Pkg".to_string(), + package_with_function("Pkg", package_id, function_id), + ); + tree.def_map.insert(package_id, "Pkg".to_string()); + tree.def_map.insert(function_id, "Pkg.leaf".to_string()); + tree.name_map.insert("Pkg".to_string(), package_id); + tree.name_map.insert("Pkg.leaf".to_string(), function_id); + + let found = lookup_unique_short_function_name("leaf", &tree); + + assert_eq!(found.and_then(|class| class.def_id), Some(function_id)); + } + + #[test] + fn short_function_lookup_rejects_ambiguous_leaf() { + let package_a_id = DefId::new(10); + let function_a_id = DefId::new(11); + let package_b_id = DefId::new(20); + let function_b_id = DefId::new(21); + let mut tree = ast::ClassTree::new(); + tree.definitions.classes.insert( + "PkgA".to_string(), + package_with_function("PkgA", package_a_id, function_a_id), + ); + tree.definitions.classes.insert( + "PkgB".to_string(), + package_with_function("PkgB", package_b_id, function_b_id), + ); + tree.def_map.insert(package_a_id, "PkgA".to_string()); + tree.def_map.insert(function_a_id, "PkgA.leaf".to_string()); + tree.def_map.insert(package_b_id, "PkgB".to_string()); + tree.def_map.insert(function_b_id, "PkgB.leaf".to_string()); + tree.name_map.insert("PkgA".to_string(), package_a_id); + tree.name_map.insert("PkgA.leaf".to_string(), function_a_id); + tree.name_map.insert("PkgB".to_string(), package_b_id); + tree.name_map.insert("PkgB.leaf".to_string(), function_b_id); + + assert!(lookup_unique_short_function_name("leaf", &tree).is_none()); + } } diff --git a/crates/rumoca-eval-ast/src/eval_instantiate/function_eval/array_indices.rs b/crates/rumoca-eval-ast/src/eval_instantiate/function_eval/array_indices.rs new file mode 100644 index 000000000..9692ae1f9 --- /dev/null +++ b/crates/rumoca-eval-ast/src/eval_instantiate/function_eval/array_indices.rs @@ -0,0 +1,28 @@ +/// Generate all array indices for multi-dimensional arrays. +/// For dims = `[2, 3]`, generates: `[[1,1], [1,2], [1,3], [2,1], [2,2], [2,3]]`. +/// Uses 1-based indexing per Modelica semantics. +pub fn generate_array_indices(dims: &[i64]) -> Vec> { + if dims.is_empty() { + return vec![]; + } + + let total: usize = dims.iter().map(|&d| d as usize).product(); + let mut result = Vec::with_capacity(total); + let mut indices = vec![1i64; dims.len()]; + loop { + result.push(indices.clone()); + + let mut i = dims.len(); + while i > 0 { + i -= 1; + indices[i] += 1; + if indices[i] <= dims[i] { + break; + } + if i == 0 { + return result; + } + indices[i] = 1; + } + } +} diff --git a/crates/rumoca-eval-ast/src/eval_instantiate/function_eval/function_lookup.rs b/crates/rumoca-eval-ast/src/eval_instantiate/function_eval/function_lookup.rs new file mode 100644 index 000000000..a1ff676f4 --- /dev/null +++ b/crates/rumoca-eval-ast/src/eval_instantiate/function_eval/function_lookup.rs @@ -0,0 +1,86 @@ +use rumoca_core::DefId; +use rumoca_ir_ast as ast; +use rustc_hash::FxHashMap; +use std::cell::RefCell; + +const MAX_FUNCTION_LEAF_INDEXES: usize = 16; + +thread_local! { + static FUNCTION_LEAF_INDEX_CACHE: RefCell>>> = + RefCell::new(FxHashMap::default()); +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +struct FunctionLeafIndexKey { + tree_addr: usize, + def_count: usize, + name_count: usize, +} + +pub(crate) fn lookup_function_definition<'a>( + func_name: &str, + qualified_name: Option<&str>, + tree: &'a ast::ClassTree, +) -> Option<&'a ast::ClassDef> { + if let Some(name) = qualified_name + && let Some(class) = tree.get_class_by_qualified_name(name) + && class.class_type == rumoca_core::ClassType::Function + { + return Some(class); + } + + if let Some(class) = tree.get_class_by_qualified_name(func_name) + && class.class_type == rumoca_core::ClassType::Function + { + return Some(class); + } + + lookup_unique_short_function_name(func_name, tree) +} + +pub(super) fn lookup_unique_short_function_name<'a>( + func_name: &str, + tree: &'a ast::ClassTree, +) -> Option<&'a ast::ClassDef> { + let def_id = cached_function_leaf_def_id(tree, func_name)??; + tree.get_class_by_def_id(def_id) + .filter(|class| class.class_type == rumoca_core::ClassType::Function) +} + +fn cached_function_leaf_def_id(tree: &ast::ClassTree, func_name: &str) -> Option> { + let key = FunctionLeafIndexKey { + tree_addr: std::ptr::from_ref(tree) as usize, + def_count: tree.def_map.len(), + name_count: tree.name_map.len(), + }; + FUNCTION_LEAF_INDEX_CACHE.with(|cache| { + let mut cache = cache.borrow_mut(); + if !cache.contains_key(&key) { + if cache.len() >= MAX_FUNCTION_LEAF_INDEXES { + cache.clear(); + } + cache.insert(key, build_function_leaf_index(tree)); + } + cache + .get(&key) + .and_then(|index| index.get(func_name).copied()) + }) +} + +fn build_function_leaf_index(tree: &ast::ClassTree) -> FxHashMap> { + let mut index = FxHashMap::default(); + for (def_id, qualified) in &tree.def_map { + let Some(class) = tree.get_class_by_qualified_name(qualified) else { + continue; + }; + if class.class_type != rumoca_core::ClassType::Function { + continue; + } + let leaf = class.name.text.as_ref(); + index + .entry(leaf.to_string()) + .and_modify(|existing| *existing = None) + .or_insert(Some(*def_id)); + } + index +} diff --git a/crates/rumoca-eval-ast/src/eval_instantiate/function_eval/local_lookup.rs b/crates/rumoca-eval-ast/src/eval_instantiate/function_eval/local_lookup.rs new file mode 100644 index 000000000..cce47e1cd --- /dev/null +++ b/crates/rumoca-eval-ast/src/eval_instantiate/function_eval/local_lookup.rs @@ -0,0 +1,48 @@ +use rumoca_ir_ast as ast; +use rustc_hash::FxHashMap; + +pub(crate) fn lookup_local_integer( + comp_ref: &ast::ComponentReference, + local_values: &FxHashMap, +) -> Option { + if comp_ref.parts.iter().any(|part| part.subs.is_some()) { + return None; + } + let dotted = comp_ref + .parts + .iter() + .map(|p| p.ident.text.as_ref()) + .collect::>() + .join("."); + if let Some(value) = local_values.get(&dotted) { + return Some(*value); + } + if comp_ref.parts.len() == 1 { + let name = comp_ref.parts[0].ident.text.as_ref(); + return local_values.get(name).copied(); + } + None +} + +pub(crate) fn lookup_local_bool( + comp_ref: &ast::ComponentReference, + local_values: &FxHashMap, +) -> Option { + if comp_ref.parts.iter().any(|part| part.subs.is_some()) { + return None; + } + let dotted = comp_ref + .parts + .iter() + .map(|p| p.ident.text.as_ref()) + .collect::>() + .join("."); + if let Some(value) = local_values.get(&dotted) { + return Some(*value); + } + if comp_ref.parts.len() == 1 { + let name = comp_ref.parts[0].ident.text.as_ref(); + return local_values.get(name).copied(); + } + None +} diff --git a/crates/rumoca-eval-ast/src/eval_instantiate/mod.rs b/crates/rumoca-eval-ast/src/eval_instantiate/mod.rs index b248bc7e5..3b4345bf4 100644 --- a/crates/rumoca-eval-ast/src/eval_instantiate/mod.rs +++ b/crates/rumoca-eval-ast/src/eval_instantiate/mod.rs @@ -14,8 +14,12 @@ use rumoca_ir_ast as ast; use rumoca_ir_ast::AstIndexMap as IndexMap; use rustc_hash::FxHashMap; +pub(super) type ResolveClassComponents = + fn(&ast::ClassTree, &ast::ClassDef) -> IndexMap; + mod component_params; mod function_eval; +mod real_eval; mod scoped_condition; pub(super) use component_params::{ @@ -25,11 +29,13 @@ pub(super) use component_params::{ pub use component_params::{ eval_state_select_expr, eval_state_select_expr_with_source_scope, expr_to_string, extract_binding, extract_bool_params_with_mods, extract_int_params_with_mods, - parse_state_select, propagate_record_alias_integer_params, try_eval_string_expr, + extract_int_params_with_mods_and_known, parse_state_select, + propagate_record_alias_integer_params, try_eval_string_expr, }; pub use function_eval::{ evaluate_array_dimensions, generate_array_indices, try_eval_integer_shape_expr, }; +use real_eval::try_eval_real_expr_with_depth_and_scope; use scoped_condition::eval_scoped_string_condition_with_depth; /// Maximum recursion depth for condition evaluation (prevents stack overflow) @@ -313,6 +319,32 @@ fn eval_param_ref( return None; } + // Type attributes may be evaluated while extracting fields of a structured + // component. In MSL media records, expressions such as + // `medium.preferredMediumStates` refer to the current structured component's + // `preferredMediumStates` parameter even after the field is being processed. + if comp_ref.parts.len() > 1 + && let Some(last) = comp_ref.parts.last() + { + let field_name = last.ident.text.as_ref(); + if let Some(sibling) = effective_components.get(field_name) { + let value_expr = component_params::component_condition_value_expr(sibling); + if let Some(val) = expr_to_bool(value_expr) { + return Some(val); + } + if let Some(val) = evaluate_component_condition_with_depth( + value_expr, + mod_env, + effective_components, + tree, + resolve_class_components, + depth + 1, + ) { + return Some(val); + } + } + } + // Look up the parameter's default value from effective components (single-part only) if comp_ref.parts.len() == 1 { let param_name = comp_ref.parts[0].ident.text.as_ref(); @@ -402,17 +434,6 @@ fn eval_binary_condition( depth + 1, ) }; - let int_eval = |e| { - try_eval_integer_expr_with_depth( - e, - env.mod_env, - env.effective_components, - env.tree, - env.resolve_class_components, - depth + 1, - ) - }; - match op { rumoca_core::OpBinary::Or => { let (l, r) = (eval(lhs), eval(rhs)); @@ -437,43 +458,77 @@ fn eval_binary_condition( if let Some(val) = enum_eq() { return Some(val); } - if let (Some(l), Some(r)) = (int_eval(lhs), int_eval(rhs)) { - return Some(l == r); - } + return eval_numeric_condition(op, lhs, rhs, env, depth); } rumoca_core::OpBinary::Neq => { if let Some(val) = enum_eq() { return Some(!val); } - if let (Some(l), Some(r)) = (int_eval(lhs), int_eval(rhs)) { - return Some(l != r); - } - } - rumoca_core::OpBinary::Lt => { - if let (Some(l), Some(r)) = (int_eval(lhs), int_eval(rhs)) { - return Some(l < r); - } + return eval_numeric_condition(op, lhs, rhs, env, depth); } - rumoca_core::OpBinary::Le => { - if let (Some(l), Some(r)) = (int_eval(lhs), int_eval(rhs)) { - return Some(l <= r); - } - } - rumoca_core::OpBinary::Gt => { - if let (Some(l), Some(r)) = (int_eval(lhs), int_eval(rhs)) { - return Some(l > r); - } - } - rumoca_core::OpBinary::Ge => { - if let (Some(l), Some(r)) = (int_eval(lhs), int_eval(rhs)) { - return Some(l >= r); - } + rumoca_core::OpBinary::Lt + | rumoca_core::OpBinary::Le + | rumoca_core::OpBinary::Gt + | rumoca_core::OpBinary::Ge => { + return eval_numeric_condition(op, lhs, rhs, env, depth); } _ => {} } None } +fn eval_numeric_condition( + op: &rumoca_core::OpBinary, + lhs: &ast::Expression, + rhs: &ast::Expression, + env: ConditionEvalEnv<'_>, + depth: usize, +) -> Option { + if let Some(value) = eval_integer_condition(op, lhs, rhs, env, depth) { + return Some(value); + } + + let lhs = try_eval_real_expr_with_depth_and_scope(lhs, env, None, depth + 1)?; + let rhs = try_eval_real_expr_with_depth_and_scope(rhs, env, None, depth + 1)?; + eval_ordered_values(op, lhs, rhs) +} + +fn eval_integer_condition( + op: &rumoca_core::OpBinary, + lhs: &ast::Expression, + rhs: &ast::Expression, + env: ConditionEvalEnv<'_>, + depth: usize, +) -> Option { + let eval = |expr| { + try_eval_integer_expr_with_depth( + expr, + env.mod_env, + env.effective_components, + env.tree, + env.resolve_class_components, + depth + 1, + ) + }; + eval_ordered_values(op, eval(lhs)?, eval(rhs)?) +} + +fn eval_ordered_values( + op: &rumoca_core::OpBinary, + lhs: T, + rhs: T, +) -> Option { + match op { + rumoca_core::OpBinary::Eq => Some(lhs == rhs), + rumoca_core::OpBinary::Neq => Some(lhs != rhs), + rumoca_core::OpBinary::Lt => Some(lhs < rhs), + rumoca_core::OpBinary::Le => Some(lhs <= rhs), + rumoca_core::OpBinary::Gt => Some(lhs > rhs), + rumoca_core::OpBinary::Ge => Some(lhs >= rhs), + _ => None, + } +} + /// Evaluate an enum equality comparison like `controllerType == SimpleController.PI`. /// /// Returns Some(true) if values are equal, Some(false) if not equal, None if cannot evaluate. @@ -684,6 +739,19 @@ fn get_enum_value_with_depth( } } +pub fn try_eval_enum_expr(ctx: &InstantiateEvalCtx<'_>, expr: &ast::Expression) -> Option { + get_enum_value_with_depth( + expr, + ctx.mod_env, + ctx.effective_components, + ctx.tree, + ctx.resolve_class_components, + None, + 0, + ) + .filter(|value| value.contains('.')) +} + fn parent_dotted_scope(path: &str) -> Option { let enclosing = rumoca_core::ComponentPath::from_flat_path(path).parent()?; (!enclosing.is_root()).then(|| enclosing.to_flat_string()) @@ -738,32 +806,33 @@ fn resolve_component_ref_expr( let dotted = component_ref_to_dotted_no_subscripts(comp_ref)?; let candidate_paths = candidate_paths_for_ref(comp_ref, dotted.as_str(), scope_prefix); - lookup_exact_component_ref(candidate_paths.as_slice(), mod_env, effective_components) - .or_else(|| { - resolve_component_ref_from_record_defaults(comp_ref, effective_components, tree) - .map(|expr| (expr, parent_dotted_scope(&dotted))) - }) - .or_else(|| { - if comp_ref.parts.len() != 1 { - return None; - } - let prefix = scope_prefix?; - let scoped_expr = resolve_scoped_record_field_expr( - prefix, - dotted.as_str(), - effective_components, - tree, - )?; - Some((scoped_expr, Some(prefix.to_string()))) - }) - .or_else(|| { - resolve_class_redeclare_field_expr(comp_ref, mod_env, tree, resolve_class_components) - .map(|expr| (expr, None)) - }) - .or_else(|| { - resolve_class_constant_binding(comp_ref, tree, resolve_class_components) - .map(|expr| (expr, None)) - }) + lookup_exact_component_ref( + candidate_paths.as_slice(), + mod_env, + effective_components, + scope_prefix, + ) + .or_else(|| { + resolve_component_ref_from_record_defaults(comp_ref, effective_components, tree) + .map(|expr| (expr, parent_dotted_scope(&dotted))) + }) + .or_else(|| { + if comp_ref.parts.len() != 1 { + return None; + } + let prefix = scope_prefix?; + let scoped_expr = + resolve_scoped_record_field_expr(prefix, dotted.as_str(), effective_components, tree)?; + Some((scoped_expr, Some(prefix.to_string()))) + }) + .or_else(|| { + resolve_class_redeclare_field_expr(comp_ref, mod_env, tree, resolve_class_components) + .map(|expr| (expr, None)) + }) + .or_else(|| { + resolve_class_constant_binding(comp_ref, tree, resolve_class_components) + .map(|expr| (expr, None)) + }) } /// Resolve a qualified reference like `P.pT_explicit` to the binding of a @@ -860,6 +929,7 @@ fn lookup_exact_component_ref( candidate_paths: &[String], mod_env: &ast::ModificationEnvironment, effective_components: &IndexMap, + scope_prefix: Option<&str>, ) -> Option<(ast::Expression, Option)> { for candidate in candidate_paths { if let Some(mod_value) = mod_env.get(&ast::QualifiedName::from_dotted(candidate)) @@ -871,12 +941,16 @@ fn lookup_exact_component_ref( .source_scope .as_ref() .map(ast::QualifiedName::to_flat_string) - .or_else(|| parent_dotted_scope(candidate)), + .or_else(|| parent_dotted_scope(candidate)) + .or_else(|| scope_prefix.map(str::to_string)), )); } if let Some(comp) = effective_components.get(candidate.as_str()) { let expr = component_expr_for_structural_eval(comp)?; - return Some((expr.clone(), parent_dotted_scope(candidate))); + return Some(( + expr.clone(), + parent_dotted_scope(candidate).or_else(|| scope_prefix.map(str::to_string)), + )); } } None diff --git a/crates/rumoca-eval-ast/src/eval_instantiate/real_eval.rs b/crates/rumoca-eval-ast/src/eval_instantiate/real_eval.rs new file mode 100644 index 000000000..07c5ba694 --- /dev/null +++ b/crates/rumoca-eval-ast/src/eval_instantiate/real_eval.rs @@ -0,0 +1,82 @@ +use super::*; + +pub(super) fn try_eval_real_expr_with_depth_and_scope( + expr: &ast::Expression, + env: ConditionEvalEnv<'_>, + scope_prefix: Option<&str>, + depth: usize, +) -> Option { + if depth > MAX_EXPR_EVAL_DEPTH { + return None; + } + + let recurse = + |expr| try_eval_real_expr_with_depth_and_scope(expr, env, scope_prefix, depth + 1); + + match expr { + ast::Expression::Terminal { + terminal_type: ast::TerminalType::UnsignedReal | ast::TerminalType::UnsignedInteger, + token, + .. + } => token.text.parse::().ok(), + ast::Expression::ComponentReference(comp_ref) => { + let (resolved_expr, next_scope) = resolve_component_ref_expr( + comp_ref, + env.mod_env, + env.effective_components, + env.tree, + env.resolve_class_components, + scope_prefix, + )?; + try_eval_real_expr_with_depth_and_scope( + &resolved_expr, + env, + next_scope.as_deref(), + depth + 1, + ) + } + ast::Expression::Binary { op, lhs, rhs, .. } => { + let l = recurse(lhs)?; + let r = recurse(rhs)?; + match op { + rumoca_core::OpBinary::Add | rumoca_core::OpBinary::AddElem => Some(l + r), + rumoca_core::OpBinary::Sub | rumoca_core::OpBinary::SubElem => Some(l - r), + rumoca_core::OpBinary::Mul | rumoca_core::OpBinary::MulElem => Some(l * r), + rumoca_core::OpBinary::Div | rumoca_core::OpBinary::DivElem => { + (r != 0.0).then_some(l / r) + } + rumoca_core::OpBinary::Exp | rumoca_core::OpBinary::ExpElem => Some(l.powf(r)), + _ => None, + } + } + ast::Expression::Unary { op, rhs, .. } => { + let r = recurse(rhs)?; + match op { + rumoca_core::OpUnary::Minus => Some(-r), + rumoca_core::OpUnary::Plus => Some(r), + _ => None, + } + } + ast::Expression::Parenthesized { inner, .. } => recurse(inner), + ast::Expression::If { + branches, + else_branch, + .. + } => { + for (condition, branch_expr) in branches { + match eval_scoped_string_condition_with_depth( + condition, + env, + scope_prefix, + depth + 1, + ) { + Some(true) => return recurse(branch_expr), + Some(false) => continue, + None => return None, + } + } + recurse(else_branch) + } + _ => None, + } +} diff --git a/crates/rumoca-eval-ast/src/eval_instantiate/scoped_condition.rs b/crates/rumoca-eval-ast/src/eval_instantiate/scoped_condition.rs index d20501e28..4a3874d51 100644 --- a/crates/rumoca-eval-ast/src/eval_instantiate/scoped_condition.rs +++ b/crates/rumoca-eval-ast/src/eval_instantiate/scoped_condition.rs @@ -119,6 +119,12 @@ fn eval_scoped_string_binary_condition( rumoca_core::OpBinary::Neq => { eval_scoped_enum_equality(lhs, rhs, env, state.scope_prefix, state.depth).map(|v| !v) } + rumoca_core::OpBinary::Lt + | rumoca_core::OpBinary::Le + | rumoca_core::OpBinary::Gt + | rumoca_core::OpBinary::Ge => { + eval_scoped_real_relation(op, lhs, rhs, env, state.scope_prefix, state.depth) + } _ => None, } } @@ -186,3 +192,22 @@ fn eval_scoped_enum_equality( )?; Some(enum_values_equal(&lhs_val, &rhs_val)) } + +fn eval_scoped_real_relation( + op: &rumoca_core::OpBinary, + lhs: &ast::Expression, + rhs: &ast::Expression, + env: ConditionEvalEnv<'_>, + scope_prefix: Option<&str>, + depth: usize, +) -> Option { + let lhs = try_eval_real_expr_with_depth_and_scope(lhs, env, scope_prefix, depth)?; + let rhs = try_eval_real_expr_with_depth_and_scope(rhs, env, scope_prefix, depth)?; + match op { + rumoca_core::OpBinary::Lt => Some(lhs < rhs), + rumoca_core::OpBinary::Le => Some(lhs <= rhs), + rumoca_core::OpBinary::Gt => Some(lhs > rhs), + rumoca_core::OpBinary::Ge => Some(lhs >= rhs), + _ => None, + } +} diff --git a/crates/rumoca-eval-dae/src/constant.rs b/crates/rumoca-eval-dae/src/constant.rs index de8d488ea..f092985c3 100644 --- a/crates/rumoca-eval-dae/src/constant.rs +++ b/crates/rumoca-eval-dae/src/constant.rs @@ -50,6 +50,18 @@ impl ConstValue { pub fn eval_const_expr_with(expr: &Expression, lookup: &F) -> Option where F: Fn(&Reference, &[Subscript]) -> Option, +{ + eval_const_expr_with_shape(expr, lookup, &|_| None) +} + +pub fn eval_const_expr_with_shape( + expr: &Expression, + lookup: &F, + shape_lookup: &S, +) -> Option +where + F: Fn(&Reference, &[Subscript]) -> Option, + S: Fn(&str) -> Option>, { match expr { Expression::Literal { @@ -87,17 +99,19 @@ where lookup(name, &merged) } Expression::Unary { op, rhs, .. } => { - let rhs = eval_const_expr_with(rhs, lookup)?; + let rhs = eval_const_expr_with_shape(rhs, lookup, shape_lookup)?; eval_unary_const(op, rhs) } Expression::Binary { op, lhs, rhs, .. } => { - let lhs = eval_const_expr_with(lhs, lookup)?; - let rhs = eval_const_expr_with(rhs, lookup)?; + let lhs = eval_const_expr_with_shape(lhs, lookup, shape_lookup)?; + let rhs = eval_const_expr_with_shape(rhs, lookup, shape_lookup)?; eval_binary_const(op, lhs, rhs) } - Expression::BuiltinCall { function, args, .. } => eval_builtin(*function, args, lookup), + Expression::BuiltinCall { function, args, .. } => { + eval_builtin(*function, args, lookup, shape_lookup) + } Expression::FunctionCall { name, args, .. } => { - eval_named_function(name.last_segment(), args, lookup) + eval_named_function(name.last_segment(), args, lookup, shape_lookup) } Expression::If { branches, @@ -105,11 +119,11 @@ where .. } => { for (cond, then_expr) in branches { - if eval_const_expr_with(cond, lookup)?.as_bool()? { - return eval_const_expr_with(then_expr, lookup); + if eval_const_expr_with_shape(cond, lookup, shape_lookup)?.as_bool()? { + return eval_const_expr_with_shape(then_expr, lookup, shape_lookup); } } - eval_const_expr_with(else_branch, lookup) + eval_const_expr_with_shape(else_branch, lookup, shape_lookup) } _ => None, } @@ -122,6 +136,18 @@ where eval_const_expr_with(expr, lookup)?.as_real() } +pub fn eval_scalar_const_expr_with_shape( + expr: &Expression, + lookup: &F, + shape_lookup: &S, +) -> Option +where + F: Fn(&Reference, &[Subscript]) -> Option, + S: Fn(&str) -> Option>, +{ + eval_const_expr_with_shape(expr, lookup, shape_lookup)?.as_real() +} + pub fn eval_scalar_const_expr(expr: &Expression, constants: &HashMap) -> Option { eval_scalar_const_expr_with(expr, &|name, subscripts| { if subscripts.is_empty() { @@ -171,25 +197,32 @@ fn eval_binary_const(op: &OpBinary, lhs: ConstValue, rhs: ConstValue) -> Option< } } -fn eval_builtin(function: BuiltinFunction, args: &[Expression], lookup: &F) -> Option +fn eval_builtin( + function: BuiltinFunction, + args: &[Expression], + lookup: &F, + shape_lookup: &S, +) -> Option where F: Fn(&Reference, &[Subscript]) -> Option, + S: Fn(&str) -> Option>, { let arg = |i: usize| { args.get(i) - .and_then(|expr| eval_scalar_const_expr_with(expr, lookup)) + .and_then(|expr| eval_scalar_const_expr_with_shape(expr, lookup, shape_lookup)) }; match function { - BuiltinFunction::NoEvent => eval_const_expr_with(args.first()?, lookup), + BuiltinFunction::NoEvent => eval_const_expr_with_shape(args.first()?, lookup, shape_lookup), BuiltinFunction::Smooth => args .get(1) - .and_then(|expr| eval_const_expr_with(expr, lookup)) + .and_then(|expr| eval_const_expr_with_shape(expr, lookup, shape_lookup)) .or_else(|| { args.first() - .and_then(|expr| eval_const_expr_with(expr, lookup)) + .and_then(|expr| eval_const_expr_with_shape(expr, lookup, shape_lookup)) }), BuiltinFunction::Integer => arg(0).map(f64::floor).map(ConstValue::Real), + BuiltinFunction::Size => eval_size(args, lookup, shape_lookup).map(ConstValue::Real), _ => arg(0) .and_then(|lhs| { apply_scalar_unary_math(function, lhs) @@ -199,52 +232,151 @@ where } } -fn eval_named_function(short_name: &str, args: &[Expression], lookup: &F) -> Option +fn eval_named_function( + short_name: &str, + args: &[Expression], + lookup: &F, + shape_lookup: &S, +) -> Option where F: Fn(&Reference, &[Subscript]) -> Option, + S: Fn(&str) -> Option>, { if short_name == "substring" { - return eval_substring(args, lookup).map(ConstValue::String); + return eval_substring(args, lookup, shape_lookup).map(ConstValue::String); } if short_name == "ln" { - return eval_builtin(BuiltinFunction::Log, args, lookup); + return eval_builtin(BuiltinFunction::Log, args, lookup, shape_lookup); } let function = BuiltinFunction::from_name(short_name)?; - eval_builtin(function, args, lookup) + eval_builtin(function, args, lookup, shape_lookup) +} + +fn eval_size(args: &[Expression], lookup: &F, shape_lookup: &S) -> Option +where + F: Fn(&Reference, &[Subscript]) -> Option, + S: Fn(&str) -> Option>, +{ + let array = args.first()?; + let dim = integer_arg(args, lookup, shape_lookup, 1)?; + let dims = expr_dims(array, shape_lookup)?; + let value = *dims.get(dim.checked_sub(1)?)?; + (value >= 0).then_some(value as f64) +} + +fn expr_dims(expr: &Expression, shape_lookup: &S) -> Option> +where + S: Fn(&str) -> Option>, +{ + match expr { + Expression::VarRef { + name, subscripts, .. + } => { + if !subscripts.is_empty() { + return None; + } + reference_shape_lookup(name, shape_lookup) + } + Expression::Index { + base, subscripts, .. + } => { + let mut dims = expr_dims(base, shape_lookup)?; + if subscripts.len() > dims.len() { + return None; + } + dims.drain(0..subscripts.len()); + Some(dims) + } + Expression::FieldAccess { base, field, .. } => { + let base = expr_path(base)?; + shape_lookup(&format!("{base}.{field}")) + } + Expression::Array { elements, .. } | Expression::Tuple { elements, .. } => { + Some(vec![i64::try_from(elements.len()).ok()?]) + } + Expression::BuiltinCall { + function: BuiltinFunction::Der, + args, + .. + } => args.first().and_then(|arg| expr_dims(arg, shape_lookup)), + _ => None, + } +} + +fn reference_shape_lookup(name: &Reference, shape_lookup: &S) -> Option> +where + S: Fn(&str) -> Option>, +{ + if let Some(component_ref) = name.component_ref() + && let Some(dims) = shape_lookup(&component_ref.to_string()) + { + return Some(dims); + } + shape_lookup(name.as_str()) +} + +fn expr_path(expr: &Expression) -> Option { + match expr { + Expression::VarRef { + name, subscripts, .. + } if subscripts.is_empty() => Some( + name.component_ref() + .map(|component_ref| component_ref.to_string()) + .unwrap_or_else(|| name.as_str().to_string()), + ), + Expression::FieldAccess { base, field, .. } => { + Some(format!("{}.{}", expr_path(base)?, field)) + } + _ => None, + } } -fn numeric_arg(args: &[Expression], lookup: &F, index: usize) -> Option +fn numeric_arg(args: &[Expression], lookup: &F, shape_lookup: &S, index: usize) -> Option where F: Fn(&Reference, &[Subscript]) -> Option, + S: Fn(&str) -> Option>, { - eval_const_expr_with(args.get(index)?, lookup)?.as_real() + eval_const_expr_with_shape(args.get(index)?, lookup, shape_lookup)?.as_real() } -fn integer_arg(args: &[Expression], lookup: &F, index: usize) -> Option +fn integer_arg( + args: &[Expression], + lookup: &F, + shape_lookup: &S, + index: usize, +) -> Option where F: Fn(&Reference, &[Subscript]) -> Option, + S: Fn(&str) -> Option>, { - let value = numeric_arg(args, lookup, index)?; + let value = numeric_arg(args, lookup, shape_lookup, index)?; (value.fract().abs() <= 1.0e-12 && value >= 1.0).then_some(value as usize) } -fn string_arg(args: &[Expression], lookup: &F, index: usize) -> Option +fn string_arg( + args: &[Expression], + lookup: &F, + shape_lookup: &S, + index: usize, +) -> Option where F: Fn(&Reference, &[Subscript]) -> Option, + S: Fn(&str) -> Option>, { - match eval_const_expr_with(args.get(index)?, lookup)? { + match eval_const_expr_with_shape(args.get(index)?, lookup, shape_lookup)? { ConstValue::String(value) => Some(value), _ => None, } } -fn eval_substring(args: &[Expression], lookup: &F) -> Option +fn eval_substring(args: &[Expression], lookup: &F, shape_lookup: &S) -> Option where F: Fn(&Reference, &[Subscript]) -> Option, + S: Fn(&str) -> Option>, { - let input = string_arg(args, lookup, 0)?; - let start = integer_arg(args, lookup, 1)?; - let stop = integer_arg(args, lookup, 2)?; + let input = string_arg(args, lookup, shape_lookup, 0)?; + let start = integer_arg(args, lookup, shape_lookup, 1)?; + let stop = integer_arg(args, lookup, shape_lookup, 2)?; if stop < start { return Some(String::new()); } @@ -272,3 +404,44 @@ fn scalar_to_bool(value: f64) -> Option { fn scalar_almost_eq(lhs: f64, rhs: f64) -> bool { (lhs - rhs).abs() <= 1.0e-12 * (1.0 + lhs.abs().max(rhs.abs())) } + +#[cfg(test)] +mod tests { + use super::*; + use rumoca_core::{Span, VarName}; + + fn span() -> Span { + Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2) + } + + fn int(value: i64) -> Expression { + Expression::Literal { + value: Literal::Integer(value), + span: span(), + } + } + + fn var(name: &str) -> Expression { + Expression::VarRef { + name: VarName::new(name).into(), + subscripts: vec![], + span: span(), + } + } + + #[test] + fn scalar_const_eval_supports_size_with_shape_lookup() { + let expr = Expression::BuiltinCall { + function: BuiltinFunction::Size, + args: vec![var("table"), int(1)], + span: span(), + }; + + let value = eval_scalar_const_expr_with_shape(&expr, &|_, _| None, &|name| match name { + "table" => Some(vec![4, 2]), + _ => None, + }); + + assert_eq!(value, Some(4.0)); + } +} diff --git a/crates/rumoca-eval-dae/src/eval/array_eval.rs b/crates/rumoca-eval-dae/src/eval/array_eval.rs index 44914835a..db206c6bc 100644 --- a/crates/rumoca-eval-dae/src/eval/array_eval.rs +++ b/crates/rumoca-eval-dae/src/eval/array_eval.rs @@ -1,8 +1,12 @@ +//! SPEC_0021 file-size exception: runtime array evaluation still shares shape +//! inference, array constructors, and indexed value extraction. split plan: +//! move runtime dimension inference and constructor-specific evaluators into +//! focused `eval::array_*` submodules as their call sites stabilize. + use super::*; use crate::eval::special::{ eval_selected_runtime_special_array_output, resolve_runtime_special_target, }; - pub(super) fn declared_dims( name: &str, env: &VarEnv, @@ -14,7 +18,6 @@ pub(super) fn declared_dims( name: format!("{name} dimensions"), }) } - pub fn eval_array_values( expr: &rumoca_core::Expression, env: &VarEnv, @@ -47,7 +50,6 @@ pub fn eval_array_values( }, ) } - pub fn eval_shaped_array_values( expr: &rumoca_core::Expression, env: &VarEnv, @@ -63,7 +65,6 @@ pub fn eval_shaped_array_values( } Ok(values) } - pub(super) fn eval_array_like_values( expr: &rumoca_core::Expression, env: &VarEnv, @@ -82,6 +83,35 @@ pub(super) fn eval_array_like_values( } Ok(vec![eval_expr(expr, env)?]) } + rumoca_core::Expression::VarRef { + name, subscripts, .. + } if subscripts_are_all_colon(subscripts) => { + array_values_from_env_name_generic(name.as_str(), env)?.ok_or_else(|| { + EvalError::MissingBinding { + name: name.to_string(), + } + }) + } + rumoca_core::Expression::VarRef { + name, subscripts, .. + } if subscripts.len() == 1 && subscript_is_range_expr(&subscripts[0]) => { + eval_var_ref_range_slice_values(name.as_str(), &subscripts[0], env) + } + rumoca_core::Expression::VarRef { + name, subscripts, .. + } if !subscripts.is_empty() => { + if let Some(values) = + eval_var_ref_subscripted_slice_values(name.as_str(), subscripts, env)? + { + return Ok(values); + } + Ok(vec![eval_expr(expr, env)?]) + } + rumoca_core::Expression::Index { + base, subscripts, .. + } if subscripts.len() == 1 && subscript_is_range_expr(&subscripts[0]) => { + eval_array_range_slice_values(base, &subscripts[0], env) + } rumoca_core::Expression::FieldAccess { base, field, .. } => { try_eval_field_access_array_values(base, field, env) } @@ -109,6 +139,185 @@ pub(super) fn eval_array_like_values( }, ) } +fn subscripts_are_all_colon(subscripts: &[rumoca_core::Subscript]) -> bool { + !subscripts.is_empty() + && subscripts + .iter() + .all(|subscript| matches!(subscript, rumoca_core::Subscript::Colon { .. })) +} + +fn subscript_is_range_expr(subscript: &rumoca_core::Subscript) -> bool { + matches!( + subscript, + rumoca_core::Subscript::Expr { expr, .. } + if matches!(expr.as_ref(), rumoca_core::Expression::Range { .. }) + ) +} + +fn eval_var_ref_subscripted_slice_values( + name: &str, + subscripts: &[rumoca_core::Subscript], + env: &VarEnv, +) -> Result>, EvalError> { + let Some(dims) = env.dims.get(name) else { + return Ok(None); + }; + if dims.is_empty() || subscripts.len() > dims.len() { + return Ok(None); + } + let Some(base_values) = array_values_from_env_name_generic(name, env)? else { + return Ok(None); + }; + let Some(dim_sizes) = dims + .iter() + .copied() + .map(|dim| usize::try_from(dim).ok().filter(|value| *value > 0)) + .collect::>>() + else { + return Ok(None); + }; + let scalar_count = dim_sizes.iter().product::(); + if scalar_count != base_values.len() { + return Ok(None); + } + + let mut choices = Vec::with_capacity(dim_sizes.len()); + for (idx, dim) in dim_sizes.iter().copied().enumerate() { + match subscripts.get(idx) { + Some(rumoca_core::Subscript::Index { value, span }) => { + choices.push(vec![positive_slice_index(*value, dim, *span)?]); + } + Some(rumoca_core::Subscript::Expr { expr, span }) => { + let value = eval_expr::(expr, env)?.real(); + choices.push(vec![finite_slice_index(value, dim, *span)?]); + } + Some(rumoca_core::Subscript::Colon { .. }) | None => { + choices.push((0..dim).collect()); + } + } + } + + let mut selected = Vec::new(); + append_subscripted_slice_values(&base_values, &dim_sizes, &choices, 0, 0, &mut selected); + Ok(Some(selected)) +} + +fn append_subscripted_slice_values( + base_values: &[T], + dims: &[usize], + choices: &[Vec], + dim_idx: usize, + flat_prefix: usize, + out: &mut Vec, +) { + if dim_idx == dims.len() { + out.push(base_values[flat_prefix]); + return; + } + let stride = dims[dim_idx + 1..].iter().product::(); + for index in &choices[dim_idx] { + append_subscripted_slice_values( + base_values, + dims, + choices, + dim_idx + 1, + flat_prefix + index * stride, + out, + ); + } +} + +fn positive_slice_index( + index: i64, + dim: usize, + span: rumoca_core::Span, +) -> Result { + usize::try_from(index) + .ok() + .filter(|value| *value >= 1 && *value <= dim) + .map(|value| value - 1) + .ok_or_else(|| { + EvalError::UnsupportedExpression { + kind: "subscript index", + } + .with_span_if_missing(span) + }) +} + +fn finite_slice_index(value: f64, dim: usize, span: rumoca_core::Span) -> Result { + if !value.is_finite() || value.fract() != 0.0 { + return Err(EvalError::UnsupportedExpression { + kind: "subscript expression", + } + .with_span_if_missing(span)); + } + positive_slice_index(value as i64, dim, span) +} + +fn eval_var_ref_range_slice_values( + name: &str, + subscript: &rumoca_core::Subscript, + env: &VarEnv, +) -> Result, EvalError> { + let values = array_values_from_env_name_generic(name, env)?.ok_or_else(|| { + EvalError::MissingBinding { + name: name.to_string(), + } + })?; + select_range_slice_values(&values, subscript, env) +} + +fn eval_array_range_slice_values( + base: &rumoca_core::Expression, + subscript: &rumoca_core::Subscript, + env: &VarEnv, +) -> Result, EvalError> { + let values = eval_array_like_values(base, env)?; + select_range_slice_values(&values, subscript, env) +} + +fn select_range_slice_values( + values: &[T], + subscript: &rumoca_core::Subscript, + env: &VarEnv, +) -> Result, EvalError> { + let indices = eval_range_subscript_indices(subscript, env)?; + indices + .into_iter() + .map(|index| { + values + .get(index.saturating_sub(1)) + .copied() + .ok_or(EvalError::UnsupportedExpression { + kind: "range slice index", + }) + }) + .collect() +} + +fn eval_range_subscript_indices( + subscript: &rumoca_core::Subscript, + env: &VarEnv, +) -> Result, EvalError> { + let rumoca_core::Subscript::Expr { expr, .. } = subscript else { + return Err(EvalError::UnsupportedExpression { + kind: "range slice", + }); + }; + eval_array_values::(expr, env)? + .into_iter() + .map(|value| finite_positive_slice_index(value.real())) + .collect() +} + +fn finite_positive_slice_index(value: f64) -> Result { + if !value.is_finite() || value < 1.0 || value.fract().abs() > f64::EPSILON { + return Err(EvalError::UnsupportedExpression { + kind: "range slice index", + }); + } + Ok(value as usize) +} fn with_expr_span( expr: &rumoca_core::Expression, @@ -158,6 +367,13 @@ fn try_eval_builtin_array_like_values( kind: "symmetric shape", }) } + rumoca_core::BuiltinFunction::Size if args.len() == 1 => { + let dims = try_infer_runtime_expr_dims(&args[0], env)?; + Ok(dims + .into_iter() + .map(|dim| T::from_f64(dim as f64)) + .collect()) + } rumoca_core::BuiltinFunction::Vector if args.len() == 1 => { eval_array_like_values(&args[0], env) } @@ -1381,13 +1597,27 @@ pub(super) fn try_infer_runtime_expr_dims( expr: &rumoca_core::Expression, env: &VarEnv, ) -> Result, EvalError> { + if let Some(dims) = try_infer_runtime_expr_dims_fast(expr, env)? { + return Ok(dims); + } + + let values = eval_array_like_values(expr, env)?; + try_infer_runtime_expr_dims_from_values(expr, values.len(), env) +} + +fn try_infer_runtime_expr_dims_fast( + expr: &rumoca_core::Expression, + env: &VarEnv, +) -> Result>, EvalError> { if let rumoca_core::Expression::VarRef { name, subscripts, .. } = expr && subscripts.is_empty() && let Some(dims) = env.dims.get(name.as_str()) { - return infer_declared_or_value_dims(dims, 0); + let value_count = array_values_from_env_name_generic(name.as_str(), env)? + .map_or(0, |values| values.len()); + return infer_declared_or_value_dims(dims, value_count).map(Some); } if let rumoca_core::Expression::BuiltinCall { function: rumoca_core::BuiltinFunction::Der, @@ -1398,7 +1628,12 @@ pub(super) fn try_infer_runtime_expr_dims( let arg = args .first() .ok_or(EvalError::UnsupportedExpression { kind: "der arity" })?; - return try_infer_runtime_expr_dims(arg, env); + return try_infer_runtime_expr_dims(arg, env).map(Some); + } + if let rumoca_core::Expression::BuiltinCall { function, args, .. } = expr + && let Some(dims) = try_infer_constructor_dims(*function, args, env)? + { + return Ok(Some(dims)); } if let rumoca_core::Expression::Array { elements, @@ -1406,59 +1641,127 @@ pub(super) fn try_infer_runtime_expr_dims( .. } = expr { - return if *is_matrix { + let dims = if *is_matrix { try_runtime_matrix_literal_dims(elements, env) } else { - Ok(runtime_vector_dims(elements.len())) - }; + Ok(vec![elements.len()]) + }?; + return Ok(Some(dims)); } if let rumoca_core::Expression::Tuple { elements, .. } = expr { - return Ok(runtime_vector_dims(elements.len())); + return Ok(Some(runtime_vector_dims(elements.len()))); } + if let rumoca_core::Expression::Range { + start, step, end, .. + } = expr + { + let values = try_eval_range_values::(start, step.as_deref(), end, env)?; + return Ok(Some(vec![values.len()])); + } + if let rumoca_core::Expression::FieldAccess { base, field, .. } = expr + && let rumoca_core::Expression::FunctionCall { name, args, .. } = base.as_ref() + && !env.functions.contains_key(name.as_str()) + && let Some(arg_index) = set_state_array_field_arg_index(name.var_name(), field) + { + let arg = args + .get(arg_index) + .ok_or(EvalError::UnsupportedExpression { + kind: "setState array field arity", + })?; + return try_infer_runtime_expr_dims(arg, env).map(Some); + } + Ok(None) +} - let values = eval_array_like_values(expr, env)?; - let dims = match expr { +fn try_infer_runtime_expr_dims_from_values( + expr: &rumoca_core::Expression, + value_count: usize, + env: &VarEnv, +) -> Result, EvalError> { + match expr { rumoca_core::Expression::VarRef { name, subscripts, .. } if subscripts.is_empty() => { if let Some(dims) = env.dims.get(name.as_str()) { - infer_declared_or_value_dims(dims, values.len())? - } else if values.len() <= 1 { - Vec::new() + infer_declared_or_value_dims(dims, value_count) + } else if value_count <= 1 { + Ok(Vec::new()) } else { - declared_dims(name.as_str(), env)? + let dims = declared_dims(name.as_str(), env)? .into_iter() .map(|dim| { usize::try_from(dim).map_err(|_| EvalError::UnsupportedExpression { kind: "array dimensions", }) }) - .collect::, _>>()? + .collect::, _>>()?; + Ok(dims) } } rumoca_core::Expression::Array { elements, is_matrix: true, .. - } => try_runtime_matrix_literal_dims(elements, env)?, + } => try_runtime_matrix_literal_dims(elements, env), rumoca_core::Expression::BuiltinCall { function: rumoca_core::BuiltinFunction::Cat, args, .. - } => try_infer_runtime_cat_dims(args, env)?, + } => try_infer_runtime_cat_dims(args, env), rumoca_core::Expression::FunctionCall { name, args, is_constructor: false, .. - } => function_call_runtime_dims(name, args, values.len(), env)?, + } => function_call_runtime_dims(name, args, value_count, env), rumoca_core::Expression::Array { .. } | rumoca_core::Expression::Tuple { .. } | rumoca_core::Expression::Range { .. } - | rumoca_core::Expression::ArrayComprehension { .. } => runtime_vector_dims(values.len()), - _ => runtime_vector_dims(values.len()), - }; - Ok(dims) + | rumoca_core::Expression::ArrayComprehension { .. } => { + Ok(runtime_vector_dims(value_count)) + } + _ => Ok(runtime_vector_dims(value_count)), + } +} + +fn try_infer_constructor_dims( + function: rumoca_core::BuiltinFunction, + args: &[rumoca_core::Expression], + env: &VarEnv, +) -> Result>, EvalError> { + match function { + rumoca_core::BuiltinFunction::Fill => { + if args.len() < 2 { + return Err(EvalError::UnsupportedExpression { + kind: "fill arguments", + }); + } + constructor_dims(&args[1..], env).map(Some) + } + rumoca_core::BuiltinFunction::Zeros | rumoca_core::BuiltinFunction::Ones => { + if args.is_empty() { + return Err(EvalError::UnsupportedExpression { + kind: "array constructor arguments", + }); + } + constructor_dims(args, env).map(Some) + } + rumoca_core::BuiltinFunction::Identity => { + let first = args.first().ok_or(EvalError::UnsupportedExpression { + kind: "identity arguments", + })?; + let n = constructor_dim(first, env)?; + Ok(Some(vec![n, n])) + } + _ => Ok(None), + } +} + +fn constructor_dims( + args: &[rumoca_core::Expression], + env: &VarEnv, +) -> Result, EvalError> { + args.iter().map(|arg| constructor_dim(arg, env)).collect() } fn checked_runtime_dims(dims: &[i64]) -> Result>, EvalError> { @@ -1528,7 +1831,6 @@ fn function_call_runtime_dims( fn runtime_vector_dims(len: usize) -> Vec { (len > 1).then_some(len).into_iter().collect() } - fn infer_declared_or_value_dims(dims: &[i64], value_count: usize) -> Result, EvalError> { if value_count == 0 && !dims.is_empty() { return dims @@ -1703,7 +2005,6 @@ pub(super) fn eval_linspace_values( } Ok(out) } - pub(super) fn eval_array_like_f64_values( expr: &rumoca_core::Expression, env: &VarEnv, @@ -1713,7 +2014,6 @@ pub(super) fn eval_array_like_f64_values( .map(|v| v.real()) .collect()) } - pub(super) fn eval_columns_arg( expr: Option<&rumoca_core::Expression>, env: &VarEnv, diff --git a/crates/rumoca-eval-dae/src/eval/array_helpers.rs b/crates/rumoca-eval-dae/src/eval/array_helpers.rs index b87f3e75c..637e061b2 100644 --- a/crates/rumoca-eval-dae/src/eval/array_helpers.rs +++ b/crates/rumoca-eval-dae/src/eval/array_helpers.rs @@ -132,6 +132,30 @@ fn collect_dense_indexed_values_generic( Some(values) } +fn collect_dense_declared_values_generic( + name: &str, + dims: &[i64], + scalar_count: usize, + env: &VarEnv, +) -> Option> { + if let Some(values) = collect_dense_indexed_values_generic(name, scalar_count, env) { + return Some(values); + } + if dims.len() <= 1 { + return None; + } + let mut values = Vec::with_capacity(scalar_count); + for flat_index in 0..scalar_count { + let subscripts = dae::flat_index_to_subscripts(dims, flat_index)?; + values.push( + env.vars + .get(dae::format_subscript_key(name, &subscripts).as_str()) + .copied()?, + ); + } + Some(values) +} + fn collect_dense_record_field_indexed_values_generic( name: &str, scalar_count: usize, @@ -150,12 +174,19 @@ pub(super) fn array_values_from_env_name_generic( name: &str, env: &VarEnv, ) -> Result>, EvalError> { - if let Some(dims) = env.dims.get(name) { + let declared_zero_count = if let Some(dims) = env.dims.get(name) { let scalar_count = dims.iter().map(|&d| d.max(0) as usize).product::(); - if scalar_count > 1 { - if let Some(values) = collect_dense_indexed_values_generic(name, scalar_count, env) { + if scalar_count > 0 { + if let Some(values) = + collect_dense_declared_values_generic(name, dims, scalar_count, env) + { return Ok(Some(values)); } + if scalar_count == 1 + && let Some(value) = env.vars.get(name).copied() + { + return Ok(Some(vec![value])); + } if let Some(values) = collect_dense_record_field_indexed_values_generic(name, scalar_count, env) { @@ -172,7 +203,10 @@ pub(super) fn array_values_from_env_name_generic( { return Ok(Some(values)); } - } + scalar_count == 0 + } else { + false + }; if let Some(dims) = env.dims.get(name) && !dims.is_empty() @@ -193,6 +227,10 @@ pub(super) fn array_values_from_env_name_generic( } } + if declared_zero_count { + return Ok(Some(Vec::new())); + } + Ok(collect_indexed_array_values_generic(name, env) .or_else(|| collect_record_field_indexed_values_generic(name, env))) } @@ -277,19 +315,38 @@ pub(super) fn try_eval_field_access_array_values( field: &str, env: &VarEnv, ) -> Result, EvalError> { + if let Some(name) = eval_field_access_array_path(base, field, env)? + && let Some(values) = array_values_from_env_name_generic(name.as_str(), env)? + { + return Ok(values); + } if let Some(name) = flattened_field_access_name(base, field) && let Some(values) = array_values_from_env_name_generic(name.as_str(), env)? { return Ok(values); } - match base { Expression::FunctionCall { name, args, is_constructor: false, .. - } => try_eval_function_record_field_array_values(name, args, field, env), + } => match try_eval_function_record_field_array_values(name, args, field, env) { + Ok(values) => Ok(values), + Err(EvalError::MissingFunction { .. } | EvalError::UnsupportedExpression { .. }) + if set_state_array_field_arg_index(name.var_name(), field).is_some() => + { + let arg_index = set_state_array_field_arg_index(name.var_name(), field) + .expect("checked by guard"); + let arg = args + .get(arg_index) + .ok_or(EvalError::UnsupportedExpression { + kind: "setState array field arity", + })?; + eval_array_like_values(arg, env) + } + Err(err) => Err(err), + }, Expression::VarRef { name, subscripts, .. } if subscripts.is_empty() @@ -355,6 +412,79 @@ pub(super) fn try_eval_field_access_array_values( } } +fn eval_field_access_array_path( + base: &Expression, + field: &str, + env: &VarEnv, +) -> Result, EvalError> { + let Some(prefix) = eval_array_field_base_path(base, env)? else { + return Ok(None); + }; + Ok(Some(format!("{prefix}.{field}"))) +} + +fn eval_array_field_base_path( + expr: &Expression, + env: &VarEnv, +) -> Result, EvalError> { + match expr { + Expression::VarRef { + name, subscripts, .. + } => { + if subscripts.is_empty() { + return Ok(Some(name.as_str().to_string())); + } + if subscripts + .iter() + .any(|subscript| matches!(subscript, Subscript::Colon { .. })) + { + return Ok(None); + } + let indices = eval_array_field_subscripts(subscripts, env)?; + Ok(Some(dae::format_subscript_key(name.as_str(), &indices))) + } + Expression::FieldAccess { base, field, .. } => { + let Some(prefix) = eval_array_field_base_path(base, env)? else { + return Ok(None); + }; + Ok(Some(format!("{prefix}.{field}"))) + } + _ => Ok(None), + } +} + +fn eval_array_field_subscripts( + subscripts: &[Subscript], + env: &VarEnv, +) -> Result, EvalError> { + subscripts + .iter() + .map(|subscript| eval_array_field_subscript(subscript, env)) + .collect() +} + +fn eval_array_field_subscript( + subscript: &Subscript, + env: &VarEnv, +) -> Result { + match subscript { + Subscript::Index { value, .. } => positive_array_field_index(*value as f64), + Subscript::Expr { expr, .. } => { + positive_array_field_index(eval_expr::(expr, env)?.real()) + } + Subscript::Colon { .. } => unreachable!("colon subscripts are filtered before indexing"), + } +} + +fn positive_array_field_index(value: f64) -> Result { + if !value.is_finite() || value.fract() != 0.0 || value <= 0.0 { + return Err(EvalError::UnsupportedExpression { + kind: "field access array subscript", + }); + } + Ok(value as usize) +} + fn try_eval_function_record_field_array_values( name: &rumoca_core::Reference, args: &[Expression], @@ -382,6 +512,15 @@ fn try_eval_function_record_field_array_values( kind: "record function output field shape", })?; let output_name = format!("{}.{}", output.name, field); + if total == 0 { + return eval_unknown_size_function_record_field_array_values( + name, + args, + &output_name, + &field_param.dims, + env, + ); + } let mut values = Vec::with_capacity(total); for flat_index in 0..total { let output_path = function_output_path(&output_name, &field_param.dims, flat_index).ok_or( @@ -399,6 +538,29 @@ fn try_eval_function_record_field_array_values( Ok(values) } +fn eval_unknown_size_function_record_field_array_values( + name: &rumoca_core::Reference, + args: &[Expression], + output_name: &str, + dims: &[i64], + env: &VarEnv, +) -> Result, EvalError> { + if !matches!(dims, [0]) { + return Ok(Vec::new()); + } + + let mut values = Vec::new(); + for flat_index in 0.. { + let output_path = dae::format_subscript_key(output_name, &[flat_index + 1]); + match eval_user_function_output_path_pub(name.var_name(), args, output_path.as_str(), env) { + Ok(value) => values.push(value), + Err(EvalError::MissingBinding { .. }) => break, + Err(err) => return Err(err), + } + } + Ok(values) +} + pub(super) fn record_constructor_field_param<'a>( function: &rumoca_core::Function, output: &rumoca_core::FunctionParam, diff --git a/crates/rumoca-eval-dae/src/eval/builtin_runtime.rs b/crates/rumoca-eval-dae/src/eval/builtin_runtime.rs index 7b7f62b8d..b699aaa3e 100644 --- a/crates/rumoca-eval-dae/src/eval/builtin_runtime.rs +++ b/crates/rumoca-eval-dae/src/eval/builtin_runtime.rs @@ -306,4 +306,18 @@ mod tests { assert_eq!(value, Ok(4.0)); } + + #[test] + fn eval_builtin_size_preserves_singleton_array_constructor_rank() { + let env = VarEnv::::new(); + let singleton_array = rumoca_core::Expression::Array { + elements: vec![int_literal(7)], + is_matrix: false, + span: rumoca_core::Span::DUMMY, + }; + + let value = eval_builtin_size(&[singleton_array, int_literal(1)], &env); + + assert_eq!(value, Ok(1.0)); + } } diff --git a/crates/rumoca-eval-dae/src/eval/eval_expr_impl.rs b/crates/rumoca-eval-dae/src/eval/eval_expr_impl.rs index ab3a20fb6..e5e5ee4ac 100644 --- a/crates/rumoca-eval-dae/src/eval/eval_expr_impl.rs +++ b/crates/rumoca-eval-dae/src/eval/eval_expr_impl.rs @@ -118,14 +118,12 @@ fn validate_expr( match expr { rumoca_core::Expression::Literal { value, .. } => validate_literal(value), rumoca_core::Expression::VarRef { - name, subscripts, .. + name, + subscripts, + span, } => { - for subscript in subscripts { - if let rumoca_core::Subscript::Expr { expr, .. } = subscript { - validate_expr(expr, env)?; - } - } - try_eval_var_ref(name.var_name(), subscripts, env).map(|_| ()) + validate_subscript_exprs(subscripts, env)?; + try_eval_var_ref(name.var_name(), subscripts, *span, env).map(|_| ()) } rumoca_core::Expression::Binary { op, lhs, rhs, .. } => { validate_binary_expr(op, lhs, rhs, env) @@ -177,6 +175,23 @@ fn validate_expr( } } +fn validate_subscript_exprs( + subscripts: &[rumoca_core::Subscript], + env: &VarEnv, +) -> Result<(), EvalError> { + for subscript in subscripts { + let rumoca_core::Subscript::Expr { expr, .. } = subscript else { + continue; + }; + if matches!(expr.as_ref(), rumoca_core::Expression::Range { .. }) { + validate_array_argument(expr, env)?; + continue; + } + validate_expr(expr, env)?; + } + Ok(()) +} + fn validate_function_call_expr( name: &rumoca_core::Reference, args: &[rumoca_core::Expression], @@ -186,6 +201,21 @@ fn validate_function_call_expr( if let Some(result) = validate_external_table_call(name, args, env) { return result; } + if name + .as_str() + .starts_with(rumoca_core::NAMED_FUNCTION_ARG_PREFIX) + && args.iter().any(|arg| { + matches!( + arg, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String(_), + .. + } + ) + }) + { + return Ok(()); + } if is_constructor || name.as_str() == "Complex" { return validate_expr_slice_checked(args, env); } @@ -242,9 +272,7 @@ fn validate_literal(lit: &rumoca_core::Literal) -> Result<(), EvalError> { rumoca_core::Literal::Real(_) | rumoca_core::Literal::Integer(_) | rumoca_core::Literal::Boolean(_) => Ok(()), - rumoca_core::Literal::String(_) => Err(EvalError::UnsupportedExpression { - kind: "string literal", - }), + rumoca_core::Literal::String(_) => Ok(()), } } @@ -253,6 +281,15 @@ fn validate_expr_slice_checked( env: &VarEnv, ) -> Result<(), EvalError> { for arg in args { + if matches!( + arg, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String(_), + .. + } | rumoca_core::Expression::Array { .. } + ) { + continue; + } validate_expr(arg, env)?; } Ok(()) @@ -359,14 +396,20 @@ fn validate_size_call( } fn size_arg_has_known_shape(expr: &rumoca_core::Expression, env: &VarEnv) -> bool { - matches!( - expr, + match expr { rumoca_core::Expression::VarRef { - name, - subscripts, + name, subscripts, .. + } => subscripts.is_empty() && env.dims.contains_key(name.as_str()), + rumoca_core::Expression::Array { .. } | rumoca_core::Expression::Tuple { .. } => true, + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Der, + args, .. - } if subscripts.is_empty() && env.dims.contains_key(name.as_str()) - ) + } => args + .first() + .is_some_and(|arg| size_arg_has_known_shape(arg, env)), + _ => false, + } } fn validate_external_table_call( @@ -403,16 +446,24 @@ fn validate_external_table_constructor( kind: "external table constructor", }, )?; - validate_array_argument(table_arg, env)?; - let table_matrix = validate_external_table_data_arg(table_arg, env)?; + if !matches!(table_arg, rumoca_core::Expression::Empty { .. }) { + validate_array_argument(table_arg, env)?; + } + let table_matrix = validate_external_table_constructor_data(args, env, is_time_table)?.ok_or( + EvalError::UnsupportedExpression { + kind: "external table data", + }, + )?; let columns_idx = if is_time_table { 4 } else { 3 }; - if let Some(columns) = external_table_constructor_arg(args, "columns", columns_idx) { + if let Some(columns) = external_table_constructor_arg(args, "columns", columns_idx) + && let Some(data_col_count) = table_matrix.first().map(Vec::len) + { if matches!(columns, rumoca_core::Expression::Empty { .. }) { return Ok(()); } validate_array_argument(columns, env)?; - validate_external_table_columns_arg(columns, env, table_matrix[0].len())?; + validate_external_table_columns_arg(columns, env, data_col_count)?; } let smoothness_idx = if is_time_table { 5 } else { 4 }; if let Some(smoothness) = external_table_constructor_arg(args, "smoothness", smoothness_idx) { @@ -511,6 +562,10 @@ fn bind_user_function_input_for_validation( .insert(format!("{function_name}.{}", input.name), closure); return Ok(true); } + if function_param_is_string(input) { + bind_string_function_input_shape_for_validation(local_env, input, arg, env)?; + return Ok(true); + } if function_param_is_aggregate(input) { validate_array_argument(arg, env)?; let values = if input.shape_expr.is_empty() @@ -520,7 +575,7 @@ fn bind_user_function_input_for_validation( } else { eval_array_values::(arg, env)? }; - bind_aggregate_input_for_validation(local_env, &input.name, &input.dims, values)?; + bind_aggregate_input_for_validation(local_env, input, values)?; return Ok(true); } if copy_record_function_output_fields(local_env, input, arg, env)? { @@ -547,10 +602,38 @@ fn function_param_is_function(input: &rumoca_core::FunctionParam) -> bool { input.type_name.to_ascii_lowercase().contains("function") } +pub(super) fn function_param_is_string(input: &rumoca_core::FunctionParam) -> bool { + rumoca_core::qualified_type_name_matches(&input.type_name, "String") +} + fn function_param_is_aggregate(input: &rumoca_core::FunctionParam) -> bool { !input.dims.is_empty() || !input.shape_expr.is_empty() } +pub(super) fn bind_string_function_input_shape_for_validation( + local_env: &mut VarEnv, + input: &rumoca_core::FunctionParam, + arg: &rumoca_core::Expression, + env: &VarEnv, +) -> Result<(), EvalError> { + let dims = if function_param_is_aggregate(input) { + try_infer_runtime_expr_dims(arg, env)? + .into_iter() + .map(|dim| { + i64::try_from(dim).map_err(|_| EvalError::UnsupportedExpression { + kind: "array dimensions", + }) + }) + .collect::, _>>()? + } else { + Vec::new() + }; + if !dims.is_empty() { + std::sync::Arc::make_mut(&mut local_env.dims).insert(input.name.clone(), dims); + } + Ok(()) +} + fn concrete_function_param_size(dims: &[i64]) -> Option { if dims.is_empty() { return None; @@ -564,14 +647,13 @@ fn concrete_function_param_size(dims: &[i64]) -> Option { fn bind_aggregate_input_for_validation( local_env: &mut VarEnv, - input_name: &str, - declared_dims: &[i64], + input: &rumoca_core::FunctionParam, values: Vec, ) -> Result<(), EvalError> { if values.is_empty() { return Ok(()); } - let inferred_dims = infer_dims_from_values(declared_dims, values.len())?; + let inferred_dims = infer_dims_from_values(&input.dims, values.len())?; let dims = inferred_dims .iter() .map(|dim| { @@ -580,8 +662,40 @@ fn bind_aggregate_input_for_validation( }) }) .collect::, _>>()?; - set_array_entries(local_env, input_name, &dims, &values); - std::sync::Arc::make_mut(&mut local_env.dims).insert(input_name.to_string(), dims); + set_array_entries(local_env, &input.name, &dims, &values); + std::sync::Arc::make_mut(&mut local_env.dims).insert(input.name.clone(), dims.clone()); + seed_function_input_shape_bindings_for_validation(local_env, input, &dims)?; + Ok(()) +} + +fn seed_function_input_shape_bindings_for_validation( + local_env: &mut VarEnv, + input: &rumoca_core::FunctionParam, + dims: &[i64], +) -> Result<(), EvalError> { + if input.shape_expr.len() != dims.len() { + return Ok(()); + } + for (subscript, dim) in input.shape_expr.iter().zip(dims.iter().copied()) { + if dim < 0 { + return Err(EvalError::UnsupportedExpression { + kind: "negative function shape dimension", + }); + } + let rumoca_core::Subscript::Expr { expr, .. } = subscript else { + continue; + }; + let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = expr.as_ref() + else { + continue; + }; + if !subscripts.is_empty() { + continue; + }; + local_env.set(name.as_str(), T::from_f64(dim as f64)); + } Ok(()) } @@ -624,15 +738,15 @@ fn copy_selected_input_fields_for_validation( Ok(copied) } -fn validate_external_table_data_arg( - table_arg: &rumoca_core::Expression, +fn validate_external_table_constructor_data( + args: &[rumoca_core::Expression], env: &VarEnv, -) -> Result>, EvalError> { - let table_matrix = - eval_table_matrix_arg(table_arg, env)?.ok_or(EvalError::UnsupportedExpression { - kind: "external table data", - })?; - validate_external_table_matrix(&table_matrix)?; + is_time_table: bool, +) -> Result>>, EvalError> { + let table_matrix = eval_external_table_data_matrix(args, env, is_time_table)?; + if let Some(table_matrix) = &table_matrix { + validate_external_table_matrix(table_matrix)?; + } Ok(table_matrix) } @@ -658,9 +772,7 @@ fn validate_external_table_columns_arg( fn validate_external_table_matrix(table_matrix: &[Vec]) -> Result<(), EvalError> { let Some(first_row) = table_matrix.first() else { - return Err(EvalError::UnsupportedExpression { - kind: "external table data", - }); + return Ok(()); }; let data_col_count = first_row.len(); if data_col_count < 2 { @@ -907,6 +1019,9 @@ pub(in crate::eval) fn validate_array_argument( { Ok(()) } + rumoca_core::Expression::VarRef { subscripts, .. } if !subscripts.is_empty() => { + eval_array_like_values::(expr, env).map(|_| ()) + } rumoca_core::Expression::FieldAccess { base, field, .. } if try_eval_field_access_array_values(base, field, env).is_ok() => { @@ -1091,6 +1206,11 @@ pub(super) fn eval_index_from_env_path( return env.vars.get(base_path).copied(); } + let subscript_key = dae::format_subscript_key(base_path, indices); + if let Some(value) = env.vars.get(&subscript_key).copied() { + return Some(value); + } + let dims = env.dims.get(base_path)?; if dims.len() != indices.len() { return None; @@ -1101,11 +1221,6 @@ pub(super) fn eval_index_from_env_path( .collect::>>()?; let flat_index = flat_index_from_dims(&dims, indices)?; - let subscript_key = dae::format_subscript_key(base_path, indices); - if let Some(value) = env.vars.get(&subscript_key).copied() { - return Some(value); - } - let flat_key = dae::format_subscript_key(base_path, &[flat_index + 1]); if flat_key != subscript_key && let Some(value) = env.vars.get(&flat_key).copied() @@ -1252,6 +1367,15 @@ pub(super) fn try_eval_field_access_path( }; Ok(Some(format!("{prefix}.{field}"))) } + rumoca_core::Expression::Index { + base, subscripts, .. + } => { + let Some(prefix) = try_eval_field_access_path(base, env)? else { + return Ok(None); + }; + let indices = eval_subscript_indices(subscripts, env)?; + Ok(Some(dae::format_subscript_key(&prefix, &indices))) + } _ => Ok(None), } } @@ -1297,19 +1421,38 @@ pub(super) fn bind_constructor_inputs( constructor: &rumoca_core::Function, args: &[rumoca_core::Expression], env: &VarEnv, -) -> Result<(VarEnv, Vec), EvalError> { +) -> Result<(VarEnv, Vec>), EvalError> { let mut local_env = env.clone(); let mut input_values = Vec::with_capacity(constructor.inputs.len()); let (named_args, positional_args) = split_named_and_positional_call_args(args); let mut positional_idx = 0usize; for input in &constructor.inputs { - let value = if let Some(arg_expr) = named_args.get(input.name.as_str()) { - eval_expr::(arg_expr, &local_env)? + let arg = if let Some(arg_expr) = named_args.get(input.name.as_str()) { + Some(*arg_expr) } else if let Some(arg_expr) = positional_args.get(positional_idx) { positional_idx += 1; - eval_expr::(arg_expr, &local_env)? + Some(*arg_expr) } else if let Some(default_expr) = &input.default { - eval_expr::(default_expr, &local_env)? + Some(default_expr) + } else { + None + }; + if function_param_is_string(input) { + if let Some(arg) = arg { + bind_string_function_input_shape_for_validation(&mut local_env, input, arg, env)?; + } + input_values.push(None); + continue; + } + if function_param_is_aggregate(input) { + if let Some(arg_expr) = arg { + bind_constructor_aggregate_input(&mut local_env, input, arg_expr, env)?; + } + input_values.push(None); + continue; + } + let value = if let Some(arg_expr) = arg { + eval_expr::(arg_expr, &local_env)? } else if let Some(existing) = local_env.vars.get(&input.name).copied() { existing } else { @@ -1323,11 +1466,50 @@ pub(super) fn bind_constructor_inputs( &input.name, value, ); - input_values.push(value); + input_values.push(Some(value)); } Ok((local_env, input_values)) } +fn bind_constructor_aggregate_input( + local_env: &mut VarEnv, + input: &rumoca_core::FunctionParam, + arg_expr: &rumoca_core::Expression, + env: &VarEnv, +) -> Result<(), EvalError> { + let values = eval_array_like_values::(arg_expr, env)?; + if values.is_empty() { + return Ok(()); + } + let dims = match try_infer_runtime_expr_dims(arg_expr, env) { + Ok(dims) if !dims.is_empty() => dims + .into_iter() + .map(|dim| { + i64::try_from(dim).map_err(|_| EvalError::UnsupportedExpression { + kind: "array dimensions", + }) + }) + .collect::, _>>()?, + Ok(_) + | Err(EvalError::UnsupportedExpression { .. }) + | Err(EvalError::MissingBinding { .. }) => { + infer_dims_from_values(&input.dims, values.len())? + .into_iter() + .map(|dim| { + i64::try_from(dim).map_err(|_| EvalError::UnsupportedExpression { + kind: "array dimensions", + }) + }) + .collect::, _>>()? + } + Err(err) => return Err(err), + }; + set_array_entries(local_env, &input.name, &dims, &values); + std::sync::Arc::make_mut(&mut local_env.dims).insert(input.name.clone(), dims.clone()); + seed_function_input_shape_bindings_for_validation(local_env, input, &dims)?; + Ok(()) +} + pub(super) fn eval_field_access_constructor_by_signature( base_name: &rumoca_core::VarName, args: &[rumoca_core::Expression], @@ -1345,7 +1527,7 @@ pub(super) fn eval_field_access_constructor_by_signature( .enumerate() .find(|(_, input)| input.name == field) { - return Ok(input_values.get(idx).copied()); + return Ok(input_values.get(idx).copied().flatten()); } if let Some(output) = constructor @@ -1390,6 +1572,24 @@ fn try_eval_field_access( return Err(EvalError::MissingBinding { name: key }); } + if let rumoca_core::Expression::If { + branches, + else_branch, + .. + } = base + { + for (cond, then_expr) in branches { + if eval_expr::(cond, env)?.real() != 0.0 { + return try_eval_field_access(then_expr, field, env); + } + } + return try_eval_field_access(else_branch, field, env); + } + + if let Some(projected) = projected_record_field_expr(base) { + return try_eval_field_access(&projected, field, env); + } + if let rumoca_core::Expression::FunctionCall { name, args, @@ -1397,6 +1597,19 @@ fn try_eval_field_access( .. } = base { + if let Some(arg_index) = set_state_array_field_arg_index(name.var_name(), field) { + let arg = named_or_positional_set_state_array_arg(args, field, arg_index).ok_or( + EvalError::UnsupportedExpression { + kind: "setState array field arity", + }, + )?; + return eval_array_like_values::(arg, env)? + .into_iter() + .next() + .ok_or(EvalError::UnsupportedExpression { + kind: "setState array field scalar projection", + }); + } if *is_constructor && let Some(value) = eval_field_access_constructor_by_signature(name.var_name(), args, field, env)? @@ -1419,6 +1632,47 @@ fn try_eval_field_access( }) } +fn named_or_positional_set_state_array_arg<'a>( + args: &'a [rumoca_core::Expression], + field: &str, + positional_idx: usize, +) -> Option<&'a rumoca_core::Expression> { + args.iter() + .find_map(|arg| { + let (name, value) = decode_named_constructor_arg(arg)?; + (name == field).then_some(value) + }) + .or_else(|| args.get(positional_idx)) +} + +fn projected_record_field_expr(expr: &rumoca_core::Expression) -> Option { + let rumoca_core::Expression::FieldAccess { base, field, .. } = expr else { + return None; + }; + constructor_named_field_expr(base, field) +} + +fn constructor_named_field_expr( + expr: &rumoca_core::Expression, + field: &str, +) -> Option { + let rumoca_core::Expression::FunctionCall { + args, + is_constructor, + .. + } = expr + else { + return None; + }; + if !is_constructor { + return None; + } + args.iter().find_map(|arg| { + let (name, value) = decode_named_constructor_arg(arg)?; + (name == field).then(|| value.clone()) + }) +} + fn try_eval_function_record_scalar_field( name: &rumoca_core::Reference, args: &[rumoca_core::Expression], @@ -1524,6 +1778,7 @@ pub(super) fn eval_literal(lit: &rumoca_core::Literal) -> T { fn try_eval_var_ref( name: &rumoca_core::VarName, subscripts: &[rumoca_core::Subscript], + span: rumoca_core::Span, env: &VarEnv, ) -> Result { if subscripts.is_empty() { @@ -1541,6 +1796,25 @@ fn try_eval_var_ref( kind: "colon subscript", }); } + if subscripts.len() == 1 + && let rumoca_core::Subscript::Expr { expr, .. } = &subscripts[0] + && matches!(expr.as_ref(), rumoca_core::Expression::Range { .. }) + { + let values = eval_array_values( + &rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new(name.as_str()), + subscripts: subscripts.to_vec(), + span, + }, + env, + )?; + if let [value] = values.as_slice() { + return Ok(*value); + } + return Err(EvalError::UnsupportedExpression { + kind: "range slice scalar value", + }); + } let indices = try_eval_index_subscripts(subscripts, env)?; if let Some(value) = eval_index_from_env_path(name.as_str(), &indices, env) { return Ok(value); diff --git a/crates/rumoca-eval-dae/src/eval/eval_expr_impl/builtin_eval.rs b/crates/rumoca-eval-dae/src/eval/eval_expr_impl/builtin_eval.rs index 619f3f4b7..958150477 100644 --- a/crates/rumoca-eval-dae/src/eval/eval_expr_impl/builtin_eval.rs +++ b/crates/rumoca-eval-dae/src/eval/eval_expr_impl/builtin_eval.rs @@ -44,6 +44,7 @@ fn left_limit_time_value(time: f64) -> f64 { fn eval_var_ref_from_pre_store( name: &rumoca_core::VarName, subscripts: &[rumoca_core::Subscript], + span: rumoca_core::Span, env: &VarEnv, ) -> Result, EvalError> { let mut pre_env = env.clone(); @@ -54,7 +55,7 @@ fn eval_var_ref_from_pre_store( } pre_env.set(key.as_str(), T::from_f64(value)); } - match try_eval_var_ref(name, subscripts, &pre_env) { + match try_eval_var_ref(name, subscripts, span, &pre_env) { Ok(value) => Ok(Some(value)), Err(EvalError::MissingBinding { .. }) => Ok(None), Err(err) => Err(err), @@ -182,7 +183,9 @@ pub(in crate::eval) fn eval_builtin_pre( } if let rumoca_core::Expression::VarRef { - name, subscripts, .. + name, + subscripts, + span, } = arg0 { if subscripts.is_empty() @@ -190,7 +193,7 @@ pub(in crate::eval) fn eval_builtin_pre( { return Ok(T::from_f64(value)); } - if let Some(value) = eval_var_ref_from_pre_store(name.var_name(), subscripts, env)? { + if let Some(value) = eval_var_ref_from_pre_store(name.var_name(), subscripts, *span, env)? { return Ok(value); } } @@ -214,7 +217,9 @@ pub(in crate::eval) fn eval_builtin_previous( }; if let rumoca_core::Expression::VarRef { - name, subscripts, .. + name, + subscripts, + span, } = arg0 { if subscripts.is_empty() @@ -222,7 +227,7 @@ pub(in crate::eval) fn eval_builtin_previous( { return Ok(T::from_f64(value)); } - if let Some(value) = eval_var_ref_from_pre_store(name.var_name(), subscripts, env)? { + if let Some(value) = eval_var_ref_from_pre_store(name.var_name(), subscripts, *span, env)? { return Ok(value); } // MLS §16.5.1 / §16.4: at the first clock tick, previous(v) reads the @@ -255,6 +260,7 @@ pub(in crate::eval) fn eval_builtin( rumoca_core::BuiltinFunction::Ceil => Ok(eval_builtin_arg(args, 0, env)?.ceil()), rumoca_core::BuiltinFunction::Min => eval_builtin_min(args, env), rumoca_core::BuiltinFunction::Max => eval_builtin_max(args, env), + rumoca_core::BuiltinFunction::Size => eval_builtin_size(args, env), rumoca_core::BuiltinFunction::Div => try_eval_div_mod_rem(args, env, DivKind::Div), rumoca_core::BuiltinFunction::Mod => try_eval_div_mod_rem(args, env, DivKind::Mod), rumoca_core::BuiltinFunction::Rem => try_eval_div_mod_rem(args, env, DivKind::Rem), @@ -273,6 +279,31 @@ pub(in crate::eval) fn eval_builtin( } } +fn eval_builtin_size( + args: &[rumoca_core::Expression], + env: &VarEnv, +) -> Result { + let array_arg = require_builtin_arg(args, 0)?; + let dim = eval_builtin_arg(args, 1, env)?.real(); + if dim.fract().abs() > 1.0e-12 || dim < 1.0 { + return Err(EvalError::UnsupportedExpression { + kind: "size dimension", + }); + } + let dim_index = (dim as usize) + .checked_sub(1) + .ok_or(EvalError::UnsupportedExpression { + kind: "size dimension", + })?; + let dims = try_infer_runtime_expr_dims(array_arg, env)?; + let value = dims + .get(dim_index) + .ok_or(EvalError::UnsupportedExpression { + kind: "size dimension", + })?; + Ok(T::from_f64(*value as f64)) +} + fn try_eval_der( args: &[rumoca_core::Expression], env: &VarEnv, diff --git a/crates/rumoca-eval-dae/src/eval/eval_expr_impl/checked_eval.rs b/crates/rumoca-eval-dae/src/eval/eval_expr_impl/checked_eval.rs index 88b5a9a1e..e9000d988 100644 --- a/crates/rumoca-eval-dae/src/eval/eval_expr_impl/checked_eval.rs +++ b/crates/rumoca-eval-dae/src/eval/eval_expr_impl/checked_eval.rs @@ -7,10 +7,12 @@ pub fn eval_expr( env: &VarEnv, ) -> Result { let value = match expr { - rumoca_core::Expression::Literal { value: lit, .. } => try_eval_literal::(lit), + rumoca_core::Expression::Literal { value: lit, .. } => try_eval_literal::(lit, env), rumoca_core::Expression::VarRef { - name, subscripts, .. - } => try_eval_var_ref::(name.var_name(), subscripts, env), + name, + subscripts, + span, + } => try_eval_var_ref::(name.var_name(), subscripts, *span, env), rumoca_core::Expression::Binary { op, lhs, rhs, .. } => { try_eval_binary::(op, lhs, rhs, env) } @@ -62,7 +64,16 @@ pub fn eval_expr( }) } -fn try_eval_literal(lit: &rumoca_core::Literal) -> Result { +fn try_eval_literal( + lit: &rumoca_core::Literal, + env: &VarEnv, +) -> Result { + let _ = env; + if matches!(lit, rumoca_core::Literal::String(_)) { + return Err(EvalError::UnsupportedExpression { + kind: "string literal", + }); + } validate_literal(lit)?; Ok(eval_literal::(lit)) } diff --git a/crates/rumoca-eval-dae/src/eval/external_table.rs b/crates/rumoca-eval-dae/src/eval/external_table.rs index c43eaae76..94d474aa9 100644 --- a/crates/rumoca-eval-dae/src/eval/external_table.rs +++ b/crates/rumoca-eval-dae/src/eval/external_table.rs @@ -1,4 +1,4 @@ -use std::collections::HashMap; +use std::collections::{HashMap, hash_map::Entry}; use std::sync::Mutex; use rumoca_core as core; @@ -17,7 +17,6 @@ pub(super) struct ExternalTableSpec { #[derive(Debug, Default)] pub(super) struct ExternalTableRegistry { - next_id: u64, by_hash: HashMap, tables: HashMap, } @@ -46,6 +45,12 @@ fn stable_u64_from_hash(hash: blake3::Hash) -> u64 { u64::from_le_bytes(bytes) } +fn stable_table_id_from_hash(hash: u64) -> u64 { + const MAX_EXACT_F64_INTEGER: u64 = (1u64 << 53) - 1; + let id = hash & MAX_EXACT_F64_INTEGER; + if id == 0 { 1 } else { id } +} + fn register_external_table_in_registry( registry: &Mutex, spec: ExternalTableSpec, @@ -54,13 +59,24 @@ fn register_external_table_in_registry( let mut reg = registry .lock() .unwrap_or_else(|poisoned| poisoned.into_inner()); - if let Some(existing_id) = reg.by_hash.get(&hash).copied() - && reg.tables.get(&existing_id) == Some(&spec) - { - return existing_id; + if let Some(existing_id) = reg.by_hash.get(&hash).copied() { + if reg.tables.get(&existing_id) == Some(&spec) { + return existing_id; + } + let mut candidate = existing_id; + loop { + candidate = if candidate >= ((1u64 << 53) - 1) { + 1 + } else { + candidate + 1 + }; + if let Entry::Vacant(entry) = reg.tables.entry(candidate) { + entry.insert(spec); + return candidate; + } + } } - reg.next_id = reg.next_id.saturating_add(1); - let id = reg.next_id; + let id = stable_table_id_from_hash(hash); reg.by_hash.insert(hash, id); reg.tables.insert(id, spec); id @@ -124,6 +140,21 @@ pub(super) fn external_table_data_for_values_in( tables } +pub(super) fn all_external_table_data_in( + registry: &Mutex, +) -> Vec { + let Ok(reg) = registry.lock() else { + return Vec::new(); + }; + let mut tables = reg + .tables + .iter() + .map(|(id, spec)| external_table_data(*id, spec)) + .collect::>(); + tables.sort_by_key(|table| table.id); + tables +} + pub(super) fn lookup_external_table_in_registry( registry: &Mutex, id_real: f64, diff --git a/crates/rumoca-eval-dae/src/eval/mod.rs b/crates/rumoca-eval-dae/src/eval/mod.rs index b14e1f7ec..5c34d9752 100644 --- a/crates/rumoca-eval-dae/src/eval/mod.rs +++ b/crates/rumoca-eval-dae/src/eval/mod.rs @@ -50,9 +50,9 @@ pub(super) fn with_expression_span_if_available(err: EvalError, expr: &Expressio mod external_table; use external_table::{ - ExternalTableRegistry, ExternalTableSpec, external_table_data_for_values, - external_table_data_for_values_in, lookup_external_table_in, lookup_external_table_in_registry, - register_external_table_in, + ExternalTableRegistry, ExternalTableSpec, all_external_table_data_in, + external_table_data_for_values, external_table_data_for_values_in, lookup_external_table_in, + lookup_external_table_in_registry, register_external_table_in, }; mod pre_seed; use pre_seed::try_seed_var_from_pre_store; @@ -79,8 +79,8 @@ use clock_eval::{ infer_clock_timing_from_expr, }; use table_eval::{ - eval_table_1d_lookup, eval_table_1d_lookup_with_runtime, eval_table_constructor, - eval_table_matrix_arg, table_x_bounds, + eval_external_table_data_matrix, eval_table_1d_lookup, eval_table_1d_lookup_with_runtime, + eval_table_constructor, table_x_bounds, }; macro_rules! warn_once { @@ -98,7 +98,7 @@ mod table_eval; mod special; use special::{ copy_record_function_output_fields, eval_function_call, function_closure_from_arg, - resolve_user_function_target, + resolve_user_function_target, set_state_array_field_arg_index, }; pub use special::{ deterministic_automatic_global_seed, eval_builtin_pub, eval_condition_as_root, @@ -511,6 +511,7 @@ pub struct VarEnv { pub functions: Arc>, pub dims: Arc>>, pub start_exprs: Arc>, + pub nonnumeric_names: Arc>, pub clock_intervals: Arc>, pub enum_literal_ordinals: Arc>, pub function_closures: IndexMap, @@ -526,6 +527,7 @@ impl Clone for VarEnv { functions: self.functions.clone(), dims: self.dims.clone(), start_exprs: self.start_exprs.clone(), + nonnumeric_names: self.nonnumeric_names.clone(), clock_intervals: self.clock_intervals.clone(), enum_literal_ordinals: self.enum_literal_ordinals.clone(), function_closures: self.function_closures.clone(), @@ -543,6 +545,7 @@ impl Default for VarEnv { functions: Arc::new(IndexMap::new()), dims: Arc::new(IndexMap::new()), start_exprs: Arc::new(IndexMap::new()), + nonnumeric_names: Arc::new(HashSet::new()), clock_intervals: Arc::new(IndexMap::new()), enum_literal_ordinals: Arc::new(IndexMap::new()), function_closures: IndexMap::new(), @@ -1002,6 +1005,13 @@ fn configure_env_metadata(env: &mut VarEnv, dae: &Dae) { } env.dims = Arc::new(collect_var_dims(dae)); env.start_exprs = Arc::new(collect_var_starts(dae)); + env.nonnumeric_names = Arc::new( + dae.metadata + .nonnumeric_variable_names + .iter() + .cloned() + .collect(), + ); env.clock_intervals = Arc::new(dae.clocks.intervals.clone()); env.enum_literal_ordinals = Arc::new(dae.symbols.enum_literal_ordinals.clone()); } @@ -1051,7 +1061,11 @@ fn bind_start_value( return Ok(()); } if size <= 1 && var.dims.is_empty() { - let value = eval_expr::(start, env)?; + let value = match eval_expr::(start, env) { + Ok(value) => value, + Err(err) if start_eval_error_uses_default(&err) => 0.0, + Err(err) => return Err(err), + }; env.set(name, value); return Ok(()); } @@ -1074,6 +1088,16 @@ fn bind_start_value( Ok(()) } +fn start_eval_error_uses_default(err: &EvalError) -> bool { + match err { + EvalError::UnsupportedExpression { + kind: "external table data" | "external table bounds" | "empty", + } => true, + EvalError::Spanned { source, .. } => start_eval_error_uses_default(source), + _ => false, + } +} + pub fn start_expr_is_nonnumeric(expr: &rumoca_core::Expression, env: &VarEnv) -> bool { start_expr_is_nonnumeric_inner(expr, env, &mut HashSet::new()) } @@ -1106,10 +1130,11 @@ fn start_expr_is_nonnumeric_inner( .iter() .any(|arg| start_expr_is_nonnumeric_inner(arg, env, visited)), rumoca_core::Expression::FunctionCall { name, args, .. } => { - !is_runtime_special_function_name(name.var_name()) - && args - .iter() - .any(|arg| start_expr_is_nonnumeric_inner(arg, env, visited)) + string_valued_function_call_name(name.var_name()) + || (!is_runtime_special_function_name(name.var_name()) + && args + .iter() + .any(|arg| start_expr_is_nonnumeric_inner(arg, env, visited))) } rumoca_core::Expression::If { branches, @@ -1152,8 +1177,10 @@ fn start_expr_is_nonnumeric_inner( } }) } - rumoca_core::Expression::FieldAccess { base, .. } => { - start_expr_is_nonnumeric_inner(base, env, visited) + rumoca_core::Expression::FieldAccess { base, field, .. } => { + flattened_field_access_name(base, field) + .is_some_and(|name| nonnumeric_name_matches(env, &name)) + || start_expr_is_nonnumeric_inner(base, env, visited) } rumoca_core::Expression::Range { start, step, end, .. @@ -1166,22 +1193,71 @@ fn start_expr_is_nonnumeric_inner( } rumoca_core::Expression::VarRef { name, subscripts, .. - } if subscripts.is_empty() => { - let key = name.as_str(); - if env.enum_literal_ordinals.contains_key(key) || !visited.insert(key.to_string()) { - return false; - } - env.start_exprs - .get(key) - .is_some_and(|start| start_expr_is_nonnumeric_inner(start, env, visited)) - } - rumoca_core::Expression::VarRef { .. } | rumoca_core::Expression::Empty { .. } => false, + } if subscripts.is_empty() => unsubscripted_start_var_ref_is_nonnumeric(name, env, visited), + rumoca_core::Expression::VarRef { + name, subscripts, .. + } => subscripted_start_var_ref_is_nonnumeric(name, subscripts, env, visited), + rumoca_core::Expression::Empty { .. } => false, } } +fn unsubscripted_start_var_ref_is_nonnumeric( + name: &rumoca_core::Reference, + env: &VarEnv, + visited: &mut HashSet, +) -> bool { + let key = name.as_str(); + if nonnumeric_name_matches(env, key) { + return true; + } + if env.enum_literal_ordinals.contains_key(key) || !visited.insert(key.to_string()) { + return false; + } + env.start_exprs + .get(key) + .is_some_and(|start| start_expr_is_nonnumeric_inner(start, env, visited)) +} + +fn subscripted_start_var_ref_is_nonnumeric( + name: &rumoca_core::Reference, + subscripts: &[rumoca_core::Subscript], + env: &VarEnv, + visited: &mut HashSet, +) -> bool { + nonnumeric_name_matches(env, name.as_str()) + || subscripts.iter().any(|subscript| match subscript { + rumoca_core::Subscript::Expr { expr, .. } => { + start_expr_is_nonnumeric_inner(expr, env, visited) + } + rumoca_core::Subscript::Index { .. } | rumoca_core::Subscript::Colon { .. } => false, + }) +} + +fn nonnumeric_name_matches(env: &VarEnv, name: &str) -> bool { + env.nonnumeric_names.contains(name) + || env + .nonnumeric_names + .iter() + .any(|candidate| name.ends_with(&format!(".{candidate}"))) +} + +fn string_valued_function_call_name(name: &rumoca_core::VarName) -> bool { + matches!(name.last_segment(), "String") + || matches!( + rumoca_core::modelica_string_intrinsic_short_name(name.last_segment()), + Some(rumoca_core::ModelicaStringIntrinsic::RequiresLowering) + ) +} + pub fn can_broadcast_start_value(expr: &rumoca_core::Expression, env: &VarEnv) -> bool { match expr { rumoca_core::Expression::Literal { .. } => true, + rumoca_core::Expression::Array { elements, .. } + | rumoca_core::Expression::Tuple { elements, .. } + if elements.len() == 1 => + { + can_broadcast_start_value(&elements[0], env) + } rumoca_core::Expression::VarRef { name, subscripts, .. } if subscripts.is_empty() && env.enum_literal_ordinals.contains_key(name.as_str()) => true, @@ -1191,10 +1267,21 @@ pub fn can_broadcast_start_value(expr: &rumoca_core::Expression, env: &VarEnv { - branches - .iter() - .all(|(_, value)| can_broadcast_start_value(value, env)) - && can_broadcast_start_value(else_branch, env) + for (condition, value) in branches { + match eval_expr::(condition, env) { + Ok(selected) if selected != 0.0 => { + return can_broadcast_start_value(value, env); + } + Ok(_) => {} + Err(_) => { + return branches + .iter() + .all(|(_, value)| can_broadcast_start_value(value, env)) + && can_broadcast_start_value(else_branch, env); + } + } + } + can_broadcast_start_value(else_branch, env) } _ => false, } @@ -1409,5 +1496,11 @@ pub fn external_table_data_for_parameter_values_in( external_table_data_for_values_in(&env.runtime.external_tables, values) } +pub fn all_external_table_data_in_env( + env: &VarEnv, +) -> Vec { + all_external_table_data_in(&env.runtime.external_tables) +} + #[cfg(test)] mod tests; diff --git a/crates/rumoca-eval-dae/src/eval/special.rs b/crates/rumoca-eval-dae/src/eval/special.rs index 826c05135..855007aad 100644 --- a/crates/rumoca-eval-dae/src/eval/special.rs +++ b/crates/rumoca-eval-dae/src/eval/special.rs @@ -1,3 +1,7 @@ +//! SPEC_0021 file-size exception: special-function evaluation still owns +//! Modelica builtins, runtime intrinsics, and shape-sensitive dispatch. split plan: +//! move table, state accessor, and distribution special cases apart. + use super::builtin_table::{ eval_external_table_function, resolve_function_closure, resolve_user_function, }; @@ -9,12 +13,16 @@ use super::distribution_clock::{ use super::*; mod runtime_specials; +mod shape_bindings; mod state_accessors; pub(super) use runtime_specials::*; pub use runtime_specials::{ deterministic_automatic_global_seed, is_runtime_special_function_name, is_runtime_special_function_short_name, modelica_strings_hash_string, }; +use shape_bindings::{ + seed_function_input_shape_bindings, seed_function_input_shape_bindings_from_arg, +}; pub(super) use state_accessors::*; #[derive(Clone)] @@ -191,6 +199,14 @@ fn resolved_function_param_dims( param: &FunctionParam, env: &VarEnv, ) -> Result>, EvalError> { + if let Some(dims) = env.dims.get(¶m.name) { + return Ok(Some(dims.clone())); + } + if let Some(default) = ¶m.default + && let Some(dims) = literal_array_shape(default) + { + return Ok(Some(dims)); + } if !param.shape_expr.is_empty() { return eval_function_shape_expr(¶m.shape_expr, env).map(Some); } @@ -207,6 +223,30 @@ fn resolved_function_param_dims( Ok(Some(param.dims.clone())) } +fn literal_array_shape(expr: &Expression) -> Option> { + let Expression::Array { + elements, + is_matrix, + .. + } = expr + else { + return None; + }; + if *is_matrix { + let rows = i64::try_from(elements.len()).ok()?; + let cols = elements + .iter() + .map(|row| match row { + Expression::Array { elements, .. } => i64::try_from(elements.len()).ok(), + _ => Some(1), + }) + .collect::>>()?; + let max_cols = cols.into_iter().max().unwrap_or(0); + return Some(vec![rows, max_cols]); + } + Some(vec![i64::try_from(elements.len()).ok()?]) +} + fn eval_function_shape_expr( shape_expr: &[Subscript], env: &VarEnv, @@ -264,6 +304,20 @@ fn bind_user_function_inputs( if bind_function_input_alias(local_env, function_name, param, arg_expr, caller_env)? { continue; } + if super::eval_expr_impl::function_param_is_string(param) { + super::eval_expr_impl::bind_string_function_input_shape_for_validation( + local_env, param, arg_expr, caller_env, + )?; + continue; + } + if !param.dims.is_empty() || !param.shape_expr.is_empty() { + seed_function_input_shape_bindings_from_arg_if_supported( + local_env, param, arg_expr, caller_env, + )?; + copy_array_input_entries(local_env, param, arg_expr, caller_env)?; + continue; + } + seed_function_input_shape_bindings_from_arg(local_env, param, arg_expr, caller_env)?; if copy_record_constructor_input_fields(local_env, param, arg_expr, caller_env)? { continue; } @@ -286,12 +340,47 @@ fn bind_user_function_inputs( if bind_function_input_alias(local_env, function_name, param, default_expr, caller_env)? { continue; } + if super::eval_expr_impl::function_param_is_string(param) { + super::eval_expr_impl::bind_string_function_input_shape_for_validation( + local_env, + param, + default_expr, + caller_env, + )?; + continue; + } + if !param.dims.is_empty() || !param.shape_expr.is_empty() { + seed_function_input_shape_bindings_from_arg_if_supported( + local_env, + param, + default_expr, + caller_env, + )?; + copy_array_input_entries(local_env, param, default_expr, caller_env)?; + continue; + } + seed_function_input_shape_bindings_from_arg(local_env, param, default_expr, caller_env)?; let val = eval_expr::(default_expr, local_env)?; bind_function_scalar_input(local_env, function_name, ¶m.name, val); } Ok(()) } +fn seed_function_input_shape_bindings_from_arg_if_supported( + local_env: &mut VarEnv, + param: &FunctionParam, + arg_expr: &Expression, + caller_env: &VarEnv, +) -> Result<(), EvalError> { + match seed_function_input_shape_bindings_from_arg(local_env, param, arg_expr, caller_env) { + Ok(()) => Ok(()), + Err(EvalError::UnsupportedExpression { + kind: "range" | "range slice" | "dynamic function shape colon", + }) => Ok(()), + Err(err) => Err(err), + } +} + fn copy_record_constructor_input_fields( local_env: &mut VarEnv, param: &FunctionParam, @@ -318,6 +407,21 @@ fn copy_record_constructor_input_fields( let Some((field, value_expr)) = decode_named_constructor_arg(arg) else { continue; }; + if expression_contains_string_literal(value_expr) { + continue; + } + if let Some(shape) = literal_array_shape(value_expr) { + let values = eval_array_like_values::(value_expr, caller_env)?; + set_array_entries( + local_env, + &format!("{}.{field}", param.name), + &shape, + &values, + ); + std::sync::Arc::make_mut(&mut local_env.dims) + .insert(format!("{}.{field}", param.name), shape); + continue; + } let value = eval_expr::(value_expr, caller_env)?; local_env.set(&format!("{}.{field}", param.name), value); explicit_fields.insert(field.to_string(), value); @@ -434,6 +538,7 @@ fn eval_record_start_field( match eval_expr::(start_expr, env) { Ok(value) => Ok(Some(value)), Err(err) if err.missing_binding_name().is_some() => Ok(None), + Err(err) if eval_error_is_unsupported_string_literal(&err) => Ok(None), Err(err) => Err(err), } } @@ -448,6 +553,7 @@ fn copy_array_literal_vector_entries( return Ok(false); } if param.shape_expr.is_empty() + && param.dims.iter().all(|dim| *dim > 0) && let Some(expected) = concrete_param_size(¶m.dims) && expected != elements.len() { @@ -478,6 +584,7 @@ fn copy_array_literal_vector_entries( if let Some(field) = selection_field { dims.insert(format!("{}.{field}", param.name), shape); } + seed_function_input_shape_bindings(local_env, param, &[elements.len() as i64])?; Ok(true) } @@ -535,6 +642,7 @@ fn copy_array_literal_matrix_entries( if let Some(field) = selection_field { dims.insert(format!("{}.{field}", param.name), shape); } + seed_function_input_shape_bindings(local_env, param, &[rows.len() as i64, max_cols as i64])?; Ok(true) } @@ -636,6 +744,11 @@ fn source_array_dims( if let Some(dims) = caller_env.dims.get(source_name).cloned() { return Ok(dims); } + if !param.dims.is_empty() && param.dims.iter().any(|dim| *dim <= 0) { + return Err(EvalError::UnsupportedExpression { + kind: "function array input argument shape", + }); + } if concrete_param_size(¶m.dims).is_some() { return Ok(param.dims.clone()); } @@ -667,6 +780,7 @@ fn copy_array_input_entries( Expression::VarRef { name, subscripts, .. } if subscripts.is_empty() => Some(name.as_str().to_string()), + Expression::VarRef { subscripts, .. } if !subscripts.is_empty() => None, _ => try_eval_field_access_path(arg_expr, caller_env)?, }, }; @@ -687,6 +801,7 @@ fn copy_array_input_entries( validate_array_input_dims(&dims, values.len())?; set_array_entries(local_env, param_name, &dims, &values); std::sync::Arc::make_mut(&mut local_env.dims).insert(param_name.to_string(), dims.clone()); + seed_function_input_shape_bindings(local_env, param, &dims)?; if trace_array_bind && source_name.contains("timeTable.table") { let t11 = env_array_sample(caller_env, &source_name, &[1, 1]); @@ -737,7 +852,8 @@ fn bind_evaluated_array_input( }); } set_array_entries(local_env, ¶m.name, &dims, &values); - std::sync::Arc::make_mut(&mut local_env.dims).insert(param.name.clone(), dims); + std::sync::Arc::make_mut(&mut local_env.dims).insert(param.name.clone(), dims.clone()); + seed_function_input_shape_bindings(local_env, param, &dims)?; Ok(()) } @@ -758,12 +874,15 @@ fn resolved_array_input_dims( ) .map(Some); } - if concrete_param_size(¶m.dims).is_some() { - return Ok(Some(param.dims.clone())); - } if !param.dims.is_empty() && param.dims.iter().any(|dim| *dim < 0) { return infer_dynamic_array_input_dims_from_declared(¶m.dims, value_count).map(Some); } + if !param.dims.is_empty() && param.dims.contains(&0) && value_count > 0 { + return infer_dynamic_array_input_dims_from_declared(¶m.dims, value_count).map(Some); + } + if concrete_param_size(¶m.dims).is_some() { + return Ok(Some(param.dims.clone())); + } Ok(None) } @@ -970,21 +1089,91 @@ fn initialize_user_function_scope_values( .get(param.name.as_str()) .cloned() .unwrap_or_else(|| param.dims.clone()); - let Some(size) = concrete_param_size(&dims) else { - if let Some(default) = param.default.as_ref() { - let val = eval_expr::(default, local_env)?; - local_env.set(¶m.name, val); - } + let Some(default) = param.default.as_ref() else { continue; }; - if let Some(default) = param.default.as_ref() { + if expression_contains_string_literal(default) { + continue; + } + if let Some(size) = concrete_param_size(&dims) { let values = eval_shaped_array_values::(default, local_env, size)?; set_array_entries(local_env, ¶m.name, &dims, &values); + } else { + let val = eval_expr::(default, local_env)?; + local_env.set(¶m.name, val); } } Ok(()) } +fn expression_contains_string_literal(expr: &Expression) -> bool { + match expr { + Expression::Literal { + value: rumoca_core::Literal::String(_), + .. + } => true, + Expression::Literal { .. } | Expression::Empty { .. } => false, + Expression::Unary { rhs, .. } => expression_contains_string_literal(rhs), + Expression::Binary { lhs, rhs, .. } => { + expression_contains_string_literal(lhs) || expression_contains_string_literal(rhs) + } + Expression::If { + branches, + else_branch, + .. + } => { + branches.iter().any(|(cond, value)| { + expression_contains_string_literal(cond) + || expression_contains_string_literal(value) + }) || expression_contains_string_literal(else_branch) + } + Expression::Array { elements, .. } | Expression::Tuple { elements, .. } => { + elements.iter().any(expression_contains_string_literal) + } + Expression::Range { + start, step, end, .. + } => { + expression_contains_string_literal(start) + || step + .as_ref() + .is_some_and(|step| expression_contains_string_literal(step)) + || expression_contains_string_literal(end) + } + Expression::FunctionCall { args, .. } | Expression::BuiltinCall { args, .. } => { + args.iter().any(expression_contains_string_literal) + } + Expression::Index { + base, subscripts, .. + } => { + expression_contains_string_literal(base) + || subscripts.iter().any(|subscript| match subscript { + rumoca_core::Subscript::Expr { expr, .. } => { + expression_contains_string_literal(expr) + } + rumoca_core::Subscript::Index { .. } | rumoca_core::Subscript::Colon { .. } => { + false + } + }) + } + Expression::FieldAccess { base, .. } => expression_contains_string_literal(base), + Expression::ArrayComprehension { + expr, + indices, + filter, + .. + } => { + expression_contains_string_literal(expr) + || indices + .iter() + .any(|index| expression_contains_string_literal(&index.range)) + || filter + .as_ref() + .is_some_and(|filter| expression_contains_string_literal(filter)) + } + Expression::VarRef { .. } => false, + } +} + fn selected_output_name(selection: &OutputSelection) -> Result { if selection.indices.is_empty() { return Ok(selection.output_name.clone()); @@ -1280,6 +1469,7 @@ fn eval_user_function_call( // so selected function calls (`*.re` / `*.im`) propagate correctly. let eval_result = with_function_call_stack(&env.runtime, name.as_str(), || { bind_user_function_inputs(&mut local_env, name.as_str(), &inputs, args, env)?; + seed_resolved_function_scope_dims(&mut local_env, &inputs)?; seed_resolved_function_scope_dims(&mut local_env, &outputs)?; seed_resolved_function_scope_dims(&mut local_env, &locals)?; trace_function_call_inputs(trace_call, &local_env, &inputs); @@ -1358,6 +1548,7 @@ fn eval_user_function_local_env( seed_static_function_scope_dims(&mut local_env, &inputs, &outputs, &locals); with_function_call_stack(&env.runtime, name.as_str(), || { bind_user_function_inputs(&mut local_env, name.as_str(), &inputs, args, env)?; + seed_resolved_function_scope_dims(&mut local_env, &inputs)?; seed_resolved_function_scope_dims(&mut local_env, &outputs)?; seed_resolved_function_scope_dims(&mut local_env, &locals)?; initialize_user_function_scope_values(&mut local_env, &outputs, &locals)?; @@ -1404,6 +1595,7 @@ pub fn eval_user_function_array_output_pub( seed_static_function_scope_dims(&mut local_env, &inputs, &outputs, &locals); with_function_call_stack(&env.runtime, name.as_str(), || { bind_user_function_inputs(&mut local_env, name.as_str(), &inputs, args, env)?; + seed_resolved_function_scope_dims(&mut local_env, &inputs)?; seed_resolved_function_scope_dims(&mut local_env, &outputs)?; seed_resolved_function_scope_dims(&mut local_env, &locals)?; initialize_user_function_scope_values(&mut local_env, &outputs, &locals)?; @@ -1527,12 +1719,23 @@ fn copy_selected_input_start_fields( match eval_expr::(start_expr, env) { Ok(value) => local_env.set(&dst, value), Err(err) if err.missing_binding_name().is_some() => {} + Err(err) if eval_error_is_unsupported_string_literal(&err) => {} Err(err) => return Err(err), } } Ok(()) } +fn eval_error_is_unsupported_string_literal(err: &EvalError) -> bool { + match err { + EvalError::UnsupportedExpression { + kind: "string literal", + } => true, + EvalError::Spanned { source, .. } => eval_error_is_unsupported_string_literal(source), + _ => false, + } +} + pub(in crate::eval) fn copy_record_function_output_fields( local_env: &mut VarEnv, param: &FunctionParam, @@ -1749,7 +1952,22 @@ fn eval_constructor_call( } let _ = name; - eval_expr::(&args[0], env) + let Some(arg) = args + .iter() + .find(|arg| constructor_scalar_fallback_arg_is_numeric(arg)) + else { + return Err(EvalError::UnsupportedExpression { + kind: "string literal", + }); + }; + eval_expr::(arg, env) +} + +fn constructor_scalar_fallback_arg_is_numeric(arg: &Expression) -> bool { + let value = decode_named_constructor_arg(arg) + .map(|(_, value)| value) + .unwrap_or(arg); + !expression_contains_string_literal(value) && !matches!(value, Expression::Array { .. }) } pub(super) fn eval_function_call( diff --git a/crates/rumoca-eval-dae/src/eval/special/runtime_specials.rs b/crates/rumoca-eval-dae/src/eval/special/runtime_specials.rs index ae79c1ec7..9e4a91043 100644 --- a/crates/rumoca-eval-dae/src/eval/special/runtime_specials.rs +++ b/crates/rumoca-eval-dae/src/eval/special/runtime_specials.rs @@ -1,3 +1,4 @@ +use super::super::table_eval::resolve_modelica_resource_path; use super::*; use std::sync::Mutex; @@ -69,6 +70,31 @@ fn eval_string_expr(expr: &Expression, env: &VarEnv) -> Result Ok(value.clone()), + Expression::VarRef { + name, subscripts, .. + } if subscripts.is_empty() => { + let Some(start) = env.start_exprs.get(name.as_str()) else { + return Err(with_expression_span_if_available( + EvalError::MissingBinding { + name: name.as_str().to_string(), + }, + expr, + )); + }; + eval_string_expr(start, env) + } + Expression::If { + branches, + else_branch, + .. + } => { + for (condition, value) in branches { + if eval_expr::(condition, env)?.to_bool() { + return eval_string_expr(value, env); + } + } + eval_string_expr(else_branch, env) + } Expression::FunctionCall { name, args, span, .. } => eval_string_function_call(name.as_str(), args, env) @@ -85,12 +111,23 @@ fn eval_string_expr(expr: &Expression, env: &VarEnv) -> Result( name: &str, args: &[Expression], - _env: &VarEnv, + env: &VarEnv, ) -> Result { match name { "getInstanceName" => Err(EvalError::UnsupportedExpression { kind: "getInstanceName must be lowered to a string literal before DAE evaluation", }), + "loadResource" | "Modelica.Utilities.Files.loadResource" => { + let raw = args + .first() + .ok_or(EvalError::UnsupportedExpression { + kind: "loadResource path", + }) + .and_then(|arg| eval_string_expr(arg, env))?; + Ok(resolve_modelica_resource_path(&raw) + .map(|path| path.to_string_lossy().into_owned()) + .unwrap_or(raw)) + } _ => { if args.iter().any(|arg| { matches!( @@ -161,9 +198,7 @@ fn eval_misc_intrinsic_function( None => T::one(), }; if cond.to_bool() { - Err(EvalError::UnsupportedExpression { - kind: "assert statement in scalar expression", - }) + Ok(Some(T::zero())) } else { Err(EvalError::UnsupportedExpression { kind: "assert" }) } @@ -226,21 +261,25 @@ fn eval_string_misc_intrinsic_function( }); } Some(rumoca_core::ModelicaStringIntrinsic::IsEmpty) => { - return eval_literal_string_arg(args, |s| T::from_bool(s.trim().is_empty())); + let value = eval_required_string_arg(args, env, 0, "string argument")?; + return Ok(Some(T::from_bool(value.trim().is_empty()))); } Some(rumoca_core::ModelicaStringIntrinsic::HashString) => { - return eval_literal_string_arg(args, |s| { - T::from_f64(modelica_strings_hash_string(s) as f64) - }); + let value = eval_required_string_arg(args, env, 0, "string argument")?; + return Ok(Some(T::from_f64( + modelica_strings_hash_string(&value) as f64 + ))); } Some(rumoca_core::ModelicaStringIntrinsic::Length) => { - return eval_literal_string_arg(args, |s| T::from_f64(s.chars().count() as f64)); + let value = eval_required_string_arg(args, env, 0, "string argument")?; + return Ok(Some(T::from_f64(value.chars().count() as f64))); } Some(intrinsic @ rumoca_core::ModelicaStringIntrinsic::Find) | Some(intrinsic @ rumoca_core::ModelicaStringIntrinsic::FindLast) => { return eval_string_find_intrinsic( intrinsic == rumoca_core::ModelicaStringIntrinsic::FindLast, args, + env, ); } None => {} @@ -258,50 +297,159 @@ fn eval_string_misc_intrinsic_function( } } -fn eval_literal_string_arg( +fn eval_required_string_arg( args: &[Expression], - eval: impl FnOnce(&str) -> T, -) -> Result, EvalError> { - if let Some(Expression::Literal { - value: Literal::String(s), - .. - }) = args.first() - { - Ok(Some(eval(s))) - } else { - Err(EvalError::UnsupportedExpression { - kind: "literal string argument", - }) - } + env: &VarEnv, + idx: usize, + kind: &'static str, +) -> Result { + let expr = args + .get(idx) + .ok_or(EvalError::UnsupportedExpression { kind })?; + eval_string_expr(expr, env) } fn eval_string_find_intrinsic( find_last: bool, args: &[Expression], + env: &VarEnv, ) -> Result, EvalError> { - let ( - Some(Expression::Literal { - value: Literal::String(haystack), - .. - }), - Some(Expression::Literal { - value: Literal::String(needle), - .. - }), - ) = (args.first(), args.get(1)) + let (named_args, positional_args) = split_named_and_positional_call_args(args); + let haystack = eval_string_find_string_arg( + &named_args, + &positional_args, + "string", + 0, + env, + "string search arguments", + )?; + let needle = eval_string_find_string_arg( + &named_args, + &positional_args, + "searchString", + 1, + env, + "string search arguments", + )?; + let case_sensitive = + eval_string_find_bool_arg(&named_args, &positional_args, "caseSensitive", 3, env)? + .unwrap_or(true); + let default_start = if find_last { + haystack.chars().count().max(1) as i64 + } else { + 1 + }; + let start_index = + eval_string_find_integer_arg(&named_args, &positional_args, "startIndex", 2, env)? + .unwrap_or(default_start); + let idx = modelica_string_find(&haystack, &needle, start_index, find_last, case_sensitive); + Ok(Some(T::from_f64(idx as f64))) +} + +fn eval_string_find_string_arg( + named_args: &std::collections::HashMap<&str, &Expression>, + positional_args: &[&Expression], + name: &str, + positional_idx: usize, + env: &VarEnv, + kind: &'static str, +) -> Result { + let expr = named_args + .get(name) + .copied() + .or_else(|| positional_args.get(positional_idx).copied()) + .ok_or(EvalError::UnsupportedExpression { kind })?; + eval_string_expr(expr, env) +} + +fn eval_string_find_integer_arg( + named_args: &std::collections::HashMap<&str, &Expression>, + positional_args: &[&Expression], + name: &str, + positional_idx: usize, + env: &VarEnv, +) -> Result, EvalError> { + let Some(expr) = named_args + .get(name) + .copied() + .or_else(|| positional_args.get(positional_idx).copied()) else { - return Err(EvalError::UnsupportedExpression { - kind: "literal string search arguments", - }); + return Ok(None); + }; + Ok(Some(eval_expr::(expr, env)?.real().round() as i64)) +} + +fn eval_string_find_bool_arg( + named_args: &std::collections::HashMap<&str, &Expression>, + positional_args: &[&Expression], + name: &str, + positional_idx: usize, + env: &VarEnv, +) -> Result, EvalError> { + let Some(expr) = named_args + .get(name) + .copied() + .or_else(|| positional_args.get(positional_idx).copied()) + else { + return Ok(None); }; - let idx = if find_last { - haystack.rfind(needle) + Ok(Some(eval_expr::(expr, env)?.to_bool())) +} + +fn modelica_string_find( + haystack: &str, + needle: &str, + start_index: i64, + find_last: bool, + case_sensitive: bool, +) -> usize { + if needle.is_empty() || start_index < 1 { + return 0; + } + let source = if case_sensitive { + haystack.to_string() } else { - haystack.find(needle) + haystack.to_lowercase() }; - Ok(Some(T::from_f64( - idx.map_or(0.0, |index| index.saturating_add(1) as f64), - ))) + let target = if case_sensitive { + needle.to_string() + } else { + needle.to_lowercase() + }; + if find_last { + find_last_modelica_index(&source, &target, start_index) + } else { + find_first_modelica_index(&source, &target, start_index) + } +} + +fn find_first_modelica_index(source: &str, target: &str, start_index: i64) -> usize { + let start = (start_index - 1) as usize; + source + .char_indices() + .nth(start) + .and_then(|(byte_start, _)| { + source[byte_start..] + .find(target) + .map(|offset| source[..byte_start + offset].chars().count() + 1) + }) + .unwrap_or(0) +} + +fn find_last_modelica_index(source: &str, target: &str, start_index: i64) -> usize { + let char_count = source.chars().count(); + if char_count == 0 { + return 0; + } + let end_char = start_index.min(char_count as i64).max(1) as usize; + source + .char_indices() + .nth(end_char) + .map(|(byte_end, _)| &source[..byte_end]) + .unwrap_or(source) + .rfind(target) + .map(|byte_index| source[..byte_index].chars().count() + 1) + .unwrap_or(0) } fn eval_qualified_special_function( @@ -553,6 +701,7 @@ pub fn is_runtime_special_function_short_name(short_name: &str) -> bool { | "getTable1DValueNoDer2" | "getTable1DValue" | "anyTrue" + | "allTrue" | "andTrue" | "firstTrueIndex" | "distribution" diff --git a/crates/rumoca-eval-dae/src/eval/special/shape_bindings.rs b/crates/rumoca-eval-dae/src/eval/special/shape_bindings.rs new file mode 100644 index 000000000..4264f821f --- /dev/null +++ b/crates/rumoca-eval-dae/src/eval/special/shape_bindings.rs @@ -0,0 +1,58 @@ +use super::*; + +pub(super) fn seed_function_input_shape_bindings( + local_env: &mut VarEnv, + param: &FunctionParam, + dims: &[i64], +) -> Result<(), EvalError> { + if param.shape_expr.len() != dims.len() { + return Ok(()); + } + for (subscript, dim) in param.shape_expr.iter().zip(dims.iter().copied()) { + if dim < 0 { + return Err(EvalError::UnsupportedExpression { + kind: "negative function shape dimension", + }); + } + let Subscript::Expr { expr, .. } = subscript else { + continue; + }; + let Expression::VarRef { + name, subscripts, .. + } = expr.as_ref() + else { + continue; + }; + if !subscripts.is_empty() { + continue; + } + local_env.set(name.as_str(), T::from_f64(dim as f64)); + } + Ok(()) +} + +pub(super) fn seed_function_input_shape_bindings_from_arg( + local_env: &mut VarEnv, + param: &FunctionParam, + arg_expr: &Expression, + caller_env: &VarEnv, +) -> Result<(), EvalError> { + if param.shape_expr.is_empty() { + return Ok(()); + } + let dims = match try_infer_runtime_expr_dims(arg_expr, caller_env) { + Ok(dims) => dims + .into_iter() + .map(|dim| { + i64::try_from(dim).map_err(|_| EvalError::UnsupportedExpression { + kind: "array dimensions", + }) + }) + .collect::, _>>()?, + Err(EvalError::UnsupportedExpression { .. }) | Err(EvalError::MissingBinding { .. }) => { + return Ok(()); + } + Err(err) => return Err(err), + }; + seed_function_input_shape_bindings(local_env, param, &dims) +} diff --git a/crates/rumoca-eval-dae/src/eval/special/state_accessors.rs b/crates/rumoca-eval-dae/src/eval/special/state_accessors.rs index a29b23342..498a5772a 100644 --- a/crates/rumoca-eval-dae/src/eval/special/state_accessors.rs +++ b/crates/rumoca-eval-dae/src/eval/special/state_accessors.rs @@ -13,7 +13,7 @@ pub(super) fn eval_boolean_vector_function( }; match short_name { "anyTrue" => Ok(Some(T::from_bool(values()?.iter().any(|v| *v != 0.0)))), - "andTrue" => { + "allTrue" | "andTrue" => { let vals = values()?; Ok(Some(T::from_bool( !vals.is_empty() && vals.iter().all(|v| *v != 0.0), @@ -306,6 +306,19 @@ pub(in crate::eval) fn eval_state_accessor_from_set_state( } } +pub(in crate::eval) fn set_state_array_field_arg_index( + name: &VarName, + field: &str, +) -> Option { + if field != "X" { + return None; + } + match name.last_segment() { + "setState_pTX" | "setState_dTX" | "setState_phX" | "setState_psX" => Some(2), + _ => None, + } +} + fn optional_eval_arg( expr: Option<&Expression>, env: &VarEnv, diff --git a/crates/rumoca-eval-dae/src/eval/table_eval.rs b/crates/rumoca-eval-dae/src/eval/table_eval.rs index ed6d7d560..64078146d 100644 --- a/crates/rumoca-eval-dae/src/eval/table_eval.rs +++ b/crates/rumoca-eval-dae/src/eval/table_eval.rs @@ -1,4 +1,5 @@ use super::*; +use std::path::{Path, PathBuf}; pub(super) fn eval_table_matrix_arg( expr: &rumoca_core::Expression, @@ -61,7 +62,18 @@ pub(super) fn eval_table_matrix_arg( else_branch, .. } => eval_table_if_matrix_arg(branches, else_branch, env), - _ => Ok(None), + _ => eval_array_values::(expr, env) + .map(|values| { + Some(if values.is_empty() { + Vec::new() + } else { + vec![values.iter().map(|value| value.real()).collect()] + }) + }) + .or_else(|err| match err { + EvalError::UnsupportedExpression { .. } => Ok(None), + err => Err(err), + }), } } @@ -121,6 +133,161 @@ fn richer_table_start_matrix( Ok((start_len > flat_len).then_some(start_matrix)) } +pub(super) fn eval_external_table_data_matrix( + args: &[rumoca_core::Expression], + env: &VarEnv, + is_time_table: bool, +) -> Result>>, EvalError> { + let table_arg_idx = 2usize; + let Some(table_arg) = external_table_constructor_arg(args, "table", table_arg_idx) else { + return Ok(None); + }; + let table_matrix = eval_table_matrix_arg(table_arg, env)?; + if table_matrix + .as_ref() + .is_some_and(|matrix| !matrix.is_empty()) + { + return Ok(table_matrix); + } + + let file_name = external_table_constructor_arg(args, "fileName", 1) + .and_then(|expr| eval_string_expr(expr, env)) + .filter(|name| name != "NoName" && !name.trim().is_empty()); + let Some(file_name) = file_name else { + return Ok(table_matrix); + }; + + let table_name = external_table_constructor_arg(args, "tableName", 0) + .and_then(|expr| eval_string_expr(expr, env)) + .filter(|name| name != "NoName" && !name.trim().is_empty()) + .unwrap_or_else(|| { + if is_time_table { + "tab".to_string() + } else { + "tab1".to_string() + } + }); + + read_modelica_text_table(&file_name, &table_name) +} + +fn eval_string_expr( + expr: &rumoca_core::Expression, + env: &VarEnv, +) -> Option { + match expr { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String(value), + .. + } => Some(value.clone()), + rumoca_core::Expression::VarRef { + name, subscripts, .. + } if subscripts.is_empty() => env + .start_exprs + .get(name.as_str()) + .and_then(|start| eval_string_expr(start, env)), + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + for (condition, value) in branches { + if eval_expr::(condition, env).ok()?.to_bool() { + return eval_string_expr(value, env); + } + } + eval_string_expr(else_branch, env) + } + rumoca_core::Expression::FunctionCall { name, args, .. } + if name.last_segment() == "loadResource" => + { + let raw = args.first().and_then(|arg| eval_string_expr(arg, env))?; + resolve_modelica_resource_path(&raw) + .map(|path| path.to_string_lossy().into_owned()) + .or(Some(raw)) + } + _ => None, + } +} + +fn read_modelica_text_table( + file_name: &str, + table_name: &str, +) -> Result>>, EvalError> { + let Some(path) = resolve_modelica_resource_path(file_name) else { + return Ok(None); + }; + let text = std::fs::read_to_string(path).map_err(|_| EvalError::UnsupportedExpression { + kind: "external table file", + })?; + parse_modelica_text_table(&text, table_name) +} + +fn parse_modelica_text_table( + text: &str, + table_name: &str, +) -> Result>>, EvalError> { + let needle = format!("double {table_name}("); + let mut lines = text.lines(); + let Some(header) = lines.find(|line| line.trim_start().starts_with(&needle)) else { + return Ok(None); + }; + let dims = header + .split_once('(') + .and_then(|(_, tail)| tail.split_once(')')) + .map(|(dims, _)| dims) + .ok_or(EvalError::UnsupportedExpression { + kind: "external table header", + })? + .split(',') + .map(|part| part.trim().parse::().ok()) + .collect::>>() + .ok_or(EvalError::UnsupportedExpression { + kind: "external table header", + })?; + if dims.len() != 2 || dims[0] == 0 || dims[1] < 2 { + return Err(EvalError::UnsupportedExpression { + kind: "external table header", + }); + } + + let mut rows = Vec::with_capacity(dims[0]); + for line in lines { + let trimmed = line.trim(); + if trimmed.is_empty() || trimmed.starts_with('#') { + continue; + } + let row = trimmed + .split(|ch: char| ch == ',' || ch == ';' || ch.is_whitespace()) + .filter(|part| !part.is_empty()) + .map(|part| part.parse::().ok()) + .collect::>>() + .ok_or(EvalError::UnsupportedExpression { + kind: "external table row", + })?; + if row.len() != dims[1] { + return Err(EvalError::UnsupportedExpression { + kind: "external table row", + }); + } + rows.push(row); + if rows.len() == dims[0] { + return Ok(Some(rows)); + } + } + Err(EvalError::UnsupportedExpression { + kind: "external table data", + }) +} + +pub(super) fn resolve_modelica_resource_path(raw: &str) -> Option { + let raw_path = Path::new(raw); + if raw_path.is_file() { + return Some(raw_path.to_path_buf()); + } + None +} + fn map_selected_table_column( columns: &[usize], requested_output_col: usize, @@ -259,7 +426,7 @@ pub(super) fn eval_table_1d_lookup_with_runtime( } let last_idx = spec.data.len() - 1; - let k = if x_real <= spec.data[0][0] { + let k = if (out_of_range && x.real() < x_min) || x_real < spec.data[0][0] { 0usize } else if x_real >= spec.data[last_idx][0] { last_idx.saturating_sub(1) @@ -296,20 +463,13 @@ pub(super) fn eval_table_constructor( env: &VarEnv, is_time_table: bool, ) -> Result, EvalError> { - let table_arg_idx = 2usize; let columns_arg_idx = if is_time_table { 4 } else { 3 }; let smoothness_idx = if is_time_table { 5 } else { 4 }; let extrapolation_idx = if is_time_table { 6 } else { 5 }; - let Some(table_arg) = external_table_constructor_arg(args, "table", table_arg_idx) else { + let Some(table_matrix) = eval_external_table_data_matrix(args, env, is_time_table)? else { return Ok(None); }; - let Some(table_matrix) = eval_table_matrix_arg(table_arg, env)? else { - return Ok(None); - }; - if table_matrix.is_empty() { - return Ok(None); - } let columns = eval_columns_arg( external_table_constructor_arg(args, "columns", columns_arg_idx), diff --git a/crates/rumoca-eval-dae/src/eval/tests/clock_and_tables.rs b/crates/rumoca-eval-dae/src/eval/tests/clock_and_tables.rs index c7f2fb5fd..c97efd843 100644 --- a/crates/rumoca-eval-dae/src/eval/tests/clock_and_tables.rs +++ b/crates/rumoca-eval-dae/src/eval/tests/clock_and_tables.rs @@ -606,6 +606,115 @@ fn test_table1d_constructor_uses_start_expr_fallback_for_dynamic_dims() { assert!((y - 12.0).abs() < 1e-12); } +#[test] +fn test_table1d_constructor_loads_modelica_text_table_file() { + let mut env = VarEnv::::new(); + let path = std::env::temp_dir().join(format!( + "rumoca-table-{}-{}.txt", + std::process::id(), + "table1d" + )); + std::fs::write( + &path, + "#1\n\ +double tab1(3,2)\n\ +0,10\n\ +1,12\n\ +2,14\n", + ) + .expect("write table fixture"); + + let constructor = fn_call( + "ExternalCombiTable1D", + vec![ + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("tab1".to_string()), + span: rumoca_core::Span::DUMMY, + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String(path.to_string_lossy().into_owned()), + span: rumoca_core::Span::DUMMY, + }, + rumoca_core::Expression::Empty { + span: rumoca_core::Span::DUMMY, + }, + columns_expr(), + int_lit(1), + int_lit(1), + ], + ); + let table_id = eval_expr::(&constructor, &env).expect("file table should register"); + assert!(table_id > 0.0); + + env.set("table_id", table_id); + let lookup = fn_call( + "getTable1DValueNoDer", + vec![var("table_id"), int_lit(1), lit(1.5)], + ); + let y = eval_expr::(&lookup, &env).expect("file table lookup should evaluate"); + std::fs::remove_file(path).ok(); + + assert!((y - 13.0).abs() < 1e-12); +} + +#[test] +fn test_table1d_constructor_loads_file_when_zero_row_table_varref_has_no_values() { + let mut env = VarEnv::::new(); + env.dims = Arc::new(IndexMap::from([("tbl".to_string(), vec![0, 2])])); + env.start_exprs = Arc::new(IndexMap::from([( + "tbl".to_string(), + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Fill, + args: vec![lit(0.0), int_lit(0), int_lit(2)], + span: rumoca_core::Span::DUMMY, + }, + )])); + let path = std::env::temp_dir().join(format!( + "rumoca-table-{}-{}.txt", + std::process::id(), + "zero-row-varref" + )); + std::fs::write( + &path, + "#1\n\ +double tab1(3,2)\n\ +0,10\n\ +1,12\n\ +2,14\n", + ) + .expect("write table fixture"); + + let constructor = fn_call( + "ExternalCombiTable1D", + vec![ + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("tab1".to_string()), + span: rumoca_core::Span::DUMMY, + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String(path.to_string_lossy().into_owned()), + span: rumoca_core::Span::DUMMY, + }, + var("tbl"), + columns_expr(), + int_lit(1), + int_lit(1), + ], + ); + let table_id = eval_expr::(&constructor, &env).expect("file table should register"); + assert!(table_id > 0.0); + + env.set("table_id", table_id); + let lookup = fn_call( + "getTable1DValueNoDer", + vec![var("table_id"), int_lit(1), lit(1.5)], + ); + let y = eval_expr::(&lookup, &env).expect("file table lookup should evaluate"); + std::fs::remove_file(path).ok(); + + assert!((y - 13.0).abs() < 1e-12); +} + #[test] fn test_table1d_constructor_accepts_flattened_field_access_matrix() { let mut env = VarEnv::::new(); diff --git a/crates/rumoca-eval-dae/src/eval/tests/complex_array_selection.rs b/crates/rumoca-eval-dae/src/eval/tests/complex_array_selection.rs index 485a8b05c..3a99d3512 100644 --- a/crates/rumoca-eval-dae/src/eval/tests/complex_array_selection.rs +++ b/crates/rumoca-eval-dae/src/eval/tests/complex_array_selection.rs @@ -163,6 +163,33 @@ fn test_eval_array_values_record_field_varref_reads_indexed_record_elements() { assert!((values[1] - 0.1).abs() < 1.0e-12); } +#[test] +fn test_eval_array_values_nested_indexed_record_field_path() { + let mut env = VarEnv::::new(); + std::sync::Arc::make_mut(&mut env.dims).insert("source[1].medium.X".to_string(), vec![2]); + env.set("source[1].medium.X[1]", 0.73); + env.set("source[1].medium.X[2]", 0.27); + + let expr = Expression::FieldAccess { + base: Box::new(Expression::FieldAccess { + base: Box::new(Expression::VarRef { + name: Reference::new("source"), + subscripts: vec![Subscript::generated_index(1, rumoca_core::Span::DUMMY)], + span: rumoca_core::Span::DUMMY, + }), + field: "medium".to_string(), + span: rumoca_core::Span::DUMMY, + }), + field: "X".to_string(), + span: rumoca_core::Span::DUMMY, + }; + + assert_eq!( + eval_array_values::(&expr, &env).expect("nested indexed record field evaluates"), + vec![0.73, 0.27] + ); +} + #[test] fn test_eval_builtin_sum_record_field_varref_reads_indexed_record_elements() { let mut env = VarEnv::::new(); diff --git a/crates/rumoca-eval-dae/src/eval/tests/env_refresh.rs b/crates/rumoca-eval-dae/src/eval/tests/env_refresh.rs index ac2d2f606..eeabe0330 100644 --- a/crates/rumoca-eval-dae/src/eval/tests/env_refresh.rs +++ b/crates/rumoca-eval-dae/src/eval/tests/env_refresh.rs @@ -233,6 +233,91 @@ fn build_runtime_parameter_tail_env_skips_string_parameter_alias_chain() { assert!(!env.vars.contains_key("fixedTranslation.shape.shapeType")); } +#[test] +fn build_runtime_parameter_tail_env_skips_get_instance_name_string_alias_chain() { + let mut dae = dae::Dae::default(); + + let mut building_name = dae::Variable::new( + VarName::new("building.modelicaNameBuilding"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ); + building_name.start = Some(rumoca_core::Expression::FunctionCall { + name: Reference::new("getInstanceName"), + args: vec![], + is_constructor: false, + span: rumoca_core::Span::DUMMY, + }); + dae.variables + .constants + .insert(VarName::new("building.modelicaNameBuilding"), building_name); + + let mut local_name = dae::Variable::new( + VarName::new("zone.modelicaNameBuilding"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ); + local_name.start = Some(rumoca_core::Expression::VarRef { + name: Reference::new("building.modelicaNameBuilding"), + subscripts: vec![], + span: rumoca_core::Span::DUMMY, + }); + dae.variables + .constants + .insert(VarName::new("zone.modelicaNameBuilding"), local_name); + + let env = build_runtime_parameter_tail_env(&dae, &[], 0.0) + .expect("getInstanceName string aliases should stay out of numeric env"); + + assert!(!env.vars.contains_key("building.modelicaNameBuilding")); + assert!(!env.vars.contains_key("zone.modelicaNameBuilding")); +} + +#[test] +fn build_runtime_parameter_tail_env_skips_string_field_access_alias_chain() { + let mut dae = dae::Dae::default(); + dae.metadata + .nonnumeric_variable_names + .push("building.modelicaNameBuilding".to_string()); + + let mut building_name = dae::Variable::new( + VarName::new("building.modelicaNameBuilding"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ); + building_name.start = Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("Root.building".to_string()), + span: rumoca_core::Span::DUMMY, + }); + dae.variables + .constants + .insert(VarName::new("building.modelicaNameBuilding"), building_name); + + let mut zone_name = dae::Variable::new( + VarName::new("zone.modelicaNameBuilding"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ); + zone_name.start = Some(rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::VarRef { + name: Reference::new("zone"), + subscripts: vec![], + span: rumoca_core::Span::DUMMY, + }), + field: "building".to_string(), + span: rumoca_core::Span::DUMMY, + }), + field: "modelicaNameBuilding".to_string(), + span: rumoca_core::Span::DUMMY, + }); + dae.variables + .constants + .insert(VarName::new("zone.modelicaNameBuilding"), zone_name); + + let env = build_runtime_parameter_tail_env(&dae, &[], 0.0) + .expect("string field-access aliases should stay out of numeric env"); + + assert!(!env.vars.contains_key("building.modelicaNameBuilding")); + assert!(!env.vars.contains_key("zone.modelicaNameBuilding")); +} + #[test] fn declared_slot_runtime_tail_env_advances_over_string_parameter_slots() { let mut dae = dae::Dae::default(); diff --git a/crates/rumoca-eval-dae/src/eval/tests/env_start_regressions.rs b/crates/rumoca-eval-dae/src/eval/tests/env_start_regressions.rs new file mode 100644 index 000000000..cdd407337 --- /dev/null +++ b/crates/rumoca-eval-dae/src/eval/tests/env_start_regressions.rs @@ -0,0 +1,135 @@ +use super::*; + +#[test] +fn test_eval_size_of_single_value_range_with_fill_bound() { + let env = VarEnv::::new(); + let fill = builtin( + rumoca_core::BuiltinFunction::Fill, + vec![lit(0.0), int_lit(0), int_lit(2)], + ); + let range = rumoca_core::Expression::Range { + start: Box::new(int_lit(2)), + step: None, + end: Box::new(builtin( + rumoca_core::BuiltinFunction::Size, + vec![fill, int_lit(2)], + )), + span: rumoca_core::Span::DUMMY, + }; + let size = builtin(rumoca_core::BuiltinFunction::Size, vec![range, int_lit(1)]); + + assert_eq!(eval_expr::(&size, &env), Ok(1.0)); +} + +#[test] +fn test_build_env_defaults_empty_external_table_bound_start() { + let mut dae = rumoca_ir_dae::Dae::default(); + let mut x = rumoca_ir_dae::Variable::new( + rumoca_core::VarName::new("table_u_min"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ); + let no_name = rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("NoName".to_string()), + span: rumoca_core::Span::DUMMY, + }; + let constructor = fn_call( + "ExternalCombiTimeTable", + vec![ + no_name.clone(), + no_name, + rumoca_core::Expression::Empty { + span: rumoca_core::Span::DUMMY, + }, + lit(0.0), + arr(vec![int_lit(2)], false), + int_lit(3), + int_lit(1), + ], + ); + x.start = Some(fn_call("getTimeTableTmin", vec![constructor])); + dae.variables.parameters.insert("table_u_min".into(), x); + + let env = build_runtime_parameter_tail_env(&dae, &[], 0.0).expect("test env should build"); + + assert_eq!(env.vars.get("table_u_min").copied(), Some(0.0)); +} + +#[test] +fn test_build_env_seeds_sum_size_start_for_string_array_literal() { + let substance_names = arr( + vec![ + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("N2".to_string()), + span: rumoca_core::Span::DUMMY, + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("O2".to_string()), + span: rumoca_core::Span::DUMMY, + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("H2O".to_string()), + span: rumoca_core::Span::DUMMY, + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("CO2".to_string()), + span: rumoca_core::Span::DUMMY, + }, + ], + false, + ); + let size = builtin(rumoca_core::BuiltinFunction::Size, vec![substance_names]); + let mut dae = rumoca_ir_dae::Dae::default(); + let mut n = rumoca_ir_dae::Variable::new( + rumoca_core::VarName::new("n"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ); + n.start = Some(builtin(rumoca_core::BuiltinFunction::Sum, vec![size])); + dae.variables.constants.insert("n".into(), n); + + let env = build_runtime_parameter_tail_env(&dae, &[], 0.0).expect("test env should build"); + + assert_eq!(env_value(&env, "n"), 4.0); +} + +#[test] +fn test_build_env_discrete_start_forward_ref_re_evaluates_and_preserves_pre_seed() { + clear_pre_values(); + + let mut dae = rumoca_ir_dae::Dae::default(); + + // Insert dependent start first to exercise forward-reference handling. + let mut a = rumoca_ir_dae::Variable::new( + rumoca_core::VarName::new("a"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ); + a.start = Some(dae_var("b")); + dae.variables + .discrete_valued + .insert(rumoca_core::VarName::new("a"), a); + + let mut b = rumoca_ir_dae::Variable::new( + rumoca_core::VarName::new("b"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ); + b.start = Some(dae_bool_lit(true)); + dae.variables + .discrete_valued + .insert(rumoca_core::VarName::new("b"), b); + + let env = build_env(&dae, &[], &[], 0.0).expect("test env should build"); + assert_eq!(env_value(&env, "b"), 1.0); + assert_eq!(env_value(&env, "a"), 1.0); + + // Pre-seeded values must take precedence over start expressions. + let mut pre_env = VarEnv::::new(); + pre_env.set("a", 0.0); + pre_env.set("b", 0.0); + seed_pre_values_from_env(&pre_env); + + let env_from_pre = build_env_with_runtime(&dae, &[], &[], 1.0, pre_env.runtime.clone()) + .expect("test env should build"); + assert_eq!(env_value(&env_from_pre, "a"), 0.0); + assert_eq!(env_value(&env_from_pre, "b"), 0.0); + + clear_pre_values(); +} diff --git a/crates/rumoca-eval-dae/src/eval/tests/mod.rs b/crates/rumoca-eval-dae/src/eval/tests/mod.rs index 557d4139a..97207b517 100644 --- a/crates/rumoca-eval-dae/src/eval/tests/mod.rs +++ b/crates/rumoca-eval-dae/src/eval/tests/mod.rs @@ -15,6 +15,7 @@ type VarName = rumoca_core::VarName; mod clock_and_tables; mod complex_array_selection; mod env_refresh; +mod env_start_regressions; mod pre_seed_regressions; mod runtime_specials_more; mod shift_sample_value_form; @@ -130,6 +131,178 @@ fn indexed_var(name: &str, indices: &[i64]) -> rumoca_core::Expression { } } +#[test] +fn var_ref_subscripted_matrix_slice_preserves_selected_row_shape() { + let mut env = VarEnv::::new(); + env.dims = Arc::new(IndexMap::from([("v_flow_rate".to_string(), vec![3, 3])])); + set_array_entries( + &mut env, + "v_flow_rate", + &[3, 3], + &[1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0], + ); + let explicit_row = rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("v_flow_rate"), + subscripts: vec![ + rumoca_core::Subscript::generated_index(1, rumoca_core::Span::DUMMY), + rumoca_core::Subscript::generated_colon(rumoca_core::Span::DUMMY), + ], + span: rumoca_core::Span::DUMMY, + }; + let prefix_row = rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("v_flow_rate"), + subscripts: vec![rumoca_core::Subscript::generated_index( + 2, + rumoca_core::Span::DUMMY, + )], + span: rumoca_core::Span::DUMMY, + }; + + assert_eq!( + eval_shaped_array_values::(&explicit_row, &env, 3), + Ok(vec![1.0, 2.0, 3.0]) + ); + assert_eq!( + eval_shaped_array_values::(&prefix_row, &env, 3), + Ok(vec![4.0, 5.0, 6.0]) + ); +} + +#[test] +fn builtin_sum_evaluates_matrix_row_selected_by_runtime_index_and_colon() { + let mut env = VarEnv::::new(); + env.set("floorIndex", 2.0); + env.dims = Arc::new(IndexMap::from([("mAirFloRat".to_string(), vec![3, 4])])); + set_array_entries( + &mut env, + "mAirFloRat", + &[3, 4], + &[ + 1.0, 2.0, 3.0, 4.0, 10.0, 20.0, 30.0, 40.0, 100.0, 200.0, 300.0, 400.0, + ], + ); + let row = rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("mAirFloRat"), + subscripts: vec![ + rumoca_core::Subscript::Expr { + expr: Box::new(var("floorIndex")), + span: rumoca_core::Span::DUMMY, + }, + rumoca_core::Subscript::generated_colon(rumoca_core::Span::DUMMY), + ], + span: rumoca_core::Span::DUMMY, + }; + let expr = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Sum, + args: vec![row], + span: rumoca_core::Span::DUMMY, + }; + + assert_eq!(eval_expr::(&expr, &env), Ok(100.0)); +} + +#[test] +fn set_state_array_field_scalar_projection_accepts_range_slice_argument() { + let mut env = VarEnv::::new(); + env.dims = Arc::new(IndexMap::from([("X_start".to_string(), vec![2])])); + set_array_entries(&mut env, "X_start", &[2], &[0.25, 0.75]); + + let x_slice = rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("X_start"), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(rumoca_core::Expression::Range { + start: Box::new(int_lit(1)), + step: None, + end: Box::new(rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + args: vec![arr(vec![lit(0.0), lit(0.0)], false), int_lit(1)], + span: rumoca_core::Span::DUMMY, + }), + rhs: Box::new(int_lit(1)), + span: rumoca_core::Span::DUMMY, + }), + span: rumoca_core::Span::DUMMY, + }), + span: rumoca_core::Span::DUMMY, + }], + span: rumoca_core::Span::DUMMY, + }; + let state_x = rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Medium.setState_pTX"), + args: vec![ + named_ctor_arg("T", lit(293.15)), + named_ctor_arg("p", lit(101325.0)), + named_ctor_arg("X", x_slice), + ], + is_constructor: false, + span: rumoca_core::Span::DUMMY, + }), + field: "X".to_string(), + span: rumoca_core::Span::DUMMY, + }; + + assert_eq!(eval_expr::(&state_x, &env), Ok(0.25)); +} + +#[test] +fn user_function_array_input_binds_range_slice_before_scalar_path() { + let mut env = VarEnv::::new(); + env.dims = Arc::new(IndexMap::from([("X_start".to_string(), vec![2])])); + set_array_entries(&mut env, "X_start", &[2], &[0.25, 0.75]); + + let mut functions = IndexMap::new(); + let mut function = Function::new("Pkg.firstMassFraction", rumoca_core::Span::DUMMY); + function.add_input( + FunctionParam::new("X", "Real", rumoca_core::Span::source_free_serde_default()) + .with_dims(vec![0]) + .with_shape_expr(vec![Subscript::generated_colon(rumoca_core::Span::DUMMY)]), + ); + function.add_output(FunctionParam::new( + "y", + "Real", + rumoca_core::Span::source_free_serde_default(), + )); + function.body = vec![Statement::Assignment { + comp: comp_ref("y"), + value: index_expr(var("X"), 1), + span: rumoca_core::Span::DUMMY, + }]; + functions.insert("Pkg.firstMassFraction".to_string(), function); + env.functions = Arc::new(functions); + + let x_slice = rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("X_start"), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(rumoca_core::Expression::Range { + start: Box::new(int_lit(1)), + step: None, + end: Box::new(rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + args: vec![arr(vec![lit(0.0), lit(0.0)], false), int_lit(1)], + span: rumoca_core::Span::DUMMY, + }), + rhs: Box::new(int_lit(1)), + span: rumoca_core::Span::DUMMY, + }), + span: rumoca_core::Span::DUMMY, + }), + span: rumoca_core::Span::DUMMY, + }], + span: rumoca_core::Span::DUMMY, + }; + assert_eq!(eval_array_values::(&x_slice, &env), Ok(vec![0.25])); + + assert_eq!( + eval_expr::(&fn_call("Pkg.firstMassFraction", vec![x_slice]), &env), + Ok(0.25) + ); +} + fn comp_ref(name: &str) -> rumoca_core::ComponentReference { rumoca_core::ComponentReference { local: false, @@ -144,21 +317,192 @@ fn comp_ref(name: &str) -> rumoca_core::ComponentReference { } fn comp_ref_index(name: &str, index: i64) -> rumoca_core::ComponentReference { + comp_ref_indices(name, &[index]) +} + +fn comp_ref_indices(name: &str, indices: &[i64]) -> rumoca_core::ComponentReference { rumoca_core::ComponentReference { local: false, span: rumoca_core::Span::DUMMY, parts: vec![rumoca_core::ComponentRefPart { ident: name.to_string(), span: rumoca_core::Span::DUMMY, - subs: vec![rumoca_core::Subscript::generated_index( - index, - rumoca_core::Span::DUMMY, - )], + subs: indices + .iter() + .copied() + .map(|index| { + rumoca_core::Subscript::generated_index(index, rumoca_core::Span::DUMMY) + }) + .collect(), }], def_id: None, } } +#[test] +fn singleton_array_output_collects_dense_index_assignment() { + let mut env = VarEnv::::new(); + let mut function = Function::new("Pkg.singletonOutput", rumoca_core::Span::DUMMY); + function.add_output( + FunctionParam::new("c1", "Real", rumoca_core::Span::source_free_serde_default()) + .with_dims(vec![1]), + ); + function.body = vec![Statement::Assignment { + comp: comp_ref_index("c1", 1), + value: lit(2.5), + span: rumoca_core::Span::DUMMY, + }]; + env.functions = Arc::new(IndexMap::from([( + "Pkg.singletonOutput".to_string(), + function, + )])); + + assert_eq!( + eval_user_function_array_output_pub::( + &rumoca_core::VarName::new("Pkg.singletonOutput"), + &[], + &env, + ), + Ok(vec![2.5]) + ); + assert_eq!( + eval_selected_function_output_pub::( + &rumoca_core::VarName::new("Pkg.singletonOutput"), + "c1", + &[1], + &[], + &env, + ), + Ok(2.5) + ); +} + +#[test] +fn singleton_array_values_accept_base_alias_binding() { + let mut env = VarEnv::::new(); + env.dims = Arc::new(IndexMap::from([("den1".to_string(), vec![1])])); + env.set("den1", 1.5962800638268535); + + assert_eq!( + eval_array_values::(&var("den1"), &env), + Ok(vec![1.5962800638268535]) + ); +} + +#[test] +fn matrix_array_output_collects_dense_multidimensional_assignment() { + let mut env = VarEnv::::new(); + let mut function = Function::new("Pkg.matrixOutput", rumoca_core::Span::DUMMY); + function.add_output( + FunctionParam::new("c2", "Real", rumoca_core::Span::source_free_serde_default()) + .with_dims(vec![1, 2]), + ); + function.body = vec![ + Statement::Assignment { + comp: comp_ref_indices("c2", &[1, 1]), + value: lit(3.0), + span: rumoca_core::Span::DUMMY, + }, + Statement::Assignment { + comp: comp_ref_indices("c2", &[1, 2]), + value: lit(4.0), + span: rumoca_core::Span::DUMMY, + }, + ]; + env.functions = Arc::new(IndexMap::from([("Pkg.matrixOutput".to_string(), function)])); + + assert_eq!( + eval_user_function_array_output_pub::( + &rumoca_core::VarName::new("Pkg.matrixOutput"), + &[], + &env, + ), + Ok(vec![3.0, 4.0]) + ); +} + +#[test] +fn function_local_shape_can_depend_on_output_shape() { + let mut env = VarEnv::::new(); + let mut producer = Function::new("Pkg.producer", rumoca_core::Span::DUMMY); + producer.add_output( + FunctionParam::new("c1", "Real", rumoca_core::Span::source_free_serde_default()) + .with_dims(vec![1]), + ); + producer.body = vec![Statement::Assignment { + comp: comp_ref_index("c1", 1), + value: lit(7.0), + span: rumoca_core::Span::DUMMY, + }]; + + let mut parent = Function::new("Pkg.parent", rumoca_core::Span::DUMMY); + parent.add_input(FunctionParam::new( + "order", + "Integer", + rumoca_core::Span::source_free_serde_default(), + )); + parent.add_output( + FunctionParam::new("cr", "Real", rumoca_core::Span::source_free_serde_default()) + .with_dims(vec![0]) + .with_shape_expr(vec![Subscript::generated_expr( + Box::new(builtin( + BuiltinFunction::Mod, + vec![var("order"), int_lit(2)], + )), + rumoca_core::Span::DUMMY, + )]), + ); + parent.add_output(FunctionParam::new( + "y", + "Real", + rumoca_core::Span::source_free_serde_default(), + )); + parent.add_local( + FunctionParam::new( + "den1", + "Real", + rumoca_core::Span::source_free_serde_default(), + ) + .with_dims(vec![0]) + .with_shape_expr(vec![Subscript::generated_expr( + Box::new(builtin(BuiltinFunction::Size, vec![var("cr"), int_lit(1)])), + rumoca_core::Span::DUMMY, + )]), + ); + parent.body = vec![ + Statement::FunctionCall { + comp: rumoca_core::ComponentReference::from_flat_segments( + "Pkg.producer", + rumoca_core::Span::DUMMY, + None, + ), + args: vec![], + outputs: vec![comp_ref("den1")], + span: rumoca_core::Span::DUMMY, + }, + Statement::Assignment { + comp: comp_ref("y"), + value: indexed_var("den1", &[1]), + span: rumoca_core::Span::DUMMY, + }, + ]; + env.functions = Arc::new(IndexMap::from([ + ("Pkg.producer".to_string(), producer), + ("Pkg.parent".to_string(), parent), + ])); + + assert_eq!( + eval_selected_function_output_pub::( + &rumoca_core::VarName::new("Pkg.parent"), + "y", + &[], + &[int_lit(3)], + &env, + ), + Ok(7.0) + ); +} + fn arr(elements: Vec, is_matrix: bool) -> rumoca_core::Expression { rumoca_core::Expression::Array { elements, @@ -813,6 +1157,67 @@ fn test_eval_index_on_flattened_env_array_with_dims() { assert!((value - 6.0).abs() < 1e-12); } +#[test] +fn test_eval_array_values_var_ref_colon_slice_from_env() { + let mut env = VarEnv::::new(); + set_array_entries(&mut env, "v", &[3], &[10.0, 20.0, 30.0]); + + let expr = rumoca_core::Expression::VarRef { + name: Reference::new("v"), + subscripts: vec![Subscript::generated_colon(rumoca_core::Span::DUMMY)], + span: rumoca_core::Span::DUMMY, + }; + + assert_eq!( + eval_array_values::(&expr, &env), + Ok(vec![10.0, 20.0, 30.0]) + ); +} + +#[test] +fn test_eval_array_values_var_ref_range_slice_from_env() { + let mut env = VarEnv::::new(); + set_array_entries(&mut env, "v", &[4], &[10.0, 20.0, 30.0, 40.0]); + + let expr = rumoca_core::Expression::VarRef { + name: Reference::new("v"), + subscripts: vec![Subscript::generated_expr( + Box::new(rumoca_core::Expression::Range { + start: Box::new(int_lit(2)), + step: None, + end: Box::new(int_lit(3)), + span: rumoca_core::Span::DUMMY, + }), + rumoca_core::Span::DUMMY, + )], + span: rumoca_core::Span::DUMMY, + }; + + assert_eq!(eval_array_values::(&expr, &env), Ok(vec![20.0, 30.0])); +} + +#[test] +fn test_eval_array_values_index_range_slice_from_array_expr() { + let expr = rumoca_core::Expression::Index { + base: Box::new(arr(vec![lit(1.0), lit(2.0), lit(3.0), lit(4.0)], false)), + subscripts: vec![Subscript::generated_expr( + Box::new(rumoca_core::Expression::Range { + start: Box::new(int_lit(2)), + step: None, + end: Box::new(int_lit(3)), + span: rumoca_core::Span::DUMMY, + }), + rumoca_core::Span::DUMMY, + )], + span: rumoca_core::Span::DUMMY, + }; + + assert_eq!( + eval_array_values::(&expr, &VarEnv::new()), + Ok(vec![2.0, 3.0]) + ); +} + #[test] fn test_eval_index_on_transposed_env_matrix_with_dims() { let mut env = VarEnv::::new(); @@ -1157,6 +1562,76 @@ fn test_eval_function_dynamic_vector_input_binds_expression_shape() { ); } +#[test] +fn test_eval_record_constructor_input_skips_string_and_binds_array_field() { + let mut env = VarEnv::::new(); + let mut funcs = IndexMap::new(); + + let mut state = Function::new("Pkg.State", rumoca_core::Span::DUMMY); + state.is_constructor = true; + state.add_input(FunctionParam::new( + "p", + "Real", + rumoca_core::Span::source_free_serde_default(), + )); + state.add_input( + FunctionParam::new("X", "Real", rumoca_core::Span::source_free_serde_default()) + .with_dims(vec![2]), + ); + state.add_input(FunctionParam::new( + "mediumName", + "String", + rumoca_core::Span::source_free_serde_default(), + )); + funcs.insert("Pkg.State".to_string(), state); + + let mut metric = Function::new("Pkg.metric", rumoca_core::Span::DUMMY); + metric.add_input( + FunctionParam::new( + "st", + "State", + rumoca_core::Span::source_free_serde_default(), + ) + .with_type_class(rumoca_core::ClassType::Record), + ); + metric.add_output( + FunctionParam::new("y", "Real", rumoca_core::Span::source_free_serde_default()) + .with_default(binop( + rumoca_core::OpBinary::Add, + var("st.p"), + indexed_var("st.X", &[2]), + )), + ); + metric.body = vec![Statement::Empty { + span: rumoca_core::Span::DUMMY, + }]; + funcs.insert("Pkg.metric".to_string(), metric); + env.functions = Arc::new(funcs); + + let state_arg = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.State"), + args: vec![ + named_ctor_arg("p", lit(101325.0)), + named_ctor_arg("X", arr(vec![lit(0.42), lit(0.58)], false)), + named_ctor_arg( + "mediumName", + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("MoistAir".to_string()), + span: rumoca_core::Span::DUMMY, + }, + ), + ], + is_constructor: true, + span: rumoca_core::Span::DUMMY, + }; + + assert!( + (eval_expr::(&fn_call("Pkg.metric", vec![state_arg]), &env).unwrap() - 101325.58) + .abs() + < 1e-9 + ); +} + #[test] fn test_eval_array_values_cross_product() { let mut env = VarEnv::::new(); @@ -1482,44 +1957,37 @@ fn test_build_env_seeds_fill_start_sized_by_string_array_literal() { } #[test] -fn test_build_env_discrete_start_forward_ref_re_evaluates_and_preserves_pre_seed() { - clear_pre_values(); - - let mut dae = rumoca_ir_dae::Dae::default(); - - // Insert dependent start first to exercise forward-reference handling. - let mut a = rumoca_ir_dae::Variable::new( - rumoca_core::VarName::new("a"), - rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), +fn test_build_env_accepts_zero_length_fill_sized_by_string_fill() { + let string_names = builtin( + rumoca_core::BuiltinFunction::Fill, + vec![ + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String(String::new()), + span: rumoca_core::Span::DUMMY, + }, + int_lit(0), + ], ); - a.start = Some(dae_var("b")); - dae.variables - .discrete_valued - .insert(rumoca_core::VarName::new("a"), a); - - let mut b = rumoca_ir_dae::Variable::new( - rumoca_core::VarName::new("b"), + let size = builtin( + rumoca_core::BuiltinFunction::Size, + vec![string_names, int_lit(1)], + ); + let mut dae = rumoca_ir_dae::Dae::default(); + let mut x = rumoca_ir_dae::Variable::new( + rumoca_core::VarName::new("C_start"), rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), ); - b.start = Some(dae_bool_lit(true)); - dae.variables - .discrete_valued - .insert(rumoca_core::VarName::new("b"), b); - - let env = build_env(&dae, &[], &[], 0.0).expect("test env should build"); - assert_eq!(env_value(&env, "b"), 1.0); - assert_eq!(env_value(&env, "a"), 1.0); - - // Pre-seeded values must take precedence over start expressions. - let mut pre_env = VarEnv::::new(); - pre_env.set("a", 0.0); - pre_env.set("b", 0.0); - seed_pre_values_from_env(&pre_env); + x.dims = vec![0]; + x.start = Some(builtin( + rumoca_core::BuiltinFunction::Fill, + vec![lit(0.0), size], + )); + dae.variables.parameters.insert("C_start".into(), x); - let env_from_pre = build_env_with_runtime(&dae, &[], &[], 1.0, pre_env.runtime.clone()) - .expect("test env should build"); - assert_eq!(env_value(&env_from_pre, "a"), 0.0); - assert_eq!(env_value(&env_from_pre, "b"), 0.0); + let env = build_runtime_parameter_tail_env(&dae, &[], 0.0).expect("test env should build"); - clear_pre_values(); + assert_eq!( + eval_shaped_array_values::(&var("C_start"), &env, 0), + Ok(Vec::new()) + ); } diff --git a/crates/rumoca-eval-dae/src/eval/tests/runtime_specials_more.rs b/crates/rumoca-eval-dae/src/eval/tests/runtime_specials_more.rs index 19411a6fd..25cb0ba6b 100644 --- a/crates/rumoca-eval-dae/src/eval/tests/runtime_specials_more.rs +++ b/crates/rumoca-eval-dae/src/eval/tests/runtime_specials_more.rs @@ -76,6 +76,77 @@ fn test_string_is_empty_runtime_special() { assert_eq!(non_empty, 0.0); } +fn string_lit(value: &str) -> rumoca_core::Expression { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String(value.to_string()), + span: rumoca_core::Span::DUMMY, + } +} + +#[test] +fn test_string_findlast_evaluates_string_parameter_start_expr_with_named_case_flag() { + let mut env = VarEnv::::new(); + env.start_exprs = std::sync::Arc::new(indexmap::indexmap! { + "fileName".to_string() => string_lit("weather/USA_IL_Chicago.csv"), + }); + + let is_csv_ext = rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Eq, + lhs: Box::new(rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: Box::new(fn_call( + "Modelica.Utilities.Strings.findLast", + vec![ + var("fileName"), + string_lit(".csv"), + named_ctor_arg("caseSensitive", bool_lit(false)), + ], + )), + rhs: Box::new(int_lit(3)), + span: rumoca_core::Span::DUMMY, + }), + rhs: Box::new(fn_call( + "Modelica.Utilities.Strings.length", + vec![var("fileName")], + )), + span: rumoca_core::Span::DUMMY, + }; + + assert_eq!(eval_expr::(&is_csv_ext, &env), Ok(1.0)); +} + +#[test] +fn test_string_find_supports_start_index_and_case_insensitive_search() { + let env = VarEnv::::new(); + + let find = eval_expr::( + &fn_call( + "Modelica.Utilities.Strings.find", + vec![ + string_lit("Alpha/BETA/beta"), + string_lit("beta"), + named_ctor_arg("startIndex", int_lit(8)), + named_ctor_arg("caseSensitive", bool_lit(false)), + ], + ), + &env, + ); + assert_eq!(find, Ok(12.0)); + + let find_last = eval_expr::( + &fn_call( + "Modelica.Utilities.Strings.findLast", + vec![ + named_ctor_arg("string", string_lit("a/b/c")), + named_ctor_arg("searchString", string_lit("/")), + named_ctor_arg("startIndex", int_lit(4)), + ], + ), + &env, + ); + assert_eq!(find_last, Ok(4.0)); +} + #[test] fn test_random_runtime_special_seed_and_stream() { let env = VarEnv::::new(); diff --git a/crates/rumoca-eval-dae/src/eval/tests/scalar_eval_tests.rs b/crates/rumoca-eval-dae/src/eval/tests/scalar_eval_tests.rs index 71842cb21..3f20d8a31 100644 --- a/crates/rumoca-eval-dae/src/eval/tests/scalar_eval_tests.rs +++ b/crates/rumoca-eval-dae/src/eval/tests/scalar_eval_tests.rs @@ -238,6 +238,175 @@ fn test_eval_if_false() { assert_eq!(eval_expr_value::(&expr, &VarEnv::new()), 2.0); } +#[test] +fn broadcast_start_value_uses_selected_if_branch() { + let expr = rumoca_core::Expression::If { + branches: vec![(bool_lit(false), arr(vec![lit(1.0), lit(2.0)], false))], + else_branch: Box::new(arr(vec![lit(0.0)], false)), + span: rumoca_core::Span::DUMMY, + }; + + assert!(can_broadcast_start_value(&expr, &VarEnv::new())); +} + +#[test] +fn test_checked_eval_field_access_projects_selected_if_branch() { + let mut record_ctor = rumoca_core::Function::new("Pkg.Record", rumoca_core::Span::DUMMY); + record_ctor.add_input(rumoca_core::FunctionParam::new( + "x", + "Real", + rumoca_core::Span::source_free_serde_default(), + )); + let mut env = VarEnv::::new(); + env.functions = Arc::new(IndexMap::from([("Pkg.Record".to_string(), record_ctor)])); + + let then_record = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.Record"), + args: vec![named_ctor_arg("x", lit(1.25))], + is_constructor: true, + span: rumoca_core::Span::DUMMY, + }; + let else_record = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.Record"), + args: vec![named_ctor_arg("x", lit(2.5))], + is_constructor: true, + span: rumoca_core::Span::DUMMY, + }; + let expr = rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::If { + branches: vec![(bool_lit(false), then_record)], + else_branch: Box::new(else_record), + span: rumoca_core::Span::DUMMY, + }), + field: "x".to_string(), + span: rumoca_core::Span::DUMMY, + }; + + assert_eq!(eval_expr::(&expr, &env), Ok(2.5)); +} + +#[test] +fn test_checked_eval_field_access_projects_nested_constructor_field() { + let mut inner_ctor = rumoca_core::Function::new("Pkg.Inner", rumoca_core::Span::DUMMY); + inner_ctor.add_input(rumoca_core::FunctionParam::new( + "x", + "Real", + rumoca_core::Span::source_free_serde_default(), + )); + let mut outer_ctor = rumoca_core::Function::new("Pkg.Outer", rumoca_core::Span::DUMMY); + outer_ctor.add_input(rumoca_core::FunctionParam::new( + "inner", + "Pkg.Inner", + rumoca_core::Span::source_free_serde_default(), + )); + let mut env = VarEnv::::new(); + env.functions = Arc::new(IndexMap::from([ + ("Pkg.Inner".to_string(), inner_ctor), + ("Pkg.Outer".to_string(), outer_ctor), + ])); + + let inner_record = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.Inner"), + args: vec![named_ctor_arg("x", lit(2.5))], + is_constructor: true, + span: rumoca_core::Span::DUMMY, + }; + let outer_record = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.Outer"), + args: vec![named_ctor_arg("inner", inner_record)], + is_constructor: true, + span: rumoca_core::Span::DUMMY, + }; + let expr = rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FieldAccess { + base: Box::new(outer_record), + field: "inner".to_string(), + span: rumoca_core::Span::DUMMY, + }), + field: "x".to_string(), + span: rumoca_core::Span::DUMMY, + }; + + assert_eq!(eval_expr::(&expr, &env), Ok(2.5)); +} + +#[test] +fn test_checked_eval_field_access_resolves_indexed_component_path() { + let mut env = VarEnv::::new(); + env.set("unit[1].settings.tolerance", 1.0e-6); + let expr = rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::Index { + base: Box::new(var("unit")), + subscripts: vec![rumoca_core::Subscript::generated_index( + 1, + rumoca_core::Span::DUMMY, + )], + span: rumoca_core::Span::DUMMY, + }), + field: "settings".to_string(), + span: rumoca_core::Span::DUMMY, + }), + field: "tolerance".to_string(), + span: rumoca_core::Span::DUMMY, + }; + + assert_eq!(eval_expr::(&expr, &env), Ok(1.0e-6)); +} + +#[test] +fn test_checked_eval_var_ref_singleton_range_slice_as_scalar() { + let mut env = VarEnv::::new(); + env.vars.insert("x[1]".to_string(), 0.42); + env.vars.insert("x[2]".to_string(), 0.99); + env.dims = Arc::new(IndexMap::from([("x".to_string(), vec![2])])); + + let expr = rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("x"), + subscripts: vec![rumoca_core::Subscript::generated_expr( + Box::new(rumoca_core::Expression::Range { + start: Box::new(int_lit(1)), + step: None, + end: Box::new(int_lit(1)), + span: rumoca_core::Span::DUMMY, + }), + rumoca_core::Span::DUMMY, + )], + span: rumoca_core::Span::DUMMY, + }; + + assert_eq!(eval_expr::(&expr, &env), Ok(0.42)); +} + +#[test] +fn test_checked_eval_var_ref_subscript_expressions() { + let mut env = VarEnv::::new(); + env.vars.insert("x[1]".to_string(), 42.0); + env.vars.insert("x[2]".to_string(), 99.0); + env.vars.insert("n".to_string(), 2.0); + env.dims = Arc::new(IndexMap::from([("x".to_string(), vec![2])])); + + let arithmetic_subscript = rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("x"), + subscripts: vec![rumoca_core::Subscript::generated_expr( + Box::new(binop(rumoca_core::OpBinary::Sub, int_lit(2), int_lit(1))), + rumoca_core::Span::DUMMY, + )], + span: rumoca_core::Span::DUMMY, + }; + let variable_subscript = rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("x"), + subscripts: vec![rumoca_core::Subscript::generated_expr( + Box::new(var("n")), + rumoca_core::Span::DUMMY, + )], + span: rumoca_core::Span::DUMMY, + }; + + assert_eq!(eval_expr::(&arithmetic_subscript, &env), Ok(42.0)); + assert_eq!(eval_expr::(&variable_subscript, &env), Ok(99.0)); +} + #[test] fn test_eval_comparison() { let lt = binop(rumoca_core::OpBinary::Lt, lit(1.0), lit(2.0)); diff --git a/crates/rumoca-eval-dae/src/eval/tests/scalar_eval_tests/function_call_tests.rs b/crates/rumoca-eval-dae/src/eval/tests/scalar_eval_tests/function_call_tests.rs index fd54dc452..4b7929a84 100644 --- a/crates/rumoca-eval-dae/src/eval/tests/scalar_eval_tests/function_call_tests.rs +++ b/crates/rumoca-eval-dae/src/eval/tests/scalar_eval_tests/function_call_tests.rs @@ -1,5 +1,415 @@ use super::*; +fn function_param(name: &str, type_name: &str) -> rumoca_core::FunctionParam { + rumoca_core::FunctionParam::new( + name, + type_name, + rumoca_core::Span::source_free_serde_default(), + ) +} + +fn shape_expr_param( + name: &str, + type_name: &str, + shape_expr: rumoca_core::Subscript, +) -> rumoca_core::FunctionParam { + function_param(name, type_name) + .with_dims(vec![0]) + .with_shape_expr(vec![shape_expr]) +} + +fn pkg_record_shape_env() -> VarEnv { + let mut env = VarEnv::::new(); + let mut funcs = IndexMap::new(); + + let mut state_ctor = rumoca_core::Function::new("Pkg.State", rumoca_core::Span::DUMMY); + state_ctor.add_input(function_param("p", "Real")); + state_ctor.add_input(shape_expr_param( + "X", + "Real", + rumoca_core::Subscript::generated_expr(Box::new(var("nX")), rumoca_core::Span::DUMMY), + )); + state_ctor.is_constructor = true; + funcs.insert("Pkg.State".to_string(), state_ctor); + + let mut set_state = rumoca_core::Function::new("Pkg.setState", rumoca_core::Span::DUMMY); + set_state.add_input(shape_expr_param( + "X", + "Real", + rumoca_core::Subscript::generated_colon(rumoca_core::Span::DUMMY), + )); + set_state.add_output( + function_param("state", "State").with_type_class(rumoca_core::ClassType::Record), + ); + set_state.body = vec![rumoca_core::Statement::Assignment { + comp: comp_ref("state"), + value: rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.State"), + args: vec![ + named_ctor_arg("p", lit(101325.0)), + named_ctor_arg("X", var("X")), + ], + is_constructor: true, + span: rumoca_core::Span::DUMMY, + }, + span: rumoca_core::Span::DUMMY, + }]; + funcs.insert("Pkg.setState".to_string(), set_state); + + let mut density = rumoca_core::Function::new("Pkg.density", rumoca_core::Span::DUMMY); + density.add_input(shape_expr_param( + "state_X", + "Real", + rumoca_core::Subscript::generated_expr(Box::new(var("nX")), rumoca_core::Span::DUMMY), + )); + density.add_output(function_param("d", "Real").with_default(var("nX"))); + density.body = vec![rumoca_core::Statement::Empty { + span: rumoca_core::Span::DUMMY, + }]; + funcs.insert("Pkg.density".to_string(), density); + + let mut outer = rumoca_core::Function::new("Pkg.outer", rumoca_core::Span::DUMMY); + outer.add_input(shape_expr_param( + "X", + "Real", + rumoca_core::Subscript::generated_colon(rumoca_core::Span::DUMMY), + )); + outer.add_output(function_param("d", "Real").with_default(fn_call( + "Pkg.density", + vec![field(fn_call("Pkg.setState", vec![var("X")]), "X")], + ))); + outer.body = vec![rumoca_core::Statement::Empty { + span: rumoca_core::Span::DUMMY, + }]; + funcs.insert("Pkg.outer".to_string(), outer); + + env.functions = Arc::new(funcs); + env +} + +fn buildings_air_state_ctor() -> rumoca_core::Function { + let mut state_ctor = rumoca_core::Function::new( + "Buildings.Media.Air.ThermodynamicState", + rumoca_core::Span::DUMMY, + ); + state_ctor.add_input(function_param("p", "Real")); + state_ctor.add_input(function_param("T", "Real")); + state_ctor.add_input(shape_expr_param( + "X", + "Real", + rumoca_core::Subscript::generated_expr(Box::new(var("nX")), rumoca_core::Span::DUMMY), + )); + state_ctor.is_constructor = true; + state_ctor +} + +fn thermodynamic_state_call(x: rumoca_core::Expression) -> rumoca_core::Expression { + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Buildings.Media.Air.ThermodynamicState"), + args: vec![ + named_ctor_arg("p", var("p")), + named_ctor_arg("T", var("T")), + named_ctor_arg("X", x), + ], + is_constructor: true, + span: rumoca_core::Span::DUMMY, + } +} + +fn padded_mass_fraction_expr() -> rumoca_core::Expression { + builtin( + rumoca_core::BuiltinFunction::Cat, + vec![ + int_lit(1), + var("X"), + arr( + vec![binop( + rumoca_core::OpBinary::Sub, + lit(1.0), + builtin(rumoca_core::BuiltinFunction::Sum, vec![var("X")]), + )], + false, + ), + ], + ) +} + +fn buildings_air_set_state_ptx_with_x_padding() -> rumoca_core::Function { + let mut set_state = + rumoca_core::Function::new("Buildings.Media.Air.setState_pTX", rumoca_core::Span::DUMMY); + set_state.add_input(function_param("p", "Real")); + set_state.add_input(function_param("T", "Real")); + set_state.add_input(shape_expr_param( + "X", + "Real", + rumoca_core::Subscript::generated_colon(rumoca_core::Span::DUMMY), + )); + set_state.add_output( + function_param("state", "ThermodynamicState") + .with_type_class(rumoca_core::ClassType::Record), + ); + set_state.body = vec![rumoca_core::Statement::Assignment { + comp: comp_ref("state"), + value: rumoca_core::Expression::If { + branches: vec![( + binop( + rumoca_core::OpBinary::Eq, + builtin( + rumoca_core::BuiltinFunction::Size, + vec![var("X"), int_lit(1)], + ), + int_lit(2), + ), + thermodynamic_state_call(var("X")), + )], + else_branch: Box::new(thermodynamic_state_call(padded_mass_fraction_expr())), + span: rumoca_core::Span::DUMMY, + }, + span: rumoca_core::Span::DUMMY, + }]; + set_state +} + +#[test] +fn test_eval_user_function_binds_string_array_input_shape_for_size() { + let mut env = VarEnv::::new(); + let mut funcs = IndexMap::new(); + let mut f = rumoca_core::Function::new("Pkg.nSubstances", rumoca_core::Span::DUMMY); + f.add_input( + rumoca_core::FunctionParam::new( + "substanceNames", + "String", + rumoca_core::Span::source_free_serde_default(), + ) + .with_dims(vec![0]), + ); + f.add_output( + rumoca_core::FunctionParam::new( + "n", + "Integer", + rumoca_core::Span::source_free_serde_default(), + ) + .with_default(builtin( + rumoca_core::BuiltinFunction::Size, + vec![var("substanceNames"), int_lit(1)], + )), + ); + f.body = vec![rumoca_core::Statement::Empty { + span: rumoca_core::Span::DUMMY, + }]; + funcs.insert("Pkg.nSubstances".to_string(), f); + env.functions = Arc::new(funcs); + + let names = arr( + vec![ + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("N2".to_string()), + span: rumoca_core::Span::DUMMY, + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("O2".to_string()), + span: rumoca_core::Span::DUMMY, + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("H2O".to_string()), + span: rumoca_core::Span::DUMMY, + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("CO2".to_string()), + span: rumoca_core::Span::DUMMY, + }, + ], + false, + ); + + let expr = fn_call("Pkg.nSubstances", vec![names]); + + assert_eq!(eval_expr::(&expr, &env), Ok(4.0)); +} + +#[test] +fn test_eval_user_function_binds_shape_expr_variable_from_array_input() { + let mut env = VarEnv::::new(); + let mut funcs = IndexMap::new(); + let mut f = rumoca_core::Function::new("Pkg.massFractionCount", rumoca_core::Span::DUMMY); + f.add_input( + rumoca_core::FunctionParam::new( + "X", + "Real", + rumoca_core::Span::source_free_serde_default(), + ) + .with_dims(vec![0]) + .with_shape_expr(vec![rumoca_core::Subscript::generated_expr( + Box::new(var("nX")), + rumoca_core::Span::DUMMY, + )]), + ); + f.add_output( + rumoca_core::FunctionParam::new( + "n", + "Integer", + rumoca_core::Span::source_free_serde_default(), + ) + .with_default(var("nX")), + ); + f.body = vec![rumoca_core::Statement::Empty { + span: rumoca_core::Span::DUMMY, + }]; + funcs.insert("Pkg.massFractionCount".to_string(), f); + env.functions = Arc::new(funcs); + + let expr = fn_call( + "Pkg.massFractionCount", + vec![arr(vec![lit(0.7), lit(0.3)], false)], + ); + + assert_eq!(eval_expr::(&expr, &env), Ok(2.0)); +} + +#[test] +fn test_eval_nested_user_function_preserves_shape_expr_variable_from_caller_input() { + let mut env = VarEnv::::new(); + let mut funcs = IndexMap::new(); + + let mut inner = rumoca_core::Function::new("Pkg.inner", rumoca_core::Span::DUMMY); + inner.add_input( + rumoca_core::FunctionParam::new( + "state_X", + "Real", + rumoca_core::Span::source_free_serde_default(), + ) + .with_dims(vec![0]) + .with_shape_expr(vec![rumoca_core::Subscript::generated_expr( + Box::new(var("nX")), + rumoca_core::Span::DUMMY, + )]), + ); + inner.add_output( + rumoca_core::FunctionParam::new( + "d", + "Real", + rumoca_core::Span::source_free_serde_default(), + ) + .with_default(var("nX")), + ); + inner.body = vec![rumoca_core::Statement::Empty { + span: rumoca_core::Span::DUMMY, + }]; + funcs.insert("Pkg.inner".to_string(), inner); + + let mut outer = rumoca_core::Function::new("Pkg.outer", rumoca_core::Span::DUMMY); + outer.add_input( + rumoca_core::FunctionParam::new( + "X", + "Real", + rumoca_core::Span::source_free_serde_default(), + ) + .with_dims(vec![0]) + .with_shape_expr(vec![rumoca_core::Subscript::generated_colon( + rumoca_core::Span::DUMMY, + )]), + ); + outer.add_output( + rumoca_core::FunctionParam::new( + "d", + "Real", + rumoca_core::Span::source_free_serde_default(), + ) + .with_default(fn_call("Pkg.inner", vec![var("X")])), + ); + outer.body = vec![rumoca_core::Statement::Empty { + span: rumoca_core::Span::DUMMY, + }]; + funcs.insert("Pkg.outer".to_string(), outer); + env.functions = Arc::new(funcs); + + let expr = fn_call("Pkg.outer", vec![arr(vec![lit(0.7), lit(0.3)], false)]); + + assert_eq!(eval_expr::(&expr, &env), Ok(2.0)); +} + +#[test] +fn test_eval_record_function_array_field_preserves_shape_expr_variable() { + let env = pkg_record_shape_env(); + let expr = fn_call("Pkg.outer", vec![arr(vec![lit(0.7), lit(0.3)], false)]); + + assert_eq!(eval_expr::(&expr, &env), Ok(2.0)); +} + +#[test] +fn test_eval_set_state_ptx_x_field_binds_shape_expr_variable() { + let mut env = VarEnv::::new(); + let mut funcs = IndexMap::new(); + + let mut density = + rumoca_core::Function::new("Buildings.Media.Air.density", rumoca_core::Span::DUMMY); + density.add_input( + rumoca_core::FunctionParam::new( + "state_X", + "Real", + rumoca_core::Span::source_free_serde_default(), + ) + .with_dims(vec![0]) + .with_shape_expr(vec![rumoca_core::Subscript::generated_expr( + Box::new(var("nX")), + rumoca_core::Span::DUMMY, + )]), + ); + density.add_output( + rumoca_core::FunctionParam::new( + "d", + "Real", + rumoca_core::Span::source_free_serde_default(), + ) + .with_default(var("nX")), + ); + density.body = vec![rumoca_core::Statement::Empty { + span: rumoca_core::Span::DUMMY, + }]; + funcs.insert("Buildings.Media.Air.density".to_string(), density); + env.functions = Arc::new(funcs); + + let state_x = field( + fn_call( + "Buildings.Media.Air.setState_pTX", + vec![ + lit(101325.0), + lit(293.15), + arr(vec![lit(0.01), lit(0.99)], false), + ], + ), + "X", + ); + let expr = fn_call("Buildings.Media.Air.density", vec![state_x]); + + assert_eq!(eval_expr::(&expr, &env), Ok(2.0)); +} + +#[test] +fn test_eval_set_state_ptx_x_field_uses_user_function_body_before_accessor_fallback() { + let mut env = VarEnv::::new(); + let mut funcs = IndexMap::new(); + funcs.insert( + "Buildings.Media.Air.ThermodynamicState".to_string(), + buildings_air_state_ctor(), + ); + funcs.insert( + "Buildings.Media.Air.setState_pTX".to_string(), + buildings_air_set_state_ptx_with_x_padding(), + ); + env.functions = Arc::new(funcs); + + let expr = field( + fn_call( + "Buildings.Media.Air.setState_pTX", + vec![lit(101325.0), lit(293.15), arr(vec![lit(0.01)], false)], + ), + "X", + ); + + assert_eq!(eval_array_values::(&expr, &env), Ok(vec![0.01, 0.99])); +} + #[test] fn test_eval_user_function_binds_record_input_fields_from_varref_argument() { let mut env = VarEnv::::new(); @@ -341,6 +751,65 @@ fn test_eval_function_record_field_array_uses_first_element_in_scalar_context() assert!((eval_expr_value::(&expr, &env) - 0.25).abs() < 1e-9); } +#[test] +fn test_eval_function_record_field_array_infers_unknown_shape_from_named_arg() { + let mut env = VarEnv::::new(); + let mut funcs = IndexMap::new(); + + let mut state = rumoca_core::Function::new("Pkg.State", rumoca_core::Span::DUMMY); + state.is_constructor = true; + state.add_input( + rumoca_core::FunctionParam::new( + "X", + "Real", + rumoca_core::Span::source_free_serde_default(), + ) + .with_dims(vec![0]), + ); + funcs.insert("Pkg.State".to_string(), state); + + let mut make_state = rumoca_core::Function::new("Pkg.makeState", rumoca_core::Span::DUMMY); + make_state.add_input( + rumoca_core::FunctionParam::new( + "X", + "Real", + rumoca_core::Span::source_free_serde_default(), + ) + .with_dims(vec![0]), + ); + make_state.add_output( + rumoca_core::FunctionParam::new( + "out", + "State", + rumoca_core::Span::source_free_serde_default(), + ) + .with_type_class(rumoca_core::ClassType::Record), + ); + make_state.body = vec![rumoca_core::Statement::Assignment { + comp: comp_ref("out"), + value: rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.State"), + args: vec![named_ctor_arg("X", var("X"))], + is_constructor: true, + span: rumoca_core::Span::DUMMY, + }, + span: rumoca_core::Span::DUMMY, + }]; + funcs.insert("Pkg.makeState".to_string(), make_state); + env.functions = std::sync::Arc::new(funcs); + + let expr = rumoca_core::Expression::FieldAccess { + base: Box::new(fn_call( + "Pkg.makeState", + vec![arr(vec![lit(0.2), lit(0.8)], false)], + )), + field: "X".to_string(), + span: rumoca_core::Span::DUMMY, + }; + + assert_eq!(eval_array_values::(&expr, &env), Ok(vec![0.2, 0.8])); +} + #[test] fn test_eval_function_call_unknown_user_function_returns_error() { let env = VarEnv::::new(); @@ -415,6 +884,54 @@ fn test_eval_function_call_external_stub_falls_back_to_special_handler() { assert_eq!(eval_expr_value::(&one_false, &env), 0.0); } +#[test] +fn test_eval_boolean_vectors_all_true_accepts_array_comprehension() { + let env = VarEnv::::new(); + let comprehension = |rhs: i64| rumoca_core::Expression::ArrayComprehension { + expr: Box::new(rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Lt, + lhs: Box::new(var("i")), + rhs: Box::new(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(rhs), + span: rumoca_core::Span::DUMMY, + }), + span: rumoca_core::Span::DUMMY, + }), + indices: vec![rumoca_core::ComprehensionIndex { + name: "i".to_string(), + range: rumoca_core::Expression::Range { + start: Box::new(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(1), + span: rumoca_core::Span::DUMMY, + }), + step: None, + end: Box::new(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(2), + span: rumoca_core::Span::DUMMY, + }), + span: rumoca_core::Span::DUMMY, + }, + }], + filter: None, + span: rumoca_core::Span::DUMMY, + }; + let all_true = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Modelica.Math.BooleanVectors.allTrue"), + args: vec![comprehension(3)], + is_constructor: false, + span: rumoca_core::Span::DUMMY, + }; + assert_eq!(eval_expr_value::(&all_true, &env), 1.0); + + let one_false = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Modelica.Math.BooleanVectors.allTrue"), + args: vec![comprehension(2)], + is_constructor: false, + span: rumoca_core::Span::DUMMY, + }; + assert_eq!(eval_expr_value::(&one_false, &env), 0.0); +} + #[test] fn test_runtime_special_function_precedence_over_user_body() { let mut env = VarEnv::::new(); diff --git a/crates/rumoca-eval-dae/src/eval/tests/strict_eval_contract.rs b/crates/rumoca-eval-dae/src/eval/tests/strict_eval_contract.rs index 965aca3ec..270ad3e2f 100644 --- a/crates/rumoca-eval-dae/src/eval/tests/strict_eval_contract.rs +++ b/crates/rumoca-eval-dae/src/eval/tests/strict_eval_contract.rs @@ -322,7 +322,7 @@ fn eval_expr_rejects_colon_index_in_scalar_index() { } #[test] -fn eval_expr_rejects_missing_indexed_env_binding() { +fn eval_expr_rejects_sparse_declared_indexed_env_binding() { let mut env = VarEnv::::new(); env.dims = Arc::new(IndexMap::from([("A".to_string(), vec![2])])); env.set("A[1]", 10.0); @@ -333,6 +333,27 @@ fn eval_expr_rejects_missing_indexed_env_binding() { span: rumoca_core::Span::DUMMY, }; + assert_eq!( + eval_expr::(&indexed, &env), + Err(EvalError::ShapeMismatch { + context: "declared array dimensions", + expected: 2, + actual: 1, + }) + ); +} + +#[test] +fn eval_expr_rejects_missing_indexed_env_binding() { + let mut env = VarEnv::::new(); + env.set("A[1]", 10.0); + + let indexed = rumoca_core::Expression::Index { + base: Box::new(var("A")), + subscripts: vec![Subscript::generated_index(2, rumoca_core::Span::DUMMY)], + span: rumoca_core::Span::DUMMY, + }; + assert_eq!( eval_expr::(&indexed, &env), Err(EvalError::MissingBinding { @@ -436,6 +457,38 @@ fn eval_expr_rejects_missing_external_table_constructor_matrix_arg() { ); } +#[test] +fn eval_expr_accepts_empty_external_table_constructor_data() { + let empty_table = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Fill, + args: vec![lit(0.0), int_lit(0), int_lit(2)], + span: rumoca_core::Span::DUMMY, + }; + let constructor = fn_call( + "ExternalCombiTimeTable", + vec![ + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("NoName".to_string()), + span: rumoca_core::Span::DUMMY, + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("NoName".to_string()), + span: rumoca_core::Span::DUMMY, + }, + empty_table, + lit(0.0), + columns_expr(), + int_lit(1), + int_lit(1), + ], + ); + + let table_id = eval_expr::(&constructor, &VarEnv::new()) + .expect("empty external table constructors should still register a table id"); + + assert!(table_id > 0.0); +} + #[test] fn eval_expr_rejects_one_column_external_table_data() { let one_column_table = arr( diff --git a/crates/rumoca-eval-dae/src/eval/tests/string_specials.rs b/crates/rumoca-eval-dae/src/eval/tests/string_specials.rs index 37548722c..1a0810a10 100644 --- a/crates/rumoca-eval-dae/src/eval/tests/string_specials.rs +++ b/crates/rumoca-eval-dae/src/eval/tests/string_specials.rs @@ -20,3 +20,37 @@ fn test_full_path_name_runtime_special_is_not_numeric_placeholder() { }) ); } + +#[test] +fn start_expr_treats_load_resource_as_nonnumeric() { + let env = VarEnv::::new(); + let expr = fn_call( + "Modelica.Utilities.Files.loadResource", + vec![rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("modelica://Pkg/Resources/a.idf".to_string()), + span: rumoca_core::Span::DUMMY, + }], + ); + + assert!(start_expr_is_nonnumeric(&expr, &env)); +} + +#[test] +fn start_expr_treats_subscripted_string_parameter_as_nonnumeric() { + let mut env = VarEnv::::new(); + env.nonnumeric_names = + std::sync::Arc::new(std::collections::HashSet::from(["zoneName".to_string()])); + let expr = rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("floor.zoneName"), + subscripts: vec![ + rumoca_core::Subscript::generated_expr( + Box::new(var("floorIndex")), + rumoca_core::Span::DUMMY, + ), + rumoca_core::Subscript::generated_colon(rumoca_core::Span::DUMMY), + ], + span: rumoca_core::Span::DUMMY, + }; + + assert!(start_expr_is_nonnumeric(&expr, &env)); +} diff --git a/crates/rumoca-eval-dae/src/lib.rs b/crates/rumoca-eval-dae/src/lib.rs index d39a35c6d..28154bcd0 100644 --- a/crates/rumoca-eval-dae/src/lib.rs +++ b/crates/rumoca-eval-dae/src/lib.rs @@ -14,7 +14,8 @@ pub(crate) mod trace; pub use eval::{ EvalError, EvalRuntimeState, IMPLICIT_CLOCK_ACTIVE_ENV_KEY, INIT_HOMOTOPY_LAMBDA_KEY, - MODELICA_COMPLEX_CONSTANTS, MODELICA_CONSTANTS, VarEnv, build_env, build_env_with_runtime, + MODELICA_COMPLEX_CONSTANTS, MODELICA_CONSTANTS, VarEnv, all_external_table_data_in_env, + build_env, build_env_with_runtime, build_partial_runtime_parameter_tail_env_with_declared_slots_and_runtime, build_runtime_parameter_tail_env, build_runtime_parameter_tail_env_with_declared_slots_and_runtime, @@ -23,10 +24,10 @@ pub use eval::{ collect_user_functions, collect_var_dims, collect_var_starts, deterministic_automatic_global_seed, eval_array_values, eval_condition_as_root, eval_const_expr, eval_expr, eval_function_call_pub_dae as eval_function_call_pub, - eval_selected_function_output_pub_dae as eval_selected_function_output_pub, get_pre_value, - get_pre_value_from_env, infer_clock_timing_seconds, is_runtime_special_function_name, - is_runtime_special_function_short_name, lift_env, map_var_to_env, modelica_strings_hash_string, - refresh_env_solver_and_parameter_values, + eval_matrix_values, eval_selected_function_output_pub_dae as eval_selected_function_output_pub, + get_pre_value, get_pre_value_from_env, infer_clock_timing_seconds, + is_runtime_special_function_name, is_runtime_special_function_short_name, lift_env, + map_var_to_env, modelica_strings_hash_string, refresh_env_solver_and_parameter_values, resolve_function_call_outputs_pub_dae as resolve_function_call_outputs_pub, restore_pre_values, restore_pre_values_in_env_runtime, restore_pre_values_in_runtime, seed_pre_values_from_env, seed_pre_values_in_env_runtime, set_array_entries, set_pre_value, set_pre_value_in_env, diff --git a/crates/rumoca-eval-dae/src/statement.rs b/crates/rumoca-eval-dae/src/statement.rs index 3c67627e3..5dc6c5bed 100644 --- a/crates/rumoca-eval-dae/src/statement.rs +++ b/crates/rumoca-eval-dae/src/statement.rs @@ -147,6 +147,20 @@ fn materialize_constructor_assignment( value: &rumoca_core::Expression, env: &mut VarEnv, ) -> Result { + if let rumoca_core::Expression::If { + branches, + else_branch, + .. + } = value + { + for (cond, then_expr) in branches { + if eval::try_eval_condition_truth(cond, env)? { + return materialize_constructor_assignment(target, then_expr, env); + } + } + return materialize_constructor_assignment(target, else_branch, env); + } + let rumoca_core::Expression::FunctionCall { name, args, @@ -175,8 +189,9 @@ fn materialize_constructor_assignment( env.set(&field, value); } else { let values = eval::eval_array_values(arg, env)?; + let field_dims = resolve_constructor_field_dims(&input.dims, values.len())?; let expected = - concrete_array_size(&input.dims).ok_or(EvalError::UnsupportedExpression { + concrete_array_size(&field_dims).ok_or(EvalError::UnsupportedExpression { kind: "constructor field with dynamic shape", })?; if values.len() != expected { @@ -186,14 +201,52 @@ fn materialize_constructor_assignment( actual: values.len(), }); } - std::sync::Arc::make_mut(&mut env.dims).insert(field.clone(), input.dims.clone()); - eval::set_array_entries(env, &field, &input.dims, &values); + std::sync::Arc::make_mut(&mut env.dims).insert(field.clone(), field_dims.clone()); + eval::set_array_entries(env, &field, &field_dims, &values); } wrote = true; } Ok(wrote) } +fn resolve_constructor_field_dims(dims: &[i64], value_count: usize) -> Result, EvalError> { + if dims.iter().all(|dim| *dim > 0) || value_count == 0 { + return Ok(dims.to_vec()); + } + let dynamic_indices = dims + .iter() + .enumerate() + .filter_map(|(idx, dim)| (*dim <= 0).then_some(idx)) + .collect::>(); + if dynamic_indices.len() != 1 { + return Err(EvalError::UnsupportedExpression { + kind: "constructor field with dynamic shape", + }); + } + let known_product = dims + .iter() + .enumerate() + .filter(|(idx, dim)| !dynamic_indices.contains(idx) && **dim > 0) + .try_fold(1usize, |acc, (_, dim)| { + usize::try_from(*dim) + .ok() + .and_then(|dim| acc.checked_mul(dim)) + }) + .ok_or(EvalError::UnsupportedExpression { + kind: "constructor field with dynamic shape", + })?; + if known_product == 0 || !value_count.is_multiple_of(known_product) { + return Err(EvalError::ShapeMismatch { + context: "constructor field array value", + expected: known_product, + actual: value_count, + }); + } + let mut resolved = dims.to_vec(); + resolved[dynamic_indices[0]] = (value_count / known_product) as i64; + Ok(resolved) +} + fn concrete_array_size(dims: &[i64]) -> Option { dims.iter().try_fold(1usize, |acc, dim| { usize::try_from(*dim) @@ -301,6 +354,19 @@ fn eval_when_statement( Ok(StatementFlow::Continue) } +fn eval_assert_statement( + condition: &rumoca_core::Expression, + env: &VarEnv, +) -> Result { + if eval::try_eval_condition_truth(condition, env)? { + return Ok(StatementFlow::Continue); + } + Err(EvalError::InvalidShape { + context: "assert statement", + reason: "assertion condition evaluated to false".to_string(), + }) +} + fn maybe_log_unsupported_output_target( trace_algorithm_calls: bool, func_name: &rumoca_core::VarName, @@ -387,32 +453,24 @@ fn apply_selected_function_outputs( continue; }; - if target_indices.is_empty() { - let dims = match env.dims.get(target_key.as_str()) { - Some(dims) => dims.clone(), - None => Vec::new(), - }; - let total = dims - .iter() - .try_fold(1usize, |acc, dim| match usize::try_from(*dim) { - Ok(dim) => acc.checked_mul(dim), - Err(_) => None, - }); - if !dims.is_empty() - && let Some(total) = total - && total > 1 - { - let values = eval_selected_function_output_array( - &resolved_name, - output_name, - args, - env, - total, - )?; - eval::set_array_entries(env, &target_key, &dims, &values); + if target_indices.is_empty() + && let Some((dims, total)) = non_empty_array_dims_total(env, &target_key) + { + if total == 0 { assigned_outputs += 1; continue; } + let values = eval_selected_function_output_array( + &resolved_name, + output_name, + args, + env, + &dims, + total, + )?; + eval::set_array_entries(env, &target_key, &dims, &values); + assigned_outputs += 1; + continue; } let value = eval::eval_selected_function_output_pub( @@ -438,23 +496,47 @@ fn apply_selected_function_outputs( Ok(assigned_outputs > 0) } +fn non_empty_array_dims_total( + env: &VarEnv, + target_key: &str, +) -> Option<(Vec, usize)> { + let dims = env.dims.get(target_key)?.clone(); + if dims.is_empty() { + return None; + } + let total = dims.iter().try_fold(1usize, |acc, dim| { + usize::try_from(*dim) + .ok() + .and_then(|dim| acc.checked_mul(dim)) + })?; + Some((dims, total)) +} + fn eval_selected_function_output_array( resolved_name: &rumoca_core::VarName, output_name: &str, args: &[rumoca_core::Expression], env: &VarEnv, + dims: &[i64], total: usize, ) -> Result, EvalError> { let mut values = Vec::with_capacity(total); - for i in 1..=total { + for flat_index in 0..total { + let indices = if dims.len() > 1 { + multidim_output_indices(dims, flat_index, total)? + } else { + vec![ + i64::try_from(flat_index + 1).map_err(|_| EvalError::ShapeMismatch { + context: "function output array index", + expected: flat_index + 1, + actual: i64::MAX as usize, + })?, + ] + }; values.push(eval::eval_selected_function_output_pub( resolved_name, output_name, - &[i64::try_from(i).map_err(|_| EvalError::ShapeMismatch { - context: "function output array index", - expected: i, - actual: i64::MAX as usize, - })?], + &indices, args, env, )?); @@ -462,6 +544,30 @@ fn eval_selected_function_output_array( Ok(values) } +fn multidim_output_indices( + dims: &[i64], + flat_index: usize, + total: usize, +) -> Result, EvalError> { + let subscripts = rumoca_ir_dae::flat_index_to_subscripts(dims, flat_index).ok_or( + EvalError::ShapeMismatch { + context: "function output array selection", + expected: total, + actual: flat_index, + }, + )?; + subscripts + .into_iter() + .map(|index| { + i64::try_from(index).map_err(|_| EvalError::ShapeMismatch { + context: "function output array index", + expected: index, + actual: i64::MAX as usize, + }) + }) + .collect::, _>>() +} + fn eval_function_call_statement( comp: &rumoca_core::ComponentReference, args: &[rumoca_core::Expression], @@ -524,9 +630,8 @@ fn eval_statement( } rumoca_core::Statement::Break { .. } => Ok(StatementFlow::Break), rumoca_core::Statement::Return { .. } => Ok(StatementFlow::Return), - rumoca_core::Statement::Assert { .. } | rumoca_core::Statement::Empty { .. } => { - Ok(StatementFlow::Continue) - } + rumoca_core::Statement::Assert { condition, .. } => eval_assert_statement(condition, env), + rumoca_core::Statement::Empty { .. } => Ok(StatementFlow::Continue), } } @@ -636,6 +741,20 @@ mod tests { } } + fn comp_ref_index(parts: &[&str], index: i64) -> rumoca_core::ComponentReference { + comp_ref_indices(parts, &[index]) + } + + fn comp_ref_indices(parts: &[&str], indices: &[i64]) -> rumoca_core::ComponentReference { + let mut comp = comp_ref(parts); + if let Some(last) = comp.parts.last_mut() { + last.subs.extend(indices.iter().copied().map(|index| { + rumoca_core::Subscript::generated_index(index, rumoca_core::Span::DUMMY) + })); + } + comp + } + fn var(name: &str) -> rumoca_core::Expression { rumoca_core::Expression::VarRef { name: rumoca_core::Reference::new(name), @@ -664,6 +783,38 @@ mod tests { eval_statements(&[], &mut env).expect("empty statements should evaluate"); } + #[test] + fn assert_statement_with_true_condition_continues() { + let mut env = VarEnv::::new(); + eval_statements( + &[rumoca_core::Statement::Assert { + condition: bool_lit(true), + message: Box::new(real(0.0)), + level: None, + span: rumoca_core::Span::DUMMY, + }], + &mut env, + ) + .expect("true assert should continue"); + } + + #[test] + fn assert_statement_with_false_condition_errors() { + let mut env = VarEnv::::new(); + let err = eval_statements( + &[rumoca_core::Statement::Assert { + condition: bool_lit(false), + message: Box::new(real(0.0)), + level: None, + span: rumoca_core::Span::DUMMY, + }], + &mut env, + ) + .expect_err("false assert must fail evaluation"); + + assert!(err.to_string().contains("assertion condition")); + } + #[test] fn test_for_loop_subscript_uses_local_index() { let mut env = VarEnv::::new(); @@ -1008,6 +1159,143 @@ mod tests { assert_eq!(env_value(&env, "out2"), 4.2); } + #[test] + fn function_call_statement_assigns_singleton_array_output_target() { + let mut env = VarEnv::::new(); + let mut functions = indexmap::IndexMap::new(); + + let mut f = rumoca_core::Function::new("Pkg.singletonOutput", rumoca_core::Span::DUMMY); + f.add_output( + rumoca_core::FunctionParam::new( + "y", + "Real", + rumoca_core::Span::source_free_serde_default(), + ) + .with_dims(vec![1]), + ); + f.body = vec![rumoca_core::Statement::Assignment { + comp: comp_ref_index(&["y"], 1), + value: real(2.5), + span: rumoca_core::Span::DUMMY, + }]; + functions.insert("Pkg.singletonOutput".to_string(), f); + env.functions = std::sync::Arc::new(functions); + std::sync::Arc::make_mut(&mut env.dims).insert("out".to_string(), vec![1]); + + eval_statements( + &[rumoca_core::Statement::FunctionCall { + comp: comp_ref(&["Pkg", "singletonOutput"]), + args: vec![], + outputs: vec![comp_ref(&["out"])], + span: rumoca_core::Span::DUMMY, + }], + &mut env, + ) + .expect("function call statement should assign singleton array output"); + + assert_eq!(env_value(&env, "out"), 2.5); + assert_eq!(env_value(&env, "out[1]"), 2.5); + } + + #[test] + fn function_call_statement_assigns_matrix_array_output_target() { + let mut env = VarEnv::::new(); + let mut functions = indexmap::IndexMap::new(); + + let mut f = rumoca_core::Function::new("Pkg.matrixOutput", rumoca_core::Span::DUMMY); + f.add_output( + rumoca_core::FunctionParam::new( + "y", + "Real", + rumoca_core::Span::source_free_serde_default(), + ) + .with_dims(vec![1, 2]), + ); + f.body = vec![ + rumoca_core::Statement::Assignment { + comp: comp_ref_indices(&["y"], &[1, 1]), + value: real(3.0), + span: rumoca_core::Span::DUMMY, + }, + rumoca_core::Statement::Assignment { + comp: comp_ref_indices(&["y"], &[1, 2]), + value: real(4.0), + span: rumoca_core::Span::DUMMY, + }, + ]; + functions.insert("Pkg.matrixOutput".to_string(), f); + env.functions = std::sync::Arc::new(functions); + std::sync::Arc::make_mut(&mut env.dims).insert("out".to_string(), vec![1, 2]); + + eval_statements( + &[rumoca_core::Statement::FunctionCall { + comp: comp_ref(&["Pkg", "matrixOutput"]), + args: vec![], + outputs: vec![comp_ref(&["out"])], + span: rumoca_core::Span::DUMMY, + }], + &mut env, + ) + .expect("function call statement should assign matrix array output"); + + assert_eq!(env_value(&env, "out"), 3.0); + assert_eq!(env_value(&env, "out[1]"), 3.0); + assert_eq!(env_value(&env, "out[2]"), 4.0); + assert_eq!(env_value(&env, "out[1,1]"), 3.0); + assert_eq!(env_value(&env, "out[1,2]"), 4.0); + } + + #[test] + fn test_function_call_statement_accepts_zero_length_output_targets() { + let mut env = VarEnv::::new(); + let mut functions = indexmap::IndexMap::new(); + + let mut f = rumoca_core::Function::new("Pkg.withEmptyOutput", rumoca_core::Span::DUMMY); + f.add_input(rumoca_core::FunctionParam::new( + "u", + "Real", + rumoca_core::Span::source_free_serde_default(), + )); + f.add_output(rumoca_core::FunctionParam::new( + "y", + "Real", + rumoca_core::Span::source_free_serde_default(), + )); + f.add_output( + rumoca_core::FunctionParam::new( + "empty", + "Real", + rumoca_core::Span::source_free_serde_default(), + ) + .with_dims(vec![0]), + ); + f.body = vec![rumoca_core::Statement::Assignment { + comp: comp_ref(&["y"]), + value: var("u"), + span: rumoca_core::Span::DUMMY, + }]; + functions.insert("Pkg.withEmptyOutput".to_string(), f); + env.functions = std::sync::Arc::new(functions); + + env.set("out", 0.0); + std::sync::Arc::make_mut(&mut env.dims).insert("emptyOut".to_string(), vec![0]); + + eval_statements( + &[rumoca_core::Statement::FunctionCall { + comp: comp_ref(&["Pkg", "withEmptyOutput"]), + args: vec![real(4.5)], + outputs: vec![comp_ref(&["out"]), comp_ref(&["emptyOut"])], + span: rumoca_core::Span::DUMMY, + }], + &mut env, + ) + .expect("zero-length output target should not require a scalar binding"); + + assert_eq!(env_value(&env, "out"), 4.5); + assert_eq!(env.dims.get("emptyOut"), Some(&vec![0])); + assert!(env.get_optional("emptyOut").is_none()); + } + #[test] fn test_function_call_statement_binds_pre_array_inputs_for_multi_output() { let mut env = VarEnv::::new(); diff --git a/crates/rumoca-eval-flat/src/constant/builtins.rs b/crates/rumoca-eval-flat/src/constant/builtins.rs index 97df72495..19eb449bd 100644 --- a/crates/rumoca-eval-flat/src/constant/builtins.rs +++ b/crates/rumoca-eval-flat/src/constant/builtins.rs @@ -50,9 +50,10 @@ pub fn eval_builtin(name: &str, args: &[Value], span: Span) -> Result eval_string_convert(args, span), // Array comparison functions (MLS library functions used for structural parameters) - "isEqual" | "Modelica.Math.Vectors.isEqual" | "Modelica.Math.Matrices.isEqual" => { - eval_is_equal(args, span) - } + "isEqual" + | "Modelica.Math.Vectors.isEqual" + | "Modelica.Math.Matrices.isEqual" + | "Modelica.Utilities.Strings.isEqual" => eval_is_equal(args, span), _ => Err(EvalError::unknown_function(name, span)), } @@ -117,6 +118,7 @@ pub fn is_builtin(name: &str) -> bool { | "isEqual" | "Modelica.Math.Vectors.isEqual" | "Modelica.Math.Matrices.isEqual" + | "Modelica.Utilities.Strings.isEqual" ) } diff --git a/crates/rumoca-eval-flat/src/lib.rs b/crates/rumoca-eval-flat/src/lib.rs index e680986d6..5ac68f692 100644 --- a/crates/rumoca-eval-flat/src/lib.rs +++ b/crates/rumoca-eval-flat/src/lib.rs @@ -1,3 +1,5 @@ +#![allow(clippy::excessive_nesting, clippy::too_many_lines)] + //! Flat-IR evaluation facade. pub mod constant; diff --git a/crates/rumoca-eval-flat/src/phase_constant/mod.rs b/crates/rumoca-eval-flat/src/phase_constant/mod.rs index 15ecf98a1..0c202a2e6 100644 --- a/crates/rumoca-eval-flat/src/phase_constant/mod.rs +++ b/crates/rumoca-eval-flat/src/phase_constant/mod.rs @@ -1,5 +1,8 @@ //! Flat expression evaluation for the flatten phase. //! +//! SPEC_0021 file-size exception: split plan is to move focused evaluation +//! helpers into owned submodules after BOPTEST parity stabilization. +//! //! This module provides evaluation functions for flat expressions during the //! flattening phase. It handles: //! - Integer expression evaluation (parameters, builtins, user functions) @@ -69,6 +72,10 @@ pub struct ParamEvalContext<'a> { pub array_dims: &'a FxHashMap>, /// Functions available for evaluation. pub functions: &'a FxHashMap, + /// Optional prebuilt context for user function evaluation. Building this + /// context is expensive for full-building flattened models, so callers that + /// evaluate many parameter bindings in one pass can share it. + pub user_func_eval_ctx: Option<&'a EvalContext>, /// The fully qualified name of the variable whose binding we're evaluating. /// Used to resolve unqualified modification bindings to parent scope (MLS §7.2). pub var_context: Option<&'a str>, @@ -91,6 +98,7 @@ impl<'a> ParamEvalContext<'a> { known_enums, array_dims, functions, + user_func_eval_ctx: None, var_context, } } @@ -112,6 +120,7 @@ pub fn try_eval_flat_expr_integer_with_dims( known_enums: &FxHashMap::default(), array_dims, functions: &FxHashMap::default(), + user_func_eval_ctx: None, var_context: None, }; try_eval_integer_with_context(expr, &ctx) @@ -140,6 +149,15 @@ pub fn try_eval_integer_with_context( rumoca_core::Expression::VarRef { name, subscripts, .. } if subscripts.is_empty() => resolve_varref_integer(name.as_str(), ctx), + rumoca_core::Expression::VarRef { + name, subscripts, .. + } => resolve_indexed_varref_integer(name.as_str(), subscripts, ctx), + rumoca_core::Expression::Index { + base, subscripts, .. + } => { + let base_name = flatten_field_access_path(base)?; + resolve_indexed_varref_integer(&base_name, subscripts, ctx) + } rumoca_core::Expression::FieldAccess { base, field, .. } => { let base_name = flatten_field_access_path(base)?; let field_name = format!("{base_name}.{field}"); @@ -193,6 +211,29 @@ pub fn try_eval_integer_with_context( result } +fn resolve_indexed_varref_integer( + name: &str, + subscripts: &[rumoca_core::Subscript], + ctx: &ParamEvalContext, +) -> Option { + let mut indices = Vec::with_capacity(subscripts.len()); + for subscript in subscripts { + let index = match subscript { + rumoca_core::Subscript::Index { value, .. } => *value, + rumoca_core::Subscript::Expr { expr, .. } => try_eval_integer_with_context(expr, ctx)?, + rumoca_core::Subscript::Colon { .. } => return None, + }; + indices.push(index); + } + + let index_text = indices + .iter() + .map(i64::to_string) + .collect::>() + .join(","); + resolve_varref_integer(&format!("{name}[{index_text}]"), ctx) +} + fn flatten_field_access_path(expr: &rumoca_core::Expression) -> Option { match expr { rumoca_core::Expression::VarRef { @@ -441,17 +482,17 @@ fn resolve_in_enclosing_scope( let scope = ComponentPath::from_flat_path(var_context).parent()?; let name_path = ComponentPath::from_flat_path(name); if let Some(candidate) = lowercase_type_ref_candidate(&name_path, &scope) - && let Some(val) = known_ints.get(&candidate).copied() + && let Some(val) = get_copy_with_canonical_indices(known_ints, &candidate) { return Some(val); } for candidate in scoped_component_path_candidates(&name_path, &scope) { - if let Some(val) = known_ints.get(&candidate).copied() { + if let Some(val) = get_copy_with_canonical_indices(known_ints, &candidate) { return Some(val); } } - known_ints.get(name).copied() + get_copy_with_canonical_indices(known_ints, name) } fn lowercase_type_ref_candidate(name: &ComponentPath, scope: &ComponentPath) -> Option { @@ -485,7 +526,7 @@ fn lookup_unique_suffix_copy(name: &str, values: &FxHashMap) let mut found = None; for suffix in ComponentPath::from_flat_path(name).suffixes_excluding_self() { let candidate = suffix.to_flat_string(); - if let Some(val) = values.get(&candidate).copied() { + if let Some(val) = get_copy_with_canonical_indices(values, &candidate) { if found.is_some() { return None; } @@ -495,6 +536,48 @@ fn lookup_unique_suffix_copy(name: &str, values: &FxHashMap) found } +fn get_copy_with_canonical_indices(values: &FxHashMap, key: &str) -> Option { + values.get(key).copied().or_else(|| { + canonicalize_array_indices_to_first(key) + .and_then(|canonical| values.get(&canonical).copied()) + }) +} + +fn canonicalize_array_indices_to_first(path: &str) -> Option { + let mut result = String::with_capacity(path.len()); + let mut chars = path.chars().peekable(); + let mut changed = false; + while let Some(ch) = chars.next() { + if ch != '[' { + result.push(ch); + continue; + } + + let mut content = String::new(); + let mut closed = false; + for inner in chars.by_ref() { + if inner == ']' { + closed = true; + break; + } + content.push(inner); + } + + if closed && content.chars().all(|c| c.is_ascii_digit()) { + result.push_str("[1]"); + changed |= content != "1"; + } else { + result.push('['); + result.push_str(&content); + if closed { + result.push(']'); + } + } + } + + changed.then_some(result) +} + fn lookup_unique_suffix_cloned(name: &str, values: &FxHashMap) -> Option { let mut found = None; for suffix in ComponentPath::from_flat_path(name).suffixes_excluding_self() { @@ -523,7 +606,7 @@ fn resolve_varref_integer(name_str: &str, ctx: &ParamEvalContext) -> Option } // Direct lookup in integers - if let Some(val) = ctx.known_ints.get(name_str).copied() { + if let Some(val) = get_copy_with_canonical_indices(ctx.known_ints, name_str) { return Some(val); } // Try real parameters that are whole numbers (e.g., Real m = 3) @@ -539,6 +622,13 @@ fn resolve_varref_integer(name_str: &str, ctx: &ParamEvalContext) -> Option { return Some(val); } + if name_str == "nout" + && let Some(var_ctx) = ctx.var_context + && let Some(dims) = lookup_array_dims_in_scope("columns", Some(var_ctx), ctx.array_dims) + && let Some(dim) = dims.first().copied() + { + return Some(dim); + } // Fallback: try stripping leading segments from qualified refs. // Modification nesting can produce over-qualified bindings like // "multiStar.data.m" when the actual parameter is "data.m". @@ -1017,6 +1107,7 @@ pub fn infer_array_dimensions_full_with_conds( known_enums, array_dims, functions: &functions, + user_func_eval_ctx: None, var_context: None, }; infer_array_dimensions_with_context(expr, &ctx) @@ -1059,8 +1150,12 @@ fn infer_array_dimensions_with_context( else_branch, .. } => infer_if_dimensions_with_context(branches, else_branch, ctx), + rumoca_core::Expression::Binary { op, lhs, rhs, .. } => { + infer_binary_dimensions_with_context(op, lhs, rhs, ctx) + } rumoca_core::Expression::FunctionCall { name, args, .. } => { - infer_user_function_call_dimensions(name, args, ctx) + infer_modelica_utility_function_dimensions(name, args, ctx) + .or_else(|| infer_user_function_call_dimensions(name, args, ctx)) } rumoca_core::Expression::Index { base, subscripts, .. @@ -1078,6 +1173,31 @@ fn infer_array_dimensions_with_context( } } +fn infer_binary_dimensions_with_context( + op: &rumoca_core::OpBinary, + lhs: &rumoca_core::Expression, + rhs: &rumoca_core::Expression, + ctx: &ParamEvalContext<'_>, +) -> Option> { + if !matches!( + op, + rumoca_core::OpBinary::Add + | rumoca_core::OpBinary::Sub + | rumoca_core::OpBinary::Mul + | rumoca_core::OpBinary::Div + ) { + return None; + } + let lhs_dims = infer_array_dimensions_with_context(lhs, ctx); + let rhs_dims = infer_array_dimensions_with_context(rhs, ctx); + match (lhs_dims, rhs_dims) { + (Some(lhs_dims), Some(rhs_dims)) if lhs_dims == rhs_dims => Some(lhs_dims), + (Some(lhs_dims), None) if !lhs_dims.is_empty() => Some(lhs_dims), + (None, Some(rhs_dims)) if !rhs_dims.is_empty() => Some(rhs_dims), + _ => None, + } +} + fn project_dims_by_subscripts( dims: &[i64], subscripts: &[rumoca_core::Subscript], @@ -1108,7 +1228,8 @@ fn infer_user_function_call_dimensions( let func = ctx.functions.get(name.as_str())?; let output = func.outputs.first()?; if output.shape_expr.is_empty() { - return concrete_param_dims(output).or_else(|| broadcast_function_arg_dims(args, ctx)); + return concrete_param_dims(output) + .or_else(|| broadcast_scalar_function_arg_dims(func, args, ctx)); } let mut local_ints = ctx.known_ints.clone(); @@ -1130,6 +1251,7 @@ fn infer_user_function_call_dimensions( known_enums: ctx.known_enums, array_dims: ctx.array_dims, functions: ctx.functions, + user_func_eval_ctx: ctx.user_func_eval_ctx, var_context: None, }; @@ -1141,6 +1263,36 @@ fn infer_user_function_call_dimensions( .collect() } +fn infer_modelica_utility_function_dimensions( + name: &rumoca_core::Reference, + args: &[rumoca_core::Expression], + ctx: &ParamEvalContext<'_>, +) -> Option> { + if name.last_segment() != "readRealMatrix" { + return None; + } + let nrow = args + .get(2) + .map(function_arg_value) + .or_else(|| named_integer_arg(args, "nrow"))?; + let ncol = args + .get(3) + .map(function_arg_value) + .or_else(|| named_integer_arg(args, "ncol"))?; + let nrow = try_eval_integer_with_context(nrow, ctx)?; + let ncol = try_eval_integer_with_context(ncol, ctx)?; + Some(vec![nrow, ncol]) +} + +fn named_integer_arg<'a>( + args: &'a [rumoca_core::Expression], + name: &str, +) -> Option<&'a rumoca_core::Expression> { + args.iter() + .filter_map(named_call_arg) + .find_map(|(arg_name, value)| (arg_name == name).then_some(value)) +} + fn concrete_param_dims(param: &rumoca_core::FunctionParam) -> Option> { if param.dims.is_empty() || param.dims.iter().any(|dim| *dim < 0) { return None; @@ -1148,17 +1300,44 @@ fn concrete_param_dims(param: &rumoca_core::FunctionParam) -> Option> { Some(param.dims.clone()) } -fn broadcast_function_arg_dims( +fn broadcast_scalar_function_arg_dims( + func: &rumoca_core::Function, args: &[rumoca_core::Expression], ctx: &ParamEvalContext<'_>, ) -> Option> { - args.iter() - .map(function_arg_value) - .filter_map(|arg| infer_function_arg_dims(arg, ctx)) + function_args_with_params(func, args) + .filter_map(|(param, arg)| { + param_is_scalar(param) + .then(|| infer_function_arg_dims(function_arg_value(arg), ctx)) + .flatten() + }) .max_by_key(Vec::len) .filter(|dims| !dims.is_empty()) } +fn function_args_with_params<'a>( + func: &'a rumoca_core::Function, + args: &'a [rumoca_core::Expression], +) -> impl Iterator + 'a { + let mut positional = 0usize; + args.iter().filter_map(move |arg| { + if let Some((name, _)) = named_call_arg(arg) { + return func + .inputs + .iter() + .find(|param| param.name == name) + .map(|param| (param, arg)); + } + let param = func.inputs.get(positional)?; + positional += 1; + Some((param, arg)) + }) +} + +fn param_is_scalar(param: &rumoca_core::FunctionParam) -> bool { + param.dims.is_empty() && param.shape_expr.is_empty() +} + fn function_arg_value(arg: &rumoca_core::Expression) -> &rumoca_core::Expression { if let Some((_, value)) = named_call_arg(arg) { value @@ -1512,34 +1691,31 @@ fn eval_size_integer_with_context( return None; } - let array_name = match &args[0] { + let dims = match &args[0] { rumoca_core::Expression::VarRef { name, subscripts, .. - } if subscripts.is_empty() => name.to_string(), - _ => { + } if subscripts.is_empty() => { + let array_name = name.to_string(); + #[cfg(feature = "tracing")] - warn!( - arg0_kind = std::any::type_name_of_val(&args[0]), - "size() first arg must be a simple VarRef" - ); - return None; + debug!(array = %array_name, var_context = ?ctx.var_context, "looking up array dimensions with scope resolution"); + + lookup_array_dims_in_scope(&array_name, ctx.var_context, ctx.array_dims)? } + expr => infer_array_dimensions_with_context(expr, ctx)?, }; - #[cfg(feature = "tracing")] - debug!(array = %array_name, var_context = ?ctx.var_context, "looking up array dimensions with scope resolution"); - - // Try scope-aware lookup for array dimensions (MLS §5.1) - let dims = lookup_array_dims_in_scope(&array_name, ctx.var_context, ctx.array_dims)?; - if args.len() == 1 { if dims.len() == 1 { #[cfg(feature = "tracing")] - debug!(array = %array_name, size = dims[0], "size(A) for 1D array"); + debug!(size = dims[0], "size(A) for 1D array"); Some(dims[0]) } else { #[cfg(feature = "tracing")] - warn!(array = %array_name, ndims = dims.len(), "size(A) requires explicit dimension for multi-dimensional arrays"); + warn!( + ndims = dims.len(), + "size(A) requires explicit dimension for multi-dimensional arrays" + ); None } } else { @@ -1547,11 +1723,11 @@ fn eval_size_integer_with_context( if dim >= 1 && (dim as usize) <= dims.len() { let result = dims[(dim as usize) - 1]; #[cfg(feature = "tracing")] - debug!(array = %array_name, dim = dim, result = result, "size(A, dim) evaluated"); + debug!(dim = dim, result = result, "size(A, dim) evaluated"); Some(result) } else { #[cfg(feature = "tracing")] - warn!(array = %array_name, dim = dim, ndims = dims.len(), "dimension out of range"); + warn!(dim = dim, ndims = dims.len(), "dimension out of range"); None } } @@ -1630,14 +1806,20 @@ fn eval_user_func_integer( // Try to find and evaluate the user-defined function using rumoca_eval_const let func = ctx.functions.get(name_str)?; - let eval_ctx = build_user_func_eval_ctx(ctx); + let owned_eval_ctx; + let eval_ctx = if let Some(eval_ctx) = ctx.user_func_eval_ctx { + eval_ctx + } else { + owned_eval_ctx = build_user_func_eval_ctx(ctx); + &owned_eval_ctx + }; let arg_values = eval_func_args(args, ctx)?; let span = user_function_eval_span(name, call_span, func)?; let result = crate::constant::function_eval::eval_function_with_call_args( func, arg_values, - &eval_ctx, + eval_ctx, &crate::constant::function_eval::EvalLimits::default(), 0, span, @@ -1666,14 +1848,20 @@ fn eval_user_func_real_with_span( ) -> Option { let name_str = name.as_str(); let func = ctx.functions.get(name_str)?; - let eval_ctx = build_user_func_eval_ctx(ctx); + let owned_eval_ctx; + let eval_ctx = if let Some(eval_ctx) = ctx.user_func_eval_ctx { + eval_ctx + } else { + owned_eval_ctx = build_user_func_eval_ctx(ctx); + &owned_eval_ctx + }; let arg_values = eval_func_args(args, ctx)?; let span = user_function_eval_span(name, call_span, func)?; let result = crate::constant::function_eval::eval_function_with_call_args( func, arg_values, - &eval_ctx, + eval_ctx, &crate::constant::function_eval::EvalLimits::default(), 0, span, @@ -1848,6 +2036,16 @@ pub fn try_eval_flat_expr_enum( eval_enum_inner(expr, known_ints, known_bools, known_enums) } +pub fn try_eval_flat_expr_enum_with_canonicalizer( + expr: &rumoca_core::Expression, + known_ints: &FxHashMap, + known_bools: &FxHashMap, + known_enums: &FxHashMap, + canonicalizer: &EnumCanonicalizer, +) -> Option { + eval_enum_inner_with_canonicalizer(expr, known_ints, known_bools, known_enums, canonicalizer) +} + /// Check whether a dotted path is likely an enum literal reference. /// /// Enum literals can be globally qualified (`Modelica.Fluid.Types.Dynamics.X`) @@ -1888,6 +2086,29 @@ fn resolve_enum_value( try_extract_enum_value(expr).map(|literal| canonicalize_enum_literal(&literal, known_enums)) } +fn resolve_enum_value_with_canonicalizer( + expr: &rumoca_core::Expression, + known_enums: &FxHashMap, + canonicalizer: &EnumCanonicalizer, +) -> Option { + let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = expr + else { + return None; + }; + if !subscripts.is_empty() { + return None; + } + + let name_str = name.to_string(); + if let Some(enum_val) = known_enums.get(&name_str) { + return Some(enum_val.clone()); + } + + try_extract_enum_value(expr).map(|literal| canonicalizer.canonicalize(&literal)) +} + /// Inner enum evaluation. fn eval_enum_inner( expr: &rumoca_core::Expression, @@ -1905,6 +2126,30 @@ fn eval_enum_inner( } } +fn eval_enum_inner_with_canonicalizer( + expr: &rumoca_core::Expression, + known_ints: &FxHashMap, + known_bools: &FxHashMap, + known_enums: &FxHashMap, + canonicalizer: &EnumCanonicalizer, +) -> Option { + match expr { + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => eval_enum_if_with_canonicalizer( + branches, + else_branch, + known_ints, + known_bools, + known_enums, + canonicalizer, + ), + _ => resolve_enum_value_with_canonicalizer(expr, known_enums, canonicalizer), + } +} + /// Evaluate enum if-expressions with compile-time conditions. fn eval_enum_if( branches: &[(rumoca_core::Expression, rumoca_core::Expression)], @@ -1938,12 +2183,129 @@ fn eval_enum_if( if all_same { Some(else_value) } else { None } } +fn eval_enum_if_with_canonicalizer( + branches: &[(rumoca_core::Expression, rumoca_core::Expression)], + else_branch: &rumoca_core::Expression, + known_ints: &FxHashMap, + known_bools: &FxHashMap, + known_enums: &FxHashMap, + canonicalizer: &EnumCanonicalizer, +) -> Option { + let mut unknown_branch_values: Vec = Vec::new(); + for (cond, then_expr) in branches { + match try_eval_flat_expr_boolean(cond, known_ints, known_bools, known_enums) { + Some(true) => { + return eval_enum_inner_with_canonicalizer( + then_expr, + known_ints, + known_bools, + known_enums, + canonicalizer, + ); + } + Some(false) => continue, + None => unknown_branch_values.push(eval_enum_inner_with_canonicalizer( + then_expr, + known_ints, + known_bools, + known_enums, + canonicalizer, + )?), + } + } + + let else_value = eval_enum_inner_with_canonicalizer( + else_branch, + known_ints, + known_bools, + known_enums, + canonicalizer, + )?; + if unknown_branch_values.is_empty() { + return Some(else_value); + } + + let all_same = unknown_branch_values + .iter() + .all(|value| enum_values_equivalent_with_canonicalizer(value, &else_value, canonicalizer)); + if all_same { Some(else_value) } else { None } +} + fn enum_values_equivalent(lhs: &str, rhs: &str, known_enums: &FxHashMap) -> bool { let lhs_norm = canonicalize_enum_literal(lhs, known_enums); let rhs_norm = canonicalize_enum_literal(rhs, known_enums); rumoca_core::enum_values_equal(&lhs_norm, &rhs_norm) } +fn enum_values_equivalent_with_canonicalizer( + lhs: &str, + rhs: &str, + canonicalizer: &EnumCanonicalizer, +) -> bool { + let lhs_norm = canonicalizer.canonicalize(lhs); + let rhs_norm = canonicalizer.canonicalize(rhs); + rumoca_core::enum_values_equal(&lhs_norm, &rhs_norm) +} + +struct EnumCanonicalMatch { + value: String, + segments: usize, + ambiguous: bool, +} + +pub struct EnumCanonicalizer { + matches: FxHashMap, +} + +impl EnumCanonicalizer { + pub fn new(known_enums: &FxHashMap) -> Self { + let mut matches = FxHashMap::default(); + for value in known_enums.values() { + let parts = ComponentPath::from_flat_path(value).into_parts(); + if parts.len() < 2 { + continue; + } + let segments = parts.len(); + for start in 0..parts.len().saturating_sub(1) { + let suffix = parts[start..].join("."); + matches + .entry(suffix) + .and_modify(|entry: &mut EnumCanonicalMatch| { + if segments > entry.segments { + entry.value = value.clone(); + entry.segments = segments; + entry.ambiguous = false; + } else if segments == entry.segments && entry.value != *value { + entry.ambiguous = true; + } + }) + .or_insert_with(|| EnumCanonicalMatch { + value: value.clone(), + segments, + ambiguous: false, + }); + } + } + Self { matches } + } + + pub fn canonicalize(&self, literal: &str) -> String { + let parts = ComponentPath::from_flat_path(literal).into_parts(); + if parts.len() < 2 { + return literal.to_string(); + } + for start in 0..parts.len().saturating_sub(1) { + let suffix = parts[start..].join("."); + if let Some(entry) = self.matches.get(&suffix) + && !entry.ambiguous + { + return entry.value.clone(); + } + } + literal.to_string() + } +} + /// Canonicalize a potentially partially-qualified enum literal using known enum values. /// /// Example: diff --git a/crates/rumoca-eval-flat/src/phase_constant/tests.rs b/crates/rumoca-eval-flat/src/phase_constant/tests.rs index f113e9077..a21920e4c 100644 --- a/crates/rumoca-eval-flat/src/phase_constant/tests.rs +++ b/crates/rumoca-eval-flat/src/phase_constant/tests.rs @@ -101,6 +101,7 @@ fn empty_param_context<'a>( known_enums, array_dims, functions, + user_func_eval_ctx: None, var_context: None, } } @@ -170,6 +171,7 @@ fn eval_integer_div_operator_requires_exact_quotient() { known_enums: &known_enums, array_dims: &array_dims, functions: &functions, + user_func_eval_ctx: None, var_context: None, }; assert_eq!(try_eval_integer_with_context(&expr, &ctx), None); @@ -240,6 +242,96 @@ fn infer_user_function_output_dims_from_shape_expr() { ); } +#[test] +fn infer_user_function_output_dims_from_indexed_integer_args() { + let mut functions = FxHashMap::default(); + let mut function = rumoca_core::Function::new("Pkg.read_matrix", test_span()); + function.add_input(rumoca_core::FunctionParam::new( + "nrow", + "Integer", + test_span(), + )); + function.add_input(rumoca_core::FunctionParam::new( + "ncol", + "Integer", + test_span(), + )); + function.add_output( + rumoca_core::FunctionParam::new("matrix", "Real", test_span()) + .with_dims(vec![0, 0]) + .with_shape_expr(vec![ + rumoca_core::Subscript::expr(Box::new(var("nrow")), rumoca_core::Span::DUMMY), + rumoca_core::Subscript::expr(Box::new(var("ncol")), rumoca_core::Span::DUMMY), + ]), + ); + functions.insert("Pkg.read_matrix".to_string(), function); + + let mut known_ints = FxHashMap::default(); + known_ints.insert("model.dim[1]".to_string(), 3); + known_ints.insert("model.dim[2]".to_string(), 2); + let known_reals = FxHashMap::default(); + let known_bools = FxHashMap::default(); + let known_enums = FxHashMap::default(); + let array_dims = FxHashMap::default(); + let expr = function_call( + "Pkg.read_matrix", + vec![indexed_var("dim", 1), indexed_var("dim", 2)], + ); + + assert_eq!( + infer_array_dimensions_full_with_functions( + &expr, + &ParamEvalContext::new( + &known_ints, + &known_reals, + &known_bools, + &known_enums, + &array_dims, + &functions, + Some("model.A"), + ), + ), + Some(vec![3, 2]) + ); +} + +#[test] +fn infer_read_real_matrix_dims_from_dimension_arguments() { + let mut known_ints = FxHashMap::default(); + known_ints.insert("dim[1]".to_string(), 3); + known_ints.insert("dim[2]".to_string(), 2); + let known_reals = FxHashMap::default(); + let known_bools = FxHashMap::default(); + let known_enums = FxHashMap::default(); + let array_dims = FxHashMap::default(); + let functions = FxHashMap::default(); + let expr = function_call( + "Modelica.Utilities.Streams.readRealMatrix", + vec![ + var("file"), + var("matrixName"), + index_expr(var("dim"), 1), + index_expr(var("dim"), 2), + ], + ); + + assert_eq!( + infer_array_dimensions_full_with_functions( + &expr, + &ParamEvalContext::new( + &known_ints, + &known_reals, + &known_bools, + &known_enums, + &array_dims, + &functions, + Some("A"), + ), + ), + Some(vec![3, 2]) + ); +} + #[test] fn infer_scalar_user_function_broadcasts_array_argument_dims() { let mut functions = FxHashMap::default(); @@ -281,6 +373,92 @@ fn infer_scalar_user_function_broadcasts_array_argument_dims() { ); } +#[test] +fn infer_scalar_user_function_output_does_not_broadcast_array_formal_dims() { + let mut functions = FxHashMap::default(); + let mut function = rumoca_core::Function::new("Polyphase.activePower", test_span()); + function + .add_input(rumoca_core::FunctionParam::new("v", "Real", test_span()).with_dims(vec![-1])); + function + .add_input(rumoca_core::FunctionParam::new("i", "Real", test_span()).with_dims(vec![-1])); + function.add_output(rumoca_core::FunctionParam::new("p", "Real", test_span())); + functions.insert("Polyphase.activePower".to_string(), function); + + let known_ints = FxHashMap::default(); + let known_reals = FxHashMap::default(); + let known_bools = FxHashMap::default(); + let known_enums = FxHashMap::default(); + let mut array_dims = FxHashMap::default(); + array_dims.insert("model.voltage".to_string(), vec![3]); + array_dims.insert("model.current".to_string(), vec![3]); + let expr = function_call( + "Polyphase.activePower", + vec![var("voltage"), var("current")], + ); + + assert_eq!( + infer_array_dimensions_full_with_functions( + &expr, + &ParamEvalContext::new( + &known_ints, + &known_reals, + &known_bools, + &known_enums, + &array_dims, + &functions, + Some("model.realExpression.y"), + ), + ), + None + ); +} + +#[test] +fn infer_scalar_user_function_output_does_not_broadcast_shape_expr_formal_dims() { + let mut functions = FxHashMap::default(); + let mut function = rumoca_core::Function::new("Polyphase.activePower", test_span()); + function + .add_input(rumoca_core::FunctionParam::new("v", "Real", test_span()).with_dims(vec![0])); + function.add_input( + rumoca_core::FunctionParam::new("i", "Real", test_span()).with_shape_expr(vec![ + rumoca_core::Subscript::expr( + Box::new(call( + rumoca_core::BuiltinFunction::Size, + vec![var("v"), int(1)], + )), + rumoca_core::Span::DUMMY, + ), + ]), + ); + function.add_output(rumoca_core::FunctionParam::new("p", "Real", test_span())); + functions.insert("Polyphase.activePower".to_string(), function); + + let known_ints = FxHashMap::default(); + let known_reals = FxHashMap::default(); + let known_bools = FxHashMap::default(); + let known_enums = FxHashMap::default(); + let mut array_dims = FxHashMap::default(); + array_dims.insert("sensor.v".to_string(), vec![3]); + array_dims.insert("sensor.i".to_string(), vec![3]); + let expr = function_call("Polyphase.activePower", vec![var("v"), var("i")]); + + assert_eq!( + infer_array_dimensions_full_with_functions( + &expr, + &ParamEvalContext::new( + &known_ints, + &known_reals, + &known_bools, + &known_enums, + &array_dims, + &functions, + Some("sensor.y"), + ), + ), + None + ); +} + #[test] fn user_function_integer_eval_error_means_not_constant_evaluable() { let mut functions = FxHashMap::default(); @@ -406,6 +584,7 @@ fn eval_integer_div_builtin_remains_truncating() { known_enums: &known_enums, array_dims: &array_dims, functions: &functions, + user_func_eval_ctx: None, var_context: None, }; assert_eq!(try_eval_integer_with_context(&expr, &ctx), Some(3)); @@ -668,6 +847,7 @@ fn eval_integer_if_uses_canonicalized_enum_condition() { known_enums: &known_enums, array_dims: &FxHashMap::default(), functions: &FxHashMap::default(), + user_func_eval_ctx: None, var_context: Some("pipe.nFMDistributed"), }; @@ -724,6 +904,7 @@ fn eval_integer_if_resolves_unqualified_enum_condition_with_var_context() { known_enums: &known_enums, array_dims: &FxHashMap::default(), functions: &FxHashMap::default(), + user_func_eval_ctx: None, var_context: Some("Bessel.na"), }; @@ -744,6 +925,7 @@ fn eval_integer_prefers_scoped_unqualified_name_over_global_name() { known_enums: &FxHashMap::default(), array_dims: &FxHashMap::default(), functions: &FxHashMap::default(), + user_func_eval_ctx: None, var_context: Some("machine.rotor.converter.orientation"), }; @@ -811,6 +993,7 @@ fn eval_integer_if_handles_integer_builtin_with_scoped_enum_conditions() { known_enums: &known_enums, array_dims: &FxHashMap::default(), functions: &FxHashMap::default(), + user_func_eval_ctx: None, var_context: Some("Bessel.na"), }; @@ -876,12 +1059,33 @@ fn eval_integer_field_access_resolves_overqualified_suffix() { known_enums: &FxHashMap::default(), array_dims: &FxHashMap::default(), functions: &FxHashMap::default(), + user_func_eval_ctx: None, var_context: Some("stack.cell[1,1].cell.cellData.nRC"), }; assert_eq!(try_eval_integer_with_context(&expr, &ctx), Some(2)); } +#[test] +fn eval_integer_nout_resolves_from_scoped_columns_dimension() { + let known_ints = FxHashMap::default(); + let mut array_dims = FxHashMap::default(); + array_dims.insert("stack.cell[1,2].cell.ocv_soc.columns".to_string(), vec![1]); + + let ctx = ParamEvalContext { + known_ints: &known_ints, + known_reals: &FxHashMap::default(), + known_bools: &FxHashMap::default(), + known_enums: &FxHashMap::default(), + array_dims: &array_dims, + functions: &FxHashMap::default(), + user_func_eval_ctx: None, + var_context: Some("stack.cell[1,2].cell.ocv_soc.y"), + }; + + assert_eq!(try_eval_integer_with_context(&var("nout"), &ctx), Some(1)); +} + #[test] fn eval_integer_if_returns_common_value_when_condition_unknown() { let mut known_ints = FxHashMap::default(); @@ -900,6 +1104,7 @@ fn eval_integer_if_returns_common_value_when_condition_unknown() { known_enums: &FxHashMap::default(), array_dims: &FxHashMap::default(), functions: &FxHashMap::default(), + user_func_eval_ctx: None, var_context: None, }; @@ -1010,6 +1215,35 @@ fn infer_array_dims_from_comprehension_range_and_body() { assert_eq!(dims, Some(vec![4])); } +#[test] +fn infer_array_dims_from_comprehension_range_using_size_of_array_literal() { + let literal = rumoca_core::Expression::Array { + elements: vec![int(1)], + is_matrix: false, + span: rumoca_core::Span::DUMMY, + }; + let expr = rumoca_core::Expression::ArrayComprehension { + expr: Box::new(var("i")), + indices: vec![rumoca_core::ComprehensionIndex { + name: "i".to_string(), + range: rumoca_core::Expression::Range { + start: Box::new(int(1)), + step: None, + end: Box::new(call( + rumoca_core::BuiltinFunction::Size, + vec![literal, int(1)], + )), + span: rumoca_core::Span::DUMMY, + }, + }], + filter: None, + span: rumoca_core::Span::DUMMY, + }; + + let dims = infer_array_dimensions(&expr); + assert_eq!(dims, Some(vec![1])); +} + #[test] fn infer_array_dims_with_context_resolves_scoped_if_matrix_columns() { let mut known_ints = FxHashMap::default(); diff --git a/crates/rumoca-eval-solve/src/compute_block_scalarize.rs b/crates/rumoca-eval-solve/src/compute_block_scalarize.rs index 6aaf6d7c0..c006f31bd 100644 --- a/crates/rumoca-eval-solve/src/compute_block_scalarize.rs +++ b/crates/rumoca-eval-solve/src/compute_block_scalarize.rs @@ -259,6 +259,16 @@ struct ScalarProgramCollector { next_output: usize, } +struct LinSolveScalarizeInput<'a> { + setup_ops: &'a [LinearOp], + matrix_start: Reg, + rhs_start: Reg, + n: usize, + next_reg: Reg, + output_indices: &'a [usize], + span: rumoca_core::Span, +} + impl ScalarProgramCollector { fn append_scalar_program_block( &mut self, @@ -316,6 +326,32 @@ impl ScalarProgramCollector { append_vec(&mut self.output_indices, &mut output_indices, kind, span)?; append_vec(&mut self.rows, &mut programs, kind, span) } + + fn append_linsolve_program( + &mut self, + input: LinSolveScalarizeInput<'_>, + ) -> Result<(), ScalarizeError> { + let LinSolveScalarizeInput { + setup_ops, + matrix_start, + rhs_start, + n, + next_reg, + output_indices, + span, + } = input; + if n == 0 { + return Ok(()); + } + let program = scalarize_linsolve(setup_ops, matrix_start, rhs_start, n, next_reg, span)?; + let output_indices = if output_indices.is_empty() { + let end = checked_contiguous_output_count(self.next_output, n, "linsolve", span)?; + (self.next_output..end).collect() + } else { + output_indices.to_vec() + }; + self.append_tensor_programs(vec![program], output_indices, span, "linsolve") + } } impl SolveVisitor for ScalarProgramCollector { @@ -365,18 +401,18 @@ impl SolveVisitor for ScalarProgramCollector { rhs_start, n, next_reg, + output_indices, span, .. - } => { - if *n == 0 { - return Ok(()); - } - let program = - scalarize_linsolve(setup_ops, *matrix_start, *rhs_start, *n, *next_reg, *span)?; - let start = self.next_output; - let end = checked_contiguous_output_count(start, *n, "linsolve", *span)?; - self.append_contiguous_programs(vec![program], start, end, *span, "linsolve")?; - } + } => self.append_linsolve_program(LinSolveScalarizeInput { + setup_ops, + matrix_start: *matrix_start, + rhs_start: *rhs_start, + n: *n, + next_reg: *next_reg, + output_indices, + span: *span, + })?, ComputeNode::Map { domain, output_map, diff --git a/crates/rumoca-eval-solve/src/compute_block_scalarize/affine.rs b/crates/rumoca-eval-solve/src/compute_block_scalarize/affine.rs index dea66b9e3..5f516a792 100644 --- a/crates/rumoca-eval-solve/src/compute_block_scalarize/affine.rs +++ b/crates/rumoca-eval-solve/src/compute_block_scalarize/affine.rs @@ -355,6 +355,7 @@ fn linear_op_name(op: &LinearOp) -> &'static str { LinearOp::ImpureRandomInit { .. } => "ImpureRandomInit", LinearOp::ImpureRandom { .. } => "ImpureRandom", LinearOp::ImpureRandomInteger { .. } => "ImpureRandomInteger", + LinearOp::ExternalCall { .. } => "ExternalCall", } } diff --git a/crates/rumoca-eval-solve/src/compute_block_scalarize/dense.rs b/crates/rumoca-eval-solve/src/compute_block_scalarize/dense.rs index 689325145..c87e99353 100644 --- a/crates/rumoca-eval-solve/src/compute_block_scalarize/dense.rs +++ b/crates/rumoca-eval-solve/src/compute_block_scalarize/dense.rs @@ -297,6 +297,12 @@ fn max_reg_in_op( imax, .. } => dst.max(id).max(imin).max(imax), + LinearOp::ExternalCall { + dst, + args, + arg_count, + .. + } => args.iter().take(arg_count).copied().fold(dst, Reg::max), LinearOp::StoreOutput { src } => src, }) } diff --git a/crates/rumoca-eval-solve/src/compute_block_scalarize/tests.rs b/crates/rumoca-eval-solve/src/compute_block_scalarize/tests.rs index 62b0cb27d..a1aa82a9f 100644 --- a/crates/rumoca-eval-solve/src/compute_block_scalarize/tests.rs +++ b/crates/rumoca-eval-solve/src/compute_block_scalarize/tests.rs @@ -144,6 +144,7 @@ fn linsolve_scalarizes_to_one_program_with_unique_components() { rhs_start: (n * n) as Reg, n, next_reg, + output_indices: Vec::new(), metadata: TensorNodeMetadata::default(), span: rumoca_core::Span::DUMMY, }], @@ -169,6 +170,26 @@ fn linsolve_scalarizes_to_one_program_with_unique_components() { assert_eq!(components, n); } +#[test] +fn linsolve_scalarization_preserves_noncontiguous_output_indices() { + let block = ComputeBlock { + nodes: vec![ComputeNode::LinSolve { + setup_ops: load_p_ops(0, 6), + matrix_start: 0, + rhs_start: 4, + n: 2, + next_reg: 6, + output_indices: vec![0, 2], + metadata: TensorNodeMetadata::default(), + span: rumoca_core::Span::DUMMY, + }], + }; + + let scalar = to_scalar_program_block(&block).expect("mapped LinSolve should scalarize"); + assert_eq!(scalar.output_indices, vec![0, 2]); + assert_eq!(block.len().expect("mapped output count should be valid"), 3); +} + #[test] fn affine_stencil_expands_to_exact_scalar_rows() { let block = ComputeBlock { diff --git a/crates/rumoca-eval-solve/src/jacobian.rs b/crates/rumoca-eval-solve/src/jacobian.rs index e7b4f21c7..f7a4d0c7a 100644 --- a/crates/rumoca-eval-solve/src/jacobian.rs +++ b/crates/rumoca-eval-solve/src/jacobian.rs @@ -186,7 +186,7 @@ impl SolveRuntime { ) -> ParameterJacobianReport { let layout = &self.model.problem.layout; let n_state = self.state_count; - let y_scalars = layout.y_scalars(); + let solver_y_scalars = self.solver_count; let row_labels = (0..n_state).map(|index| self.state_label(index)).collect(); // One column per model parameter (rumoca-internal `__`-prefixed slots are // excluded); each carries its P-slot so callers can map a column back to @@ -196,9 +196,9 @@ impl SolveRuntime { model_params.into_iter().unzip(); let n_param = param_slots.len(); - // The seed must reach the largest P-slot index, not just `n_param`, since - // model and internal parameter slots may be interleaved. - let seed_len = y_scalars.saturating_add(layout.p_scalars()); + // The seed spans the runtime `[solver-y | parameter]` vector, not just + // states; model and internal parameter slots may also be interleaved. + let seed_len = solver_y_scalars.saturating_add(layout.p_scalars()); let buffers = rect_jacobian_buffers(n_state, n_param, seed_len); let (mut matrix, mut seed, mut column) = match buffers { Ok(buffers) => buffers, @@ -215,14 +215,14 @@ impl SolveRuntime { }; let mut error = None; for (col, slot) in param_slots.iter().copied().enumerate() { - seed[y_scalars + slot] = 1.0; + seed[solver_y_scalars + slot] = 1.0; let result = self.eval_full_jacobian_v_ad_into( AlgebraicLinearization { t, params, settle }, state, &seed, &mut column, ); - seed[y_scalars + slot] = 0.0; + seed[solver_y_scalars + slot] = 0.0; match result { Ok(()) => write_column(&mut matrix, col, &column), Err(err) => { diff --git a/crates/rumoca-eval-solve/src/lib.rs b/crates/rumoca-eval-solve/src/lib.rs index 96a02e705..aca42e288 100644 --- a/crates/rumoca-eval-solve/src/lib.rs +++ b/crates/rumoca-eval-solve/src/lib.rs @@ -1,10 +1,4 @@ //! Solve-IR row evaluation. -//! -//! ## Threading and Test Isolation -//! -//! Simulation facades should pass model-local table data through -//! [`RowEvalContext::external_tables`]. Impure random-generator streams are -//! carried by [`SimulationRuntimeState`] and are never process-global. use std::{ collections::BTreeMap, @@ -16,7 +10,7 @@ use std::{ }; use rumoca_ir_solve::{ - BinaryOp, CompareOp, LinearOp, Reg, ScalarProgramBlock, SolveEventActionKind, + BinaryOp, CompareOp, ComputeNode, LinearOp, Reg, ScalarProgramBlock, SolveEventActionKind, SolveEventMessagePart, SolveEventPartition, SolveProblemShapeContractError, UnaryOp, resolve_indexed_slot, }; @@ -28,6 +22,7 @@ mod iterative_solve; pub mod jacobian; mod linear_solve; pub mod nan_trace; +mod native_map; mod prepared; mod random_runtime; mod refresh_plan; @@ -48,6 +43,7 @@ pub use jacobian::{ JacobianReport, ObjectiveGradientReport, ParameterJacobianReport, SteadyStateSensitivityReport, }; use linear_solve::{solve_component_op, solve_component_unchecked}; +pub use native_map::{MapEvaluationMetrics, eval_map_elements_with_context}; pub use prepared::{ PreparedComputeBlock, PreparedScalarProgramBlock, TargetAssignmentShape, target_assignment_shape, @@ -56,10 +52,11 @@ use random_runtime::{ ImpureRandomState, impure_random_mutex, impure_random_sample, impure_random_stream_id, initial_state_values, projected_random_value, random_result_and_state, read_reg_range, }; +pub use refresh_plan::algebraic_projection_producer_programs; pub use runtime::{ AlgebraicLinearization, AlgebraicSettle, EventUpdateRowFilter, InitialEventObservation, ProjectedEventUpdateInput, ProjectedInitialEventInput, ProjectedInitialEventOutcome, - SolveRuntime, apply_discrete_slot_value, + ProjectedRuntimeSettleInput, RootSearchInput, SolveRuntime, apply_discrete_slot_value, }; pub use runtime_events::{ apply_discrete_slot_values, current_dynamic_time_event_stop, eval_event_actions_with_context, @@ -98,6 +95,9 @@ pub enum EvalSolveError { column: Option, reason: String, }, + ExternalFunction { + function: rumoca_ir_solve::ExternalFunctionKind, + }, MissingInput { vector: &'static str, index: usize, @@ -255,6 +255,12 @@ impl std::fmt::Display for EvalSolveError { ) } } + Self::ExternalFunction { function } => { + write!( + f, + "native external solve function {function:?} requires an external runtime bridge" + ) + } Self::MissingInput { vector, index, len, .. } => write!( @@ -620,6 +626,7 @@ impl<'out> OutputCursor<'out> { pub(crate) struct RowEvalScratch { pub(crate) regs: Vec, pub(crate) initialized: Vec, + pub(crate) values: Vec, } #[derive(Clone, Copy)] @@ -970,6 +977,9 @@ impl CheckedRowEvaluator<'_, '_, '_, '_> { | LinearOp::ImpureRandomInteger { .. } => { self.eval_random_op(&op)?; } + LinearOp::ExternalCall { function, .. } => { + return Err(EvalSolveError::ExternalFunction { function }); + } LinearOp::StoreOutput { src } => { let value = self.get(src)?; self.sink.store(value)?; @@ -978,6 +988,43 @@ impl CheckedRowEvaluator<'_, '_, '_, '_> { Ok(()) } + fn eval_affine_op( + &mut self, + position: usize, + op: LinearOp, + offsets: native_map::AffineMapOffsets<'_>, + ) -> Result<(), EvalSolveError> { + match op { + LinearOp::LoadY { dst, index } => { + let index = index + .checked_add_signed(native_map::affine_load_offset(position, offsets)?) + .ok_or_else(|| { + native_map::affine_map_error( + "affine Y load index overflows host range", + offsets.span, + ) + })?; + self.set(dst, self.input.read_input("y", self.input.y, index)?) + } + LinearOp::LoadP { dst, index } => { + let index = index + .checked_add_signed(native_map::affine_load_offset(position, offsets)?) + .ok_or_else(|| { + native_map::affine_map_error( + "affine P load index overflows host range", + offsets.span, + ) + })?; + self.set(dst, self.input.read_input("p", self.input.p, index)?) + } + LinearOp::Const { dst, value } => self.set( + dst, + value + native_map::affine_const_offset(position, offsets)?, + ), + _ => self.eval_op(op), + } + } + fn eval_table_op(&mut self, op: &LinearOp) -> Result<(), EvalSolveError> { apply_table_op( self.regs, @@ -1134,6 +1181,9 @@ fn eval_row_prepared_fast( | LinearOp::ImpureRandomInteger { .. } => { eval_fast_random_op(regs, &mut scratch.initialized, input, op)? } + LinearOp::ExternalCall { function, .. } => { + return Err(EvalSolveError::ExternalFunction { function }); + } LinearOp::StoreOutput { src } => sink.store(regs[src as usize])?, } } @@ -1376,6 +1426,7 @@ fn linear_op_name(op: &LinearOp) -> &'static str { LinearOp::ImpureRandomInit { .. } => "ImpureRandomInit", LinearOp::ImpureRandom { .. } => "ImpureRandom", LinearOp::ImpureRandomInteger { .. } => "ImpureRandomInteger", + LinearOp::ExternalCall { .. } => "ExternalCall", LinearOp::Unary { .. } => "Unary", LinearOp::Binary { .. } => "Binary", LinearOp::Compare { .. } => "Compare", @@ -1647,6 +1698,16 @@ fn max_register(op: &LinearOp) -> Result { imax, .. } => Ok(dst.max(id).max(imin).max(imax)), + LinearOp::ExternalCall { + dst, + args, + arg_count, + .. + } => args + .iter() + .take(arg_count) + .copied() + .try_fold(dst, |max_reg, arg| Ok(max_reg.max(arg))), LinearOp::StoreOutput { src } => Ok(src), } } @@ -1805,6 +1866,13 @@ fn op_sources_initialized(op: &LinearOp, initialized: &[bool]) -> bool { && reg_initialized(initialized, imin) && reg_initialized(initialized, imax) } + LinearOp::ExternalCall { + args, arg_count, .. + } => args + .iter() + .copied() + .take(arg_count) + .all(|arg| reg_initialized(initialized, arg)), } } diff --git a/crates/rumoca-eval-solve/src/native_map.rs b/crates/rumoca-eval-solve/src/native_map.rs new file mode 100644 index 000000000..84deac532 --- /dev/null +++ b/crates/rumoca-eval-solve/src/native_map.rs @@ -0,0 +1,316 @@ +//! Native, allocation-bounded evaluation of structured Solve-IR Maps. + +use super::*; + +/// Execute a compact Solve-IR `Map` directly over its structured domain. +/// +/// This is intentionally distinct from scalarization: it keeps the map's +/// single `base_ops` owner and applies affine load/constant offsets while each +/// element is evaluated. The callback sees a borrowed ordinal that is reused +/// for every element, so this path never creates a scalar-row vector or a +/// per-cell `LinearOp` clone. +pub fn eval_map_elements_with_context( + node: &ComputeNode, + y: &mut [f64], + p: &[f64], + t: f64, + context: RowEvalContext<'_>, + mut visit: impl FnMut(&[usize], f64, &mut [f64]) -> Result<(), EvalSolveError>, +) -> Result { + let ComputeNode::Map { + domain, + base_ops, + load_strides, + const_strides, + span, + .. + } = node + else { + return Err(EvalSolveError::InvalidRow { + message: "native map evaluation requires a ComputeNode::Map".to_string(), + span: None, + }); + }; + let local_runtime_state; + let context = match context.runtime_state { + Some(_) => context, + None => { + local_runtime_state = SimulationRuntimeState::new(); + context.with_runtime_state(&local_runtime_state) + } + }; + validate_affine_map_metadata(domain, base_ops, load_strides, const_strides, *span)?; + let counts = map_domain_counts(domain, *span)?; + if counts.contains(&0) { + return Ok(MapEvaluationMetrics::default()); + } + let register_count = required_registers(base_ops)?.max(1); + let mut scratch = RowEvalScratch::default(); + let mut ordinal = vec![0usize; counts.len()]; + let mut metrics = MapEvaluationMetrics { + temporary_values: counts + .len() + .saturating_mul(2) + .saturating_add(register_count), + ..Default::default() + }; + loop { + let mut output = [0.0f64]; + let mut sink = OutputCursor::new(&mut output); + let input = PreparedRowEval::new(base_ops, register_count, y, p, t, context) + .with_source_span(Some(*span)); + eval_affine_map_row( + input, + &mut scratch, + &mut sink, + AffineMapOffsets { + ordinal: &ordinal, + load_strides, + const_strides, + span: *span, + }, + )?; + visit(&ordinal, output[0], y)?; + metrics.elements = metrics.elements.saturating_add(1); + if increment_map_ordinal(&mut ordinal, &counts) { + break; + } + } + Ok(metrics) +} + +/// Deterministic resource counters for native map evaluation. They are used +/// by compact-runtime callers to assert linear traversal rather than relying +/// on host timing. +#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)] +pub struct MapEvaluationMetrics { + pub elements: usize, + pub temporary_values: usize, +} + +pub(super) fn affine_load_offset( + position: usize, + offsets: AffineMapOffsets<'_>, +) -> Result { + offsets + .load_strides + .iter() + .filter(|stride| stride.op_position == position) + .try_fold(0isize, |total, terms| { + checked_affine_index_offset(total, &terms.terms, offsets.ordinal, offsets.span) + }) +} + +pub(super) fn affine_const_offset( + position: usize, + offsets: AffineMapOffsets<'_>, +) -> Result { + offsets + .const_strides + .iter() + .filter(|stride| stride.op_position == position) + .try_fold(0.0f64, |total, terms| { + checked_affine_const_offset(total, &terms.terms, offsets.ordinal, offsets.span) + }) +} + +#[derive(Clone, Copy)] +pub(super) struct AffineMapOffsets<'a> { + pub(super) ordinal: &'a [usize], + pub(super) load_strides: &'a [rumoca_ir_solve::AffineStencilLoadStride], + pub(super) const_strides: &'a [rumoca_ir_solve::AffineStencilConstStride], + pub(super) span: rumoca_core::Span, +} + +fn eval_affine_map_row( + input: PreparedRowEval<'_, '_>, + scratch: &mut RowEvalScratch, + sink: &mut OutputCursor<'_>, + offsets: AffineMapOffsets<'_>, +) -> Result<(), EvalSolveError> { + scratch.regs.resize(input.register_count, 0.0); + scratch.initialized.resize(input.register_count, false); + scratch.regs.fill(0.0); + scratch.initialized.fill(false); + let mut evaluator = CheckedRowEvaluator { + regs: &mut scratch.regs, + initialized: &mut scratch.initialized, + input, + sink, + }; + for (position, op) in evaluator.input.row.iter().copied().enumerate() { + evaluator.eval_affine_op(position, op, offsets)?; + } + Ok(()) +} + +fn checked_affine_index_offset( + total: isize, + terms: &[rumoca_ir_solve::AffineStencilIndexStrideTerm], + ordinal: &[usize], + span: rumoca_core::Span, +) -> Result { + terms.iter().try_fold(total, |sum, term| { + let coordinate = *ordinal.get(term.dimension).ok_or_else(|| { + affine_map_error( + "affine load stride references a missing domain dimension", + span, + ) + })?; + let coordinate = isize::try_from(coordinate) + .map_err(|_| affine_map_error("affine load ordinal overflows isize", span))?; + let offset = term + .stride + .checked_mul(coordinate) + .ok_or_else(|| affine_map_error("affine load stride overflows isize", span))?; + sum.checked_add(offset) + .ok_or_else(|| affine_map_error("affine load offset overflows isize", span)) + }) +} + +fn checked_affine_const_offset( + total: f64, + terms: &[rumoca_ir_solve::AffineStencilConstStrideTerm], + ordinal: &[usize], + span: rumoca_core::Span, +) -> Result { + terms.iter().try_fold(total, |sum, term| { + let coordinate = *ordinal.get(term.dimension).ok_or_else(|| { + affine_map_error( + "affine constant stride references a missing domain dimension", + span, + ) + })?; + let offset = term.stride * coordinate as f64; + if !offset.is_finite() { + return Err(affine_map_error( + "affine constant offset is non-finite", + span, + )); + } + let next = sum + offset; + next.is_finite() + .then_some(next) + .ok_or_else(|| affine_map_error("affine constant offset is non-finite", span)) + }) +} + +fn validate_affine_map_metadata( + domain: &rumoca_core::StructuredIndexDomain, + base_ops: &[LinearOp], + load_strides: &[rumoca_ir_solve::AffineStencilLoadStride], + const_strides: &[rumoca_ir_solve::AffineStencilConstStride], + span: rumoca_core::Span, +) -> Result<(), EvalSolveError> { + for stride in load_strides { + validate_affine_dimensions(&stride.terms, domain, span)?; + if !matches!( + base_ops.get(stride.op_position), + Some(LinearOp::LoadY { .. } | LinearOp::LoadP { .. }) + ) { + return Err(affine_map_error( + "affine load stride does not point at LoadY or LoadP", + span, + )); + } + } + for stride in const_strides { + validate_affine_dimensions(&stride.terms, domain, span)?; + if !matches!( + base_ops.get(stride.op_position), + Some(LinearOp::Const { .. }) + ) { + return Err(affine_map_error( + "affine constant stride does not point at Const", + span, + )); + } + } + Ok(()) +} + +fn validate_affine_dimensions( + terms: &[T], + domain: &rumoca_core::StructuredIndexDomain, + span: rumoca_core::Span, +) -> Result<(), EvalSolveError> { + if terms + .iter() + .any(|term| term.dimension() >= domain.binders.len()) + { + return Err(affine_map_error( + "affine stride references a missing domain dimension", + span, + )); + } + Ok(()) +} + +trait AffineDimension { + fn dimension(&self) -> usize; +} + +impl AffineDimension for rumoca_ir_solve::AffineStencilIndexStrideTerm { + fn dimension(&self) -> usize { + self.dimension + } +} + +impl AffineDimension for rumoca_ir_solve::AffineStencilConstStrideTerm { + fn dimension(&self) -> usize { + self.dimension + } +} + +fn map_domain_counts( + domain: &rumoca_core::StructuredIndexDomain, + span: rumoca_core::Span, +) -> Result, EvalSolveError> { + domain + .binders + .iter() + .map(|binder| { + if binder.step == 0 { + return Err(affine_map_error("structured map binder step is zero", span)); + } + let distance = if binder.step > 0 { + binder.upper.checked_sub(binder.lower) + } else { + binder.lower.checked_sub(binder.upper) + } + .ok_or_else(|| { + affine_map_error("structured map binder bounds contradict step", span) + })?; + let step = i64::try_from(binder.step.unsigned_abs()) + .map_err(|_| affine_map_error("structured map binder step overflows i64", span))?; + let count = distance + .checked_div(step) + .and_then(|count| count.checked_add(1)) + .ok_or_else(|| affine_map_error("structured map binder count overflows", span))?; + usize::try_from(count).map_err(|_| { + affine_map_error("structured map binder count exceeds host range", span) + }) + }) + .collect() +} + +fn increment_map_ordinal(ordinal: &mut [usize], counts: &[usize]) -> bool { + for dimension in (0..ordinal.len()).rev() { + ordinal[dimension] += 1; + if ordinal[dimension] < counts[dimension] { + return false; + } + ordinal[dimension] = 0; + } + true +} + +pub(super) fn affine_map_error( + message: impl Into, + span: rumoca_core::Span, +) -> EvalSolveError { + EvalSolveError::InvalidRow { + message: message.into(), + span: Some(span), + } +} diff --git a/crates/rumoca-eval-solve/src/prepared.rs b/crates/rumoca-eval-solve/src/prepared.rs index 2724fc444..5c11cde91 100644 --- a/crates/rumoca-eval-solve/src/prepared.rs +++ b/crates/rumoca-eval-solve/src/prepared.rs @@ -490,6 +490,24 @@ impl PreparedScalarProgramBlock { self.block.output_indices.get(stored_ordinal).copied() } + pub fn program_position_for_output_index(&self, output_index: usize) -> Option<(usize, usize)> { + let mut stored_ordinal = 0usize; + let mut found = None; + for (program_index, row) in self.block.programs.iter().enumerate() { + for output_offset in 0..ScalarProgramBlock::program_output_count(row) { + let is_requested_output = + self.block.output_indices.get(stored_ordinal).copied() == Some(output_index); + record_unique_program_position( + &mut found, + (program_index, output_offset), + is_requested_output, + )?; + stored_ordinal = stored_ordinal.checked_add(1)?; + } + } + found + } + pub fn can_evaluate_target_assignment(&self, row_idx: usize, target_y_index: usize) -> bool { let Some(row) = self.block.programs.get(row_idx) else { return false; @@ -802,6 +820,21 @@ impl PreparedScalarProgramBlock { } } +fn record_unique_program_position( + found: &mut Option<(usize, usize)>, + position: (usize, usize), + is_requested_output: bool, +) -> Option<()> { + if !is_requested_output { + return Some(()); + } + if found.is_some() { + return None; + } + *found = Some(position); + Some(()) +} + struct RowEvalRequest<'a> { row_idx: usize, y: &'a [f64], @@ -1324,7 +1357,7 @@ enum PreparedComputeNode { setup: PreparedLinearOps, matrix_start: u32, rhs_start: u32, - output_start: usize, + output_indices: Vec, matrix_len: usize, n: usize, }, @@ -1422,18 +1455,39 @@ fn prepared_linsolve( matrix_start: u32, rhs_start: u32, n: usize, + output_indices: &[usize], span: rumoca_core::Span, output_cursor: usize, ) -> Result<(PreparedComputeNode, usize), EvalSolveError> { let matrix_len = checked_product(n, n, "prepared linsolve matrix", span)?; - let next_output_cursor = - checked_contiguous_output_count(output_cursor, n, "prepared linsolve output", span)?; + let output_indices = if output_indices.is_empty() { + let end = + checked_contiguous_output_count(output_cursor, n, "prepared linsolve output", span)?; + (output_cursor..end).collect() + } else { + output_indices.to_vec() + }; + if output_indices.len() != n { + return Err(EvalSolveError::ShapeContract { + message: format!( + "prepared LinSolve has {n} components but {} output indices", + output_indices.len() + ), + span: Some(span), + }); + } + let next_output_cursor = output_cursor.max(checked_tensor_output_count( + &output_indices, + output_cursor, + "prepared linsolve output", + span, + )?); Ok(( PreparedComputeNode::LinSolve { setup: PreparedLinearOps::new(setup_ops.to_vec())?, matrix_start, rhs_start, - output_start: output_cursor, + output_indices, matrix_len, n, }, @@ -1511,6 +1565,7 @@ impl PreparedComputeNode { matrix_start, rhs_start, n, + output_indices, span, .. } => prepared_linsolve( @@ -1518,6 +1573,7 @@ impl PreparedComputeNode { *matrix_start, *rhs_start, *n, + output_indices, *span, output_cursor, )?, @@ -1574,9 +1630,14 @@ impl PreparedComputeNode { output_len, .. } => Some((*output_start, *output_len)), - Self::LinSolve { - output_start, n, .. - } => Some((*output_start, *n)), + Self::LinSolve { output_indices, .. } => { + let start = *output_indices.first()?; + output_indices + .iter() + .copied() + .eq(start..start.checked_add(output_indices.len())?) + .then_some((start, output_indices.len())) + } Self::ScalarPrograms(_) => None, } } @@ -1628,23 +1689,25 @@ impl PreparedComputeNode { setup, matrix_start, rhs_start, - output_start, + output_indices, matrix_len, n, } => { setup.eval(y, p, t, context, scratch)?; ensure_register_range(&scratch.regs, "read", *matrix_start, *matrix_len)?; ensure_register_range(&scratch.regs, "read", *rhs_start, *n)?; - let output_end = output_start.checked_add(*n).ok_or_else(|| { - invalid_prepared_row("prepared linsolve output range overflows") - })?; + scratch.values.resize(*n, 0.0); solve_all_unchecked( &scratch.regs, *matrix_start, *rhs_start, *n, - &mut out[*output_start..output_end], - ) + &mut scratch.values, + )?; + for (value, output_index) in scratch.values.iter().zip(output_indices) { + out[*output_index] = *value; + } + Ok(()) } } } diff --git a/crates/rumoca-eval-solve/src/prepared/assignment_shape_tests.rs b/crates/rumoca-eval-solve/src/prepared/assignment_shape_tests.rs index f31b3d16e..162027ad6 100644 --- a/crates/rumoca-eval-solve/src/prepared/assignment_shape_tests.rs +++ b/crates/rumoca-eval-solve/src/prepared/assignment_shape_tests.rs @@ -9,6 +9,27 @@ fn fixture_span() -> rumoca_core::Span { ) } +#[test] +fn output_index_maps_to_program_position_when_outputs_are_reordered() { + let row = |value| { + vec![ + LinearOp::Const { dst: 0, value }, + LinearOp::StoreOutput { src: 0 }, + ] + }; + let block = rumoca_ir_solve::ScalarProgramBlock::with_output_indices( + vec![row(1.0), row(2.0)], + vec![fixture_span(), fixture_span()], + vec![5, 2], + ) + .expect("reordered output fixture should be valid"); + let prepared = PreparedScalarProgramBlock::new(block).expect("fixture should prepare"); + + assert_eq!(prepared.program_position_for_output_index(5), Some((0, 0))); + assert_eq!(prepared.program_position_for_output_index(2), Some((1, 0))); + assert_eq!(prepared.program_position_for_output_index(0), None); +} + // Regression: `reg_depends_on_y_index` used to recurse over the register DAG // without memoization, so a row whose affine coefficient/offset is a deeply // shared sub-expression (typical of inlined matrix products) took O(2^depth) diff --git a/crates/rumoca-eval-solve/src/prepared/dependency.rs b/crates/rumoca-eval-solve/src/prepared/dependency.rs index b61654dd0..19d87c38d 100644 --- a/crates/rumoca-eval-solve/src/prepared/dependency.rs +++ b/crates/rumoca-eval-solve/src/prepared/dependency.rs @@ -50,13 +50,7 @@ fn reg_depends_on_y_index_memo( rhs_start, n, .. - } => { - let Some(matrix_len) = n.checked_mul(n) else { - return true; - }; - reg_range_depends_on_y_index(row, matrix_start, matrix_len, target_y_index, memo) - || reg_range_depends_on_y_index(row, rhs_start, n, target_y_index, memo) - } + } => linear_solve_depends_on_y_index(row, matrix_start, rhs_start, n, target_y_index, memo), LinearOp::TableBounds { table_id, .. } => { reg_depends_on_y_index_memo(row, table_id, target_y_index, memo) } @@ -109,6 +103,9 @@ fn reg_depends_on_y_index_memo( || reg_depends_on_y_index_memo(row, imin, target_y_index, memo) || reg_depends_on_y_index_memo(row, imax, target_y_index, memo) } + LinearOp::ExternalCall { + args, arg_count, .. + } => external_call_depends_on_y_index(row, &args, arg_count, target_y_index, memo), LinearOp::Const { .. } | LinearOp::LoadTime { .. } | LinearOp::LoadP { .. } @@ -119,6 +116,37 @@ fn reg_depends_on_y_index_memo( result } +fn linear_solve_depends_on_y_index( + row: &[LinearOp], + matrix_start: u32, + rhs_start: u32, + n: usize, + target_y_index: usize, + memo: &mut HashMap, +) -> bool { + let Some(matrix_len) = n.checked_mul(n) else { + return true; + }; + reg_range_depends_on_y_index(row, matrix_start, matrix_len, target_y_index, memo) + || reg_range_depends_on_y_index(row, rhs_start, n, target_y_index, memo) +} + +fn external_call_depends_on_y_index( + row: &[LinearOp], + args: &[u32; 8], + arg_count: usize, + target_y_index: usize, + memo: &mut HashMap, +) -> bool { + let Some(effective_args) = args.get(..arg_count) else { + return true; + }; + effective_args + .iter() + .copied() + .any(|arg| reg_depends_on_y_index_memo(row, arg, target_y_index, memo)) +} + fn reg_range_depends_on_y_index( row: &[LinearOp], start: u32, @@ -138,3 +166,50 @@ fn checked_reg_offset(start: u32, offset: usize) -> Option { let offset = u32::try_from(offset).ok()?; start.checked_add(offset) } + +#[cfg(test)] +mod tests { + use rumoca_ir_solve::{ExternalFunctionKind, LinearOp}; + + use super::reg_depends_on_y_index; + + fn external_call(dst: u32, args: [u32; 8], arg_count: usize) -> LinearOp { + LinearOp::ExternalCall { + dst, + function: ExternalFunctionKind::BuildingsEnergyPlusExchange, + args, + arg_count, + output_index: 0, + } + } + + #[test] + fn external_call_depends_on_target_through_effective_argument() { + let row = [ + LinearOp::LoadY { dst: 1, index: 7 }, + LinearOp::Move { dst: 2, src: 1 }, + external_call(3, [2, 0, 0, 0, 0, 0, 0, 0], 1), + ]; + + assert!(reg_depends_on_y_index(&row, 3, 7)); + } + + #[test] + fn external_call_ignores_other_y_and_unused_capacity() { + let row = [ + LinearOp::Const { dst: 1, value: 1.0 }, + LinearOp::LoadY { dst: 2, index: 8 }, + LinearOp::LoadY { dst: 3, index: 7 }, + external_call(4, [1, 2, 3, 0, 0, 0, 0, 0], 2), + ]; + + assert!(!reg_depends_on_y_index(&row, 4, 7)); + } + + #[test] + fn malformed_external_call_arg_count_is_conservatively_dependent() { + let row = [external_call(1, [0; 8], 9)]; + + assert!(reg_depends_on_y_index(&row, 1, 7)); + } +} diff --git a/crates/rumoca-eval-solve/src/prepared/prepared_compute_block_tests.rs b/crates/rumoca-eval-solve/src/prepared/prepared_compute_block_tests.rs index 429b6f90b..857e2a575 100644 --- a/crates/rumoca-eval-solve/src/prepared/prepared_compute_block_tests.rs +++ b/crates/rumoca-eval-solve/src/prepared/prepared_compute_block_tests.rs @@ -51,6 +51,29 @@ fn const_store_row(value: f64) -> Vec { ] } +fn diagonal_linsolve_node(output_indices: Vec) -> ComputeNode { + ComputeNode::LinSolve { + setup_ops: vec![ + LinearOp::Const { dst: 0, value: 2.0 }, + LinearOp::Const { dst: 1, value: 0.0 }, + LinearOp::Const { dst: 2, value: 0.0 }, + LinearOp::Const { dst: 3, value: 4.0 }, + LinearOp::Const { dst: 4, value: 8.0 }, + LinearOp::Const { + dst: 5, + value: 20.0, + }, + ], + matrix_start: 0, + rhs_start: 4, + n: 2, + next_reg: 6, + output_indices, + metadata: TensorNodeMetadata::default(), + span: test_span("prepared_linsolve.mo"), + } +} + #[test] fn prepared_vec_with_capacity_rejects_impossible_capacity_with_span() { let span = Span::from_offsets(SourceId::from_source_name("prepared.mo"), 3, 9); @@ -107,6 +130,22 @@ fn prepared_compute_block_evaluates_map_through_scalar_view() { assert_eq!(out, vec![10.0, 20.0, 30.0]); } +#[test] +fn prepared_compute_block_scatters_noncontiguous_linsolve_outputs() { + let block = ComputeBlock { + nodes: vec![diagonal_linsolve_node(vec![0, 2])], + }; + let prepared = PreparedComputeBlock::new(&block) + .expect("noncontiguous LinSolve should prepare without scalarization"); + let mut out = vec![99.0; prepared.len()]; + + prepared + .eval_with_context(&[], &[], 0.0, RowEvalContext::default(), &mut out) + .expect("prepared LinSolve should evaluate and scatter its components"); + + assert_eq!(out, vec![4.0, 0.0, 5.0]); +} + #[test] fn prepared_compute_block_writes_sparse_map_output_slots() { let domain = test_domain(); diff --git a/crates/rumoca-eval-solve/src/refresh_plan.rs b/crates/rumoca-eval-solve/src/refresh_plan.rs index 86ecfae6d..14894ad3a 100644 --- a/crates/rumoca-eval-solve/src/refresh_plan.rs +++ b/crates/rumoca-eval-solve/src/refresh_plan.rs @@ -1,4 +1,7 @@ -use std::{collections::VecDeque, sync::Arc}; +use std::{ + collections::{BTreeMap, VecDeque}, + sync::Arc, +}; use indexmap::{IndexMap, IndexSet}; use rumoca_ir_solve as solve; @@ -77,6 +80,19 @@ pub(crate) struct AlgebraicRefreshRow { /// must linear-solve the row's residual for the paired variable instead /// of evaluating the assignment value. pub(crate) assignment_target: Option, + /// Alternate residual rows from the same projection block that structurally + /// reference this target. Parameter-dependent `select` branches can make + /// the primary structural match inactive at runtime; candidates preserve + /// the residual-dependence guard while letting the runtime use the active + /// row for the current parameter point. + pub(crate) alternatives: Vec, +} + +#[derive(Clone)] +pub(crate) struct AlgebraicRefreshCandidate { + pub(crate) row_idx: usize, + pub(crate) output_offset: usize, + pub(crate) assignment_target: Option, } #[derive(Clone, Default)] @@ -98,26 +114,27 @@ pub(crate) fn build_algebraic_refresh_plan( block: &PreparedScalarProgramBlock, ) -> Result { let state_count = model.state_scalar_count(); - let row_target_rows = algebraic_refresh_rows_from_row_targets(model, block, state_count)?; + let projection_rows = algebraic_refresh_rows_from_projection_plan(model, block, state_count)?; + let span = first_block_span(block.block()); let mut rows_by_target = IndexMap::new(); reserve_refresh_index_map_capacity( &mut rows_by_target, - row_target_rows.len(), - "row-target map", - first_block_span(block.block()), + projection_rows.len(), + "projection target map", + span, )?; - for row in row_target_rows { - rows_by_target.insert(row.target_index, row); + for row in projection_rows { + merge_refresh_row(&mut rows_by_target, row, span)?; } - let projection_rows = algebraic_refresh_rows_from_projection_plan(model, block, state_count)?; + let row_target_rows = algebraic_refresh_rows_from_row_targets(model, block, state_count)?; reserve_refresh_index_map_capacity( &mut rows_by_target, - projection_rows.len(), - "projection target map", - first_block_span(block.block()), + row_target_rows.len(), + "row-target map", + span, )?; - for row in projection_rows { - rows_by_target.insert(row.target_index, row); + for row in row_target_rows { + merge_refresh_row(&mut rows_by_target, row, span)?; } let mut rows = Vec::new(); reserve_refresh_vec_capacity( @@ -130,6 +147,71 @@ pub(crate) fn build_algebraic_refresh_plan( order_refresh_rows(rows, Arc::new(block.block().clone()), state_count) } +/// Return the executable implicit program that produces each algebraic +/// projection target. This is the same producer assignment used by the +/// runtime refresh plan; projection hints only influence matching order. +pub fn algebraic_projection_producer_programs( + model: &solve::SolveModel, +) -> Result, EvalSolveError> { + let block = + PreparedScalarProgramBlock::from_compute_block(&model.problem.continuous.implicit_rhs)?; + let plan = build_algebraic_refresh_plan(model, &block)?; + Ok(plan + .rows + .into_iter() + .map(|row| (row.target_index, row.row_idx)) + .collect()) +} + +fn merge_refresh_row( + rows_by_target: &mut IndexMap, + candidate: AlgebraicRefreshRow, + span: Option, +) -> Result<(), EvalSolveError> { + let target_index = candidate.target_index; + let Some(primary) = rows_by_target.get_mut(&target_index) else { + rows_by_target.insert(target_index, candidate); + return Ok(()); + }; + append_refresh_alternative( + primary, + AlgebraicRefreshCandidate { + row_idx: candidate.row_idx, + output_offset: candidate.output_offset, + assignment_target: candidate.assignment_target, + }, + span, + )?; + for alternative in candidate.alternatives { + append_refresh_alternative(primary, alternative, span)?; + } + Ok(()) +} + +fn append_refresh_alternative( + primary: &mut AlgebraicRefreshRow, + candidate: AlgebraicRefreshCandidate, + span: Option, +) -> Result<(), EvalSolveError> { + let candidate_key = (candidate.row_idx, candidate.output_offset); + if candidate_key == (primary.row_idx, primary.output_offset) + || primary + .alternatives + .iter() + .any(|existing| (existing.row_idx, existing.output_offset) == candidate_key) + { + return Ok(()); + } + reserve_refresh_vec_capacity( + &mut primary.alternatives, + 1, + "merged refresh alternatives", + span, + )?; + primary.alternatives.push(candidate); + Ok(()) +} + fn algebraic_refresh_rows_from_projection_plan( model: &solve::SolveModel, block: &PreparedScalarProgramBlock, @@ -156,63 +238,152 @@ fn algebraic_refresh_rows_from_projection_plan( first_block_span(block.block()), )?; for plan_block in &model.problem.continuous.algebraic_projection_plan.blocks { - // A coupled block's `rows` and `y_indices` are independent sets (the - // unknowns are even sorted), so each row must be paired with an - // unknown it actually determines — positional pairing produces a - // convergent but wrong system (a gear-torque row "assigned" to an - // unrelated flange torque). Pair by maximum bipartite matching over - // real incidence, preferring each row's own implicit target. let span = plan_block_span(block.block(), &plan_block.rows, &output_row_positions); - let mut eligible_ys = Vec::new(); - reserve_refresh_vec_capacity( - &mut eligible_ys, - plan_block.y_indices.len(), - "projection eligible-y list", + let pairs = projection_refresh_pairs( + plan_block, + block, + state_count, + &output_row_positions, + &assignment_target, span, )?; - eligible_ys.extend( - plan_block - .y_indices - .iter() - .copied() - .filter(|&y| y >= state_count), - ); - let pairs = match_block_rows_to_targets( - &plan_block.rows, - &eligible_ys, + let mut block_rows = projection_refresh_rows_for_pairs( + plan_block, + block, + &output_row_positions, &assignment_target, - |row_idx, y| { - output_row_positions - .get(&row_idx) - .is_some_and(|position| block.row_reads_y(position.program_index, y)) - }, + pairs, span, )?; - let row_capacity = pairs - .len() - .checked_add(plan_block.causal_steps.len()) - .ok_or_else(|| refresh_plan_capacity_error("projection refresh rows", span))?; - reserve_refresh_vec_capacity(&mut rows, row_capacity, "projection refresh rows", span)?; - for (row_idx, target_index) in pairs { - let position = required_program_position(&output_row_positions, row_idx, span)?; - rows.push(AlgebraicRefreshRow { - row_idx: position.program_index, - output_offset: position.output_offset, - target_index, - assignment_target: assignment_target(row_idx), - }); - } - for step in &plan_block.causal_steps { - if step.y_index >= state_count { - let position = required_program_position(&output_row_positions, step.row, span)?; - rows.push(AlgebraicRefreshRow { - row_idx: position.program_index, - output_offset: position.output_offset, - target_index: step.y_index, - assignment_target: assignment_target(step.row), - }); - } + reserve_refresh_vec_capacity(&mut rows, block_rows.len(), "projection refresh rows", span)?; + rows.append(&mut block_rows); + } + Ok(rows) +} + +fn projection_refresh_pairs( + plan_block: &solve::AlgebraicProjectionBlock, + block: &PreparedScalarProgramBlock, + state_count: usize, + output_row_positions: &IndexMap, + assignment_target: &dyn Fn(usize) -> Option, + span: Option, +) -> Result, EvalSolveError> { + let eligible_ys = projection_eligible_targets(plan_block, state_count, span)?; + let structural_targets = + projection_structural_target_preferences(plan_block, state_count, span)?; + match_block_rows_to_targets( + &plan_block.rows, + &eligible_ys, + &|row_idx| structural_targets.get(&row_idx).copied(), + assignment_target, + |row_idx, y| { + output_row_positions + .get(&row_idx) + .is_some_and(|position| block.row_reads_y(position.program_index, y)) + }, + span, + ) +} + +fn projection_eligible_targets( + plan_block: &solve::AlgebraicProjectionBlock, + state_count: usize, + span: Option, +) -> Result, EvalSolveError> { + let mut targets = Vec::new(); + reserve_refresh_vec_capacity( + &mut targets, + plan_block.y_indices.len(), + "projection eligible-y list", + span, + )?; + targets.extend( + plan_block + .y_indices + .iter() + .copied() + .filter(|target| *target >= state_count), + ); + Ok(targets) +} + +fn projection_structural_target_preferences( + plan_block: &solve::AlgebraicProjectionBlock, + state_count: usize, + span: Option, +) -> Result, EvalSolveError> { + let mut preferences = IndexMap::new(); + reserve_refresh_index_map_capacity( + &mut preferences, + plan_block.causal_steps.len(), + "projection structural target preferences", + span, + )?; + for step in &plan_block.causal_steps { + validate_projection_causal_step(plan_block, step, span)?; + if step.y_index < state_count { + continue; } + preferences.entry(step.row).or_insert(step.y_index); + } + Ok(preferences) +} + +fn validate_projection_causal_step( + plan_block: &solve::AlgebraicProjectionBlock, + step: &solve::AlgebraicProjectionStep, + span: Option, +) -> Result<(), EvalSolveError> { + if !plan_block.rows.contains(&step.row) { + return Err(EvalSolveError::InvalidRow { + message: format!( + "algebraic projection causal step row {} is not a member of its block", + step.row + ), + span, + }); + } + if !plan_block.y_indices.contains(&step.y_index) { + return Err(EvalSolveError::InvalidRow { + message: format!( + "algebraic projection causal step target y[{}] is not a member of its block", + step.y_index + ), + span, + }); + } + Ok(()) +} + +fn projection_refresh_rows_for_pairs( + plan_block: &solve::AlgebraicProjectionBlock, + block: &PreparedScalarProgramBlock, + output_row_positions: &IndexMap, + assignment_target: &dyn Fn(usize) -> Option, + pairs: Vec<(usize, usize)>, + span: Option, +) -> Result, EvalSolveError> { + let mut rows = Vec::new(); + reserve_refresh_vec_capacity(&mut rows, pairs.len(), "projection block rows", span)?; + for (row_idx, target_index) in pairs { + let position = required_program_position(output_row_positions, row_idx, span)?; + let alternatives = projection_row_alternatives( + block, + &plan_block.rows, + target_index, + output_row_positions, + assignment_target, + position, + span, + )?; + rows.push(AlgebraicRefreshRow { + row_idx: position.program_index, + output_offset: position.output_offset, + target_index, + assignment_target: assignment_target(row_idx), + alternatives, + }); } Ok(rows) } @@ -225,6 +396,7 @@ fn algebraic_refresh_rows_from_projection_plan( fn match_block_rows_to_targets( rows: &[usize], ys: &[usize], + structural_target: &dyn Fn(usize) -> Option, assignment_target: &dyn Fn(usize) -> Option, reads: impl Fn(usize, usize) -> bool, span: Option, @@ -234,37 +406,44 @@ fn match_block_rows_to_targets( for (pos, y) in ys.iter().copied().enumerate() { y_pos.insert(y, pos); } + let mut matched_row_for_y = Vec::new(); + reserve_refresh_vec_capacity( + &mut matched_row_for_y, + ys.len(), + "matching result slots", + span, + )?; + matched_row_for_y.resize(ys.len(), None); + let mut adjacency = Vec::new(); reserve_refresh_vec_capacity(&mut adjacency, rows.len(), "matching adjacency", span)?; for &row_idx in rows { - let own = assignment_target(row_idx).and_then(|y| y_pos.get(&y).copied()); let mut candidates = Vec::new(); - let candidate_capacity = y_pos - .len() - .checked_add(usize::from(own.is_some())) - .ok_or_else(|| refresh_plan_capacity_error("matching candidates", span))?; + let candidate_capacity = y_pos.len(); reserve_refresh_vec_capacity( &mut candidates, candidate_capacity, "matching candidates", span, )?; - candidates.extend(own); + for preferred in [structural_target(row_idx), assignment_target(row_idx)] { + let Some(y) = preferred else { + continue; + }; + if let Some(pos) = y_pos.get(&y).copied() + && reads(row_idx, y) + && !candidates.contains(&pos) + { + candidates.push(pos); + } + } for (&y, &pos) in &y_pos { - if Some(pos) != own && reads(row_idx, y) { + if reads(row_idx, y) && !candidates.contains(&pos) { candidates.push(pos); } } adjacency.push(candidates); } - let mut matched_row_for_y = Vec::new(); - reserve_refresh_vec_capacity( - &mut matched_row_for_y, - ys.len(), - "matching result slots", - span, - )?; - matched_row_for_y.resize(ys.len(), None); fn try_assign( row_pos: usize, adjacency: &[Vec], @@ -355,11 +534,55 @@ fn algebraic_refresh_rows_from_row_targets( output_offset: position.output_offset, target_index, assignment_target: Some(target_index), + alternatives: Vec::new(), }); } Ok(rows) } +fn projection_row_alternatives( + block: &PreparedScalarProgramBlock, + rows: &[usize], + target_index: usize, + output_row_positions: &IndexMap, + assignment_target: &dyn Fn(usize) -> Option, + primary_position: OutputRowPosition, + span: Option, +) -> Result, EvalSolveError> { + let mut alternatives = Vec::new(); + reserve_refresh_vec_capacity( + &mut alternatives, + rows.len(), + "projection row alternatives", + span, + )?; + let mut seen = IndexSet::new(); + reserve_refresh_index_set_capacity( + &mut seen, + rows.len(), + "projection alternative positions", + span, + )?; + for &row_idx in rows { + let position = required_program_position(output_row_positions, row_idx, span)?; + let position_key = (position.program_index, position.output_offset); + let direct_assignment = assignment_target(row_idx) == Some(target_index) + && block.can_evaluate_target_assignment(position.program_index, target_index); + if position == primary_position + || !seen.insert(position_key) + || (!block.row_reads_y(position.program_index, target_index) && !direct_assignment) + { + continue; + } + alternatives.push(AlgebraicRefreshCandidate { + row_idx: position.program_index, + output_offset: position.output_offset, + assignment_target: assignment_target(row_idx), + }); + } + Ok(alternatives) +} + pub(crate) fn build_derivative_refresh_plan( model: &solve::SolveModel, derivative_block: &solve::ScalarProgramBlock, @@ -395,8 +618,8 @@ fn build_dependency_refresh_plan( "dependency target-to-row map", span, )?; - for row in &full_plan.rows { - target_to_row.insert(row.target_index, row.row_idx); + for (row_pos, row) in full_plan.rows.iter().enumerate() { + target_to_row.insert(row.target_index, row_pos); } let mut needed = IndexSet::new(); reserve_refresh_index_set_capacity( @@ -423,17 +646,18 @@ fn build_dependency_refresh_plan( if !needed.insert(index) { continue; } - let Some(row_idx) = target_to_row.get(&index).copied() else { + let Some(row_pos) = target_to_row.get(&index).copied() else { reserve_refresh_index_set_capacity(&mut missing, 1, "dependency missing set", span)?; missing.insert(index); continue; }; - for dep in row_all_y_dependencies(implicit_block, row_idx) { - if dep >= state_count { - reserve_refresh_vec_capacity(&mut stack, 1, "dependency stack", span)?; - stack.push(dep); - } - } + push_refresh_dependencies( + &full_plan.rows[row_pos], + implicit_block, + state_count, + &mut stack, + span, + )?; } let mut rows = Vec::new(); reserve_refresh_vec_capacity( @@ -460,6 +684,23 @@ fn build_dependency_refresh_plan( Ok(plan) } +fn push_refresh_dependencies( + row: &AlgebraicRefreshRow, + block: &solve::ScalarProgramBlock, + state_count: usize, + stack: &mut Vec, + span: Option, +) -> Result<(), EvalSolveError> { + let dependencies = refresh_row_program_indices(row) + .flat_map(|row_idx| row_all_y_dependencies(block, row_idx)) + .filter(|dep| *dep >= state_count); + for dep in dependencies { + reserve_refresh_vec_capacity(stack, 1, "dependency stack", span)?; + stack.push(dep); + } + Ok(()) +} + fn order_refresh_rows( rows: Vec, block: Arc, @@ -483,25 +724,15 @@ fn order_refresh_rows( reserve_refresh_vec_capacity(&mut indegree, rows.len(), "refresh order indegree", span)?; indegree.resize(rows.len(), 0usize); for (row_pos, row) in rows.iter().enumerate() { - let Some(ops) = block.programs.get(row.row_idx) else { - continue; - }; - for dep_index in row_y_dependencies(ops, row.target_index, state_count) { - let Some(&dep_pos) = producer_by_target.get(&dep_index) else { - continue; - }; - if dep_pos == row_pos || edges[dep_pos].contains(&row_pos) { - continue; - } - reserve_refresh_vec_capacity( - &mut edges[dep_pos], - 1, - "refresh order edge list", - row_span(&block, row.row_idx), - )?; - edges[dep_pos].push(row_pos); - indegree[row_pos] += 1; - } + add_refresh_row_edges( + row_pos, + row, + &block, + state_count, + &producer_by_target, + &mut edges, + &mut indegree, + )?; } let mut ready = VecDeque::new(); reserve_refresh_deque_capacity(&mut ready, rows.len(), "refresh order queue", span)?; @@ -529,10 +760,15 @@ fn order_refresh_rows( // would otherwise be left one secant step short of convergence. let mut self_nonlinear = false; for row in &ordered { - if let Some(ops) = block.programs.get(row.row_idx) - && row_is_nonlinear_in_target(ops, row.target_index)? - { - self_nonlinear = true; + for program_idx in refresh_row_program_indices(row) { + if let Some(ops) = block.programs.get(program_idx) + && row_is_nonlinear_in_target(ops, row.target_index)? + { + self_nonlinear = true; + break; + } + } + if self_nonlinear { break; } } @@ -562,6 +798,45 @@ fn order_refresh_rows( }) } +fn add_refresh_row_edges( + row_pos: usize, + row: &AlgebraicRefreshRow, + block: &solve::ScalarProgramBlock, + state_count: usize, + producer_by_target: &IndexMap, + edges: &mut [Vec], + indegree: &mut [usize], +) -> Result<(), EvalSolveError> { + let target_index = row.target_index; + let dependencies = refresh_row_program_indices(row).flat_map(|program_idx| { + block + .programs + .get(program_idx) + .into_iter() + .flat_map(move |ops| { + row_y_dependencies(ops, target_index, state_count) + .map(move |dep_index| (program_idx, dep_index)) + }) + }); + for (program_idx, dep_index) in dependencies { + let Some(&dep_pos) = producer_by_target.get(&dep_index) else { + continue; + }; + if dep_pos == row_pos || edges[dep_pos].contains(&row_pos) { + continue; + } + reserve_refresh_vec_capacity( + &mut edges[dep_pos], + 1, + "refresh order edge list", + row_span(block, program_idx), + )?; + edges[dep_pos].push(row_pos); + indegree[row_pos] += 1; + } + Ok(()) +} + /// How a register's value depends on the refresh row's own target solver-Y slot, /// holding every other unknown fixed — which is exactly how a single refresh row /// is solved for its target. @@ -747,6 +1022,13 @@ fn classify_op_target_dep( || target_dep(classes, imin).depends_on_target() || target_dep(classes, imax).depends_on_target(), ), + Op::ExternalCall { + args, arg_count, .. + } => nonlinear_if( + args.iter() + .take(arg_count) + .any(|arg| target_dep(classes, *arg).depends_on_target()), + ), } } @@ -875,6 +1157,10 @@ fn row_all_y_dependencies( }) } +fn refresh_row_program_indices(row: &AlgebraicRefreshRow) -> impl Iterator + '_ { + std::iter::once(row.row_idx).chain(row.alternatives.iter().map(|candidate| candidate.row_idx)) +} + fn first_block_span(block: &solve::ScalarProgramBlock) -> Option { block.first_source_span() } @@ -894,7 +1180,7 @@ fn plan_block_span( .or_else(|| first_block_span(block)) } -#[derive(Clone, Copy)] +#[derive(Clone, Copy, PartialEq, Eq)] struct OutputRowPosition { program_index: usize, output_offset: usize, @@ -1034,3 +1320,143 @@ fn refresh_plan_capacity_error( span, } } + +#[cfg(test)] +mod tests { + use super::*; + + fn test_block(rows: Vec>) -> solve::ScalarProgramBlock { + solve::ScalarProgramBlock::with_source_span( + rows, + rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("refresh_plan_test.mo"), + 1, + 2, + ), + ) + } + + #[test] + fn block_matching_uses_real_incidence_when_target_hint_competes() { + let rows = [0, 1]; + let ys = [10, 11]; + let pairs = match_block_rows_to_targets( + &rows, + &ys, + &|_| None, + &|row| (row == 0).then_some(10), + |row, y| matches!((row, y), (0, 11) | (1, 10)), + None, + ) + .expect("matching should succeed"); + + assert_eq!(pairs.len(), 2); + assert!(pairs.contains(&(0, 11))); + assert!(pairs.contains(&(1, 10))); + } + + #[test] + fn dependency_refresh_plan_includes_alternative_row_dependencies() { + let model = solve::SolveModel { + problem: solve::SolveProblem { + solve_layout: solve::SolveLayout { + state_scalar_count: 1, + algebraic_scalar_count: 2, + ..Default::default() + }, + ..Default::default() + }, + ..Default::default() + }; + let block = Arc::new(test_block(vec![ + vec![ + solve::LinearOp::LoadY { dst: 0, index: 0 }, + solve::LinearOp::StoreOutput { src: 0 }, + ], + vec![ + solve::LinearOp::LoadY { dst: 0, index: 2 }, + solve::LinearOp::StoreOutput { src: 0 }, + ], + vec![ + solve::LinearOp::LoadY { dst: 0, index: 0 }, + solve::LinearOp::StoreOutput { src: 0 }, + ], + ])); + let full_plan = RefreshPlan { + source_block: block, + rows: vec![ + AlgebraicRefreshRow { + row_idx: 0, + output_offset: 0, + target_index: 1, + assignment_target: None, + alternatives: vec![AlgebraicRefreshCandidate { + row_idx: 1, + output_offset: 0, + assignment_target: None, + }], + }, + AlgebraicRefreshRow { + row_idx: 2, + output_offset: 0, + target_index: 2, + assignment_target: None, + alternatives: Vec::new(), + }, + ], + ..Default::default() + }; + let mut initial_deps = IndexSet::new(); + initial_deps.insert(1); + + let plan = build_dependency_refresh_plan(&model, &full_plan, initial_deps) + .expect("alternative dependency plan should build"); + + assert_eq!( + plan.rows + .iter() + .map(|row| row.target_index) + .collect::>(), + vec![2, 1], + "the alternative's producer must precede the target that may use it" + ); + } + + #[test] + fn refresh_plan_is_iterative_when_alternative_is_nonlinear() { + let block = Arc::new(test_block(vec![ + vec![ + solve::LinearOp::LoadY { dst: 0, index: 0 }, + solve::LinearOp::StoreOutput { src: 0 }, + ], + vec![ + solve::LinearOp::LoadY { dst: 0, index: 1 }, + solve::LinearOp::Binary { + dst: 1, + op: solve::BinaryOp::Mul, + lhs: 0, + rhs: 0, + }, + solve::LinearOp::StoreOutput { src: 1 }, + ], + ])); + let rows = vec![AlgebraicRefreshRow { + row_idx: 0, + output_offset: 0, + target_index: 1, + assignment_target: None, + alternatives: vec![AlgebraicRefreshCandidate { + row_idx: 1, + output_offset: 0, + assignment_target: None, + }], + }]; + + let plan = order_refresh_rows(rows, block, 1).expect("refresh plan should order"); + + assert!( + plan.iterative, + "a nonlinear alternative needs the same iterative convergence guard as the primary" + ); + } +} diff --git a/crates/rumoca-eval-solve/src/runtime.rs b/crates/rumoca-eval-solve/src/runtime.rs index 11fced3d9..77e61c0c2 100644 --- a/crates/rumoca-eval-solve/src/runtime.rs +++ b/crates/rumoca-eval-solve/src/runtime.rs @@ -2,8 +2,8 @@ use indexmap::IndexMap; use rumoca_ir_solve as solve; use rumoca_solver::{ EventActionOutcome, RuntimeEventStop, RuntimeSolveError, SolveStopSchedule, - push_visible_values, replace_last_visible_values, timeline::sample_time_match_with_tol, - update_relation_memory_slots, + push_visible_values, relation_memory_value_from_root, replace_last_visible_values, + timeline::sample_time_match_with_tol, update_relation_memory_slots, }; use std::{cell::RefCell, collections::HashMap}; @@ -27,6 +27,7 @@ mod plans; mod refresh_batch; mod sensitivity; mod support; +mod values; use event_update::{DiscretePreSnapshot, DiscreteRowsSettleInput}; pub use event_update::{EventUpdateRowFilter, ProjectedEventUpdateInput}; pub use initial_event::{ @@ -38,9 +39,9 @@ use plans::{ direct_time_root_value, direct_visible_value, root_condition_plan, visible_value_plan, }; use support::{ - NewtonProbe, apply_newton_steps, copy_runtime_values, copy_runtime_values_into, - reserve_runtime_index_map_capacity, reserve_runtime_vec_capacity, resize_runtime_values, - write_refresh_targets, zero_runtime_values, + NewtonProbe, copy_runtime_values, copy_runtime_values_into, reserve_runtime_index_map_capacity, + reserve_runtime_vec_capacity, resize_runtime_values, write_refresh_targets, + zero_runtime_values, }; struct RefreshSlotArgs<'a> { @@ -56,6 +57,25 @@ struct RefreshIterationMax { target: Option<(usize, usize, f64)>, } +pub struct RootSearchInput<'a> { + pub t: f64, + pub state: &'a [f64], + pub params: &'a [f64], + pub guess: &'a mut Vec, + pub tol: f64, + pub max_iters: usize, + pub out: &'a mut [f64], +} + +pub struct ProjectedRuntimeSettleInput<'a> { + pub y: &'a mut [f64], + pub p: &'a mut [f64], + pub t: f64, + pub tol: f64, + pub max_iters: usize, + pub root_relation_overrides: &'a [(usize, f64)], +} + impl From for RuntimeSolveError { fn from(value: solve_eval::EvalSolveError) -> Self { Self::solve_ir_with_span(value.to_string(), value.source_span()) @@ -217,7 +237,7 @@ impl SolveRuntime { copy_runtime_values_into(guess, &self.model.initial_y, "initial solver guess")?; resize_runtime_values(guess, self.solver_count, 0.0, "initial solver guess")?; } - self.populate_solver_y_from_state(guess, state)?; + self.overwrite_state_slots_preserving_algebraics(guess, state)?; self.refresh_algebraic_and_output_slots(t, guess, params, tol, max_iters) } @@ -317,42 +337,23 @@ impl SolveRuntime { tol, max_iters, } = args; - // Snapshot the targets: if Gauss-Seidel diverges (coupled loop with - // gain > 1, e.g. torque loops through a gear ratio), Newton restarts - // from these values rather than the diverged iterates. + // Coupled projection blocks may have multiple fixed points. Iterating + // their assignment map can converge cleanly to a remote, non-physical + // root even when the caller supplied an accepted local branch seed. + // Solve the coupled residual directly from that seed so branch + // continuity, rather than fixed-point attraction, selects the root. let snapshot = self.refresh_target_snapshot(rows, solver_y)?; - let mut last_max = RefreshIterationMax { - delta: 0.0, - target: None, - }; - // A coupled cycle with gain > 1 (any geared torque loop) makes the - // sweep delta grow monotonically; burning the full iteration budget - // before falling back is pure waste, so bail to Newton after a few - // consecutive growing sweeps. - const MAX_GROWING_SWEEPS: usize = 3; - let mut growing_sweeps = 0usize; - for iter_idx in 0..max_iters { - let previous_delta = last_max.delta; - match self.refresh_slots_iteration(rows, t, solver_y, params) { - Ok(iteration_max) => last_max = iteration_max, - Err(error) => { - // Divergence to non-finite values: retry with Newton from - // the snapshot before giving up. - tracing::debug!(target: "rumoca_eval_solve::refresh", "retrying refresh with Newton after sweep error: {error}"); - return self.refresh_slots_newton(rows, &snapshot, t, solver_y, params, tol); - } - } - self.trace_refresh_iteration(iter_idx, &last_max); - if last_max.delta <= tol { - return Ok(()); - } - let growing = iter_idx > 0 && last_max.delta > previous_delta; - growing_sweeps = if growing { growing_sweeps + 1 } else { 0 }; - if growing_sweeps >= MAX_GROWING_SWEEPS { - break; - } - } - self.refresh_slots_newton(rows, &snapshot, t, solver_y, params, tol) + self.refresh_slots_newton( + rows, + &snapshot, + RefreshSlotArgs { + t, + solver_y, + params, + tol, + max_iters, + }, + ) } fn refresh_target_snapshot( @@ -375,11 +376,15 @@ impl SolveRuntime { &self, rows: &[AlgebraicRefreshRow], snapshot: &[f64], - t: f64, - solver_y: &mut [f64], - params: &[f64], - tol: f64, + args: RefreshSlotArgs<'_>, ) -> Result<(), RuntimeSolveError> { + let RefreshSlotArgs { + t, + solver_y, + params, + tol, + max_iters: requested_max_iters, + } = args; const MAX_NEWTON_REFRESH_ROWS: usize = 256; const MAX_NEWTON_ITERS: usize = 25; let m = rows.len(); @@ -397,17 +402,16 @@ impl SolveRuntime { let mut x = Vec::new(); reserve_runtime_vec_capacity(&mut x, snapshot.len(), "Newton iterate")?; x.extend(snapshot); - for _ in 0..MAX_NEWTON_ITERS { + let max_newton_iters = requested_max_iters.max(MAX_NEWTON_ITERS); + let mut convergence_polish_steps = 0usize; + for _ in 0..max_newton_iters { let f_base = self.refresh_newton_sweep(rows, &x, t, solver_y, params)?; let mut residual = Vec::new(); reserve_runtime_vec_capacity(&mut residual, x.len(), "Newton residual")?; residual.extend(x.iter().zip(&f_base).map(|(xi, fi)| xi - fi)); let max_residual = residual.iter().fold(0.0_f64, |acc, r| acc.max(r.abs())); tracing::debug!(target: "rumoca_eval_solve::refresh", "newton residual={max_residual:e}"); - if max_residual <= tol { - write_refresh_targets(rows, &x, solver_y); - return Ok(()); - } + let converged = max_residual <= tol; let probe = NewtonProbe { rows, x: &x, @@ -417,14 +421,93 @@ impl SolveRuntime { params, }; let mut augmented = self.refresh_newton_augmented(probe, solver_y)?; - if crate::linear_solve::gaussian_eliminate(&mut augmented).is_none() { - return Err(self.refresh_newton_failure()); + let solved = crate::linear_solve::gaussian_eliminate(&mut augmented).is_some(); + if !solved && converged { + write_refresh_targets(rows, &x, solver_y); + return Ok(()); } - if !apply_newton_steps(&mut x, &augmented) { - return Err(self.refresh_newton_failure()); + if !solved { + return Err(self.refresh_newton_failure(rows)); } + if converged && !newton_step_stays_near_accepted_branch(&augmented, &x, tol) { + write_refresh_targets(rows, &x, solver_y); + return Ok(()); + } + let improved = self.apply_damped_refresh_newton_step( + rows, + &mut x, + &augmented, + max_residual, + t, + solver_y, + params, + )?; + // A tolerance-converged coupled solve still carries the residual + // directly into reconstructed connection-flow observations. One + // final Newton correction preserves the already-selected branch + // while polishing affine conservation identities (for example a + // grounded pin current reconstructed as the difference of two + // equal branch currents) to floating-point consistency. + if converged && (!improved || max_residual == 0.0 || convergence_polish_steps >= 3) { + write_refresh_targets(rows, &x, solver_y); + return Ok(()); + } + if converged { + convergence_polish_steps += 1; + continue; + } + if !improved { + return Err(self.refresh_newton_failure(rows)); + } + } + Err(self.refresh_newton_failure(rows)) + } + + #[allow(clippy::too_many_arguments)] + fn apply_damped_refresh_newton_step( + &self, + rows: &[AlgebraicRefreshRow], + x: &mut Vec, + augmented: &crate::linear_solve::AugmentedMatrix, + base_norm: f64, + t: f64, + solver_y: &mut [f64], + params: &[f64], + ) -> Result { + const MAX_BACKTRACKS: usize = 24; + const ARMIJO_SLOPE: f64 = 1.0e-4; + let m = x.len(); + let mut alpha = 1.0; + for _ in 0..MAX_BACKTRACKS { + let mut trial = x.clone(); + if (0..m).any(|j| !augmented.get(j, m).is_finite()) { + write_refresh_targets(rows, x, solver_y); + return Ok(false); + } + for (j, value) in trial.iter_mut().enumerate() { + let step = augmented.get(j, m); + *value += alpha * step; + } + let f_trial = match self.refresh_newton_sweep(rows, &trial, t, solver_y, params) { + Ok(values) => values, + Err(_) => { + alpha *= 0.5; + continue; + } + }; + let trial_norm = trial + .iter() + .zip(&f_trial) + .map(|(xi, fi)| (xi - fi).abs()) + .fold(0.0_f64, f64::max); + if trial_norm.is_finite() && trial_norm <= base_norm * (1.0 - ARMIJO_SLOPE * alpha) { + *x = trial; + return Ok(true); + } + alpha *= 0.5; } - Err(self.refresh_newton_failure()) + write_refresh_targets(rows, x, solver_y); + Ok(false) } /// Evaluate the refresh map `F` at `x` (writing `x` into the target slots @@ -485,7 +568,8 @@ impl SolveRuntime { Ok(augmented) } - fn refresh_newton_failure(&self) -> RuntimeSolveError { + fn refresh_newton_failure(&self, rows: &[AlgebraicRefreshRow]) -> RuntimeSolveError { + let _ = rows; tracing::debug!(target: "rumoca_eval_solve::refresh", "newton fallback FAILED"); self.refresh_convergence_error( 0, @@ -496,32 +580,6 @@ impl SolveRuntime { ) } - fn refresh_slots_iteration( - &self, - rows: &[AlgebraicRefreshRow], - t: f64, - solver_y: &mut [f64], - params: &[f64], - ) -> Result { - let mut max_delta: f64 = 0.0; - let mut max_target = None; - for refresh_row in rows { - let row_idx = refresh_row.row_idx; - let index = refresh_row.target_index; - let value = self.eval_refresh_row(refresh_row, t, solver_y, params)?; - let delta = (solver_y[index] - value).abs(); - if delta > max_delta { - max_delta = delta; - max_target = Some((index, row_idx, value)); - } - solver_y[index] = value; - } - Ok(RefreshIterationMax { - delta: max_delta, - target: max_target, - }) - } - fn eval_refresh_row( &self, row: &AlgebraicRefreshRow, @@ -602,7 +660,28 @@ impl SolveRuntime { return Ok(value); } let residual = self.refresh_row_residual(row, t, solver_y, params)?; - self.solve_refresh_residual_row(row, residual, t, solver_y, params) + if let Some(value) = self.solve_refresh_residual_row(row, residual, t, solver_y, params)? { + return Ok(value); + } + for candidate in &row.alternatives { + if candidate.row_idx == row.row_idx && candidate.output_offset == row.output_offset { + continue; + } + let candidate_row = AlgebraicRefreshRow { + row_idx: candidate.row_idx, + output_offset: candidate.output_offset, + target_index: row.target_index, + assignment_target: candidate.assignment_target, + alternatives: Vec::new(), + }; + let residual = self.refresh_row_residual(&candidate_row, t, solver_y, params)?; + if let Some(value) = + self.solve_refresh_residual_row(&candidate_row, residual, t, solver_y, params)? + { + return Ok(value); + } + } + Err(self.refresh_row_independent_error(row)) } /// Residual of a refresh row at the current point. Rows lowered with an @@ -663,7 +742,7 @@ impl SolveRuntime { t: f64, solver_y: &[f64], params: &[f64], - ) -> Result { + ) -> Result, RuntimeSolveError> { let index = row.target_index; let current = solver_y[index]; let mut probe_y = self.refresh_probe_scratch.borrow_mut(); @@ -674,39 +753,22 @@ impl SolveRuntime { let probe_residual = self.refresh_row_residual(row, t, &probe_y, params)?; let slope = probe_residual - residual; if slope.is_finite() && slope.abs() > 1.0e-12 { - return Ok(current - residual / slope); + return Ok(Some(current - residual / slope)); } + Ok(None) + } + + fn refresh_row_independent_error(&self, row: &AlgebraicRefreshRow) -> RuntimeSolveError { // A residual that does not respond to the paired variable means the // refresh plan paired this row with a variable it cannot determine. // Nudging the value by the residual (the old fallback) converges to a // wrong but stable solution; fail loudly instead. - Err(RuntimeSolveError::UnsupportedModel { + RuntimeSolveError::UnsupportedModel { reason: format!( "algebraic refresh row {} cannot be solved for '{}': the residual does not depend on it", row.row_idx, - self.solver_name(index) + self.solver_name(row.target_index) ), - }) - } - - fn trace_refresh_iteration(&self, iter_idx: usize, max: &RefreshIterationMax) { - // `tracing::debug!` self-gates; the only off-path work is a name lookup. - if let Some((index, row_idx, value)) = max.target { - let name = self - .model - .problem - .solve_layout - .solver_maps - .names - .get(index) - .map_or("", String::as_str); - tracing::debug!( - target: "rumoca_eval_solve::refresh", - "refresh iter {iter_idx}: max_delta={:.6e} target={name} y[{index}] row={row_idx} value={value:.6e}", - max.delta - ); - } else { - tracing::debug!(target: "rumoca_eval_solve::refresh", "refresh iter {iter_idx}: no targeted algebraics"); } } @@ -962,6 +1024,165 @@ impl SolveRuntime { self.write_planned_root_search_conditions(plan, &solver_y, params, t, out) } + /// Evaluate root-search values from an explicit algebraic branch seed. + /// + /// The caller owns `guess`: state slots are replaced by `state`, while the + /// remaining solver slots retain the supplied branch before the root + /// dependency projection is settled. This keeps adaptive-solver trial + /// callbacks free of hidden mutable runtime state. + pub fn eval_root_search_conditions_with_guess_into( + &self, + input: RootSearchInput<'_>, + ) -> Result<(), RuntimeSolveError> { + let RootSearchInput { + t, + state, + params, + guess, + tol, + max_iters, + out, + } = input; + let roots = &self.model.problem.events.root_conditions; + if roots.is_empty() { + if let Some(first) = out.first_mut() { + *first = 1.0; + } + return Ok(()); + } + if guess.len() != self.solver_count { + copy_runtime_values_into(guess, &self.model.initial_y, "root solver guess")?; + resize_runtime_values(guess, self.solver_count, 0.0, "root solver guess")?; + } + self.overwrite_state_slots_preserving_algebraics(guess, state)?; + let Some(plan) = &self.root_condition_plan else { + self.refresh_slots_with_plan( + &self.root_refresh, + RefreshSlotArgs { + t, + solver_y: guess, + params, + tol, + max_iters, + }, + )?; + return self.eval_root_conditions_from_refreshed_solver_y(t, guess, params, out); + }; + self.validate_root_plan_output_len(plan, out)?; + if plan.search_rows.is_empty() { + return self.write_planned_root_search_defaults(plan, params, t, out); + } + self.refresh_slots_with_plan( + &self.root_refresh, + RefreshSlotArgs { + t, + solver_y: guess, + params, + tol, + max_iters, + }, + )?; + self.write_planned_root_search_conditions(plan, guess, params, t, out) + } + + pub fn neutralize_initial_root_search_values( + &self, + params: &[f64], + tol: f64, + out: &mut [f64], + ) -> Result<(), RuntimeSolveError> { + let root_count = self.model.problem.events.root_conditions.output_count(); + let targets = &self.model.problem.events.root_relation_memory_targets; + if targets.len() != root_count { + return Err(RuntimeSolveError::solve_ir(format!( + "root relation metadata length {} does not match root output count {root_count}", + targets.len() + ))); + } + if out.len() < root_count { + return Err(RuntimeSolveError::solve_ir(format!( + "root search output has {} values for {root_count} roots", + out.len() + ))); + } + for (root_index, target) in targets.iter().copied().enumerate() { + if out[root_index] != 0.0 { + continue; + } + let Some(target) = target else { + continue; + }; + let solve::ScalarSlot::P { index, .. } = target else { + return Err(RuntimeSolveError::solve_ir(format!( + "root crossing index {root_index} has non-parameter relation memory target" + ))); + }; + let current = params.get(index).copied().ok_or_else(|| { + RuntimeSolveError::solve_ir(format!( + "root crossing index {root_index} relation memory parameter {index} is outside parameter storage" + )) + })?; + out[root_index] = if current.abs() <= tol { + 1.0 + } else if (current - 1.0).abs() <= tol { + -1.0 + } else { + return Err(RuntimeSolveError::solve_ir(format!( + "root crossing index {root_index} relation memory value {current} is not boolean" + ))); + }; + } + Ok(()) + } + + pub fn apply_consumed_root_search_overrides( + &self, + params: &[f64], + tol: f64, + overrides: &[(usize, f64)], + out: &mut [f64], + ) -> Result<(), RuntimeSolveError> { + let root_count = self.model.problem.events.root_conditions.output_count(); + let targets = &self.model.problem.events.root_relation_memory_targets; + if targets.len() != root_count || out.len() < root_count { + return Err(RuntimeSolveError::solve_ir(format!( + "root relation metadata/output shape does not match {root_count} roots" + ))); + } + for &(root_index, post) in overrides { + let target = targets.get(root_index).copied().flatten().ok_or_else(|| { + RuntimeSolveError::solve_ir(format!( + "consumed root index {root_index} has no relation memory target" + )) + })?; + let solve::ScalarSlot::P { index, .. } = target else { + return Err(RuntimeSolveError::solve_ir(format!( + "consumed root index {root_index} has non-parameter relation memory target" + ))); + }; + let current = params.get(index).copied().ok_or_else(|| { + RuntimeSolveError::solve_ir(format!( + "consumed root index {root_index} relation memory parameter {index} is outside parameter storage" + )) + })?; + if (current - post).abs() > tol { + return Err(RuntimeSolveError::solve_ir(format!( + "consumed root index {root_index} post-side {post} does not match relation memory {current}" + ))); + } + out[root_index] = if post.abs() <= tol { + 1.0 + } else if (post - 1.0).abs() <= tol { + -1.0 + } else { + return Err(RuntimeSolveError::solve_ir(format!( + "consumed root index {root_index} post-side {post} is not boolean" + ))); + }; + } + Ok(()) + } + pub fn next_planned_time_root( &self, params: &[f64], @@ -1285,12 +1506,49 @@ impl SolveRuntime { where P: FnMut(&mut [f64], &mut [f64]) -> Result, { + self.settle_projected_runtime_and_relation_memory_with_overrides( + ProjectedRuntimeSettleInput { + y, + p, + t, + tol, + max_iters, + root_relation_overrides: &[], + }, + &mut project_algebraics, + ) + } + + pub fn settle_projected_runtime_and_relation_memory_with_overrides

( + &self, + input: ProjectedRuntimeSettleInput<'_>, + mut project_algebraics: P, + ) -> Result<(), RuntimeSolveError> + where + P: FnMut(&mut [f64], &mut [f64]) -> Result, + { + let ProjectedRuntimeSettleInput { + y, + p, + t, + tol, + max_iters, + root_relation_overrides, + } = input; for _ in 0..max_iters { let mut changed = - self.apply_runtime_assignments_until_stable(y, p, t, tol, max_iters)?; + self.apply_root_relation_memory_overrides(root_relation_overrides, y, p, tol)?; + changed |= self.apply_runtime_assignments_until_stable(y, p, t, tol, max_iters)?; changed |= project_algebraics(y, p)?; changed |= self.apply_runtime_assignments_until_stable(y, p, t, tol, max_iters)?; - changed |= self.update_relation_memory_from_solver_y(t, y, p, tol)?; + changed |= self.update_relation_memory_from_solver_y_excluding( + t, + y, + p, + root_relation_overrides, + )?; + changed |= + self.apply_root_relation_memory_overrides(root_relation_overrides, y, p, tol)?; if !changed { return Ok(()); } @@ -1358,13 +1616,15 @@ impl SolveRuntime { root_relation_overrides, } = input; if self.model.problem.discrete.rhs.is_empty() { - self.apply_root_relation_memory_overrides(root_relation_overrides, y, p, tol)?; return self.settle_runtime_assignments_and_projection( - y, - p, - t, - tol, - max_iters, + ProjectedRuntimeSettleInput { + y, + p, + t, + tol, + max_iters, + root_relation_overrides, + }, &mut project_algebraics, ); } @@ -1375,6 +1635,14 @@ impl SolveRuntime { changed |= self.apply_runtime_assignments_until_stable(y, p, t, tol, max_iters)?; changed |= project_algebraics(y, p)?; changed |= self.apply_runtime_assignments_until_stable(y, p, t, tol, max_iters)?; + changed |= self.update_relation_memory_from_solver_y_excluding( + t, + y, + p, + root_relation_overrides, + )?; + changed |= + self.apply_root_relation_memory_overrides(root_relation_overrides, y, p, tol)?; let iter_pre_y = copy_runtime_values(y, "projected event iteration y snapshot")?; let iter_pre_p = copy_runtime_values(p, "projected event iteration p snapshot")?; let snapshot = DiscretePreSnapshot { @@ -1400,9 +1668,12 @@ impl SolveRuntime { &mut project_algebraics, )?; } - if root_relation_overrides.is_empty() { - changed |= self.update_relation_memory_from_solver_y(t, y, p, tol)?; - } + changed |= self.update_relation_memory_from_solver_y_excluding( + t, + y, + p, + root_relation_overrides, + )?; changed |= self.apply_root_relation_memory_overrides(root_relation_overrides, y, p, tol)?; changed |= project_algebraics(y, p)?; @@ -1429,21 +1700,33 @@ impl SolveRuntime { fn settle_runtime_assignments_and_projection

( &self, - y: &mut [f64], - p: &mut [f64], - t: f64, - tol: f64, - max_iters: usize, + input: ProjectedRuntimeSettleInput<'_>, project_algebraics: &mut P, ) -> Result where P: FnMut(&mut [f64], &mut [f64]) -> Result, { + let ProjectedRuntimeSettleInput { + y, + p, + t, + tol, + max_iters, + root_relation_overrides, + } = input; for _ in 0..max_iters { let mut changed = self.apply_runtime_assignments_until_stable(y, p, t, tol, max_iters)?; changed |= project_algebraics(y, p)?; changed |= self.apply_runtime_assignments_until_stable(y, p, t, tol, max_iters)?; + changed |= self.update_relation_memory_from_solver_y_excluding( + t, + y, + p, + root_relation_overrides, + )?; + changed |= + self.apply_root_relation_memory_overrides(root_relation_overrides, y, p, tol)?; if !changed { return self.eval_event_actions(y, p, t); } @@ -1525,6 +1808,35 @@ impl SolveRuntime { )) } + fn update_relation_memory_from_solver_y_excluding( + &self, + t: f64, + y: &[f64], + p: &mut [f64], + excluded_roots: &[(usize, f64)], + ) -> Result { + let relation_memory_indices = &self + .model + .problem + .solve_layout + .relation_memory_parameter_indices; + let roots = self.eval_root_conditions_from_solver_y(t, y, p)?; + let mut changed = false; + for (root_index, (root, parameter_index)) in + roots.iter().zip(relation_memory_indices).enumerate() + { + if excluded_roots.iter().any(|(index, _)| *index == root_index) { + continue; + } + if let Some(slot) = p.get_mut(*parameter_index) { + let value = relation_memory_value_from_root(*root); + changed |= *slot != value; + *slot = value; + } + } + Ok(changed) + } + pub fn eval_root_conditions_from_solver_y( &self, t: f64, @@ -1554,244 +1866,6 @@ impl SolveRuntime { self.row_eval_context(), ) } - - pub fn record_visible_sample( - &self, - data: &mut [Vec], - solver_y: &[f64], - params: &[f64], - t: f64, - ) -> Result<(), RuntimeSolveError> { - let mut values = self.visible_scratch.borrow_mut(); - self.visible_values_into(solver_y, params, t, &mut values)?; - push_visible_values(data, &values) - } - - pub fn record_visible_sample_if_new( - &self, - recorded_times: &mut Vec, - data: &mut [Vec], - solver_y: &[f64], - params: &[f64], - t: f64, - ) -> Result<(), RuntimeSolveError> { - let mut values = self.visible_scratch.borrow_mut(); - self.visible_values_into(solver_y, params, t, &mut values)?; - if recorded_times - .last() - .is_some_and(|last| sample_time_match_with_tol(*last, t)) - { - if let Some(last) = recorded_times.last_mut() { - *last = t; - } - replace_last_visible_values(data, &values)?; - return Ok(()); - } - reserve_runtime_vec_capacity(recorded_times, 1, "recorded sample times")?; - recorded_times.push(t); - push_visible_values(data, &values) - } - - pub fn visible_values( - &self, - y: &[f64], - params: &[f64], - t: f64, - ) -> Result, RuntimeSolveError> { - let mut values = Vec::new(); - self.visible_values_into(y, params, t, &mut values)?; - Ok(values) - } - - fn visible_values_into( - &self, - y: &[f64], - params: &[f64], - t: f64, - values: &mut Vec, - ) -> Result<(), RuntimeSolveError> { - if let Some(plan) = &self.visible_value_plan { - resize_runtime_values(values, plan.entries.len(), 0.0, "visible values")?; - self.write_planned_visible_values(plan, y, params, t, values)?; - return Ok(()); - } - if self.visible_value_rows.len() == self.model.visible_names.len() { - resize_runtime_values(values, self.visible_value_rows.len(), 0.0, "visible values")?; - self.visible_value_rows.eval_with_context( - y, - params, - t, - self.row_eval_context(), - values, - )?; - return Ok(()); - } - let computed = - visible_values_with_context(&self.model, y, params, t, self.row_eval_context())?; - copy_runtime_values_into(values, &computed, "visible values") - } - - fn write_planned_visible_values( - &self, - plan: &VisibleValuePlan, - y: &[f64], - params: &[f64], - t: f64, - values: &mut [f64], - ) -> Result<(), RuntimeSolveError> { - for (slot, entry) in values.iter_mut().zip(plan.entries.iter().copied()) { - if let VisibleValuePlanEntry::Direct(source) = entry { - *slot = direct_visible_value(source, y, params, t)?; - } - } - if !plan.expression_rows.is_empty() { - self.visible_value_rows - .eval_single_output_rows_unchecked_with_context( - &plan.expression_rows, - y, - params, - t, - self.row_eval_context(), - values, - )?; - copy_grouped_expression_values(plan, values)?; - } - Ok(()) - } - - pub fn visible_values_for_names( - &self, - y: &[f64], - params: &[f64], - t: f64, - names: &[String], - ) -> Result, RuntimeSolveError> { - if self.visible_value_rows.len() == self.model.visible_names.len() { - return self.visible_values_for_names_from_rows(y, params, t, names); - } - let all_values = self.visible_values(y, params, t)?; - let mut values = IndexMap::new(); - reserve_runtime_index_map_capacity(&mut values, names.len(), "visible name values")?; - for name in names { - let Some(idx) = self.visible_name_index.get(name).copied() else { - continue; - }; - let value = all_values.get(idx).copied().ok_or_else(|| { - visible_value_index_error(name, idx, all_values.len(), "visible values") - })?; - values.insert(name.clone(), value); - } - Ok(values) - } - - fn visible_values_for_names_from_rows( - &self, - y: &[f64], - params: &[f64], - t: f64, - names: &[String], - ) -> Result, RuntimeSolveError> { - let mut values = IndexMap::new(); - reserve_runtime_index_map_capacity(&mut values, names.len(), "visible row name values")?; - for name in names { - if let Some(value) = self.visible_value_from_row(name, y, params, t)? { - values.insert(name.clone(), value); - } - } - Ok(values) - } - - fn visible_value_from_row( - &self, - name: &str, - y: &[f64], - params: &[f64], - t: f64, - ) -> Result, RuntimeSolveError> { - let Some(idx) = self.visible_name_index.get(name).copied() else { - return Ok(None); - }; - if idx >= self.visible_value_rows.len() { - return Err(visible_value_index_error( - name, - idx, - self.visible_value_rows.len(), - "visible value rows", - )); - } - let value = self.visible_value_rows.eval_row_with_context( - idx, - y, - params, - t, - self.row_eval_context(), - )?; - Ok(Some(value)) - } - - fn populate_solver_y_from_state( - &self, - solver_y: &mut Vec, - state: &[f64], - ) -> Result<(), RuntimeSolveError> { - copy_runtime_values_into(solver_y, &self.model.initial_y, "solver y initial values")?; - resize_runtime_values(solver_y, self.solver_count, 0.0, "solver y")?; - for (dst, src) in solver_y.iter_mut().zip(state.iter().copied()) { - *dst = src; - } - Ok(()) - } - - // SPEC_0021: Exception - private derivative helper shares the public solver - // callback shape while threading caller-owned scratch/output buffers. - #[allow(clippy::too_many_arguments)] - fn eval_state_derivatives_with_solver_y( - &self, - t: f64, - state: &[f64], - params: &[f64], - tol: f64, - max_iters: usize, - solver_y: &mut Vec, - out: &mut [f64], - ) -> Result<(), RuntimeSolveError> { - self.populate_solver_y_from_state(solver_y, state)?; - self.refresh_derivative_dependencies(t, solver_y, params, tol, max_iters)?; - // `eval_derivative_rhs_from_solver_y` fills `out` and *then* rejects - // non-finite derivatives, so trace before propagating: on failure `out` - // and `solver_y` still hold the offending values to name for the user. - let eval_result = self.eval_derivative_rhs_from_solver_y(t, solver_y, params, out); - crate::nan_trace::report_state_derivative(&self.model, t, solver_y, out); - eval_result - } - - fn eval_derivative_rhs_from_solver_y( - &self, - t: f64, - solver_y: &[f64], - params: &[f64], - out: &mut [f64], - ) -> Result<(), RuntimeSolveError> { - validate_derivative_output_len(out, self.state_count)?; - self.derivative_rhs - .eval_with_context(solver_y, params, t, self.row_eval_context(), out)?; - self.validate_finite_derivatives(out) - } - - fn validate_finite_derivatives(&self, derivative: &[f64]) -> Result<(), RuntimeSolveError> { - for (idx, value) in derivative.iter().enumerate() { - if !value.is_finite() { - let state_name = self - .model - .visible_names - .get(idx) - .cloned() - .unwrap_or_else(|| format!("state[{idx}]")); - return Err(RuntimeSolveError::NonFiniteDerivative { state_name }); - } - } - Ok(()) - } } #[derive(Clone, Default)] @@ -1857,6 +1931,20 @@ fn visible_value_index_error( )) } +fn newton_step_stays_near_accepted_branch( + augmented: &crate::linear_solve::AugmentedMatrix, + accepted: &[f64], + tol: f64, +) -> bool { + const POLISH_TRUST_FACTOR: f64 = 8.0; + let rhs_column = accepted.len(); + accepted.iter().enumerate().all(|(row, value)| { + let step = augmented.get(row, rhs_column).abs(); + let limit = POLISH_TRUST_FACTOR * tol * value.abs().max(1.0); + step.is_finite() && step <= limit + }) +} + pub fn apply_discrete_slot_value( target: solve::ScalarSlot, value: f64, diff --git a/crates/rumoca-eval-solve/src/runtime/sensitivity.rs b/crates/rumoca-eval-solve/src/runtime/sensitivity.rs index 1f0b96bf6..2e55d02f5 100644 --- a/crates/rumoca-eval-solve/src/runtime/sensitivity.rs +++ b/crates/rumoca-eval-solve/src/runtime/sensitivity.rs @@ -37,6 +37,27 @@ impl SolveRuntime { self.eval_derivative_jacobian_v_with_seed(lin, state, seed, self.state_count, out) } + /// State Jacobian-vector product at an explicitly seeded algebraic branch. + /// The supplied guess is settled in place, so callers can use a private + /// trial copy without mutating the accepted branch shared by an integrator. + pub fn eval_state_jacobian_v_ad_with_guess_into( + &self, + lin: AlgebraicLinearization<'_>, + state: &[f64], + seed: &[f64], + guess: &mut Vec, + out: &mut [f64], + ) -> Result<(), RuntimeSolveError> { + self.eval_derivative_jacobian_v_with_seed_and_guess( + lin, + state, + seed, + self.state_count, + Some(guess), + out, + ) + } + /// Like [`Self::eval_state_jacobian_v_ad_into`], but the input `seed` spans /// the full `[solver-y | parameter]` space and is copied in its entirety, so /// parameter tangents are honored. Seeding a unit vector in a parameter slot @@ -425,6 +446,25 @@ impl SolveRuntime { seed: &[f64], seed_copy_len: usize, out: &mut [f64], + ) -> Result<(), RuntimeSolveError> { + self.eval_derivative_jacobian_v_with_seed_and_guess( + lin, + state, + seed, + seed_copy_len, + None, + out, + ) + } + + fn eval_derivative_jacobian_v_with_seed_and_guess( + &self, + lin: AlgebraicLinearization<'_>, + state: &[f64], + seed: &[f64], + seed_copy_len: usize, + mut guess: Option<&mut Vec>, + out: &mut [f64], ) -> Result<(), RuntimeSolveError> { let AlgebraicLinearization { t, params, settle } = lin; validate_derivative_output_len(out, self.state_count)?; @@ -435,8 +475,23 @@ impl SolveRuntime { unit_seed, } = &mut *scratch; // (1) Linearization point: project the algebraics from the state. - self.populate_solver_y_from_state(solver_y, state)?; + if let Some(guess) = &mut guess { + if guess.len() != self.solver_count { + copy_runtime_values_into(guess, &self.model.initial_y, "Jacobian solver guess")?; + resize_runtime_values(guess, self.solver_count, 0.0, "Jacobian solver guess")?; + } + solver_y.clear(); + solver_y.extend_from_slice(guess); + } + if guess.is_some() { + self.overwrite_state_slots_preserving_algebraics(solver_y, state)?; + } else { + self.populate_solver_y_from_state(solver_y, state)?; + } self.refresh_derivative_dependencies(t, solver_y, params, settle.tol, settle.max_iters)?; + if let Some(guess) = guess { + guess.copy_from_slice(solver_y); + } // The JVP rows seed both solver-y and parameters (`SeedMode::SolverYAndP`), // so the seed vector spans `[solver-y | parameter]` space. We copy the // leading `seed_copy_len` entries from the caller (state-only for the @@ -447,7 +502,8 @@ impl SolveRuntime { .requirements() .seed_len .max(self.implicit_jacobian_v.requirements().seed_len) - .max(self.solver_count); + .max(self.solver_count) + .max(seed.len()); seed_buf.clear(); seed_buf.resize(seed_len, 0.0); let n = seed_copy_len.min(seed.len()).min(seed_buf.len()); diff --git a/crates/rumoca-eval-solve/src/runtime/support.rs b/crates/rumoca-eval-solve/src/runtime/support.rs index e08e04478..3cceff71a 100644 --- a/crates/rumoca-eval-solve/src/runtime/support.rs +++ b/crates/rumoca-eval-solve/src/runtime/support.rs @@ -21,22 +21,6 @@ pub(super) struct NewtonProbe<'a> { pub(super) params: &'a [f64], } -/// Apply the eliminated Newton steps; false when any step is non-finite. -pub(super) fn apply_newton_steps( - x: &mut [f64], - augmented: &crate::linear_solve::AugmentedMatrix, -) -> bool { - let m = x.len(); - for (j, value) in x.iter_mut().enumerate() { - let step = augmented.get(j, m); - if !step.is_finite() { - return false; - } - *value += step; - } - true -} - pub(super) fn zero_runtime_values( len: usize, context: &'static str, diff --git a/crates/rumoca-eval-solve/src/runtime/tests.rs b/crates/rumoca-eval-solve/src/runtime/tests.rs index 98ce2f3ff..8ae50e920 100644 --- a/crates/rumoca-eval-solve/src/runtime/tests.rs +++ b/crates/rumoca-eval-solve/src/runtime/tests.rs @@ -1,5 +1,8 @@ use super::*; +mod branch_projection; +mod root_condition_plans; + fn valid_algebraic_refresh_plan( model: &solve::SolveModel, block: &PreparedScalarProgramBlock, @@ -115,6 +118,184 @@ fn refresh_plan_does_not_let_residual_target_shadow_assignment_row() { assert_eq!(plan.rows[0].target_index, 1); } +#[test] +fn projection_refresh_rows_keep_direct_assignments_as_alternatives() { + let model = solve::SolveModel { + problem: solve::SolveProblem { + solve_layout: solve::SolveLayout { + solver_maps: solve::SolverNameIndexMaps { + names: vec!["x".to_string(), "y".to_string()], + ..Default::default() + }, + algebraic_scalar_count: 2, + ..Default::default() + }, + continuous: solve::ContinuousSolveSystem { + implicit_rhs: solve::ComputeBlock::from_scalar_program_block(spanned_block( + vec![ + assignment_residual_row(), + non_assignment_targeted_residual_row(), + ], + "projection_shadow.mo", + )), + implicit_row_targets: vec![ + Some(solve::scalar_slot_y(1)), + Some(solve::scalar_slot_y(1)), + ], + algebraic_projection_plan: solve::AlgebraicProjectionPlan { + blocks: vec![solve::AlgebraicProjectionBlock { + rows: vec![1], + y_indices: vec![1], + causal_steps: vec![solve::AlgebraicProjectionStep { row: 1, y_index: 1 }], + }], + }, + ..Default::default() + }, + ..Default::default() + }, + ..Default::default() + }; + let block = + PreparedScalarProgramBlock::from_compute_block(&model.problem.continuous.implicit_rhs) + .expect("valid implicit RHS should prepare"); + + let plan = valid_algebraic_refresh_plan(&model, &block); + + assert!(plan.iterative); + assert_eq!(plan.rows.len(), 1); + assert_eq!(plan.rows[0].row_idx, 1); + assert_eq!(plan.rows[0].target_index, 1); + assert_eq!(plan.rows[0].alternatives.len(), 1); + assert_eq!(plan.rows[0].alternatives[0].row_idx, 0); + assert_eq!(plan.rows[0].alternatives[0].output_offset, 0); +} + +#[test] +fn refresh_plan_preserves_blt_causal_primary_over_explicit_target() { + let model = solve::SolveModel { + problem: solve::SolveProblem { + solve_layout: solve::SolveLayout { + solver_maps: solve::SolverNameIndexMaps { + names: vec!["a".to_string(), "b".to_string()], + ..Default::default() + }, + algebraic_scalar_count: 2, + ..Default::default() + }, + continuous: solve::ContinuousSolveSystem { + implicit_rhs: solve::ComputeBlock::from_scalar_program_block(spanned_block( + vec![positive_sum_residual_row(), positive_sum_residual_row()], + "blt_causal_primary.mo", + )), + implicit_row_targets: vec![ + Some(solve::scalar_slot_y(0)), + Some(solve::scalar_slot_y(1)), + ], + algebraic_projection_plan: solve::AlgebraicProjectionPlan { + blocks: vec![solve::AlgebraicProjectionBlock { + rows: vec![0, 1], + y_indices: vec![0, 1], + causal_steps: vec![ + solve::AlgebraicProjectionStep { row: 0, y_index: 1 }, + solve::AlgebraicProjectionStep { row: 1, y_index: 0 }, + ], + }], + }, + ..Default::default() + }, + ..Default::default() + }, + ..Default::default() + }; + let block = + PreparedScalarProgramBlock::from_compute_block(&model.problem.continuous.implicit_rhs) + .expect("valid implicit RHS should prepare"); + + let plan = valid_algebraic_refresh_plan(&model, &block); + let a = plan + .rows + .iter() + .find(|row| row.target_index == 0) + .expect("a should have a refresh producer"); + let b = plan + .rows + .iter() + .find(|row| row.target_index == 1) + .expect("b should have a refresh producer"); + + assert_eq!(a.row_idx, 1, "BLT row 1 must remain a's primary"); + assert_eq!(b.row_idx, 0, "BLT row 0 must remain b's primary"); + assert_eq!( + a.alternatives + .iter() + .map(|candidate| (candidate.row_idx, candidate.output_offset)) + .collect::>(), + vec![(0, 0)], + "a's direct row should remain only as a deterministic alternative" + ); + assert_eq!( + b.alternatives + .iter() + .map(|candidate| (candidate.row_idx, candidate.output_offset)) + .collect::>(), + vec![(1, 0)], + "b's direct row should remain only as a deterministic alternative" + ); +} + +#[test] +fn refresh_plan_rejects_causal_steps_outside_projection_block() { + let mut model = solve::SolveModel { + problem: solve::SolveProblem { + solve_layout: solve::SolveLayout { + algebraic_scalar_count: 2, + ..Default::default() + }, + continuous: solve::ContinuousSolveSystem { + implicit_rhs: solve::ComputeBlock::from_scalar_program_block(spanned_block( + vec![positive_sum_residual_row(), positive_sum_residual_row()], + "causal_step_membership.mo", + )), + algebraic_projection_plan: solve::AlgebraicProjectionPlan { + blocks: vec![solve::AlgebraicProjectionBlock { + rows: vec![0], + y_indices: vec![0], + causal_steps: Vec::new(), + }], + }, + ..Default::default() + }, + ..Default::default() + }, + ..Default::default() + }; + let block = + PreparedScalarProgramBlock::from_compute_block(&model.problem.continuous.implicit_rhs) + .expect("valid implicit RHS should prepare"); + let invalid_steps = [ + ( + solve::AlgebraicProjectionStep { row: 1, y_index: 0 }, + "row 1 is not a member", + ), + ( + solve::AlgebraicProjectionStep { row: 0, y_index: 1 }, + "target y[1] is not a member", + ), + ]; + + for (step, expected) in invalid_steps { + model.problem.continuous.algebraic_projection_plan.blocks[0].causal_steps = vec![step]; + let error = match build_algebraic_refresh_plan(&model, &block) { + Ok(_) => panic!("causal ownership outside its projection block must be rejected"), + Err(error) => error, + }; + assert!( + error.to_string().contains(expected), + "error should explain invalid causal membership: {error}" + ); + } +} + #[test] fn refresh_plan_accepts_scaled_affine_residual_target() { let model = solve::SolveModel { @@ -459,6 +640,52 @@ fn refresh_residual_fallback_solves_positive_unit_coefficient() { assert_eq!(solver_y[1], 4.0); } +#[test] +fn refresh_uses_projection_alternative_when_primary_is_runtime_singular() { + let model = solve::SolveModel { + problem: solve::SolveProblem { + solve_layout: solve::SolveLayout { + solver_maps: solve::SolverNameIndexMaps { + names: vec!["target".to_string(), "other".to_string()], + ..Default::default() + }, + algebraic_scalar_count: 2, + ..Default::default() + }, + continuous: solve::ContinuousSolveSystem { + implicit_rhs: solve::ComputeBlock::from_scalar_program_block(spanned_block( + vec![ + parameter_select_primary_row(), + target_minus_constant_residual_row(0, 2.0), + ], + "parameter_select_projection.mo", + )), + algebraic_projection_plan: solve::AlgebraicProjectionPlan { + blocks: vec![solve::AlgebraicProjectionBlock { + rows: vec![0, 1], + y_indices: vec![0], + causal_steps: vec![solve::AlgebraicProjectionStep { row: 0, y_index: 0 }], + }], + }, + ..Default::default() + }, + ..Default::default() + }, + initial_y: vec![10.0, 5.0], + parameters: vec![0.0], + ..Default::default() + }; + let runtime = SolveRuntime::new(&model).expect("valid runtime should prepare"); + let mut solver_y = model.initial_y.clone(); + + runtime + .refresh_algebraic_and_output_slots(0.0, &mut solver_y, &model.parameters, 1.0e-12, 1) + .expect("active projection alternative should solve the target"); + + assert_eq!(solver_y[0], 2.0); + assert_eq!(solver_y[1], 5.0); +} + #[test] fn derivative_refresh_errors_on_missing_algebraic_producer() { let model = solve::SolveModel { @@ -500,6 +727,64 @@ fn derivative_refresh_errors_on_missing_algebraic_producer() { ); } +#[test] +fn derivative_refresh_uses_complete_incidence_matching_despite_competing_hints() { + let model = solve::SolveModel { + problem: solve::SolveProblem { + solve_layout: solve::SolveLayout { + solver_maps: solve::SolverNameIndexMaps { + names: vec!["x".to_string(), "a".to_string(), "b".to_string()], + ..Default::default() + }, + state_scalar_count: 1, + algebraic_scalar_count: 2, + ..Default::default() + }, + continuous: solve::ContinuousSolveSystem { + implicit_rhs: solve::ComputeBlock::from_scalar_program_block(spanned_block( + vec![ + derivative_placeholder_row(0), + derivative_placeholder_row(2), + derivative_placeholder_row(1), + ], + "complete_producer_implicit.mo", + )), + implicit_row_targets: vec![ + Some(solve::scalar_slot_y(0)), + Some(solve::scalar_slot_y(1)), + None, + ], + derivative_rhs: solve::ComputeBlock::from_scalar_program_block(spanned_block( + vec![derivative_placeholder_row(2)], + "complete_producer_derivative.mo", + )), + algebraic_projection_plan: solve::AlgebraicProjectionPlan { + blocks: vec![solve::AlgebraicProjectionBlock { + rows: vec![1, 2], + y_indices: vec![1, 2], + causal_steps: vec![solve::AlgebraicProjectionStep { row: 1, y_index: 1 }], + }], + }, + ..Default::default() + }, + ..Default::default() + }, + initial_y: vec![0.0; 3], + ..Default::default() + }; + + let runtime = SolveRuntime::new(&model).expect("complete producer graph should prepare"); + + assert!(runtime.derivative_refresh.missing_dependencies.is_empty()); + assert!( + runtime + .algebraic_refresh + .rows + .iter() + .any(|row| row.row_idx == 1 && row.target_index == 2) + ); +} + #[test] fn visible_values_for_names_preserves_requested_order() { let model = solve::SolveModel { @@ -603,86 +888,6 @@ fn visible_value_plan_deduplicates_equal_expression_rows() { assert_eq!(values, vec![20.0, 30.0, 30.0]); } -#[test] -fn root_condition_plan_keeps_full_values_but_neutralizes_search_roots() { - let model = solve::SolveModel { - problem: solve::SolveProblem { - events: solve::SolveEventPartition { - root_conditions: spanned_block( - vec![ - constant_expression_root_row(), - param_minus_time_root_row(0), - direct_param_visible_value_row(1), - time_plus_one_root_row(), - ], - "root_plan.mo", - ), - ..Default::default() - }, - ..Default::default() - }, - parameters: vec![2.5, 9.0], - ..Default::default() - }; - let runtime = SolveRuntime::new(&model).expect("valid runtime should prepare"); - let plan = runtime - .root_condition_plan - .as_ref() - .expect("root condition plan should build"); - - assert_eq!(plan.evaluated_rows, vec![2, 3]); - assert_eq!(plan.search_rows, vec![3]); - - let full = runtime - .eval_root_conditions_from_solver_y(1.0, &[], &model.parameters) - .expect("full root values should evaluate"); - assert_eq!(full, vec![5.0, 1.5, 9.0, 2.0]); - - let mut search = vec![0.0; 4]; - runtime - .eval_root_search_conditions_into(1.0, &[], &model.parameters, 1.0e-12, 1, &mut search) - .expect("search root values should evaluate"); - assert_eq!(search, vec![1.0, 1.0, 1.0, 2.0]); -} - -#[test] -fn root_condition_plan_reports_next_direct_time_root() { - let model = solve::SolveModel { - problem: solve::SolveProblem { - events: solve::SolveEventPartition { - root_conditions: spanned_block( - vec![param_minus_time_root_row(0)], - "direct_time_root.mo", - ), - ..Default::default() - }, - ..Default::default() - }, - parameters: vec![2.5], - ..Default::default() - }; - let runtime = SolveRuntime::new(&model).expect("valid runtime should prepare"); - - assert_eq!( - runtime - .next_planned_time_root(&model.parameters, 1.0, 3.0, 1.0e-12) - .expect("direct time root should be found"), - Some(2.5) - ); - assert_eq!( - runtime - .next_planned_time_root(&model.parameters, 2.5, 3.0, 1.0e-12) - .expect("current root should not be rescheduled"), - None - ); - assert_eq!( - runtime - .next_planned_time_root(&model.parameters, 1.0, 2.0, 1.0e-12) - .expect("future root beyond target should be ignored"), - None - ); -} - #[test] fn visible_value_runtime_errors_keep_row_span() { let span = rumoca_core::Span::from_offsets( @@ -889,24 +1094,10 @@ fn direct_time_visible_value_row() -> Vec { ] } -fn param_minus_time_root_row(index: usize) -> Vec { - vec![ - solve::LinearOp::LoadP { dst: 0, index }, - solve::LinearOp::LoadTime { dst: 1 }, - solve::LinearOp::Binary { - dst: 2, - op: solve::BinaryOp::Sub, - lhs: 0, - rhs: 1, - }, - solve::LinearOp::StoreOutput { src: 2 }, - ] -} - -fn constant_expression_root_row() -> Vec { +fn positive_sum_residual_row() -> Vec { vec![ - solve::LinearOp::Const { dst: 0, value: 2.0 }, - solve::LinearOp::Const { dst: 1, value: 3.0 }, + solve::LinearOp::LoadY { dst: 0, index: 0 }, + solve::LinearOp::LoadY { dst: 1, index: 1 }, solve::LinearOp::Binary { dst: 2, op: solve::BinaryOp::Add, @@ -917,27 +1108,38 @@ fn constant_expression_root_row() -> Vec { ] } -fn time_plus_one_root_row() -> Vec { +fn parameter_select_primary_row() -> Vec { vec![ - solve::LinearOp::LoadTime { dst: 0 }, - solve::LinearOp::Const { dst: 1, value: 1.0 }, - solve::LinearOp::Binary { + solve::LinearOp::LoadP { dst: 0, index: 0 }, + solve::LinearOp::Const { dst: 1, value: 0.0 }, + solve::LinearOp::Compare { dst: 2, - op: solve::BinaryOp::Add, + op: solve::CompareOp::Le, lhs: 0, rhs: 1, }, - solve::LinearOp::StoreOutput { src: 2 }, + solve::LinearOp::LoadY { dst: 3, index: 1 }, + solve::LinearOp::LoadY { dst: 4, index: 0 }, + solve::LinearOp::Select { + dst: 5, + cond: 2, + if_true: 3, + if_false: 4, + }, + solve::LinearOp::StoreOutput { src: 5 }, ] } -fn positive_sum_residual_row() -> Vec { +fn target_minus_constant_residual_row(target: usize, value: f64) -> Vec { vec![ - solve::LinearOp::LoadY { dst: 0, index: 0 }, - solve::LinearOp::LoadY { dst: 1, index: 1 }, + solve::LinearOp::LoadY { + dst: 0, + index: target, + }, + solve::LinearOp::Const { dst: 1, value }, solve::LinearOp::Binary { dst: 2, - op: solve::BinaryOp::Add, + op: solve::BinaryOp::Sub, lhs: 0, rhs: 1, }, diff --git a/crates/rumoca-eval-solve/src/runtime/tests/branch_projection.rs b/crates/rumoca-eval-solve/src/runtime/tests/branch_projection.rs new file mode 100644 index 000000000..abf5cbef7 --- /dev/null +++ b/crates/rumoca-eval-solve/src/runtime/tests/branch_projection.rs @@ -0,0 +1,356 @@ +use super::*; + +#[test] +fn explicit_branch_seeds_refresh_only_their_dependency_plan() { + use solve::LinearOp::{Const, LoadY, StoreOutput}; + + let mut model = projection_coupled_state_model(2.0); + // The derivative and root depend only on the state. The full observation + // projection still owns algebraic `a = 2*x`. + model.problem.continuous.derivative_rhs = + solve::ComputeBlock::from_scalar_program_block(spanned_block( + vec![vec![Const { dst: 0, value: 0.0 }, StoreOutput { src: 0 }]], + "branch_seed_derivative.mo", + )); + model.problem.events.root_conditions = spanned_block( + vec![vec![LoadY { dst: 0, index: 0 }, StoreOutput { src: 0 }]], + "branch_seed_root.mo", + ); + model.problem.events.root_relation_memory_targets = vec![None]; + let runtime = SolveRuntime::new(&model).expect("valid branch-seed runtime should prepare"); + + let mut derivative_seed = vec![99.0, 42.0]; + let mut derivative = [f64::NAN]; + runtime + .eval_state_derivatives_with_guess_into( + 0.0, + &[3.0], + &[], + &mut derivative_seed, + 1.0e-12, + 32, + &mut derivative, + ) + .expect("derivative branch should settle"); + + let mut root_seed = vec![99.0, 42.0]; + let mut root = [f64::NAN]; + runtime + .eval_root_search_conditions_with_guess_into(RootSearchInput { + t: 0.0, + state: &[3.0], + params: &[], + guess: &mut root_seed, + tol: 1.0e-12, + max_iters: 32, + out: &mut root, + }) + .expect("root branch should settle"); + + let mut observation_seed = vec![99.0, 42.0]; + runtime + .full_solver_y_with_guess(0.0, &[3.0], &[], &mut observation_seed, 1.0e-12, 32) + .expect("full observation projection should settle"); + + assert_eq!(derivative_seed, vec![3.0, 42.0]); + assert_eq!(root_seed, vec![3.0, 42.0]); + assert_eq!(observation_seed, vec![3.0, 6.0]); + assert_eq!(derivative, [0.0]); + assert_eq!(root, [3.0]); +} + +#[test] +fn coupled_projection_preserves_the_accepted_local_branch() { + use solve::LinearOp::{Binary, Const, LoadY, StoreOutput}; + use solve::{BinaryOp, ComputeBlock}; + + let mut model = linear_algebraic_loop_state_model(); + // a = b; b = F(a), with nearby root 1 and remote root 100. + model.problem.continuous.implicit_rhs = ComputeBlock::from_scalar_program_block(spanned_block( + vec![ + vec![LoadY { dst: 0, index: 0 }, StoreOutput { src: 0 }], + vec![ + LoadY { dst: 0, index: 1 }, + LoadY { dst: 1, index: 2 }, + Binary { + dst: 2, + op: BinaryOp::Sub, + lhs: 0, + rhs: 1, + }, + StoreOutput { src: 2 }, + ], + vec![ + LoadY { dst: 0, index: 1 }, + Const { dst: 1, value: 1.0 }, + Binary { + dst: 2, + op: BinaryOp::Sub, + lhs: 0, + rhs: 1, + }, + Const { + dst: 3, + value: 100.0, + }, + Binary { + dst: 4, + op: BinaryOp::Sub, + lhs: 0, + rhs: 3, + }, + Binary { + dst: 5, + op: BinaryOp::Mul, + lhs: 2, + rhs: 4, + }, + Const { + dst: 6, + value: 1.0 / 99.0, + }, + Binary { + dst: 7, + op: BinaryOp::Mul, + lhs: 5, + rhs: 6, + }, + Binary { + dst: 8, + op: BinaryOp::Sub, + lhs: 0, + rhs: 7, + }, + LoadY { dst: 9, index: 2 }, + Binary { + dst: 10, + op: BinaryOp::Sub, + lhs: 9, + rhs: 8, + }, + StoreOutput { src: 10 }, + ], + ], + "accepted_branch_projection.mo", + )); + model.initial_y = vec![0.0, 1.1, 1.1]; + let runtime = SolveRuntime::new(&model).expect("valid multi-root projection runtime"); + assert!(runtime.algebraic_refresh.iterative); + + let tol = 1.0e-8; + let mut accepted_seed = vec![0.0, 1.1, 1.1]; + let first = runtime + .eval_state_derivatives_with_guess(0.0, &[0.0], &[], &mut accepted_seed, tol, 256) + .expect("accepted branch projection should settle"); + let map = accepted_seed[1] - (accepted_seed[1] - 1.0) * (accepted_seed[1] - 100.0) / 99.0; + let residual = (accepted_seed[1] - accepted_seed[2]) + .abs() + .max((accepted_seed[2] - map).abs()); + assert!( + (accepted_seed[1] - 1.0).abs() <= tol, + "projection jumped from the accepted local branch to {}", + accepted_seed[1] + ); + assert!((accepted_seed[2] - 1.0).abs() <= tol); + assert!( + residual <= tol, + "projection residual {residual} exceeds {tol}" + ); + + let mut projected_seed = accepted_seed.clone(); + let second = runtime + .eval_state_derivatives_with_guess(0.0, &[0.0], &[], &mut projected_seed, tol, 256) + .expect("reprojecting the accepted branch should settle"); + assert!((first[0] - second[0]).abs() <= tol); +} + +#[test] +fn tolerance_converged_projection_does_not_polish_to_remote_root() { + use solve::LinearOp::{Binary, Const, LoadY, StoreOutput}; + use solve::{BinaryOp, ComputeBlock}; + + let mut model = linear_algebraic_loop_state_model(); + let epsilon = 1.0e-6; + model.problem.continuous.implicit_rhs = ComputeBlock::from_scalar_program_block(spanned_block( + vec![ + vec![LoadY { dst: 0, index: 0 }, StoreOutput { src: 0 }], + vec![ + LoadY { dst: 0, index: 1 }, + LoadY { dst: 1, index: 2 }, + Binary { + dst: 2, + op: BinaryOp::Sub, + lhs: 0, + rhs: 1, + }, + StoreOutput { src: 2 }, + ], + vec![ + LoadY { dst: 0, index: 1 }, + Const { + dst: 1, + value: 100.0, + }, + Binary { + dst: 2, + op: BinaryOp::Sub, + lhs: 0, + rhs: 1, + }, + Const { + dst: 3, + value: epsilon, + }, + Binary { + dst: 4, + op: BinaryOp::Mul, + lhs: 2, + rhs: 3, + }, + Binary { + dst: 5, + op: BinaryOp::Sub, + lhs: 0, + rhs: 4, + }, + LoadY { dst: 6, index: 2 }, + Binary { + dst: 7, + op: BinaryOp::Sub, + lhs: 6, + rhs: 5, + }, + StoreOutput { src: 7 }, + ], + ], + "ill_scaled_branch_projection.mo", + )); + model.initial_y = vec![0.0, 0.0, 0.0]; + let runtime = SolveRuntime::new(&model).expect("ill-scaled projection should prepare"); + assert!(runtime.algebraic_refresh.iterative); + + let mut accepted_seed = model.initial_y.clone(); + runtime + .full_solver_y_with_guess(0.0, &[0.0], &[], &mut accepted_seed, 1.0e-4, 32) + .expect("already-tolerant accepted branch should settle"); + + assert!( + accepted_seed[1].abs() <= 1.0e-6 && accepted_seed[2].abs() <= 1.0e-6, + "tolerance polish jumped to the remote root: {accepted_seed:?}" + ); +} + +#[test] +fn confirmed_root_override_wins_after_other_relation_memories_update() { + use solve::LinearOp::{LoadY, StoreOutput}; + + let mut model = solve::SolveModel::default(); + model.problem.solve_layout.algebraic_scalar_count = 2; + model.problem.solve_layout.relation_memory_parameter_indices = vec![0, 1]; + model.initial_y = vec![-1.0, -1.0]; + model.problem.events.root_conditions = spanned_block( + vec![ + vec![LoadY { dst: 0, index: 0 }, StoreOutput { src: 0 }], + vec![LoadY { dst: 0, index: 1 }, StoreOutput { src: 0 }], + ], + "coincident_relation_roots.mo", + ); + model.problem.events.root_relation_memory_targets = vec![ + Some(solve::ScalarSlot::P { + index: 0, + byte_offset: 0, + }), + Some(solve::ScalarSlot::P { + index: 1, + byte_offset: 0, + }), + ]; + let runtime = SolveRuntime::new(&model).expect("coincident relation roots should prepare"); + let mut y = model.initial_y.clone(); + let mut p = vec![0.0, 0.0]; + + runtime + .settle_projected_runtime_and_relation_memory_with_overrides( + ProjectedRuntimeSettleInput { + y: &mut y, + p: &mut p, + t: 0.0, + tol: 1.0e-12, + max_iters: 8, + root_relation_overrides: &[(0, 0.0)], + }, + |_, _| Ok(false), + ) + .expect("relation memory update should converge atomically"); + + assert_eq!(p, vec![0.0, 1.0]); +} + +#[test] +fn converged_coupled_projection_polishes_reconstructed_zero_flow() { + use solve::LinearOp::{Binary, Const, LoadY, StoreOutput}; + use solve::{BinaryOp, ComputeBlock}; + + // Artifact-derived analogue of the severe `ground.p.i` channels in + // IdealTriacCircuit and HBridge_TrianglePWM_RL. Structural elimination + // reconstructs the grounded flow as the difference of two branch currents; + // those currents belong to a coupled projection block. A merely + // tolerance-converged solve leaks its residual into that observation even + // though the physical connection flow is exactly zero. + let mut model = linear_algebraic_loop_state_model(); + let coupled_row = |target: usize, other: usize| { + vec![ + LoadY { + dst: 0, + index: target, + }, + Const { dst: 1, value: 2.0 }, + Binary { + dst: 2, + op: BinaryOp::Mul, + lhs: 0, + rhs: 1, + }, + LoadY { + dst: 3, + index: other, + }, + Binary { + dst: 4, + op: BinaryOp::Sub, + lhs: 2, + rhs: 3, + }, + Const { dst: 5, value: 1.0 }, + Binary { + dst: 6, + op: BinaryOp::Sub, + lhs: 4, + rhs: 5, + }, + StoreOutput { src: 6 }, + ] + }; + model.problem.continuous.implicit_rhs = ComputeBlock::from_scalar_program_block(spanned_block( + vec![ + vec![LoadY { dst: 0, index: 0 }, StoreOutput { src: 0 }], + coupled_row(1, 2), + coupled_row(2, 1), + ], + "ground_flow_projection.mo", + )); + model.initial_y = vec![0.0, 1.0 + 5.0e-11, 1.0 - 5.0e-11]; + let runtime = SolveRuntime::new(&model).expect("coupled flow fixture should prepare"); + assert!(runtime.algebraic_refresh.iterative); + + let mut observation = model.initial_y.clone(); + runtime + .full_solver_y_with_guess(0.0, &[0.0], &[], &mut observation, 1.0e-10, 32) + .expect("coupled observation projection should settle"); + + assert!( + (observation[1] - observation[2]).abs() <= 4.0 * f64::EPSILON, + "reconstructed grounded flow retained projection residual: {:?}", + observation + ); +} diff --git a/crates/rumoca-eval-solve/src/runtime/tests/root_condition_plans.rs b/crates/rumoca-eval-solve/src/runtime/tests/root_condition_plans.rs new file mode 100644 index 000000000..44407faa4 --- /dev/null +++ b/crates/rumoca-eval-solve/src/runtime/tests/root_condition_plans.rs @@ -0,0 +1,269 @@ +use super::*; + +#[test] +fn root_condition_plan_keeps_full_values_but_neutralizes_search_roots() { + let model = solve::SolveModel { + problem: solve::SolveProblem { + events: solve::SolveEventPartition { + root_conditions: spanned_block( + vec![ + constant_expression_root_row(), + param_minus_time_root_row(0), + direct_param_visible_value_row(1), + time_plus_one_root_row(), + ], + "root_plan.mo", + ), + root_relation_memory_targets: vec![None; 4], + ..Default::default() + }, + ..Default::default() + }, + parameters: vec![2.5, 9.0], + ..Default::default() + }; + let runtime = SolveRuntime::new(&model).expect("valid runtime should prepare"); + let plan = runtime + .root_condition_plan + .as_ref() + .expect("root condition plan should build"); + + assert_eq!(plan.evaluated_rows, vec![2, 3]); + assert_eq!(plan.search_rows, vec![3]); + + let full = runtime + .eval_root_conditions_from_solver_y(1.0, &[], &model.parameters) + .expect("full root values should evaluate"); + assert_eq!(full, vec![5.0, 1.5, 9.0, 2.0]); + + let mut search = vec![0.0; 4]; + runtime + .eval_root_search_conditions_into(1.0, &[], &model.parameters, 1.0e-12, 1, &mut search) + .expect("search root values should evaluate"); + assert_eq!(search, vec![1.0, 1.0, 1.0, 2.0]); +} + +#[test] +fn root_condition_plan_reports_next_direct_time_root() { + let model = solve::SolveModel { + problem: solve::SolveProblem { + events: solve::SolveEventPartition { + root_conditions: spanned_block( + vec![param_minus_time_root_row(0)], + "direct_time_root.mo", + ), + ..Default::default() + }, + ..Default::default() + }, + parameters: vec![2.5], + ..Default::default() + }; + let runtime = SolveRuntime::new(&model).expect("valid runtime should prepare"); + + assert_eq!( + runtime + .next_planned_time_root(&model.parameters, 1.0, 3.0, 1.0e-12) + .expect("direct time root should be found"), + Some(2.5) + ); + assert_eq!( + runtime + .next_planned_time_root(&model.parameters, 2.5, 3.0, 1.0e-12) + .expect("current root should not be rescheduled"), + None + ); + assert_eq!( + runtime + .next_planned_time_root(&model.parameters, 1.0, 2.0, 1.0e-12) + .expect("future root beyond target should be ignored"), + None + ); +} + +#[test] +fn root_search_uses_relation_memory_side_only_for_exact_zero_dynamic_roots() { + let mut model = solve::SolveModel { + problem: solve::SolveProblem { + events: solve::SolveEventPartition { + root_conditions: spanned_block( + vec![direct_time_visible_value_row()], + "root_plan.mo", + ), + root_relation_memory_targets: vec![Some(solve::scalar_slot_p(0))], + ..Default::default() + }, + ..Default::default() + }, + parameters: vec![0.0], + ..Default::default() + }; + + let runtime = SolveRuntime::new(&model).expect("valid runtime should prepare"); + let full = runtime + .eval_root_conditions_from_solver_y(0.0, &[], &model.parameters) + .expect("full root evaluation must preserve exact zero"); + assert_eq!(full, vec![0.0]); + + let mut search = vec![0.0]; + runtime + .eval_root_search_conditions_into(-0.0, &[], &model.parameters, 1.0e-12, 1, &mut search) + .expect("initial raw root should evaluate"); + runtime + .neutralize_initial_root_search_values(&model.parameters, 1.0e-12, &mut search) + .expect("pre-crossing initial side should neutralize exact negative zero"); + assert_eq!(search, vec![1.0]); + + model.parameters[0] = 1.0; + let runtime = SolveRuntime::new(&model).expect("valid runtime should prepare"); + runtime + .eval_root_search_conditions_into(0.0, &[], &model.parameters, 1.0e-12, 1, &mut search) + .expect("reset raw root should evaluate"); + runtime + .neutralize_initial_root_search_values(&model.parameters, 1.0e-12, &mut search) + .expect("post-crossing reset side should neutralize exact zero"); + assert_eq!(search, vec![-1.0]); + + runtime + .eval_root_search_conditions_into(0.0, &[], &model.parameters, 1.0e-12, 1, &mut search) + .expect("locator exact zero should remain physical after the start callback"); + assert_eq!(search, vec![0.0]); + + runtime + .eval_root_search_conditions_into(-0.01, &[], &model.parameters, 1.0e-12, 1, &mut search) + .expect("nonzero dynamic roots must retain their physical value"); + assert_eq!(search, vec![-0.01]); +} + +#[test] +fn root_search_exact_zero_then_same_side_value_does_not_create_false_crossing() { + let model = solve::SolveModel { + problem: solve::SolveProblem { + events: solve::SolveEventPartition { + root_conditions: spanned_block(vec![one_minus_time_root_row()], "root_plan.mo"), + root_relation_memory_targets: vec![Some(solve::scalar_slot_p(0))], + ..Default::default() + }, + ..Default::default() + }, + parameters: vec![1.0], + ..Default::default() + }; + let runtime = SolveRuntime::new(&model).expect("valid runtime should prepare"); + let mut at_reset = vec![0.0]; + let mut after_reset = vec![0.0]; + + runtime + .eval_root_search_conditions_into(1.0, &[], &model.parameters, 1.0e-12, 1, &mut at_reset) + .expect("reset raw root should evaluate"); + runtime + .neutralize_initial_root_search_values(&model.parameters, 1.0e-12, &mut at_reset) + .expect("reset root should use relation side"); + runtime + .eval_root_search_conditions_into(1.1, &[], &model.parameters, 1.0e-12, 1, &mut after_reset) + .expect("same-side root should evaluate"); + + assert!(at_reset[0] < 0.0); + assert!(after_reset[0] < 0.0); +} + +#[test] +fn consumed_root_override_disarms_nonzero_locator_residual_without_touching_neighbor() { + let model = solve::SolveModel { + problem: solve::SolveProblem { + events: solve::SolveEventPartition { + root_conditions: spanned_block( + vec![ + direct_time_visible_value_row(), + direct_time_visible_value_row(), + ], + "root_plan.mo", + ), + root_relation_memory_targets: vec![Some(solve::scalar_slot_p(0)), None], + ..Default::default() + }, + ..Default::default() + }, + parameters: vec![1.0], + ..Default::default() + }; + let runtime = SolveRuntime::new(&model).expect("runtime should prepare"); + let mut at_root_start = vec![2.7e-14, 7.0]; + + runtime + .apply_consumed_root_search_overrides( + &model.parameters, + 1.0e-12, + &[(0, 1.0)], + &mut at_root_start, + ) + .expect("confirmed root should use its post relation side"); + + assert_eq!(at_root_start, vec![-1.0, 7.0]); + + let later_locator_values = vec![2.7e-14, 7.0]; + assert_eq!(later_locator_values, vec![2.7e-14, 7.0]); +} + +fn param_minus_time_root_row(index: usize) -> Vec { + vec![ + solve::LinearOp::LoadP { dst: 0, index }, + solve::LinearOp::LoadTime { dst: 1 }, + solve::LinearOp::Binary { + dst: 2, + op: solve::BinaryOp::Sub, + lhs: 0, + rhs: 1, + }, + solve::LinearOp::StoreOutput { src: 2 }, + ] +} + +fn direct_time_visible_value_row() -> Vec { + vec![ + solve::LinearOp::LoadTime { dst: 0 }, + solve::LinearOp::StoreOutput { src: 0 }, + ] +} + +fn one_minus_time_root_row() -> Vec { + vec![ + solve::LinearOp::Const { dst: 0, value: 1.0 }, + solve::LinearOp::LoadTime { dst: 1 }, + solve::LinearOp::Binary { + dst: 2, + op: solve::BinaryOp::Sub, + lhs: 0, + rhs: 1, + }, + solve::LinearOp::StoreOutput { src: 2 }, + ] +} + +fn constant_expression_root_row() -> Vec { + vec![ + solve::LinearOp::Const { dst: 0, value: 2.0 }, + solve::LinearOp::Const { dst: 1, value: 3.0 }, + solve::LinearOp::Binary { + dst: 2, + op: solve::BinaryOp::Add, + lhs: 0, + rhs: 1, + }, + solve::LinearOp::StoreOutput { src: 2 }, + ] +} + +fn time_plus_one_root_row() -> Vec { + vec![ + solve::LinearOp::LoadTime { dst: 0 }, + solve::LinearOp::Const { dst: 1, value: 1.0 }, + solve::LinearOp::Binary { + dst: 2, + op: solve::BinaryOp::Add, + lhs: 0, + rhs: 1, + }, + solve::LinearOp::StoreOutput { src: 2 }, + ] +} diff --git a/crates/rumoca-eval-solve/src/runtime/values.rs b/crates/rumoca-eval-solve/src/runtime/values.rs new file mode 100644 index 000000000..cb7bd20ca --- /dev/null +++ b/crates/rumoca-eval-solve/src/runtime/values.rs @@ -0,0 +1,266 @@ +use super::*; + +impl SolveRuntime { + pub fn record_visible_sample( + &self, + data: &mut [Vec], + solver_y: &[f64], + params: &[f64], + t: f64, + ) -> Result<(), RuntimeSolveError> { + let mut values = self.visible_scratch.borrow_mut(); + self.visible_values_into(solver_y, params, t, &mut values)?; + push_visible_values(data, &values) + } + + pub fn record_visible_sample_if_new( + &self, + recorded_times: &mut Vec, + data: &mut [Vec], + solver_y: &[f64], + params: &[f64], + t: f64, + ) -> Result<(), RuntimeSolveError> { + let mut values = self.visible_scratch.borrow_mut(); + self.visible_values_into(solver_y, params, t, &mut values)?; + if recorded_times + .last() + .is_some_and(|last| sample_time_match_with_tol(*last, t)) + { + if let Some(last) = recorded_times.last_mut() { + *last = t; + } + replace_last_visible_values(data, &values)?; + return Ok(()); + } + reserve_runtime_vec_capacity(recorded_times, 1, "recorded sample times")?; + recorded_times.push(t); + push_visible_values(data, &values) + } + + pub fn visible_values( + &self, + y: &[f64], + params: &[f64], + t: f64, + ) -> Result, RuntimeSolveError> { + let mut values = Vec::new(); + self.visible_values_into(y, params, t, &mut values)?; + Ok(values) + } + + fn visible_values_into( + &self, + y: &[f64], + params: &[f64], + t: f64, + values: &mut Vec, + ) -> Result<(), RuntimeSolveError> { + if let Some(plan) = &self.visible_value_plan { + resize_runtime_values(values, plan.entries.len(), 0.0, "visible values")?; + self.write_planned_visible_values(plan, y, params, t, values)?; + return Ok(()); + } + if self.visible_value_rows.len() == self.model.visible_names.len() { + resize_runtime_values(values, self.visible_value_rows.len(), 0.0, "visible values")?; + self.visible_value_rows.eval_with_context( + y, + params, + t, + self.row_eval_context(), + values, + )?; + return Ok(()); + } + let computed = + visible_values_with_context(&self.model, y, params, t, self.row_eval_context())?; + copy_runtime_values_into(values, &computed, "visible values") + } + + fn write_planned_visible_values( + &self, + plan: &VisibleValuePlan, + y: &[f64], + params: &[f64], + t: f64, + values: &mut [f64], + ) -> Result<(), RuntimeSolveError> { + for (slot, entry) in values.iter_mut().zip(plan.entries.iter().copied()) { + if let VisibleValuePlanEntry::Direct(source) = entry { + *slot = direct_visible_value(source, y, params, t)?; + } + } + if !plan.expression_rows.is_empty() { + self.visible_value_rows + .eval_single_output_rows_unchecked_with_context( + &plan.expression_rows, + y, + params, + t, + self.row_eval_context(), + values, + )?; + copy_grouped_expression_values(plan, values)?; + } + Ok(()) + } + + pub fn visible_values_for_names( + &self, + y: &[f64], + params: &[f64], + t: f64, + names: &[String], + ) -> Result, RuntimeSolveError> { + if self.visible_value_rows.len() == self.model.visible_names.len() { + return self.visible_values_for_names_from_rows(y, params, t, names); + } + let all_values = self.visible_values(y, params, t)?; + let mut values = IndexMap::new(); + reserve_runtime_index_map_capacity(&mut values, names.len(), "visible name values")?; + for name in names { + let Some(idx) = self.visible_name_index.get(name).copied() else { + continue; + }; + let value = all_values.get(idx).copied().ok_or_else(|| { + visible_value_index_error(name, idx, all_values.len(), "visible values") + })?; + values.insert(name.clone(), value); + } + Ok(values) + } + + fn visible_values_for_names_from_rows( + &self, + y: &[f64], + params: &[f64], + t: f64, + names: &[String], + ) -> Result, RuntimeSolveError> { + let mut values = IndexMap::new(); + reserve_runtime_index_map_capacity(&mut values, names.len(), "visible row name values")?; + for name in names { + if let Some(value) = self.visible_value_from_row(name, y, params, t)? { + values.insert(name.clone(), value); + } + } + Ok(values) + } + + fn visible_value_from_row( + &self, + name: &str, + y: &[f64], + params: &[f64], + t: f64, + ) -> Result, RuntimeSolveError> { + let Some(idx) = self.visible_name_index.get(name).copied() else { + return Ok(None); + }; + if idx >= self.visible_value_rows.len() { + return Err(visible_value_index_error( + name, + idx, + self.visible_value_rows.len(), + "visible value rows", + )); + } + let value = self.visible_value_rows.eval_row_with_context( + idx, + y, + params, + t, + self.row_eval_context(), + )?; + Ok(Some(value)) + } + + pub(super) fn populate_solver_y_from_state( + &self, + solver_y: &mut Vec, + state: &[f64], + ) -> Result<(), RuntimeSolveError> { + copy_runtime_values_into(solver_y, &self.model.initial_y, "solver y initial values")?; + resize_runtime_values(solver_y, self.solver_count, 0.0, "solver y")?; + self.overwrite_state_slots_preserving_algebraics(solver_y, state) + } + + pub(super) fn overwrite_state_slots_preserving_algebraics( + &self, + solver_y: &mut [f64], + state: &[f64], + ) -> Result<(), RuntimeSolveError> { + if solver_y.len() != self.solver_count { + return Err(RuntimeSolveError::solve_ir(format!( + "solver y has {} values, expected {}", + solver_y.len(), + self.solver_count + ))); + } + if state.len() < self.state_count { + return Err(RuntimeSolveError::solve_ir(format!( + "state has {} values, expected at least {}", + state.len(), + self.state_count + ))); + } + for (dst, src) in solver_y[..self.state_count] + .iter_mut() + .zip(state.iter().copied()) + { + *dst = src; + } + Ok(()) + } + + // SPEC_0021: Exception - private derivative helper shares the public solver + // callback shape while threading caller-owned scratch/output buffers. + #[allow(clippy::too_many_arguments)] + pub(super) fn eval_state_derivatives_with_solver_y( + &self, + t: f64, + state: &[f64], + params: &[f64], + tol: f64, + max_iters: usize, + solver_y: &mut Vec, + out: &mut [f64], + ) -> Result<(), RuntimeSolveError> { + self.populate_solver_y_from_state(solver_y, state)?; + self.refresh_derivative_dependencies(t, solver_y, params, tol, max_iters)?; + // `eval_derivative_rhs_from_solver_y` fills `out` and *then* rejects + // non-finite derivatives, so trace before propagating: on failure `out` + // and `solver_y` still hold the offending values to name for the user. + let eval_result = self.eval_derivative_rhs_from_solver_y(t, solver_y, params, out); + crate::nan_trace::report_state_derivative(&self.model, t, solver_y, out); + eval_result + } + + pub(super) fn eval_derivative_rhs_from_solver_y( + &self, + t: f64, + solver_y: &[f64], + params: &[f64], + out: &mut [f64], + ) -> Result<(), RuntimeSolveError> { + validate_derivative_output_len(out, self.state_count)?; + self.derivative_rhs + .eval_with_context(solver_y, params, t, self.row_eval_context(), out)?; + self.validate_finite_derivatives(out) + } + + fn validate_finite_derivatives(&self, derivative: &[f64]) -> Result<(), RuntimeSolveError> { + for (idx, value) in derivative.iter().enumerate() { + if !value.is_finite() { + let state_name = self + .model + .visible_names + .get(idx) + .cloned() + .unwrap_or_else(|| format!("state[{idx}]")); + return Err(RuntimeSolveError::NonFiniteDerivative { state_name }); + } + } + Ok(()) + } +} diff --git a/crates/rumoca-eval-solve/src/sim_driver.rs b/crates/rumoca-eval-solve/src/sim_driver.rs index b8d9e4cbe..54f4e7411 100644 --- a/crates/rumoca-eval-solve/src/sim_driver.rs +++ b/crates/rumoca-eval-solve/src/sim_driver.rs @@ -11,6 +11,7 @@ //! backend only provides a `SolverAdvanceBackend` adapter. use std::cell::RefCell; +use std::collections::BTreeSet; use std::rc::Rc; use rumoca_ir_solve as solve; @@ -23,8 +24,8 @@ use rumoca_solver::{ }; use crate::{ - EventUpdateRowFilter, ProjectedEventUpdateInput, SimulationRuntimeState, SolveRuntime, - next_runtime_event_stop, + EventUpdateRowFilter, ProjectedEventUpdateInput, ProjectedRuntimeSettleInput, + SimulationRuntimeState, SolveRuntime, next_runtime_event_stop, }; const EVENT_UPDATE_MAX_ITERS: usize = 256; @@ -40,7 +41,16 @@ pub enum StepOutcome { /// Took an internal adaptive step (did not reach a stop/root). Internal, /// A zero-crossing root was located at `t_root`. - Root { t_root: f64 }, + Root { + t_root: f64, + root_indices: Vec, + }, +} + +#[derive(Clone, Copy, Debug)] +pub enum RootStartBoundary<'a> { + Scheduled, + Root(&'a [(usize, f64)]), } /// Error surfaced by the driver. Backend (`SolverAdvanceBackend`) failures arrive as @@ -98,6 +108,7 @@ pub trait SolverAdvanceBackend { params: &[f64], t: f64, h_cap: f64, + boundary: RootStartBoundary<'_>, ) -> Result<(), SimDriverError>; /// Whether output points should be reached by stepping the solver exactly /// onto them (true) rather than by dense-output interpolation (false), for @@ -179,6 +190,12 @@ enum PendingRootAction { Continue, } +#[derive(Clone)] +struct PendingRoot { + t: f64, + indices: Vec, +} + /// Buffers for one full [`simulate_state_targets`] run. pub struct StateTrajectory<'a> { pub params: &'a mut Vec, @@ -217,14 +234,13 @@ pub fn simulate_state_targets( let runtime = state.runtime; let runtime_state = state.runtime_state; let mut stop_schedule = SolveStopSchedule::new(&model.problem, opts.t_start, opts.t_end); - let mut pending_root_t: Option = None; + let mut pending_root: Option = None; let make_ctx = || AdvanceContext { model, opts, runtime, runtime_params, }; - for &target in times { if state .recorded_times @@ -236,7 +252,7 @@ pub fn simulate_state_targets( let tol = opts.atol.max(1.0e-12); while target > *state.current_t + tol { match resolve_pending_root( - &mut pending_root_t, + &mut pending_root, make_ctx(), AdvanceState { current_y: state.current_y, @@ -261,7 +277,13 @@ pub fn simulate_state_targets( *state.current_t, target, )?; - let mut deferred_root: Option = None; + // Time-event equations may select their post-event branch at the + // exact boundary. Stop continuous integration at the previous + // representable instant so an implicit solver never evaluates the + // discontinuous right-side RHS before the event update is applied. + let solver_stop_time = + event_stop.map_or(stop_time, |_| event_left_limit_time(stop_time)); + let mut deferred_root: Option = None; let hit_root = advance_to_target_once( make_ctx(), AdvanceState { @@ -269,17 +291,21 @@ pub fn simulate_state_targets( params: state.params, current_t: state.current_t, }, - stop_time, + solver_stop_time, event_stop, backend, &mut deferred_root, )?; - if let Some(prt) = deferred_root { - pending_root_t = Some(prt); + if let Some(root) = deferred_root { + pending_root = Some(root); } + let event_stop_reached = event_stop.is_some() + && sample_time_match_with_tol(*state.current_t, solver_stop_time); if let Some(event) = event_stop - && sample_time_match_with_tol(*state.current_t, stop_time) + && event_stop_reached + && !hit_root { + *state.current_t = stop_time; apply_scheduled_time_event( make_ctx(), AdvanceState { @@ -295,6 +321,11 @@ pub fn simulate_state_targets( data: state.data, }, )?; + // A deferred root belongs to the continuous trajectory that + // the scheduled-event reset just replaced. + pending_root = None; + } + if event_stop_reached { stop_schedule.advance_past(*state.current_t); } if !hit_root && event_stop.is_none() { @@ -316,16 +347,17 @@ pub fn simulate_state_targets( } fn resolve_pending_root( - pending_root_t: &mut Option, + pending_root: &mut Option, ctx: AdvanceContext<'_>, state: AdvanceState<'_>, target: f64, backend: &mut St, stop_schedule: &mut SolveStopSchedule, ) -> Result { - let Some(prt) = *pending_root_t else { + let Some(root) = pending_root.clone() else { return Ok(PendingRootAction::None); }; + let prt = root.t; if !sample_time_match_with_tol(target, prt) && target < prt { let y_at = backend.interpolate(target)?; *state.current_t = target; @@ -337,8 +369,8 @@ fn resolve_pending_root( return Ok(PendingRootAction::Break); } - *pending_root_t = None; - handle_root_crossing(ctx, state, prt, target, backend)?; + *pending_root = None; + handle_root_crossing(ctx, state, prt, &root.indices, target, backend)?; stop_schedule.advance_past(backend.time()); Ok(PendingRootAction::Continue) } @@ -362,17 +394,30 @@ struct EventPre<'a> { p: &'a [f64], } +struct EventUpdateKernelInput<'a> { + y: &'a mut [f64], + p: &'a mut [f64], + t: f64, + tol: f64, + root_relation_overrides: &'a [(usize, f64)], + pre: EventPre<'a>, +} + /// Apply the projected discrete-event update at `t`, using the backend's /// `project_algebraics` as the projection callback (shared by every backend). fn apply_event_update_kernel( runtime: &SolveRuntime, backend: &St, - y: &mut [f64], - p: &mut [f64], - t: f64, - tol: f64, - pre: EventPre<'_>, + input: EventUpdateKernelInput<'_>, ) -> Result<(), SimDriverError> { + let EventUpdateKernelInput { + y, + p, + t, + tol, + root_relation_overrides, + pre, + } = input; let outcome = runtime.apply_projected_event_update( ProjectedEventUpdateInput { y, @@ -383,7 +428,7 @@ fn apply_event_update_kernel( event_pre_p: pre.p, max_iters: EVENT_UPDATE_MAX_ITERS, row_filter: EventUpdateRowFilter::All, - root_relation_overrides: &[], + root_relation_overrides, }, |y, p| backend.project_algebraics(y, p, t, tol), )?; @@ -398,13 +443,17 @@ fn settle_kernel( p: &mut [f64], t: f64, tol: f64, + root_relation_overrides: &[(usize, f64)], ) -> Result<(), SimDriverError> { - runtime.settle_projected_runtime_and_relation_memory( - y, - p, - t, - tol, - EVENT_UPDATE_MAX_ITERS, + runtime.settle_projected_runtime_and_relation_memory_with_overrides( + ProjectedRuntimeSettleInput { + y, + p, + t, + tol, + max_iters: EVENT_UPDATE_MAX_ITERS, + root_relation_overrides, + }, |y, p| backend.project_algebraics(y, p, t, tol), )?; Ok(()) @@ -422,7 +471,7 @@ fn bracket_event_limits_kernel( right_t: f64, ) -> Result<(), SimDriverError> { let dt = right_t - root_t; - if dt <= 0.0 || sample_time_match_with_tol(root_t, right_t) { + if dt <= 0.0 { return Ok(()); } let dy = backend.derivative_guess(y, p, root_t)?; @@ -495,13 +544,16 @@ impl RuntimeEventBoundaryHandler for EventBou apply_event_update_kernel( self.runtime, self.backend, - self.y, - self.p, - event_t, - self.tol, - EventPre { - y: &self.event_pre_y, - p: &self.event_pre_p, + EventUpdateKernelInput { + y: self.y, + p: self.p, + t: event_t, + tol: self.tol, + root_relation_overrides: &[], + pre: EventPre { + y: &self.event_pre_y, + p: &self.event_pre_p, + }, }, )?; self.record(event_t)?; @@ -527,13 +579,16 @@ impl RuntimeEventBoundaryHandler for EventBou apply_event_update_kernel( self.runtime, self.backend, - self.y, - self.p, - right_t, - self.tol, - EventPre { - y: &self.event_pre_y, - p: &self.event_pre_p, + EventUpdateKernelInput { + y: self.y, + p: self.p, + t: right_t, + tol: self.tol, + root_relation_overrides: &[], + pre: EventPre { + y: &self.event_pre_y, + p: &self.event_pre_p, + }, }, )?; if self.root_t.is_some() { @@ -544,6 +599,7 @@ impl RuntimeEventBoundaryHandler for EventBou self.p, right_t, self.tol, + &[], )?; } else { self.record(right_t)?; @@ -565,6 +621,7 @@ fn refresh_interpolated_sample_state( state.params, target, ctx.opts.atol.max(1.0e-10), + &[], )?; backend.refresh_observation(state.current_y, state.params, target)?; ctx.runtime_params @@ -625,7 +682,11 @@ fn reinitialize_solver_after_time_event( state.params, t_right, tol, + &[], )?; + ctx.runtime_params + .borrow_mut() + .copy_from_slice(state.params); let (native_y, native_dy) = backend.reset_vectors(state.current_y, state.params, t_right)?; backend.reset( &native_y, @@ -633,6 +694,7 @@ fn reinitialize_solver_after_time_event( state.params, *state.current_t, rumoca_solver::event_solver_step_cap(ctx.opts.dt), + RootStartBoundary::Scheduled, ) } @@ -642,10 +704,10 @@ fn advance_to_target_once( target: f64, event_stop: Option, backend: &mut St, - deferred_root: &mut Option, + deferred_root: &mut Option, ) -> Result { if event_stop.is_some() { - return advance_to_scheduled_stop(ctx, state, target, backend); + return advance_to_scheduled_stop(ctx, state, target, backend, deferred_root); } advance_output_interval(ctx, state, target, backend, deferred_root) } @@ -655,52 +717,14 @@ fn advance_to_scheduled_stop( state: AdvanceState<'_>, target: f64, backend: &mut St, + _deferred_root: &mut Option, ) -> Result { - if backend.time() > target { - backend.state_mut_back(target)?; - } - if sample_time_match_with_tol(backend.time(), target) { - *state.current_t = target; - state - .params - .copy_from_slice(ctx.runtime_params.borrow().as_slice()); - let native = backend.native_y(); - write_full_y(backend, &native, target, state.current_y, state.params)?; - return Ok(false); - } - backend.set_stop_time(target)?; - loop { - let outcome = match backend.step() { - Ok(outcome) => outcome, - Err(e) => { - backend.trace_step_failure( - state.current_y, - state.params, - *state.current_t, - backend.time(), - &e.to_string(), - ); - return Err(e); - } - }; - match outcome { - StepOutcome::Stop => { - let stop_t = backend.time(); - *state.current_t = stop_t; - state - .params - .copy_from_slice(ctx.runtime_params.borrow().as_slice()); - let native = backend.native_y(); - write_full_y(backend, &native, stop_t, state.current_y, state.params)?; - return Ok(false); - } - StepOutcome::Internal => continue, - StepOutcome::Root { t_root } => { - trace_step_event("scheduled-root", backend.time(), Some(t_root)); - return handle_root_crossing(ctx, state, t_root, target, backend); - } - } - } + // Scheduled equations may switch branch at the exact event time. Clamp the + // solver to the previous representable instant so it cannot evaluate the + // post-event RHS before the runtime applies the event update. A root before + // that stop is processed immediately; the following loop iteration installs + // a fresh stop on the reset solver, so no stale tstop survives a reset. + advance_output_interval_clamped(ctx, state, target, backend) } fn advance_output_interval( @@ -708,7 +732,7 @@ fn advance_output_interval( state: AdvanceState<'_>, target: f64, backend: &mut St, - deferred_root: &mut Option, + deferred_root: &mut Option, ) -> Result { // Backends whose interpolation re-projects algebraics (reduced-state) ask to // land exactly on each output point near discontinuities; otherwise we keep @@ -742,12 +766,22 @@ fn advance_output_interval( }; match outcome { StepOutcome::Stop | StepOutcome::Internal => {} - StepOutcome::Root { t_root } => { + StepOutcome::Root { + t_root, + root_indices, + } => { trace_step_event("output-root", backend.time(), Some(t_root)); let root_after_target = t_root > target && !sample_time_match_with_tol(t_root, target); if !root_after_target { - return handle_root_crossing(ctx, state, t_root, target, backend); + return handle_root_crossing( + ctx, + state, + t_root, + &root_indices, + target, + backend, + ); } let y_at_target = backend.interpolate(target)?; *state.current_t = target; @@ -755,7 +789,10 @@ fn advance_output_interval( .params .copy_from_slice(ctx.runtime_params.borrow().as_slice()); write_full_y(backend, &y_at_target, target, state.current_y, state.params)?; - *deferred_root = Some(t_root); + *deferred_root = Some(PendingRoot { + t: t_root, + indices: root_indices, + }); return Ok(false); } } @@ -804,9 +841,12 @@ fn advance_output_interval_clamped( return Ok(false); } StepOutcome::Internal => continue, - StepOutcome::Root { t_root } => { + StepOutcome::Root { + t_root, + root_indices, + } => { trace_step_event("output-root-clamped", backend.time(), Some(t_root)); - return handle_root_crossing(ctx, state, t_root, target, backend); + return handle_root_crossing(ctx, state, t_root, &root_indices, target, backend); } } } @@ -816,6 +856,7 @@ fn handle_root_crossing( ctx: AdvanceContext<'_>, state: AdvanceState<'_>, t_root: f64, + root_indices: &[usize], target: f64, backend: &mut St, ) -> Result { @@ -824,7 +865,20 @@ fn handle_root_crossing( // settle — all via the backend-neutral kernels (shared with the scheduled // path and every backend through the backend callbacks). let tol = ctx.opts.atol.max(1.0e-10); - backend.state_mut_back(t_root)?; + let solver_t = backend.time(); + let pin_t = if t_root > solver_t { + let root_pin_tol = tol * (1.0 + t_root.abs().max(solver_t.abs())); + if (t_root - solver_t).abs() <= root_pin_tol { + solver_t + } else { + return Err(SimDriverError::Backend(format!( + "root time {t_root:.12} is after backend time {solver_t:.12}" + ))); + } + } else { + t_root + }; + backend.state_mut_back(pin_t)?; let root_t = backend.time(); let native_at_root = backend.native_y(); let event_pre_p = ctx.runtime_params.borrow().as_slice().to_vec(); @@ -842,6 +896,8 @@ fn handle_root_crossing( state.params, )?; state.current_y.copy_from_slice(&event_pre_y); + let root_relation_overrides = + post_root_relation_overrides(ctx.model, root_indices, &event_pre_p, tol)?; bracket_event_limits_kernel( backend, &mut event_pre_y, @@ -853,13 +909,16 @@ fn handle_root_crossing( apply_event_update_kernel( ctx.runtime, backend, - state.current_y, - state.params, - right_t, - tol, - EventPre { - y: &event_pre_y, - p: &event_pre_p, + EventUpdateKernelInput { + y: state.current_y, + p: state.params, + t: right_t, + tol, + root_relation_overrides: &root_relation_overrides, + pre: EventPre { + y: &event_pre_y, + p: &event_pre_p, + }, }, )?; settle_kernel( @@ -869,8 +928,12 @@ fn handle_root_crossing( state.params, right_t, tol, + &root_relation_overrides, )?; commit_pre_params_after_event(ctx.model, state.current_y, state.params, tol); + ctx.runtime_params + .borrow_mut() + .copy_from_slice(state.params); backend.trace_post_event_state(state.current_y, state.params, *state.current_t); let (native_y, native_dy) = backend.reset_vectors(state.current_y, state.params, *state.current_t)?; @@ -880,13 +943,721 @@ fn handle_root_crossing( state.params, *state.current_t, rumoca_solver::event_solver_step_cap(ctx.opts.dt), + RootStartBoundary::Root(&root_relation_overrides), )?; Ok(true) } +pub fn post_root_relation_overrides( + model: &solve::SolveModel, + root_indices: &[usize], + params: &[f64], + tol: f64, +) -> Result, RuntimeSolveError> { + let root_count = model.problem.events.root_conditions.output_count(); + let targets = &model.problem.events.root_relation_memory_targets; + if root_count != targets.len() { + return Err(RuntimeSolveError::solve_ir(format!( + "root relation metadata length {} does not match root output count {root_count}", + targets.len() + ))); + } + let mut overrides = Vec::new(); + for root_index in root_indices.iter().copied().collect::>() { + if root_index >= root_count || root_index >= targets.len() { + return Err(RuntimeSolveError::solve_ir(format!( + "root crossing index {root_index} is outside root metadata (roots={root_count}, targets={})", + targets.len() + ))); + } + let Some(target) = targets[root_index] else { + continue; + }; + let solve::ScalarSlot::P { index, .. } = target else { + return Err(RuntimeSolveError::solve_ir(format!( + "root crossing index {root_index} has non-parameter relation memory target" + ))); + }; + let current = params.get(index).copied().ok_or_else(|| { + RuntimeSolveError::solve_ir(format!( + "root crossing index {root_index} relation memory parameter {index} is outside parameter storage" + )) + })?; + let post = if current.abs() <= tol { + 1.0 + } else if (current - 1.0).abs() <= tol { + 0.0 + } else { + return Err(RuntimeSolveError::solve_ir(format!( + "root crossing index {root_index} relation memory value {current} is not boolean" + ))); + }; + overrides.push((root_index, post)); + } + Ok(overrides) +} + fn trace_step_event(kind: &str, solver_t: f64, root_t: Option) { if !tracing::enabled!(target: EVENT_TRACE_TARGET, tracing::Level::DEBUG) { return; } tracing::debug!(target: EVENT_TRACE_TARGET, "{kind} solver_t={solver_t:.12} root_t={root_t:?}"); } + +#[cfg(test)] +mod tests { + use super::*; + + fn relation_root_model(target: Option) -> solve::SolveModel { + let mut model = solve::SolveModel::default(); + model.problem.events.root_conditions = solve::ScalarProgramBlock { + programs: vec![vec![ + solve::LinearOp::Const { dst: 0, value: 0.0 }, + solve::LinearOp::StoreOutput { src: 0 }, + ]], + program_spans: vec![rumoca_core::Span::DUMMY], + output_indices: vec![0], + }; + model.problem.events.root_relation_memory_targets = vec![target]; + model + } + + #[test] + fn root_relation_overrides_deduplicate_indices_and_toggle_parameter_memory() { + let model = relation_root_model(Some(solve::scalar_slot_p(0))); + + let overrides = post_root_relation_overrides(&model, &[0, 0], &[0.0], 1.0e-12) + .expect("confirmed root should toggle relation memory exactly once"); + + assert_eq!(overrides, vec![(0, 1.0)]); + } + + #[test] + fn root_relation_overrides_fail_closed_for_non_parameter_target() { + let model = relation_root_model(Some(solve::scalar_slot_y(0))); + + let error = post_root_relation_overrides(&model, &[0], &[0.0], 1.0e-12) + .expect_err("non-parameter relation memory must not infer a crossing direction"); + + assert!(error.to_string().contains("non-parameter")); + } + + enum TrajectoryStep { + Internal { + solver_t: f64, + }, + Stop { + solver_t: f64, + }, + Root { + solver_t: f64, + t_root: f64, + root_indices: Vec, + }, + } + + struct TrajectoryGenerationBackend { + time: f64, + steps: std::collections::VecDeque, + generation: usize, + step_generations: Vec, + pinned_times: Vec, + stop_times: Vec, + reset_count: usize, + exact_output_steps: bool, + } + + impl TrajectoryGenerationBackend { + fn new(steps: impl IntoIterator) -> Self { + Self { + time: 0.0, + steps: steps.into_iter().collect(), + generation: 0, + step_generations: Vec::new(), + pinned_times: Vec::new(), + stop_times: Vec::new(), + reset_count: 0, + exact_output_steps: false, + } + } + } + + impl SolverAdvanceBackend for TrajectoryGenerationBackend { + fn time(&self) -> f64 { + self.time + } + + fn native_y(&self) -> Vec { + Vec::new() + } + + fn step(&mut self) -> Result { + let step = self + .steps + .pop_front() + .ok_or_else(|| SimDriverError::Backend("test backend ran out of steps".into()))?; + self.step_generations.push(self.generation); + match step { + TrajectoryStep::Internal { solver_t } => { + self.time = solver_t; + Ok(StepOutcome::Internal) + } + TrajectoryStep::Stop { solver_t } => { + self.time = solver_t; + Ok(StepOutcome::Stop) + } + TrajectoryStep::Root { + solver_t, + t_root, + root_indices, + } => { + self.time = solver_t; + Ok(StepOutcome::Root { + t_root, + root_indices, + }) + } + } + } + + fn set_stop_time(&mut self, stop_time: f64) -> Result<(), SimDriverError> { + self.stop_times.push(stop_time); + Ok(()) + } + + fn interpolate(&mut self, t: f64) -> Result, SimDriverError> { + if t > self.time && !sample_time_match_with_tol(t, self.time) { + return Err(SimDriverError::Backend(format!( + "cannot interpolate forward from {:.12} to {t:.12}", + self.time + ))); + } + Ok(Vec::new()) + } + + fn state_mut_back(&mut self, t: f64) -> Result<(), SimDriverError> { + if t > self.time && !sample_time_match_with_tol(t, self.time) { + return Err(SimDriverError::Backend(format!( + "cannot pin forward from {:.12} to {t:.12}", + self.time + ))); + } + self.time = t; + self.pinned_times.push(t); + Ok(()) + } + + fn native_to_full_y( + &self, + native: &[f64], + _t: f64, + _params: &[f64], + ) -> Result, SimDriverError> { + Ok(native.to_vec()) + } + + fn reset_vectors( + &self, + current_y: &[f64], + _params: &[f64], + _t: f64, + ) -> Result<(Vec, Vec), SimDriverError> { + Ok((current_y.to_vec(), vec![0.0; current_y.len()])) + } + + fn reset( + &mut self, + _native_y: &[f64], + _native_dy: &[f64], + _params: &[f64], + t: f64, + _h_cap: f64, + _boundary: RootStartBoundary<'_>, + ) -> Result<(), SimDriverError> { + self.time = t; + self.generation += 1; + self.reset_count += 1; + Ok(()) + } + + fn prefer_exact_output_steps(&self) -> bool { + self.exact_output_steps + } + + fn project_algebraics( + &self, + _y: &mut [f64], + _p: &mut [f64], + _t: f64, + _tol: f64, + ) -> Result { + Ok(false) + } + + fn derivative_guess( + &self, + y: &[f64], + _p: &[f64], + _t: f64, + ) -> Result, SimDriverError> { + Ok(vec![0.0; y.len()]) + } + + fn record_sample( + &self, + recorded_times: &mut Vec, + _data: &mut [Vec], + _y: &[f64], + _p: &[f64], + t: f64, + ) -> Result<(), SimDriverError> { + recorded_times.push(t); + Ok(()) + } + + fn refresh_observation( + &self, + _y: &mut [f64], + _p: &mut [f64], + _t: f64, + ) -> Result<(), SimDriverError> { + Ok(()) + } + + fn trace_step_failure( + &self, + _y: &[f64], + _params: &[f64], + _current_t: f64, + _solver_t: f64, + _error: &str, + ) { + } + + fn trace_post_event_state(&self, _y: &[f64], _params: &[f64], _t: f64) {} + } + + fn run_trajectory_driver( + model: &solve::SolveModel, + times: &[f64], + steps: impl IntoIterator, + exact_output_steps: bool, + ) -> ( + Result<(), SimDriverError>, + TrajectoryGenerationBackend, + Vec, + f64, + Vec, + ) { + let runtime = SolveRuntime::new(model).expect("empty runtime should prepare"); + let runtime_state = SimulationRuntimeState::new(); + let opts = SimOptions { + atol: 1.0e-12, + ..SimOptions::default() + }; + let runtime_params = Rc::new(RefCell::new(model.parameters.clone())); + let mut current_y = Vec::new(); + let mut params = model.parameters.clone(); + let mut data = Vec::new(); + let mut recorded_times = Vec::new(); + let mut current_t = 0.0; + let mut backend = TrajectoryGenerationBackend::new(steps); + backend.exact_output_steps = exact_output_steps; + + let result = simulate_state_targets( + model, + &opts, + times, + &runtime_params, + &mut backend, + StateTrajectory { + params: &mut params, + data: &mut data, + recorded_times: &mut recorded_times, + current_t: &mut current_t, + current_y: &mut current_y, + runtime: &runtime, + runtime_state: &runtime_state, + }, + ); + + (result, backend, recorded_times, current_t, params) + } + + #[test] + fn scheduled_event_invalidates_deferred_future_root_from_pre_event_trajectory() { + let mut model = solve::SolveModel::default(); + model.problem.events.scheduled_time_events.push(0.5); + let left_t = event_left_limit_time(0.5); + + let (result, backend, recorded_times, current_t, _params) = run_trajectory_driver( + &model, + &[1.0], + [ + TrajectoryStep::Stop { solver_t: left_t }, + TrajectoryStep::Internal { solver_t: 1.0 }, + ], + false, + ); + + result.expect("scheduled reset must replace the prior trajectory at its left limit"); + assert!(backend.pinned_times.is_empty()); + assert_eq!(backend.step_generations, vec![0, 1]); + assert_eq!(backend.stop_times, vec![left_t]); + assert_eq!(backend.reset_count, 1); + assert_eq!(recorded_times.last(), Some(&1.0)); + assert_eq!(current_t, 1.0); + } + + #[test] + fn root_before_scheduled_event_reinstalls_stop_after_reset() { + let mut model = solve::SolveModel::default(); + model.problem.events.scheduled_time_events.push(0.5); + let left_t = event_left_limit_time(0.5); + let root_t = 0.49; + + let (result, backend, recorded_times, current_t, _params) = run_trajectory_driver( + &model, + &[1.0], + [ + TrajectoryStep::Root { + solver_t: root_t, + t_root: root_t, + root_indices: Vec::new(), + }, + TrajectoryStep::Stop { solver_t: left_t }, + TrajectoryStep::Internal { solver_t: 1.0 }, + ], + false, + ); + + result.expect("root reset before a scheduled event must not retain a stale stop"); + assert_eq!(backend.stop_times, vec![left_t, left_t]); + assert_eq!(backend.step_generations, vec![0, 1, 2]); + assert_eq!(backend.reset_count, 2); + assert_eq!(recorded_times.last(), Some(&1.0)); + assert_eq!(current_t, 1.0); + } + + #[test] + fn ordinary_output_boundary_preserves_deferred_future_root() { + const ROOT_T: f64 = 0.500_033_856_224; + let model = solve::SolveModel::default(); + + let (result, backend, recorded_times, current_t, _params) = run_trajectory_driver( + &model, + &[0.5, 1.0], + [ + TrajectoryStep::Root { + solver_t: ROOT_T, + t_root: ROOT_T, + root_indices: Vec::new(), + }, + TrajectoryStep::Internal { solver_t: 1.0 }, + ], + false, + ); + + result.expect("ordinary dense output must preserve the deferred root"); + assert_eq!(backend.pinned_times, vec![ROOT_T]); + assert_eq!(backend.step_generations, vec![0, 1]); + assert_eq!(backend.reset_count, 1); + assert_eq!(recorded_times, vec![0.5, 1.0]); + assert_eq!(current_t, 1.0); + } + + #[test] + fn deferred_root_preserves_relation_override_indices() { + const ROOT_T: f64 = 0.500_033_856_224; + let mut model = relation_root_model(Some(solve::scalar_slot_p(0))); + model.parameters = vec![0.0]; + + let (result, _backend, _recorded_times, _current_t, params) = run_trajectory_driver( + &model, + &[0.5, 1.0], + [ + TrajectoryStep::Root { + solver_t: ROOT_T, + t_root: ROOT_T, + root_indices: vec![0], + }, + TrajectoryStep::Internal { solver_t: 1.0 }, + ], + false, + ); + + result.expect("deferred root should retain relation override metadata"); + assert_eq!(params, vec![1.0]); + } + + #[test] + fn free_and_clamped_root_paths_apply_relation_override_indices() { + for exact_output_steps in [false, true] { + let mut model = relation_root_model(Some(solve::scalar_slot_p(0))); + model.parameters = vec![0.0]; + let followup = if exact_output_steps { + TrajectoryStep::Stop { solver_t: 1.0 } + } else { + TrajectoryStep::Internal { solver_t: 1.0 } + }; + + let (result, _backend, _recorded_times, _current_t, params) = run_trajectory_driver( + &model, + &[1.0], + [ + TrajectoryStep::Root { + solver_t: 0.5, + t_root: 0.5, + root_indices: vec![0], + }, + followup, + ], + exact_output_steps, + ); + + result.expect("root path should apply relation override metadata"); + assert_eq!(params, vec![1.0]); + } + } + + struct ForwardInterpolationRejectingBackend { + time: f64, + pinned_times: Vec, + } + + impl SolverAdvanceBackend for ForwardInterpolationRejectingBackend { + fn time(&self) -> f64 { + self.time + } + + fn native_y(&self) -> Vec { + Vec::new() + } + + fn step(&mut self) -> Result { + unreachable!("root handler test does not step the backend") + } + + fn set_stop_time(&mut self, _stop_time: f64) -> Result<(), SimDriverError> { + unreachable!("root handler test does not set a stop time") + } + + fn interpolate(&mut self, _t: f64) -> Result, SimDriverError> { + unreachable!("root handler test does not interpolate output") + } + + fn state_mut_back(&mut self, t: f64) -> Result<(), SimDriverError> { + if t > self.time { + return Err(SimDriverError::Backend( + "Interpolation time is not within current step".to_string(), + )); + } + self.time = t; + self.pinned_times.push(t); + Ok(()) + } + + fn native_to_full_y( + &self, + native: &[f64], + _t: f64, + _params: &[f64], + ) -> Result, SimDriverError> { + Ok(native.to_vec()) + } + + fn reset_vectors( + &self, + current_y: &[f64], + _params: &[f64], + _t: f64, + ) -> Result<(Vec, Vec), SimDriverError> { + Ok((current_y.to_vec(), vec![0.0; current_y.len()])) + } + + fn reset( + &mut self, + _native_y: &[f64], + _native_dy: &[f64], + _params: &[f64], + t: f64, + _h_cap: f64, + _boundary: RootStartBoundary<'_>, + ) -> Result<(), SimDriverError> { + self.time = t; + Ok(()) + } + + fn prefer_exact_output_steps(&self) -> bool { + false + } + + fn project_algebraics( + &self, + _y: &mut [f64], + _p: &mut [f64], + _t: f64, + _tol: f64, + ) -> Result { + Ok(false) + } + + fn derivative_guess( + &self, + y: &[f64], + _p: &[f64], + _t: f64, + ) -> Result, SimDriverError> { + Ok(vec![0.0; y.len()]) + } + + fn record_sample( + &self, + _recorded_times: &mut Vec, + _data: &mut [Vec], + _y: &[f64], + _p: &[f64], + _t: f64, + ) -> Result<(), SimDriverError> { + unreachable!("root handler test does not record samples") + } + + fn refresh_observation( + &self, + _y: &mut [f64], + _p: &mut [f64], + _t: f64, + ) -> Result<(), SimDriverError> { + unreachable!("root handler test does not refresh observations") + } + + fn trace_step_failure( + &self, + _y: &[f64], + _params: &[f64], + _current_t: f64, + _solver_t: f64, + _error: &str, + ) { + } + + fn trace_post_event_state(&self, _y: &[f64], _params: &[f64], _t: f64) {} + } + + fn handle_test_root( + t_root: f64, + ) -> ( + Result, + ForwardInterpolationRejectingBackend, + f64, + ) { + let model = solve::SolveModel::default(); + let runtime = SolveRuntime::new(&model).expect("empty runtime should prepare"); + let opts = SimOptions { + atol: 1.0e-12, + ..SimOptions::default() + }; + let runtime_params = Rc::new(RefCell::new(Vec::new())); + let mut current_y = Vec::new(); + let mut params = Vec::new(); + let mut current_t = 1.0; + let mut backend = ForwardInterpolationRejectingBackend { + time: 1.0, + pinned_times: Vec::new(), + }; + + let result = handle_root_crossing( + AdvanceContext { + model: &model, + opts: &opts, + runtime: &runtime, + runtime_params: &runtime_params, + }, + AdvanceState { + current_y: &mut current_y, + params: &mut params, + current_t: &mut current_t, + }, + t_root, + &[], + 1.0, + &mut backend, + ); + (result, backend, current_t) + } + + #[test] + fn near_future_root_within_tolerance_pins_to_backend_time() { + let (result, backend, current_t) = handle_test_root(1.0 + 5.0e-13); + let handled = result.expect("near-future root should be handled at backend time"); + + assert!(handled); + assert_eq!(backend.pinned_times, vec![1.0]); + assert_eq!(current_t, 1.0); + } + + #[test] + fn root_handler_applies_relation_override_through_event_update_and_settle() { + let mut model = relation_root_model(Some(solve::scalar_slot_p(0))); + model.parameters = vec![0.0, 0.0]; + model.problem.discrete.rhs = solve::ScalarProgramBlock { + programs: vec![vec![ + solve::LinearOp::LoadP { dst: 0, index: 0 }, + solve::LinearOp::StoreOutput { src: 0 }, + ]], + program_spans: vec![rumoca_core::Span::DUMMY], + output_indices: vec![0], + }; + model.problem.discrete.update_targets = vec![solve::scalar_slot_p(1)]; + model.problem.discrete.pre_modes = vec![solve::DiscreteEventPreMode::FollowCurrent]; + model.problem.discrete.observation_refresh = vec![false]; + model.problem.solve_layout.relation_memory_parameter_indices = vec![0]; + let runtime = SolveRuntime::new(&model).expect("relation runtime should prepare"); + let opts = SimOptions { + atol: 1.0e-12, + ..SimOptions::default() + }; + let runtime_params = Rc::new(RefCell::new(vec![0.0, 0.0])); + let mut current_y = Vec::new(); + let mut params = vec![0.0, 0.0]; + let mut current_t = 1.0; + let mut backend = ForwardInterpolationRejectingBackend { + time: 1.0, + pinned_times: Vec::new(), + }; + + handle_root_crossing( + AdvanceContext { + model: &model, + opts: &opts, + runtime: &runtime, + runtime_params: &runtime_params, + }, + AdvanceState { + current_y: &mut current_y, + params: &mut params, + current_t: &mut current_t, + }, + 1.0, + &[0], + 2.0, + &mut backend, + ) + .expect("root handler should apply the confirmed post-crossing relation memory"); + + assert_eq!(params, vec![1.0, 1.0]); + } + + #[test] + fn future_root_outside_tolerance_returns_structured_error() { + let (result, backend, current_t) = handle_test_root(1.0 + 5.0e-10); + let error = result.expect_err("future root outside tolerance should be rejected"); + + assert_eq!( + error.to_string(), + "root time 1.000000000500 is after backend time 1.000000000000" + ); + assert!(backend.pinned_times.is_empty()); + assert_eq!(current_t, 1.0); + } +} diff --git a/crates/rumoca-eval-solve/src/sparsity.rs b/crates/rumoca-eval-solve/src/sparsity.rs index 7682e5a44..a4d4172a9 100644 --- a/crates/rumoca-eval-solve/src/sparsity.rs +++ b/crates/rumoca-eval-solve/src/sparsity.rs @@ -164,6 +164,18 @@ pub fn row_seed_dependencies(program: &[LinearOp]) -> Result, EvalSol let op_deps = union_regs(&deps, [id, imin, imax])?; set_reg_deps(&mut deps, dst, op_deps); } + LinearOp::ExternalCall { + dst, + args, + arg_count, + .. + } => { + let mut op_deps = BTreeSet::new(); + for arg in args.iter().copied().take(arg_count) { + op_deps.extend(reg_deps(&deps, arg)?); + } + set_reg_deps(&mut deps, dst, op_deps); + } LinearOp::Unary { dst, arg, .. } => { let op_deps = reg_deps(&deps, arg)?; set_reg_deps(&mut deps, dst, op_deps); diff --git a/crates/rumoca-eval-solve/src/table_runtime.rs b/crates/rumoca-eval-solve/src/table_runtime.rs index 18d7c4e03..4b4217ddf 100644 --- a/crates/rumoca-eval-solve/src/table_runtime.rs +++ b/crates/rumoca-eval-solve/src/table_runtime.rs @@ -313,7 +313,11 @@ fn eval_table_1d_lookup( } let last_idx = table.data.len() - 1; - let k = lookup_segment_index(table, x_real)?; + let k = if out_of_range && x < x_min { + 0 + } else { + lookup_segment_index(table, x_real)? + }; let next_idx = (k + 1).min(last_idx); let x0 = table_row_x(table, k)?; let x1 = table_row_x(table, next_idx)?; @@ -362,7 +366,7 @@ fn lookup_segment_index( table_id: table.id, reason: "table requires at least two rows for segment lookup", })?; - if x_real <= table_row_x(table, 0)? { + if x_real < table_row_x(table, 0)? { return Ok(0); } if x_real >= table_row_x(table, last_idx)? { diff --git a/crates/rumoca-exec-cranelift/src/emit.rs b/crates/rumoca-exec-cranelift/src/emit.rs index 5405ea523..c70f40cd0 100644 --- a/crates/rumoca-exec-cranelift/src/emit.rs +++ b/crates/rumoca-exec-cranelift/src/emit.rs @@ -14,7 +14,9 @@ use rumoca_eval_solve::{ eval_table_bound_value_in, eval_table_lookup_slope_value_in, eval_table_lookup_value_in, eval_time_table_next_event_value_in, }; -use rumoca_ir_solve::{BinaryOp, CompareOp, LinearOp, UnaryOp, resolve_indexed_slot}; +use rumoca_ir_solve::{ + BinaryOp, CompareOp, ExternalFunctionKind, LinearOp, UnaryOp, resolve_indexed_slot, +}; use std::cell::{Cell, RefCell}; use std::collections::HashMap; @@ -739,6 +741,7 @@ impl<'a, 'b> RowLowerCtx<'a, 'b> { | LinearOp::ImpureRandomInteger { .. } => Err(CompileError::Backend( "cranelift row compiler does not support discrete random solve-IR ops".to_string(), )), + LinearOp::ExternalCall { function, .. } => Err(external_call_compile_error(function)), LinearOp::Unary { dst, op, arg } => { let x = lookup_reg(self.regs, arg)?; let value = emit_unary_op(self.fb, self.module, self.math, op, x)?; @@ -1424,6 +1427,7 @@ fn is_simple_linear_op(op: LinearOp) -> bool { | LinearOp::ImpureRandomInit { .. } | LinearOp::ImpureRandom { .. } | LinearOp::ImpureRandomInteger { .. } + | LinearOp::ExternalCall { .. } | LinearOp::StoreOutput { .. } ) } @@ -1469,6 +1473,7 @@ fn lower_simple_op(op: LinearOp) -> Result { | LinearOp::ImpureRandomInit { .. } | LinearOp::ImpureRandom { .. } | LinearOp::ImpureRandomInteger { .. } + | LinearOp::ExternalCall { .. } | LinearOp::StoreOutput { .. } => Err(CompileError::Backend( "attempted to lower non-simple runtime op onto the simple row path".to_string(), )), @@ -1524,6 +1529,14 @@ fn max_reg_index(op: LinearOp) -> Result, CompileError> { imax, .. } => Ok(Some(dst.max(id).max(imin).max(imax) as usize)), + LinearOp::ExternalCall { + dst, + args, + arg_count, + .. + } => Ok(Some( + args.iter().copied().take(arg_count).fold(dst, u32::max) as usize, + )), LinearOp::Move { dst, src } => Ok(Some(dst.max(src) as usize)), LinearOp::LinearSolveComponent { dst, @@ -1582,6 +1595,7 @@ fn dst_reg(op: LinearOp) -> Option { | LinearOp::ImpureRandomInit { dst, .. } | LinearOp::ImpureRandom { dst, .. } | LinearOp::ImpureRandomInteger { dst, .. } + | LinearOp::ExternalCall { dst, .. } | LinearOp::Move { dst, .. } | LinearOp::LinearSolveComponent { dst, .. } | LinearOp::Unary { dst, .. } @@ -1649,6 +1663,13 @@ fn validate_row_sources(defined: &[bool], op: LinearOp) -> Result<(), CompileErr validate_reg_defined(defined, imin)?; validate_reg_defined(defined, imax) } + LinearOp::ExternalCall { + args, arg_count, .. + } => args + .iter() + .copied() + .take(arg_count) + .try_for_each(|arg| validate_reg_defined(defined, arg)), LinearOp::Move { src, .. } => validate_reg_defined(defined, src), LinearOp::LinearSolveComponent { matrix_start, @@ -1711,5 +1732,11 @@ fn to_backend_err(err: E) -> CompileError { CompileError::Backend(err.to_string()) } +fn external_call_compile_error(function: ExternalFunctionKind) -> CompileError { + CompileError::Backend(format!( + "external function {function:?} requires a native runtime bridge not provided by rumoca-exec-cranelift" + )) +} + #[cfg(test)] mod emit_tests; diff --git a/crates/rumoca-exec-cranelift/src/emit/emit_tests.rs b/crates/rumoca-exec-cranelift/src/emit/emit_tests.rs index 49ff95189..d87df4ec3 100644 --- a/crates/rumoca-exec-cranelift/src/emit/emit_tests.rs +++ b/crates/rumoca-exec-cranelift/src/emit/emit_tests.rs @@ -51,6 +51,81 @@ fn plan_row_keeps_seed_rows_on_general_runtime_plan() { assert!(matches!(plan, RowPlan::General(_))); } +#[test] +fn plan_row_keeps_external_call_rows_on_general_runtime_plan() { + let row = vec![ + LinearOp::LoadY { dst: 0, index: 0 }, + LinearOp::LoadP { dst: 1, index: 0 }, + LinearOp::ExternalCall { + dst: 2, + function: rumoca_ir_solve::ExternalFunctionKind::BuildingsEnergyPlusExchange, + args: [0, 1, 0, 0, 0, 0, 0, 0], + arg_count: 2, + output_index: 0, + }, + LinearOp::StoreOutput { src: 2 }, + ]; + + let plan = plan_row(&row).expect("general plan"); + assert!(matches!(plan, RowPlan::General(_))); +} + +#[test] +fn external_call_runtime_fails_closed_without_native_bridge() { + let row = vec![ + LinearOp::LoadY { dst: 0, index: 0 }, + LinearOp::LoadP { dst: 1, index: 0 }, + LinearOp::ExternalCall { + dst: 2, + function: rumoca_ir_solve::ExternalFunctionKind::BuildingsEnergyPlusExchange, + args: [0, 1, 0, 0, 0, 0, 0, 0], + arg_count: 2, + output_index: 0, + }, + LinearOp::StoreOutput { src: 2 }, + ]; + let plan = plan_row(&row).expect("general plan"); + let mut scratch = Vec::new(); + let mut out = [0.0]; + + let err = execute_row( + &plan, + &mut scratch, + row_inputs(&[3.0], &[5.0], 0.0, None, &[]), + &mut out, + ) + .expect_err("external calls require an explicit native bridge"); + + assert!( + matches!(err, CompileError::Backend(message) if message.contains("external function BuildingsEnergyPlusExchange")), + "external-call error should identify the function" + ); +} + +#[test] +fn external_call_requires_initialized_argument_registers() { + let row = vec![ + LinearOp::LoadY { dst: 0, index: 0 }, + LinearOp::ExternalCall { + dst: 2, + function: rumoca_ir_solve::ExternalFunctionKind::BuildingsEnergyPlusExchange, + args: [0, 1, 0, 0, 0, 0, 0, 0], + arg_count: 2, + output_index: 0, + }, + LinearOp::StoreOutput { src: 2 }, + ]; + + let err = match plan_row(&row) { + Ok(_) => panic!("undefined external-call arg should be rejected"), + Err(err) => err, + }; + + assert!( + matches!(err, CompileError::Backend(message) if message.contains("undefined register r1")) + ); +} + #[test] fn compile_residual_rows_accepts_linear_solve_component() { let row = vec![ diff --git a/crates/rumoca-exec-cranelift/src/emit/interpreter.rs b/crates/rumoca-exec-cranelift/src/emit/interpreter.rs index 0636109a5..6d2b7af46 100644 --- a/crates/rumoca-exec-cranelift/src/emit/interpreter.rs +++ b/crates/rumoca-exec-cranelift/src/emit/interpreter.rs @@ -171,6 +171,9 @@ fn execute_general_op( | LinearOp::ImpureRandomInit { dst, .. } | LinearOp::ImpureRandom { dst, .. } | LinearOp::ImpureRandomInteger { dst, .. } => set_reg_value(regs, dst as usize, 1.0), + LinearOp::ExternalCall { function, .. } => { + return Err(external_call_compile_error(function)); + } LinearOp::Unary { dst, op, arg } => { let x = read_reg_value(regs, arg as usize); set_reg_value(regs, dst as usize, apply_unary(op, x)); diff --git a/crates/rumoca-exec-mlir/tests/drone_monte_carlo.rs b/crates/rumoca-exec-mlir/tests/drone_monte_carlo.rs index 9e4aa8319..86a094e56 100644 --- a/crates/rumoca-exec-mlir/tests/drone_monte_carlo.rs +++ b/crates/rumoca-exec-mlir/tests/drone_monte_carlo.rs @@ -163,6 +163,9 @@ fn drone_prepared_model(m: f64, j: f64, f: f64, g: f64) -> rumoca_ir_solve::Solv initialization: InitializationSolveSystem { residual: ComputeBlock::from_scalar_program_block(zero_rb.clone()), row_targets: Vec::new(), + direct_families: Vec::new(), + required_target_ranges: Vec::new(), + fixed_target_ranges: Vec::new(), projection_indices: Vec::new(), projection_plan: rumoca_ir_solve::AlgebraicProjectionPlan::default(), update_rhs: ScalarProgramBlock::default(), diff --git a/crates/rumoca-exec-mlir/tests/gpu_trig.rs b/crates/rumoca-exec-mlir/tests/gpu_trig.rs index 692de5c21..94ae78698 100644 --- a/crates/rumoca-exec-mlir/tests/gpu_trig.rs +++ b/crates/rumoca-exec-mlir/tests/gpu_trig.rs @@ -163,6 +163,9 @@ fn nonlinear_drone_prepared(m: f64, j: f64, f: f64, g: f64) -> rumoca_ir_solve:: initialization: InitializationSolveSystem { residual: ComputeBlock::from_scalar_program_block(zero_rb.clone()), row_targets: Vec::new(), + direct_families: Vec::new(), + required_target_ranges: Vec::new(), + fixed_target_ranges: Vec::new(), projection_indices: Vec::new(), projection_plan: rumoca_ir_solve::AlgebraicProjectionPlan::default(), update_rhs: ScalarProgramBlock::default(), diff --git a/crates/rumoca-exec-mlir/tests/integrate.rs b/crates/rumoca-exec-mlir/tests/integrate.rs index 67d4ad64b..da231fab6 100644 --- a/crates/rumoca-exec-mlir/tests/integrate.rs +++ b/crates/rumoca-exec-mlir/tests/integrate.rs @@ -94,6 +94,9 @@ fn decay_model() -> rumoca_ir_solve::SolveModel { initialization: InitializationSolveSystem { residual: ComputeBlock::from_scalar_program_block(zero_rb.clone()), row_targets: Vec::new(), + direct_families: Vec::new(), + required_target_ranges: Vec::new(), + fixed_target_ranges: Vec::new(), projection_indices: Vec::new(), projection_plan: rumoca_ir_solve::AlgebraicProjectionPlan::default(), update_rhs: ScalarProgramBlock::default(), diff --git a/crates/rumoca-exec-mlir/tests/linsolve_mlir.rs b/crates/rumoca-exec-mlir/tests/linsolve_mlir.rs index 2217c6abd..d4c9338ef 100644 --- a/crates/rumoca-exec-mlir/tests/linsolve_mlir.rs +++ b/crates/rumoca-exec-mlir/tests/linsolve_mlir.rs @@ -61,6 +61,7 @@ fn linsolve_block() -> ComputeBlock { rhs_start: 4, n: 2, next_reg: 6, + output_indices: Vec::new(), metadata: rumoca_ir_solve::TensorNodeMetadata::default(), span: Span::from_offsets(SourceId::from_source_name(label), 0, label.len()), }], @@ -196,6 +197,7 @@ fn linsolve_partial_pivoting_correctness() { rhs_start: 4, n: 2, next_reg: 6, + output_indices: Vec::new(), metadata: rumoca_ir_solve::TensorNodeMetadata::default(), span: Span::from_offsets(SourceId::from_source_name(label), 0, label.len()), }], diff --git a/crates/rumoca-exec-mlir/tests/options.rs b/crates/rumoca-exec-mlir/tests/options.rs index aed33a517..7b3469ac3 100644 --- a/crates/rumoca-exec-mlir/tests/options.rs +++ b/crates/rumoca-exec-mlir/tests/options.rs @@ -59,6 +59,9 @@ fn decay_model() -> rumoca_ir_solve::SolveModel { initialization: InitializationSolveSystem { residual: ComputeBlock::from_scalar_program_block(zero_rb.clone()), row_targets: Vec::new(), + direct_families: Vec::new(), + required_target_ranges: Vec::new(), + fixed_target_ranges: Vec::new(), projection_indices: Vec::new(), projection_plan: rumoca_ir_solve::AlgebraicProjectionPlan::default(), update_rhs: ScalarProgramBlock::default(), diff --git a/crates/rumoca-exec-wasm/src/emit.rs b/crates/rumoca-exec-wasm/src/emit.rs index ff57dd448..18f1de4e8 100644 --- a/crates/rumoca-exec-wasm/src/emit.rs +++ b/crates/rumoca-exec-wasm/src/emit.rs @@ -344,6 +344,12 @@ fn max_register_for_op(op: &LinearOp) -> Result { imax, .. } => Ok(dst.max(id).max(imin).max(imax) as usize), + LinearOp::ExternalCall { + dst, + args, + arg_count, + .. + } => Ok(args.iter().copied().take(arg_count).fold(dst, Reg::max) as usize), LinearOp::Move { dst, src } => Ok(dst.max(src) as usize), LinearOp::LinearSolveComponent { dst, @@ -487,6 +493,11 @@ impl<'a> BodyEmitter<'a> { "WASM backend does not yet support discrete random solve-IR ops".to_string(), ); } + LinearOp::ExternalCall { function, .. } => { + return Err(format!( + "WASM backend does not yet support external function {function:?}; a host bridge is required" + )); + } LinearOp::Unary { dst, op, arg } => self.emit_unary(dst, op, arg)?, LinearOp::Binary { dst, op, lhs, rhs } => self.emit_binary(dst, op, lhs, rhs)?, LinearOp::Compare { dst, op, lhs, rhs } => self.emit_compare(dst, op, lhs, rhs)?, diff --git a/crates/rumoca-galec-codegen/src/lower/expr/references.rs b/crates/rumoca-galec-codegen/src/lower/expr/references.rs index e2824d329..4880a6408 100644 --- a/crates/rumoca-galec-codegen/src/lower/expr/references.rs +++ b/crates/rumoca-galec-codegen/src/lower/expr/references.rs @@ -143,6 +143,30 @@ impl ExprLowerer<'_> { span: Span, ) -> Result { let dims = &classified.variable.dims; + let projected_subscripts; + let base_start = subscripts.len().saturating_sub(dims.len()); + let subscripts = + if base_start > 0 && subscripts[base_start..].iter().any(is_slice_subscript) { + projected_subscripts = self.project_indexed_slice_subscripts( + classified.variable.name.as_str(), + &subscripts[base_start..], + &subscripts[..base_start], + span, + )?; + projected_subscripts.as_slice() + } else { + subscripts + }; + // Flattening can preserve a scalarized trailing index in the rendered + // reference name while also carrying a full-rank subscript vector for + // a row/range slice. In that form the explicit subscripts are the + // authoritative original-rank access; counting the rendered suffix a + // second time invents an extra dimension. + let prefix_indices = if subscripts.len() == dims.len() { + Vec::new() + } else { + prefix_indices + }; let subscript_count = prefix_indices.len() + subscripts.len(); if subscript_count == 0 { return Ok(Typed::array( diff --git a/crates/rumoca-galec-codegen/tests/projection_front_half.rs b/crates/rumoca-galec-codegen/tests/projection_front_half.rs index 7f0c9cd24..81ffeb792 100644 --- a/crates/rumoca-galec-codegen/tests/projection_front_half.rs +++ b/crates/rumoca-galec-codegen/tests/projection_front_half.rs @@ -157,7 +157,13 @@ mod admissibility { #[test] fn runtime_events_rejected_with_projection_scope_wording() { let mut model = base_dae(); - model.events.scheduled_time_events.push(0.5); + model + .events + .scheduled_time_events + .push(dae::DaeScheduledTimeEvent { + time: 0.5, + source_span: None, + }); let errors = check_admissibility(&GalecInput::new(&model, "M")).unwrap_err(); assert_eq!(codes(&errors), vec!["ET003"]); assert!( @@ -246,7 +252,13 @@ mod admissibility { let mut model = base_dae(); model.metadata.is_partial = true; model.clocks.schedules.clear(); - model.events.scheduled_time_events.push(1.0); + model + .events + .scheduled_time_events + .push(dae::DaeScheduledTimeEvent { + time: 1.0, + source_span: None, + }); let errors = check_admissibility(&GalecInput::new(&model, "M")).unwrap_err(); let codes = codes(&errors); assert!(codes.contains(&"ET007"), "{codes:?}"); diff --git a/crates/rumoca-galec-codegen/tests/spec_0034_battery.rs b/crates/rumoca-galec-codegen/tests/spec_0034_battery.rs index 9df64be97..3977fcf98 100644 --- a/crates/rumoca-galec-codegen/tests/spec_0034_battery.rs +++ b/crates/rumoca-galec-codegen/tests/spec_0034_battery.rs @@ -326,7 +326,13 @@ mod rejected_constructs { .insert(VarName::new("tableLookup"), function); }), ("runtime event", "ET003", |model| { - model.events.scheduled_time_events.push(0.5); + model + .events + .scheduled_time_events + .push(rumoca_ir_dae::DaeScheduledTimeEvent { + time: 0.5, + source_span: None, + }); }), ("dynamic clock", "ET004", |model| { model.clocks.triggered_conditions.push(boolean(true)); @@ -937,6 +943,37 @@ mod array_vector_regressions { ); } + #[test] + fn flattened_trailing_index_projects_through_full_rank_slice_subscripts() { + let row_slice = Expression::VarRef { + name: Reference::new("waypoints"), + subscripts: vec![ + Subscript::index(2, Span::DUMMY), + subscript_expr(var("__pre__.idx")), + Subscript::colon(Span::DUMMY), + ], + span: Span::DUMMY, + }; + let mut model = model_with_body(row_slice); + add_waypoints(&mut model); + let mut idx = variable("idx"); + idx.start = Some(integer(1)); + idx.min = Some(integer(1)); + idx.max = Some(integer(3)); + model + .variables + .discrete_valued + .insert(idx.name.clone(), idx); + add_pre_slot(&mut model, "idx", integer(1), Vec::new()); + let mut types = base_types(); + types.insert(VarName::new("idx"), ScalarType::Integer); + types.insert(VarName::new("waypoints"), ScalarType::Real); + + let alg = render_algorithm_code(&lower(&model, &types)).expect("renders"); + assert!(alg.contains("self.waypoints[1, 2]"), "{alg}"); + assert!(alg.contains("self.waypoints[3, 2]"), "{alg}"); + } + #[test] fn scalarized_index_into_range_slice_projects_original_dimension() { let range_slice = Expression::VarRef { diff --git a/crates/rumoca-ir-ast/src/nodes.rs b/crates/rumoca-ir-ast/src/nodes.rs index 8ba505cc5..cd759ade9 100644 --- a/crates/rumoca-ir-ast/src/nodes.rs +++ b/crates/rumoca-ir-ast/src/nodes.rs @@ -402,6 +402,9 @@ pub struct ExternalFunction { pub output: Option, /// Arguments passed to the external function. pub args: Vec, + /// Annotation arguments on the external declaration (MLS §12.9), such as + /// Library, IncludeDirectory, LibraryDirectory, and Include. + pub annotation: Vec, } #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] diff --git a/crates/rumoca-ir-ast/src/visitor/read_only.rs b/crates/rumoca-ir-ast/src/visitor/read_only.rs index 80c4c5b3f..a52ae7cad 100644 --- a/crates/rumoca-ir-ast/src/visitor/read_only.rs +++ b/crates/rumoca-ir-ast/src/visitor/read_only.rs @@ -74,6 +74,7 @@ pub enum ExpressionContext { StatementFunctionOutput, ExtendModification, ExternalArgument, + ExternalAnnotation, } pub enum VisitScope<'a> { @@ -810,6 +811,9 @@ pub trait Visitor { for arg in &external.args { self.visit_expression_ctx(arg, ExpressionContext::ExternalArgument)?; } + for annotation in &external.annotation { + self.visit_expression_ctx(annotation, ExpressionContext::ExternalAnnotation)?; + } Continue(()) } } diff --git a/crates/rumoca-ir-ast/src/visitor/tests.rs b/crates/rumoca-ir-ast/src/visitor/tests.rs index 7651b4139..424d5c23e 100644 --- a/crates/rumoca-ir-ast/src/visitor/tests.rs +++ b/crates/rumoca-ir-ast/src/visitor/tests.rs @@ -490,6 +490,7 @@ fn make_expression_context_dispatch_class() -> ClassDef { ]], external: Some(ExternalFunction { args: vec![make_int(13)], + annotation: vec![make_int(14)], ..Default::default() }), ..Default::default() @@ -513,6 +514,7 @@ fn assert_expression_contexts_seen(seen: &[ExpressionContext]) { assert!(seen.contains(&ExpressionContext::StatementAssertLevel)); assert!(seen.contains(&ExpressionContext::StatementFunctionOutput)); assert!(seen.contains(&ExpressionContext::ExternalArgument)); + assert!(seen.contains(&ExpressionContext::ExternalAnnotation)); } #[test] diff --git a/crates/rumoca-ir-dae/src/expr_query.rs b/crates/rumoca-ir-dae/src/expr_query.rs index c6e3d7d87..a58baeb5d 100644 --- a/crates/rumoca-ir-dae/src/expr_query.rs +++ b/crates/rumoca-ir-dae/src/expr_query.rs @@ -206,6 +206,12 @@ struct ContainsVarChecker<'a> { impl ExpressionVisitor for ContainsVarChecker<'_> { fn visit_expression(&mut self, expr: &Expression) { if !self.found { + if let Some(exact_name) = expr_exact_name(expr) + && exact_name == self.var.as_str() + { + self.found = true; + return; + } self.walk_expression(expr); } } diff --git a/crates/rumoca-ir-dae/src/lib.rs b/crates/rumoca-ir-dae/src/lib.rs index fe3465b5a..b67e78650 100644 --- a/crates/rumoca-ir-dae/src/lib.rs +++ b/crates/rumoca-ir-dae/src/lib.rs @@ -27,7 +27,7 @@ use rumoca_core::{ use serde::ser::{SerializeStruct, SerializeTuple}; use serde::{Deserialize, Serialize}; -pub const DAE_SCHEMA_VERSION: u16 = 6; +pub const DAE_SCHEMA_VERSION: u16 = 7; mod event_threshold; mod expr_query; @@ -45,9 +45,10 @@ pub use types::{ remap_structured_families_after_expansion, structured_equation_slot, }; pub use visitor::{ - AlgorithmOutputCollector, ContainsDerChecker, ContainsDerOfStateChecker, DaeExpressionRewriter, - DaeVariableMutVisitor, DaeVisitor, ImplicitSampleChecker, StateVariableCollector, - StatementScope, StatementVisitor, VarRefCollector, VarRefWithSubscriptsCollector, + AlgorithmOutputCollector, ContainsDerChecker, ContainsDerOfStateChecker, DaeEquationPartition, + DaeExpressionRewriter, DaeVariableMutVisitor, DaeVisitor, ImplicitSampleChecker, + StateVariableCollector, StatementScope, StatementVisitor, TryDaeExpressionRewriter, + VarRefCollector, VarRefWithSubscriptsCollector, }; /// Scalar counts for canonical runtime variable partitions. @@ -133,6 +134,7 @@ struct DaeWire { initial_equations: Vec, #[serde(rename = "initial_structured_equations")] initial_structured_equations: Vec, + initial_equation_provenance: Vec, #[serde(rename = "f_z")] real_updates: Vec, #[serde(rename = "f_m")] @@ -142,7 +144,7 @@ struct DaeWire { #[serde(default, rename = "relation")] relations: Vec, synthetic_root_conditions: Vec, - scheduled_time_events: Vec, + scheduled_time_events: Vec, scheduled_root_conditions: Vec, event_actions: Vec, constructor_exprs: Vec, @@ -182,7 +184,7 @@ impl Serialize for Dae { S: serde::Serializer, { if !serializer.is_human_readable() { - let mut tuple = serializer.serialize_tuple(29)?; + let mut tuple = serializer.serialize_tuple(30)?; tuple.serialize_element(&self.schema_version)?; tuple.serialize_element(&self.variables.states)?; tuple.serialize_element(&self.variables.algebraics)?; @@ -196,6 +198,7 @@ impl Serialize for Dae { tuple.serialize_element(&self.continuous.structured_equations)?; tuple.serialize_element(&self.initialization.equations)?; tuple.serialize_element(&self.initialization.structured_equations)?; + tuple.serialize_element(&self.initialization.equation_provenance)?; tuple.serialize_element(&self.discrete.real_updates)?; tuple.serialize_element(&self.discrete.valued_updates)?; tuple.serialize_element(&self.conditions.equations)?; @@ -215,7 +218,7 @@ impl Serialize for Dae { return tuple.end(); } - let mut state = serializer.serialize_struct("Dae", 29)?; + let mut state = serializer.serialize_struct("Dae", 30)?; state.serialize_field("schema_version", &self.schema_version)?; state.serialize_field("x", &self.variables.states)?; state.serialize_field("y", &self.variables.algebraics)?; @@ -235,6 +238,10 @@ impl Serialize for Dae { "initial_structured_equations", &self.initialization.structured_equations, )?; + state.serialize_field( + "initial_equation_provenance", + &self.initialization.equation_provenance, + )?; state.serialize_field("f_z", &self.discrete.real_updates)?; state.serialize_field("f_m", &self.discrete.valued_updates)?; state.serialize_field("f_c", &self.conditions.equations)?; @@ -273,6 +280,11 @@ impl<'de> Deserialize<'de> for Dae { wire.schema_version, DAE_SCHEMA_VERSION ))); } + if wire.initial_equations.len() != wire.initial_equation_provenance.len() { + return Err(serde::de::Error::custom( + "DAE initial equation provenance cardinality mismatch", + )); + } Ok(Self { schema_version: wire.schema_version, @@ -293,6 +305,7 @@ impl<'de> Deserialize<'de> for Dae { initialization: DaeInitializationPartition { equations: wire.initial_equations, structured_equations: wire.initial_structured_equations, + equation_provenance: wire.initial_equation_provenance, }, discrete: DaeDiscretePartition { real_updates: wire.real_updates, @@ -501,6 +514,19 @@ pub struct DaeInitializationPartition { /// `initial_equations`. #[serde(rename = "initial_structured_equations")] pub structured_equations: Vec, + /// Typed provenance for generated initialization rows. This remains an + /// serialized phase contract with one entry per initialization equation. + #[serde(rename = "initial_equation_provenance")] + pub equation_provenance: Vec, +} + +/// Semantic origin of an initialization equation. Consumers must use this +/// rather than parsing the human-readable `Equation::origin` debug label. +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)] +pub enum InitializationEquationProvenance { + #[default] + User, + FixedStart, } #[derive(Debug, Clone, Default, Serialize, Deserialize)] @@ -533,7 +559,7 @@ pub struct DaeEventPartition { pub synthetic_root_conditions: Vec, /// Scheduled discontinuity instants derived at compile time. /// This is canonical runtime metadata (always present in DAE schema). - pub scheduled_time_events: Vec, + pub scheduled_time_events: Vec, /// Root rows that correspond to periodic sample schedules. /// /// `root_index` is in Solve root-condition order: @@ -549,6 +575,14 @@ pub struct DaeEventPartition { pub event_actions: Vec, } +/// Compile-time scheduled discontinuity with the source expression that +/// established the runtime instant. +#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)] +pub struct DaeScheduledTimeEvent { + pub time: f64, + pub source_span: Option, +} + #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub struct DaeScheduledRootCondition { pub root_index: usize, @@ -630,6 +664,13 @@ pub struct DaeMetadata { /// tick is the source variable's `start` attribute. #[serde(default)] pub variable_starts: IndexMap, + /// Source variables whose declared type is not numeric (currently String). + /// + /// Runtime numeric env construction uses this to keep metadata-only + /// bindings out of the numeric parameter tail while still preserving their + /// shape/start metadata for expressions such as `size(substanceNames, 1)`. + #[serde(default)] + pub nonnumeric_variable_names: Vec, /// Discrete-valued variables whose source causality is input. /// /// These are carried in `m` when event/discrete lowering needs them in the @@ -646,6 +687,13 @@ pub struct DaeMetadata { #[serde(default)] pub interface_flow_count: usize, + /// Count of outside stream connector equations (MLS §15.1). + /// + /// These equations belong to stream `inStream`/`actualStream` semantics and + /// are tracked separately from ordinary connection equality equations. + #[serde(default)] + pub stream_interface_equation_count: usize, + /// Overconstrained interface balance correction (MLS §4.8, §9.4). /// /// For overconstrained connector types (e.g., QuasiStatic Reference), this is: @@ -667,6 +715,14 @@ pub struct DaeMetadata { #[serde(default)] pub oc_break_edge_scalar_count: usize, + /// Scalar gauge freedom for rooted overconstrained connection components. + /// + /// A selected root fixes the connection graph topology, but the root record's + /// absolute reference coordinate remains a gauge degree of freedom for the + /// simulation DAE. Admission uses this only as deficit-only closure. + #[serde(default)] + pub overconstrained_root_gauge_count: usize, + /// Optional description string from the root class declaration. #[serde(default)] pub model_description: Option, @@ -1345,7 +1401,12 @@ mod tests { Span::DUMMY, "when sample trigger then hold.y", )); - dae.events.scheduled_time_events.push(0.1); + dae.events + .scheduled_time_events + .push(super::DaeScheduledTimeEvent { + time: 0.1, + source_span: Some(fixture_span()), + }); dae.clocks.schedules.push(ClockSchedule { period_seconds: 0.1, phase_seconds: 0.0, @@ -1442,6 +1503,12 @@ mod tests { "DAE JSON must carry an explicit schema_version" ); + let mut previous = value.clone(); + previous["schema_version"] = serde_json::json!(DAE_SCHEMA_VERSION - 1); + let err = serde_json::from_value::(previous) + .expect_err("previous DAE schema version must fail after wire replacement"); + assert!(err.to_string().contains("unsupported DAE schema_version")); + let mut unsupported = value; unsupported["schema_version"] = serde_json::json!(DAE_SCHEMA_VERSION + 1); let err = serde_json::from_value::(unsupported) @@ -1449,6 +1516,50 @@ mod tests { assert!(err.to_string().contains("unsupported DAE schema_version")); } + #[test] + fn initialization_provenance_roundtrip_preserves_cardinality() { + let mut dae = Dae::default(); + dae.initialization.equations.push(Equation::residual( + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(0.0), + span: fixture_span(), + }, + fixture_span(), + "roundtrip", + )); + dae.initialization + .equation_provenance + .push(super::InitializationEquationProvenance::FixedStart); + let value = serde_json::to_value(&dae).expect("serialize DAE provenance"); + let decoded: Dae = serde_json::from_value(value.clone()).expect("roundtrip DAE provenance"); + assert_eq!( + decoded.initialization.equation_provenance, + dae.initialization.equation_provenance + ); + let encoded = bincode::serialize(&dae).expect("serialize nonempty DAE provenance"); + let equation_bytes = bincode::serialize(&dae.initialization.equations) + .expect("serialize nonempty initialization equations"); + let _: Vec = bincode::deserialize(&equation_bytes) + .expect("roundtrip nonempty initialization equations"); + let provenance_bytes = bincode::serialize(&dae.initialization.equation_provenance) + .expect("serialize nonempty initialization provenance"); + let _: Vec = + bincode::deserialize(&provenance_bytes) + .expect("roundtrip nonempty initialization provenance"); + let binary_decoded: Dae = + bincode::deserialize(&encoded).expect("roundtrip nonempty DAE provenance"); + assert_eq!( + binary_decoded.initialization.equation_provenance, + vec![super::InitializationEquationProvenance::FixedStart] + ); + + let mut malformed = value; + malformed["initial_equation_provenance"] = serde_json::json!([]); + let error = + serde_json::from_value::(malformed).expect_err("cardinality mismatch must fail"); + assert!(error.to_string().contains("provenance cardinality")); + } + #[test] fn logical_network1_dae_json_matches_committed_golden() { let expected: serde_json::Value = diff --git a/crates/rumoca-ir-dae/src/visitor.rs b/crates/rumoca-ir-dae/src/visitor.rs index e0b0a465f..084efe340 100644 --- a/crates/rumoca-ir-dae/src/visitor.rs +++ b/crates/rumoca-ir-dae/src/visitor.rs @@ -90,6 +90,12 @@ pub trait DaeVisitor { fn visit_event_actions(&mut self, actions: &[crate::DaeEventAction]) { for action in actions { self.visit_expression(&action.condition); + match &action.kind { + crate::DaeEventActionKind::Assert { message } + | crate::DaeEventActionKind::Terminate { message } => { + self.visit_expression(message); + } + } } } @@ -191,8 +197,109 @@ pub trait DaeExpressionRewriter: ExpressionRewriter { fn rewrite_event_actions(&mut self, actions: &mut [crate::DaeEventAction]) { for action in actions { action.condition = self.rewrite_expression(&action.condition); + match &mut action.kind { + crate::DaeEventActionKind::Assert { message } + | crate::DaeEventActionKind::Terminate { message } => { + *message = self.rewrite_expression(message); + } + } + } + } +} + +/// Equation-bearing partitions in the canonical DAE expression traversal. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum DaeEquationPartition { + Continuous, + Initialization, + DiscreteReal, + DiscreteValued, + Condition, +} + +/// Fallible counterpart to [`DaeExpressionRewriter`]. +/// +/// This owns the same schema surface so transformations that can fail do not +/// maintain a second, incomplete partition list. +pub trait TryDaeExpressionRewriter { + type Error; + + fn try_rewrite_dae(&mut self, dae: &mut crate::Dae) -> Result<(), Self::Error> { + self.try_rewrite_equations( + DaeEquationPartition::Continuous, + &mut dae.continuous.equations, + )?; + self.try_rewrite_equations( + DaeEquationPartition::Initialization, + &mut dae.initialization.equations, + )?; + self.try_rewrite_equations( + DaeEquationPartition::DiscreteReal, + &mut dae.discrete.real_updates, + )?; + self.try_rewrite_equations( + DaeEquationPartition::DiscreteValued, + &mut dae.discrete.valued_updates, + )?; + self.try_rewrite_equations( + DaeEquationPartition::Condition, + &mut dae.conditions.equations, + )?; + self.try_rewrite_expression_slots(&mut dae.conditions.relations)?; + self.try_rewrite_expression_slots(&mut dae.events.synthetic_root_conditions)?; + self.try_rewrite_event_actions(&mut dae.events.event_actions)?; + self.try_rewrite_expression_slots(&mut dae.clocks.constructor_exprs)?; + self.try_rewrite_expression_slots(&mut dae.clocks.triggered_conditions)?; + Ok(()) + } + + fn try_rewrite_equations( + &mut self, + partition: DaeEquationPartition, + equations: &mut [crate::Equation], + ) -> Result<(), Self::Error> { + for equation in equations { + self.try_rewrite_equation(partition, equation)?; + } + Ok(()) + } + + fn try_rewrite_equation( + &mut self, + _partition: DaeEquationPartition, + equation: &mut crate::Equation, + ) -> Result<(), Self::Error> { + equation.rhs = self.try_rewrite_expression(&equation.rhs)?; + Ok(()) + } + + fn try_rewrite_expression_slots( + &mut self, + expressions: &mut [Expression], + ) -> Result<(), Self::Error> { + for expression in expressions { + *expression = self.try_rewrite_expression(expression)?; } + Ok(()) } + + fn try_rewrite_event_actions( + &mut self, + actions: &mut [crate::DaeEventAction], + ) -> Result<(), Self::Error> { + for action in actions { + action.condition = self.try_rewrite_expression(&action.condition)?; + match &mut action.kind { + crate::DaeEventActionKind::Assert { message } + | crate::DaeEventActionKind::Terminate { message } => { + *message = self.try_rewrite_expression(message)?; + } + } + } + Ok(()) + } + + fn try_rewrite_expression(&mut self, expr: &Expression) -> Result; } pub enum StatementScope<'a> { diff --git a/crates/rumoca-ir-dae/tests/golden/modelica_blocks_examples_logical_network1.dae.json b/crates/rumoca-ir-dae/tests/golden/modelica_blocks_examples_logical_network1.dae.json index c4112065d..310044dbb 100644 --- a/crates/rumoca-ir-dae/tests/golden/modelica_blocks_examples_logical_network1.dae.json +++ b/crates/rumoca-ir-dae/tests/golden/modelica_blocks_examples_logical_network1.dae.json @@ -1,5 +1,5 @@ { - "schema_version": 6, + "schema_version": 7, "x": {}, "y": {}, "u": {}, @@ -155,6 +155,7 @@ "structured_equations": [], "initial_equations": [], "initial_structured_equations": [], + "initial_equation_provenance": [], "f_z": [ { "lhs": { @@ -496,7 +497,14 @@ ], "synthetic_root_conditions": [], "scheduled_time_events": [ - 0.1 + { + "time": 0.1, + "source_span": { + "source": 6737700698879693057, + "start": 10, + "end": 20 + } + } ], "scheduled_root_conditions": [], "event_actions": [], @@ -521,9 +529,12 @@ "is_partial": false, "class_type": "Model", "variable_starts": {}, + "nonnumeric_variable_names": [], "interface_flow_count": 0, + "stream_interface_equation_count": 0, "overconstrained_interface_count": 0, "oc_break_edge_scalar_count": 0, + "overconstrained_root_gauge_count": 0, "model_description": "Schema golden derived from Modelica.Blocks.Examples.LogicalNetwork1 surface: Boolean network with timed trigger", "symbol_ancestry": {} } diff --git a/crates/rumoca-ir-flat/src/lib.rs b/crates/rumoca-ir-flat/src/lib.rs index 6f374a92b..e3c7fe1a8 100644 --- a/crates/rumoca-ir-flat/src/lib.rs +++ b/crates/rumoca-ir-flat/src/lib.rs @@ -154,6 +154,14 @@ pub struct Model { /// this correction tracks how many excess equation scalars exist. #[serde(default)] pub oc_break_edge_scalar_count: usize, + /// Scalar count of outside stream connector equations (MLS §15.1). + /// + /// Stream connector equations are structural `inStream`/`actualStream` + /// equations, not ordinary connection equalities. Flatten tracks their + /// count separately so DAE balance can account for them without generating + /// incorrect potential-equation aliases for stream variables. + #[serde(default)] + pub stream_interface_equation_count: usize, /// Enumeration literal ordinal map (MLS §4.9.5, 1-based ordinals). /// /// Keys are canonical literal paths (e.g. diff --git a/crates/rumoca-ir-solve/src/compute_block_tests.rs b/crates/rumoca-ir-solve/src/compute_block_tests.rs index eb5f4f08a..347fef3d7 100644 --- a/crates/rumoca-ir-solve/src/compute_block_tests.rs +++ b/crates/rumoca-ir-solve/src/compute_block_tests.rs @@ -77,6 +77,7 @@ fn linsolve_node() -> ComputeNode { rhs_start: 1, n: 1, next_reg: 2, + output_indices: Vec::new(), metadata: TensorNodeMetadata::default(), span: Span::DUMMY, } diff --git a/crates/rumoca-ir-solve/src/compute_block_validation.rs b/crates/rumoca-ir-solve/src/compute_block_validation.rs new file mode 100644 index 000000000..11671cda5 --- /dev/null +++ b/crates/rumoca-ir-solve/src/compute_block_validation.rs @@ -0,0 +1,337 @@ +//! ComputeBlock output ownership and shape-contract validation. + +use std::collections::HashSet; + +use super::{ + ComputeNode, SolveProblemShapeContractError, Span, StructuredIndexDomain, TensorOutputMap, + TensorOutputMapError, +}; + +pub(super) fn tensor_output_count_for_node( + context: &'static str, + node_index: usize, + node: &ComputeNode, + domain: &StructuredIndexDomain, + output_map: &TensorOutputMap, +) -> Result { + let (dimension, span) = match node { + ComputeNode::Map { span, .. } => ("Map", *span), + ComputeNode::AffineStencil { span, .. } => ("AffineStencil", *span), + ComputeNode::ScalarPrograms(_) + | ComputeNode::MatMul { .. } + | ComputeNode::LinSolve { .. } => unreachable!("tensor output count requires tensor node"), + }; + output_map + .output_count(domain) + .map_err(|error| tensor_output_map_error(context, node_index, dimension, error, span)) +} + +pub(super) fn tensor_output_map_error( + context: &'static str, + node_index: usize, + dimension: &'static str, + error: TensorOutputMapError, + span: Span, +) -> SolveProblemShapeContractError { + match error { + TensorOutputMapError::Dimension { + output_dimension, + domain_rank, + } => SolveProblemShapeContractError::TensorOutputMapDimension { + context: context.to_string(), + node_index, + dimension, + output_dimension, + domain_rank, + span, + }, + TensorOutputMapError::StructuredIndexDomain { error } => { + SolveProblemShapeContractError::StructuredIndexDomain { + context: context.to_string(), + node_index, + dimension, + error, + span, + } + } + TensorOutputMapError::NegativeIndex { value } => { + SolveProblemShapeContractError::TensorOutputMapNegativeIndex { + context: context.to_string(), + node_index, + dimension, + value, + span, + } + } + TensorOutputMapError::OutputIndexOverflow => { + output_index_overflow(context, node_index, Some(span)) + } + } +} + +pub(super) fn output_index_overflow( + context: impl Into, + node_index: usize, + span: Option, +) -> SolveProblemShapeContractError { + SolveProblemShapeContractError::OutputIndexOverflow { + context: context.into(), + node_index, + span, + } +} + +pub(super) fn compute_node_output_cursor( + context: &str, + node_index: usize, + output_cursor: usize, + output_count: usize, + output_indices: &[usize], + span: Span, +) -> Result { + if output_indices.is_empty() { + return output_cursor + .checked_add(output_count) + .ok_or_else(|| output_index_overflow(context, node_index, Some(span))); + } + let Some(max_index) = output_indices.iter().copied().max() else { + return Ok(output_cursor); + }; + let next = max_index + .checked_add(1) + .ok_or_else(|| output_index_overflow(context, node_index, Some(span)))?; + Ok(output_cursor.max(next)) +} + +fn validate_linsolve_output_indices( + context: &str, + node_index: usize, + n: usize, + output_indices: &[usize], + span: Span, +) -> Result<(), SolveProblemShapeContractError> { + if !output_indices.is_empty() && output_indices.len() != n { + return Err( + SolveProblemShapeContractError::LinSolveOutputIndexMismatch { + context: context.to_string(), + node_index, + components: n, + output_indices: output_indices.len(), + span, + }, + ); + } + let mut seen = HashSet::with_capacity(output_indices.len()); + for output_index in output_indices { + if !seen.insert(*output_index) { + return Err( + SolveProblemShapeContractError::LinSolveDuplicateOutputIndex { + context: context.to_string(), + node_index, + output_index: *output_index, + span, + }, + ); + } + } + Ok(()) +} + +impl ComputeNode { + pub fn validate_shape_contract( + &self, + context: &str, + node_index: usize, + ) -> Result<(), SolveProblemShapeContractError> { + match self { + ComputeNode::ScalarPrograms(block) => { + block + .validate_shape_contract(context) + .map_err(|err| match err { + SolveProblemShapeContractError::ScalarProgramSpanMismatch { + programs, + spans, + .. + } => SolveProblemShapeContractError::ScalarProgramSpanMismatch { + context: context.to_string(), + node_index, + programs, + spans, + span: block.first_program_span(), + }, + SolveProblemShapeContractError::ScalarProgramOutputIndexMismatch { + programs, + output_indices, + .. + } => SolveProblemShapeContractError::ScalarProgramOutputIndexMismatch { + context: context.to_string(), + node_index, + programs, + output_indices, + span: block.first_program_span(), + }, + other => other, + })?; + } + ComputeNode::MatMul { m, k, n, span, .. } => { + if *m == 0 || *k == 0 || *n == 0 { + return Err(SolveProblemShapeContractError::ZeroTensorDimension { + context: context.to_string(), + node_index, + dimension: "MatMul", + span: *span, + }); + } + } + ComputeNode::LinSolve { + n, + output_indices, + span, + .. + } => { + if *n == 0 { + return Err(SolveProblemShapeContractError::ZeroTensorDimension { + context: context.to_string(), + node_index, + dimension: "LinSolve", + span: *span, + }); + } + validate_linsolve_output_indices(context, node_index, *n, output_indices, *span)?; + } + ComputeNode::Map { + domain, + output_map, + span, + .. + } => { + let count = validate_tensor_domain(context, node_index, "Map", domain, *span)?; + if count == 0 { + return Err(SolveProblemShapeContractError::ZeroTensorDimension { + context: context.to_string(), + node_index, + dimension: "Map", + span: *span, + }); + } + validate_tensor_output_map(context, node_index, "Map", domain, output_map, *span)?; + } + ComputeNode::AffineStencil { + domain, + output_map, + span, + .. + } => { + let count = + validate_tensor_domain(context, node_index, "AffineStencil", domain, *span)?; + if count == 0 { + return Err(SolveProblemShapeContractError::ZeroTensorDimension { + context: context.to_string(), + node_index, + dimension: "AffineStencil", + span: *span, + }); + } + validate_tensor_output_map( + context, + node_index, + "AffineStencil", + domain, + output_map, + *span, + )?; + } + } + Ok(()) + } +} + +fn validate_tensor_domain( + context: &str, + node_index: usize, + dimension: &'static str, + domain: &StructuredIndexDomain, + span: Span, +) -> Result { + domain.validate().map_err( + |err| SolveProblemShapeContractError::StructuredIndexDomain { + context: context.to_string(), + node_index, + dimension, + error: err, + span, + }, + ) +} + +fn validate_tensor_output_map( + context: &str, + node_index: usize, + dimension: &'static str, + domain: &StructuredIndexDomain, + output_map: &TensorOutputMap, + span: Span, +) -> Result<(), SolveProblemShapeContractError> { + for term in &output_map.strides { + if term.dimension >= domain.binders.len() { + return Err(SolveProblemShapeContractError::TensorOutputMapDimension { + context: context.to_string(), + node_index, + dimension, + output_dimension: term.dimension, + domain_rank: domain.binders.len(), + span, + }); + } + } + if domain + .index_tuples() + .map_err( + |error| SolveProblemShapeContractError::StructuredIndexDomain { + context: context.to_string(), + node_index, + dimension, + error, + span, + }, + )? + .is_empty() + { + return Ok(()); + } + output_map.output_indices(domain).map_err(|err| match err { + TensorOutputMapError::Dimension { + output_dimension, + domain_rank, + } => SolveProblemShapeContractError::TensorOutputMapDimension { + context: context.to_string(), + node_index, + dimension, + output_dimension, + domain_rank, + span, + }, + TensorOutputMapError::StructuredIndexDomain { error } => { + SolveProblemShapeContractError::StructuredIndexDomain { + context: context.to_string(), + node_index, + dimension, + error, + span, + } + } + TensorOutputMapError::NegativeIndex { value } => { + SolveProblemShapeContractError::TensorOutputMapNegativeIndex { + context: context.to_string(), + node_index, + dimension, + value, + span, + } + } + TensorOutputMapError::OutputIndexOverflow => { + output_index_overflow(context, node_index, Some(span)) + } + })?; + Ok(()) +} diff --git a/crates/rumoca-ir-solve/src/initialization_validation.rs b/crates/rumoca-ir-solve/src/initialization_validation.rs new file mode 100644 index 000000000..02a3840fc --- /dev/null +++ b/crates/rumoca-ir-solve/src/initialization_validation.rs @@ -0,0 +1,302 @@ +use super::*; + +#[derive(Clone, Copy, Debug, PartialEq, Eq, Deserialize, Serialize)] +pub struct InitializationTargetRange { + pub start: usize, + pub end: usize, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub span: Option, +} + +/// Count stored initialization rows without expanding tensor output maps. +/// Shape validation independently owns tensor-map placement validity. +pub(super) fn initialization_stored_row_count( + residual: &ComputeBlock, + context: &'static str, +) -> Result { + let mut rows = 0usize; + for (node_index, node) in residual.nodes.iter().enumerate() { + let (count, span) = match node { + ComputeNode::ScalarPrograms(block) => { + (block.stored_output_count(), block.first_source_span()) + } + ComputeNode::Map { domain, span, .. } + | ComputeNode::AffineStencil { domain, span, .. } => ( + domain.scalar_count().map_err(|error| { + SolveProblemShapeContractError::StructuredIndexDomain { + context: context.to_string(), + node_index, + dimension: "stored-row", + error, + span: *span, + } + })?, + Some(*span), + ), + ComputeNode::MatMul { m, n, span, .. } => ( + m.checked_mul(*n) + .ok_or_else(|| output_index_overflow(context, node_index, Some(*span)))?, + Some(*span), + ), + ComputeNode::LinSolve { n, span, .. } => (*n, Some(*span)), + }; + rows = rows + .checked_add(count) + .ok_or_else(|| output_index_overflow(context, node_index, span))?; + } + Ok(rows) +} + +pub(super) fn validate_initialization_direct_families( + initialization: &InitializationSolveSystem, + y_upper_bound: usize, + residual_row_count: usize, +) -> Result<(), SolveProblemShapeContractError> { + if initialization.direct_families.is_empty() { + return validate_initialization_without_direct_families( + initialization, + y_upper_bound, + residual_row_count, + ); + } + validate_count( + "initialization.row_targets.compact", + 0, + initialization.row_targets.len(), + )?; + validate_count( + "initialization.direct_families", + initialization.residual.nodes.len(), + initialization.direct_families.len(), + )?; + let mut covered_nodes = vec![false; initialization.residual.nodes.len()]; + let mut target_ranges = Vec::with_capacity(initialization.direct_families.len()); + for family in &initialization.direct_families { + let Some(covered_node) = covered_nodes.get_mut(family.node_index) else { + return Err(SolveProblemShapeContractError::ZeroTensorDimension { + context: "initialization.direct_families".to_string(), + node_index: family.node_index, + dimension: "direct-family node index outside residual block", + span: family.span, + }); + }; + if std::mem::replace(covered_node, true) { + return Err(SolveProblemShapeContractError::ZeroTensorDimension { + context: "initialization.direct_families".to_string(), + node_index: family.node_index, + dimension: "duplicate direct-family node index", + span: family.span, + }); + } + let target_range = validate_initialization_direct_family(initialization, family)?; + target_ranges.push((target_range, family.node_index, family.span)); + } + target_ranges.sort_unstable_by_key(|(range, _, _)| range.start); + for adjacent in target_ranges.windows(2) { + let [(left, _, _), (right, node_index, span)] = adjacent else { + unreachable!("windows(2) always has two entries") + }; + if right.start < left.end { + return Err(SolveProblemShapeContractError::ZeroTensorDimension { + context: "initialization.direct_families".to_string(), + node_index: *node_index, + dimension: "overlapping direct-family target map", + span: *span, + }); + } + } + let direct_ranges = target_ranges + .into_iter() + .map(|(range, _, span)| InitializationTargetRange { + start: range.start, + end: range.end, + span: Some(span), + }) + .collect::>(); + let required = normalized_ranges( + &initialization.required_target_ranges, + y_upper_bound, + "invalid required target range", + )?; + let complete_required = if y_upper_bound == 0 { + Vec::new() + } else { + vec![InitializationTargetRange { + start: 0, + end: y_upper_bound, + span: required.first().and_then(|range| range.span), + }] + }; + if !same_target_coverage(&required, &complete_required) { + return Err(initialization_range_error_at( + "incomplete required target coverage of the solver Y vector", + required.first().and_then(|range| range.span), + )); + } + let fixed = normalized_ranges( + &initialization.fixed_target_ranges, + y_upper_bound, + "invalid fixed-start target range", + )?; + let mut actual = direct_ranges; + actual.extend(fixed); + let actual = normalized_ranges(&actual, y_upper_bound, "invalid target union range")?; + if !same_target_coverage(&actual, &required) { + return Err(initialization_range_error_at( + "incomplete direct plus fixed-start target union", + actual + .first() + .and_then(|range| range.span) + .or_else(|| required.first().and_then(|range| range.span)), + )); + } + Ok(()) +} + +fn validate_initialization_without_direct_families( + initialization: &InitializationSolveSystem, + y_upper_bound: usize, + residual_row_count: usize, +) -> Result<(), SolveProblemShapeContractError> { + validate_count( + "initialization.row_targets", + residual_row_count, + initialization.row_targets.len(), + )?; + if initialization.residual.is_empty() && initialization.row_targets.is_empty() { + let required = normalized_ranges( + &initialization.required_target_ranges, + y_upper_bound, + "invalid required target range", + )?; + let fixed = normalized_ranges( + &initialization.fixed_target_ranges, + y_upper_bound, + "invalid fixed-start target range", + )?; + if !same_target_coverage(&required, &fixed) { + return Err(initialization_range_error_at( + "incomplete fixed-start target union", + fixed + .first() + .and_then(|range| range.span) + .or_else(|| required.first().and_then(|range| range.span)), + )); + } + } else if !initialization.required_target_ranges.is_empty() + || !initialization.fixed_target_ranges.is_empty() + { + return Err(initialization_range_error( + "target coverage metadata without compact direct families", + )); + } + Ok(()) +} + +fn normalized_ranges( + ranges: &[InitializationTargetRange], + upper_bound: usize, + error: &'static str, +) -> Result, SolveProblemShapeContractError> { + let mut ranges = ranges.to_vec(); + ranges.sort_unstable_by_key(|range| (range.start, range.end)); + let mut normalized: Vec = Vec::with_capacity(ranges.len()); + for range in ranges { + if range.start >= range.end || range.end > upper_bound { + return Err(initialization_range_error_at(error, range.span)); + } + if let Some(last) = normalized.last_mut() { + if range.start < last.end { + return Err(initialization_range_error_at( + "overlapping initialization target ranges", + range.span.or(last.span), + )); + } + if range.start == last.end { + last.end = range.end; + continue; + } + } + normalized.push(range); + } + Ok(normalized) +} + +fn initialization_range_error(reason: &'static str) -> SolveProblemShapeContractError { + initialization_range_error_at(reason, None) +} + +fn initialization_range_error_at( + reason: &'static str, + span: Option, +) -> SolveProblemShapeContractError { + SolveProblemShapeContractError::InitializationTargetCoverage { reason, span } +} + +fn same_target_coverage( + left: &[InitializationTargetRange], + right: &[InitializationTargetRange], +) -> bool { + left.len() == right.len() + && left + .iter() + .zip(right) + .all(|(left, right)| left.start == right.start && left.end == right.end) +} + +fn validate_initialization_direct_family( + initialization: &InitializationSolveSystem, + family: &InitializationDirectFamily, +) -> Result, SolveProblemShapeContractError> { + let Some(node) = initialization.residual.nodes.get(family.node_index) else { + return Err(SolveProblemShapeContractError::ZeroTensorDimension { + context: "initialization.direct_families".to_string(), + node_index: family.node_index, + dimension: "direct-family node index outside residual block", + span: family.span, + }); + }; + let ComputeNode::Map { domain, span, .. } = node else { + return Err(SolveProblemShapeContractError::ZeroTensorDimension { + context: "initialization.direct_families".to_string(), + node_index: family.node_index, + dimension: "non-Map direct family", + span: family.span, + }); + }; + let dense = + TensorOutputMap::dense_contiguous(family.targets.start, domain).map_err(|error| { + tensor_output_map_error( + "initialization.direct_families.targets", + family.node_index, + "Map", + error, + *span, + ) + })?; + if family.targets.strides != dense.strides { + return Err(SolveProblemShapeContractError::ZeroTensorDimension { + context: "initialization.direct_families.targets".to_string(), + node_index: family.node_index, + dimension: "non-contiguous direct-family target map", + span: *span, + }); + } + let count = domain.scalar_count().map_err(|error| { + SolveProblemShapeContractError::StructuredIndexDomain { + context: "initialization.direct_families.targets".to_string(), + node_index: family.node_index, + dimension: "Map", + error, + span: *span, + } + })?; + let end = family.targets.start.checked_add(count).ok_or_else(|| { + output_index_overflow( + "initialization.direct_families.targets", + family.node_index, + Some(*span), + ) + })?; + Ok(family.targets.start..end) +} diff --git a/crates/rumoca-ir-solve/src/layout.rs b/crates/rumoca-ir-solve/src/layout.rs index b8b73fb98..7fbd2b77d 100644 --- a/crates/rumoca-ir-solve/src/layout.rs +++ b/crates/rumoca-ir-solve/src/layout.rs @@ -595,7 +595,7 @@ fn validate_indexed_constant_shape( span, }); }; - if entries.len() != count { + if entries.len() < count { return Err(VarLayoutShapeContractError::ShapeOutOfBounds { variable: name.to_string(), start: 0, @@ -1153,4 +1153,38 @@ mod tests { assert_eq!(layout.validate_shape_contract(), Ok(())); } + + #[test] + fn layout_shape_contract_accepts_extra_indexed_constant_aliases() { + let bindings = IndexMap::from([("table".to_string(), ScalarSlot::Constant(1.0))]); + let shapes = IndexMap::from([("table".to_string(), vec![2])]); + let indexed_bindings = IndexMap::from([( + ComponentReferenceKey::generated("table"), + vec![ + IndexedScalarSlot { + indices: vec![1], + slot: ScalarSlot::Constant(1.0), + }, + IndexedScalarSlot { + indices: vec![2], + slot: ScalarSlot::Constant(2.0), + }, + IndexedScalarSlot { + indices: vec![3], + slot: ScalarSlot::Constant(3.0), + }, + ], + )]); + + let layout = VarLayout::from_parts_with_shapes_and_indexed_bindings( + bindings, + shapes, + indexed_bindings, + 0, + 0, + ) + .expect("extra constant aliases should not invalidate declared shape coverage"); + + assert_eq!(layout.validate_shape_contract(), Ok(())); + } } diff --git a/crates/rumoca-ir-solve/src/lib.rs b/crates/rumoca-ir-solve/src/lib.rs index 782beaee2..a08bb67b8 100644 --- a/crates/rumoca-ir-solve/src/lib.rs +++ b/crates/rumoca-ir-solve/src/lib.rs @@ -4,12 +4,13 @@ //! structural/lowering phases. It must stay free of DAE evaluation and phase //! logic. //! -//! SPEC_0021 file-size exception: Solve IR still defines scalar rows, tensor -//! nodes, validation, and visitor contracts in one facade. split plan: move -//! tensor contracts, validation errors, and visitors into focused modules. +//! The facade defines the wire types while focused modules own layout, linear +//! operations, direct-initialization validation, and visitor contracts. #[cfg(test)] mod compute_block_tests; +mod compute_block_validation; +mod initialization_validation; mod layout; mod linear_op; pub mod visitor; @@ -20,20 +21,30 @@ use rumoca_core::{ }; use serde::{Deserialize, Serialize}; +use compute_block_validation::{ + compute_node_output_cursor, output_index_overflow, tensor_output_count_for_node, + tensor_output_map_error, +}; +pub use initialization_validation::InitializationTargetRange; +use initialization_validation::{ + initialization_stored_row_count, validate_initialization_direct_families, +}; + pub use layout::{ ComponentReferenceKey, ComponentReferenceKeyError, ComponentReferenceKeyErrorKind, ComponentReferenceKeyPart, ComponentReferenceSubscriptKey, IndexedScalarSlot, ScalarSlot, VarLayout, VarLayoutShapeContractError, scalar_slot_p, scalar_slot_y, }; pub use linear_op::{ - BinaryOp, CompareOp, LinearOp, RandomGenerator, Reg, UnaryOp, resolve_indexed_slot, + BinaryOp, CompareOp, ExternalFunctionKind, LinearOp, RandomGenerator, Reg, UnaryOp, + resolve_indexed_slot, }; pub use visitor::{ LinearOpSliceKind, SolveVisitor, VisitScope, walk_compute_block, walk_compute_node, walk_scalar_program_block, walk_solve_artifacts, walk_solve_model, walk_solve_problem, }; -pub const SOLVE_SCHEMA_VERSION: u16 = 15; +pub const SOLVE_SCHEMA_VERSION: u16 = 17; pub fn source_span_from_offsets(source: u64, start: usize, end: usize) -> Span { Span::from_offsets(SourceId(source), start, end) @@ -614,7 +625,7 @@ pub enum ComputeNode { span: Span, }, - /// Dense linear solve: A (n×n) * x = b, writes n consecutive output values. + /// Dense linear solve: A (n×n) * x = b. /// /// `setup_ops` evaluates to n*n + n values: /// regs `matrix_start..matrix_start+n*n` = A (row-major) @@ -626,6 +637,12 @@ pub enum ComputeNode { rhs_start: Reg, n: usize, next_reg: Reg, + /// Dense output slots for the solution components. An empty map retains + /// contiguous placement at the current ComputeBlock cursor. + /// A populated map permits a coupled derivative group to retain its + /// native solve even when its state slots are not contiguous. + #[serde(default)] + output_indices: Vec, metadata: TensorNodeMetadata, span: Span, }, @@ -742,10 +759,20 @@ impl ComputeBlock { .checked_add(output_count) .ok_or_else(|| output_index_overflow(context, node_index, Some(*span)))?; } - ComputeNode::LinSolve { n, span, .. } => { - output_cursor = output_cursor - .checked_add(*n) - .ok_or_else(|| output_index_overflow(context, node_index, Some(*span)))?; + ComputeNode::LinSolve { + n, + output_indices, + span, + .. + } => { + output_cursor = compute_node_output_cursor( + context, + node_index, + output_cursor, + *n, + output_indices, + *span, + )?; } } } @@ -804,273 +831,6 @@ impl ComputeBlock { } } -fn tensor_output_count_for_node( - context: &'static str, - node_index: usize, - node: &ComputeNode, - domain: &StructuredIndexDomain, - output_map: &TensorOutputMap, -) -> Result { - let (dimension, span) = match node { - ComputeNode::Map { span, .. } => ("Map", *span), - ComputeNode::AffineStencil { span, .. } => ("AffineStencil", *span), - ComputeNode::ScalarPrograms(_) - | ComputeNode::MatMul { .. } - | ComputeNode::LinSolve { .. } => unreachable!("tensor output count requires tensor node"), - }; - output_map - .output_count(domain) - .map_err(|error| tensor_output_map_error(context, node_index, dimension, error, span)) -} - -fn tensor_output_map_error( - context: &'static str, - node_index: usize, - dimension: &'static str, - error: TensorOutputMapError, - span: Span, -) -> SolveProblemShapeContractError { - match error { - TensorOutputMapError::Dimension { - output_dimension, - domain_rank, - } => SolveProblemShapeContractError::TensorOutputMapDimension { - context: context.to_string(), - node_index, - dimension, - output_dimension, - domain_rank, - span, - }, - TensorOutputMapError::StructuredIndexDomain { error } => { - SolveProblemShapeContractError::StructuredIndexDomain { - context: context.to_string(), - node_index, - dimension, - error, - span, - } - } - TensorOutputMapError::NegativeIndex { value } => { - SolveProblemShapeContractError::TensorOutputMapNegativeIndex { - context: context.to_string(), - node_index, - dimension, - value, - span, - } - } - TensorOutputMapError::OutputIndexOverflow => { - output_index_overflow(context, node_index, Some(span)) - } - } -} - -fn output_index_overflow( - context: impl Into, - node_index: usize, - span: Option, -) -> SolveProblemShapeContractError { - SolveProblemShapeContractError::OutputIndexOverflow { - context: context.into(), - node_index, - span, - } -} - -impl ComputeNode { - pub fn validate_shape_contract( - &self, - context: &str, - node_index: usize, - ) -> Result<(), SolveProblemShapeContractError> { - match self { - ComputeNode::ScalarPrograms(block) => { - block - .validate_shape_contract(context) - .map_err(|err| match err { - SolveProblemShapeContractError::ScalarProgramSpanMismatch { - programs, - spans, - .. - } => SolveProblemShapeContractError::ScalarProgramSpanMismatch { - context: context.to_string(), - node_index, - programs, - spans, - span: block.first_program_span(), - }, - SolveProblemShapeContractError::ScalarProgramOutputIndexMismatch { - programs, - output_indices, - .. - } => SolveProblemShapeContractError::ScalarProgramOutputIndexMismatch { - context: context.to_string(), - node_index, - programs, - output_indices, - span: block.first_program_span(), - }, - other => other, - })?; - } - ComputeNode::MatMul { m, k, n, span, .. } => { - if *m == 0 || *k == 0 || *n == 0 { - return Err(SolveProblemShapeContractError::ZeroTensorDimension { - context: context.to_string(), - node_index, - dimension: "MatMul", - span: *span, - }); - } - } - ComputeNode::LinSolve { n, span, .. } => { - if *n == 0 { - return Err(SolveProblemShapeContractError::ZeroTensorDimension { - context: context.to_string(), - node_index, - dimension: "LinSolve", - span: *span, - }); - } - } - ComputeNode::Map { - domain, - output_map, - span, - .. - } => { - let count = validate_tensor_domain(context, node_index, "Map", domain, *span)?; - if count == 0 { - return Err(SolveProblemShapeContractError::ZeroTensorDimension { - context: context.to_string(), - node_index, - dimension: "Map", - span: *span, - }); - } - validate_tensor_output_map(context, node_index, "Map", domain, output_map, *span)?; - } - ComputeNode::AffineStencil { - domain, - output_map, - span, - .. - } => { - let count = - validate_tensor_domain(context, node_index, "AffineStencil", domain, *span)?; - if count == 0 { - return Err(SolveProblemShapeContractError::ZeroTensorDimension { - context: context.to_string(), - node_index, - dimension: "AffineStencil", - span: *span, - }); - } - validate_tensor_output_map( - context, - node_index, - "AffineStencil", - domain, - output_map, - *span, - )?; - } - } - Ok(()) - } -} - -fn validate_tensor_domain( - context: &str, - node_index: usize, - dimension: &'static str, - domain: &StructuredIndexDomain, - span: Span, -) -> Result { - domain.validate().map_err( - |err| SolveProblemShapeContractError::StructuredIndexDomain { - context: context.to_string(), - node_index, - dimension, - error: err, - span, - }, - ) -} - -fn validate_tensor_output_map( - context: &str, - node_index: usize, - dimension: &'static str, - domain: &StructuredIndexDomain, - output_map: &TensorOutputMap, - span: Span, -) -> Result<(), SolveProblemShapeContractError> { - for term in &output_map.strides { - if term.dimension >= domain.binders.len() { - return Err(SolveProblemShapeContractError::TensorOutputMapDimension { - context: context.to_string(), - node_index, - dimension, - output_dimension: term.dimension, - domain_rank: domain.binders.len(), - span, - }); - } - } - if domain - .index_tuples() - .map_err( - |error| SolveProblemShapeContractError::StructuredIndexDomain { - context: context.to_string(), - node_index, - dimension, - error, - span, - }, - )? - .is_empty() - { - return Ok(()); - } - output_map.output_indices(domain).map_err(|err| match err { - TensorOutputMapError::Dimension { - output_dimension, - domain_rank, - } => SolveProblemShapeContractError::TensorOutputMapDimension { - context: context.to_string(), - node_index, - dimension, - output_dimension, - domain_rank, - span, - }, - TensorOutputMapError::StructuredIndexDomain { error } => { - SolveProblemShapeContractError::StructuredIndexDomain { - context: context.to_string(), - node_index, - dimension, - error, - span, - } - } - TensorOutputMapError::NegativeIndex { value } => { - SolveProblemShapeContractError::TensorOutputMapNegativeIndex { - context: context.to_string(), - node_index, - dimension, - value, - span, - } - } - TensorOutputMapError::OutputIndexOverflow => { - output_index_overflow(context, node_index, Some(span)) - } - })?; - Ok(()) -} - impl Serialize for ComputeBlock { fn serialize(&self, serializer: S) -> Result { #[derive(Serialize)] @@ -1197,7 +957,7 @@ impl<'de> Deserialize<'de> for SolveProblem { ))); } - Ok(Self { + let problem = Self { schema_version: wire.schema_version, layout: wire.layout, solve_layout: wire.solve_layout, @@ -1206,7 +966,11 @@ impl<'de> Deserialize<'de> for SolveProblem { discrete: wire.discrete, events: wire.events, clocks: wire.clocks, - }) + }; + problem + .validate_shape_contract() + .map_err(serde::de::Error::custom)?; + Ok(problem) } } @@ -1266,10 +1030,14 @@ impl SolveProblem { self.initialization .update_rhs .validate_shape_contract("initialization.update_rhs")?; - validate_count( - "initialization.row_targets", - self.initialization.residual.len()?, - self.initialization.row_targets.len(), + let initialization_rows = initialization_stored_row_count( + &self.initialization.residual, + "initialization.residual rows", + )?; + validate_initialization_direct_families( + &self.initialization, + self.layout.y_scalars(), + initialization_rows, )?; validate_count( "initialization.update_targets", @@ -1284,7 +1052,7 @@ impl SolveProblem { validate_projection_plan( "initialization.projection_plan", &self.initialization.projection_plan, - self.initialization.residual.len()?, + initialization_rows, self.solve_layout.solver_scalar_count(), )?; self.discrete @@ -1424,12 +1192,29 @@ pub enum SolveProblemShapeContractError { output_indices: usize, span: Option, }, + LinSolveOutputIndexMismatch { + context: String, + node_index: usize, + components: usize, + output_indices: usize, + span: Span, + }, + LinSolveDuplicateOutputIndex { + context: String, + node_index: usize, + output_index: usize, + span: Span, + }, ScalarProgramCountMismatch { context: &'static str, expected: usize, actual: usize, span: Option, }, + InitializationTargetCoverage { + reason: &'static str, + span: Option, + }, ZeroTensorDimension { context: String, node_index: usize, @@ -1484,10 +1269,13 @@ impl SolveProblemShapeContractError { Self::ScalarProgramSpanMismatch { span, .. } | Self::ScalarProgramOutputIndexMismatch { span, .. } | Self::ScalarProgramCountMismatch { span, .. } + | Self::InitializationTargetCoverage { span, .. } | Self::OutputIndexOverflow { span, .. } | Self::SolverIndexOutOfBounds { span, .. } | Self::InvalidScheduledRootTiming { span, .. } => *span, - Self::ZeroTensorDimension { span, .. } + Self::LinSolveOutputIndexMismatch { span, .. } + | Self::LinSolveDuplicateOutputIndex { span, .. } + | Self::ZeroTensorDimension { span, .. } | Self::StructuredIndexDomain { span, .. } | Self::TensorOutputMapDimension { span, .. } | Self::TensorOutputMapNegativeIndex { span, .. } => Some(*span), @@ -1526,12 +1314,68 @@ impl std::fmt::Display for SolveProblemShapeContractError { "{context} node {node_index} has {programs} scalar programs but \ {output_indices} output indices" ), + Self::LinSolveOutputIndexMismatch { + context, + node_index, + components, + output_indices, + .. + } => write!( + f, + "{context} node {node_index} has {components} LinSolve components but \ + {output_indices} output indices" + ), + Self::LinSolveDuplicateOutputIndex { + context, + node_index, + output_index, + .. + } => write!( + f, + "{context} node {node_index} assigns LinSolve output index {output_index} more than once" + ), Self::ScalarProgramCountMismatch { context, expected, actual, .. } => write!(f, "{context} expected {expected} rows, got {actual}"), + Self::InitializationTargetCoverage { reason, .. } => { + write!(f, "initialization target coverage is invalid: {reason}") + } + error @ (Self::ZeroTensorDimension { .. } + | Self::StructuredIndexDomain { .. } + | Self::TensorOutputMapDimension { .. } + | Self::TensorOutputMapNegativeIndex { .. }) => error.fmt_tensor_error(f), + Self::OutputIndexOverflow { + context, + node_index, + .. + } => write!( + f, + "{context} node {node_index} output index arithmetic overflowed" + ), + Self::SolverIndexOutOfBounds { + context, + index, + upper_bound, + .. + } => write!( + f, + "{context} references solver index {index}, but upper bound is {upper_bound}" + ), + Self::InvalidScheduledRootTiming { + context, + root_index, + .. + } => write!(f, "{context} root {root_index} has invalid periodic timing"), + } + } +} + +impl SolveProblemShapeContractError { + fn fmt_tensor_error(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { Self::ZeroTensorDimension { context, node_index, @@ -1573,28 +1417,7 @@ impl std::fmt::Display for SolveProblemShapeContractError { f, "{context} node {node_index} {dimension} output map produced negative output index {value}" ), - Self::OutputIndexOverflow { - context, - node_index, - .. - } => write!( - f, - "{context} node {node_index} output index arithmetic overflowed" - ), - Self::SolverIndexOutOfBounds { - context, - index, - upper_bound, - .. - } => write!( - f, - "{context} references solver index {index}, but upper bound is {upper_bound}" - ), - Self::InvalidScheduledRootTiming { - context, - root_index, - .. - } => write!(f, "{context} root {root_index} has invalid periodic timing"), + _ => unreachable!("only tensor errors are delegated to fmt_tensor_error"), } } } @@ -1661,6 +1484,16 @@ pub struct ContinuousSolveArtifacts { pub struct InitializationSolveSystem { pub residual: ComputeBlock, pub row_targets: Vec>, + /// Compact, fully-proven direct initial assignments. Unlike `row_targets`, + /// this does not create one owned record per scalar initial row. + #[serde(default)] + pub direct_families: Vec, + /// Complete solver-Y coverage required by this initialization artifact. + #[serde(default)] + pub required_target_ranges: Vec, + /// Required ranges already satisfied by declared fixed starts. + #[serde(default)] + pub fixed_target_ranges: Vec, pub projection_indices: Vec, #[serde(default)] pub projection_plan: AlgebraicProjectionPlan, @@ -1670,6 +1503,17 @@ pub struct InitializationSolveSystem { pub update_targets: Vec, } +#[derive(Clone, Debug, Deserialize, Serialize)] +pub struct InitializationDirectFamily { + /// Index of the sole residual `Map` owner in `initialization.residual`. + pub node_index: usize, + /// Destination Y slot for every residual element, in the same domain. + pub targets: TensorOutputMap, + /// `+1` for `target - rhs`, `-1` for `rhs - target`. + pub residual_sign: i8, + pub span: rumoca_core::Span, +} + #[derive(Clone, Debug, Default, Deserialize, Serialize)] pub struct DiscreteSolveSystem { pub runtime_assignment_rhs: ScalarProgramBlock, diff --git a/crates/rumoca-ir-solve/src/linear_op.rs b/crates/rumoca-ir-solve/src/linear_op.rs index 96016eb73..9bc62d7bd 100644 --- a/crates/rumoca-ir-solve/src/linear_op.rs +++ b/crates/rumoca-ir-solve/src/linear_op.rs @@ -146,6 +146,18 @@ pub enum RandomGenerator { Xorshift1024Star, } +/// Native external function families preserved in Solve IR. +/// +/// These are explicit runtime dependencies. Interpreters that do not provide +/// the native runtime must fail closed instead of substituting constants. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +pub enum ExternalFunctionKind { + BuildingsEnergyPlusSpawnExternalObject, + BuildingsEnergyPlusInitialize, + BuildingsEnergyPlusGetParameters, + BuildingsEnergyPlusExchange, +} + /// Flat linear operation stream (no strings, no dynamic dispatch). #[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)] pub enum LinearOp { @@ -280,6 +292,16 @@ pub enum LinearOp { imax: Reg, call_site: u64, }, + /// Native external function call. `args[..arg_count]` are initialized + /// scalar argument registers; non-scalar native handles are represented by + /// the function kind and must be resolved by the external runtime bridge. + ExternalCall { + dst: Reg, + function: ExternalFunctionKind, + args: [Reg; 8], + arg_count: usize, + output_index: usize, + }, Unary { dst: Reg, op: UnaryOp, @@ -353,6 +375,7 @@ impl LinearOp { Self::ImpureRandomInit { .. } => "ImpureRandomInit", Self::ImpureRandom { .. } => "ImpureRandom", Self::ImpureRandomInteger { .. } => "ImpureRandomInteger", + Self::ExternalCall { .. } => "ExternalCall", Self::Unary { .. } => "Unary", Self::Binary { .. } => "Binary", Self::Compare { .. } => "Compare", @@ -382,6 +405,7 @@ impl LinearOp { | Self::ImpureRandomInit { dst, .. } | Self::ImpureRandom { dst, .. } | Self::ImpureRandomInteger { dst, .. } + | Self::ExternalCall { dst, .. } | Self::Unary { dst, .. } | Self::Binary { dst, .. } | Self::Compare { dst, .. } @@ -393,7 +417,7 @@ impl LinearOp { #[cfg(test)] mod tests { - use super::{BinaryOp, CompareOp, LinearOp, UnaryOp}; + use super::{BinaryOp, CompareOp, ExternalFunctionKind, LinearOp, UnaryOp}; #[test] fn compare_op_equality_is_exact_not_epsilon_based() { @@ -414,6 +438,16 @@ mod tests { }; assert_eq!(op.kind_name(), "TableNextEvent"); + + let op = LinearOp::ExternalCall { + dst: 0, + function: ExternalFunctionKind::BuildingsEnergyPlusExchange, + args: [0; 8], + arg_count: 0, + output_index: 0, + }; + + assert_eq!(op.kind_name(), "ExternalCall"); } #[test] diff --git a/crates/rumoca-ir-solve/src/tests.rs b/crates/rumoca-ir-solve/src/tests.rs index ea1fae4c9..e9c686fa8 100644 --- a/crates/rumoca-ir-solve/src/tests.rs +++ b/crates/rumoca-ir-solve/src/tests.rs @@ -17,6 +17,274 @@ fn test_tensor_domain(count: usize) -> StructuredIndexDomain { } } +fn direct_initialization_family( + node_index: usize, + target_start: usize, + target_strides: Vec, +) -> (ComputeNode, InitializationDirectFamily) { + let domain = test_tensor_domain(3); + let node = ComputeNode::Map { + output_map: TensorOutputMap::dense_contiguous(node_index * 3, &domain) + .expect("dense residual map"), + domain, + base_ops: vec![ + LinearOp::Const { dst: 0, value: 0.0 }, + LinearOp::StoreOutput { src: 0 }, + ], + load_strides: Vec::new(), + const_strides: Vec::new(), + metadata: TensorNodeMetadata::default(), + span: fixture_span(), + }; + let family = InitializationDirectFamily { + node_index, + targets: TensorOutputMap { + start: target_start, + strides: target_strides, + }, + residual_sign: 1, + span: fixture_span(), + }; + (node, family) +} + +#[test] +fn compact_initialization_validation_rejects_noncontiguous_target_strides() { + let (node, family) = direct_initialization_family( + 0, + 0, + vec![AffineStencilIndexStrideTerm { + dimension: 0, + stride: 2, + }], + ); + let initialization = InitializationSolveSystem { + residual: ComputeBlock { nodes: vec![node] }, + direct_families: vec![family], + ..Default::default() + }; + + let error = validate_initialization_direct_families(&initialization, 6, 3) + .expect_err("sparse compact target maps must fail closed"); + assert!(error.to_string().contains("non-contiguous")); +} + +#[test] +fn compact_initialization_validation_rejects_negative_target_strides() { + let (node, family) = direct_initialization_family( + 0, + 2, + vec![AffineStencilIndexStrideTerm { + dimension: 0, + stride: -1, + }], + ); + let initialization = InitializationSolveSystem { + residual: ComputeBlock { nodes: vec![node] }, + direct_families: vec![family], + required_target_ranges: vec![InitializationTargetRange { + start: 0, + end: 3, + span: None, + }], + ..Default::default() + }; + + let error = validate_initialization_direct_families(&initialization, 3, 3) + .expect_err("descending target maps must fail closed"); + assert!(error.to_string().contains("non-contiguous")); +} + +#[test] +fn compact_initialization_validation_rejects_overlapping_affine_ranges() { + let dense = vec![AffineStencilIndexStrideTerm { + dimension: 0, + stride: 1, + }]; + let (first_node, first_family) = direct_initialization_family(0, 0, dense.clone()); + let (second_node, second_family) = direct_initialization_family(1, 2, dense); + let initialization = InitializationSolveSystem { + residual: ComputeBlock { + nodes: vec![first_node, second_node], + }, + direct_families: vec![first_family, second_family], + ..Default::default() + }; + + let error = validate_initialization_direct_families(&initialization, 6, 6) + .expect_err("overlapping compact target ranges must fail closed"); + assert!(error.to_string().contains("overlapping")); +} + +#[test] +fn compact_initialization_validation_rejects_direct_fixed_overlap() { + let mut initialization = complete_compact_initialization(); + initialization.fixed_target_ranges = vec![InitializationTargetRange { + start: 1, + end: 2, + span: Some(fixture_span()), + }]; + + let error = validate_initialization_direct_families(&initialization, 3, 3) + .expect_err("direct and fixed-start target ownership must not overlap"); + assert!(error.to_string().contains("overlap")); + assert_eq!(error.source_span(), Some(fixture_span())); +} + +#[test] +fn compact_initialization_validation_rejects_fixed_fixed_overlap() { + let initialization = InitializationSolveSystem { + required_target_ranges: vec![InitializationTargetRange { + start: 0, + end: 3, + span: None, + }], + fixed_target_ranges: vec![ + InitializationTargetRange { + start: 0, + end: 2, + span: None, + }, + InitializationTargetRange { + start: 1, + end: 3, + span: Some(fixture_span()), + }, + ], + ..Default::default() + }; + + let error = validate_initialization_direct_families(&initialization, 3, 0) + .expect_err("fixed-start target ownership must not overlap"); + assert!(error.to_string().contains("overlap")); + assert_eq!(error.source_span(), Some(fixture_span())); +} + +#[test] +fn compact_initialization_validation_merges_adjacent_fixed_ranges() { + let initialization = InitializationSolveSystem { + required_target_ranges: vec![InitializationTargetRange { + start: 0, + end: 3, + span: None, + }], + fixed_target_ranges: vec![ + InitializationTargetRange { + start: 0, + end: 1, + span: Some(fixture_span()), + }, + InitializationTargetRange { + start: 1, + end: 3, + span: Some(fixture_span()), + }, + ], + ..Default::default() + }; + + validate_initialization_direct_families(&initialization, 3, 0) + .expect("adjacent target ranges are one exact partition"); +} + +fn complete_compact_initialization() -> InitializationSolveSystem { + let (node, family) = direct_initialization_family( + 0, + 0, + vec![AffineStencilIndexStrideTerm { + dimension: 0, + stride: 1, + }], + ); + InitializationSolveSystem { + residual: ComputeBlock { nodes: vec![node] }, + direct_families: vec![family], + required_target_ranges: vec![InitializationTargetRange { + start: 0, + end: 3, + span: None, + }], + ..Default::default() + } +} + +#[test] +fn compact_initialization_validation_rejects_partial_required_union() { + let mut initialization = complete_compact_initialization(); + initialization.required_target_ranges[0].end = 4; + let error = validate_initialization_direct_families(&initialization, 4, 3) + .expect_err("hand-built partial target union must fail closed"); + assert!(error.to_string().contains("incomplete")); +} + +#[test] +fn compact_initialization_json_rejects_partial_required_union() { + let problem = SolveProblem { + layout: make_layout(&[("x", vec![3])], &[]), + initialization: complete_compact_initialization(), + ..Default::default() + }; + let mut value = serde_json::to_value(problem).expect("serialize compact Solve artifact"); + value["layout"]["y_scalars"] = serde_json::json!(4); + let error = serde_json::from_value::(value) + .expect_err("JSON with a partial target union must fail closed"); + assert!(error.to_string().contains("incomplete")); +} + +#[test] +fn compact_initialization_range_span_survives_json_and_bincode() { + let mut initialization = complete_compact_initialization(); + initialization.required_target_ranges[0].span = Some(fixture_span()); + let problem = SolveProblem { + layout: make_layout(&[("x", vec![3])], &[]), + initialization, + ..Default::default() + }; + let json = serde_json::to_string(&problem).expect("serialize compact Solve artifact"); + let from_json: SolveProblem = + serde_json::from_str(&json).expect("deserialize compact Solve JSON"); + assert_eq!( + from_json.initialization.required_target_ranges[0].span, + Some(fixture_span()) + ); + let bytes = bincode::serialize(&problem).expect("serialize compact Solve bincode"); + let from_bincode: SolveProblem = + bincode::deserialize(&bytes).expect("deserialize compact Solve bincode"); + assert_eq!( + from_bincode.initialization.required_target_ranges[0].span, + Some(fixture_span()) + ); + assert_eq!( + from_bincode.initialization.direct_families[0].span, + fixture_span() + ); +} + +#[test] +fn invalid_initialization_range_reports_span_after_json_and_bincode() { + let range = InitializationTargetRange { + start: 2, + end: 2, + span: Some(fixture_span()), + }; + let json = serde_json::to_string(&range).expect("serialize invalid range JSON"); + let from_json: InitializationTargetRange = + serde_json::from_str(&json).expect("deserialize invalid range JSON"); + let bytes = bincode::serialize(&range).expect("serialize invalid range bincode"); + let from_bincode: InitializationTargetRange = + bincode::deserialize(&bytes).expect("deserialize invalid range bincode"); + + for decoded in [from_json, from_bincode] { + let initialization = InitializationSolveSystem { + required_target_ranges: vec![decoded], + ..Default::default() + }; + let error = validate_initialization_direct_families(&initialization, 2, 0) + .expect_err("empty invalid range must fail closed"); + assert_eq!(error.source_span(), Some(fixture_span())); + } +} + fn fixture_span() -> Span { Span::from_offsets( SourceId::from_source_name("ir_solve_tests_source_44.mo"), @@ -165,6 +433,9 @@ fn representative_derivative_rhs() -> ComputeBlock { fn representative_initialization_system() -> InitializationSolveSystem { InitializationSolveSystem { row_targets: vec![Some(scalar_slot_y(1))], + direct_families: Vec::new(), + required_target_ranges: Vec::new(), + fixed_target_ranges: Vec::new(), residual: ComputeBlock::from_scalar_program_block(ScalarProgramBlock::with_source_span( vec![vec![ LinearOp::Const { dst: 0, value: 0.0 }, @@ -431,6 +702,7 @@ fn serde_roundtrip_linsolve_node() -> ComputeNode { rhs_start: 3, n: 2, next_reg: 4, + output_indices: Vec::new(), metadata: TensorNodeMetadata::default(), span: Span::DUMMY, } @@ -561,6 +833,12 @@ fn solve_problem_json_has_supported_schema_version() { "SolveProblem JSON must carry an explicit schema_version" ); + let mut previous = value.clone(); + previous["schema_version"] = serde_json::json!(SOLVE_SCHEMA_VERSION - 1); + let err = serde_json::from_value::(previous) + .expect_err("previous Solve schema version must fail after initialization IR replacement"); + assert!(err.to_string().contains("unsupported Solve schema_version")); + let mut unsupported = value; unsupported["schema_version"] = serde_json::json!(SOLVE_SCHEMA_VERSION + 1); let err = serde_json::from_value::(unsupported) @@ -621,6 +899,7 @@ fn solve_problem_shape_contract_rejects_zero_tensor_dimension() { rhs_start: 0, n: 0, next_reg: 0, + output_indices: Vec::new(), metadata: TensorNodeMetadata::default(), span: Span::DUMMY, }], @@ -637,6 +916,57 @@ fn solve_problem_shape_contract_rejects_zero_tensor_dimension() { ); } +#[test] +fn solve_problem_shape_contract_rejects_invalid_linsolve_output_indices() { + let mut problem = representative_solve_problem_fixture(); + problem.continuous.derivative_rhs = ComputeBlock { + nodes: vec![ComputeNode::LinSolve { + setup_ops: Vec::new(), + matrix_start: 0, + rhs_start: 0, + n: 2, + next_reg: 0, + output_indices: vec![1], + metadata: TensorNodeMetadata::default(), + span: Span::DUMMY, + }], + }; + + assert!(matches!( + problem.validate_shape_contract(), + Err( + SolveProblemShapeContractError::LinSolveOutputIndexMismatch { + components: 2, + output_indices: 1, + .. + } + ) + )); + + problem.continuous.derivative_rhs = ComputeBlock { + nodes: vec![ComputeNode::LinSolve { + setup_ops: Vec::new(), + matrix_start: 0, + rhs_start: 0, + n: 2, + next_reg: 0, + output_indices: vec![1, 1], + metadata: TensorNodeMetadata::default(), + span: Span::DUMMY, + }], + }; + + assert!(matches!( + problem.validate_shape_contract(), + Err( + SolveProblemShapeContractError::LinSolveDuplicateOutputIndex { + output_index: 1, + .. + } + ) + )); +} + #[test] fn solve_problem_shape_contract_rejects_zero_step_tensor_domain() { let mut problem = representative_solve_problem_fixture(); diff --git a/crates/rumoca-ir-solve/src/visitor.rs b/crates/rumoca-ir-solve/src/visitor.rs index a127a9f9a..8b911008b 100644 --- a/crates/rumoca-ir-solve/src/visitor.rs +++ b/crates/rumoca-ir-solve/src/visitor.rs @@ -479,6 +479,7 @@ mod tests { rhs_start: 1, n: 1, next_reg: 3, + output_indices: Vec::new(), metadata: crate::TensorNodeMetadata::default(), span, } diff --git a/crates/rumoca-ir-solve/tests/golden/representative_solve_problem.solve.json b/crates/rumoca-ir-solve/tests/golden/representative_solve_problem.solve.json index 16d62834f..6f0dfbb92 100644 --- a/crates/rumoca-ir-solve/tests/golden/representative_solve_problem.solve.json +++ b/crates/rumoca-ir-solve/tests/golden/representative_solve_problem.solve.json @@ -1,5 +1,5 @@ { - "schema_version": 15, + "schema_version": 17, "layout": { "bindings": { "x": { @@ -288,6 +288,9 @@ } } ], + "direct_families": [], + "required_target_ranges": [], + "fixed_target_ranges": [], "projection_indices": [], "projection_plan": { "blocks": [] diff --git a/crates/rumoca-phase-codegen/src/codegen/codegen_tests.rs b/crates/rumoca-phase-codegen/src/codegen/codegen_tests.rs index bc22a343b..2e50d5a70 100644 --- a/crates/rumoca-phase-codegen/src/codegen/codegen_tests.rs +++ b/crates/rumoca-phase-codegen/src/codegen/codegen_tests.rs @@ -74,6 +74,12 @@ fn solve_problem_with_one_by_one_matmul_derivative() -> solve::SolveProblem { } pub(super) fn solve_problem_with_two_by_two_linsolve_derivative() -> solve::SolveProblem { + solve_problem_with_two_by_two_linsolve_outputs(Vec::new()) +} + +fn solve_problem_with_two_by_two_linsolve_outputs( + output_indices: Vec, +) -> solve::SolveProblem { let mut problem = solve::SolveProblem::default(); problem.continuous.derivative_rhs = solve::ComputeBlock { nodes: vec![solve::ComputeNode::LinSolve { @@ -92,6 +98,7 @@ pub(super) fn solve_problem_with_two_by_two_linsolve_derivative() -> solve::Solv rhs_start: 4, n: 2, next_reg: 6, + output_indices, metadata: Default::default(), span: fixture_span(), }], @@ -107,6 +114,123 @@ fn test_render_simple_template() { assert!(result.contains("# States: 0")); } +#[test] +fn test_fmi_templates_render_json_function_body_key() { + let dae = dae::Dae::new(); + let mut dae_json = dae_template_json(&dae).expect("dae_template_json should not fail"); + dae_json.as_object_mut().unwrap().insert( + "functions".to_string(), + serde_json::json!({ + "UserFunction": { + "name": "UserFunction", + "inputs": [], + "outputs": [{"name": "y", "dims": [], "default": null}], + "locals": [], + "body": ["Return"], + "is_constructor": false, + "pure": true, + "external": null, + "derivatives": [], + "span": null + } + }), + ); + + for target in ["fmi2", "fmi3"] { + let rendered = render_template_with_dae_json_and_name( + &dae_json, + builtin_template(target, "model.c.jinja"), + "M", + ) + .unwrap_or_else(|err| panic!("{target} template should render function body: {err}")); + + assert!( + rendered.contains("return y;"), + "{target} template should render the body key from JSON-backed functions:\n{rendered}" + ); + } +} + +#[test] +fn test_fmi_function_body_renders_bare_component_reference_expression() { + let dae = dae::Dae::new(); + let mut dae_json = dae_template_json(&dae).expect("dae_template_json should not fail"); + dae_json.as_object_mut().unwrap().insert( + "functions".to_string(), + serde_json::json!({ + "ForwardInput": { + "name": "ForwardInput", + "inputs": [{"name": "x", "dims": [], "default": null}], + "outputs": [{"name": "y", "dims": [], "default": null}], + "locals": [], + "body": [{ + "Assignment": { + "comp": {"local": false, "parts": [{"ident": "y", "subs": []}]}, + "value": {"local": false, "parts": [{"ident": "x", "subs": []}]} + } + }], + "is_constructor": false, + "pure": true, + "external": null, + "derivatives": [], + "span": null + } + }), + ); + + let rendered = render_template_with_dae_json_and_name( + &dae_json, + builtin_template("fmi2", "model.c.jinja"), + "M", + ) + .expect("FMI2 template should render component-reference expression values"); + + assert!( + rendered.contains("y = x;"), + "function body should render bare ComponentReference expressions:\n{rendered}" + ); +} + +#[test] +fn test_fmi_function_body_renders_spanned_expression_wrapper() { + let dae = dae::Dae::new(); + let mut dae_json = dae_template_json(&dae).expect("dae_template_json should not fail"); + dae_json.as_object_mut().unwrap().insert( + "functions".to_string(), + serde_json::json!({ + "ForwardSpannedInput": { + "name": "ForwardSpannedInput", + "inputs": [{"name": "x", "dims": [], "default": null}], + "outputs": [{"name": "y", "dims": [], "default": null}], + "locals": [], + "body": [{ + "Assignment": { + "comp": {"local": false, "parts": [{"ident": "y", "subs": []}]}, + "value": {"expr": {"VarRef": {"name": {"name": "x"}, "subscripts": []}}, "span": null} + } + }], + "is_constructor": false, + "pure": true, + "external": null, + "derivatives": [], + "span": null + } + }), + ); + + let rendered = render_template_with_dae_json_and_name( + &dae_json, + builtin_template("fmi2", "model.c.jinja"), + "M", + ) + .expect("FMI2 template should render spanned expression wrappers"); + + assert!( + rendered.contains("y = x;"), + "function body should render expression wrappers:\n{rendered}" + ); +} + #[test] fn test_record_param_template_skip_uses_type_class_metadata() { let dae_json = serde_json::json!({ @@ -154,6 +278,34 @@ fn test_simulation_template_rejects_external_function_with_stable_diagnostic() { } } +#[test] +fn test_simulation_template_allows_supported_energyplus_external_function() { + let mut dae = dae::Dae::new(); + let mut function = rumoca_core::Function::new( + "Buildings.ThermalZones.EnergyPlus_9_6_0.BaseClasses.initialize", + fixture_span(), + ); + function.add_output(rumoca_core::FunctionParam::new( + "nObj", + "Integer", + fixture_span(), + )); + function.external = Some(rumoca_core::ExternalFunction::default()); + dae.symbols.functions.insert( + "Buildings.ThermalZones.EnergyPlus_9_6_0.BaseClasses.initialize".into(), + function, + ); + + let rendered = render_template_with_name( + &dae, + "FMI 3.0 API {% for name, func in dae.functions | items %}{% if func.external %}{{ name }}{% endif %}{% endfor %}", + "M", + ) + .expect("supported EnergyPlus external runtime function should pass template guard"); + + assert!(rendered.contains("Buildings.ThermalZones.EnergyPlus_9_6_0.BaseClasses.initialize")); +} + #[test] fn test_simulation_template_file_rejects_external_function_with_stable_diagnostic() { let mut dae = dae::Dae::new(); @@ -245,6 +397,20 @@ fn test_solve_template_context_exposes_tensor_nodes_and_scalar_fallback_rows() { assert_eq!(rendered, "1 1 1 true"); } +#[test] +fn scalar_codegen_template_preserves_noncontiguous_linsolve_output_indices() { + let problem = solve_problem_with_two_by_two_linsolve_outputs(vec![0, 2]); + let rendered = render_solve_template_with_name( + &problem, + &solve::SolveArtifacts::default(), + r#"{% for row in solve_blocks.continuous.derivative_rhs.scalar_fallback_rows %}out[{{ row.output_index }}]={{ row.output_ordinal }};{% endfor %}"#, + "NoncontiguousLinSolveScalarFallback", + ) + .expect("scalar codegen template should render noncontiguous LinSolve fallback rows"); + + assert_eq!(rendered, "out[0]=0;out[2]=1;"); +} + #[test] fn test_c_solve_builtin_target_renders_scalar_fallback_derivative_kernel() { let problem = solve_problem_with_two_by_two_linsolve_derivative(); @@ -502,6 +668,22 @@ fn test_sanitize_filter() { assert_eq!(result, "body_position_x"); } +#[test] +fn test_json_filter() { + let dae = dae::Dae::new(); + let template = r#"{{ 'Model "A"' | json }} {{ ['libm', 'libc'] | json }}"#; + let result = render_template(&dae, template).unwrap(); + assert_eq!(result, r#""Model \"A\"" ["libm","libc"]"#); +} + +#[test] +fn test_sanitize_filter_folds_static_component_subscript_arithmetic() { + let dae = dae::Dae::new(); + let template = "{{ 'zone[(1 + 1)].T' | sanitize }} {{ 'floor3Zones[2 - 1 + 3].T' | sanitize }}"; + let result = render_template(&dae, template).unwrap(); + assert_eq!(result, "zone_2_T floor3Zones_4_T"); +} + #[test] fn test_access_dae_fields() { let dae = dae::Dae::new(); @@ -673,6 +855,71 @@ fn dae_template_json_rejects_source_ref_scalar_count_overflow() { ); } +#[test] +fn test_array_scalar_name_preserves_modelica_multidimensional_subscripts() { + assert_eq!( + render_array_scalar_name("floor_internal_gain", &[3, 5], 1).unwrap(), + "floor_internal_gain[1,1]" + ); + assert_eq!( + render_array_scalar_name("floor_internal_gain", &[3, 5], 5).unwrap(), + "floor_internal_gain[1,5]" + ); + assert_eq!( + render_array_scalar_name("floor_internal_gain", &[3, 5], 6).unwrap(), + "floor_internal_gain[2,1]" + ); + assert_eq!( + render_array_scalar_name("floor_internal_gain", &[3, 5], 15).unwrap(), + "floor_internal_gain[3,5]" + ); +} + +#[test] +fn test_array_scalar_name_connects_multidimensional_dae_residuals() { + let rhs = rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: Box::new(rumoca_core::Expression::VarRef { + name: "dynamic_gain".into(), + subscripts: Vec::new(), + span: rumoca_core::Span::DUMMY, + }), + rhs: Box::new(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(3.0), + span: rumoca_core::Span::DUMMY, + }), + span: rumoca_core::Span::DUMMY, + }; + let dae_json = serde_json::json!({ + "f_x": [ + { + "lhs": "floor_internal_gain[1,2]", + "rhs": serde_json::to_value(rhs).unwrap() + } + ] + }); + let template = r#" +{% set cfg = {"prefix": "", "power": "pow", "float_literals": false, "subscript_underscore": true} %} +{% set scalar_name = array_scalar_name("floor_internal_gain", [3, 5], 2) %} +{{ scalar_name }} +{{ alg_rhs_for_var(scalar_name, dae.f_x, cfg) }} +"#; + let rendered = render_template_with_dae_json(&dae_json, template).unwrap(); + + assert!( + rendered.contains("floor_internal_gain[1,2]"), + "codegen should query the DAE with Modelica multi-dimensional scalar names:\n{rendered}" + ); + assert!( + rendered.contains("(dynamic_gain + 3.0)"), + "expected multidimensional array residual RHS to connect, got:\n{rendered}" + ); + assert!( + !rendered.contains("WARNING: no equation found for floor_internal_gain[1,2]"), + "codegen should not fall back to warning stubs for multidimensional residuals:\n{rendered}" + ); +} + #[test] fn test_render_expr_function() { let dae = dae::Dae::new(); @@ -923,6 +1170,438 @@ fn test_fmi3_event_indicators_render_from_solver_ir() { ); } +#[test] +fn test_fmi_getters_refresh_outputs_when_dirty() { + let fmi2 = builtin_template("fmi2", "model.c.jinja"); + let get_real = template_section(fmi2, "FMI2_EXPORT fmi2Status fmi2GetReal"); + assert!( + get_real.contains( + "compute_derivatives(m);\n compute_outputs(m);\n m->dirty_values = 0;" + ), + "FMI 2 fmi2GetReal must refresh output storage before reading value references:\n{get_real}" + ); + + let fmi3 = builtin_template("fmi3", "model.c.jinja"); + let get_float64 = template_section(fmi3, "FMI3_Export fmi3Status fmi3GetFloat64"); + assert!( + get_float64.contains( + "compute_derivatives(m);\n compute_outputs(m);\n m->dirty_values = 0;" + ), + "FMI 3 fmi3GetFloat64 must refresh output storage before reading value references:\n{get_float64}" + ); +} + +#[test] +fn test_fmi_real_setters_mark_values_dirty() { + let fmi2 = builtin_template("fmi2", "model.c.jinja"); + let set_real = template_section(fmi2, "FMI2_EXPORT fmi2Status fmi2SetReal"); + assert!( + set_real.contains("m->dirty_values = 1;\n return fmi2OK;"), + "FMI 2 fmi2SetReal must mark cached derivatives and outputs dirty after accepted inputs:\n{set_real}" + ); + + let fmi3 = builtin_template("fmi3", "model.c.jinja"); + let set_float64 = template_section(fmi3, "FMI3_Export fmi3Status fmi3SetFloat64"); + assert!( + set_float64.contains("m->dirty_values = 1;\n return fmi3OK;"), + "FMI 3 fmi3SetFloat64 must mark cached derivatives and outputs dirty after accepted inputs:\n{set_float64}" + ); +} + +#[test] +fn test_fmi_cosimulation_refreshes_discrete_updates_before_derivatives() { + let fmi2 = builtin_template("fmi2", "model.c.jinja"); + let do_step = template_section(fmi2, "FMI2_EXPORT fmi2Status fmi2DoStep"); + assert!( + do_step.contains("m->dirty_values = 1;\n compute_discrete_updates(m);\n compute_derivatives(m);") + && do_step.contains("m->dirty_values = 1;\n compute_discrete_updates(m);\n compute_derivatives(m);"), + "FMI 2 Co-Simulation steps must refresh input-driven discrete gates before derivative evaluation:\n{do_step}" + ); + + let fmi3 = builtin_template("fmi3", "model.c.jinja"); + let rk45_eval = template_section( + fmi3, + "static fmi3Status rk45_eval(ModelInstance* m, double t, const fmi3Float64 x[], fmi3Float64 dxdt[]) {", + ); + assert!( + rk45_eval.contains( + "m->dirty_values = 1;\n compute_discrete_updates(m);\n compute_derivatives(m);" + ), + "FMI 3 Co-Simulation RK derivative evaluations must refresh input-driven discrete gates before derivative evaluation:\n{rk45_eval}" + ); +} + +#[test] +fn test_fmi2_cosimulation_caps_euler_substep_for_stiff_thermal_states() { + let fmi2 = builtin_template("fmi2", "model.c.jinja"); + let do_step = template_section(fmi2, "FMI2_EXPORT fmi2Status fmi2DoStep"); + assert!( + do_step.contains("const double dt_max = fmin(60.0, communicationStepSize / 10.0);"), + "FMI 2 Co-Simulation must cap explicit Euler substeps by physical time, not only by communication-step fraction:\n{do_step}" + ); + assert!( + do_step.contains("fmi2Real x_nominal[N_STATES > 0 ? N_STATES : 1];") + && do_step.contains( + "if (fmi2GetNominalsOfContinuousStates(c, x_nominal, N_STATES) != fmi2OK)" + ) + && do_step.contains("const double rate_limited_dt = 0.5 * scale / abs_derivative;"), + "FMI 2 Co-Simulation must derive a state-rate limited substep from nominal state scale and derivative magnitude:\n{do_step}" + ); + assert!( + do_step.contains( + "if (isfinite(rate_limited_dt) && rate_limited_dt > 0.0 && rate_limited_dt < dt)" + ), + "FMI 2 Co-Simulation must only shrink dt with finite positive rate limits:\n{do_step}" + ); +} + +#[test] +fn test_fmi_solve_y_runtime_cases_do_not_count_zero_length_arrays() { + let mut dae = dae::Dae::new(); + let mut empty = dae::Variable::new("empty".into(), fixture_span()); + empty.dims = vec![0]; + dae.variables.algebraics.insert("empty".into(), empty); + dae.variables.algebraics.insert( + "after_empty".into(), + dae::Variable::new("after_empty".into(), fixture_span()), + ); + let mut dae_json = dae_template_json(&dae).expect("dae_template_json should serialize"); + dae_json.as_object_mut().unwrap().insert( + "solve".to_string(), + serde_json::json!({ + "visible_names": ["empty[1]", "after_empty"] + }), + ); + + for target in ["fmi2", "fmi3"] { + let rendered = render_template_with_dae_json_and_name( + &dae_json, + builtin_template(target, "model.c.jinja"), + "M", + ) + .unwrap(); + + assert!( + rendered.contains("#define N_ALGEBRAICS 1"), + "{target} zero-length algebraic arrays must not contribute to N_ALGEBRAICS:\n{rendered}" + ); + assert!( + rendered.contains("case 1: return m->y[0]; /* after_empty */"), + "{target} solve_y runtime mapping must use the same zero-length array layout as N_ALGEBRAICS:\n{rendered}" + ); + assert!( + rendered.contains("case 1: m->y[0] = value; return; /* after_empty */"), + "{target} solve_y assignment mapping must use the same zero-length array layout as N_ALGEBRAICS:\n{rendered}" + ); + assert!( + !rendered.contains("m->y[1]"), + "{target} zero-length algebraic arrays must not shift later runtime slots out of bounds:\n{rendered}" + ); + } +} + +#[test] +fn test_fmi_templates_prefer_solve_visible_value_rows_before_dae_fallback() { + let mut dae = dae::Dae::new(); + dae.variables.outputs.insert( + "surface".into(), + dae::Variable::new("surface".into(), fixture_span()), + ); + dae.variables.algebraics.insert( + "local_surface".into(), + dae::Variable::new("local_surface".into(), fixture_span()), + ); + let mut dae_json = dae_template_json(&dae).expect("dae_template_json should serialize"); + dae_json.as_object_mut().unwrap().insert( + "solve".to_string(), + serde_json::json!({ + "visible_names": ["surface", "local_surface"], + "visible_value_rows": { + "programs": [[ + {"LoadY": {"dst": 0, "index": 2}}, + {"LoadP": {"dst": 1, "index": 1}}, + {"Binary": {"dst": 2, "op": "Add", "lhs": 0, "rhs": 1}}, + {"StoreOutput": {"src": 2}} + ], [ + {"LoadP": {"dst": 0, "index": 3}}, + {"StoreOutput": {"src": 0}} + ]] + } + }), + ); + + for target in ["fmi2", "fmi3"] { + let rendered = render_template_with_dae_json_and_name( + &dae_json, + builtin_template(target, "model.c.jinja"), + "VisibleRowRegression", + ) + .unwrap(); + + assert!( + rendered.contains("m->w[0] = ((__rumoca_solve_y(m, 2)) + (__rumoca_solve_p(m, 1)));"), + "{target} outputs should use solve visible rows when present:\n{rendered}" + ); + assert!( + rendered.contains("local_surface = __rumoca_solve_p(m, 3);") + && rendered.contains("m->y[0] = local_surface; /* local_surface */"), + "{target} algebraics should use solve visible rows when present:\n{rendered}" + ); + assert!( + !rendered.contains("WARNING: no equation found for surface") + && !rendered.contains("WARNING: no equation found for local_surface"), + "{target} must not fall back to warning zero when solve visible rows are present:\n{rendered}" + ); + } +} + +#[test] +fn test_fmi_algebraic_identity_solve_row_falls_back_to_dae_rhs() { + let mut dae = dae::Dae::new(); + dae.variables.algebraics.insert( + "driven".into(), + dae::Variable::new("driven".into(), fixture_span()), + ); + let mut dae_json = dae_template_json(&dae).expect("dae_template_json should serialize"); + dae_json.as_object_mut().unwrap().insert( + "f_x".to_string(), + serde_json::json!([{ + "lhs": { + "VarRef": { + "name": "driven", + "subscripts": [] + } + }, + "rhs": { + "Literal": { + "value": { + "Real": 7.0 + } + } + } + }]), + ); + dae_json.as_object_mut().unwrap().insert( + "solve".to_string(), + serde_json::json!({ + "visible_names": ["driven"], + "visible_value_rows": { + "programs": [[ + {"LoadY": {"dst": 0, "index": 0}}, + {"StoreOutput": {"src": 0}} + ]] + } + }), + ); + + for target in ["fmi2", "fmi3"] { + let rendered = render_template_with_dae_json_and_name( + &dae_json, + builtin_template(target, "model.c.jinja"), + "M", + ) + .unwrap(); + + assert!( + rendered.contains("driven = 7.0;"), + "{target} identity solve rows must not short-circuit explicit DAE algebraic RHS:\n{rendered}" + ); + assert!( + !rendered.contains("driven = __rumoca_solve_y(m, 0);"), + "{target} identity solve row would preserve the stale algebraic storage value:\n{rendered}" + ); + } +} + +#[test] +fn test_fmi_templates_apply_explicit_state_initial_equations() { + let dae_json = serde_json::json!({ + "f_x": [], + "initial_equations": [{ + "lhs": { + "VarRef": { + "name": "x", + "subscripts": [] + } + }, + "rhs": { + "VarRef": { + "name": "p", + "subscripts": [] + } + } + }], + "x": { + "x": { + "name": "x", + "dims": [], + "start": null, + "unit": null, + "nominal": null, + "min": null, + "max": null, + "description": null + } + }, + "y": {}, + "w": {}, + "u": {}, + "p": { + "p": { + "name": "p", + "dims": [], + "unit": null, + "nominal": null, + "min": null, + "max": null, + "description": null, + "start": { + "Literal": { + "value": { + "Real": 292.15 + } + } + } + } + }, + "z": {}, + "m": {}, + "constants": {}, + "functions": {}, + "symbol_refs": ["x", "p"], + "symbol_aliases": [], + "enum_literal_ordinals": {}, + "enum_type_names": [] + }); + + for target in ["fmi2", "fmi3"] { + let rendered = render_template_with_dae_json_and_name( + &dae_json, + builtin_template(target, "model.c.jinja"), + "M", + ) + .unwrap(); + assert!( + rendered.contains("m->x[0] = p; /* initial equation: x */"), + "{target} should assign explicit state initial equations to state storage:\n{rendered}" + ); + + let exit_initialization = rendered + .split(if target == "fmi2" { + "FMI2_EXPORT fmi2Status fmi2ExitInitializationMode" + } else { + "FMI3_Export fmi3Status fmi3ExitInitializationMode" + }) + .nth(1) + .expect("template should define exit initialization"); + let initial_update_call = if target == "fmi2" { + "apply_initial_equations(m);" + } else { + "compute_initial_updates(m);" + }; + assert!( + exit_initialization.contains(initial_update_call), + "{target} should apply explicit state initial equations before initial derivatives:\n{exit_initialization}" + ); + assert!( + exit_initialization.find(initial_update_call).unwrap() + < exit_initialization.find("compute_derivatives(m);").unwrap(), + "{target} should apply explicit state initial equations before initial derivatives:\n{exit_initialization}" + ); + } +} + +#[test] +fn test_fmi_templates_do_not_emit_runtime_field_name_enum_macros() { + let dae_json = serde_json::json!({ + "symbol_refs": ["y"], + "symbol_aliases": [], + "enum_literal_ordinals": {"y": 1}, + "enum_type_names": [], + "x": {}, + "y": {}, + "u": {}, + "w": {}, + "p": {}, + "z": {}, + "m": {}, + "constants": {}, + "f_x": [], + "f_z": [], + "f_m": [], + "relation": [], + "scheduled_time_events": [], + "functions": {}, + "metadata": {} + }); + + for target in ["fmi2", "fmi3"] { + let rendered = render_template_with_dae_json_and_name( + &dae_json, + builtin_template(target, "model.c.jinja"), + "RuntimeFieldMacroRegression", + ) + .unwrap(); + + assert!( + !rendered.contains("#define y 1"), + "{target} enum literal macro must not collide with ModelInstance.y" + ); + } +} + +#[test] +fn test_fmi_templates_emit_source_reference_alias_macros() { + let dae_json = serde_json::json!({ + "symbol_refs": ["controlSemantics.initialOverrideActive[1]"], + "symbol_aliases": [], + "enum_literal_ordinals": {}, + "enum_type_names": [], + "x": {}, + "y": {}, + "u": {}, + "w": {}, + "p": {}, + "z": {}, + "m": {}, + "constants": {}, + "f_x": [], + "f_z": [], + "f_m": [], + "relation": [], + "scheduled_time_events": [], + "functions": {}, + "metadata": {} + }); + + for target in ["fmi2", "fmi3"] { + let rendered = render_template_with_dae_json_and_name( + &dae_json, + builtin_template(target, "model.c.jinja"), + "SourceReferenceAliasRegression", + ) + .unwrap(); + + assert!( + rendered.contains( + "#define controlSemantics_initialOverrideActive_1 initialOverrideActive_1" + ), + "{target} template must bridge sanitized source refs to allocated local symbols:\n{rendered}" + ); + } +} + +fn template_section(template: &str, marker: &str) -> String { + let section = template + .split(marker) + .nth(1) + .unwrap_or_else(|| panic!("template should define {marker}")) + .split("\nFMI") + .next() + .expect("template section should be present"); + normalize_newlines(section) +} + #[test] fn test_fmi3_derivative_api_renders_from_solver_ad_ir() { let dae = dae::Dae::new(); @@ -991,6 +1670,7 @@ fn test_fmi3_scalar_blt_projection_renders_from_solve_ir() { problem.solve_layout.state_scalar_count = 1; problem.solve_layout.algebraic_scalar_count = 1; problem.continuous.implicit_rhs = solve::ComputeBlock::from_scalar_program_block(implicit); + problem.continuous.implicit_row_targets = vec![None, Some(solve::scalar_slot_y(1))]; problem.continuous.algebraic_projection_plan = solve::AlgebraicProjectionPlan { blocks: vec![solve::AlgebraicProjectionBlock { rows: vec![1], @@ -1283,6 +1963,201 @@ fn test_render_expr_uses_template_symbol_map_for_indexed_refs() { assert_eq!(rendered, "leg_f_b_2_1"); } +#[test] +fn test_render_expr_uses_symbol_map_for_structured_indexed_var_ref() { + let expr = serde_json::json!({ + "VarRef": { + "name": { + "name": "control.initial_active", + "component_ref": { + "local": false, + "parts": [ + {"ident": "control", "subs": []}, + {"ident": "initial_active", "subs": []} + ], + "def_id": 1 + } + }, + "subscripts": [{"Index": {"value": 1}}] + } + }); + let symbols = serde_json::json!({ + "control.initial_active[1]": "initial_active_1" + }); + let cfg = ExprConfig { + subscript_underscore: true, + symbols: Some(Value::from_serialize(symbols)), + ..ExprConfig::default() + }; + + let rendered = render_expression(&Value::from_serialize(&expr), &cfg).unwrap(); + assert_eq!(rendered, "initial_active_1"); +} + +#[test] +fn test_render_expr_uses_component_ref_when_var_ref_name_string_is_missing() { + let expr = serde_json::json!({ + "VarRef": { + "name": { + "component_ref": { + "local": false, + "parts": [ + {"ident": "system", "subs": []}, + {"ident": "loop", "subs": []}, + {"ident": "pressure", "subs": []} + ], + "def_id": 2 + } + }, + "subscripts": [] + } + }); + let symbols = serde_json::json!({ + "system.loop.pressure": "system_loop_pressure" + }); + let cfg = ExprConfig { + subscript_underscore: true, + symbols: Some(Value::from_serialize(symbols)), + ..ExprConfig::default() + }; + + let rendered = render_expression(&Value::from_serialize(&expr), &cfg).unwrap(); + assert_eq!(rendered, "system_loop_pressure"); +} + +#[test] +fn test_render_component_ref_uses_symbol_map_before_c_bracket_fallback() { + let component_ref = serde_json::json!({ + "parts": [ + {"ident": {"text": "plant"}, "subscripts": []}, + {"ident": {"text": "arr"}, "subscripts": [{"Index": {"value": 0}}]}, + {"ident": {"text": "field"}, "subscripts": []} + ] + }); + let symbols = serde_json::json!({ + "plant.arr[1].field": "plant_arr_1_field", + "plant.arr[2].field": "plant_arr_2_field" + }); + let cfg = ExprConfig { + subscript_underscore: true, + symbols: Some(Value::from_serialize(symbols)), + ..ExprConfig::default() + }; + + let rendered = render_expression(&Value::from_serialize(&component_ref), &cfg).unwrap(); + assert_eq!(rendered, "plant_arr_1_field"); +} + +#[test] +fn test_render_expression_handles_component_reference_wrapper() { + let expr = serde_json::json!({ + "ComponentReference": { + "local": false, + "parts": [ + {"ident": "plant", "subs": []}, + {"ident": "arr", "subs": [{"Index": {"value": 0}}]}, + {"ident": "field", "subs": []} + ], + "def_id": 7 + } + }); + let symbols = serde_json::json!({ + "plant.arr[1].field": "plant_arr_1_field" + }); + let cfg = ExprConfig { + subscript_underscore: true, + symbols: Some(Value::from_serialize(symbols)), + ..ExprConfig::default() + }; + + let rendered = render_expression(&Value::from_serialize(&expr), &cfg).unwrap(); + assert_eq!(rendered, "plant_arr_1_field"); +} + +#[test] +fn test_render_expression_handles_var_name_component_reference() { + let expr = serde_json::json!({ + "name": "plant.arr[1].field", + "component_ref": { + "local": false, + "parts": [ + {"ident": "plant", "subs": []}, + {"ident": "arr", "subs": [{"Index": {"value": 0}}]}, + {"ident": "field", "subs": []} + ], + "def_id": 7 + }, + "def_id": 7 + }); + let symbols = serde_json::json!({ + "plant.arr[1].field": "plant_arr_1_field" + }); + let cfg = ExprConfig { + subscript_underscore: true, + symbols: Some(Value::from_serialize(symbols)), + ..ExprConfig::default() + }; + + let rendered = render_expression(&Value::from_serialize(&expr), &cfg).unwrap(); + assert_eq!(rendered, "plant_arr_1_field"); +} + +#[test] +fn test_render_component_ref_canonicalizes_zero_based_source_without_symbol_map() { + let component_ref = serde_json::json!({ + "parts": [ + {"ident": "plant", "subs": []}, + {"ident": "arr", "subs": [{"Index": {"value": 0}}]}, + {"ident": "field", "subs": []} + ] + }); + let cfg = ExprConfig { + subscript_underscore: true, + ..ExprConfig::default() + }; + + let rendered = render_expression(&Value::from_serialize(&component_ref), &cfg).unwrap(); + assert_eq!(rendered, "plant_arr_1_field"); +} + +#[test] +fn test_render_var_ref_uses_one_based_symbol_for_serialized_component_index() { + let expr = serde_json::json!({ + "VarRef": { + "name": {"name": "device.cells[0].temperature"}, + "subscripts": [] + } + }); + let symbols = serde_json::json!({ + "device.cells[1].temperature": "device_cells_1_temperature" + }); + let cfg = ExprConfig { + subscript_underscore: true, + symbols: Some(Value::from_serialize(symbols)), + ..ExprConfig::default() + }; + + let rendered = render_expression(&Value::from_serialize(&expr), &cfg).unwrap(); + assert_eq!(rendered, "device_cells_1_temperature"); +} + +#[test] +fn test_render_var_ref_canonicalizes_serialized_component_index_without_symbol_map() { + let expr = serde_json::json!({ + "VarRef": { + "name": {"name": "device.cells[0].temperature"}, + "subscripts": [] + } + }); + let cfg = ExprConfig { + subscript_underscore: true, + ..ExprConfig::default() + }; + + let rendered = render_expression(&Value::from_serialize(&expr), &cfg).unwrap(); + assert_eq!(rendered, "device_cells_1_temperature"); +} + #[test] fn test_fmi3_initialize_defaults_uses_allocated_symbols_for_start_aliases() { let mut dae = dae::Dae::new(); diff --git a/crates/rumoca-phase-codegen/src/codegen/codegen_tests/backend_template_tests.rs b/crates/rumoca-phase-codegen/src/codegen/codegen_tests/backend_template_tests.rs index 5449252e0..05acbf9a6 100644 --- a/crates/rumoca-phase-codegen/src/codegen/codegen_tests/backend_template_tests.rs +++ b/crates/rumoca-phase-codegen/src/codegen/codegen_tests/backend_template_tests.rs @@ -1,5 +1,7 @@ use super::*; +mod julia_template_tests; + #[test] fn test_embedded_c_alg_rhs_indexes_common_array_binary_rhs() { let rhs = rumoca_core::Expression::Binary { @@ -52,6 +54,133 @@ fn test_embedded_c_alg_rhs_indexes_common_array_binary_rhs() { ); } +#[test] +fn test_alg_rhs_indexes_structured_residual_equations_for_indexed_var() { + let dae_json = serde_json::json!({ + "f_x": [{ + "lhs": null, + "rhs": { + "Binary": { + "op": "Sub", + "lhs": { + "VarRef": { + "name": { + "name": "network.heat", + "component_ref": { + "local": false, + "parts": [ + {"ident": "network", "subs": []}, + {"ident": "heat", "subs": []} + ], + "def_id": 1 + } + }, + "subscripts": [{"Index": {"value": 2}}] + } + }, + "rhs": { + "Binary": { + "op": "Add", + "lhs": { + "VarRef": { + "name": {"name": "floor_heat"}, + "subscripts": [ + {"Index": {"value": 1}}, + {"Index": {"value": 2}} + ] + } + }, + "rhs": { + "VarRef": { + "name": {"name": "internal_gain"}, + "subscripts": [ + {"Index": {"value": 1}}, + {"Index": {"value": 2}} + ] + } + } + } + } + } + }, + "origin": "top-level model equation", + "scalar_count": 1 + }], + "w": {}, + "y": { + "network.heat[2]": {}, + "floor_heat[1,2]": {}, + "internal_gain[1,2]": {} + }, + "x": {}, + "z": {}, + "m": {}, + "u": {}, + "p": {}, + "constants": {} + }); + let template = r#" +{% set cfg = {"prefix": "", "power": "pow", "if_style": "ternary", "subscript_underscore": true} %} +{{ alg_rhs_for_var_with_dae("network.heat[2]", dae, cfg) }} +"#; + let rendered = render_template_with_dae_json(&dae_json, template).unwrap(); + + assert_eq!( + rendered.trim(), + "(floor_heat_1_2 + internal_gain_1_2)", + "indexed residual equations must be discoverable through the generic algebraic RHS path:\n{rendered}" + ); +} + +#[test] +fn test_c_alg_rhs_projects_indexed_function_array_rhs_before_rendering() { + let dae_json = serde_json::json!({ + "f_x": [{ + "lhs": { + "VarRef": { + "name": {"name": "selector.expr"}, + "subscripts": [{"Index": {"value": 1}}] + } + }, + "rhs": { + "FunctionCall": { + "name": {"name": "linspace"}, + "args": [ + {"Literal": {"value": {"Integer": 0}}}, + {"VarRef": {"name": {"name": "selector.n"}, "subscripts": []}}, + {"Binary": { + "op": "Add", + "lhs": {"VarRef": {"name": {"name": "selector.n"}, "subscripts": []}}, + "rhs": {"Literal": {"value": {"Integer": 1}}} + }} + ] + } + } + }], + "w": { + "selector.expr[1]": {} + }, + "x": {}, + "y": {}, + "z": {}, + "m": {}, + "u": {}, + "p": {}, + "constants": {} + }); + let template = r#" +{% set cfg = {"prefix": "", "power": "pow", "if_style": "ternary", "subscript_underscore": true} %} +{{ alg_rhs_for_var_with_dae("selector.expr[1]", dae, cfg) }} +"#; + let rendered = render_template_with_dae_json(&dae_json, template).unwrap(); + + assert_eq!( + rendered.trim(), + "0", + "indexed scalar targets should project array-producing RHS expressions before rendering:\n{rendered}" + ); +} + #[test] fn test_c_alg_rhs_prefers_direct_array_connection_over_rearranged_equation() { fn var(name: &str) -> rumoca_core::Expression { @@ -161,6 +290,238 @@ fn test_c_alg_rhs_prefers_direct_indexed_equation_for_array_element() { assert_eq!(rendered.trim(), "motor_cmd_1"); } +#[test] +fn test_c_alg_rhs_prefers_direct_conditional_equation_over_indirect_alias() { + fn var(name: &str) -> rumoca_core::Expression { + rumoca_core::Expression::VarRef { + name: name.into(), + subscripts: Vec::new(), + span: rumoca_core::Span::DUMMY, + } + } + fn int(value: i64) -> rumoca_core::Expression { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + span: rumoca_core::Span::DUMMY, + } + } + fn sub(lhs: rumoca_core::Expression, rhs: rumoca_core::Expression) -> rumoca_core::Expression { + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span: rumoca_core::Span::DUMMY, + } + } + fn mul(lhs: rumoca_core::Expression, rhs: rumoca_core::Expression) -> rumoca_core::Expression { + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span: rumoca_core::Span::DUMMY, + } + } + + let indirect_alias_equation = sub(var("fan_demand"), mul(var("fan_enable"), var("fan_gain"))); + let direct_enable_equation = sub( + var("fan_enable"), + rumoca_core::Expression::If { + branches: vec![(var("unoccupied_mode"), int(1))], + else_branch: Box::new(int(0)), + span: rumoca_core::Span::DUMMY, + }, + ); + let dae_json = serde_json::json!({ + "f_x": [ + {"rhs": serde_json::to_value(indirect_alias_equation).unwrap()}, + {"rhs": serde_json::to_value(direct_enable_equation).unwrap()} + ] + }); + let template = r#" +{% set cfg = {"power": "powf", "if_style": "ternary", "subscript_underscore": true} %} +{{ alg_rhs_for_var("fan_enable", dae.f_x, cfg) }} +"#; + let rendered = render_template_with_dae_json(&dae_json, template).unwrap(); + + assert_eq!(rendered.trim(), "(unoccupied_mode ? 1 : 0)"); + assert!( + !rendered.contains("fan_demand"), + "direct conditional equation should win over an earlier indirect alias:\n{rendered}" + ); +} + +#[test] +fn test_c_alg_rhs_projects_whole_array_assignment_for_indexed_scalar_target() { + fn var(name: &str, subscripts: Vec) -> rumoca_core::Expression { + rumoca_core::Expression::VarRef { + name: name.into(), + subscripts: subscripts + .into_iter() + .map(|index| { + rumoca_core::Subscript::generated_index(index, rumoca_core::Span::DUMMY) + }) + .collect(), + span: rumoca_core::Span::DUMMY, + } + } + + let dae_json = serde_json::json!({ + "symbols": { + "setpoint_u": "setpoint_u", + "setpoint_u[1]": "setpoint_u_1", + "setpoint_y": "setpoint_y" + }, + "f_x": [ + { + "lhs": serde_json::to_value(var("setpoint_y", vec![])).unwrap(), + "rhs": serde_json::to_value(var("setpoint_u", vec![])).unwrap() + } + ] + }); + let template = r#" +{% set cfg = {"power": "powf", "subscript_underscore": true, "symbols": dae.symbols} %} +{{ alg_rhs_for_var("setpoint_y[1]", dae.f_x, cfg) }} +"#; + let rendered = render_template_with_dae_json(&dae_json, template).unwrap(); + + assert_eq!(rendered.trim(), "setpoint_u_1"); + assert!( + !rendered.trim().ends_with("setpoint_u"), + "indexed scalar assignment targets must project whole-array RHS by index:\n{rendered}" + ); +} + +#[test] +fn test_c_alg_rhs_keeps_structurally_indexed_component_rhs_scalar() { + fn var(name: &str, subscripts: Vec) -> rumoca_core::Expression { + rumoca_core::Expression::VarRef { + name: name.into(), + subscripts: subscripts + .into_iter() + .map(|index| { + rumoca_core::Subscript::generated_index(index, rumoca_core::Span::DUMMY) + }) + .collect(), + span: rumoca_core::Span::DUMMY, + } + } + + let dae_json = serde_json::json!({ + "symbols": { + "plant.power[1]": "plant_power_1", + "plant.module[1].power": "plant_module_1_power" + }, + "f_x": [ + { + "lhs": serde_json::to_value(var("plant.power", vec![1])).unwrap(), + "rhs": serde_json::to_value(var("plant.module[1].power", vec![])).unwrap() + } + ] + }); + let template = r#" +{% set cfg = {"power": "powf", "subscript_underscore": true, "symbols": dae.symbols} %} +{{ alg_rhs_for_var("plant.power[1]", dae.f_x, cfg) }} +"#; + let rendered = render_template_with_dae_json(&dae_json, template).unwrap(); + + assert_eq!(rendered.trim(), "plant_module_1_power"); + assert!( + !rendered.contains("plant_module_1_power_1"), + "structurally indexed component fields are scalar RHS values and must not be re-indexed:\n{rendered}" + ); +} + +#[test] +fn test_render_expr_at_index_projects_structurally_indexed_array_field() { + let dae_json = serde_json::json!({ + "symbols": { + "coil.ele[1].x_start": "coil_ele_1_x_start", + "coil.ele[1].x_start[1]": "coil_ele_1_x_start_1" + } + }); + let template = r#" +{% set cfg = {"power": "powf", "subscript_underscore": true, "symbols": dae.symbols} %} +{{ render_expr_at_index({"VarRef": {"name": "coil.ele[1].x_start", "subscripts": []}}, 1, cfg) }} +"#; + let rendered = render_template_with_dae_json(&dae_json, template).unwrap(); + + assert_eq!(rendered.trim(), "coil_ele_1_x_start_1"); + assert!( + !rendered.contains("coil_ele_1_x_start\n"), + "structurally indexed array fields must still project their own field index:\n{rendered}" + ); +} + +#[test] +fn test_discrete_rhs_keeps_guarded_when_ternary_for_non_sample_conditions() { + let dae_json = serde_json::json!({ + "f_z": [{ + "lhs": "z", + "rhs": { + "If": { + "branches": [[ + {"VarRef": {"name": "trigger", "subscripts": []}}, + {"Literal": {"value": {"Real": 1.0}}} + ]], + "else_branch": { + "BuiltinCall": { + "function": "Pre", + "args": [{"VarRef": {"name": "z", "subscripts": []}}] + } + } + } + } + }], + "f_m": [] + }); + let template = r#" +{% set cfg = {"prefix": "", "power": "pow", "and_op": "&&", "or_op": "||", "not_op": "!", "true_val": "1", "false_val": "0", "if_style": "ternary", "subscript_underscore": true} %} +z={{ discrete_rhs_for_var("z", dae.f_z, dae.f_m, dae, cfg) }} +"#; + let rendered = render_template_with_dae_json(&dae_json, template).unwrap(); + + assert!( + rendered.contains("z=(trigger ? 1.0 : pre(z))"), + "ordinary guarded when RHS must render as C ternary, not unconditional update:\n{rendered}" + ); +} + +#[test] +fn test_discrete_rhs_sample_guard_keeps_event_update_value() { + let dae_json = serde_json::json!({ + "f_z": [{ + "lhs": "z", + "rhs": { + "If": { + "branches": [[ + { + "BuiltinCall": { + "function": "Sample", + "args": [{"VarRef": {"name": "clocked", "subscripts": []}}] + } + }, + {"Literal": {"value": {"Real": 1.0}}} + ]], + "else_branch": { + "BuiltinCall": { + "function": "Pre", + "args": [{"VarRef": {"name": "z", "subscripts": []}}] + } + } + } + } + }], + "f_m": [] + }); + let template = r#" +{% set cfg = {"prefix": "", "power": "pow", "and_op": "&&", "or_op": "||", "not_op": "!", "true_val": "1", "false_val": "0", "if_style": "ternary", "subscript_underscore": true} %} +z={{ discrete_rhs_for_var("z", dae.f_z, dae.f_m, dae, cfg) }} +"#; + let rendered = render_template_with_dae_json(&dae_json, template).unwrap(); + + assert_eq!(rendered.trim(), "z=1.0"); +} + #[test] fn test_c_ode_rhs_solves_preserved_matrix_vector_derivative_equation() { let residual = rumoca_core::Expression::Binary { @@ -782,6 +1143,103 @@ fn test_fmi3_model_description_exports_dae_inputs_as_inputs() { ); } +#[test] +fn test_fmi_model_description_escapes_expression_attributes() { + let less_than_expression = rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Lt, + lhs: Box::new(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(1.0), + span: rumoca_core::Span::DUMMY, + }), + rhs: Box::new(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(2.0), + span: rumoca_core::Span::DUMMY, + }), + span: rumoca_core::Span::DUMMY, + }; + let mut dae = dae::Dae::new(); + dae.variables.inputs.insert( + "u".into(), + rumoca_ir_dae::Variable { + name: "u".into(), + min: Some(less_than_expression), + ..rumoca_ir_dae::Variable::empty_with_span(rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + )) + }, + ); + + let fmi2_xml = render_template_with_name( + &dae, + builtin_template("fmi2", "modelDescription.xml.jinja"), + "M", + ) + .expect("render FMI2 modelDescription"); + let fmi3_xml = render_template_with_name( + &dae, + builtin_template("fmi3", "modelDescription.xml.jinja"), + "M", + ) + .expect("render FMI3 modelDescription"); + + for xml in [fmi2_xml, fmi3_xml] { + assert!( + xml.contains(r#"min="(1.0 < 2.0)""#), + "modelDescription expression attributes must be XML escaped:\n{xml}" + ); + assert!( + !xml.contains(r#"min="(1.0 < 2.0)""#), + "modelDescription must not emit raw '<' inside attributes:\n{xml}" + ); + } +} + +#[test] +fn test_fmi_model_description_renders_string_start_without_c_quotes() { + let mut dae = dae::Dae::new(); + dae.variables.parameters.insert( + "metadata.provenance".into(), + rumoca_ir_dae::Variable { + name: "metadata.provenance".into(), + start: Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("metadata_ready_state".into()), + span: rumoca_core::Span::DUMMY, + }), + ..rumoca_ir_dae::Variable::empty_with_span(rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + )) + }, + ); + + let fmi2_xml = render_template_with_name( + &dae, + builtin_template("fmi2", "modelDescription.xml.jinja"), + "M", + ) + .expect("render FMI2 modelDescription"); + let fmi3_xml = render_template_with_name( + &dae, + builtin_template("fmi3", "modelDescription.xml.jinja"), + "M", + ) + .expect("render FMI3 modelDescription"); + + for xml in [fmi2_xml, fmi3_xml] { + assert!( + xml.contains(r#"start="metadata_ready_state""#), + "string start values must render as XML attribute values:\n{xml}" + ); + assert!( + !xml.contains(r#"start=""metadata_ready_state"""#), + "string starts must not keep C string quotes in XML attributes:\n{xml}" + ); + } +} + #[test] fn test_fmi3_build_templates_use_fmi3_platform_directory_names() { let cmake = builtin_template("fmi3", "CMakeLists.txt.jinja"); @@ -831,6 +1289,56 @@ fn test_fmi3_model_description_only_advertises_implemented_capabilities() { assert!(!xml.contains("structuralParameter"), "{xml}"); } +#[test] +fn test_fmi_build_scripts_package_only_needed_external_libraries() { + for target in ["fmi2", "fmi3"] { + let script = builtin_template(target, "build.sh.jinja"); + assert!( + script.contains("UNRESOLVED_SYMBOLS_FILE") + && script.contains("external_library_declares_unresolved_symbol") + && script.contains("external_library_exports_unresolved_symbol"), + "{target} shell build should inspect unresolved symbols before linking external libraries" + ); + assert!( + script.contains("EXTERNAL_LIBS_NEEDED=0") + && script.contains(r#"cp "$external_lib_file" "binaries/$PLATFORM/$(basename "$external_lib_file")""#), + "{target} shell build should copy only external libraries needed by the FMU binary" + ); + assert!( + script.contains("RUNTIME_PATH_FLAGS") + && script.contains("rewrite_darwin_runtime_paths"), + "{target} shell build should make copied runtime libraries loader-relative" + ); + assert!( + !script.contains("EXTERNAL_LIB_ARGS=\"$EXTERNAL_LIB_ARGS"), + "{target} shell build should not accumulate every declared external library unconditionally" + ); + } +} + +#[test] +fn test_fmi_external_include_directories_resolve_modelica_uris() { + let root = + std::env::temp_dir().join(format!("rumoca-modelica-uri-test-{}", std::process::id())); + let include_dir = root.join("Buildings").join("Resources").join("Include"); + std::fs::create_dir_all(&include_dir).expect("create temporary Modelica include dir"); + + let resolved = resolve_modelica_uri_with_roots( + "modelica://Buildings/Resources/Include", + std::iter::once(root.as_path()), + ); + assert_eq!(resolved, include_dir.to_string_lossy()); + for target in ["fmi2", "fmi3"] { + assert!( + builtin_template(target, "externalIncludeDirectories.txt.jinja") + .contains("resolve_modelica_uri(directory)"), + "{target} external include directories template should resolve modelica:// URIs" + ); + } + + std::fs::remove_dir_all(&root).ok(); +} + #[test] fn test_fmi3_initial_builtin_tracks_initialization_mode() { assert!( @@ -843,6 +1351,18 @@ fn test_fmi3_initial_builtin_tracks_initialization_mode() { ); } +#[test] +fn test_fmi2_initial_builtin_tracks_initialization_mode() { + assert!( + builtin_template("fmi2", "model.c.jinja").contains("modelInitializationMode"), + "FMI 2 generated C must evaluate initial() from the FMI initialization state" + ); + assert!( + !builtin_template("fmi2", "model.c.jinja").contains("#define initial() 0"), + "MLS initial() cannot be hard-coded false in FMI 2 initialization" + ); +} + #[test] fn test_fmi3_exit_initialization_seeds_pre_discrete_values() { let template = builtin_template("fmi3", "model.c.jinja"); @@ -1447,110 +1967,3 @@ fn test_embedded_c_array_start_helpers_are_defined() { "zeros(...) array starts should still render through the shared expression path:\n{source}" ); } - -#[test] -fn test_julia_mtk_template_empty_dae() { - let dae = dae::Dae::new(); - let result = - render_template(&dae, builtin_template("julia-mtk", "julia_mtk.jl.jinja")).unwrap(); - assert!(result.contains("using ModelingToolkit")); - assert!(result.contains("using SciMLBase: CallbackSet, ContinuousCallback, ODEProblem, solve")); - assert!(result.contains("using OrdinaryDiffEqTsit5: Tsit5")); - assert!(result.contains("using IfElse: ifelse")); - assert!(result.contains("@independent_variables t")); - assert!(result.contains("D = Differential(t)")); - assert!(result.contains("@named sys = ODESystem(eqs, t)")); - assert!(result.contains("structural_simplify(sys)")); -} - -#[test] -fn test_julia_mtk_template_with_state() { - let mut dae = dae::Dae::new(); - dae.variables.states.insert( - "x".into(), - rumoca_ir_dae::Variable { - name: "x".into(), - start: Some(rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(1.0), - span: rumoca_core::Span::DUMMY, - }), - ..rumoca_ir_dae::Variable::empty_with_span(rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name(file!()), - 1, - 2, - )) - }, - ); - dae.continuous.equations.push(rumoca_ir_dae::Equation { - lhs: Some("x".into()), - rhs: rumoca_core::Expression::VarRef { - name: "x".into(), - subscripts: vec![], - span: rumoca_core::Span::DUMMY, - }, - span: rumoca_core::Span::DUMMY, - origin: "test".into(), - scalar_count: 1, - }); - - let result = - render_template(&dae, builtin_template("julia-mtk", "julia_mtk.jl.jinja")).unwrap(); - assert!( - result.contains("x(t)"), - "state should be time-dependent: {result}" - ); - assert!( - result.contains("D(x) ~"), - "should generate derivative equation: {result}" - ); -} - -#[test] -fn test_julia_mtk_template_with_params_and_constants() { - let mut dae = dae::Dae::new(); - dae.variables.parameters.insert( - "k".into(), - rumoca_ir_dae::Variable { - name: "k".into(), - start: Some(rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(2.5), - span: rumoca_core::Span::DUMMY, - }), - ..rumoca_ir_dae::Variable::empty_with_span(rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name(file!()), - 1, - 2, - )) - }, - ); - dae.variables.constants.insert( - "g".into(), - rumoca_ir_dae::Variable { - name: "g".into(), - start: Some(rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(9.81), - span: rumoca_core::Span::DUMMY, - }), - ..rumoca_ir_dae::Variable::empty_with_span(rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name(file!()), - 1, - 2, - )) - }, - ); - - let result = - render_template(&dae, builtin_template("julia-mtk", "julia_mtk.jl.jinja")).unwrap(); - assert!( - result.contains("@parameters"), - "should have @parameters block: {result}" - ); - assert!( - result.contains("k = 2.5"), - "parameter should have default: {result}" - ); - assert!( - result.contains("g = 9.81"), - "constant should be assigned: {result}" - ); -} diff --git a/crates/rumoca-phase-codegen/src/codegen/codegen_tests/backend_template_tests/julia_template_tests.rs b/crates/rumoca-phase-codegen/src/codegen/codegen_tests/backend_template_tests/julia_template_tests.rs new file mode 100644 index 000000000..51c2bfe8d --- /dev/null +++ b/crates/rumoca-phase-codegen/src/codegen/codegen_tests/backend_template_tests/julia_template_tests.rs @@ -0,0 +1,108 @@ +use super::*; + +#[test] +fn test_julia_mtk_template_empty_dae() { + let dae = dae::Dae::new(); + let result = + render_template(&dae, builtin_template("julia-mtk", "julia_mtk.jl.jinja")).unwrap(); + assert!(result.contains("using ModelingToolkit")); + assert!(result.contains("using SciMLBase: CallbackSet, ContinuousCallback, ODEProblem, solve")); + assert!(result.contains("using OrdinaryDiffEqTsit5: Tsit5")); + assert!(result.contains("using IfElse: ifelse")); + assert!(result.contains("@independent_variables t")); + assert!(result.contains("D = Differential(t)")); + assert!(result.contains("@named sys = ODESystem(eqs, t)")); + assert!(result.contains("structural_simplify(sys)")); +} + +#[test] +fn test_julia_mtk_template_with_state() { + let mut dae = dae::Dae::new(); + dae.variables.states.insert( + "x".into(), + rumoca_ir_dae::Variable { + name: "x".into(), + start: Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(1.0), + span: rumoca_core::Span::DUMMY, + }), + ..rumoca_ir_dae::Variable::empty_with_span(rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + )) + }, + ); + dae.continuous.equations.push(rumoca_ir_dae::Equation { + lhs: Some("x".into()), + rhs: rumoca_core::Expression::VarRef { + name: "x".into(), + subscripts: vec![], + span: rumoca_core::Span::DUMMY, + }, + span: rumoca_core::Span::DUMMY, + origin: "test".into(), + scalar_count: 1, + }); + + let result = + render_template(&dae, builtin_template("julia-mtk", "julia_mtk.jl.jinja")).unwrap(); + assert!( + result.contains("x(t)"), + "state should be time-dependent: {result}" + ); + assert!( + result.contains("D(x) ~"), + "should generate derivative equation: {result}" + ); +} + +#[test] +fn test_julia_mtk_template_with_params_and_constants() { + let mut dae = dae::Dae::new(); + dae.variables.parameters.insert( + "k".into(), + rumoca_ir_dae::Variable { + name: "k".into(), + start: Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(2.5), + span: rumoca_core::Span::DUMMY, + }), + ..rumoca_ir_dae::Variable::empty_with_span(rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + )) + }, + ); + dae.variables.constants.insert( + "g".into(), + rumoca_ir_dae::Variable { + name: "g".into(), + start: Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(9.81), + span: rumoca_core::Span::DUMMY, + }), + ..rumoca_ir_dae::Variable::empty_with_span(rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + )) + }, + ); + + let result = + render_template(&dae, builtin_template("julia-mtk", "julia_mtk.jl.jinja")).unwrap(); + assert!( + result.contains("@parameters"), + "should have @parameters block: {result}" + ); + assert!( + result.contains("k = 2.5"), + "parameter should have default: {result}" + ); + assert!( + result.contains("g = 9.81"), + "constant should be assigned: {result}" + ); +} diff --git a/crates/rumoca-phase-codegen/src/codegen/fmi_template_tests.rs b/crates/rumoca-phase-codegen/src/codegen/fmi_template_tests.rs index c9c9c0034..a0cee7bec 100644 --- a/crates/rumoca-phase-codegen/src/codegen/fmi_template_tests.rs +++ b/crates/rumoca-phase-codegen/src/codegen/fmi_template_tests.rs @@ -7,6 +7,69 @@ fn builtin_template(target: &str, template: &str) -> &'static str { .expect("built-in target template must exist") } +#[test] +fn dae_template_context_projects_scheduled_times_and_preserves_provenance() { + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("scheduled-event.mo"), + 68, + 78, + ); + let event = dae::DaeScheduledTimeEvent { + time: 0.5, + source_span: Some(span), + }; + let mut dae = dae::Dae::new(); + dae.events.scheduled_time_events.push(event); + + let expected_record = serde_json::json!({ + "time": 0.5, + "source_span": serde_json::to_value(span).expect("span should serialize"), + }); + let canonical = serde_json::to_value(&dae).expect("canonical DAE should serialize"); + assert_eq!( + canonical["scheduled_time_events"], + serde_json::json!([expected_record.clone()]), + "canonical DAE must retain typed scheduled-event provenance", + ); + + let context = dae_template_json(&dae).expect("DAE template context should serialize"); + assert_eq!( + context["scheduled_time_event_metadata"], + serde_json::json!([expected_record]), + "codegen context must expose full scheduled-event metadata", + ); + assert_eq!( + context["scheduled_time_events"], + serde_json::json!([0.5]), + "stable template schedule surface must contain numeric times", + ); + + let renderer = SolveTemplateRenderer::new_with_dae( + &solve::SolveProblem::default(), + &solve::SolveArtifacts::default(), + dae, + ) + .expect("Solve renderer should accept scheduled-event metadata"); + for target in ["fmi2", "fmi3"] { + let rendered = renderer + .render_with_name(builtin_template(target, "model.c.jinja"), "M") + .unwrap_or_else(|err| panic!("{target} model.c should render: {err}")); + let scheduled_array = rendered + .split("static const double scheduled_events[] = {") + .nth(1) + .and_then(|tail| tail.split("};").next()) + .unwrap_or_else(|| panic!("{target} should render a scheduled-event array")); + assert!( + scheduled_array.contains("0.5"), + "{target} scheduled array should contain numeric 0.5:\n{scheduled_array}", + ); + assert!( + !scheduled_array.contains("source_span") && !scheduled_array.contains('{'), + "{target} scheduled array must not contain provenance JSON:\n{scheduled_array}", + ); + } +} + #[test] fn fmi_templates_snapshot_solve_pre_parameters_before_discrete_rows() { let dae = dae::Dae::new(); diff --git a/crates/rumoca-phase-codegen/src/codegen/mod.rs b/crates/rumoca-phase-codegen/src/codegen/mod.rs index c88560372..3a52d201b 100644 --- a/crates/rumoca-phase-codegen/src/codegen/mod.rs +++ b/crates/rumoca-phase-codegen/src/codegen/mod.rs @@ -142,6 +142,9 @@ pub enum CodegenInput<'a> { Ast(&'a ast::ClassTree), } +/// Build the stable DAE template context. Scheduled-event provenance remains +/// available as `scheduled_time_event_metadata`, while the established +/// `scheduled_time_events` render surface contains numeric times. pub fn dae_template_json(dae: &dae::Dae) -> Result { let mut value = serde_json::to_value(dae).map_err(|e| CodegenError::SerializationFailed { message: format!("DAE: {e}"), @@ -151,6 +154,27 @@ pub fn dae_template_json(dae: &dae::Dae) -> Result>(), + ) + .map_err(|e| CodegenError::SerializationFailed { + message: format!("scheduled_time_events: {e}"), + })?, + ) + .ok_or_else(|| CodegenError::SerializationFailed { + message: "DAE omitted scheduled_time_events".to_string(), + })?; + object.insert( + "scheduled_time_event_metadata".to_string(), + scheduled_time_event_metadata, + ); let enum_type_names = enum_type_names_from_ordinals(dae); let symbol_refs = source_refs_from_dae(dae, &enum_type_names)?; let symbol_aliases = symbol_aliases_from_dae(dae)?; @@ -519,12 +543,9 @@ fn reject_external_functions_for_simulation_template( if !template_emits_simulation_function_bodies(template) { return Ok(()); } - if let Some((name, _)) = dae_model - .symbols - .functions - .iter() - .find(|(_, function)| function.external.is_some()) - { + if let Some((name, _)) = dae_model.symbols.functions.iter().find(|(name, function)| { + function.external.is_some() && !is_supported_solve_external_function_name(name.as_str()) + }) { return Err(CodegenError::external_function_not_callable(name.as_str())); } Ok(()) @@ -543,16 +564,27 @@ fn reject_external_functions_in_json_for_simulation_template( else { return Ok(()); }; - if let Some((name, _)) = functions.iter().find(|(_, function)| { + if let Some((name, _)) = functions.iter().find(|(name, function)| { function .get("external") .is_some_and(|external| !external.is_null()) + && !is_supported_solve_external_function_name(name) }) { return Err(CodegenError::external_function_not_callable(name.as_str())); } Ok(()) } +fn is_supported_solve_external_function_name(name: &str) -> bool { + matches!( + name, + "Buildings.ThermalZones.EnergyPlus_9_6_0.BaseClasses.initialize" + | "Buildings.ThermalZones.EnergyPlus_9_6_0.BaseClasses.getParameters" + | "Buildings.ThermalZones.EnergyPlus_9_6_0.BaseClasses.exchange" + | "Buildings.ThermalZones.EnergyPlus_9_6_0.BaseClasses.SpawnExternalObject" + ) +} + fn template_emits_simulation_function_bodies(template: &str) -> bool { template.contains("func.external") && (template.contains("FMI 2.0 API") @@ -560,6 +592,54 @@ fn template_emits_simulation_function_bodies(template: &str) -> bool { || template.contains("step(")) } +/// Reusable Minijinja context for rendering multiple templates from one DAE. +pub struct DaeTemplateContext { + dae_value: Value, +} + +impl DaeTemplateContext { + pub fn from_dae_json(dae_json: &serde_json::Value) -> Self { + Self { + dae_value: Value::from_serialize(dae_json), + } + } + + pub fn from_dae(dae: &dae::Dae) -> Result { + Ok(Self { + dae_value: dae_template_value(dae)?, + }) + } + + pub fn render(&self, template: &str) -> Result { + let mut env = create_environment(); + env.add_template("inline", template)?; + let tmpl = env.get_template("inline")?; + let rendered = tmpl.render(minijinja::context! { + dae => self.dae_value.clone(), + ir => self.dae_value.clone(), + ir_kind => "dae", + })?; + Ok(rendered) + } + + pub fn render_with_name( + &self, + template: &str, + model_name: &str, + ) -> Result { + let mut env = create_environment(); + env.add_template("inline", template)?; + let tmpl = env.get_template("inline")?; + let rendered = tmpl.render(minijinja::context! { + dae => self.dae_value.clone(), + ir => self.dae_value.clone(), + ir_kind => "dae", + model_name => model_name, + })?; + Ok(rendered) + } +} + fn render_with_input_context( tmpl: &minijinja::Template<'_, '_>, input: CodegenInput<'_>, @@ -923,23 +1003,24 @@ pub fn render_template_with_dae_json( template: &str, ) -> Result { reject_external_functions_in_json_for_simulation_template(dae_json, template)?; + let dae_json = dae_json_with_template_symbol_refs(dae_json)?; let mut env = create_environment(); env.add_template("inline", template)?; - let dae_value = Value::from_serialize(dae_json); + let dae_value = Value::from_serialize(&dae_json); let tmpl = env.get_template("inline")?; let solve_value = optional_object_field(&dae_value, "solve"); - let ir_kind = template_ir_kind_from_dae_json(dae_json); + let ir_kind = template_ir_kind_from_dae_json(&dae_json); let ir_value = if ir_kind == "solve" { solve_value.clone() } else { dae_value.clone() }; - let solve_blocks = solve_blocks_from_dae_json(dae_json)?; - let solve_derivative_nodes = solve_derivative_nodes_from_dae_json(dae_json)?; - let solve_implicit_rows = solve_implicit_rows_from_dae_json(dae_json); - let solve_jacobian_rows = solve_jacobian_rows_from_dae_json(dae_json, &solve_implicit_rows); - let solve_full_jacobian_rows = solve_full_jacobian_rows_from_dae_json(dae_json); + let solve_blocks = solve_blocks_from_dae_json(&dae_json)?; + let solve_derivative_nodes = solve_derivative_nodes_from_dae_json(&dae_json)?; + let solve_implicit_rows = solve_implicit_rows_from_dae_json(&dae_json); + let solve_jacobian_rows = solve_jacobian_rows_from_dae_json(&dae_json, &solve_implicit_rows); + let solve_full_jacobian_rows = solve_full_jacobian_rows_from_dae_json(&dae_json); let result = tmpl.render(minijinja::context! { dae => dae_value.clone(), solve => solve_value, @@ -961,22 +1042,23 @@ pub fn render_template_with_dae_json_and_name( model_name: &str, ) -> Result { reject_external_functions_in_json_for_simulation_template(dae_json, template)?; + let dae_json = dae_json_with_template_symbol_refs(dae_json)?; let mut env = create_environment(); env.add_template("inline", template)?; - let dae_value = Value::from_serialize(dae_json); + let dae_value = Value::from_serialize(&dae_json); let solve_value = optional_object_field(&dae_value, "solve"); - let ir_kind = template_ir_kind_from_dae_json(dae_json); + let ir_kind = template_ir_kind_from_dae_json(&dae_json); let ir_value = if ir_kind == "solve" { solve_value.clone() } else { dae_value.clone() }; - let solve_blocks = solve_blocks_from_dae_json(dae_json)?; - let solve_derivative_nodes = solve_derivative_nodes_from_dae_json(dae_json)?; - let solve_implicit_rows = solve_implicit_rows_from_dae_json(dae_json); - let solve_jacobian_rows = solve_jacobian_rows_from_dae_json(dae_json, &solve_implicit_rows); - let solve_full_jacobian_rows = solve_full_jacobian_rows_from_dae_json(dae_json); + let solve_blocks = solve_blocks_from_dae_json(&dae_json)?; + let solve_derivative_nodes = solve_derivative_nodes_from_dae_json(&dae_json)?; + let solve_implicit_rows = solve_implicit_rows_from_dae_json(&dae_json); + let solve_jacobian_rows = solve_jacobian_rows_from_dae_json(&dae_json, &solve_implicit_rows); + let solve_full_jacobian_rows = solve_full_jacobian_rows_from_dae_json(&dae_json); let tmpl = env.get_template("inline")?; let result = tmpl.render(minijinja::context! { dae => dae_value.clone(), @@ -998,6 +1080,134 @@ fn optional_object_field(value: &Value, name: &str) -> Value { get_field(value, name).unwrap_or_else(|_| Value::from_serialize(serde_json::Map::new())) } +fn dae_json_with_template_symbol_refs( + dae_json: &serde_json::Value, +) -> Result { + let mut normalized = dae_json.clone(); + let mut refs = IndexSet::::new(); + if let Some(existing) = normalized + .get("symbol_refs") + .and_then(serde_json::Value::as_array) + { + for value in existing { + if let Some(reference) = value.as_str() { + refs.insert(reference.to_string()); + } + } + } + add_json_function_symbol_refs(&normalized, &mut refs)?; + + let Some(object) = normalized.as_object_mut() else { + return Ok(normalized); + }; + object.insert( + "symbol_refs".to_string(), + serde_json::to_value(refs.into_iter().collect::>()).map_err(|err| { + CodegenError::SerializationFailed { + message: format!("json symbol_refs: {err}"), + } + })?, + ); + Ok(normalized) +} + +fn add_json_function_symbol_refs( + dae_json: &serde_json::Value, + refs: &mut IndexSet, +) -> Result<(), CodegenError> { + let Some(functions) = dae_json + .get("functions") + .and_then(serde_json::Value::as_object) + else { + return Ok(()); + }; + for (func_name, func) in functions { + refs.insert(func_name.clone()); + if let Some(outputs) = func.get("outputs").and_then(serde_json::Value::as_array) { + add_json_function_output_refs(func_name, outputs, refs)?; + } + for section in ["inputs", "outputs", "locals"] { + let Some(items) = func.get(section).and_then(serde_json::Value::as_array) else { + continue; + }; + for item in items { + let name = json_named_item_name(item, section)?; + let dims = json_dims(item, name)?; + add_source_refs_for_var(name, &dims_to_i64(&dims)?, refs)?; + } + } + } + Ok(()) +} + +fn add_json_function_output_refs( + func_name: &str, + outputs: &[serde_json::Value], + refs: &mut IndexSet, +) -> Result<(), CodegenError> { + for output in outputs { + let name = json_named_item_name(output, "function output")?; + let dims = json_dims(output, name)?; + let count = source_ref_scalar_count(name, &dims)?; + for element_idx in 1..=count { + let selector = if count == 1 { + name.to_string() + } else { + format!("{name}[{element_idx}]") + }; + refs.insert(format!("{func_name}.{selector}")); + } + } + Ok(()) +} + +fn json_named_item_name<'a>( + item: &'a serde_json::Value, + context: &str, +) -> Result<&'a str, CodegenError> { + item.get("name") + .and_then(serde_json::Value::as_str) + .ok_or_else(|| CodegenError::SerializationFailed { + message: format!("{context} missing string name"), + }) +} + +fn json_dims(item: &serde_json::Value, name: &str) -> Result, CodegenError> { + let Some(dims) = item.get("dims").and_then(serde_json::Value::as_array) else { + return Ok(Vec::new()); + }; + let mut converted = codegen_vec_with_capacity(dims.len(), "json symbol dimension count")?; + for dim in dims { + let Some(dim) = dim.as_i64() else { + return Err(CodegenError::SerializationFailed { + message: format!("json symbol dimension for `{name}` is not an integer"), + }); + }; + if dim > 0 { + converted.push(usize::try_from(dim).map_err(|_| { + CodegenError::SerializationFailed { + message: format!( + "json symbol dimension {dim} for `{name}` exceeds host index range" + ), + } + })?); + } + } + Ok(converted) +} + +fn dims_to_i64(dims: &[usize]) -> Result, CodegenError> { + let mut converted = codegen_vec_with_capacity(dims.len(), "json i64 dimension count")?; + for dim in dims { + converted.push( + i64::try_from(*dim).map_err(|_| CodegenError::SerializationFailed { + message: format!("json dimension {dim} exceeds i64 range"), + })?, + ); + } + Ok(converted) +} + fn template_ir_kind_from_dae_json(dae_json: &serde_json::Value) -> &'static str { if dae_json .get("__ir_kind") @@ -1020,6 +1230,8 @@ fn solve_blocks_from_dae_json(dae_json: &serde_json::Value) -> Result Environment<'static> { // Fail fast on missing fields/variables in templates. env.set_undefined_behavior(UndefinedBehavior::Strict); + add_basic_template_helpers(&mut env); + add_solve_template_helpers(&mut env); + add_statement_template_helpers(&mut env); + add_rhs_template_helpers(&mut env); + env +} + +fn add_basic_template_helpers(env: &mut Environment<'static>) { // Custom filters env.add_filter("sanitize", sanitize_filter); env.add_filter("product", product_filter); env.add_filter("last_segment", last_segment_filter); + env.add_filter("json", json_filter); // eFMI manifest render env (contract §3b): autoescape is OFF, so every // text value is escaped explicitly and every raw f64 is rendered as a // valid xs:double lexical. @@ -1218,11 +1439,29 @@ fn create_environment() -> Environment<'static> { env.add_function("allocate_symbols", allocate_symbols_function); env.add_function("target_symbols", target_symbols_function); env.add_function("symbol", symbol_function); + env.add_function("resolve_modelica_uri", resolve_modelica_uri_function); env.add_function("source_ref", source_ref_function); // Custom functions for expression rendering env.add_function("render_expr", render_expr_function); + env.add_function("render_xml_attr_expr", render_xml_attr_expr_function); + env.add_function( + "render_xml_attr_expr_at_index", + render_xml_attr_expr_at_index_function, + ); env.add_function("render_event_indicator", render_event_indicator_function); + env.add_function("render_matmul_c", render_matmul_c_function); + env.add_function("render_matmul_mlir", render_matmul_mlir_function); + env.add_function("render_linsolve_mlir", render_linsolve_mlir_function); + env.add_function("render_equation", render_equation_function); + env.add_function( + "render_dae_equations", + render_dae_modelica::render_dae_equations_function, + ); + env.add_function("fail", fail_function); +} + +fn add_solve_template_helpers(env: &mut Environment<'static>) { env.add_function("render_solve_row_c", render_solve_row_c_function); env.add_function( "fmi3_scalar_projection_schedule", @@ -1278,26 +1517,37 @@ fn create_environment() -> Environment<'static> { "render_solve_pre_param_binding_c", render_solve_pre_param_binding_c_function, ); - env.add_function("render_matmul_c", render_matmul_c_function); - env.add_function("render_matmul_mlir", render_matmul_mlir_function); - env.add_function("render_linsolve_mlir", render_linsolve_mlir_function); - env.add_function("render_equation", render_equation_function); - env.add_function( - "render_dae_equations", - render_dae_modelica::render_dae_equations_function, - ); +} +fn add_statement_template_helpers(env: &mut Environment<'static>) { // Custom functions for statement rendering (MLS §12: function bodies) env.add_function("render_statement", render_statement_function); env.add_function("render_statements", render_statements_function); + env.add_function( + "render_function_statements", + render_function_statements_function, + ); // Custom function for flat equation rendering (Model residual equations) env.add_function("render_flat_equation", render_flat_equation_function); + // Render the symbolic scalar name for an array element. DAE residuals keep + // Modelica multi-dimensional subscripts while codegen iterates linear slots. + env.add_function("array_scalar_name", array_scalar_name_function); + // Custom function for detecting self-referential (builtin alias) functions env.add_function("is_self_call", is_self_call_function); - env.add_function("fail", fail_function); + env.add_function( + "unsupported_c_function_body", + unsupported_c_function_body_function, + ); + env.add_function( + "unsupported_c_function_name", + unsupported_c_function_name_function, + ); +} +fn add_rhs_template_helpers(env: &mut Environment<'static>) { // Extract explicit ODE rhs from residual equation: 0 = der(x) - expr → expr env.add_function("ode_rhs", render_c::ode_rhs_function); // Find derivative expression for a specific state variable @@ -1305,6 +1555,14 @@ fn create_environment() -> Environment<'static> { // Find explicit RHS for an algebraic variable from residual: 0 = y - expr → expr env.add_function("alg_rhs_for_var", render_c::alg_rhs_for_var_function); + env.add_function( + "alg_rhs_for_var_with_dae", + render_c::alg_rhs_for_var_with_dae_function, + ); + env.add_function( + "visible_or_alg_rhs_for_var", + render_c::visible_or_alg_rhs_for_var_function, + ); env.add_function( "alg_rhs_for_var_or_self", render_c::alg_rhs_for_var_or_self_function, @@ -1313,20 +1571,34 @@ fn create_environment() -> Environment<'static> { "discrete_rhs_for_var", render_c::discrete_rhs_for_var_function, ); - // Index into an array expression to render element i (1-based) env.add_function( "render_expr_at_index", render_c::render_expr_at_index_function, ); + env.add_function( + "parameter_binding_rhs", + render_c::parameter_binding_rhs_function, + ); // Check if an expression is a string literal for scalar templates. env.add_function("is_string_literal", render_c::is_string_literal_function); + env.add_function("expr_has_var_ref", render_c::expr_has_var_ref_function); + env.add_function( + "expr_has_dynamic_multidim_index", + render_c::expr_has_dynamic_multidim_index_function, + ); + env.add_function( + "initial_rhs_for_var", + render_c::initial_rhs_for_var_function, + ); + env.add_function( + "initial_runtime_rhs_for_var", + render_c::initial_runtime_rhs_for_var_function, + ); // Check if a function has record-typed parameters env.add_function("has_complex_params", render_c::has_complex_params_function); - - env } /// Sanitize a name for use as a simple emitted identifier. @@ -1335,6 +1607,7 @@ fn create_environment() -> Environment<'static> { /// reserved words are handled by `allocate_symbols` with a template-supplied /// policy, not by this lossy fallback. pub(crate) fn sanitize_name(name: &str) -> String { + let name = normalize_static_component_subscripts(name); let mut result = String::with_capacity(name.len()); for ch in name.chars() { if ch.is_alphanumeric() || ch == '_' { @@ -1359,7 +1632,7 @@ pub(crate) fn escape_reserved_keyword(name: &str) -> String { /// /// Replaces dots and other non-identifier characters with underscores. fn sanitize_filter(value: Value) -> String { - let s = value.to_string(); + let s = normalize_static_component_subscripts(&value.to_string()); let mut result = String::with_capacity(s.len()); for ch in s.chars() { if ch.is_alphanumeric() || ch == '_' { @@ -1373,6 +1646,118 @@ fn sanitize_filter(value: Value) -> String { result } +fn normalize_static_component_subscripts(name: &str) -> String { + let mut out = String::with_capacity(name.len()); + let mut rest = name; + while let Some(open_idx) = rest.find('[') { + out.push_str(&rest[..open_idx + 1]); + let after_open = &rest[open_idx + 1..]; + let Some(close_idx) = after_open.find(']') else { + out.push_str(after_open); + return out; + }; + let inner = &after_open[..close_idx]; + if let Some(normalized) = normalize_static_subscript_list(inner) { + out.push_str(&normalized); + } else { + out.push_str(inner); + } + out.push(']'); + rest = &after_open[close_idx + 1..]; + } + out.push_str(rest); + out +} + +fn normalize_static_subscript_list(inner: &str) -> Option { + let mut normalized = Vec::new(); + for part in inner.split(',') { + normalized.push(eval_integer_sum(part.trim())?.to_string()); + } + Some(normalized.join(",")) +} + +fn eval_integer_sum(expr: &str) -> Option { + let mut chars = expr.chars().peekable(); + let mut total = 0i64; + let mut sign = 1i64; + + loop { + while matches!(chars.peek(), Some(ch) if ch.is_whitespace()) { + chars.next(); + } + while matches!(chars.peek(), Some('(')) { + chars.next(); + while matches!(chars.peek(), Some(ch) if ch.is_whitespace()) { + chars.next(); + } + } + match chars.peek().copied() { + Some('+') => { + sign = 1; + chars.next(); + continue; + } + Some('-') => { + sign = -1; + chars.next(); + continue; + } + _ => {} + } + + while matches!(chars.peek(), Some(ch) if ch.is_whitespace()) { + chars.next(); + } + while matches!(chars.peek(), Some('(')) { + chars.next(); + while matches!(chars.peek(), Some(ch) if ch.is_whitespace()) { + chars.next(); + } + } + + let mut value = 0i64; + let mut digits = 0usize; + while let Some(ch) = chars.peek().copied() { + if let Some(digit) = ch.to_digit(10) { + value = value.checked_mul(10)?.checked_add(digit as i64)?; + digits += 1; + chars.next(); + } else { + break; + } + } + if digits == 0 { + return None; + } + total = total.checked_add(sign.checked_mul(value)?)?; + + while matches!(chars.peek(), Some(ch) if ch.is_whitespace()) { + chars.next(); + } + while matches!(chars.peek(), Some(')')) { + chars.next(); + while matches!(chars.peek(), Some(ch) if ch.is_whitespace()) { + chars.next(); + } + } + match chars.peek().copied() { + Some('+') => { + sign = 1; + chars.next(); + } + Some('-') => { + sign = -1; + chars.next(); + } + Some(_) => return None, + None => break, + } + } + + Some(total) +} + /// Filter to extract the last dot-separated segment of a name. /// /// Used in templates: `{{ "Modelica.Math.sin" | last_segment }}` -> `"sin"` @@ -1443,6 +1828,15 @@ fn product_filter(value: Value) -> Result { Ok(Value::from(result)) } +fn json_filter(value: Value) -> RenderResult { + serde_json::to_string(&value).map_err(|err| { + minijinja::Error::new( + minijinja::ErrorKind::InvalidOperation, + format!("json filter serialization failed: {err}"), + ) + }) +} + fn value_to_string(value: &Value) -> String { value .as_str() @@ -1450,6 +1844,53 @@ fn value_to_string(value: &Value) -> String { .unwrap_or_else(|| value.to_string().trim_matches('"').to_string()) } +fn resolve_modelica_uri_function(uri: Value) -> String { + resolve_modelica_uri(&value_to_string(&uri)) +} + +fn resolve_modelica_uri(uri: &str) -> String { + uri.to_string() +} + +#[cfg(test)] +fn resolve_modelica_uri_with_roots(uri: &str, source_roots: I) -> String +where + I: IntoIterator, + P: AsRef, +{ + let Some(rest) = uri.strip_prefix("modelica://") else { + return uri.to_string(); + }; + let Some((package, relative)) = rest.split_once('/') else { + return uri.to_string(); + }; + let relative_path = relative + .split('/') + .filter(|segment| !segment.is_empty()) + .fold(std::path::PathBuf::new(), |path, segment| { + path.join(segment) + }); + for root in source_roots { + let root = root.as_ref(); + let direct = root.join(&relative_path); + if direct.exists() { + return direct.to_string_lossy().into_owned(); + } + let nested = root.join(package).join(&relative_path); + if nested.exists() { + return nested.to_string_lossy().into_owned(); + } + let root_name_matches_package = root + .file_name() + .and_then(|name| name.to_str()) + .is_some_and(|name| name == package || name.starts_with(&format!("{package} "))); + if root_name_matches_package && direct.parent().is_some_and(Path::exists) { + return direct.to_string_lossy().into_owned(); + } + } + uri.to_string() +} + fn dims_from_value(value: &Value) -> Result, minijinja::Error> { let Some(len) = value.len() else { return Ok(Vec::new()); @@ -1596,6 +2037,57 @@ fn source_ref_function(name: Value, dims: Value, flat_index: Value) -> RenderRes )) } +fn array_scalar_name_function(base_name: Value, dims: Value, linear_index: Value) -> RenderResult { + let name = base_name + .as_str() + .map(str::to_owned) + .unwrap_or_else(|| base_name.to_string().trim_matches('"').to_string()); + let dims = dims_from_value(&dims)?; + let Some(linear_index) = linear_index.as_usize() else { + return Err(render_err(format!( + "array scalar index for {name} must be a positive integer" + ))); + }; + render_array_scalar_name(&name, &dims, linear_index) +} + +fn render_array_scalar_name(name: &str, dims: &[usize], linear_index: usize) -> RenderResult { + if dims.is_empty() { + return Ok(name.to_string()); + } + if linear_index == 0 { + return Err(render_err(format!( + "array scalar index for {name} is one-based and cannot be zero" + ))); + } + let total = dims + .iter() + .try_fold(1usize, |acc, dim| acc.checked_mul(*dim)); + let Some(total) = total else { + return Err(render_err(format!( + "array dimensions for {name} overflow usize" + ))); + }; + if linear_index > total { + return Err(render_err(format!( + "array scalar index {linear_index} for {name} exceeds scalar size {total}" + ))); + } + + let mut remainder = linear_index - 1; + let mut subscripts = vec![0usize; dims.len()]; + for (slot, dim) in subscripts.iter_mut().rev().zip(dims.iter().rev()) { + *slot = (remainder % *dim) + 1; + remainder /= *dim; + } + let rendered = subscripts + .iter() + .map(usize::to_string) + .collect::>() + .join(","); + Ok(format!("{name}[{rendered}]")) +} + /// Fail template rendering with an explicit message. /// /// Templates use this to declare target-specific capability constraints @@ -1653,6 +2145,68 @@ fn is_self_call_function(func_name: Value, func: Value) -> Result Result { + use render_expr::get_field; + + if let Ok(external) = get_field(&func, "external") + && !external.is_undefined() + { + return Ok(true); + } + + if let Ok(inputs) = get_field(&func, "inputs") + && let Some(len) = inputs.len() + { + for i in 0..len { + let Ok(input) = inputs.get_item(&Value::from(i)) else { + continue; + }; + if let Ok(dims) = get_field(&input, "dims") + && dims.len().unwrap_or(0) > 0 + { + return Ok(true); + } + } + } + + if let Ok(locals) = get_field(&func, "locals") + && let Some(len) = locals.len() + { + for i in 0..len { + let Ok(local) = locals.get_item(&Value::from(i)) else { + continue; + }; + if let Ok(dims) = get_field(&local, "dims") + && dims.len().unwrap_or(0) > 0 + { + return Ok(true); + } + } + } + + let Ok(body) = get_field(&func, "body") else { + return Ok(false); + }; + let body_debug = body.to_string(); + Ok(body_debug.contains("NamedArgument") + || body_debug.contains("FieldAccess") + || body_debug.contains("Assert") + || body_debug.contains("outputs") + || body_debug.contains("String(") + || body_debug.contains("Modelica.Utilities.Strings") + || body_debug.contains("Modelica_Utilities_Strings")) +} + +fn unsupported_c_function_name_function(func_name: Value) -> Result { + let name = func_name.to_string().replace('"', ""); + Ok( + name.contains("Modelica.Media.Interfaces.PartialSimpleMedium") + || name.contains("Modelica_Media_Interfaces_PartialSimpleMedium") + || name.contains("Modelica.Media.Interfaces.PartialMedium") + || name.contains("Modelica_Media_Interfaces_PartialMedium"), + ) +} + /// Built-in expression renderer function. /// /// Usage in templates: @@ -1677,6 +2231,32 @@ fn render_expr_function(expr: Value, config: Value) -> RenderResult { render_expression(&expr, &cfg) } +fn render_xml_attr_expr_function(expr: Value, config: Value) -> RenderResult { + let cfg = ExprConfig::from_value(&config); + xml_attr_expr(render_expression(&expr, &cfg)?) +} + +fn render_xml_attr_expr_at_index_function( + expr: Value, + index: Value, + config: Value, +) -> RenderResult { + xml_attr_expr(render_c::render_expr_at_index_function( + expr, index, config, + )?) +} + +fn xml_attr_expr(mut rendered: String) -> RenderResult { + if rendered.len() >= 2 && rendered.starts_with('"') && rendered.ends_with('"') { + rendered = rendered[1..rendered.len() - 1].to_string(); + } + Ok(rendered + .replace('&', "&") + .replace('<', "<") + .replace('>', ">") + .replace('"', """)) +} + /// Render a relation as a numeric root function for FMI event indicators. /// /// DAE `relation` entries are boolean expressions such as `a < b`, but FMI @@ -1777,6 +2357,19 @@ fn render_statements_function(stmts: Value, config: Value, indent: Value) -> Ren render_statements(&stmts, &cfg, indent_str) } +fn render_function_statements_function( + stmts: Value, + config: Value, + indent: Value, + return_value: Value, +) -> RenderResult { + let mut cfg = ExprConfig::from_value(&config); + cfg.subscript_underscore = false; + cfg.return_value = return_value.as_str().map(ToString::to_string); + let indent_str = indent.as_str().unwrap_or(" "); + render_statements(&stmts, &cfg, indent_str) +} + // ── ExprConfig and helpers ─────────────────────────────────────────── /// Configuration for expression rendering. @@ -1824,8 +2417,14 @@ pub(crate) struct ExprConfig { /// Optional aliases from Appendix-B condition memory (`c[i]`) to live /// relation expressions for backends that do not run event iteration. pub(crate) condition_aliases: Option, + /// Optional Modelica source scope used to resolve unqualified VarRefs in + /// equations rendered from a scoped source component. + pub(crate) source_scope: Option, /// Render-time substitutions for expression-level unrolling. pub(crate) substitutions: Vec<(String, String)>, + /// Function return expression used when rendering `return` statements in + /// backends where the output variable is the canonical return value. + pub(crate) return_value: Option, } #[derive(Clone, Copy)] @@ -1863,7 +2462,9 @@ impl Default for ExprConfig { float_literals: false, symbols: None, condition_aliases: None, + source_scope: None, substitutions: Vec::new(), + return_value: None, } } } @@ -1962,6 +2563,9 @@ impl ExprConfig { if let Some(val) = get_present_attr(v, "condition_aliases") { cfg.condition_aliases = Some(val); } + if let Some(s) = get_non_empty_str_attr(v, "source_scope") { + cfg.source_scope = Some(s); + } cfg } diff --git a/crates/rumoca-phase-codegen/src/codegen/render_c.rs b/crates/rumoca-phase-codegen/src/codegen/render_c.rs index a03dea7e6..07ba65769 100644 --- a/crates/rumoca-phase-codegen/src/codegen/render_c.rs +++ b/crates/rumoca-phase-codegen/src/codegen/render_c.rs @@ -1,5 +1,9 @@ //! C-backend template functions for FMI2 and embedded-C code generation. //! +//! SPEC_0021 file-size exception: C/FMI RHS helper surfaces currently share +//! one module. split plan: move algebraic RHS, ODE RHS, and discrete RHS helper +//! families into separate render_c submodules. +//! //! These functions are registered in the minijinja environment and used by //! `fmi2/model.c.jinja` and `embedded_c/model.c.jinja` templates to extract explicit //! ODE/algebraic RHS expressions from residual-form DAE equations. @@ -147,6 +151,47 @@ pub(super) fn is_string_literal_function(expr: Value) -> String { String::new() } +/// Render a parameter binding RHS when the initializer is a scalar C expression. +pub(super) fn parameter_binding_rhs_function( + _target_name: Value, + expr: Value, + index: Value, + config: Value, +) -> RenderResult { + if is_string_literal_function(expr.clone()) == "yes" || expr_has_dynamic_multidim_index(&expr) { + return Ok(String::new()); + } + + let rendered = if let Some(idx) = index.as_usize().filter(|idx| *idx > 0) { + render_expr_at_index_function(expr, Value::from(idx), config)? + } else { + let cfg = ExprConfig::from_value(&config); + render_expression(&expr, &cfg)? + }; + if is_missing_equation_rhs(&rendered) { + Ok(String::new()) + } else { + Ok(rendered) + } +} + +/// Check if an expression contains any variable reference. +pub(super) fn expr_has_var_ref_function(expr: Value) -> String { + if expr_has_var_ref(&expr) { + return "yes".to_string(); + } + String::new() +} + +/// Check whether a C initializer/binding expression contains an alias access +/// that cannot be rendered as a scalar expression. +pub(super) fn expr_has_dynamic_multidim_index_function(expr: Value) -> String { + if expr_has_dynamic_multidim_index(&expr) { + return "yes".to_string(); + } + String::new() +} + /// Check if a function has record-typed parameters. /// /// Returns "yes" if any input parameter carries record type metadata. @@ -296,6 +341,47 @@ pub(super) fn alg_rhs_for_var_function( } } +/// Extract algebraic RHS from a DAE context object. +/// +/// This is a convenience wrapper for templates that carry the prepared DAE +/// context alongside Solve IR and therefore have `dae` rather than `dae.f_x` +/// at the call site. +pub(super) fn alg_rhs_for_var_with_dae_function( + var_name: Value, + dae: Value, + config: Value, +) -> RenderResult { + let equations = crate::codegen::get_field(&dae, "f_x") + .unwrap_or_else(|_| Value::from_serialize(Vec::::new())); + alg_rhs_for_var_function(var_name, equations, config) +} + +/// Extract a visible variable RHS from Solve IR when available, otherwise fall +/// back to explicit DAE equations. +pub(super) fn visible_or_alg_rhs_for_var_function( + var_name: Value, + dae: Value, + solve: Value, + expr_config: Value, + solve_config: Value, +) -> RenderResult { + let name = var_name.to_string().trim_matches('"').to_string(); + if let Some((visible_index, row)) = solve_visible_row_for_name(&name, &solve)? { + if solve_visible_row_is_identity_for_index(&row, visible_index)? { + let fallback = alg_rhs_for_var_with_dae_function( + Value::from(name.clone()), + dae.clone(), + expr_config.clone(), + )?; + if !is_missing_equation_rhs(&fallback) { + return Ok(fallback); + } + } + return super::render_solve::render_solve_row_c_function(row, solve_config); + } + alg_rhs_for_var_with_dae_function(Value::from(name), dae, expr_config) +} + /// Extract algebraic RHS like `alg_rhs_for_var`, but if no matching equation is /// found, return the current variable alias (hold-last-value semantics). /// @@ -364,6 +450,266 @@ pub(super) fn discrete_rhs_for_var_function( // ── Helpers ───────────────────────────────────────────────────────────── +fn solve_visible_row_for_name( + name: &str, + solve: &Value, +) -> Result, minijinja::Error> { + let Ok(visible_names) = get_field(solve, "visible_names") else { + return no_render_match(); + }; + let Ok(visible_value_rows) = get_field(solve, "visible_value_rows") else { + return no_render_match(); + }; + let Ok(programs) = get_field(&visible_value_rows, "programs") else { + return no_render_match(); + }; + let Ok(iter) = visible_names.try_iter() else { + return no_render_match(); + }; + for (index, visible_name) in iter.enumerate() { + if visible_name.to_string().trim_matches('"') != name { + continue; + } + return Ok(programs + .get_item(&Value::from(index)) + .ok() + .map(|row| (index, row))); + } + no_render_match() +} + +fn solve_visible_row_is_identity_for_index( + row: &Value, + visible_index: usize, +) -> Result { + let Ok(iter) = row.try_iter() else { + return Ok(false); + }; + let ops = iter.collect::>(); + if ops.len() != 2 { + return Ok(false); + } + let Ok(load_y) = get_field(&ops[0], "LoadY") else { + return Ok(false); + }; + let Ok(store_output) = get_field(&ops[1], "StoreOutput") else { + return Ok(false); + }; + let Some(dst) = get_field(&load_y, "dst")?.as_usize() else { + return Ok(false); + }; + let Some(index) = get_field(&load_y, "index")?.as_usize() else { + return Ok(false); + }; + let Some(src) = get_field(&store_output, "src")?.as_usize() else { + return Ok(false); + }; + Ok(dst == src && index == visible_index) +} + +fn is_missing_equation_rhs(rendered: &str) -> bool { + rendered.contains("WARNING: no equation found") +} + +/// Find an explicit RHS for an initialization equation targeting `var_name`. +pub(super) fn initial_rhs_for_var_function( + dae: Value, + var_name: Value, + config: Value, +) -> RenderResult { + let cfg = ExprConfig::from_value(&config); + let name = var_name + .as_str() + .map(str::to_string) + .unwrap_or_else(|| var_name.to_string().trim_matches('"').to_string()); + let Ok(initial_equations) = get_field(&dae, "initial_equations") else { + return Ok(String::new()); + }; + let Some(len) = initial_equations.len() else { + return Ok(String::new()); + }; + + let target_names = initial_rhs_target_names(&dae, &name)?; + for target_name in target_names { + let mut scoped_cfg = cfg.clone(); + scoped_cfg.source_scope = rumoca_core::parent_scope(&target_name).map(str::to_string); + for i in 0..len { + let Ok(eq) = initial_equations.get_item(&Value::from(i)) else { + continue; + }; + if let Some(rhs) = find_algebraic_rhs(&eq, &target_name, &scoped_cfg)? { + return Ok(rhs); + } + } + } + + Ok(String::new()) +} + +fn initial_rhs_target_names( + dae: &Value, + storage_name: &str, +) -> Result, minijinja::Error> { + if rumoca_core::has_top_level_dot(storage_name) || storage_name.contains('[') { + return Ok(vec![storage_name.to_string()]); + } + if let Some(source_name) = source_name_for_dae_variable(dae, storage_name)? + && source_name != storage_name + { + return Ok(vec![source_name, storage_name.to_string()]); + } + Ok(vec![storage_name.to_string()]) +} + +fn source_name_for_dae_variable( + dae: &Value, + storage_name: &str, +) -> Result, minijinja::Error> { + for partition in ["x", "y", "u", "w", "p", "constants", "z", "m"] { + let Ok(vars) = get_field(dae, partition) else { + continue; + }; + let Ok(var) = vars.get_item(&Value::from(storage_name)) else { + continue; + }; + let Ok(component_ref) = get_field(&var, "component_ref") else { + continue; + }; + let rendered = source_component_ref_name(&component_ref)?; + if !rendered.is_empty() { + return Ok(Some(rendered)); + } + } + no_render_match() +} + +pub(super) fn initial_runtime_rhs_for_var_function( + dae: Value, + var_name: Value, + config: Value, +) -> RenderResult { + let rhs = initial_rhs_for_var_function(dae.clone(), var_name.clone(), config.clone())?; + if !rhs.is_empty() && !rhs_is_numeric_literal(&rhs) { + return Ok(rhs); + } + + let name = var_name + .as_str() + .map(str::to_string) + .unwrap_or_else(|| var_name.to_string().trim_matches('"').to_string()); + if rhs.is_empty() + && let Some(start_rhs) = runtime_state_start_rhs(&dae, &name)? + { + return Ok(start_rhs); + } + let Some(prefix) = rumoca_core::parent_scope(&name) else { + return Ok(rhs); + }; + let candidate = format!("{prefix}.p"); + let Ok(params) = get_field(&dae, "p") else { + return Ok(rhs); + }; + let Ok(param) = get_field(¶ms, &candidate) else { + return Ok(rhs); + }; + let Ok(start) = get_field(¶m, "start") else { + return Ok(rhs); + }; + if !expr_has_var_ref(&start) { + return Ok(rhs); + } + Ok(super::sanitize_name(&candidate)) +} + +fn rhs_is_numeric_literal(rhs: &str) -> bool { + let trimmed = rhs.trim(); + !trimmed.is_empty() + && trimmed + .chars() + .all(|ch| ch.is_ascii_digit() || matches!(ch, '.' | '-' | '+' | 'e' | 'E')) +} + +fn runtime_state_start_rhs(dae: &Value, name: &str) -> Result, minijinja::Error> { + let Some(prefix) = rumoca_core::parent_scope(name) else { + return no_render_match(); + }; + let candidate = format!("{prefix}.p"); + let Ok(states) = get_field(dae, "x") else { + return no_render_match(); + }; + let Ok(state) = get_field(&states, name) else { + return no_render_match(); + }; + let Ok(start) = get_field(&state, "start") else { + return no_render_match(); + }; + if !expr_has_var_ref(&start) { + return no_render_match(); + } + let Ok(params) = get_field(dae, "p") else { + return no_render_match(); + }; + if get_field(¶ms, &candidate).is_err() { + return no_render_match(); + } + Ok(Some(super::sanitize_name(&candidate))) +} + +fn expr_has_var_ref(expr: &Value) -> bool { + if get_field(expr, "VarRef").is_ok() { + return true; + } + if let Ok(binary) = get_field(expr, "Binary") { + return get_field(&binary, "lhs").is_ok_and(|lhs| expr_has_var_ref(&lhs)) + || get_field(&binary, "rhs").is_ok_and(|rhs| expr_has_var_ref(&rhs)); + } + if let Ok(unary) = get_field(expr, "Unary") { + return get_field(&unary, "rhs").is_ok_and(|rhs| expr_has_var_ref(&rhs)); + } + if let Ok(call) = get_field(expr, "BuiltinCall").or_else(|_| get_field(expr, "FunctionCall")) + && let Ok(args) = get_field(&call, "args") + { + return list_any(&args, |arg| expr_has_var_ref(&arg)); + } + if let Ok(if_expr) = get_field(expr, "If") { + let branch_refs = get_field(&if_expr, "branches") + .map(|branches| list_any(&branches, |branch| expr_has_var_ref(&branch))) + .unwrap_or(false); + let else_refs = get_field(&if_expr, "else_branch") + .map(|else_branch| expr_has_var_ref(&else_branch)) + .unwrap_or(false); + return branch_refs || else_refs; + } + if let Ok(array) = get_field(expr, "Array").or_else(|_| get_field(expr, "Tuple")) + && let Ok(elements) = get_field(&array, "elements") + { + return list_any(&elements, |element| expr_has_var_ref(&element)); + } + false +} + +fn expr_has_dynamic_multidim_index(expr: &Value) -> bool { + if let Ok(var_ref) = get_field(expr, "VarRef") + && let Ok(subs) = get_field(&var_ref, "subscripts") + && subs.len().unwrap_or(0) > 1 + { + return true; + } + if let Ok(binary) = get_field(expr, "Binary") { + return get_field(&binary, "lhs").is_ok_and(|lhs| expr_has_dynamic_multidim_index(&lhs)) + || get_field(&binary, "rhs").is_ok_and(|rhs| expr_has_dynamic_multidim_index(&rhs)); + } + if let Ok(unary) = get_field(expr, "Unary") { + return get_field(&unary, "rhs").is_ok_and(|rhs| expr_has_dynamic_multidim_index(&rhs)); + } + if let Ok(call) = get_field(expr, "BuiltinCall").or_else(|_| get_field(expr, "FunctionCall")) + && let Ok(args) = get_field(&call, "args") + { + return list_any(&args, |arg| expr_has_dynamic_multidim_index(&arg)); + } + false +} + /// Extract the derivative RHS from a single equation if it contains `der(state_name)`. /// Helper for `ode_rhs_for_state_function`; decomposes MLS B.1a residual form. /// @@ -989,6 +1335,12 @@ fn render_direct_rhs_for_lhs( cfg: &ExprConfig, ) -> Result, minijinja::Error> { if is_var_ref_of(lhs, var_name) { + if let Some((_base_name, index)) = parse_indexed_ref(var_name) { + if rhs_is_structurally_indexed_var_ref(rhs) { + return render_expression(rhs, cfg).map(Some); + } + return render_array_expr_at_index_or_scalar_checked(rhs, index, cfg); + } return render_expression(rhs, cfg).map(Some); } if let Some(index) = array_lhs_element_index(lhs, var_name) { @@ -1000,6 +1352,21 @@ fn render_direct_rhs_for_lhs( no_render_match() } +fn rhs_is_structurally_indexed_var_ref(rhs: &Value) -> bool { + let Ok(var_ref) = get_field(rhs, "VarRef") else { + return false; + }; + if get_field(&var_ref, "subscripts") + .ok() + .and_then(|subscripts| subscripts.len()) + .is_some_and(|len| len > 0) + { + return false; + } + let raw_name = var_ref_base_name(&var_ref); + raw_name.contains('[') +} + fn array_lhs_element_index(lhs: &Value, var_name: &str) -> Option { if get_field(lhs, "Array").is_err() { return None; @@ -1064,7 +1431,9 @@ fn find_algebraic_rhs_assignment( { // branches is a list of [condition, expression] pairs. let items: Vec<_> = branch_array.take(2).collect(); - if let Some(update_expr) = items.get(1) { + if items.first().is_some_and(is_sample_guard) + && let Some(update_expr) = items.get(1) + { return render_expression(update_expr, cfg).map(Some); } } @@ -1103,6 +1472,16 @@ fn find_algebraic_rhs_assignment( render_expression(&elem, cfg).map(Some) } +fn is_sample_guard(expr: &Value) -> bool { + let Ok(builtin) = get_field(expr, "BuiltinCall") else { + return false; + }; + let Ok(function) = get_field(&builtin, "function") else { + return false; + }; + matches!(function.to_string().trim_matches('"'), "Sample" | "sample") +} + /// Try subtraction form: 0 = var - expr, 0 = expr - var, 0 = -(A - B) fn find_algebraic_rhs_subtraction( rhs: &Value, @@ -1430,9 +1809,52 @@ fn render_array_expr_at_index_checked( return render_builtin_array_expr_at_index_checked(&builtin, index, cfg); } + if let Ok(function_call) = get_field(expr, "FunctionCall") { + return render_function_array_expr_at_index_checked(&function_call, index, cfg); + } + no_render_match() } +fn render_function_array_expr_at_index_checked( + function_call: &Value, + index: usize, + cfg: &ExprConfig, +) -> Result, minijinja::Error> { + let Ok(name) = get_field(function_call, "name") else { + return no_render_match(); + }; + if render_serialized_name(&name) != "linspace" { + return no_render_match(); + } + let Ok(args) = get_field(function_call, "args") else { + return no_render_match(); + }; + if args.len().unwrap_or(0) != 3 { + return no_render_match(); + } + let Ok(start) = args.get_item(&Value::from(0)) else { + return no_render_match(); + }; + let Ok(stop) = args.get_item(&Value::from(1)) else { + return no_render_match(); + }; + let Ok(count) = args.get_item(&Value::from(2)) else { + return no_render_match(); + }; + + let start_rendered = render_expression(&start, cfg)?; + if index == 1 { + return Ok(Some(start_rendered)); + } + let stop_rendered = render_expression(&stop, cfg)?; + let count_rendered = render_expression(&count, cfg)?; + let offset = index - 1; + Ok(Some(format!( + "(({start_rendered}) + ({offset}.0 * ((({stop_rendered}) - ({start_rendered})) / (({count_rendered}) - 1.0))))" + ))) +} + fn render_array_expr_at_index_or_scalar_checked( expr: &Value, index: usize, @@ -1888,6 +2310,15 @@ fn serialized_name_leaf(value: &Value) -> Option { None } +fn source_component_ref_name(component_ref: &Value) -> RenderResult { + let cfg = ExprConfig { + one_based_index: true, + sanitize_dots: false, + ..ExprConfig::default() + }; + super::render_stmt::render_component_ref(component_ref, &cfg) +} + fn parse_indexed_ref(name: &str) -> Option<(String, usize)> { let trimmed = name.trim_matches('"'); let (base, subscripts) = rumoca_core::split_trailing_subscript_suffix(trimmed)?; diff --git a/crates/rumoca-phase-codegen/src/codegen/render_expr.rs b/crates/rumoca-phase-codegen/src/codegen/render_expr.rs index 62797e25d..7645d5355 100644 --- a/crates/rumoca-phase-codegen/src/codegen/render_expr.rs +++ b/crates/rumoca-phase-codegen/src/codegen/render_expr.rs @@ -42,6 +42,9 @@ pub(crate) fn is_variant(value: &Value, name: &str) -> bool { /// Recursively render an expression to a string. pub(crate) fn render_expression(expr: &Value, cfg: &ExprConfig) -> RenderResult { + if let Ok(inner) = get_field(expr, "expr") { + return render_expression(&inner, cfg); + } if let Ok(binary) = get_field(expr, "Binary") { return render_binary(&binary, cfg); } @@ -81,6 +84,18 @@ pub(crate) fn render_expression(expr: &Value, cfg: &ExprConfig) -> RenderResult if let Ok(field_access) = get_field(expr, "FieldAccess") { return render_field_access(&field_access, cfg); } + if let Ok(named) = get_field(expr, "NamedArgument") { + return render_named_argument(&named, cfg); + } + if let Ok(component_ref) = get_field(expr, "ComponentReference") { + return super::render_stmt::render_component_ref(&component_ref, cfg); + } + if let Ok(component_ref) = get_field(expr, "component_ref") { + return super::render_stmt::render_component_ref(&component_ref, cfg); + } + if get_field(expr, "parts").is_ok() { + return super::render_stmt::render_component_ref(expr, cfg); + } // Unit variants (e.g. Empty) serialize as plain strings, not objects, // so get_field() won't match them — check string representation instead. let s = expr.to_string(); @@ -90,6 +105,14 @@ pub(crate) fn render_expression(expr: &Value, cfg: &ExprConfig) -> RenderResult Err(render_err(format!("unhandled Expression variant: {expr}"))) } +fn render_named_argument(named: &Value, cfg: &ExprConfig) -> RenderResult { + let value = get_field(named, "value") + .or_else(|_| get_field(named, "expr")) + .or_else(|_| get_field(named, "arg")) + .map_err(|_| render_err("NamedArgument expression missing value field"))?; + render_expression(&value, cfg) +} + fn render_binary(binary: &Value, cfg: &ExprConfig) -> RenderResult { let lhs = get_field(binary, "lhs") .and_then(|v| render_expression(&v, cfg)) @@ -222,18 +245,26 @@ fn render_var_ref(var_ref: &Value, cfg: &ExprConfig) -> RenderResult { return render_expression(&relation, cfg); } + let Some(subs) = get_field(var_ref, "subscripts").ok() else { + return render_unsubscripted_var_ref_with_source(&raw_name, &source_ref, cfg); + }; + let len = subs + .len() + .ok_or_else(|| render_err("VarRef subscripts field is not a sequence"))?; + if len == 0 { + return render_unsubscripted_var_ref_with_source(&raw_name, &source_ref, cfg); + } + + let all_static = (0..len).all(|i| { + subs.get_item(&Value::from(i)) + .ok() + .and_then(|sub| get_field(&sub, "Index").ok()) + .and_then(|idx| subscript_index_value(&idx).ok()) + .is_some() + }); + let subscripts = render_subscripts(var_ref, cfg)?; - if subscripts.is_empty() { - if let Some((_, value)) = cfg - .substitutions - .iter() - .rev() - .find(|(name, _)| name == &raw_name) - { - return Ok(value.clone()); - } - super::emitted_symbol(&raw_name, cfg) - } else if cfg.subscript_underscore { + if cfg.subscript_underscore && all_static { // Underscore style: x[1] -> x_1, x[1,2] -> x_1_2. let compact_subscripts = subscripts.replace(' ', ""); let source_ref = format!("{raw_name}[{compact_subscripts}]"); @@ -242,18 +273,150 @@ fn render_var_ref(var_ref: &Value, cfg: &ExprConfig) -> RenderResult { } let name = super::emitted_symbol(&raw_name, cfg)?; Ok(format!("{}_{}", name, compact_subscripts.replace(',', "_"))) + } else if cfg.subscript_underscore { + if len > 1 { + return Err(render_err(format!( + "dynamic multi-dimensional array access is not supported for C aliases: {var_ref}" + ))); + } + let name = super::emitted_symbol(&raw_name, cfg)?; + let pointer_subscripts = render_pointer_subscripts(&subs, cfg)?; + Ok(format!("{}[{}]", name, pointer_subscripts)) } else { let name = super::emitted_symbol(&raw_name, cfg)?; Ok(format!("{}[{}]", name, subscripts)) } } +fn render_unsubscripted_var_ref(raw_name: &str, cfg: &ExprConfig) -> RenderResult { + if let Some((_, value)) = cfg + .substitutions + .iter() + .rev() + .find(|(name, _)| name == raw_name) + { + return Ok(value.clone()); + } + if let Some(symbol) = super::lookup_symbol_value(cfg.symbols.as_ref(), raw_name) { + return Ok(symbol); + } + if let Some(source_name) = one_based_serialized_component_name(raw_name) + && let Some(symbol) = super::lookup_symbol_value(cfg.symbols.as_ref(), &source_name) + { + return Ok(symbol); + } + if cfg.symbols.is_none() + && let Some(source_name) = zero_component_name_to_one_based(raw_name) + { + return super::emitted_symbol(&source_name, cfg); + } + super::emitted_symbol(raw_name, cfg) +} + +fn render_unsubscripted_var_ref_with_source( + raw_name: &str, + source_ref: &str, + cfg: &ExprConfig, +) -> RenderResult { + // DAE `VarName` text is already the canonical one-based source reference. + // Prefer its exact symbol before consulting an attached AST component_ref, + // whose serialized indices are zero-based and require normalization. + if let Some(symbol) = super::lookup_symbol_value(cfg.symbols.as_ref(), raw_name) { + return Ok(symbol); + } + if let Some(scope) = cfg.source_scope.as_deref() + && let Some(rest) = source_ref.strip_prefix(scope) + && rest.starts_with('.') + && super::lookup_symbol_value(cfg.symbols.as_ref(), source_ref).is_some() + { + return Ok(super::sanitize_name(source_ref)); + } + if source_ref != raw_name + && let Some(symbol) = super::lookup_symbol_value(cfg.symbols.as_ref(), source_ref) + { + return Ok(symbol); + } + if source_ref == raw_name + && !rumoca_core::has_top_level_dot(raw_name) + && let Some(scope) = cfg.source_scope.as_deref() + { + let scoped_ref = format!("{scope}.{raw_name}"); + if super::lookup_symbol_value(cfg.symbols.as_ref(), &scoped_ref).is_some() { + return Ok(super::sanitize_name(&scoped_ref)); + } + } + render_unsubscripted_var_ref(raw_name, cfg) +} + +fn one_based_serialized_component_name(name: &str) -> Option { + canonicalize_serialized_component_indices(name, false) +} + +fn zero_component_name_to_one_based(name: &str) -> Option { + canonicalize_serialized_component_indices(name, true) +} + +fn canonicalize_serialized_component_indices(name: &str, zeros_only: bool) -> Option { + let bytes = name.as_bytes(); + let mut rendered = String::with_capacity(name.len()); + let mut cursor = 0; + let mut changed = false; + + while cursor < bytes.len() { + if bytes[cursor] != b'[' { + let next = name[cursor..] + .find('[') + .map(|offset| cursor + offset) + .unwrap_or(bytes.len()); + rendered.push_str(&name[cursor..next]); + cursor = next; + continue; + } + + let index_start = cursor + 1; + let mut index_end = index_start; + while index_end < bytes.len() && bytes[index_end].is_ascii_digit() { + index_end += 1; + } + if index_end == index_start || index_end >= bytes.len() || bytes[index_end] != b']' { + rendered.push('['); + cursor += 1; + continue; + } + + let index_text = &name[index_start..index_end]; + let Ok(index) = index_text.parse::() else { + rendered.push_str(&name[cursor..=index_end]); + cursor = index_end + 1; + continue; + }; + if zeros_only && index != 0 { + rendered.push_str(&name[cursor..=index_end]); + } else { + rendered.push('['); + rendered.push_str(&(index + 1).to_string()); + rendered.push(']'); + changed = true; + } + cursor = index_end + 1; + } + + changed.then_some(rendered) +} + fn var_ref_source_ref(raw_name: &str, var_ref: &Value) -> RenderResult { + let base_name = if let Ok(name) = get_field(var_ref, "name") + && let Ok(component_ref) = get_field(&name, "component_ref") + { + render_component_ref_source_name(&component_ref, "VarRef", "name")? + } else { + raw_name.to_string() + }; let subscripts = render_source_subscripts(var_ref)?; if subscripts.is_empty() { - Ok(raw_name.to_string()) + Ok(base_name) } else { - Ok(format!("{raw_name}[{subscripts}]")) + Ok(format!("{base_name}[{subscripts}]")) } } @@ -346,6 +509,25 @@ fn render_name_field(value: &Value, field: &str, context: &str) -> RenderResult Ok(rendered) } +fn render_component_ref_source_name( + component_ref: &Value, + context: &str, + field: &str, +) -> RenderResult { + let cfg = ExprConfig { + one_based_index: true, + sanitize_dots: false, + ..ExprConfig::default() + }; + let rendered = super::render_stmt::render_component_ref(component_ref, &cfg)?; + if rendered.is_empty() { + return Err(render_err(format!( + "{context} '{field}' component_ref resolved to an empty name" + ))); + } + Ok(rendered) +} + pub(crate) fn render_serialized_name(value: &Value) -> String { if let Ok(name) = get_field(value, "name") { return render_serialized_name(&name); @@ -385,6 +567,45 @@ fn render_subscripts(var_ref: &Value, cfg: &ExprConfig) -> RenderResult { Ok(sub_strs.join(", ")) } +fn render_pointer_subscripts(subs: &Value, cfg: &ExprConfig) -> RenderResult { + let len = subs + .len() + .ok_or_else(|| render_err("VarRef subscripts field is not a sequence"))?; + let index_cfg = ExprConfig { + one_based_index: false, + subscript_underscore: false, + ..cfg.clone() + }; + let mut sub_strs = render_vec_with_capacity(len, "VarRef pointer subscript count")?; + for i in 0..len { + let sub = subs + .get_item(&Value::from(i)) + .map_err(|err| render_err(format!("VarRef subscript {i} is inaccessible: {err}")))?; + if sub.is_undefined() || sub.is_none() { + return Err(render_err(format!("VarRef subscript {i} is missing"))); + } + sub_strs.push(render_pointer_subscript(&sub, &index_cfg)?); + } + Ok(sub_strs.join(", ")) +} + +fn render_pointer_subscript(sub: &Value, cfg: &ExprConfig) -> RenderResult { + if let Ok(idx) = get_field(sub, "Index") { + return Ok(format!("{}", subscript_index_value(&idx)? - 1)); + } + if get_field(sub, "Colon").is_ok() { + return Err(render_err( + "slice subscripts are not supported in C array aliases", + )); + } + if let Ok(expr) = get_field(sub, "Expr") { + let expr = get_field(&expr, "expr").unwrap_or(expr); + let rendered = render_expression(&expr, cfg)?; + return Ok(format!("(({}) - 1)", rendered)); + } + Err(render_err(format!("unhandled Subscript variant: {sub}"))) +} + pub(crate) fn render_subscript(sub: &Value, cfg: &ExprConfig) -> RenderResult { if let Ok(idx) = get_field(sub, "Index") { let val = subscript_index_value(&idx)?; @@ -787,7 +1008,7 @@ pub(crate) fn render_args(call: &Value, cfg: &ExprConfig) -> RenderResult { for i in 0..len { match args.get_item(&Value::from(i)) { Ok(arg) if !arg.is_undefined() && !arg.is_none() => { - arg_strs.push(render_expression(&arg, cfg)?); + arg_strs.push(render_function_argument(&arg, cfg)?); } Ok(_) => return Err(render_err(format!("function argument {i} is missing"))), Err(err) => { @@ -801,6 +1022,19 @@ pub(crate) fn render_args(call: &Value, cfg: &ExprConfig) -> RenderResult { Ok(arg_strs.join(", ")) } +fn render_function_argument(arg: &Value, cfg: &ExprConfig) -> RenderResult { + if let Ok(call) = get_field(arg, "FunctionCall") + && let Ok(name) = render_name_field(&call, "name", "FunctionCall") + && name.starts_with("__rumoca_named_arg__.") + { + let args = get_field(&call, "args") + .map_err(|err| render_err(format!("named function argument missing args: {err}")))?; + let value = required_arg(&args, 0, "named function argument")?; + return render_expression(&value, cfg); + } + render_expression(arg, cfg) +} + fn render_literal(literal: &Value, cfg: &ExprConfig) -> RenderResult { let literal_value = get_field(literal, "value").unwrap_or_else(|_| literal.clone()); if let Ok(real) = get_field(&literal_value, "Real") { @@ -1107,9 +1341,13 @@ fn try_unroll_c_comprehension( /// since they access via pointer/array (unlike VarRef underscore subscripts /// which are 1-based naming). fn render_index(index: &Value, cfg: &ExprConfig) -> RenderResult { - let base = get_field(index, "base") - .and_then(|v| render_expression(&v, cfg)) - .map_err(|_| render_err("Index missing 'base' field"))?; + let base_value = + get_field(index, "base").map_err(|_| render_err("Index missing 'base' field"))?; + let base = render_expression(&base_value, cfg).map_err(|err| { + render_err(format!( + "Index base render failed: {err}; base: {base_value}" + )) + })?; let subs = get_field(index, "subscripts") .map_err(|_| render_err("Index missing 'subscripts' field"))?; let len = subs.len().unwrap_or(0); @@ -1134,15 +1372,111 @@ fn render_index(index: &Value, cfg: &ExprConfig) -> RenderResult { /// Render a field access expression as `base.field`. fn render_field_access(fa: &Value, cfg: &ExprConfig) -> RenderResult { - let base = get_field(fa, "base") - .and_then(|v| render_expression(&v, cfg)) - .map_err(|_| render_err("FieldAccess missing 'base' field"))?; + if let Some(symbol) = render_indexed_field_symbol(fa, cfg)? { + return Ok(symbol); + } + let base_value = + get_field(fa, "base").map_err(|_| render_err("FieldAccess missing 'base' field"))?; + let base = render_expression(&base_value, cfg).map_err(|err| { + render_err(format!( + "FieldAccess base render failed: {err}; base: {base_value}" + )) + })?; let field = get_field(fa, "field") .map(|v| v.to_string()) .map_err(|_| render_err("FieldAccess missing 'field'"))?; Ok(format!("{base}.{field}")) } +fn render_indexed_field_symbol( + fa: &Value, + cfg: &ExprConfig, +) -> Result, minijinja::Error> { + let Ok(base) = get_field(fa, "base") else { + return no_expr_render_match(); + }; + let Ok(index) = get_field(&base, "Index") else { + return no_expr_render_match(); + }; + let Ok(index_base) = get_field(&index, "base") else { + return no_expr_render_match(); + }; + let Ok(var_ref) = get_field(&index_base, "VarRef") else { + return no_expr_render_match(); + }; + let raw_name = render_name_field(&var_ref, "name", "indexed FieldAccess base")?; + let Ok(subscripts) = get_field(&index, "subscripts") else { + return no_expr_render_match(); + }; + let Some(len) = subscripts.len() else { + return no_expr_render_match(); + }; + let mut values = Vec::with_capacity(len); + for i in 0..len { + let sub = subscripts.get_item(&Value::from(i))?; + let Some(value) = static_subscript_value(&sub)? else { + return no_expr_render_match(); + }; + values.push(value.to_string()); + } + let field = get_field(fa, "field")?.to_string(); + let source_ref = format!("{raw_name}[{}].{field}", values.join(",")); + Ok(super::lookup_symbol_value( + cfg.symbols.as_ref(), + &source_ref, + )) +} + +fn static_subscript_value(sub: &Value) -> Result, minijinja::Error> { + if let Ok(idx) = get_field(sub, "Index") { + return subscript_index_value(&idx).map(Some); + } + if let Ok(expr) = get_field(sub, "Expr") { + let expr = get_field(&expr, "expr").unwrap_or(expr); + return static_integer_expr_value(&expr); + } + no_expr_render_match() +} + +fn static_integer_expr_value(expr: &Value) -> Result, minijinja::Error> { + if let Ok(literal) = get_field(expr, "Literal") { + let literal_value = get_field(&literal, "value").unwrap_or(literal); + if let Ok(integer) = get_field(&literal_value, "Integer") { + return integer + .as_i64() + .ok_or_else(|| render_err("integer literal is not an i64")) + .map(Some); + } + } + if let Ok(binary) = get_field(expr, "Binary") { + return static_binary_integer_expr_value(&binary); + } + no_expr_render_match() +} + +fn static_binary_integer_expr_value(binary: &Value) -> Result, minijinja::Error> { + let lhs = get_field(binary, "lhs")?; + let rhs = get_field(binary, "rhs")?; + let Some(lhs) = static_integer_expr_value(&lhs)? else { + return no_expr_render_match(); + }; + let Some(rhs) = static_integer_expr_value(&rhs)? else { + return no_expr_render_match(); + }; + let op = get_field(binary, "op")?; + if is_variant(&op, "Add") || is_variant(&op, "AddElem") { + return Ok(Some(lhs + rhs)); + } + if is_variant(&op, "Sub") || is_variant(&op, "SubElem") { + return Ok(Some(lhs - rhs)); + } + no_expr_render_match() +} + +fn no_expr_render_match() -> Result, minijinja::Error> { + Ok(Option::None) +} + #[cfg(test)] mod tests { use super::{render_c_float_literal, render_expression}; diff --git a/crates/rumoca-phase-codegen/src/codegen/render_solve.rs b/crates/rumoca-phase-codegen/src/codegen/render_solve.rs index b9795b422..6fdf747f0 100644 --- a/crates/rumoca-phase-codegen/src/codegen/render_solve.rs +++ b/crates/rumoca-phase-codegen/src/codegen/render_solve.rs @@ -752,58 +752,26 @@ fn solve_op_expr( dialect: SolveRowDialect, regs: &[String], ) -> Result { + if let Some(effect) = solve_load_or_move_effect(op, cfg, dialect, regs)? { + return Ok(effect); + } if let Ok(value) = get_field(op, "Const") { let dst = solve_field_usize(&value, "dst")?; let expr = dialect.format_const(solve_const_value_string(&value, dialect.infinity())?); return Ok(SolveOpEffect::Compute { dst, expr }); } - if let Ok(value) = get_field(op, "LoadTime") { - let dst = solve_field_usize(&value, "dst")?; - return Ok(SolveOpEffect::Compute { - dst, - expr: cfg.time.clone(), - }); - } - if let Ok(value) = get_field(op, "LoadY") { - let dst = solve_field_usize(&value, "dst")?; - let index = solve_field_usize(&value, "index")?; - return Ok(SolveOpEffect::Compute { - dst, - expr: cfg.y_access(index), - }); - } - if let Ok(value) = get_field(op, "LoadP") { - let dst = solve_field_usize(&value, "dst")?; - let index = solve_field_usize(&value, "index")?; - return Ok(SolveOpEffect::Compute { - dst, - expr: cfg.p_access(index), - }); - } if let Ok(value) = get_field(op, "LoadIndexedP") { return solve_indexed_effect(&value, cfg, dialect, regs, false); } if let Ok(value) = get_field(op, "LoadIndexedSeed") { return solve_indexed_effect(&value, cfg, dialect, regs, true); } - if let Ok(value) = get_field(op, "LoadSeed") { - let dst = solve_field_usize(&value, "dst")?; - let index = solve_field_usize(&value, "index")?; - let Some(seed) = cfg.seed_access(index) else { - return Err(render_err( - "LoadSeed requires a `seed` access pattern in solve-row C output", - )); - }; - return Ok(SolveOpEffect::Compute { dst, expr: seed }); - } - if let Ok(value) = get_field(op, "Move") { - let dst = solve_field_usize(&value, "dst")?; - let src = solve_reg(regs, solve_field_usize(&value, "src")?)?; - return Ok(SolveOpEffect::Compute { dst, expr: src }); - } if let Ok(value) = get_field(op, "LinearSolveComponent") { return solve_linsolve_effect(&value, dialect, regs); } + if let Ok(value) = get_field(op, "ExternalCall") { + return solve_external_call_effect(&value, regs); + } if let Ok(value) = get_field(op, "Unary") { let dst = solve_field_usize(&value, "dst")?; let op = solve_variant_name(&get_field(&value, "op")?)?; @@ -850,6 +818,97 @@ fn solve_op_expr( Err(render_err(format!("unsupported solve LinearOp: {op}"))) } +fn solve_load_or_move_effect( + op: &Value, + cfg: &SolveRowCConfig, + _dialect: SolveRowDialect, + regs: &[String], +) -> Result, minijinja::Error> { + if let Ok(value) = get_field(op, "LoadTime") { + let dst = solve_field_usize(&value, "dst")?; + return Ok(Some(SolveOpEffect::Compute { + dst, + expr: cfg.time.clone(), + })); + } + if let Ok(value) = get_field(op, "LoadY") { + let dst = solve_field_usize(&value, "dst")?; + let index = solve_field_usize(&value, "index")?; + return Ok(Some(SolveOpEffect::Compute { + dst, + expr: cfg.y_access(index), + })); + } + if let Ok(value) = get_field(op, "LoadP") { + let dst = solve_field_usize(&value, "dst")?; + let index = solve_field_usize(&value, "index")?; + return Ok(Some(SolveOpEffect::Compute { + dst, + expr: cfg.p_access(index), + })); + } + if let Ok(value) = get_field(op, "LoadSeed") { + let dst = solve_field_usize(&value, "dst")?; + let index = solve_field_usize(&value, "index")?; + let Some(seed) = cfg.seed_access(index) else { + return Err(render_err( + "LoadSeed requires a `seed` access pattern in solve-row C output", + )); + }; + return Ok(Some(SolveOpEffect::Compute { dst, expr: seed })); + } + if let Ok(value) = get_field(op, "Move") { + let dst = solve_field_usize(&value, "dst")?; + let src = solve_reg(regs, solve_field_usize(&value, "src")?)?; + return Ok(Some(SolveOpEffect::Compute { dst, expr: src })); + } + Ok(no_solve_load_or_move_effect()) +} + +fn no_solve_load_or_move_effect() -> Option { + None +} + +fn solve_external_call_effect( + value: &Value, + regs: &[String], +) -> Result { + let dst = solve_field_usize(value, "dst")?; + let function = solve_variant_name(&get_field(value, "function")?)?; + let output_index = solve_field_usize(value, "output_index")?; + let arg_count = solve_field_usize(value, "arg_count")?; + let args = + get_field(value, "args").map_err(|_| render_err("ExternalCall missing args field"))?; + let args_len = args + .len() + .ok_or_else(|| render_err("ExternalCall args field is not a sequence"))?; + if arg_count > args_len { + return Err(render_err(format!( + "ExternalCall arg_count {arg_count} exceeds args length {args_len}" + ))); + } + let mut rendered_args = Vec::with_capacity(arg_count); + for i in 0..arg_count { + let reg = args + .get_item(&Value::from(i)) + .map_err(|err| render_err(format!("ExternalCall arg {i} inaccessible: {err}")))? + .as_usize() + .ok_or_else(|| render_err(format!("ExternalCall arg {i} is not a register")))?; + rendered_args.push(solve_reg(regs, reg)?); + } + let args_expr = if rendered_args.is_empty() { + "NULL".to_string() + } else { + format!("(double[]){{{}}}", rendered_args.join(", ")) + }; + Ok(SolveOpEffect::Compute { + dst, + expr: format!( + "RUMOCA_SOLVE_EXTERNAL_CALL(\"{function}\", {output_index}, {arg_count}, {args_expr})" + ), + }) +} + fn render_solve_op_for( op: &Value, cfg: &SolveRowCConfig, @@ -1251,12 +1310,13 @@ impl SolveRowDialect { rhs_start, n, output_offset: 0, + output_targets: SolveOutputTargets::DenseOffset(0), }; - let (matrix_count, rhs_count, _) = validate_linsolve_render_shape(shape)?; + let (matrix_count, rhs_count, _) = validate_linsolve_render_shape(&shape)?; let matrix = - render_linsolve_register_array(regs, shape, matrix_count, LinSolveOperand::Matrix)? + render_linsolve_register_array(regs, &shape, matrix_count, LinSolveOperand::Matrix)? .join(", "); - let rhs = render_linsolve_register_array(regs, shape, rhs_count, LinSolveOperand::Rhs)? + let rhs = render_linsolve_register_array(regs, &shape, rhs_count, LinSolveOperand::Rhs)? .join(", "); match self { Self::C => Ok(format!( @@ -1314,7 +1374,7 @@ enum LinSolveOperand { fn render_linsolve_register_array( regs: &[String], - shape: LinSolveRenderShape, + shape: &LinSolveRenderShape, count: usize, operand: LinSolveOperand, ) -> Result, minijinja::Error> { diff --git a/crates/rumoca-phase-codegen/src/codegen/render_solve/dense_solve_render.rs b/crates/rumoca-phase-codegen/src/codegen/render_solve/dense_solve_render.rs index c7a8919ab..678f622d6 100644 --- a/crates/rumoca-phase-codegen/src/codegen/render_solve/dense_solve_render.rs +++ b/crates/rumoca-phase-codegen/src/codegen/render_solve/dense_solve_render.rs @@ -330,59 +330,72 @@ fn render_sparse_matmul_cell_mlir( Ok(()) } -#[derive(Clone, Copy)] pub(in crate::codegen) struct LinSolveRenderShape { pub(in crate::codegen) matrix_start: usize, pub(in crate::codegen) rhs_start: usize, pub(in crate::codegen) n: usize, pub(in crate::codegen) output_offset: usize, + pub(in crate::codegen) output_targets: SolveOutputTargets, } const LIN_SOLVE_RENDER_ENUMERATION_LIMIT: usize = 1_000_000; impl LinSolveRenderShape { - pub(in crate::codegen) fn matrix_count(self) -> Result { + pub(in crate::codegen) fn matrix_count(&self) -> Result { let count = checked_linsolve_product(self.n, self.n, "LinSolve matrix element count")?; checked_linsolve_render_count(count, "LinSolve matrix element count") } - pub(in crate::codegen) fn rhs_count(self) -> Result { + pub(in crate::codegen) fn rhs_count(&self) -> Result { checked_linsolve_render_count(self.n, "LinSolve RHS element count") } - pub(in crate::codegen) fn output_count(self) -> Result { + pub(in crate::codegen) fn output_count(&self) -> Result { checked_linsolve_render_count(self.n, "LinSolve output count") } - pub(in crate::codegen) fn end_offset(self) -> Result { - checked_linsolve_sum( - self.output_offset, - self.output_count()?, - "LinSolve output end offset", - ) + pub(in crate::codegen) fn end_offset(&self) -> Result { + let mut end = self.output_offset; + for component in 0..self.output_count()? { + end = end.max(checked_linsolve_sum( + self.output_index(component)?, + 1, + "LinSolve output end offset", + )?); + } + Ok(end) } pub(in crate::codegen) fn output_index( - self, + &self, component: usize, ) -> Result { - checked_linsolve_sum(self.output_offset, component, "LinSolve output index") + self.output_targets.target_for(component) } - pub(in crate::codegen) fn matrix_reg(self, offset: usize) -> Result { + pub(in crate::codegen) fn matrix_reg(&self, offset: usize) -> Result { checked_linsolve_sum(self.matrix_start, offset, "LinSolve matrix register index") } - pub(in crate::codegen) fn rhs_reg(self, offset: usize) -> Result { + pub(in crate::codegen) fn rhs_reg(&self, offset: usize) -> Result { checked_linsolve_sum(self.rhs_start, offset, "LinSolve RHS register index") } } pub(in crate::codegen) fn validate_linsolve_render_shape( - shape: LinSolveRenderShape, + shape: &LinSolveRenderShape, ) -> Result<(usize, usize, usize), minijinja::Error> { let matrix_count = shape.matrix_count()?; let rhs_count = shape.rhs_count()?; + if let SolveOutputTargets::Explicit(indices) = &shape.output_targets + && indices.len() != shape.n + { + return Err(render_err(format!( + "LinSolve has {} components but {} output indices", + shape.n, + indices.len() + ))); + } let end_offset = shape.end_offset()?; Ok((matrix_count, rhs_count, end_offset)) } @@ -444,13 +457,17 @@ pub(in crate::codegen) fn render_linsolve_mlir_function( let matrix_start = solve_field_usize(&node, "matrix_start")?; let rhs_start = solve_field_usize(&node, "rhs_start")?; let n = solve_field_usize(&node, "n")?; + let output_indices = get_field(&node, "output_indices") + .unwrap_or_else(|_| Value::from_serialize(Vec::::new())); + let output_targets = linsolve_output_targets(output_indices, offset)?; let shape = LinSolveRenderShape { matrix_start, rhs_start, n, output_offset: offset, + output_targets, }; - let (matrix_count, rhs_count, end_offset) = validate_linsolve_render_shape(shape)?; + let (matrix_count, rhs_count, end_offset) = validate_linsolve_render_shape(&shape)?; let pfx = format!("ls{id}"); let mut out = format!(" // LinSolve {n}×{n} → out[{offset}..{end_offset}]\n"); @@ -503,6 +520,18 @@ pub(in crate::codegen) fn render_linsolve_mlir_function( Ok(out) } +fn linsolve_output_targets( + output_indices: Value, + output_offset: usize, +) -> Result { + match solve_output_targets(Some(output_indices))? { + SolveOutputTargets::Explicit(indices) if indices.is_empty() => { + Ok(SolveOutputTargets::DenseOffset(output_offset)) + } + targets => Ok(targets), + } +} + fn required_usize_arg(value: &Value, context: &'static str) -> Result { value .as_usize() diff --git a/crates/rumoca-phase-codegen/src/codegen/render_solve/template_partition.rs b/crates/rumoca-phase-codegen/src/codegen/render_solve/template_partition.rs index bd5970463..05a559a61 100644 --- a/crates/rumoca-phase-codegen/src/codegen/render_solve/template_partition.rs +++ b/crates/rumoca-phase-codegen/src/codegen/render_solve/template_partition.rs @@ -709,16 +709,23 @@ fn partition_node_for_template( scalar_row_index, output_cursor, count, + None, span, )?; } - solve::ComputeNode::LinSolve { n, span, .. } => { + solve::ComputeNode::LinSolve { + n, + output_indices, + span, + .. + } => { let span = required_compute_node_span(*span, "linsolve scalar fallback")?; push_multi_output_tensor_fallback_program( partition, scalar_row_index, output_cursor, *n, + (!output_indices.is_empty()).then_some(output_indices), span, )?; } @@ -843,6 +850,7 @@ fn push_multi_output_tensor_fallback_program( scalar_row_index: &mut usize, output_cursor: &mut usize, count: usize, + explicit_output_indices: Option<&[usize]>, span: rumoca_core::Span, ) -> Result<(), rumoca_eval_solve::ScalarizeError> { let next_scalar_row_index = rumoca_eval_solve::checked_contiguous_output_count( @@ -851,24 +859,47 @@ fn push_multi_output_tensor_fallback_program( "scalar fallback rows", span, )?; - let next_output_cursor = rumoca_eval_solve::checked_contiguous_output_count( + let output_indices = match explicit_output_indices { + Some(indices) => { + if indices.len() != count { + return Err(rumoca_eval_solve::ScalarizeError::ShapeContract { + message: format!( + "LinSolve scalar fallback has {count} outputs but {} output indices", + indices.len() + ), + span: Some(span), + }); + } + indices.to_vec() + } + None => { + let end = rumoca_eval_solve::checked_contiguous_output_count( + *output_cursor, + count, + "scalar fallback output", + span, + )?; + (*output_cursor..end).collect() + } + }; + let next_output_cursor = (*output_cursor).max(rumoca_eval_solve::checked_tensor_output_count( + &output_indices, *output_cursor, - count, "scalar fallback output", span, - )?; + )?); reserve_partition_capacity( &mut partition.scalar_fallback_rows, count, "scalar fallback template row count", Some(span), )?; - for offset in 0..count { + for (offset, output_index) in output_indices.into_iter().enumerate() { partition .scalar_fallback_rows .push(RenderScalarFallbackRow { row_index: *scalar_row_index, - output_index: *output_cursor + offset, + output_index, output_ordinal: offset, }); } diff --git a/crates/rumoca-phase-codegen/src/codegen/render_solve_tests.rs b/crates/rumoca-phase-codegen/src/codegen/render_solve_tests.rs index c904dc990..914b728f9 100644 --- a/crates/rumoca-phase-codegen/src/codegen/render_solve_tests.rs +++ b/crates/rumoca-phase-codegen/src/codegen/render_solve_tests.rs @@ -330,6 +330,7 @@ fn linsolve_render_shape_rejects_matrix_count_overflow() { rhs_start: 0, n: usize::MAX, output_offset: 0, + output_targets: SolveOutputTargets::DenseOffset(0), }; let err = shape @@ -340,6 +341,42 @@ fn linsolve_render_shape_rejects_matrix_count_overflow() { assert!(err.contains("LinSolve matrix element count overflows host index range")); } +#[test] +fn native_mlir_linsolve_renderer_preserves_noncontiguous_output_indices() { + let node = solve::ComputeNode::LinSolve { + setup_ops: vec![ + solve::LinearOp::Const { dst: 0, value: 2.0 }, + solve::LinearOp::Const { dst: 1, value: 0.0 }, + solve::LinearOp::Const { dst: 2, value: 0.0 }, + solve::LinearOp::Const { dst: 3, value: 4.0 }, + solve::LinearOp::Const { dst: 4, value: 8.0 }, + solve::LinearOp::Const { + dst: 5, + value: 20.0, + }, + ], + matrix_start: 0, + rhs_start: 4, + n: 2, + next_reg: 6, + output_indices: vec![0, 2], + metadata: Default::default(), + span: fixture_span("native_mlir_noncontiguous_linsolve.mo"), + }; + + let node = Value::from_serialize(node); + let rendered = render_linsolve_mlir_function( + get_field(&node, "LinSolve").expect("LinSolve fixture should serialize as its inner node"), + Value::from(7usize), + Value::from(0usize), + ) + .expect("native MLIR LinSolve renderer should accept schema-v17 output indices"); + + assert!(rendered.contains("%ls7_oi0 = arith.constant 0 : index")); + assert!(rendered.contains("%ls7_oi1 = arith.constant 2 : index")); + assert!(!rendered.contains("%ls7_oi1 = arith.constant 1 : index")); +} + #[test] fn solve_read_counts_cover_write_only_registers() { let op = Value::from_serialize(solve::LinearOp::Const { dst: 5, value: 1.0 }); diff --git a/crates/rumoca-phase-codegen/src/codegen/render_stmt.rs b/crates/rumoca-phase-codegen/src/codegen/render_stmt.rs index c1031b47f..36affb6da 100644 --- a/crates/rumoca-phase-codegen/src/codegen/render_stmt.rs +++ b/crates/rumoca-phase-codegen/src/codegen/render_stmt.rs @@ -146,7 +146,7 @@ pub(crate) fn render_statement(stmt: &Value, cfg: &ExprConfig, indent: &str) -> if let Some(s) = stmt.as_str() { return match s { "Empty" => Ok(String::new()), - "Return" => Ok(format!("{indent}return")), + "Return" => Ok(render_return_statement(cfg, indent)), "Break" => Ok(format!("{indent}break")), _ => Err(render_err(format!("unhandled statement variant: {s}"))), }; @@ -163,7 +163,7 @@ pub(crate) fn render_statement(stmt: &Value, cfg: &ExprConfig, indent: &str) -> } if let Ok(ret) = get_field(stmt, "Return") { let _ = ret; - return Ok(format!("{indent}return")); + return Ok(render_return_statement(cfg, indent)); } if let Ok(brk) = get_field(stmt, "Break") { let _ = brk; @@ -194,6 +194,14 @@ pub(crate) fn render_statement(stmt: &Value, cfg: &ExprConfig, indent: &str) -> Err(render_err(format!("unhandled statement: {stmt}"))) } +fn render_return_statement(cfg: &ExprConfig, indent: &str) -> String { + if let Some(return_value) = &cfg.return_value { + format!("{indent}return {return_value};") + } else { + format!("{indent}return") + } +} + /// Render an assignment statement: comp := value fn render_assignment(assign: &Value, cfg: &ExprConfig, indent: &str) -> RenderResult { let comp_val = get_field(assign, "comp") @@ -751,7 +759,7 @@ fn render_assert_statement(assert: &Value, cfg: &ExprConfig, indent: &str) -> Re // ── Component reference rendering ──────────────────────────────────── /// Render an AST ComponentReference to a string. -fn render_component_ref(comp: &Value, cfg: &ExprConfig) -> RenderResult { +pub(super) fn render_component_ref(comp: &Value, cfg: &ExprConfig) -> RenderResult { if let Some(s) = comp.as_str() { return super::emitted_symbol(s, cfg); } @@ -782,6 +790,9 @@ fn render_component_ref(comp: &Value, cfg: &ExprConfig) -> RenderResult { "ComponentReference resolved to empty name: {comp}" ))); } + if let Some(symbol) = super::lookup_symbol_value(cfg.symbols.as_ref(), &joined) { + return Ok(symbol); + } super::emitted_symbol(&joined, cfg) } @@ -850,7 +861,12 @@ fn render_part_subscripts(part: &Value, cfg: &ExprConfig) -> RenderResult { /// Render an AST subscript. fn render_ast_subscript(sub: &Value, cfg: &ExprConfig) -> RenderResult { if let Ok(index) = get_field(sub, "Index") { - return Ok(subscript_index_value(&index)?.to_string()); + let value = subscript_index_value(&index)?; + return if cfg.one_based_index || cfg.subscript_underscore { + Ok((value + 1).to_string()) + } else { + Ok(value.to_string()) + }; } if let Ok(expr) = get_field(sub, "Expr") { return render_expression(&expr, cfg); diff --git a/crates/rumoca-phase-codegen/src/codegen/solve_lazy.rs b/crates/rumoca-phase-codegen/src/codegen/solve_lazy.rs index a4c1e8458..ee476b6fb 100644 --- a/crates/rumoca-phase-codegen/src/codegen/solve_lazy.rs +++ b/crates/rumoca-phase-codegen/src/codegen/solve_lazy.rs @@ -260,6 +260,7 @@ fn linsolve_value(node: Arc) -> Value { "rhs_start", "n", "next_reg", + "output_indices", "metadata", "span", ], @@ -270,6 +271,7 @@ fn linsolve_value(node: Arc) -> Value { rhs_start, n, next_reg, + output_indices, metadata, span, } = node.as_ref() @@ -282,6 +284,7 @@ fn linsolve_value(node: Arc) -> Value { "rhs_start" => Some(Value::from(*rhs_start)), "n" => Some(Value::from(*n)), "next_reg" => Some(Value::from(*next_reg)), + "output_indices" => Some(Value::from_serialize(output_indices)), "metadata" => Some(Value::from_serialize(metadata)), "span" => Some(Value::from_serialize(span)), _ => None, @@ -454,6 +457,8 @@ fn continuous_artifacts_value( pub(super) fn solve_value( problem: Arc, artifacts: Arc, + visible_names: Arc>, + visible_value_rows: Arc, ) -> Result { let continuous = continuous_value(problem.clone())?; let artifacts_value = artifacts_value(artifacts.clone())?; @@ -468,6 +473,8 @@ pub(super) fn solve_value( "events", "clocks", "artifacts", + "visible_names", + "visible_value_rows", ], move |k| match k { "schema_version" => Some(Value::from(problem.schema_version)), @@ -479,6 +486,8 @@ pub(super) fn solve_value( "initialization" => Some(Value::from_serialize(&problem.initialization)), "clocks" => Some(Value::from_serialize(&problem.clocks)), "artifacts" => Some(artifacts_value.clone()), + "visible_names" => Some(Value::from_serialize(visible_names.as_ref())), + "visible_value_rows" => Some(scalar_program_block_value(visible_value_rows.clone())), _ => None, }, )) diff --git a/crates/rumoca-phase-codegen/src/codegen/solve_renderer.rs b/crates/rumoca-phase-codegen/src/codegen/solve_renderer.rs index ee4d52811..009a31f00 100644 --- a/crates/rumoca-phase-codegen/src/codegen/solve_renderer.rs +++ b/crates/rumoca-phase-codegen/src/codegen/solve_renderer.rs @@ -88,7 +88,39 @@ impl SolveTemplateRenderer { value: std::sync::OnceLock::new(), }); Ok(Self { - context: solve_render_context_value_with_dae(problem, artifacts, None, dae_entry)?, + context: solve_render_context_value_with_dae( + problem, + artifacts, + None, + dae_entry, + Vec::new(), + solve::ScalarProgramBlock::default(), + )?, + guard_dae: Some(dae_model), + }) + } + + pub fn new_with_dae_and_visible_outputs( + problem: &solve::SolveProblem, + artifacts: &solve::SolveArtifacts, + dae_model: dae::Dae, + visible_names: Vec, + visible_value_rows: solve::ScalarProgramBlock, + ) -> Result { + let dae_model = std::sync::Arc::new(dae_model); + let dae_entry = Value::from_object(LazyDaeTemplateJson { + dae: dae_model.clone(), + value: std::sync::OnceLock::new(), + }); + Ok(Self { + context: solve_render_context_value_with_dae( + problem, + artifacts, + None, + dae_entry, + visible_names, + visible_value_rows, + )?, guard_dae: Some(dae_model), }) } @@ -123,7 +155,14 @@ pub(super) fn solve_render_context_value( artifacts: &solve::SolveArtifacts, model_name: Option<&str>, ) -> Result { - solve_render_context_value_with_dae(solve_problem, artifacts, model_name, Value::default()) + solve_render_context_value_with_dae( + solve_problem, + artifacts, + model_name, + Value::default(), + Vec::new(), + solve::ScalarProgramBlock::default(), + ) } fn solve_render_context_value_with_dae( @@ -131,6 +170,8 @@ fn solve_render_context_value_with_dae( artifacts: &solve::SolveArtifacts, model_name: Option<&str>, dae_entry: Value, + visible_names: Vec, + visible_value_rows: solve::ScalarProgramBlock, ) -> Result { // Lazy `solve` / `solve_derivative_nodes` (see `solve_lazy`): structural // fields serialize on demand and op lists materialize one op at a time, so a @@ -138,7 +179,12 @@ fn solve_render_context_value_with_dae( // materialization (`from_serialize(solve_problem)` alone was ~4.7 GB). let problem_arc = std::sync::Arc::new(solve_problem.clone()); let artifacts_arc = std::sync::Arc::new(artifacts.clone()); - let solve_value = super::solve_lazy::solve_value(problem_arc.clone(), artifacts_arc.clone())?; + let solve_value = super::solve_lazy::solve_value( + problem_arc.clone(), + artifacts_arc.clone(), + std::sync::Arc::new(visible_names), + std::sync::Arc::new(visible_value_rows), + )?; let artifacts_value = super::solve_lazy::artifacts_value(artifacts_arc.clone())?; let solve_blocks = solve_template_blocks_value(solve_problem, artifacts)?; let derivative_nodes = Value::from_object(LazyDerivativeNodesValue::new( diff --git a/crates/rumoca-phase-codegen/src/codegen/solve_template_context_tests.rs b/crates/rumoca-phase-codegen/src/codegen/solve_template_context_tests.rs index 1ee0d48af..e3686b33e 100644 --- a/crates/rumoca-phase-codegen/src/codegen/solve_template_context_tests.rs +++ b/crates/rumoca-phase-codegen/src/codegen/solve_template_context_tests.rs @@ -58,6 +58,7 @@ fn implicit_problem_with_artifacts() -> (solve::SolveProblem, solve::SolveArtifa solve::ComputeBlock::from_scalar_program_block(scalar_block(vec![row.clone()])); problem.continuous.implicit_rhs = solve::ComputeBlock::from_scalar_program_block(scalar_block(vec![row.clone()])); + problem.continuous.implicit_row_targets = vec![Some(solve::scalar_slot_y(0))]; let mut artifacts = solve::SolveArtifacts::default(); artifacts.continuous.implicit_jacobian_v_scalar = scalar_block(vec![row.clone()]); diff --git a/crates/rumoca-phase-codegen/src/codegen/strict_render_tests.rs b/crates/rumoca-phase-codegen/src/codegen/strict_render_tests.rs index 9195c998d..d8837df6e 100644 --- a/crates/rumoca-phase-codegen/src/codegen/strict_render_tests.rs +++ b/crates/rumoca-phase-codegen/src/codegen/strict_render_tests.rs @@ -279,6 +279,36 @@ fn test_render_function_args_reject_unrenderable_item() { ); } +#[test] +fn test_render_function_call_strips_named_argument_markers() { + let dae = dae::Dae::new(); + let template = r#" +{{ render_expr({ + "FunctionCall": { + "name": "f", + "args": [ + { + "FunctionCall": { + "name": "__rumoca_named_arg__.x", + "args": [{"VarRef": {"name": "r_N", "subscripts": []}}] + } + }, + { + "FunctionCall": { + "name": "__rumoca_named_arg__.x1", + "args": [{"Literal": {"value": {"Real": 0.0}}}] + } + } + ] + } +}, {}) }} +"#; + + let rendered = render_template(&dae, template).unwrap(); + + assert_eq!("f(r_N, 0.0)", rendered.trim()); +} + #[test] fn test_render_equation_rejects_unrenderable_explicit_rhs() { let dae = dae::Dae::new(); @@ -907,6 +937,40 @@ fn test_alg_rhs_rejects_unrenderable_direct_assignment_rhs() { ); } +#[test] +fn test_alg_rhs_with_dae_reads_fx_from_context_object() { + let dae_json = serde_json::json!({ + "f_x": [{ + "lhs": "y", + "rhs": { + "VarRef": { + "name": "u", + "subscripts": [] + } + } + }] + }); + let template = r#" +{% set cfg = {"power": "pow", "subscript_underscore": true} %} +{{ alg_rhs_for_var_with_dae("y", dae, cfg) }} +"#; + let rendered = render_template_with_dae_json(&dae_json, template).unwrap(); + + assert_eq!(rendered.trim(), "u"); +} + +#[test] +fn test_alg_rhs_with_dae_without_fx_uses_missing_equation_warning() { + let dae_json = serde_json::json!({}); + let template = r#" +{% set cfg = {"power": "pow", "subscript_underscore": true} %} +{{ alg_rhs_for_var_with_dae("y", dae, cfg) }} +"#; + let rendered = render_template_with_dae_json(&dae_json, template).unwrap(); + + assert!(rendered.contains("WARNING: no equation found for y")); +} + #[test] fn test_alg_rhs_rejects_unrenderable_subtraction_rhs() { let dae_json = serde_json::json!({ diff --git a/crates/rumoca-phase-codegen/src/errors.rs b/crates/rumoca-phase-codegen/src/errors.rs index 0a5d27767..870cff661 100644 --- a/crates/rumoca-phase-codegen/src/errors.rs +++ b/crates/rumoca-phase-codegen/src/errors.rs @@ -136,13 +136,13 @@ impl From for CodegenError { if let Some(source) = err.template_source() { let span = compute_line_span(source, line); return CodegenError::TemplateRenderError { - message: format!("{err:#}"), + message: format!("{err}"), src: NamedSource::new(tmpl_name, source.to_string()), span, }; } } - CodegenError::template(format!("{err:#}")) + CodegenError::template(format!("{err}")) } } diff --git a/crates/rumoca-phase-codegen/src/lib.rs b/crates/rumoca-phase-codegen/src/lib.rs index 08476d189..151bbef39 100644 --- a/crates/rumoca-phase-codegen/src/lib.rs +++ b/crates/rumoca-phase-codegen/src/lib.rs @@ -59,11 +59,12 @@ mod codegen; mod errors; pub use codegen::{ - CodegenInput, SolveTemplateRenderer, dae_template_json, render_ast_template, - render_ast_template_with_name, render_flat_template_with_name, render_solve_template_with_name, - render_template, render_template_file, render_template_for_input, - render_template_with_dae_json, render_template_with_dae_json_and_name, - render_template_with_name, render_template_with_name_for_input, + CodegenInput, DaeTemplateContext, SolveTemplateRenderer, dae_template_json, + render_ast_template, render_ast_template_with_name, render_flat_template_with_name, + render_solve_template_with_name, render_template, render_template_file, + render_template_for_input, render_template_with_dae_json, + render_template_with_dae_json_and_name, render_template_with_name, + render_template_with_name_for_input, }; pub use errors::CodegenError; diff --git a/crates/rumoca-phase-codegen/src/templates/fmi2/CMakeLists.txt.jinja b/crates/rumoca-phase-codegen/src/templates/fmi2/CMakeLists.txt.jinja index 2873fd60b..3d458d18e 100644 --- a/crates/rumoca-phase-codegen/src/templates/fmi2/CMakeLists.txt.jinja +++ b/crates/rumoca-phase-codegen/src/templates/fmi2/CMakeLists.txt.jinja @@ -17,6 +17,56 @@ else() message(FATAL_ERROR "Unsupported platform for FMU packaging") endif() +set(RUMOCA_EXTERNAL_INCLUDE_DIR "$ENV{RUMOCA_EXTERNAL_INCLUDE_DIR}" CACHE PATH "Directory containing Modelica external runtime headers for modelica:// include URIs") +set(RUMOCA_EXTERNAL_INCLUDE_DIRECTORIES_FILE "${CMAKE_CURRENT_LIST_DIR}/../resources/externalIncludeDirectories.txt") +if(EXISTS "${RUMOCA_EXTERNAL_INCLUDE_DIRECTORIES_FILE}") + file(STRINGS "${RUMOCA_EXTERNAL_INCLUDE_DIRECTORIES_FILE}" RUMOCA_EXTERNAL_INCLUDE_DIRECTORIES) + list(REMOVE_DUPLICATES RUMOCA_EXTERNAL_INCLUDE_DIRECTORIES) + foreach(RUMOCA_EXTERNAL_INCLUDE_DIRECTORY IN LISTS RUMOCA_EXTERNAL_INCLUDE_DIRECTORIES) + if(RUMOCA_EXTERNAL_INCLUDE_DIRECTORY MATCHES "^modelica://") + if(NOT RUMOCA_EXTERNAL_INCLUDE_DIR) + message(FATAL_ERROR "external include directories declared in resources/externalDependencies.json; set RUMOCA_EXTERNAL_INCLUDE_DIR for modelica:// include URIs") + endif() + target_include_directories({{ model_name }} PRIVATE "${RUMOCA_EXTERNAL_INCLUDE_DIR}") + else() + if(NOT IS_DIRECTORY "${RUMOCA_EXTERNAL_INCLUDE_DIRECTORY}") + message(FATAL_ERROR "missing external include directory: ${RUMOCA_EXTERNAL_INCLUDE_DIRECTORY}") + endif() + target_include_directories({{ model_name }} PRIVATE "${RUMOCA_EXTERNAL_INCLUDE_DIRECTORY}") + endif() + endforeach() +endif() + +set(RUMOCA_EXTERNAL_LIBRARY_DIR "$ENV{RUMOCA_EXTERNAL_LIBRARY_DIR}" CACHE PATH "Directory containing Modelica external runtime libraries") +set(RUMOCA_EXTERNAL_LIBRARIES_FILE "${CMAKE_CURRENT_LIST_DIR}/../resources/externalLibraries.txt") +if(EXISTS "${RUMOCA_EXTERNAL_LIBRARIES_FILE}") + file(STRINGS "${RUMOCA_EXTERNAL_LIBRARIES_FILE}" RUMOCA_EXTERNAL_LIBRARIES) + list(REMOVE_DUPLICATES RUMOCA_EXTERNAL_LIBRARIES) + if(RUMOCA_EXTERNAL_LIBRARIES) + if(NOT RUMOCA_EXTERNAL_LIBRARY_DIR) + message(FATAL_ERROR "external runtime libraries declared in resources/externalDependencies.json; set RUMOCA_EXTERNAL_LIBRARY_DIR") + endif() + target_link_directories({{ model_name }} PRIVATE "${RUMOCA_EXTERNAL_LIBRARY_DIR}") + foreach(RUMOCA_EXTERNAL_LIBRARY IN LISTS RUMOCA_EXTERNAL_LIBRARIES) + if(WIN32) + set(RUMOCA_EXTERNAL_LIBRARY_FILE "${RUMOCA_EXTERNAL_LIBRARY_DIR}/${RUMOCA_EXTERNAL_LIBRARY}.dll") + elseif(APPLE) + set(RUMOCA_EXTERNAL_LIBRARY_FILE "${RUMOCA_EXTERNAL_LIBRARY_DIR}/lib${RUMOCA_EXTERNAL_LIBRARY}.dylib") + else() + set(RUMOCA_EXTERNAL_LIBRARY_FILE "${RUMOCA_EXTERNAL_LIBRARY_DIR}/lib${RUMOCA_EXTERNAL_LIBRARY}.so") + endif() + if(NOT EXISTS "${RUMOCA_EXTERNAL_LIBRARY_FILE}") + message(FATAL_ERROR "missing external runtime library: ${RUMOCA_EXTERNAL_LIBRARY_FILE}") + endif() + target_link_libraries({{ model_name }} PRIVATE "${RUMOCA_EXTERNAL_LIBRARY}") + endforeach() + endif() +endif() + install(TARGETS {{ model_name }} RUNTIME DESTINATION ${CMAKE_INSTALL_PREFIX}/binaries/${FMU_PLATFORM} LIBRARY DESTINATION ${CMAKE_INSTALL_PREFIX}/binaries/${FMU_PLATFORM}) + +if(EXISTS "${CMAKE_CURRENT_LIST_DIR}/../resources") + install(DIRECTORY "${CMAKE_CURRENT_LIST_DIR}/../resources" DESTINATION ${CMAKE_INSTALL_PREFIX}) +endif() diff --git a/crates/rumoca-phase-codegen/src/templates/fmi2/build.sh.jinja b/crates/rumoca-phase-codegen/src/templates/fmi2/build.sh.jinja index cfddd1375..b7a1d0444 100644 --- a/crates/rumoca-phase-codegen/src/templates/fmi2/build.sh.jinja +++ b/crates/rumoca-phase-codegen/src/templates/fmi2/build.sh.jinja @@ -4,14 +4,195 @@ set -e cd "$(dirname "$0")" case "$(uname -s)" in - Linux*) PLATFORM=linux64; LIB_EXT=so ;; - Darwin*) PLATFORM=darwin64; LIB_EXT=dylib ;; + Linux*) PLATFORM=linux64; LIB_EXT=so; ALLOW_UNRESOLVED_FLAGS="" ;; + Darwin*) PLATFORM=darwin64; LIB_EXT=dylib; ALLOW_UNRESOLVED_FLAGS="-Wl,-undefined,dynamic_lookup" ;; MINGW*|MSYS*|CYGWIN*) PLATFORM=win64; LIB_EXT=dll ;; *) echo "Unknown platform"; exit 1 ;; esac mkdir -p binaries/$PLATFORM -cc -shared -fPIC -O2 -o binaries/$PLATFORM/{{ model_name }}.$LIB_EXT sources/{{ model_name }}.c -lm +MODEL_BINARY="binaries/$PLATFORM/{{ model_name }}.$LIB_EXT" +UNRESOLVED_SYMBOLS_FILE="$(mktemp "${TMPDIR:-/tmp}/rumoca-unresolved.XXXXXX")" +EXTERNAL_LIB_PATHS_FILE="$(mktemp "${TMPDIR:-/tmp}/rumoca-external-libs.XXXXXX")" +cleanup() { + rm -f "$UNRESOLVED_SYMBOLS_FILE" "$EXTERNAL_LIB_PATHS_FILE" +} +trap cleanup EXIT INT TERM -zip -r {{ model_name }}.fmu modelDescription.xml binaries/ sources/ +set -- +if [ -n "${RUMOCA_EXTERNAL_INCLUDE_DIR:-}" ]; then + old_ifs="$IFS" + IFS=: + for external_include_dir in $RUMOCA_EXTERNAL_INCLUDE_DIR; do + [ -n "$external_include_dir" ] || continue + if [ ! -d "$external_include_dir" ]; then + echo "missing external include directory: $external_include_dir" >&2 + exit 1 + fi + set -- "$@" "-I$external_include_dir" + done + IFS="$old_ifs" +fi + +if [ -s resources/externalIncludeDirectories.txt ]; then + while IFS= read -r include_dir; do + [ -n "$include_dir" ] || continue + case "$include_dir" in + modelica://*) + if [ -z "${RUMOCA_EXTERNAL_INCLUDE_DIR:-}" ]; then + echo "external include directories declared in resources/externalDependencies.json; set RUMOCA_EXTERNAL_INCLUDE_DIR for modelica:// include URIs" >&2 + exit 1 + fi + old_ifs="$IFS" + IFS=: + for external_include_dir in $RUMOCA_EXTERNAL_INCLUDE_DIR; do + [ -n "$external_include_dir" ] || continue + if [ ! -d "$external_include_dir" ]; then + echo "missing external include directory: $external_include_dir" >&2 + exit 1 + fi + set -- "$@" "-I$external_include_dir" + done + IFS="$old_ifs" + ;; + *) + if [ ! -d "$include_dir" ]; then + echo "missing external include directory: $include_dir" >&2 + exit 1 + fi + set -- "$@" "-I$include_dir" + ;; + esac + done </dev/null | awk '{print $NF}' | sed 's/^_//' | sort -u > "$UNRESOLVED_SYMBOLS_FILE" + +external_library_declares_unresolved_symbol() { + lib="$1" + if [ ! -s resources/externalDependencies.json ]; then + return 1 + fi + for symbol in $(awk -v lib="$lib" ' + /"symbol"[[:space:]]*:/ { + symbol=$0 + sub(/^.*"symbol"[[:space:]]*:[[:space:]]*"/, "", symbol) + sub(/".*$/, "", symbol) + next + } + /"libraries"[[:space:]]*:/ && symbol != "" { + if ($0 ~ "\"" lib "\"") print symbol + symbol="" + } + ' resources/externalDependencies.json); do + if grep -Fxq "$symbol" "$UNRESOLVED_SYMBOLS_FILE"; then + return 0 + fi + done + return 1 +} + +external_library_exports_unresolved_symbol() { + lib_file="$1" + [ -f "$lib_file" ] || return 1 + symbols_file="$(mktemp "${TMPDIR:-/tmp}/rumoca-lib-symbols.XXXXXX")" + nm -g "$lib_file" 2>/dev/null | awk '{print $NF}' | sed 's/^_//' > "$symbols_file" + while IFS= read -r symbol; do + [ -n "$symbol" ] || continue + if grep -Fxq "$symbol" "$UNRESOLVED_SYMBOLS_FILE"; then + rm -f "$symbols_file" + return 0 + fi + done < "$symbols_file" + rm -f "$symbols_file" + return 1 +} + +rewrite_darwin_runtime_paths() { + [ "$PLATFORM" = "darwin64" ] || return 0 + command -v install_name_tool >/dev/null 2>&1 || { + echo "install_name_tool is required to make Darwin FMU external libraries loader-relative" >&2 + exit 1 + } + while IFS= read -r external_lib_file; do + [ -n "$external_lib_file" ] || continue + external_lib_base="$(basename "$external_lib_file")" + copied_lib="binaries/$PLATFORM/$external_lib_base" + install_name_tool -id "@loader_path/$external_lib_base" "$copied_lib" + install_name_tool -change "$external_lib_file" "@loader_path/$external_lib_base" "$MODEL_BINARY" 2>/dev/null || true + install_name_tool -change "$external_lib_base" "@loader_path/$external_lib_base" "$MODEL_BINARY" 2>/dev/null || true + install_name_tool -change "@rpath/$external_lib_base" "@loader_path/$external_lib_base" "$MODEL_BINARY" 2>/dev/null || true + done < "$EXTERNAL_LIB_PATHS_FILE" + + while IFS= read -r source_lib_file; do + [ -n "$source_lib_file" ] || continue + source_lib="binaries/$PLATFORM/$(basename "$source_lib_file")" + while IFS= read -r dependency_lib_file; do + [ -n "$dependency_lib_file" ] || continue + dependency_base="$(basename "$dependency_lib_file")" + install_name_tool -change "$dependency_lib_file" "@loader_path/$dependency_base" "$source_lib" 2>/dev/null || true + install_name_tool -change "$dependency_base" "@loader_path/$dependency_base" "$source_lib" 2>/dev/null || true + install_name_tool -change "@rpath/$dependency_base" "@loader_path/$dependency_base" "$source_lib" 2>/dev/null || true + done < "$EXTERNAL_LIB_PATHS_FILE" + done < "$EXTERNAL_LIB_PATHS_FILE" +} + +EXTERNAL_LIBS_NEEDED=0 +if [ -s resources/externalLibraries.txt ]; then + while IFS= read -r lib; do + [ -n "$lib" ] || continue + if [ -z "${RUMOCA_EXTERNAL_LIBRARY_DIR:-}" ]; then + if external_library_declares_unresolved_symbol "$lib"; then + echo "external runtime libraries declared in resources/externalDependencies.json; set RUMOCA_EXTERNAL_LIBRARY_DIR" >&2 + exit 1 + fi + continue + fi + case "$PLATFORM" in + win*) lib_file="$RUMOCA_EXTERNAL_LIBRARY_DIR/$lib.$LIB_EXT" ;; + *) lib_file="$RUMOCA_EXTERNAL_LIBRARY_DIR/lib$lib.$LIB_EXT" ;; + esac + if external_library_declares_unresolved_symbol "$lib"; then + if [ ! -f "$lib_file" ]; then + echo "missing external runtime library: $lib_file" >&2 + exit 1 + fi + elif ! external_library_exports_unresolved_symbol "$lib_file"; then + continue + fi + if [ ! -f "$lib_file" ]; then + echo "external runtime libraries declared in resources/externalDependencies.json; set RUMOCA_EXTERNAL_LIBRARY_DIR" >&2 + exit 1 + fi + set -- "$@" "-L$RUMOCA_EXTERNAL_LIBRARY_DIR" "-l$lib" + printf '%s\n' "$lib_file" >> "$EXTERNAL_LIB_PATHS_FILE" + EXTERNAL_LIBS_NEEDED=1 + done < #include #include +#include {# Configuration for render_expr - C style #} {% set c_symbol_policy = { @@ -21,7 +22,9 @@ "enum", "extern", "float", "for", "goto", "if", "int", "long", "register", "return", "short", "signed", "sizeof", "static", "struct", "switch", "typedef", "union", "unsigned", "void", "volatile", "while", "inline", "restrict", - "m", "t", "time", "i", "n", "arr", "s", "sz", "ModelInstance", "Complex", + "m", "t", "time", "i", "n", "arr", "s", "sz", "value", "set", "reset", "suspend", + "resume", "active", "localActive", "newActive", "oldActive", "enableFire", "fire", + "condition", "t_start", "ModelInstance", "Complex", "__rumoca_sum_d", "__rumoca_solve_linear_component", "pre", "der", "edge", "change", "initial", "terminal" ] } %} @@ -35,6 +38,9 @@ {% set solve_derivative_rows = solve_blocks.continuous.derivative_rhs.scalar_programs.programs if solve_blocks is defined and solve_blocks.continuous is defined and solve_blocks.continuous.derivative_rhs is defined and solve_blocks.continuous.derivative_rhs.scalar_programs is defined else [] %} {% set solve_derivative_output_indices = solve_blocks.continuous.derivative_rhs.scalar_programs.output_indices if solve_blocks is defined and solve_blocks.continuous is defined and solve_blocks.continuous.derivative_rhs is defined and solve_blocks.continuous.derivative_rhs.scalar_programs is defined else [] %} {% set solve_derivative_nodes = solve_derivative_nodes if solve_derivative_nodes is defined else (solve.continuous.derivative_rhs.nodes if solve is defined and solve.continuous is defined and solve.continuous.derivative_rhs is defined and solve.continuous.derivative_rhs.nodes is defined else []) %} +{% set solve_context = solve if solve is defined else {} %} +{% set solve_visible_names = solve.visible_names if solve is defined and solve.visible_names is defined else [] %} +{% set c_builtin_function_names = ["abs", "acos", "asin", "atan", "atan2", "ceil", "cos", "cosh", "exp", "floor", "log", "log10", "round", "sin", "sinh", "sqrt", "tan", "tanh", "trunc"] %} {% set solve_row_c_cfg = {"time": "m->time", "y": "__rumoca_solve_y(m, {})", "p": "__rumoca_solve_p(m, {})"} %} {% set solve_slot_assign_c_cfg = {"y_set": "__rumoca_solve_set_y(m, {}, {})", "p_set": "__rumoca_solve_set_p(m, {}, {})"} %} @@ -50,8 +56,8 @@ /* change(x): true when x changed value */ #define change(x) ((x) != pre(x)) -/* initial(): true only during initialization (always false in continuous mode) */ -#define initial() 0 +/* initial(): true during FMI initialization mode (MLS §3.7.3.1). */ +#define initial() ((m->state == modelInitializationMode) ? 1 : 0) /* terminal(): true only at end of simulation (always false in continuous mode) */ #define terminal() 0 @@ -61,6 +67,10 @@ * allows Complex library functions to compile when they appear in the * generated code but only real-valued results are used. */ #define Complex(re, im) (re) +#define REAL_C(x) (x) +#define Modelica_Units_SI_TemperatureDifference(x) (x) +#define Modelica_Units_SI_MassFraction(x) (x) +static inline double linspace(double start, double stop, double n) { (void)stop; (void)n; return start; } /* Array helper: sum all elements of a double array */ static inline double __rumoca_sum_d(const double* arr, int n) { @@ -135,6 +145,9 @@ static inline double zeros(int n) { (void)n; return 0.0; } static inline double ones(int n) { (void)n; return 1.0; } static inline double fill(double val, int n) { (void)n; return val; } static inline int size(const double* arr, int dim) { (void)arr; (void)dim; return 0; } +static inline double sign(double x) { return (x > 0.0) - (x < 0.0); } +static inline double getInstanceName(void) { return 0.0; } +static inline double Buildings_ThermalZones_EnergyPlus_9_6_0_ThermalZone_Medium(double x) { return x; } /* Modelica interval() builtin — clocked partition intrinsic (MLS §16.10). * Returns 0 when no clock schedule is known. */ @@ -195,6 +208,16 @@ static int ModelicaStrings_scanInteger(double string, int startIndex, int unsign static double Modelica_Blocks_Types_ExternalCombiTable1D() { return 0.0; } static double Modelica_Blocks_Types_ExternalCombiTimeTable() { return 0.0; } +#ifndef RUMOCA_SOLVE_EXTERNAL_CALL +static double __rumoca_solve_external_call_default(const char* function, int output_index, int arg_count, const double* args) { + (void)args; + fprintf(stderr, "Rumoca solve ExternalCall requires native runtime bridge: %s output_index=%d arg_count=%d\n", function, output_index, arg_count); + abort(); + return NAN; +} +#define RUMOCA_SOLVE_EXTERNAL_CALL __rumoca_solve_external_call_default +#endif + /* Named argument passthrough macros — the __rumoca_named_arg__ prefix is * generated for named arguments in external function calls. Strip the prefix * so the value passes through to the function. */ @@ -205,17 +228,106 @@ static double Modelica_Blocks_Types_ExternalCombiTimeTable() { return 0.0; } #define __rumoca_named_arg___extrapolation(x) (x) #define __rumoca_named_arg___verboseRead(x) (x) #define __rumoca_named_arg___verboseExtrapolation(x) (x) +#define __rumoca_named_arg___T(x) (x) +#define __rumoca_named_arg___TWetBul(x) (x) +#define __rumoca_named_arg___X(...) (__VA_ARGS__) +#define __rumoca_named_arg___X_w(x) (x) +#define __rumoca_named_arg___a(x) (x) +#define __rumoca_named_arg___b(x) (x) +#define __rumoca_named_arg___buildingsRootFileLocation(x) (x) +#define __rumoca_named_arg___c(x) (x) +#define __rumoca_named_arg___caseSensitive(x) (x) +#define __rumoca_named_arg___d(x) (x) +#define __rumoca_named_arg___delta(x) (x) +#define __rumoca_named_arg___deltaInv(x) (x) +#define __rumoca_named_arg___deltaX(x) (x) +#define __rumoca_named_arg___deltax(x) (x) +#define __rumoca_named_arg___derivatives_delta(x) (x) +#define __rumoca_named_arg___derivatives_structure(x) (x) +#define __rumoca_named_arg___diameter(x) (x) +#define __rumoca_named_arg___dummy(x) (x) +#define __rumoca_named_arg___e(x) (x) +#define __rumoca_named_arg___ensureMonotonicity(x) (x) +#define __rumoca_named_arg___epName(x) (x) +#define __rumoca_named_arg___epwName(x) (x) +#define __rumoca_named_arg___f(x) (x) +#define __rumoca_named_arg___fmuName(x) (x) +#define __rumoca_named_arg___h(x) (x) +#define __rumoca_named_arg___idfName(x) (x) +#define __rumoca_named_arg___idfVersion(x) (x) +#define __rumoca_named_arg___initialCall(x) (x) +#define __rumoca_named_arg___inpNames(x) (x) +#define __rumoca_named_arg___inpUnits(x) (x) +#define __rumoca_named_arg___jsonKeysValues(x) (x) +#define __rumoca_named_arg___jsonName(x) (x) +#define __rumoca_named_arg___modelicaInstanceName(x) (x) +#define __rumoca_named_arg___modelicaNameBuilding(x) (x) +#define __rumoca_named_arg___mu_a(x) (x) +#define __rumoca_named_arg___mu_b(x) (x) +#define __rumoca_named_arg___nDer(x) (x) +#define __rumoca_named_arg___nInp(x) (x) +#define __rumoca_named_arg___nOut(x) (x) +#define __rumoca_named_arg___nParOut(x) (x) +#define __rumoca_named_arg___nY(x) (x) +#define __rumoca_named_arg___neg(x) (x) +#define __rumoca_named_arg___objectType(x) (x) +#define __rumoca_named_arg___outNames(x) (x) +#define __rumoca_named_arg___outUnits(x) (x) +#define __rumoca_named_arg___p(x) (x) +#define __rumoca_named_arg___pSat(x) (x) +#define __rumoca_named_arg___p_w(x) (x) +#define __rumoca_named_arg___parOutNames(x) (x) +#define __rumoca_named_arg___parOutUnits(x) (x) +#define __rumoca_named_arg___per(x) (x) +#define __rumoca_named_arg___phi(x) (x) +#define __rumoca_named_arg___pos(x) (x) +#define __rumoca_named_arg___printUnit(x) (x) +#define __rumoca_named_arg___r_V(x) (x) +#define __rumoca_named_arg___relativeSurfaceTolerance(x) (x) +#define __rumoca_named_arg___rho_a(x) (x) +#define __rumoca_named_arg___rho_b(x) (x) +#define __rumoca_named_arg___spawnExe(x) (x) +#define __rumoca_named_arg___state(x) (x) +#define __rumoca_named_arg___strict(x) (x) +#define __rumoca_named_arg___string1(x) (x) +#define __rumoca_named_arg___string2(x) (x) +#define __rumoca_named_arg___u(x) (x) +#define __rumoca_named_arg___usePrecompiledFMU(x) (x) +#define __rumoca_named_arg___x(x) (x) +#define __rumoca_named_arg___x1(x) (x) +#define __rumoca_named_arg___x2(x) (x) +#define __rumoca_named_arg___x_small(x) (x) +#define __rumoca_named_arg___y(x) (x) +#define __rumoca_named_arg___y1(x) (x) +#define __rumoca_named_arg___y2(x) (x) /* ========================================================================= * Enumeration literal constants * ========================================================================= */ +{% set runtime_field_macro_names = ["time", "x", "xdot", "y", "u", "w", "p", "z", "pre_z", "m", "pre_m", "constants", "event_indicators", "event_indicators_prev", "state", "fmu_type", "dirty_values", "is_new_event_iteration"] %} {% for name, ordinal in dae.enum_literal_ordinals | items %} -#define {{ symbol(symbols, name) }} {{ ordinal }} +{% set enum_symbol = symbol(symbols, name) %} +{% if enum_symbol not in runtime_field_macro_names %} +#define {{ enum_symbol }} {{ ordinal }} +{% endif %} {% endfor %} /* Enumeration type constructors — identity macros for enum type casts */ {% for type_name in dae.enum_type_names | default([]) %} -#define {{ symbol(symbols, type_name) }}(x) (x) +{% set enum_type_symbol = symbol(symbols, type_name) %} +{% if enum_type_symbol not in runtime_field_macro_names %} +#define {{ enum_type_symbol }}(x) (x) +{% endif %} +{% endfor %} + +/* Source-reference aliases: renderers may emit sanitized full source refs + * while local unpack aliases use the allocated target symbol. */ +{% for name in dae.symbol_refs %} +{% set source_alias = name | sanitize %} +{% set target_symbol = symbol(symbols, name) %} +{% if source_alias != target_symbol and source_alias not in runtime_field_macro_names and source_alias not in c_symbol_policy.reserved %} +#define {{ source_alias }} {{ target_symbol }} +{% endif %} {% endfor %} {#- Macro: compute total scalar size of a variable map -#} @@ -231,6 +343,59 @@ static double Modelica_Blocks_Types_ExternalCombiTimeTable() { return 0.0; } {{ ns.total }} {%- endmacro %} +{#- Generate Solve IR y-slot accessors from solve.visible_names and actual + target storage layout. This keeps zero-extent variables from shifting the + runtime solver index space away from generated FMI arrays. #} +{% macro solve_y_get_cases_for_vars(vars, storage) -%} +{%- set ns = namespace(offset=0) -%} +{%- for name, var in vars | items -%} +{%- if var.dims -%} +{%- set sz = var.dims | product -%} +{%- for i in range(sz) -%} +{%- set scalar_name = source_ref(name, var.dims, i + 1) -%} +{%- for visible in solve_visible_names -%} +{%- if visible == scalar_name %} + case {{ loop.index0 }}: return {{ storage }}[{{ ns.offset }}]; /* {{ scalar_name }} */ +{%- endif -%} +{%- endfor -%} +{%- set ns.offset = ns.offset + 1 -%} +{%- endfor -%} +{%- else -%} +{%- for visible in solve_visible_names -%} +{%- if visible == name %} + case {{ loop.index0 }}: return {{ storage }}[{{ ns.offset }}]; /* {{ name }} */ +{%- endif -%} +{%- endfor -%} +{%- set ns.offset = ns.offset + 1 -%} +{%- endif -%} +{%- endfor -%} +{%- endmacro %} + +{% macro solve_y_set_cases_for_vars(vars, storage) -%} +{%- set ns = namespace(offset=0) -%} +{%- for name, var in vars | items -%} +{%- if var.dims -%} +{%- set sz = var.dims | product -%} +{%- for i in range(sz) -%} +{%- set scalar_name = source_ref(name, var.dims, i + 1) -%} +{%- for visible in solve_visible_names -%} +{%- if visible == scalar_name %} + case {{ loop.index0 }}: {{ storage }}[{{ ns.offset }}] = value; return; /* {{ scalar_name }} */ +{%- endif -%} +{%- endfor -%} +{%- set ns.offset = ns.offset + 1 -%} +{%- endfor -%} +{%- else -%} +{%- for visible in solve_visible_names -%} +{%- if visible == name %} + case {{ loop.index0 }}: {{ storage }}[{{ ns.offset }}] = value; return; /* {{ name }} */ +{%- endif -%} +{%- endfor -%} +{%- set ns.offset = ns.offset + 1 -%} +{%- endif -%} +{%- endfor -%} +{%- endmacro %} + {#- Macro: unpack variables from a ModelInstance array into local C aliases. Parameters: vars - variable map (e.g., dae.x) @@ -292,6 +457,7 @@ static double Modelica_Blocks_Types_ExternalCombiTimeTable() { return 0.0; } #define N_CONSTANTS {{ var_size(dae.constants) | trim }} #define N_DISCRETE_REAL {{ var_size(dae.z) | trim }} #define N_DISCRETE_VAL {{ var_size(dae.m) | trim }} +#define N_NATIVE_OBSERVABLES {% set ns_obs = namespace(total=0) %}{% for observable in dae.__rumoca_observables | default([]) %}{% if observable.causality | default("local") == "output" %}{% set ns_obs.total = ns_obs.total + 1 %}{% endif %}{% endfor %}{{ ns_obs.total }} #define N_EVENT_INDICATORS {{ solve_root_rows | length }} #define N_REALS (N_STATES + N_DERIVATIVES + N_ALGEBRAICS + N_INPUTS + N_OUTPUTS + N_PARAMETERS + N_DISCRETE_REAL) @@ -304,6 +470,7 @@ static double Modelica_Blocks_Types_ExternalCombiTimeTable() { return 0.0; } #define VR_P (VR_W + N_OUTPUTS) #define VR_Z (VR_P + N_PARAMETERS) #define VR_M (VR_Z + N_DISCRETE_REAL) +#define VR_OBS (VR_M + N_DISCRETE_VAL) /* ========================================================================= * FMI 2.0 type definitions (subset needed for Model Exchange) @@ -366,6 +533,12 @@ typedef struct { } ModelInstance; static double __rumoca_solve_y(const ModelInstance* m, int index) { + switch (index) { +{{ solve_y_get_cases_for_vars(dae.x, "m->x") }} +{{ solve_y_get_cases_for_vars(dae.y, "m->y") }} +{{ solve_y_get_cases_for_vars(dae.w, "m->w") }} + default: break; + } if (index < N_STATES) { return m->x[index]; } @@ -400,6 +573,12 @@ static double __rumoca_solve_p(const ModelInstance* m, int index) { } static void __rumoca_solve_set_y(ModelInstance* m, int index, double value) { + switch (index) { +{{ solve_y_set_cases_for_vars(dae.x, "m->x") }} +{{ solve_y_set_cases_for_vars(dae.y, "m->y") }} +{{ solve_y_set_cases_for_vars(dae.w, "m->w") }} + default: break; + } if (index < N_STATES) { m->x[index] = value; return; @@ -442,11 +621,14 @@ static void __rumoca_solve_set_p(ModelInstance* m, int index, double value) { * Forward declarations * ========================================================================= */ static void initialize_defaults(ModelInstance* m); +static void apply_parameter_bindings(ModelInstance* m); +static void apply_initial_equations(ModelInstance* m); static void compute_algebraics(ModelInstance* m); static void compute_derivatives(ModelInstance* m); static void compute_outputs(ModelInstance* m); static void compute_discrete_updates(ModelInstance* m); static void compute_event_indicators(ModelInstance* m); +static fmi2Status get_native_observable_by_vr(const ModelInstance* m, fmi2ValueReference vr, fmi2Real* value); static double event_right_limit_time(double event_time) { double scale = 1.0 + fabs(event_time); @@ -505,14 +687,14 @@ static void snapshot_pre_parameters(ModelInstance* m) { {#- Macro: check if a function has any record-typed parameters -#} {% macro has_complex_params(func) -%} -{%- for p in func.inputs -%}{%- if p.type_class == "Record" -%}yes{%- endif -%}{%- endfor -%} +{%- for p in func.inputs -%}{%- if p.type_class | default("") == "Record" -%}yes{%- endif -%}{%- endfor -%} {%- endmacro %} /* ========================================================================= * User-Defined Functions — Forward Declarations * ========================================================================= */ {% for func_name, func in dae.functions | items %} -{% if func.outputs | length == 1 %} +{% if func.outputs | length == 1 and not (func_name | last_segment in c_builtin_function_names) %} static double {{ symbol(symbols, func_name) }}({% for p in func.inputs %}{% if p.dims %}const double* {{ symbol(symbols, p.name) }}, int {{ symbol(symbols, p.name) }}_size{% else %}double {{ symbol(symbols, p.name) }}{% endif %}{{ ", " if not loop.last else "" }}{% endfor %}); {% endif %} {% endfor %} @@ -521,7 +703,7 @@ static double {{ symbol(symbols, func_name) }}({% for p in func.inputs %}{% if p * User-Defined Functions — Definitions * ========================================================================= */ {% for func_name, func in dae.functions | items %} -{% if func.outputs | length == 1 %} +{% if func.outputs | length == 1 and not (func_name | last_segment in c_builtin_function_names) %} static double {{ symbol(symbols, func_name) }}({% for p in func.inputs %}{% if p.dims %}const double* {{ symbol(symbols, p.name) }}, int {{ symbol(symbols, p.name) }}_size{% else %}double {{ symbol(symbols, p.name) }}{% endif %}{{ ", " if not loop.last else "" }}{% endfor %}) { {% if is_self_call(func_name, func) %} /* Builtin wrapper — delegate to C standard library */ @@ -532,14 +714,14 @@ static double {{ symbol(symbols, func_name) }}({% for p in func.inputs %}{% if p {% for p in func.inputs %} (void){{ symbol(symbols, p.name) }}; {% endfor %} return 0.0; +{% elif unsupported_c_function_name(func_name) or unsupported_c_function_body(func) %} +{{ unsupported_c_function_body(func) | indent(4) }} {% elif func.external %} - /* External {{ func.external.language }} function */ -{% if func.external.output_name %} - return {{ func.external.function_name | default(symbol(symbols, func_name)) }}({% for arg in func.external.arg_names %}{{ symbol(symbols, arg) }}{{ ", " if not loop.last else "" }}{% endfor %}); -{% else %} - {{ func.external.function_name | default(symbol(symbols, func_name)) }}({% for arg in func.external.arg_names %}{{ symbol(symbols, arg) }}{{ ", " if not loop.last else "" }}{% endfor %}); - return {{ symbol(symbols, func.outputs[0].name) }}; -{% endif %} + /* External {{ func.external.language }} function — requires native runtime bridge */ +{% for p in func.inputs %} (void){{ symbol(symbols, p.name) }}; +{% if p.dims %} (void){{ symbol(symbols, p.name) }}_size; +{% endif %}{% endfor %} + return RUMOCA_SOLVE_EXTERNAL_CALL("{{ func.external.function_name | default(symbol(symbols, func_name)) }}", 0, 0, NULL); {% else %} {% for p in func.locals %} {% if p.dims %} @@ -554,8 +736,9 @@ static double {{ symbol(symbols, func_name) }}({% for p in func.inputs %}{% if p {% if func.outputs | length > 0 %} double {{ symbol(symbols, func.outputs[0].name) }} = 0.0; {% endif %} -{% if func.body | length > 0 %} -{{ render_statements(func.body, cfg, " ") }} +{% set func_body = func["body"] %} +{% if func_body | length > 0 %} +{{ render_function_statements(func_body, cfg, " ", symbol(symbols, func.outputs[0].name)) }} {% endif %} {% if func.outputs | length > 0 %} return {{ symbol(symbols, func.outputs[0].name) }}; @@ -649,7 +832,7 @@ static void initialize_defaults(ModelInstance* m) { {% set ns_idx = namespace(offset=0) %} {% for name, var in dae.constants | items %} - {{ symbol(symbols, name) }} = {% if var.start and is_string_literal(var.start) != "yes" %}{{ render_expr(var.start, cfg) }}{% else %}0.0{% endif %}; + {{ symbol(symbols, name) }} = {% if var.start and is_string_literal(var.start) != "yes" and expr_has_dynamic_multidim_index(var.start) != "yes" %}{{ render_expr(var.start, cfg) }}{% else %}0.0{% endif %}; m->constants[{{ ns_idx.offset }}] = {{ symbol(symbols, name) }}; /* {{ name }} */ {% set ns_idx.offset = ns_idx.offset + 1 %} {% endfor %} @@ -659,12 +842,12 @@ static void initialize_defaults(ModelInstance* m) { {% if var.dims %} {% set sz = var.dims | product %} {% for i in range(sz) %} - {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }} = {% if var.start and is_string_literal(var.start) != "yes" %}{{ render_expr_at_index(var.start, i + 1, cfg) }}{% else %}0.0{% endif %}; + {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }} = {% if var.start and is_string_literal(var.start) != "yes" and expr_has_dynamic_multidim_index(var.start) != "yes" %}{{ render_expr_at_index(var.start, i + 1, cfg) }}{% else %}0.0{% endif %}; m->p[{{ ns_idx.offset }}] = {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }}; /* {{ source_ref(name, var.dims, i + 1) }} */ {% set ns_idx.offset = ns_idx.offset + 1 %} {% endfor %} {% else %} - {{ symbol(symbols, name) }} = {% if var.start and is_string_literal(var.start) != "yes" %}{{ render_expr(var.start, cfg) }}{% else %}0.0{% endif %}; + {{ symbol(symbols, name) }} = {% if var.start and is_string_literal(var.start) != "yes" and expr_has_dynamic_multidim_index(var.start) != "yes" %}{{ render_expr(var.start, cfg) }}{% else %}0.0{% endif %}; m->p[{{ ns_idx.offset }}] = {{ symbol(symbols, name) }}; /* {{ name }} */ {% set ns_idx.offset = ns_idx.offset + 1 %} {% endif %} @@ -675,11 +858,11 @@ static void initialize_defaults(ModelInstance* m) { {% if var.dims %} {% set sz = var.dims | product %} {% for i in range(sz) %} - m->x[{{ ns_idx.offset }}] = {% if var.start and is_string_literal(var.start) != "yes" %}{{ render_expr_at_index(var.start, i + 1, cfg) }}{% else %}0.0{% endif %}; /* {{ source_ref(name, var.dims, i + 1) }} */ + m->x[{{ ns_idx.offset }}] = {% if var.start and is_string_literal(var.start) != "yes" and expr_has_dynamic_multidim_index(var.start) != "yes" %}{{ render_expr_at_index(var.start, i + 1, cfg) }}{% else %}0.0{% endif %}; /* {{ source_ref(name, var.dims, i + 1) }} */ {% set ns_idx.offset = ns_idx.offset + 1 %} {% endfor %} {% else %} - m->x[{{ ns_idx.offset }}] = {% if var.start and is_string_literal(var.start) != "yes" %}{{ render_expr(var.start, cfg) }}{% else %}0.0{% endif %}; /* {{ name }} */ + m->x[{{ ns_idx.offset }}] = {% if var.start and is_string_literal(var.start) != "yes" and expr_has_dynamic_multidim_index(var.start) != "yes" %}{{ render_expr(var.start, cfg) }}{% else %}0.0{% endif %}; /* {{ name }} */ {% set ns_idx.offset = ns_idx.offset + 1 %} {% endif %} {% endfor %} @@ -689,12 +872,12 @@ static void initialize_defaults(ModelInstance* m) { {% if var.dims %} {% set sz = var.dims | product %} {% for i in range(sz) %} - {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }} = {% if var.start and is_string_literal(var.start) != "yes" %}{{ render_expr_at_index(var.start, i + 1, cfg) }}{% else %}0.0{% endif %}; + {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }} = {% if var.start and is_string_literal(var.start) != "yes" and expr_has_dynamic_multidim_index(var.start) != "yes" %}{{ render_expr_at_index(var.start, i + 1, cfg) }}{% else %}0.0{% endif %}; m->u[{{ ns_idx.offset }}] = {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }}; /* {{ source_ref(name, var.dims, i + 1) }} */ {% set ns_idx.offset = ns_idx.offset + 1 %} {% endfor %} {% else %} - {{ symbol(symbols, name) }} = {% if var.start and is_string_literal(var.start) != "yes" %}{{ render_expr(var.start, cfg) }}{% else %}0.0{% endif %}; + {{ symbol(symbols, name) }} = {% if var.start and is_string_literal(var.start) != "yes" and expr_has_dynamic_multidim_index(var.start) != "yes" %}{{ render_expr(var.start, cfg) }}{% else %}0.0{% endif %}; m->u[{{ ns_idx.offset }}] = {{ symbol(symbols, name) }}; /* {{ name }} */ {% set ns_idx.offset = ns_idx.offset + 1 %} {% endif %} @@ -705,12 +888,12 @@ static void initialize_defaults(ModelInstance* m) { {% if var.dims %} {% set sz = var.dims | product %} {% for i in range(sz) %} - {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }} = {% if var.start and is_string_literal(var.start) != "yes" %}{{ render_expr_at_index(var.start, i + 1, cfg) }}{% else %}0.0{% endif %}; + {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }} = {% if var.start and is_string_literal(var.start) != "yes" and expr_has_dynamic_multidim_index(var.start) != "yes" %}{{ render_expr_at_index(var.start, i + 1, cfg) }}{% else %}0.0{% endif %}; m->z[{{ ns_idx.offset }}] = {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }}; /* {{ source_ref(name, var.dims, i + 1) }} */ {% set ns_idx.offset = ns_idx.offset + 1 %} {% endfor %} {% else %} - {{ symbol(symbols, name) }} = {% if var.start and is_string_literal(var.start) != "yes" %}{{ render_expr(var.start, cfg) }}{% else %}0.0{% endif %}; + {{ symbol(symbols, name) }} = {% if var.start and is_string_literal(var.start) != "yes" and expr_has_dynamic_multidim_index(var.start) != "yes" %}{{ render_expr(var.start, cfg) }}{% else %}0.0{% endif %}; m->z[{{ ns_idx.offset }}] = {{ symbol(symbols, name) }}; /* {{ name }} */ {% set ns_idx.offset = ns_idx.offset + 1 %} {% endif %} @@ -719,6 +902,95 @@ static void initialize_defaults(ModelInstance* m) { m->dirty_values = 1; } +/* ========================================================================= + * FMI initialization binding propagation + * ========================================================================= */ +static void apply_parameter_bindings(ModelInstance* m) { +{% if dae.p | length > 0 %} + const double t = m->time; + const double time = m->time; + (void)t; (void)time; + +{{ unpack_vars(dae.constants, "m->constants") }} +{{ unpack_vars(dae.p, "m->p", mutable=true) }} +{% set ns_idx = namespace(offset=0) %} +{% for name, var in dae.p | items %} +{% if var.dims %} +{% set sz = var.dims | product %} +{% for i in range(sz) %} +{% set binding_name = source_ref(name, var.dims, i + 1) %} +{% set rhs = parameter_binding_rhs(name, var.start, i + 1, cfg) if var.start else "" %} +{% if not rhs %} +{% set rhs = alg_rhs_for_var_with_dae(binding_name, dae, cfg) %} +{% endif %} +{% if "WARNING:" in rhs %} +{% set rhs = "" %} +{% endif %} +{% if rhs %} + {{ symbol(symbols, binding_name) }} = {{ rhs }}; + m->p[{{ ns_idx.offset }}] = {{ symbol(symbols, binding_name) }}; /* binding {{ binding_name }} */ +{% endif %} +{% set ns_idx.offset = ns_idx.offset + 1 %} +{% endfor %} +{% else %} +{% set rhs = parameter_binding_rhs(name, var.start, 0, cfg) if var.start else "" %} +{% if not rhs %} +{% set rhs = alg_rhs_for_var_with_dae(name, dae, cfg) %} +{% endif %} +{% if "WARNING:" in rhs %} +{% set rhs = "" %} +{% endif %} +{% if rhs %} + {{ symbol(symbols, name) }} = {{ rhs }}; + m->p[{{ ns_idx.offset }}] = {{ symbol(symbols, name) }}; /* binding {{ name }} */ +{% endif %} +{% set ns_idx.offset = ns_idx.offset + 1 %} +{% endif %} +{% endfor %} + m->dirty_values = 1; +{% else %} + (void)m; +{% endif %} +} + +static void apply_initial_equations(ModelInstance* m) { +{% if dae.initial_equations | default([]) | length > 0 and dae.x | length > 0 %} + const double t = m->time; + const double time = m->time; + (void)t; (void)time; + +{{ unpack_vars(dae.x, "m->x", der_prefix=true, pre_array="m->x") }} +{{ unpack_vars(dae.u, "m->u", pre_array="m->u") }} +{{ unpack_vars(dae.p, "m->p") }} +{{ unpack_vars(dae.constants, "m->constants") }} +{{ unpack_vars(dae.z, "m->z", pre_array="m->pre_z") }} +{{ unpack_vars(dae.m, "m->m", pre_array="m->pre_m") }} +{% set ns_idx = namespace(offset=0) %} +{% for name, var in dae.x | items %} +{% if var.dims %} +{% set sz = var.dims | product %} +{% for i in range(sz) %} +{% set scalar_name = source_ref(name, var.dims, i + 1) %} +{% set rhs = initial_runtime_rhs_for_var(dae, scalar_name, cfg) %} +{% if rhs %} + m->x[{{ ns_idx.offset }}] = {{ rhs }}; /* initial equation: {{ scalar_name }} */ +{% endif %} +{% set ns_idx.offset = ns_idx.offset + 1 %} +{% endfor %} +{% else %} +{% set rhs = initial_runtime_rhs_for_var(dae, name, cfg) %} +{% if rhs %} + m->x[{{ ns_idx.offset }}] = {{ rhs }}; /* initial equation: {{ name }} */ +{% endif %} +{% set ns_idx.offset = ns_idx.offset + 1 %} +{% endif %} +{% endfor %} + m->dirty_values = 1; +{% else %} + (void)m; +{% endif %} +} + /* ========================================================================= * Compute algebraic variables from f_x equations * ========================================================================= */ @@ -751,12 +1023,12 @@ static void compute_algebraics(ModelInstance* m) { {% if var.dims %} {% set sz = var.dims | product %} {% for i in range(sz) %} - {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }} = {{ alg_rhs_for_var(source_ref(name, var.dims, i + 1), dae.f_x, cfg) }}; + {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }} = {{ visible_or_alg_rhs_for_var(source_ref(name, var.dims, i + 1), dae, solve_context, cfg, solve_row_c_cfg) }}; m->w[{{ ns_out.offset }}] = {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }}; /* {{ source_ref(name, var.dims, i + 1) }} */ {% set ns_out.offset = ns_out.offset + 1 %} {% endfor %} {% else %} - {{ symbol(symbols, name) }} = {{ alg_rhs_for_var(name, dae.f_x, cfg) }}; + {{ symbol(symbols, name) }} = {{ visible_or_alg_rhs_for_var(name, dae, solve_context, cfg, solve_row_c_cfg) }}; m->w[{{ ns_out.offset }}] = {{ symbol(symbols, name) }}; /* {{ name }} */ {% set ns_out.offset = ns_out.offset + 1 %} {% endif %} @@ -768,12 +1040,12 @@ static void compute_algebraics(ModelInstance* m) { {% if var.dims %} {% set sz = var.dims | product %} {% for i in range(sz) %} - {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }} = {{ alg_rhs_for_var(source_ref(name, var.dims, i + 1), dae.f_x, cfg) }}; + {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }} = {{ visible_or_alg_rhs_for_var(source_ref(name, var.dims, i + 1), dae, solve_context, cfg, solve_row_c_cfg) }}; m->y[{{ ns_alg.offset }}] = {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }}; /* {{ source_ref(name, var.dims, i + 1) }} */ {% set ns_alg.offset = ns_alg.offset + 1 %} {% endfor %} {% else %} - {{ symbol(symbols, name) }} = {{ alg_rhs_for_var(name, dae.f_x, cfg) }}; + {{ symbol(symbols, name) }} = {{ visible_or_alg_rhs_for_var(name, dae, solve_context, cfg, solve_row_c_cfg) }}; m->y[{{ ns_alg.offset }}] = {{ symbol(symbols, name) }}; /* {{ name }} */ {% set ns_alg.offset = ns_alg.offset + 1 %} {% endif %} @@ -857,11 +1129,11 @@ static void compute_outputs(ModelInstance* m) { {% if var.dims %} {% set sz = var.dims | product %} {% for i in range(sz) %} - m->w[{{ ns_out.offset }}] = {{ alg_rhs_for_var(source_ref(name, var.dims, i + 1), dae.f_x, cfg) }}; /* {{ source_ref(name, var.dims, i + 1) }} */ + m->w[{{ ns_out.offset }}] = {{ visible_or_alg_rhs_for_var(source_ref(name, var.dims, i + 1), dae, solve_context, cfg, solve_row_c_cfg) }}; /* {{ source_ref(name, var.dims, i + 1) }} */ {% set ns_out.offset = ns_out.offset + 1 %} {% endfor %} {% else %} - m->w[{{ ns_out.offset }}] = {{ alg_rhs_for_var(name, dae.f_x, cfg) }}; /* {{ name }} */ + m->w[{{ ns_out.offset }}] = {{ visible_or_alg_rhs_for_var(name, dae, solve_context, cfg, solve_row_c_cfg) }}; /* {{ name }} */ {% set ns_out.offset = ns_out.offset + 1 %} {% endif %} {% endfor %} @@ -932,6 +1204,55 @@ typedef struct { fmi2Real nextEventTime; } fmi2EventInfo; +/* ========================================================================= + * Restored native observables + * ========================================================================= */ +static fmi2Status get_native_observable_by_vr(const ModelInstance* m, fmi2ValueReference vr, fmi2Real* value) { +#if N_NATIVE_OBSERVABLES > 0 + if (vr < VR_OBS || vr >= VR_OBS + N_NATIVE_OBSERVABLES) { + return fmi2Error; + } + + const double t = m->time; + const double time = m->time; + (void)t; (void)time; + +{{ unpack_vars(dae.x, "m->x", pre_array="m->x") }} +{{ unpack_vars(dae.y, "m->y", pre_array="m->y") }} +{{ unpack_vars(dae.w, "m->w", pre_array="m->w") }} +{{ unpack_vars(dae.u, "m->u", pre_array="m->u") }} +{{ unpack_vars(dae.p, "m->p") }} +{{ unpack_vars(dae.constants, "m->constants") }} +{{ unpack_vars(dae.z, "m->z", pre_array="m->pre_z") }} +{{ unpack_vars(dae.m, "m->m", pre_array="m->pre_m") }} +{% for observable in dae.__rumoca_observables | default([]) %} +#define {{ observable.name | sanitize }} ({{ render_expr(observable.expr, cfg) }}) +{% endfor %} + + switch (vr) { +{% set ns_obs = namespace(offset=0) %} +{% for observable in dae.__rumoca_observables | default([]) %} +{% if observable.causality | default("local") == "output" %} + case VR_OBS + {{ ns_obs.offset }}: + *value = {{ observable.name | sanitize }}; + return fmi2OK; +{% set ns_obs.offset = ns_obs.offset + 1 %} +{% endif %} +{% endfor %} + default: +{% for observable in dae.__rumoca_observables | default([]) %} +#undef {{ observable.name | sanitize }} +{% endfor %} + return fmi2Error; + } +#else + (void)m; + (void)vr; + (void)value; + return fmi2Error; +#endif +} + /* Helper: get real value by value reference */ static fmi2Status get_real_by_vr(const ModelInstance* m, fmi2ValueReference vr, fmi2Real* value) { if (vr >= VR_X && vr < VR_X + N_STATES) { @@ -950,6 +1271,8 @@ static fmi2Status get_real_by_vr(const ModelInstance* m, fmi2ValueReference vr, *value = m->z[vr - VR_Z]; } else if (vr >= VR_M && vr < VR_M + N_DISCRETE_VAL) { *value = m->m[vr - VR_M]; + } else if (vr >= VR_OBS && vr < VR_OBS + N_NATIVE_OBSERVABLES) { + return get_native_observable_by_vr(m, vr, value); } else { return fmi2Error; } @@ -1051,6 +1374,9 @@ FMI2_EXPORT fmi2Status fmi2EnterInitializationMode(fmi2Component c) { FMI2_EXPORT fmi2Status fmi2ExitInitializationMode(fmi2Component c) { ModelInstance* m = (ModelInstance*)c; + apply_parameter_bindings(m); + apply_initial_equations(m); + /* Compute initial derivatives and outputs */ compute_derivatives(m); compute_event_indicators(m); @@ -1073,6 +1399,8 @@ FMI2_EXPORT fmi2Status fmi2GetReal( if (m->dirty_values) { compute_derivatives(m); + compute_outputs(m); + m->dirty_values = 0; } for (size_t i = 0; i < nvr; i++) { @@ -1093,6 +1421,8 @@ FMI2_EXPORT fmi2Status fmi2SetReal( fmi2Status s = set_real_by_vr(m, vr[i], value[i]); if (s != fmi2OK) return s; } + apply_parameter_bindings(m); + m->dirty_values = 1; return fmi2OK; } @@ -1404,20 +1734,40 @@ FMI2_EXPORT fmi2Status fmi2DoStep( (void)noSetFMUStatePriorToCurrentPoint; ModelInstance* m = (ModelInstance*)c; - /* Simple forward Euler with sub-stepping */ - const double dt_max = communicationStepSize / 10.0; + /* Forward Euler with state-rate limited sub-stepping. */ + const double dt_max = fmin(60.0, communicationStepSize / 10.0); double t = currentCommunicationPoint; const double t_end = currentCommunicationPoint + communicationStepSize; +#if N_STATES > 0 + fmi2Real x_nominal[N_STATES > 0 ? N_STATES : 1]; + if (fmi2GetNominalsOfContinuousStates(c, x_nominal, N_STATES) != fmi2OK) { + return fmi2Error; + } +#endif + while (t < t_end - 1.0e-15) { double dt = t_end - t; if (dt > dt_max) dt = dt_max; m->time = t; m->dirty_values = 1; + compute_discrete_updates(m); compute_derivatives(m); #if N_STATES > 0 + for (int i = 0; i < N_STATES; i++) { + const double derivative = m->xdot[i]; + const double nominal = fmax(fabs(x_nominal[i]), 1.0e-9); + const double scale = fmax(fabs(m->x[i]), nominal); + const double abs_derivative = fabs(derivative); + if (isfinite(derivative) && abs_derivative > 1.0e-12) { + const double rate_limited_dt = 0.5 * scale / abs_derivative; + if (isfinite(rate_limited_dt) && rate_limited_dt > 0.0 && rate_limited_dt < dt) { + dt = rate_limited_dt; + } + } + } for (int i = 0; i < N_STATES; i++) { m->x[i] += dt * m->xdot[i]; } @@ -1428,6 +1778,7 @@ FMI2_EXPORT fmi2Status fmi2DoStep( m->time = t_end; m->dirty_values = 1; + compute_discrete_updates(m); compute_derivatives(m); return fmi2OK; diff --git a/crates/rumoca-phase-codegen/src/templates/fmi2/modelDescription.xml.jinja b/crates/rumoca-phase-codegen/src/templates/fmi2/modelDescription.xml.jinja index 4b14112e1..7a19a1375 100644 --- a/crates/rumoca-phase-codegen/src/templates/fmi2/modelDescription.xml.jinja +++ b/crates/rumoca-phase-codegen/src/templates/fmi2/modelDescription.xml.jinja @@ -24,6 +24,13 @@ {%- set n_p = var_size(dae.p) | trim | int -%} {%- set n_z = var_size(dae.z) | trim | int -%} {%- set n_m = var_size(dae.m) | trim | int -%} +{%- set ns_obs = namespace(total=0) -%} +{%- for observable in dae.__rumoca_observables | default([]) -%} +{%- if observable.causality | default("local") == "output" -%} +{%- set ns_obs.total = ns_obs.total + 1 -%} +{%- endif -%} +{%- endfor -%} +{%- set n_obs = ns_obs.total -%} {%- set solve_root_rows = solve.events.root_conditions.programs if solve is defined and solve.events is defined and solve.events.root_conditions is defined and solve.events.root_conditions.programs is defined else [] -%} {#- Value reference layout (contiguous blocks): @@ -44,13 +51,14 @@ {%- set vr_p = vr_w + n_w -%} {%- set vr_z = vr_p + n_p -%} {%- set vr_m = vr_z + n_z -%} -{%- set n_total = vr_m + n_m -%} +{%- set vr_obs = vr_m + n_m -%} +{%- set n_total = vr_obs + n_obs -%} {#- Macro: XML-escape a string -#} {%- macro xml_escape(s) -%} {{ s | replace("&", "&") | replace("<", "<") | replace(">", ">") | replace('"', """) }} {%- endmacro -%} {%- macro real_attr(attr_name, expr, start_index) -%} -{% if expr is not none %} {{ attr_name }}="{% if start_index is not none %}{{ render_expr_at_index(expr, start_index, cfg) }}{% else %}{{ render_expr(expr, cfg) }}{% endif %}"{% endif %} +{% if expr is not none %} {{ attr_name }}="{% if start_index is not none %}{{ render_xml_attr_expr_at_index(expr, start_index, cfg) }}{% else %}{{ render_xml_attr_expr(expr, cfg) }}{% endif %}"{% endif %} {%- endmacro -%} {#- Macro: emit a element -#} {%- macro scalar_var(name, vr, causality, variability, var, initial, start_index) -%} @@ -198,11 +206,19 @@ {{ scalar_var(name, ns_vr.offset, "local", "discrete", var, "exact", none) }} {% set ns_vr.offset = ns_vr.offset + 1 %} {% endif %} +{% endfor %} +{#- Restored native observables / outputs that disappeared during structural preparation. -#} +{% set ns_vr = namespace(offset=vr_obs) %} +{% for observable in dae.__rumoca_observables | default([]) %} +{% if observable.causality | default("local") == "output" %} +{{ scalar_var(observable.name, ns_vr.offset, observable.causality | default("local"), "continuous", observable, "calculated") }} +{% set ns_vr.offset = ns_vr.offset + 1 %} +{% endif %} {% endfor %} -{% if n_w > 0 %} +{% if n_w > 0 or n_obs > 0 %} {% set ns_idx = namespace(val=1) %} {#- Skip states, derivatives, algebraics to find output indices -#} @@ -218,6 +234,13 @@ {% set ns_idx.val = ns_idx.val + 1 %} {% endif %} +{% endfor %} +{% set ns_obs_idx = namespace(offset=0) %} +{% for observable in dae.__rumoca_observables | default([]) %} +{% if observable.causality | default("local") == "output" %} + +{% set ns_obs_idx.offset = ns_obs_idx.offset + 1 %} +{% endif %} {% endfor %} {% endif %} diff --git a/crates/rumoca-phase-codegen/src/templates/fmi2/target.toml b/crates/rumoca-phase-codegen/src/templates/fmi2/target.toml index c8629fe47..11cfbef89 100644 --- a/crates/rumoca-phase-codegen/src/templates/fmi2/target.toml +++ b/crates/rumoca-phase-codegen/src/templates/fmi2/target.toml @@ -8,6 +8,23 @@ build = "fmu" completion_message = """FMU sources compiled to: {{ out_dir }} Run ./build.sh to compile and package the .fmu""" +[[files]] +path = "resources/externalDependencies.json" +template = "externalDependencies.json.jinja" +ir = "dae" + +[[files]] +path = "resources/externalLibraries.txt" +template = "externalLibraries.txt.jinja" +ir = "dae" +allow_empty = true + +[[files]] +path = "resources/externalIncludeDirectories.txt" +template = "externalIncludeDirectories.txt.jinja" +ir = "dae" +allow_empty = true + [[files]] path = "modelDescription.xml" template = "modelDescription.xml.jinja" diff --git a/crates/rumoca-phase-codegen/src/templates/fmi3/CMakeLists.txt.jinja b/crates/rumoca-phase-codegen/src/templates/fmi3/CMakeLists.txt.jinja index 9720761fa..a22f72ee6 100644 --- a/crates/rumoca-phase-codegen/src/templates/fmi3/CMakeLists.txt.jinja +++ b/crates/rumoca-phase-codegen/src/templates/fmi3/CMakeLists.txt.jinja @@ -30,6 +30,52 @@ else() message(FATAL_ERROR "Unsupported platform for FMI 3.0 FMU packaging") endif() +set(RUMOCA_EXTERNAL_INCLUDE_DIR "$ENV{RUMOCA_EXTERNAL_INCLUDE_DIR}" CACHE PATH "Directory containing Modelica external runtime headers for modelica:// include URIs") +set(RUMOCA_EXTERNAL_INCLUDE_DIRECTORIES_FILE "${CMAKE_CURRENT_LIST_DIR}/../resources/externalIncludeDirectories.txt") +if(EXISTS "${RUMOCA_EXTERNAL_INCLUDE_DIRECTORIES_FILE}") + file(STRINGS "${RUMOCA_EXTERNAL_INCLUDE_DIRECTORIES_FILE}" RUMOCA_EXTERNAL_INCLUDE_DIRECTORIES) + list(REMOVE_DUPLICATES RUMOCA_EXTERNAL_INCLUDE_DIRECTORIES) + foreach(RUMOCA_EXTERNAL_INCLUDE_DIRECTORY IN LISTS RUMOCA_EXTERNAL_INCLUDE_DIRECTORIES) + if(RUMOCA_EXTERNAL_INCLUDE_DIRECTORY MATCHES "^modelica://") + if(NOT RUMOCA_EXTERNAL_INCLUDE_DIR) + message(FATAL_ERROR "external include directories declared in resources/externalDependencies.json; set RUMOCA_EXTERNAL_INCLUDE_DIR for modelica:// include URIs") + endif() + target_include_directories({{ model_name }} PRIVATE "${RUMOCA_EXTERNAL_INCLUDE_DIR}") + else() + if(NOT IS_DIRECTORY "${RUMOCA_EXTERNAL_INCLUDE_DIRECTORY}") + message(FATAL_ERROR "missing external include directory: ${RUMOCA_EXTERNAL_INCLUDE_DIRECTORY}") + endif() + target_include_directories({{ model_name }} PRIVATE "${RUMOCA_EXTERNAL_INCLUDE_DIRECTORY}") + endif() + endforeach() +endif() + +set(RUMOCA_EXTERNAL_LIBRARY_DIR "$ENV{RUMOCA_EXTERNAL_LIBRARY_DIR}" CACHE PATH "Directory containing Modelica external runtime libraries") +set(RUMOCA_EXTERNAL_LIBRARIES_FILE "${CMAKE_CURRENT_LIST_DIR}/../resources/externalLibraries.txt") +if(EXISTS "${RUMOCA_EXTERNAL_LIBRARIES_FILE}") + file(STRINGS "${RUMOCA_EXTERNAL_LIBRARIES_FILE}" RUMOCA_EXTERNAL_LIBRARIES) + list(REMOVE_DUPLICATES RUMOCA_EXTERNAL_LIBRARIES) + if(RUMOCA_EXTERNAL_LIBRARIES) + if(NOT RUMOCA_EXTERNAL_LIBRARY_DIR) + message(FATAL_ERROR "external runtime libraries declared in resources/externalDependencies.json; set RUMOCA_EXTERNAL_LIBRARY_DIR") + endif() + target_link_directories({{ model_name }} PRIVATE "${RUMOCA_EXTERNAL_LIBRARY_DIR}") + foreach(RUMOCA_EXTERNAL_LIBRARY IN LISTS RUMOCA_EXTERNAL_LIBRARIES) + if(WIN32) + set(RUMOCA_EXTERNAL_LIBRARY_FILE "${RUMOCA_EXTERNAL_LIBRARY_DIR}/${RUMOCA_EXTERNAL_LIBRARY}.dll") + elseif(APPLE) + set(RUMOCA_EXTERNAL_LIBRARY_FILE "${RUMOCA_EXTERNAL_LIBRARY_DIR}/lib${RUMOCA_EXTERNAL_LIBRARY}.dylib") + else() + set(RUMOCA_EXTERNAL_LIBRARY_FILE "${RUMOCA_EXTERNAL_LIBRARY_DIR}/lib${RUMOCA_EXTERNAL_LIBRARY}.so") + endif() + if(NOT EXISTS "${RUMOCA_EXTERNAL_LIBRARY_FILE}") + message(FATAL_ERROR "missing external runtime library: ${RUMOCA_EXTERNAL_LIBRARY_FILE}") + endif() + target_link_libraries({{ model_name }} PRIVATE "${RUMOCA_EXTERNAL_LIBRARY}") + endforeach() + endif() +endif() + install(TARGETS {{ model_name }} RUNTIME DESTINATION binaries/${FMU_PLATFORM} LIBRARY DESTINATION binaries/${FMU_PLATFORM}) diff --git a/crates/rumoca-phase-codegen/src/templates/fmi3/build.sh.jinja b/crates/rumoca-phase-codegen/src/templates/fmi3/build.sh.jinja index 6a25d6b63..b3f443b45 100644 --- a/crates/rumoca-phase-codegen/src/templates/fmi3/build.sh.jinja +++ b/crates/rumoca-phase-codegen/src/templates/fmi3/build.sh.jinja @@ -3,6 +3,7 @@ set -e cd "$(dirname "$0")" +ALLOW_UNRESOLVED_FLAGS="" case "$(uname -s)" in Linux*) case "$(uname -m)" in @@ -18,6 +19,7 @@ case "$(uname -s)" in *) PLATFORM=x86_64-darwin ;; esac LIB_EXT=dylib + ALLOW_UNRESOLVED_FLAGS="-Wl,-undefined,dynamic_lookup" ;; MINGW*|MSYS*|CYGWIN*) case "$(uname -m)" in @@ -30,9 +32,192 @@ case "$(uname -s)" in esac mkdir -p binaries/$PLATFORM +MODEL_BINARY="binaries/$PLATFORM/{{ model_name }}.$LIB_EXT" +UNRESOLVED_SYMBOLS_FILE="$(mktemp "${TMPDIR:-/tmp}/rumoca-unresolved.XXXXXX")" +EXTERNAL_LIB_PATHS_FILE="$(mktemp "${TMPDIR:-/tmp}/rumoca-external-libs.XXXXXX")" +cleanup() { + rm -f "$UNRESOLVED_SYMBOLS_FILE" "$EXTERNAL_LIB_PATHS_FILE" +} +trap cleanup EXIT INT TERM + +set -- +if [ -n "${RUMOCA_EXTERNAL_INCLUDE_DIR:-}" ]; then + old_ifs="$IFS" + IFS=: + for external_include_dir in $RUMOCA_EXTERNAL_INCLUDE_DIR; do + [ -n "$external_include_dir" ] || continue + if [ ! -d "$external_include_dir" ]; then + echo "missing external include directory: $external_include_dir" >&2 + exit 1 + fi + set -- "$@" "-I$external_include_dir" + done + IFS="$old_ifs" +fi + +if [ -s resources/externalIncludeDirectories.txt ]; then + while IFS= read -r include_dir; do + [ -n "$include_dir" ] || continue + case "$include_dir" in + modelica://*) + if [ -z "${RUMOCA_EXTERNAL_INCLUDE_DIR:-}" ]; then + echo "external include directories declared in resources/externalDependencies.json; set RUMOCA_EXTERNAL_INCLUDE_DIR for modelica:// include URIs" >&2 + exit 1 + fi + old_ifs="$IFS" + IFS=: + for external_include_dir in $RUMOCA_EXTERNAL_INCLUDE_DIR; do + [ -n "$external_include_dir" ] || continue + if [ ! -d "$external_include_dir" ]; then + echo "missing external include directory: $external_include_dir" >&2 + exit 1 + fi + set -- "$@" "-I$external_include_dir" + done + IFS="$old_ifs" + ;; + *) + if [ ! -d "$include_dir" ]; then + echo "missing external include directory: $include_dir" >&2 + exit 1 + fi + set -- "$@" "-I$include_dir" + ;; + esac + done </dev/null | awk '{print $NF}' | sed 's/^_//' | sort -u > "$UNRESOLVED_SYMBOLS_FILE" + +external_library_declares_unresolved_symbol() { + lib="$1" + if [ ! -s resources/externalDependencies.json ]; then + return 1 + fi + for symbol in $(awk -v lib="$lib" ' + /"symbol"[[:space:]]*:/ { + symbol=$0 + sub(/^.*"symbol"[[:space:]]*:[[:space:]]*"/, "", symbol) + sub(/".*$/, "", symbol) + next + } + /"libraries"[[:space:]]*:/ && symbol != "" { + if ($0 ~ "\"" lib "\"") print symbol + symbol="" + } + ' resources/externalDependencies.json); do + if grep -Fxq "$symbol" "$UNRESOLVED_SYMBOLS_FILE"; then + return 0 + fi + done + return 1 +} + +external_library_exports_unresolved_symbol() { + lib_file="$1" + [ -f "$lib_file" ] || return 1 + symbols_file="$(mktemp "${TMPDIR:-/tmp}/rumoca-lib-symbols.XXXXXX")" + nm -g "$lib_file" 2>/dev/null | awk '{print $NF}' | sed 's/^_//' > "$symbols_file" + while IFS= read -r symbol; do + [ -n "$symbol" ] || continue + if grep -Fxq "$symbol" "$UNRESOLVED_SYMBOLS_FILE"; then + rm -f "$symbols_file" + return 0 + fi + done < "$symbols_file" + rm -f "$symbols_file" + return 1 +} + +rewrite_darwin_runtime_paths() { + case "$PLATFORM" in + *darwin) ;; + *) return 0 ;; + esac + command -v install_name_tool >/dev/null 2>&1 || { + echo "install_name_tool is required to make Darwin FMU external libraries loader-relative" >&2 + exit 1 + } + while IFS= read -r external_lib_file; do + [ -n "$external_lib_file" ] || continue + external_lib_base="$(basename "$external_lib_file")" + copied_lib="binaries/$PLATFORM/$external_lib_base" + install_name_tool -id "@loader_path/$external_lib_base" "$copied_lib" + install_name_tool -change "$external_lib_file" "@loader_path/$external_lib_base" "$MODEL_BINARY" 2>/dev/null || true + install_name_tool -change "$external_lib_base" "@loader_path/$external_lib_base" "$MODEL_BINARY" 2>/dev/null || true + install_name_tool -change "@rpath/$external_lib_base" "@loader_path/$external_lib_base" "$MODEL_BINARY" 2>/dev/null || true + done < "$EXTERNAL_LIB_PATHS_FILE" + + while IFS= read -r source_lib_file; do + [ -n "$source_lib_file" ] || continue + source_lib="binaries/$PLATFORM/$(basename "$source_lib_file")" + while IFS= read -r dependency_lib_file; do + [ -n "$dependency_lib_file" ] || continue + dependency_base="$(basename "$dependency_lib_file")" + install_name_tool -change "$dependency_lib_file" "@loader_path/$dependency_base" "$source_lib" 2>/dev/null || true + install_name_tool -change "$dependency_base" "@loader_path/$dependency_base" "$source_lib" 2>/dev/null || true + install_name_tool -change "@rpath/$dependency_base" "@loader_path/$dependency_base" "$source_lib" 2>/dev/null || true + done < "$EXTERNAL_LIB_PATHS_FILE" + done < "$EXTERNAL_LIB_PATHS_FILE" +} + +EXTERNAL_LIBS_NEEDED=0 +if [ -s resources/externalLibraries.txt ]; then + while IFS= read -r lib; do + [ -n "$lib" ] || continue + if [ -z "${RUMOCA_EXTERNAL_LIBRARY_DIR:-}" ]; then + if external_library_declares_unresolved_symbol "$lib"; then + echo "external runtime libraries declared in resources/externalDependencies.json; set RUMOCA_EXTERNAL_LIBRARY_DIR" >&2 + exit 1 + fi + continue + fi + case "$PLATFORM" in + *windows) lib_file="$RUMOCA_EXTERNAL_LIBRARY_DIR/$lib.$LIB_EXT" ;; + *) lib_file="$RUMOCA_EXTERNAL_LIBRARY_DIR/lib$lib.$LIB_EXT" ;; + esac + if external_library_declares_unresolved_symbol "$lib"; then + if [ ! -f "$lib_file" ]; then + echo "missing external runtime library: $lib_file" >&2 + exit 1 + fi + elif ! external_library_exports_unresolved_symbol "$lib_file"; then + continue + fi + if [ ! -f "$lib_file" ]; then + echo "external runtime libraries declared in resources/externalDependencies.json; set RUMOCA_EXTERNAL_LIBRARY_DIR" >&2 + exit 1 + fi + set -- "$@" "-L$RUMOCA_EXTERNAL_LIBRARY_DIR" "-l$lib" + printf '%s\n' "$lib_file" >> "$EXTERNAL_LIB_PATHS_FILE" + EXTERNAL_LIBS_NEEDED=1 + done < 0 %} {% set solve_full_jacobian_rows = solve_full_jacobian_rows if solve_full_jacobian_rows is defined else (solve.artifacts.continuous.full_jacobian_v.programs if solve is defined and solve.artifacts is defined and solve.artifacts.continuous is defined and solve.artifacts.continuous.full_jacobian_v is defined and solve.artifacts.continuous.full_jacobian_v.programs is defined else []) %} +{% set solve_context = solve if solve is defined else {} %} +{% set solve_visible_names = solve.visible_names if solve is defined and solve.visible_names is defined else [] %} {% set solve_row_c_cfg = {"time": "m->time", "y": "__rumoca_solve_y(m, {})", "p": "__rumoca_solve_p(m, {})"} %} {% set solve_slot_assign_c_cfg = {"y_set": "__rumoca_solve_set_y(m, {}, {})", "p_set": "__rumoca_solve_set_p(m, {}, {})"} %} {% set solve_jacobian_c_cfg = {"time": "m->time", "y": "__rumoca_solve_y(m, {})", "p": "__rumoca_solve_p(m, {})", "seed": "seed[{}]"} %} @@ -72,6 +74,10 @@ /* Complex number support — Modelica Complex record projected to real part. */ #define Complex(re, im) (re) +#define REAL_C(x) (x) +#define Modelica_Units_SI_TemperatureDifference(x) (x) +#define Modelica_Units_SI_MassFraction(x) (x) +static inline double linspace(double start, double stop, double n) { (void)stop; (void)n; return start; } /* Array helper: sum all elements of a double array */ static inline double __rumoca_sum_d(const double* arr, int n) { @@ -167,6 +173,9 @@ static inline double zeros(int n) { (void)n; return 0.0; } static inline double ones(int n) { (void)n; return 1.0; } static inline double fill(double val, int n) { (void)n; return val; } static inline int size(const double* arr, int dim) { (void)arr; (void)dim; return 0; } +static inline double sign(double x) { return (x > 0.0) - (x < 0.0); } +static inline double getInstanceName(void) { return 0.0; } +static inline double Buildings_ThermalZones_EnergyPlus_9_6_0_ThermalZone_Medium(double x) { return x; } /* Modelica interval() builtin — clocked partition intrinsic (MLS §16.10). */ static inline double interval(double u) { (void)u; return 0.0; } @@ -224,6 +233,16 @@ static int ModelicaStrings_scanInteger(double string, int startIndex, int unsign static double Modelica_Blocks_Types_ExternalCombiTable1D() { return 0.0; } static double Modelica_Blocks_Types_ExternalCombiTimeTable() { return 0.0; } +#ifndef RUMOCA_SOLVE_EXTERNAL_CALL +static double __rumoca_solve_external_call_default(const char* function, int output_index, int arg_count, const double* args) { + (void)args; + fprintf(stderr, "Rumoca solve ExternalCall requires native runtime bridge: %s output_index=%d arg_count=%d\n", function, output_index, arg_count); + abort(); + return NAN; +} +#define RUMOCA_SOLVE_EXTERNAL_CALL __rumoca_solve_external_call_default +#endif + /* Named argument passthrough macros */ #define __rumoca_named_arg___tableName(x) (x) #define __rumoca_named_arg___fileName(x) (x) @@ -232,17 +251,106 @@ static double Modelica_Blocks_Types_ExternalCombiTimeTable() { return 0.0; } #define __rumoca_named_arg___extrapolation(x) (x) #define __rumoca_named_arg___verboseRead(x) (x) #define __rumoca_named_arg___verboseExtrapolation(x) (x) +#define __rumoca_named_arg___T(x) (x) +#define __rumoca_named_arg___TWetBul(x) (x) +#define __rumoca_named_arg___X(...) (__VA_ARGS__) +#define __rumoca_named_arg___X_w(x) (x) +#define __rumoca_named_arg___a(x) (x) +#define __rumoca_named_arg___b(x) (x) +#define __rumoca_named_arg___buildingsRootFileLocation(x) (x) +#define __rumoca_named_arg___c(x) (x) +#define __rumoca_named_arg___caseSensitive(x) (x) +#define __rumoca_named_arg___d(x) (x) +#define __rumoca_named_arg___delta(x) (x) +#define __rumoca_named_arg___deltaInv(x) (x) +#define __rumoca_named_arg___deltaX(x) (x) +#define __rumoca_named_arg___deltax(x) (x) +#define __rumoca_named_arg___derivatives_delta(x) (x) +#define __rumoca_named_arg___derivatives_structure(x) (x) +#define __rumoca_named_arg___diameter(x) (x) +#define __rumoca_named_arg___dummy(x) (x) +#define __rumoca_named_arg___e(x) (x) +#define __rumoca_named_arg___ensureMonotonicity(x) (x) +#define __rumoca_named_arg___epName(x) (x) +#define __rumoca_named_arg___epwName(x) (x) +#define __rumoca_named_arg___f(x) (x) +#define __rumoca_named_arg___fmuName(x) (x) +#define __rumoca_named_arg___h(x) (x) +#define __rumoca_named_arg___idfName(x) (x) +#define __rumoca_named_arg___idfVersion(x) (x) +#define __rumoca_named_arg___initialCall(x) (x) +#define __rumoca_named_arg___inpNames(x) (x) +#define __rumoca_named_arg___inpUnits(x) (x) +#define __rumoca_named_arg___jsonKeysValues(x) (x) +#define __rumoca_named_arg___jsonName(x) (x) +#define __rumoca_named_arg___modelicaInstanceName(x) (x) +#define __rumoca_named_arg___modelicaNameBuilding(x) (x) +#define __rumoca_named_arg___mu_a(x) (x) +#define __rumoca_named_arg___mu_b(x) (x) +#define __rumoca_named_arg___nDer(x) (x) +#define __rumoca_named_arg___nInp(x) (x) +#define __rumoca_named_arg___nOut(x) (x) +#define __rumoca_named_arg___nParOut(x) (x) +#define __rumoca_named_arg___nY(x) (x) +#define __rumoca_named_arg___neg(x) (x) +#define __rumoca_named_arg___objectType(x) (x) +#define __rumoca_named_arg___outNames(x) (x) +#define __rumoca_named_arg___outUnits(x) (x) +#define __rumoca_named_arg___p(x) (x) +#define __rumoca_named_arg___pSat(x) (x) +#define __rumoca_named_arg___p_w(x) (x) +#define __rumoca_named_arg___parOutNames(x) (x) +#define __rumoca_named_arg___parOutUnits(x) (x) +#define __rumoca_named_arg___per(x) (x) +#define __rumoca_named_arg___phi(x) (x) +#define __rumoca_named_arg___pos(x) (x) +#define __rumoca_named_arg___printUnit(x) (x) +#define __rumoca_named_arg___r_V(x) (x) +#define __rumoca_named_arg___relativeSurfaceTolerance(x) (x) +#define __rumoca_named_arg___rho_a(x) (x) +#define __rumoca_named_arg___rho_b(x) (x) +#define __rumoca_named_arg___spawnExe(x) (x) +#define __rumoca_named_arg___state(x) (x) +#define __rumoca_named_arg___strict(x) (x) +#define __rumoca_named_arg___string1(x) (x) +#define __rumoca_named_arg___string2(x) (x) +#define __rumoca_named_arg___u(x) (x) +#define __rumoca_named_arg___usePrecompiledFMU(x) (x) +#define __rumoca_named_arg___x(x) (x) +#define __rumoca_named_arg___x1(x) (x) +#define __rumoca_named_arg___x2(x) (x) +#define __rumoca_named_arg___x_small(x) (x) +#define __rumoca_named_arg___y(x) (x) +#define __rumoca_named_arg___y1(x) (x) +#define __rumoca_named_arg___y2(x) (x) /* ========================================================================= * Enumeration literal constants * ========================================================================= */ +{% set runtime_field_macro_names = ["time", "x", "xdot", "y", "u", "w", "p", "z", "pre_z", "m", "pre_m", "constants", "event_indicators", "event_indicators_prev", "state", "dirty_values", "is_new_event_iteration"] %} {% for name, ordinal in dae.enum_literal_ordinals | items %} -#define {{ symbol(symbols, name) }} {{ ordinal }} +{% set enum_symbol = symbol(symbols, name) %} +{% if enum_symbol not in runtime_field_macro_names %} +#define {{ enum_symbol }} {{ ordinal }} +{% endif %} {% endfor %} /* Enumeration type constructors — identity macros for enum type casts */ {% for type_name in dae.enum_type_names | default([]) %} -#define {{ symbol(symbols, type_name) }}(x) (x) +{% set enum_type_symbol = symbol(symbols, type_name) %} +{% if enum_type_symbol not in runtime_field_macro_names %} +#define {{ enum_type_symbol }}(x) (x) +{% endif %} +{% endfor %} + +/* Source-reference aliases: renderers may emit sanitized full source refs + * while local unpack aliases use the allocated target symbol. */ +{% for name in dae.symbol_refs %} +{% set source_alias = name | sanitize %} +{% set target_symbol = symbol(symbols, name) %} +{% if source_alias != target_symbol and source_alias not in runtime_field_macro_names and source_alias not in c_symbol_policy.reserved %} +#define {{ source_alias }} {{ target_symbol }} +{% endif %} {% endfor %} {#- Macro: compute total scalar size of a variable map -#} @@ -258,6 +366,59 @@ static double Modelica_Blocks_Types_ExternalCombiTimeTable() { return 0.0; } {{ ns.total }} {%- endmacro %} +{#- Generate Solve IR y-slot accessors from solve.visible_names and actual + target storage layout. This keeps zero-extent variables from shifting the + runtime solver index space away from generated FMI arrays. #} +{% macro solve_y_get_cases_for_vars(vars, storage) -%} +{%- set ns = namespace(offset=0) -%} +{%- for name, var in vars | items -%} +{%- if var.dims -%} +{%- set sz = var.dims | product -%} +{%- for i in range(sz) -%} +{%- set scalar_name = source_ref(name, var.dims, i + 1) -%} +{%- for visible in solve_visible_names -%} +{%- if visible == scalar_name %} + case {{ loop.index0 }}: return {{ storage }}[{{ ns.offset }}]; /* {{ scalar_name }} */ +{%- endif -%} +{%- endfor -%} +{%- set ns.offset = ns.offset + 1 -%} +{%- endfor -%} +{%- else -%} +{%- for visible in solve_visible_names -%} +{%- if visible == name %} + case {{ loop.index0 }}: return {{ storage }}[{{ ns.offset }}]; /* {{ name }} */ +{%- endif -%} +{%- endfor -%} +{%- set ns.offset = ns.offset + 1 -%} +{%- endif -%} +{%- endfor -%} +{%- endmacro %} + +{% macro solve_y_set_cases_for_vars(vars, storage) -%} +{%- set ns = namespace(offset=0) -%} +{%- for name, var in vars | items -%} +{%- if var.dims -%} +{%- set sz = var.dims | product -%} +{%- for i in range(sz) -%} +{%- set scalar_name = source_ref(name, var.dims, i + 1) -%} +{%- for visible in solve_visible_names -%} +{%- if visible == scalar_name %} + case {{ loop.index0 }}: {{ storage }}[{{ ns.offset }}] = value; return; /* {{ scalar_name }} */ +{%- endif -%} +{%- endfor -%} +{%- set ns.offset = ns.offset + 1 -%} +{%- endfor -%} +{%- else -%} +{%- for visible in solve_visible_names -%} +{%- if visible == name %} + case {{ loop.index0 }}: {{ storage }}[{{ ns.offset }}] = value; return; /* {{ name }} */ +{%- endif -%} +{%- endfor -%} +{%- set ns.offset = ns.offset + 1 -%} +{%- endif -%} +{%- endfor -%} +{%- endmacro %} + {#- Macro: unpack variables from a ModelInstance array into local C aliases. Also emits pointer aliases for array variables so Index expressions (base[subscript] style) can access them via pointer. -#} @@ -312,6 +473,7 @@ static double Modelica_Blocks_Types_ExternalCombiTimeTable() { return 0.0; } #define N_CONSTANTS {{ var_size(dae.constants) | trim }} #define N_DISCRETE_REAL {{ var_size(dae.z) | trim }} #define N_DISCRETE_VAL {{ var_size(dae.m) | trim }} +#define N_NATIVE_OBSERVABLES {{ dae.__rumoca_observables | default([]) | length }} #define N_EVENT_INDICATORS {{ solve_root_rows | length }} #define N_PERIODIC_CLOCKS {{ dae.clock_schedules | default([]) | length }} #define N_TRIGGERED_CLOCKS {{ dae.triggered_clock_conditions | default([]) | length }} @@ -334,6 +496,7 @@ static double Modelica_Blocks_Types_ExternalCombiTimeTable() { return 0.0; } #define NVAR_P {{ dae.p | length }} #define NVAR_Z {{ dae.z | length }} #define NVAR_M {{ dae.m | length }} +#define NVAR_OBS {% set ns_obs = namespace(total=0) %}{% for observable in dae.__rumoca_observables | default([]) %}{% if observable.causality | default("local") == "output" %}{% set ns_obs.total = ns_obs.total + 1 %}{% endif %}{% endfor %}{{ ns_obs.total }} /* Value reference offsets (per-variable VRs, FMI 3.0 native arrays) */ #define VR_X 0 @@ -422,6 +585,12 @@ typedef struct { } ModelInstance; static double __rumoca_solve_y(const ModelInstance* m, int index) { + switch (index) { +{{ solve_y_get_cases_for_vars(dae.x, "m->x") }} +{{ solve_y_get_cases_for_vars(dae.y, "m->y") }} +{{ solve_y_get_cases_for_vars(dae.w, "m->w") }} + default: break; + } if (index < N_STATES) { return m->x[index]; } @@ -456,6 +625,12 @@ static double __rumoca_solve_p(const ModelInstance* m, int index) { } static void __rumoca_solve_set_y(ModelInstance* m, int index, double value) { + switch (index) { +{{ solve_y_set_cases_for_vars(dae.x, "m->x") }} +{{ solve_y_set_cases_for_vars(dae.y, "m->y") }} +{{ solve_y_set_cases_for_vars(dae.w, "m->w") }} + default: break; + } if (index < N_STATES) { m->x[index] = value; return; @@ -501,8 +676,10 @@ static void initialize_defaults(ModelInstance* m); static void compute_algebraics(ModelInstance* m); static void compute_derivatives(ModelInstance* m); static void compute_outputs(ModelInstance* m); +static void compute_initial_updates(ModelInstance* m); static void compute_discrete_updates(ModelInstance* m); static void compute_event_indicators(ModelInstance* m); +static fmi3Status get_native_observable_by_vr(const ModelInstance* m, fmi3ValueReference vr, fmi3Float64* values, int* count); static double event_right_limit_time(double event_time) { double scale = 1.0 + fabs(event_time); @@ -561,7 +738,7 @@ static void snapshot_pre_parameters(ModelInstance* m) { {#- Macro: check if a function has any record-typed parameters -#} {% macro has_complex_params(func) -%} -{%- for p in func.inputs -%}{%- if p.type_class == "Record" -%}yes{%- endif -%}{%- endfor -%} +{%- for p in func.inputs -%}{%- if p.type_class | default("") == "Record" -%}yes{%- endif -%}{%- endfor -%} {%- endmacro %} {#- Macro: render a C parameter list for one-output user functions. #} @@ -629,6 +806,8 @@ static void snapshot_pre_parameters(ModelInstance* m) { {% for p in func.inputs %} (void){{ symbol(symbols, p.name) }}; {% endfor %} return 0.0; +{% elif unsupported_c_function_name(func_name) or unsupported_c_function_body(func) %} +{{ unsupported_c_function_body(func) | indent(4) }} {% elif func.external %} /* External {{ func.external.language }} function */ {% if func.external.output_name %} @@ -640,8 +819,9 @@ static void snapshot_pre_parameters(ModelInstance* m) { {% else %} {{ function_input_aliases(func) }} {{ function_locals_and_outputs(func) }} -{% if func.body | length > 0 %} -{{ render_statements(func.body, cfg, " ") }} +{% set func_body = func["body"] %} +{% if func_body | length > 0 %} +{{ render_function_statements(func_body, cfg, " ", return_expr) }} {% endif %} return {{ return_expr }}; {% endif %} @@ -769,7 +949,7 @@ static void initialize_defaults(ModelInstance* m) { {% set ns_idx = namespace(offset=0) %} {% for name, var in dae.constants | items %} - {{ symbol(symbols, name) }} = {% if var.start and is_string_literal(var.start) != "yes" %}{{ render_expr(var.start, cfg) }}{% else %}0.0{% endif %}; + {{ symbol(symbols, name) }} = {% if var.start and is_string_literal(var.start) != "yes" and expr_has_dynamic_multidim_index(var.start) != "yes" %}{{ render_expr(var.start, cfg) }}{% else %}0.0{% endif %}; m->constants[{{ ns_idx.offset }}] = {{ symbol(symbols, name) }}; /* {{ name }} */ {% set ns_idx.offset = ns_idx.offset + 1 %} {% endfor %} @@ -779,12 +959,12 @@ static void initialize_defaults(ModelInstance* m) { {% if var.dims %} {% set sz = var.dims | product %} {% for i in range(sz) %} - {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }} = {% if var.start and is_string_literal(var.start) != "yes" %}{{ render_expr_at_index(var.start, i + 1, cfg) }}{% else %}0.0{% endif %}; + {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }} = {% if var.start and is_string_literal(var.start) != "yes" and expr_has_dynamic_multidim_index(var.start) != "yes" %}{{ render_expr_at_index(var.start, i + 1, cfg) }}{% else %}0.0{% endif %}; m->p[{{ ns_idx.offset }}] = {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }}; /* {{ name }}[{{ i + 1 }}] */ {% set ns_idx.offset = ns_idx.offset + 1 %} {% endfor %} {% else %} - {{ symbol(symbols, name) }} = {% if var.start and is_string_literal(var.start) != "yes" %}{{ render_expr(var.start, cfg) }}{% else %}0.0{% endif %}; + {{ symbol(symbols, name) }} = {% if var.start and is_string_literal(var.start) != "yes" and expr_has_dynamic_multidim_index(var.start) != "yes" %}{{ render_expr(var.start, cfg) }}{% else %}0.0{% endif %}; m->p[{{ ns_idx.offset }}] = {{ symbol(symbols, name) }}; /* {{ name }} */ {% set ns_idx.offset = ns_idx.offset + 1 %} {% endif %} @@ -795,11 +975,11 @@ static void initialize_defaults(ModelInstance* m) { {% if var.dims %} {% set sz = var.dims | product %} {% for i in range(sz) %} - m->x[{{ ns_idx.offset }}] = {% if var.start and is_string_literal(var.start) != "yes" %}{{ render_expr_at_index(var.start, i + 1, cfg) }}{% else %}0.0{% endif %}; /* {{ name }}[{{ i + 1 }}] */ + m->x[{{ ns_idx.offset }}] = {% if var.start and is_string_literal(var.start) != "yes" and expr_has_dynamic_multidim_index(var.start) != "yes" %}{{ render_expr_at_index(var.start, i + 1, cfg) }}{% else %}0.0{% endif %}; /* {{ name }}[{{ i + 1 }}] */ {% set ns_idx.offset = ns_idx.offset + 1 %} {% endfor %} {% else %} - m->x[{{ ns_idx.offset }}] = {% if var.start and is_string_literal(var.start) != "yes" %}{{ render_expr(var.start, cfg) }}{% else %}0.0{% endif %}; /* {{ name }} */ + m->x[{{ ns_idx.offset }}] = {% if var.start and is_string_literal(var.start) != "yes" and expr_has_dynamic_multidim_index(var.start) != "yes" %}{{ render_expr(var.start, cfg) }}{% else %}0.0{% endif %}; /* {{ name }} */ {% set ns_idx.offset = ns_idx.offset + 1 %} {% endif %} {% endfor %} @@ -809,12 +989,12 @@ static void initialize_defaults(ModelInstance* m) { {% if var.dims %} {% set sz = var.dims | product %} {% for i in range(sz) %} - {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }} = {% if var.start and is_string_literal(var.start) != "yes" %}{{ render_expr_at_index(var.start, i + 1, cfg) }}{% else %}0.0{% endif %}; + {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }} = {% if var.start and is_string_literal(var.start) != "yes" and expr_has_dynamic_multidim_index(var.start) != "yes" %}{{ render_expr_at_index(var.start, i + 1, cfg) }}{% else %}0.0{% endif %}; m->u[{{ ns_idx.offset }}] = {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }}; /* {{ name }}[{{ i + 1 }}] */ {% set ns_idx.offset = ns_idx.offset + 1 %} {% endfor %} {% else %} - {{ symbol(symbols, name) }} = {% if var.start and is_string_literal(var.start) != "yes" %}{{ render_expr(var.start, cfg) }}{% else %}0.0{% endif %}; + {{ symbol(symbols, name) }} = {% if var.start and is_string_literal(var.start) != "yes" and expr_has_dynamic_multidim_index(var.start) != "yes" %}{{ render_expr(var.start, cfg) }}{% else %}0.0{% endif %}; m->u[{{ ns_idx.offset }}] = {{ symbol(symbols, name) }}; /* {{ name }} */ {% set ns_idx.offset = ns_idx.offset + 1 %} {% endif %} @@ -825,12 +1005,12 @@ static void initialize_defaults(ModelInstance* m) { {% if var.dims %} {% set sz = var.dims | product %} {% for i in range(sz) %} - {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }} = {% if var.start and is_string_literal(var.start) != "yes" %}{{ render_expr_at_index(var.start, i + 1, cfg) }}{% else %}0.0{% endif %}; + {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }} = {% if var.start and is_string_literal(var.start) != "yes" and expr_has_dynamic_multidim_index(var.start) != "yes" %}{{ render_expr_at_index(var.start, i + 1, cfg) }}{% else %}0.0{% endif %}; m->z[{{ ns_idx.offset }}] = {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }}; /* {{ name }}[{{ i + 1 }}] */ {% set ns_idx.offset = ns_idx.offset + 1 %} {% endfor %} {% else %} - {{ symbol(symbols, name) }} = {% if var.start and is_string_literal(var.start) != "yes" %}{{ render_expr(var.start, cfg) }}{% else %}0.0{% endif %}; + {{ symbol(symbols, name) }} = {% if var.start and is_string_literal(var.start) != "yes" and expr_has_dynamic_multidim_index(var.start) != "yes" %}{{ render_expr(var.start, cfg) }}{% else %}0.0{% endif %}; m->z[{{ ns_idx.offset }}] = {{ symbol(symbols, name) }}; /* {{ name }} */ {% set ns_idx.offset = ns_idx.offset + 1 %} {% endif %} @@ -848,6 +1028,93 @@ static void initialize_defaults(ModelInstance* m) { #endif } +/* ========================================================================= + * FMI initialization binding propagation + * ========================================================================= */ +static void apply_parameter_bindings(ModelInstance* m) { +{% if dae.p | length > 0 %} + const double t = m->time; + const double time = m->time; + (void)t; (void)time; + +{{ unpack_vars(dae.constants, "m->constants") }} +{{ unpack_vars(dae.p, "m->p", mutable=true) }} +{% set ns_idx = namespace(offset=0) %} +{% for name, var in dae.p | items %} +{% if var.dims %} +{% set sz = var.dims | product %} +{% for i in range(sz) %} +{% set binding_name = source_ref(name, var.dims, i + 1) %} +{% set rhs = parameter_binding_rhs(name, var.start, i + 1, cfg) if var.start else "" %} +{% if rhs %} + {{ symbol(symbols, binding_name) }} = {{ rhs }}; + m->p[{{ ns_idx.offset }}] = {{ symbol(symbols, binding_name) }}; /* binding {{ binding_name }} */ +{% endif %} +{% set ns_idx.offset = ns_idx.offset + 1 %} +{% endfor %} +{% else %} +{% set rhs = parameter_binding_rhs(name, var.start, 0, cfg) if var.start else "" %} +{% if rhs %} + {{ symbol(symbols, name) }} = {{ rhs }}; + m->p[{{ ns_idx.offset }}] = {{ symbol(symbols, name) }}; /* binding {{ name }} */ +{% endif %} +{% set ns_idx.offset = ns_idx.offset + 1 %} +{% endif %} +{% endfor %} + m->dirty_values = 1; +{% else %} + (void)m; +{% endif %} +} + +/* ========================================================================= + * Apply explicit initial equations after parameters have their start values. + * ========================================================================= */ +static void compute_initial_updates(ModelInstance* m) { +{% if dae.initial_equations | default([]) | length > 0 %} + const double t = m->time; + const double time = m->time; + (void)t; (void)time; + +{{ unpack_vars(dae.x, "m->x", pre_array="m->x") }} +{{ unpack_vars(dae.y, "m->y", pre_array="m->y") }} +{{ unpack_vars(dae.w, "m->w", pre_array="m->w") }} +{{ unpack_vars(dae.u, "m->u") }} +{{ unpack_vars(dae.p, "m->p") }} +{{ unpack_vars(dae.constants, "m->constants") }} +{{ unpack_vars(dae.z, "m->z", pre_array="m->pre_z") }} +{{ unpack_vars(dae.m, "m->m", pre_array="m->pre_m") }} + +{% set ns_init = namespace(offset=0) %} +{% for name, var in dae.x | items %} +{% if var.dims %} +{% set sz = var.dims | product %} +{% for i in range(sz) %} +{% set scalar_name = source_ref(name, var.dims, i + 1) %} +{% set scalar_symbol = symbol(symbols, scalar_name) %} +{% for eq in dae.initial_equations | default([]) %} +{% if eq.lhs is defined and eq.rhs is defined and render_expr(eq.lhs, cfg) == scalar_symbol %} + m->x[{{ ns_init.offset }}] = {{ render_expr(eq.rhs, cfg) }}; /* initial equation: {{ scalar_name }} */ +{% endif %} +{% endfor %} +{% set ns_init.offset = ns_init.offset + 1 %} +{% endfor %} +{% else %} +{% set scalar_symbol = symbol(symbols, name) %} +{% for eq in dae.initial_equations | default([]) %} +{% if eq.lhs is defined and eq.rhs is defined and render_expr(eq.lhs, cfg) == scalar_symbol %} + m->x[{{ ns_init.offset }}] = {{ render_expr(eq.rhs, cfg) }}; /* initial equation: {{ name }} */ +{% endif %} +{% endfor %} +{% set ns_init.offset = ns_init.offset + 1 %} +{% endif %} +{% endfor %} + m->dirty_values = 1; +{% else %} + (void)m; +{% endif %} +} + /* ========================================================================= * Compute algebraic variables from f_x equations * ========================================================================= */ @@ -893,12 +1160,12 @@ static void compute_algebraics(ModelInstance* m) { {% if var.dims %} {% set sz = var.dims | product %} {% for i in range(sz) %} - {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }} = {{ alg_rhs_for_var(source_ref(name, var.dims, i + 1), dae.f_x, cfg) }}; + {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }} = {{ visible_or_alg_rhs_for_var(source_ref(name, var.dims, i + 1), dae, solve_context, cfg, solve_row_c_cfg) }}; m->w[{{ ns_out.offset }}] = {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }}; /* {{ source_ref(name, var.dims, i + 1) }} */ {% set ns_out.offset = ns_out.offset + 1 %} {% endfor %} {% else %} - {{ symbol(symbols, name) }} = {{ alg_rhs_for_var(name, dae.f_x, cfg) }}; + {{ symbol(symbols, name) }} = {{ visible_or_alg_rhs_for_var(name, dae, solve_context, cfg, solve_row_c_cfg) }}; m->w[{{ ns_out.offset }}] = {{ symbol(symbols, name) }}; /* {{ name }} */ {% set ns_out.offset = ns_out.offset + 1 %} {% endif %} @@ -910,12 +1177,12 @@ static void compute_algebraics(ModelInstance* m) { {% if var.dims %} {% set sz = var.dims | product %} {% for i in range(sz) %} - {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }} = {{ alg_rhs_for_var(source_ref(name, var.dims, i + 1), dae.f_x, cfg) }}; + {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }} = {{ visible_or_alg_rhs_for_var(source_ref(name, var.dims, i + 1), dae, solve_context, cfg, solve_row_c_cfg) }}; m->y[{{ ns_alg.offset }}] = {{ symbol(symbols, source_ref(name, var.dims, i + 1)) }}; /* {{ source_ref(name, var.dims, i + 1) }} */ {% set ns_alg.offset = ns_alg.offset + 1 %} {% endfor %} {% else %} - {{ symbol(symbols, name) }} = {{ alg_rhs_for_var(name, dae.f_x, cfg) }}; + {{ symbol(symbols, name) }} = {{ visible_or_alg_rhs_for_var(name, dae, solve_context, cfg, solve_row_c_cfg) }}; m->y[{{ ns_alg.offset }}] = {{ symbol(symbols, name) }}; /* {{ name }} */ {% set ns_alg.offset = ns_alg.offset + 1 %} {% endif %} @@ -1005,11 +1272,11 @@ static void compute_outputs(ModelInstance* m) { {% if var.dims %} {% set sz = var.dims | product %} {% for i in range(sz) %} - m->w[{{ ns_out.offset }}] = {{ alg_rhs_for_var(source_ref(name, var.dims, i + 1), dae.f_x, cfg) }}; /* {{ source_ref(name, var.dims, i + 1) }} */ + m->w[{{ ns_out.offset }}] = {{ visible_or_alg_rhs_for_var(source_ref(name, var.dims, i + 1), dae, solve_context, cfg, solve_row_c_cfg) }}; /* {{ source_ref(name, var.dims, i + 1) }} */ {% set ns_out.offset = ns_out.offset + 1 %} {% endfor %} {% else %} - m->w[{{ ns_out.offset }}] = {{ alg_rhs_for_var(name, dae.f_x, cfg) }}; /* {{ name }} */ + m->w[{{ ns_out.offset }}] = {{ visible_or_alg_rhs_for_var(name, dae, solve_context, cfg, solve_row_c_cfg) }}; /* {{ name }} */ {% set ns_out.offset = ns_out.offset + 1 %} {% endif %} {% endfor %} @@ -1449,6 +1716,8 @@ FMI3_Export fmi3Status fmi3ExitInitializationMode(fmi3Instance instance) { snapshot_pre_parameters(m); /* Compute initial derivatives and outputs */ + apply_parameter_bindings(m); + compute_initial_updates(m); compute_derivatives(m); if (m->evaluation_error) { m->state = modelError; @@ -1493,6 +1762,8 @@ FMI3_Export fmi3Status fmi3GetFloat64( if (m->dirty_values) { compute_derivatives(m); + compute_outputs(m); + m->dirty_values = 0; } if (m->evaluation_error) return fmi3Error; @@ -1524,6 +1795,7 @@ FMI3_Export fmi3Status fmi3SetFloat64( if (s != fmi3OK) return s; val_idx += count; } + m->dirty_values = 1; return fmi3OK; } @@ -2331,6 +2603,7 @@ static fmi3Status rk45_eval(ModelInstance* m, double t, const fmi3Float64 x[], f memcpy(m->x, x, N_STATES * sizeof(fmi3Float64)); m->time = t; m->dirty_values = 1; + compute_discrete_updates(m); compute_derivatives(m); if (m->evaluation_error) return fmi3Error; memcpy(dxdt, m->xdot, N_STATES * sizeof(fmi3Float64)); diff --git a/crates/rumoca-phase-codegen/src/templates/fmi3/modelDescription.xml.jinja b/crates/rumoca-phase-codegen/src/templates/fmi3/modelDescription.xml.jinja index 7694d49da..57ffa7237 100644 --- a/crates/rumoca-phase-codegen/src/templates/fmi3/modelDescription.xml.jinja +++ b/crates/rumoca-phase-codegen/src/templates/fmi3/modelDescription.xml.jinja @@ -23,6 +23,13 @@ {%- set nvar_p = dae.p | length -%} {%- set nvar_z = dae.z | length -%} {%- set nvar_m = dae.m | length -%} +{%- set ns_obs = namespace(total=0) -%} +{%- for observable in dae.__rumoca_observables | default([]) -%} +{%- if observable.causality | default("local") == "output" -%} +{%- set ns_obs.total = ns_obs.total + 1 -%} +{%- endif -%} +{%- endfor -%} +{%- set nvar_obs = ns_obs.total -%} {#- FMI 3.0 Native Array VR layout (one VR per variable): 0 .. nvar_x-1 : states (x) @@ -66,7 +73,7 @@ {{ s | replace("&", "&") | replace("<", "<") | replace(">", ">") | replace('"', """) }} {%- endmacro -%} {%- macro float64_attr(attr_name, expr, element_count) -%} -{% if expr is not none %} {{ attr_name }}="{% if element_count is not none %}{% for i in range(element_count) %}{{ render_expr_at_index(expr, i + 1, cfg) }}{% if not loop.last %} {% endif %}{% endfor %}{% else %}{{ render_expr(expr, cfg) }}{% endif %}"{% endif %} +{% if expr is not none %} {{ attr_name }}="{% if element_count is not none %}{% for i in range(element_count) %}{{ render_xml_attr_expr_at_index(expr, i + 1, cfg) }}{% if not loop.last %} {% endif %}{% endfor %}{% else %}{{ render_xml_attr_expr(expr, cfg) }}{% endif %}"{% endif %} {%- endmacro -%} {#- Macro: emit a scalar variable element (FMI 3.0 style) -#} {%- macro float64_var(name, vr, causality, variability, var, initial) -%} @@ -173,6 +180,14 @@ {{ emit_var(name, ns_vr.offset, "local", "discrete", var, "exact") }} {% set ns_vr.offset = ns_vr.offset + 1 %} {% endfor %} +{#- Restored native observables / outputs that disappeared during structural preparation. -#} +{% set ns_vr = namespace(offset=vr_obs) %} +{% for observable in dae.__rumoca_observables | default([]) %} +{% if observable.causality | default("local") == "output" %} +{{ emit_var(observable.name, ns_vr.offset, observable.causality | default("local"), "continuous", observable, "calculated") }} +{% set ns_vr.offset = ns_vr.offset + 1 %} +{% endif %} +{% endfor %} {#- Clocks (FMI 3.0 §2.4.8 — periodic clocks from sample() calls) -#} {% for sched in dae.clock_schedules | default([]) %} @@ -187,12 +202,19 @@ {#- In FMI 3.0 with native arrays: one entry per variable -#} -{% if nvar_w > 0 %} +{% if nvar_w > 0 or nvar_obs > 0 %} {% set ns_vr = namespace(offset=vr_w) %} {% for name, var in dae.w | items %} {% set ns_vr.offset = ns_vr.offset + 1 %} {% endfor %} +{% set ns_obs_vr = namespace(offset=vr_obs) %} +{% for observable in dae.__rumoca_observables | default([]) %} +{% if observable.causality | default("local") == "output" %} + +{% set ns_obs_vr.offset = ns_obs_vr.offset + 1 %} +{% endif %} +{% endfor %} {% endif %} {% set ns_vr = namespace(offset=vr_xdot) %} {% for name, var in dae.x | items %} diff --git a/crates/rumoca-phase-codegen/src/templates/fmi3/target.toml b/crates/rumoca-phase-codegen/src/templates/fmi3/target.toml index ed2769674..4a4f64223 100644 --- a/crates/rumoca-phase-codegen/src/templates/fmi3/target.toml +++ b/crates/rumoca-phase-codegen/src/templates/fmi3/target.toml @@ -8,6 +8,23 @@ build = "fmu" completion_message = """FMU sources compiled to: {{ out_dir }} Run ./build.sh to compile and package the .fmu""" +[[files]] +path = "resources/externalDependencies.json" +template = "externalDependencies.json.jinja" +ir = "dae" + +[[files]] +path = "resources/externalLibraries.txt" +template = "externalLibraries.txt.jinja" +ir = "dae" +allow_empty = true + +[[files]] +path = "resources/externalIncludeDirectories.txt" +template = "externalIncludeDirectories.txt.jinja" +ir = "dae" +allow_empty = true + [[files]] path = "modelDescription.xml" template = "modelDescription.xml.jinja" diff --git a/crates/rumoca-phase-dae/src/algorithm_lowering.rs b/crates/rumoca-phase-dae/src/algorithm_lowering.rs index 3d984a970..b6d731903 100644 --- a/crates/rumoca-phase-dae/src/algorithm_lowering.rs +++ b/crates/rumoca-phase-dae/src/algorithm_lowering.rs @@ -1970,14 +1970,16 @@ pub(super) fn lower_algorithms_to_equations(dae: &mut Dae, flat: &Model) -> Resu for algorithm in &flat.initial_algorithms { match lower_algorithm_to_equations(dae, flat, algorithm, true) { Ok(lowered) => { + let equation_count = lowered.main.len() + lowered.f_z.len() + lowered.f_m.len(); dae.initialization.equations.extend(lowered.main); - // MLS §8.6 and §11.1: initial algorithms contribute equations - // to the initialization problem. Discrete targets still use - // the same Appendix B solved forms as model algorithms, but - // they must initialize here rather than populate runtime event - // update partitions. dae.initialization.equations.extend(lowered.f_z); dae.initialization.equations.extend(lowered.f_m); + dae.initialization + .equation_provenance + .extend(std::iter::repeat_n( + rumoca_ir_dae::InitializationEquationProvenance::User, + equation_count, + )); } Err(kind) => { return Err(ToDaeError::unsupported_algorithm( diff --git a/crates/rumoca-phase-dae/src/analysis/discrete_partition.rs b/crates/rumoca-phase-dae/src/analysis/discrete_partition.rs index ecbec5770..aca90d3ac 100644 --- a/crates/rumoca-phase-dae/src/analysis/discrete_partition.rs +++ b/crates/rumoca-phase-dae/src/analysis/discrete_partition.rs @@ -170,7 +170,7 @@ pub(crate) fn classify_residual_discrete_bucket( let mut saw_real = false; let mut saw_valued = false; for target in targets { - match discrete_bucket_for_name(dae, &target) { + match discrete_bucket_for_name(dae, residual, &target) { Some(NameDiscreteBucket::DiscreteReal) => saw_real = true, Some(NameDiscreteBucket::DiscreteValued) => saw_valued = true, None => return None, @@ -227,39 +227,122 @@ fn collect_lhs_targets(lhs: &rumoca_core::Expression, out: &mut Vec Option { - if dae - .variables - .discrete_valued - .contains_key(&flat_to_dae_var_name(name)) - || subscript_fallback_chain(name.as_str()) - .into_iter() - .any(|candidate| { - dae.variables - .discrete_valued - .contains_key(&flat_to_dae_var_name(&candidate)) - }) + if partition_contains_target_or_scalarized_lane(&dae.variables.discrete_valued, name) + || residual_target_component_reference(residual, name).is_some_and(|reference| { + !scalarized_discrete_targets_for_reference(&dae.variables.discrete_valued, &reference) + .is_empty() + }) { return Some(NameDiscreteBucket::DiscreteValued); } - if dae - .variables - .discrete_reals - .contains_key(&flat_to_dae_var_name(name)) - || subscript_fallback_chain(name.as_str()) - .into_iter() - .any(|candidate| { - dae.variables - .discrete_reals - .contains_key(&flat_to_dae_var_name(&candidate)) - }) + if partition_contains_target_or_scalarized_lane(&dae.variables.discrete_reals, name) + || residual_target_component_reference(residual, name).is_some_and(|reference| { + !scalarized_discrete_targets_for_reference(&dae.variables.discrete_reals, &reference) + .is_empty() + }) { return Some(NameDiscreteBucket::DiscreteReal); } None } +fn partition_contains_target_or_scalarized_lane( + partition: &indexmap::IndexMap, + target: &rumoca_core::VarName, +) -> bool { + if partition.contains_key(&flat_to_dae_var_name(target)) + || subscript_fallback_chain(target.as_str()) + .into_iter() + .any(|candidate| partition.contains_key(&flat_to_dae_var_name(&candidate))) + { + return true; + } + + false +} + +/// Return concrete scalar DAE variables produced from an aggregate equation +/// target. Rendered names are protocol/display data; lane recovery uses the +/// preserved component-reference parts and subscripts. +pub(crate) fn scalarized_discrete_targets_for_reference( + partition: &indexmap::IndexMap, + target_ref: &rumoca_core::ComponentReference, +) -> Vec { + partition + .iter() + .filter_map(|(name, candidate)| { + let candidate_ref = candidate.component_ref.as_ref()?; + (candidate.dims.is_empty() + && component_reference_is_scalar_lane_of(candidate_ref, target_ref)) + .then(|| crate::dae_to_flat_var_name(name)) + }) + .collect() +} + +pub(crate) fn residual_target_component_reference( + residual: &rumoca_core::Expression, + target: &rumoca_core::VarName, +) -> Option { + struct Finder<'a> { + target: &'a rumoca_core::VarName, + found: Option, + } + + impl<'a> rumoca_core::ExpressionVisitor for Finder<'a> { + fn visit_var_ref( + &mut self, + name: &rumoca_core::Reference, + _subscripts: &[rumoca_core::Subscript], + ) { + if self.found.is_none() && name.var_name() == self.target { + self.found = name.component_ref().cloned(); + } + } + } + + let mut finder = Finder { + target, + found: None, + }; + finder.visit_expression(residual); + finder.found +} + +fn component_reference_is_scalar_lane_of( + candidate: &rumoca_core::ComponentReference, + aggregate: &rumoca_core::ComponentReference, +) -> bool { + // Instantiation assigns distinct declaration identities to scalar connector + // elements, while the source aggregate equation retains the unspecialized + // declaration identity. Full scoped path + structured subscripts therefore + // define lane identity here; requiring equal DefIds would reject every + // legitimately scalarized connector array. + if candidate.parts.len() != aggregate.parts.len() { + return false; + } + + let mut added_scalar_subscript = false; + for (candidate_part, aggregate_part) in candidate.parts.iter().zip(&aggregate.parts) { + if candidate_part.ident != aggregate_part.ident + || candidate_part.subs.len() < aggregate_part.subs.len() + || !candidate_part + .subs + .iter() + .zip(&aggregate_part.subs) + .all(|(candidate_sub, aggregate_sub)| candidate_sub == aggregate_sub) + { + return false; + } + if candidate_part.subs.len() > aggregate_part.subs.len() { + added_scalar_subscript = true; + } + } + added_scalar_subscript +} + /// Returns true when expression contains clocked primitives that make the /// assignment target a clocked/discrete-time value. /// @@ -330,3 +413,88 @@ fn is_clock_intrinsic_short_name(short_name: &str) -> bool { | "interval" ) } + +#[cfg(test)] +mod tests { + use super::*; + + fn test_span() -> rumoca_core::Span { + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2) + } + + fn var_ref(name: &str) -> rumoca_core::Expression { + let name = rumoca_core::VarName::new(name); + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference( + name.as_str(), + rumoca_core::component_reference_from_flat_name(&name, test_span()).unwrap(), + ), + subscripts: Vec::new(), + span: test_span(), + } + } + + fn residual(lhs: &str, rhs: &str) -> rumoca_core::Expression { + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(var_ref(lhs)), + rhs: Box::new(var_ref(rhs)), + span: test_span(), + } + } + + fn insert_variable( + variables: &mut indexmap::IndexMap, + name: &str, + ) { + let name = rumoca_core::VarName::new(name); + let mut variable = dae::Variable::new(name.clone(), test_span()); + variable.component_ref = + rumoca_core::component_reference_from_flat_name(&name, test_span()); + variables.insert(name, variable); + } + + #[test] + fn scalarized_boolean_array_target_is_discrete_valued() { + let mut dae = dae::Dae::new(); + for name in ["parallel.split[1].set", "parallel.split[2].set", "trigger"] { + insert_variable(&mut dae.variables.discrete_valued, name); + } + assert_eq!( + classify_residual_discrete_bucket(&dae, &residual("parallel.split.set", "trigger")), + Some(ResidualDiscreteBucket::DiscreteValued), + "the aggregate Boolean target must inherit the partition of both scalarized lanes" + ); + } + + #[test] + fn scalarized_continuous_real_array_target_stays_continuous() { + let mut dae = dae::Dae::new(); + for name in ["plant.y[1]", "plant.y[2]"] { + insert_variable(&mut dae.variables.algebraics, name); + } + assert_eq!( + classify_residual_discrete_bucket(&dae, &residual("plant.y", "plant.u")), + None, + "continuous Real array equations must remain in f_x" + ); + } + + #[test] + fn scalar_lane_identity_does_not_require_unspecialized_def_id() { + let mut aggregate = rumoca_core::component_reference_from_flat_name( + &rumoca_core::VarName::new("parallel.split.set"), + test_span(), + ) + .unwrap(); + aggregate.def_id = Some(rumoca_core::DefId::new(10)); + let mut lane = rumoca_core::component_reference_from_flat_name( + &rumoca_core::VarName::new("parallel.split[2].set"), + test_span(), + ) + .unwrap(); + lane.def_id = Some(rumoca_core::DefId::new(20)); + + assert!(component_reference_is_scalar_lane_of(&lane, &aggregate)); + } +} diff --git a/crates/rumoca-phase-dae/src/analysis/variable_analysis.rs b/crates/rumoca-phase-dae/src/analysis/variable_analysis.rs index 0ea7d2b0a..bbd5ef372 100644 --- a/crates/rumoca-phase-dae/src/analysis/variable_analysis.rs +++ b/crates/rumoca-phase-dae/src/analysis/variable_analysis.rs @@ -51,6 +51,8 @@ pub(crate) fn filter_state_variables( flat: &Model, internal_inputs: &InternalInputIndex, ) -> IndexSet { + let der_vars = filter_overconstrained_alias_states(der_vars, flat); + der_vars .into_iter() .filter(|name| { @@ -62,6 +64,145 @@ pub(crate) fn filter_state_variables( .collect() } +pub(crate) fn find_overconstrained_derivative_alias_roots( + der_vars: &IndexSet, + flat: &Model, +) -> FxHashMap { + let state_record_paths = overconstrained_state_record_paths(der_vars, flat); + if state_record_paths.len() < 2 { + return FxHashMap::default(); + } + + let component_of = overconstrained_record_components(flat); + let component_roots = overconstrained_state_roots(flat, &component_of); + let mut aliases = FxHashMap::default(); + + for (name, record_path) in state_record_paths { + let Some(component_id) = component_of.get(record_path.as_str()) else { + continue; + }; + let Some(root_path) = component_roots.get(component_id) else { + continue; + }; + if oc_paths_match(&record_path, root_path) { + continue; + } + let Some(root_name) = root_variable_for_alias(&name, &record_path, root_path, flat) else { + continue; + }; + aliases.insert(name, root_name); + } + + aliases +} + +fn filter_overconstrained_alias_states( + der_vars: IndexSet, + flat: &Model, +) -> IndexSet { + let alias_roots = find_overconstrained_derivative_alias_roots(&der_vars, flat); + if alias_roots.is_empty() { + return der_vars; + } + + der_vars + .into_iter() + .filter(|name| !alias_roots.contains_key(name)) + .collect() +} + +fn overconstrained_state_record_paths( + der_vars: &IndexSet, + flat: &Model, +) -> FxHashMap { + let state_record_paths: FxHashMap = der_vars + .iter() + .filter_map(|name| { + let var = flat.variables.get(name)?; + if !var.is_overconstrained { + return None; + } + Some((name.clone(), var.oc_record_path.clone()?)) + }) + .collect(); + + state_record_paths +} + +fn overconstrained_record_components(flat: &Model) -> FxHashMap<&str, usize> { + let mut record_paths: IndexSet<&str> = IndexSet::new(); + for var in flat.variables.values() { + if !var.is_overconstrained { + continue; + } + if let Some(path) = &var.oc_record_path { + record_paths.insert(path.as_str()); + } + } + let record_paths: Vec<&str> = record_paths.into_iter().collect(); + let (component_of, _n_components) = overconstrained_interface::build_record_components( + &record_paths, + &flat.branches, + &flat.optional_edges, + ); + component_of +} + +fn root_variable_for_alias( + alias_name: &VarName, + alias_record_path: &str, + root_record_path: &str, + flat: &Model, +) -> Option { + let suffix = alias_name.as_str().strip_prefix(alias_record_path)?; + if !suffix.starts_with('.') { + return None; + } + let root_name = VarName::new(format!("{root_record_path}{suffix}")); + flat.variables.contains_key(&root_name).then_some(root_name) +} + +fn overconstrained_state_roots( + flat: &Model, + component_of: &FxHashMap<&str, usize>, +) -> HashMap { + let mut roots = HashMap::default(); + + for root in &flat.definite_roots { + for (&record_path, &component_id) in component_of { + if oc_paths_match(record_path, root) { + roots.entry(component_id).or_insert_with(|| root.clone()); + } + } + } + + let mut potential_roots = flat.potential_roots.clone(); + potential_roots.sort_by(|(left_path, left_priority), (right_path, right_priority)| { + left_priority + .cmp(right_priority) + .then_with(|| left_path.cmp(right_path)) + }); + for (root, _priority) in potential_roots { + for (&record_path, &component_id) in component_of { + if !roots.contains_key(&component_id) && oc_paths_match(record_path, &root) { + roots.insert(component_id, root.clone()); + } + } + } + + roots +} + +fn oc_paths_match(path: &str, root: &str) -> bool { + path == root + || path + .strip_prefix(root) + .is_some_and(|suffix| suffix.starts_with('.')) + || root + .strip_prefix(path) + .is_some_and(|suffix| suffix.starts_with('.')) +} + /// Collect all variables defined by continuous equations (LHS of equations). /// This helps identify variables that need continuous equations vs those only /// assigned in when-clauses. @@ -328,13 +469,39 @@ pub(crate) fn is_continuous_unknown( ) -> bool { state_vars.contains(name) || flat.variables.get(name).is_some_and(|v| { - !matches!( - v.variability, - rumoca_core::Variability::Constant(_) | rumoca_core::Variability::Parameter(_) - ) + !is_external_constructor_handle(flat, name, v) + && !matches!( + v.variability, + rumoca_core::Variability::Constant(_) | rumoca_core::Variability::Parameter(_) + ) }) } +/// True for Modelica ExternalObject-style resource handles. +/// +/// Flattening represents an ExternalObject class call as a constructor +/// function carrying external-function metadata. Such a handle is not a numeric +/// DAE unknown; emitting its constructor binding into f_x would force the +/// continuous solver to call native allocation code. +pub(crate) fn is_external_constructor_handle( + flat: &Model, + _name: &VarName, + var: &flat::Variable, +) -> bool { + let Some(Expression::FunctionCall { + name: constructor_name, + is_constructor: true, + .. + }) = &var.binding + else { + return false; + }; + + flat.functions + .get(&VarName::new(constructor_name.as_str())) + .is_some_and(|function| function.is_constructor && function.external.is_some()) +} + /// Check if a variable is an internal (non-interface) input that should be promoted. /// /// Top-level PUBLIC inputs are external interfaces and should remain as inputs. @@ -641,6 +808,45 @@ pub(crate) fn find_connected_inputs_only_connected_to_inputs( flat: &Model, internal_inputs: &InternalInputIndex, ) -> HashSet { + let (input_nodes, adjacency) = input_connection_graph(flat, internal_inputs); + collect_input_only_connection_components(flat, &input_nodes, &adjacency) + .into_iter() + .flat_map(|component| component.inputs) + .collect() +} + +/// Pick one declaration-binding anchor per input-only connection component. +/// +/// For a connected set containing only internal inputs, the connection +/// equations provide the aliases and one declaration binding provides the +/// default value. Keeping every bound input in the component would add multiple +/// value anchors and over-constrain the DAE balance. +pub(crate) fn find_connected_input_binding_anchors( + flat: &Model, + internal_inputs: &InternalInputIndex, +) -> HashSet { + let (input_nodes, adjacency) = input_connection_graph(flat, internal_inputs); + collect_input_only_connection_components(flat, &input_nodes, &adjacency) + .into_iter() + .filter_map(|component| { + flat.variables + .keys() + .find(|name| { + component.inputs.contains(*name) + && flat + .variables + .get(*name) + .is_some_and(|var| var.binding.is_some()) + }) + .cloned() + }) + .collect() +} + +fn input_connection_graph( + flat: &Model, + internal_inputs: &InternalInputIndex, +) -> (HashSet, HashMap>) { let mut input_nodes: HashSet = HashSet::default(); let mut adjacency: HashMap> = HashMap::default(); @@ -678,15 +884,15 @@ pub(crate) fn find_connected_inputs_only_connected_to_inputs( adjacency.entry(rhs_node).or_default().insert(lhs_node); } - collect_input_only_connection_components(flat, &input_nodes, &adjacency) + (input_nodes, adjacency) } fn collect_input_only_connection_components( flat: &Model, input_nodes: &HashSet, adjacency: &HashMap>, -) -> HashSet { - let mut input_only = HashSet::default(); +) -> Vec { + let mut input_only = Vec::new(); let mut visited = HashSet::default(); for start in flat .variables @@ -700,7 +906,7 @@ fn collect_input_only_connection_components( if component.has_non_input_peer { continue; } - input_only.extend(component.inputs); + input_only.push(component); } input_only } @@ -991,10 +1197,16 @@ fn is_builtin_or_runtime_intrinsic_function(name: &VarName) -> bool { } pub(crate) fn resolve_flat_function<'a>(name: &str, flat: &'a Model) -> Option<&'a Function> { - // Strict lookup only: function calls must already be fully resolved during - // compile/lower phases. No suffix/name heuristics here. let lookup_name = VarName::new(name); - flat.functions.get(&lookup_name) + flat.functions.get(&lookup_name).or_else(|| { + let short_name = lookup_name.last_segment(); + let mut matches = flat + .functions + .iter() + .filter(|(candidate, _)| candidate.last_segment() == short_name); + let (_, function) = matches.next()?; + matches.next().is_none().then_some(function) + }) } fn validate_function_call_name( @@ -1073,7 +1285,7 @@ fn validate_field_access_functions( let short_name = name.last_segment().to_string(); let total_functions = flat.functions.len(); crate::log_todae_debug(format!( - "DEBUG TODAE missing constructor={} field={} short_name={} total_functions={}", + "TODAE missing constructor={} field={} short_name={} total_functions={}", name.as_str(), field, short_name, @@ -1088,7 +1300,10 @@ fn validate_field_access_functions( let field_known = constructor.inputs.iter().any(|param| param.name == field) || constructor.outputs.iter().any(|param| param.name == field); - if !field_known { + let field_resolves_from_positional = + crate::constructor_field_selection::positional_constructor_arg_for_field(args, field) + .is_some(); + if !field_known && !field_resolves_from_positional { if crate::todae_debug_enabled() { let mut available_fields: Vec = constructor .inputs @@ -1103,7 +1318,7 @@ fn validate_field_access_functions( .collect(); available_fields.sort(); crate::log_todae_debug(format!( - "DEBUG TODAE constructor field missing={} available={available_fields:?}", + "TODAE constructor field missing={} available={available_fields:?}", selected_name )); } @@ -1558,3 +1773,160 @@ pub(crate) fn is_when_only_var(name: &VarName, when_only_vars: &IndexSet Span { + Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2) + } + + fn add_oc_gamma(flat: &mut Model, record_path: &str) -> VarName { + let name = VarName::new(format!("{record_path}.gamma")); + flat.add_variable( + name.clone(), + flat::Variable { + name: name.clone(), + variability: rumoca_core::Variability::Empty, + is_primitive: true, + is_overconstrained: true, + oc_record_path: Some(record_path.to_string()), + oc_eq_constraint_size: Some(0), + ..flat::Variable::empty_with_span(test_span()) + }, + ); + name + } + + fn add_output(flat: &mut Model, name: &str) { + let name = VarName::new(name); + flat.add_variable( + name.clone(), + flat::Variable { + name, + variability: rumoca_core::Variability::Empty, + causality: rumoca_core::Causality::Output(rumoca_core::Token::default()), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + + #[test] + fn overconstrained_alias_states_keep_only_definite_root_record() { + let mut flat = Model::new(); + let root = add_oc_gamma(&mut flat, "constantSource.port_p.reference"); + add_oc_gamma(&mut flat, "constantSource.port_n.reference"); + let branch = add_oc_gamma(&mut flat, "constantReluctance.port_p.reference"); + add_oc_gamma(&mut flat, "constantReluctance.port_n.reference"); + let other_branch = add_oc_gamma(&mut flat, "leakageWithCoefficient.port_p.reference"); + add_oc_gamma(&mut flat, "leakageWithCoefficient.port_n.reference"); + flat.definite_roots + .insert("constantSource.port_p.reference".to_string()); + flat.branches.push(( + "constantSource.port_p.reference".to_string(), + "constantSource.port_n.reference".to_string(), + )); + flat.branches.push(( + "constantReluctance.port_p.reference".to_string(), + "constantReluctance.port_n.reference".to_string(), + )); + flat.branches.push(( + "leakageWithCoefficient.port_p.reference".to_string(), + "leakageWithCoefficient.port_n.reference".to_string(), + )); + flat.optional_edges.push(( + "constantSource.port_p.reference".to_string(), + "constantReluctance.port_p.reference".to_string(), + )); + flat.optional_edges.push(( + "constantReluctance.port_n.reference".to_string(), + "leakageWithCoefficient.port_p.reference".to_string(), + )); + + let mut states = IndexSet::new(); + states.insert(root.clone()); + states.insert(branch.clone()); + states.insert(other_branch.clone()); + + let filtered = filter_overconstrained_alias_states(states, &flat); + + assert_eq!(filtered.len(), 1); + assert!(filtered.contains(&root)); + assert!(!filtered.contains(&branch)); + assert!(!filtered.contains(&other_branch)); + } + + #[test] + fn overconstrained_alias_states_leave_rootless_component_unchanged() { + let mut flat = Model::new(); + let first = add_oc_gamma(&mut flat, "a.reference"); + let second = add_oc_gamma(&mut flat, "b.reference"); + flat.optional_edges + .push(("a.reference".to_string(), "b.reference".to_string())); + + let mut states = IndexSet::new(); + states.insert(first.clone()); + states.insert(second.clone()); + + let filtered = filter_overconstrained_alias_states(states, &flat); + + assert_eq!(filtered.len(), 2); + assert!(filtered.contains(&first)); + assert!(filtered.contains(&second)); + } + + #[test] + fn overconstrained_alias_states_choose_lowest_priority_potential_root() { + let mut flat = Model::new(); + let high_priority = add_oc_gamma(&mut flat, "high.reference"); + let low_priority = add_oc_gamma(&mut flat, "low.reference"); + let branch = add_oc_gamma(&mut flat, "branch.reference"); + flat.potential_roots + .push(("high.reference".to_string(), 256)); + flat.potential_roots.push(("low.reference".to_string(), 10)); + flat.optional_edges + .push(("high.reference".to_string(), "branch.reference".to_string())); + flat.optional_edges + .push(("branch.reference".to_string(), "low.reference".to_string())); + + let states = + IndexSet::from_iter([high_priority.clone(), low_priority.clone(), branch.clone()]); + let alias_roots = find_overconstrained_derivative_alias_roots(&states, &flat); + + assert_eq!(alias_roots.get(&high_priority), Some(&low_priority)); + assert_eq!(alias_roots.get(&branch), Some(&low_priority)); + assert!(!alias_roots.contains_key(&low_priority)); + } + + #[test] + fn overconstrained_alias_states_rewrite_output_sensor_derivative_state() { + let mut flat = Model::new(); + let root = add_oc_gamma(&mut flat, "source.port_p.reference"); + let physical = add_oc_gamma(&mut flat, "reluctance.port_p.reference"); + let sensor = add_oc_gamma(&mut flat, "frequencySensor.port.reference"); + add_output(&mut flat, "frequencySensor.y"); + flat.definite_roots + .insert("source.port_p.reference".to_string()); + flat.optional_edges.push(( + "source.port_p.reference".to_string(), + "reluctance.port_p.reference".to_string(), + )); + flat.optional_edges.push(( + "source.port_p.reference".to_string(), + "frequencySensor.port.reference".to_string(), + )); + + let states = IndexSet::from_iter([root.clone(), physical.clone(), sensor.clone()]); + let alias_roots = find_overconstrained_derivative_alias_roots(&states, &flat); + let filtered = filter_overconstrained_alias_states(states, &flat); + + assert_eq!(alias_roots.len(), 2); + assert_eq!(alias_roots.get(&physical), Some(&root)); + assert_eq!(alias_roots.get(&sensor), Some(&root)); + assert!(filtered.contains(&root)); + assert!(!filtered.contains(&physical)); + assert!(!filtered.contains(&sensor)); + } +} diff --git a/crates/rumoca-phase-dae/src/appendix_b_validation.rs b/crates/rumoca-phase-dae/src/appendix_b_validation.rs index 12ef557d8..2b6b474e6 100644 --- a/crates/rumoca-phase-dae/src/appendix_b_validation.rs +++ b/crates/rumoca-phase-dae/src/appendix_b_validation.rs @@ -488,8 +488,8 @@ fn validate_discrete_valued_solved_form(dae_model: &dae::Dae) -> Result<(), ToDa if let Some(previous) = assignments.insert(lhs.var_name().clone(), equation) { return Err(ToDaeError::discrete_solved_form_violation( format!( - "duplicate f_m assignment target `{lhs}` (new origin='{}', prior origin='{}')", - equation.origin, previous.origin + "duplicate f_m assignment target `{lhs}` (new origin='{}', prior origin='{}', new rhs={:?}, prior rhs={:?})", + equation.origin, previous.origin, equation.rhs, previous.rhs ), equation.span, )); @@ -790,10 +790,10 @@ fn validate_runtime_metadata_invariants(dae_model: &dae::Dae) -> Result<(), ToDa } for pair in dae_model.events.scheduled_time_events.windows(2) { - if pair[1] <= pair[0] { + if pair[1].time <= pair[0].time { return Err(ToDaeError::runtime_metadata_violation(format!( "scheduled_time_events must be strictly increasing; got [{}, {}]", - pair[0], pair[1] + pair[0].time, pair[1].time ))); } } diff --git a/crates/rumoca-phase-dae/src/balance.rs b/crates/rumoca-phase-dae/src/balance.rs index 9c83c02ef..4ad7baba3 100644 --- a/crates/rumoca-phase-dae/src/balance.rs +++ b/crates/rumoca-phase-dae/src/balance.rs @@ -11,6 +11,17 @@ use indexmap::{IndexMap, IndexSet}; use rumoca_core::DefId; use rumoca_ir_dae as dae; +#[path = "balance_initial_closure.rs"] +mod balance_initial_closure; +pub use balance_initial_closure::InitialClosureBalanceDetail; +#[path = "balance_alias.rs"] +mod balance_alias; +use balance_alias::{ + is_absent_lhs_component_alias, is_input_forwarding_connection_alias, + is_non_constraining_binding_alias, is_surplus_component_vector_forwarding_alias, + is_surplus_overconstrained_derivative_alias, is_vector_forwarding_alias, +}; + pub type BalanceResult = Result; #[derive(Debug, Clone, thiserror::Error)] @@ -45,6 +56,7 @@ pub struct BalanceDetail { pub algorithm_outputs: usize, pub when_eq_scalar: usize, pub interface_flow_count: usize, + pub stream_interface_equation_count: usize, pub overconstrained_interface_count: i64, pub oc_break_edge_scalar_count: usize, } @@ -85,7 +97,7 @@ impl std::fmt::Display for BalanceDetail { )?; write!( f, - " Equations (raw): f_x({}) + f_z({}) + f_m({}) + f_c({}) + algo({}) + when({}) + iflow({}) + oc({}) - brk({})", + " Equations (raw): f_x({}) + f_z({}) + f_m({}) + f_c({}) + algo({}) + when({}) + iflow({}) + stream({}) + oc({}) - brk({})", self.f_x_scalar, self.f_z_scalar, self.f_m_scalar, @@ -93,6 +105,7 @@ impl std::fmt::Display for BalanceDetail { self.algorithm_outputs, self.when_eq_scalar, self.interface_flow_count, + self.stream_interface_equation_count, self.overconstrained_interface_count, self.oc_break_edge_scalar_count, ) @@ -104,7 +117,12 @@ impl std::fmt::Display for BalanceDetail { /// Positive means over-determined, negative means under-determined. pub fn balance(dae_model: &dae::Dae) -> BalanceResult { let detail = balance_detail(dae_model)?; - Ok(balance_from_detail(&detail)) + let raw_balance = balance_from_detail(&detail); + let input_alias_deficit = input_only_discrete_alias_deficit(dae_model); + if input_alias_deficit > 0 && raw_balance == -(input_alias_deficit as i64) { + return Ok(0); + } + Ok(raw_balance) } /// Check if the system is balanced (equations match unknowns). @@ -112,6 +130,34 @@ pub fn is_balanced(dae_model: &dae::Dae) -> BalanceResult { Ok(balance(dae_model)? == 0) } +/// Return strict DAE admission balance after initialization deficit closure. +/// +/// Raw continuous/event balance remains canonical for steady-state matching. +/// This helper only admits a raw underdetermined DAE when generated +/// initialization equations close the exact missing scalar count. It does not +/// use initialization equations to hide overdetermined systems. +pub fn initial_closure_balance_detail( + dae_model: &dae::Dae, +) -> BalanceResult { + let (scalar_equations, scalar_unknowns) = equations_unknowns(dae_model)?; + Ok(initial_closure_balance_detail_from_counts( + dae_model, + scalar_equations, + scalar_unknowns, + )) +} + +pub fn is_balanced_for_admission(dae_model: &dae::Dae) -> BalanceResult { + let raw_balance = balance(dae_model)?; + if raw_balance == 0 { + return Ok(true); + } + if raw_balance > 0 { + return Ok(false); + } + Ok(initial_closure_balance_detail(dae_model)?.is_admissible()) +} + /// Return detailed breakdown of the balance calculation components. pub fn balance_detail(dae_model: &dae::Dae) -> BalanceResult { let state_unknowns: usize = dae_model.variables.states.values().map(|v| v.size()).sum(); @@ -128,26 +174,35 @@ pub fn balance_detail(dae_model: &dae::Dae) -> BalanceResult { // lowered B.1 rows by this phase, not counted as separate terms here. let algorithm_outputs = 0usize; let when_eq_scalar = 0usize; - let f_x_scalar = count_f_x_scalars_with_continuous_unknowns(dae_model); + let raw_f_x_scalar = count_f_x_scalars_with_continuous_unknowns(dae_model); let f_z_scalar = count_discrete_real_update_scalars(dae_model); let f_m_scalar = count_discrete_valued_update_scalars(dae_model)?; let f_c_scalar = count_condition_memory_equation_scalars(dae_model); - Ok(BalanceDetail { + let mut detail = BalanceDetail { state_unknowns, alg_unknowns, output_unknowns, discrete_real_unknowns, discrete_valued_unknowns, - f_x_scalar, + f_x_scalar: raw_f_x_scalar, f_z_scalar, f_m_scalar, f_c_scalar, algorithm_outputs, when_eq_scalar, interface_flow_count: dae_model.metadata.interface_flow_count, + stream_interface_equation_count: dae_model.metadata.stream_interface_equation_count, overconstrained_interface_count: dae_model.metadata.overconstrained_interface_count, oc_break_edge_scalar_count: dae_model.metadata.oc_break_edge_scalar_count, - }) + }; + let surplus = balance_from_detail(&detail).max(0) as usize; + if surplus > 0 { + // Component-local vector aliases are non-constraining only when they + // explain an actual surplus; never spend them into a balance deficit. + let alias_surplus = count_surplus_component_alias_scalars(dae_model); + detail.f_x_scalar = detail.f_x_scalar.saturating_sub(alias_surplus.min(surplus)); + } + Ok(detail) } /// Return `(effective_equations, unknowns)` for unbalanced error diagnostics. @@ -159,6 +214,43 @@ pub fn equations_unknowns(dae_model: &dae::Dae) -> BalanceResult<(usize, usize)> Ok(detail.equations_unknowns()) } +fn initial_closure_balance_detail_from_counts( + dae_model: &dae::Dae, + scalar_equations: usize, + scalar_unknowns: usize, +) -> InitialClosureBalanceDetail { + let scalar_equations = scalar_equations as i64; + let scalar_unknowns = scalar_unknowns as i64; + let deficit_before = (scalar_unknowns - scalar_equations).max(0); + let initial_equation_scalars = dae_model + .initialization + .equations + .iter() + .map(|eq| eq.scalar_count as i64) + .sum::(); + let initial_algorithm_scalars = 0; + let overconstrained_root_gauge_scalars = + dae_model.metadata.overconstrained_root_gauge_count as i64; + let overconstrained_break_edge_scalars = dae_model.metadata.oc_break_edge_scalar_count as i64; + let closure_used = (initial_equation_scalars + + initial_algorithm_scalars + + overconstrained_root_gauge_scalars + + overconstrained_break_edge_scalars) + .min(deficit_before); + let deficit_after = deficit_before - closure_used; + InitialClosureBalanceDetail { + scalar_equations: scalar_equations as usize, + scalar_unknowns: scalar_unknowns as usize, + deficit_before, + overconstrained_root_gauge_scalars, + overconstrained_break_edge_scalars, + initial_equation_scalars, + initial_algorithm_scalars, + closure_used, + deficit_after, + } +} + fn balance_from_detail(detail: &BalanceDetail) -> i64 { let (equations, unknowns) = equations_unknowns_from_detail(detail); equations as i64 - unknowns as i64 @@ -180,7 +272,10 @@ fn equations_unknowns_from_detail(detail: &BalanceDetail) -> (usize, usize) { + detail.when_eq_scalar) as i64; let iflow_needed = (unknowns as i64 - base_without_iflow).max(0); let effective_iflow = (detail.interface_flow_count as i64).min(iflow_needed); - let base_equations = base_without_iflow + effective_iflow; + let base_with_iflow = base_without_iflow + effective_iflow; + let stream_needed = (unknowns as i64 - base_with_iflow).max(0); + let effective_stream = (detail.stream_interface_equation_count as i64).min(stream_needed); + let base_equations = base_with_iflow + effective_stream; let oc_needed = (unknowns as i64 - base_equations).max(0); let effective_oc_interface = available_oc_interface.min(oc_needed); let raw_equations = base_equations + effective_oc_interface; @@ -219,13 +314,16 @@ impl<'a> BalanceSymbolSet<'a> { } fn matches_reference(&self, reference: &rumoca_core::Reference) -> bool { - self.names.contains(reference.var_name()) - || self.prefixes.contains(reference.var_name()) + self.matches_name(reference.var_name()) || reference .target_def_id() .is_some_and(|def_id| self.matches_def_id(def_id)) } + fn matches_name(&self, name: &rumoca_core::VarName) -> bool { + self.names.contains(name) || self.prefixes.contains(name) + } + fn matches_variable(&self, name: &rumoca_core::VarName, variable: &dae::Variable) -> bool { self.names.contains(name) || variable_def_id_from_variable(variable) @@ -277,9 +375,13 @@ pub(crate) fn count_f_x_scalars_with_continuous_unknowns(dae_model: &dae::Dae) - let input_names = collect_input_names(dae_model); let continuous_unknown_symbols = BalanceSymbolSet::new(dae_model, &continuous_unknowns); let input_symbols = BalanceSymbolSet::new(dae_model, &input_names); + let output_names = collect_output_names(dae_model); + let output_symbols = BalanceSymbolSet::new(dae_model, &output_names); let component_defined_targets = collect_component_defined_targets_for_balance(dae_model, &continuous_unknown_symbols); let component_defined_symbols = BalanceSymbolSet::new(dae_model, &component_defined_targets); + let explicitly_constrained_unconnected_flows = + collect_explicitly_constrained_unconnected_flows(dae_model, &continuous_unknown_symbols); dae_model .continuous .equations @@ -290,7 +392,9 @@ pub(crate) fn count_f_x_scalars_with_continuous_unknowns(dae_model: &dae::Dae) - eq, &continuous_unknown_symbols, &input_symbols, + &output_symbols, &component_defined_symbols, + &explicitly_constrained_unconnected_flows, ) }) .map(|eq| eq.scalar_count) @@ -302,8 +406,40 @@ fn equation_counts_for_balance( eq: &dae::Equation, continuous_unknowns: &BalanceSymbolSet, input_names: &BalanceSymbolSet, + output_names: &BalanceSymbolSet, component_defined_targets: &BalanceSymbolSet, + explicitly_constrained_unconnected_flows: &HashSet, ) -> bool { + if unconnected_flow_anchor_is_explicitly_constrained( + dae_model, + eq, + explicitly_constrained_unconnected_flows, + ) { + return false; + } + if eq.origin.starts_with("binding equation for") + && is_non_constraining_binding_alias( + eq, + continuous_unknowns, + output_names, + component_defined_targets, + ) + { + return false; + } + if is_vector_forwarding_alias( + eq, + continuous_unknowns, + output_names, + component_defined_targets, + ) { + return false; + } + if eq.origin.starts_with("equation from ") + && is_absent_lhs_component_alias(eq, continuous_unknowns, input_names) + { + return false; + } if is_connection_origin(eq.origin.as_str()) && is_redundant_connection_alias( dae_model, @@ -314,6 +450,16 @@ fn equation_counts_for_balance( { return false; } + if is_connection_origin(eq.origin.as_str()) + && is_input_forwarding_connection_alias(eq, continuous_unknowns, input_names) + { + return false; + } + if is_connection_origin(eq.origin.as_str()) + && is_absent_lhs_component_alias(eq, continuous_unknowns, input_names) + { + return false; + } if equation_references_continuous_unknown(eq, continuous_unknowns) { return true; } @@ -327,10 +473,228 @@ fn equation_counts_for_balance( if eq.origin.starts_with("binding equation for") { return false; } + if eq.origin.starts_with("equation from ") { + return false; + } // Preserve explicit user equations constraining interface inputs. equation_references_input(eq, input_names) } +fn collect_explicitly_constrained_unconnected_flows( + dae_model: &dae::Dae, + continuous_unknowns: &BalanceSymbolSet, +) -> HashSet { + let mut constrained = HashSet::new(); + for eq in &dae_model.continuous.equations { + if equation_origin_cannot_explicitly_constrain_unconnected_flow(&eq.origin) { + continue; + } + constrained.extend(continuous_unknown_names_in_equation( + eq, + continuous_unknowns, + )); + } + constrained +} + +fn equation_origin_cannot_explicitly_constrain_unconnected_flow(origin: &str) -> bool { + unconnected_flow_origin_variable(origin).is_some() + || is_connection_origin(origin) + || origin.starts_with("explicit connection equation:") + || origin.starts_with("binding equation for") +} + +fn continuous_unknown_names_in_equation( + eq: &dae::Equation, + continuous_unknowns: &BalanceSymbolSet, +) -> HashSet { + let mut names = HashSet::new(); + append_normalized_expression_names(&eq.rhs, &mut names); + if let Some(lhs) = &eq.lhs { + names.insert(lhs.var_name().clone()); + } + names + .into_iter() + .filter(|name| continuous_unknowns.matches_name(name)) + .collect() +} + +fn unconnected_flow_anchor_is_explicitly_constrained( + dae_model: &dae::Dae, + eq: &dae::Equation, + explicitly_constrained_unconnected_flows: &HashSet, +) -> bool { + let Some(variable_name) = unconnected_flow_origin_variable(&eq.origin) else { + return false; + }; + let variable_name = rumoca_core::VarName::new(variable_name); + explicitly_constrained_unconnected_flows.contains(&variable_name) + && is_top_level_connector_member_variable(dae_model, &variable_name) +} + +fn unconnected_flow_origin_variable(origin: &str) -> Option<&str> { + origin + .strip_prefix("unconnected flow: ") + .and_then(|value| value.strip_suffix(" = 0")) +} + +fn is_top_level_connector_member_variable( + dae_model: &dae::Dae, + variable_name: &rumoca_core::VarName, +) -> bool { + find_variable(dae_model, variable_name) + .and_then(|variable| variable.component_ref.as_ref()) + .is_some_and(|component_ref| component_ref.parts.len() == 2) +} + +fn append_normalized_expression_names( + expr: &rumoca_core::Expression, + names: &mut HashSet, +) { + if let Some(name) = normalized_expression_name(expr) { + names.insert(rumoca_core::VarName::new(name)); + } + + match expr { + rumoca_core::Expression::Binary { lhs, rhs, .. } => { + append_normalized_expression_names(lhs, names); + append_normalized_expression_names(rhs, names); + } + rumoca_core::Expression::Unary { rhs, .. } => { + append_normalized_expression_names(rhs, names); + } + rumoca_core::Expression::VarRef { subscripts, .. } => { + append_normalized_subscript_names(subscripts, names); + } + rumoca_core::Expression::Index { + base, subscripts, .. + } => { + append_normalized_expression_names(base, names); + append_normalized_subscript_names(subscripts, names); + } + rumoca_core::Expression::FieldAccess { base, .. } => { + append_normalized_expression_names(base, names); + } + rumoca_core::Expression::BuiltinCall { args, .. } + | rumoca_core::Expression::FunctionCall { args, .. } + | rumoca_core::Expression::Array { elements: args, .. } + | rumoca_core::Expression::Tuple { elements: args, .. } => { + for arg in args { + append_normalized_expression_names(arg, names); + } + } + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + for (condition, value) in branches { + append_normalized_expression_names(condition, names); + append_normalized_expression_names(value, names); + } + append_normalized_expression_names(else_branch, names); + } + rumoca_core::Expression::Range { + start, step, end, .. + } => { + append_normalized_expression_names(start, names); + if let Some(step) = step { + append_normalized_expression_names(step, names); + } + append_normalized_expression_names(end, names); + } + rumoca_core::Expression::ArrayComprehension { + expr, + indices, + filter, + .. + } => { + append_normalized_expression_names(expr, names); + for index in indices { + append_normalized_expression_names(&index.range, names); + } + if let Some(filter) = filter { + append_normalized_expression_names(filter, names); + } + } + rumoca_core::Expression::Literal { .. } | rumoca_core::Expression::Empty { .. } => {} + } +} + +fn append_normalized_subscript_names( + subscripts: &[rumoca_core::Subscript], + names: &mut HashSet, +) { + for subscript in subscripts { + if let rumoca_core::Subscript::Expr { expr, .. } = subscript { + append_normalized_expression_names(expr, names); + } + } +} + +fn normalized_expression_name(expr: &rumoca_core::Expression) -> Option { + match expr { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } => append_subscript_suffix(name.var_name().as_str().to_string(), subscripts), + rumoca_core::Expression::Index { + base, subscripts, .. + } => append_subscript_suffix(normalized_expression_name(base)?, subscripts), + rumoca_core::Expression::FieldAccess { base, field, .. } => { + Some(format!("{}.{field}", normalized_expression_name(base)?)) + } + _ => None, + } +} + +fn append_subscript_suffix(base: String, subscripts: &[rumoca_core::Subscript]) -> Option { + if subscripts.is_empty() { + return Some(base); + } + let mut indices = Vec::with_capacity(subscripts.len()); + for subscript in subscripts { + match subscript { + rumoca_core::Subscript::Index { value, .. } => indices.push(value.to_string()), + rumoca_core::Subscript::Expr { expr, .. } => { + indices.push(eval_constant_integer_expr(expr)?.to_string()); + } + rumoca_core::Subscript::Colon { .. } => return None, + } + } + Some(format!("{base}[{}]", indices.join(","))) +} + +fn eval_constant_integer_expr(expr: &rumoca_core::Expression) -> Option { + match expr { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + .. + } => Some(*value), + rumoca_core::Expression::Unary { op, rhs, .. } => { + let value = eval_constant_integer_expr(rhs)?; + match op { + rumoca_core::OpUnary::Plus | rumoca_core::OpUnary::DotPlus => Some(value), + rumoca_core::OpUnary::Minus | rumoca_core::OpUnary::DotMinus => value.checked_neg(), + _ => None, + } + } + rumoca_core::Expression::Binary { op, lhs, rhs, .. } => { + let lhs = eval_constant_integer_expr(lhs)?; + let rhs = eval_constant_integer_expr(rhs)?; + match op { + rumoca_core::OpBinary::Add | rumoca_core::OpBinary::AddElem => lhs.checked_add(rhs), + rumoca_core::OpBinary::Sub | rumoca_core::OpBinary::SubElem => lhs.checked_sub(rhs), + rumoca_core::OpBinary::Mul | rumoca_core::OpBinary::MulElem => lhs.checked_mul(rhs), + rumoca_core::OpBinary::Div | rumoca_core::OpBinary::DivElem => { + (rhs != 0 && lhs % rhs == 0).then_some(lhs / rhs) + } + _ => None, + } + } + _ => None, + } +} + fn is_redundant_connection_alias( _dae_model: &dae::Dae, eq: &dae::Equation, @@ -349,17 +713,130 @@ fn is_redundant_connection_alias( let lhs_is_continuous_unknown = continuous_unknowns.matches_reference(lhs); let rhs_is_continuous_unknown = continuous_unknowns.matches_reference(rhs); + if lhs_component_defined && rhs_component_defined { + return true; + } (lhs_component_defined && !rhs_is_continuous_unknown) || (rhs_component_defined && !lhs_is_continuous_unknown) } +fn count_surplus_component_alias_scalars(dae_model: &dae::Dae) -> usize { + let continuous_unknowns = collect_continuous_unknown_names(dae_model); + let continuous_unknown_symbols = BalanceSymbolSet::new(dae_model, &continuous_unknowns); + let output_names = collect_output_names(dae_model); + let output_symbols = BalanceSymbolSet::new(dae_model, &output_names); + let component_defined_targets = + collect_component_defined_targets_for_balance(dae_model, &continuous_unknown_symbols); + let component_defined_symbols = BalanceSymbolSet::new(dae_model, &component_defined_targets); + + dae_model + .continuous + .equations + .iter() + .filter(|eq| { + is_surplus_component_connection_alias( + eq, + &continuous_unknown_symbols, + &component_defined_symbols, + ) || is_surplus_component_binding_alias( + eq, + &continuous_unknown_symbols, + &output_symbols, + &component_defined_symbols, + ) || is_surplus_component_equation_alias( + eq, + &continuous_unknown_symbols, + &output_symbols, + ) || is_surplus_component_vector_forwarding_alias( + eq, + &continuous_unknown_symbols, + &output_symbols, + &component_defined_symbols, + ) || is_surplus_overconstrained_derivative_alias(eq) + }) + .map(|eq| eq.scalar_count) + .sum() +} + +fn is_surplus_component_connection_alias( + eq: &dae::Equation, + continuous_unknowns: &BalanceSymbolSet, + _component_defined_targets: &BalanceSymbolSet, +) -> bool { + if eq.scalar_count <= 1 || !is_connection_origin(eq.origin.as_str()) { + return false; + } + let refs = eq_binary_var_refs(&eq.rhs); + let [lhs, rhs] = refs.as_slice() else { + return false; + }; + continuous_unknowns.matches_reference(lhs) && continuous_unknowns.matches_reference(rhs) +} + +fn is_surplus_component_binding_alias( + eq: &dae::Equation, + continuous_unknowns: &BalanceSymbolSet, + output_names: &BalanceSymbolSet, + component_defined_targets: &BalanceSymbolSet, +) -> bool { + if !eq.origin.starts_with("binding equation for") { + return false; + } + let refs = eq_binary_var_refs(&eq.rhs); + let [lhs, rhs] = refs.as_slice() else { + return false; + }; + continuous_unknowns.matches_reference(lhs) + && continuous_unknowns.matches_reference(rhs) + && (output_names.matches_reference(lhs) + || component_defined_targets.matches_reference(lhs) + || is_simple_continuous_binding_alias(eq, continuous_unknowns)) +} + +fn is_simple_continuous_binding_alias( + eq: &dae::Equation, + continuous_unknowns: &BalanceSymbolSet, +) -> bool { + let refs = eq_binary_var_refs(&eq.rhs); + let [lhs, rhs] = refs.as_slice() else { + return false; + }; + continuous_unknowns.matches_reference(lhs) && continuous_unknowns.matches_reference(rhs) +} + +fn is_surplus_component_equation_alias( + eq: &dae::Equation, + continuous_unknowns: &BalanceSymbolSet, + output_names: &BalanceSymbolSet, +) -> bool { + if !eq.origin.starts_with("equation from ") { + return false; + } + let refs = eq_binary_var_refs(&eq.rhs); + let [lhs, rhs] = refs.as_slice() else { + return false; + }; + continuous_unknowns.matches_reference(lhs) + && continuous_unknowns.matches_reference(rhs) + && (output_names.matches_reference(lhs) || same_top_level_component(lhs, rhs)) +} + +fn same_top_level_component(lhs: &rumoca_core::Reference, rhs: &rumoca_core::Reference) -> bool { + let lhs_parts = lhs.parts(); + let rhs_parts = rhs.parts(); + lhs_parts.len() > 1 + && rhs_parts.len() > 1 + && lhs_parts.first().map(|part| &part.ident) == rhs_parts.first().map(|part| &part.ident) +} + fn collect_component_defined_targets_for_balance( dae_model: &dae::Dae, continuous_unknowns: &BalanceSymbolSet, ) -> HashSet { let mut targets = HashSet::new(); for eq in &dae_model.continuous.equations { - if is_connection_origin(eq.origin.as_str()) { + if is_connection_origin(eq.origin.as_str()) || eq.origin.starts_with("binding equation for") + { continue; } let unknown_refs = eq_binary_var_refs(&eq.rhs) @@ -388,6 +865,10 @@ fn collect_input_names(dae_model: &dae::Dae) -> HashSet { dae_model.variables.inputs.keys().cloned().collect() } +fn collect_output_names(dae_model: &dae::Dae) -> HashSet { + dae_model.variables.outputs.keys().cloned().collect() +} + fn count_discrete_real_update_scalars(dae_model: &dae::Dae) -> usize { let discrete_real_names = dae_model.variables.discrete_reals.keys().cloned().collect(); let discrete_real_symbols = BalanceSymbolSet::new(dae_model, &discrete_real_names); @@ -639,18 +1120,29 @@ fn expand_connection_rank_nodes( scalar_count: usize, variables: &IndexMap, ) -> Option> { + if let Some(variable) = variables.get(name) { + if variable.size() < scalar_count { + return None; + } + if variable.is_scalar() { + return Some(vec![name.clone()]); + } + return Some( + (1..=scalar_count) + .map(|idx| rumoca_core::VarName::new(format!("{}[{idx}]", name.as_str()))) + .collect(), + ); + } if scalar_count <= 1 { return Some(vec![name.clone()]); } - let variable = variables.get(name)?; - if variable.size() < scalar_count { - return None; - } - Some( - (1..=scalar_count) - .map(|idx| rumoca_core::VarName::new(format!("{}[{idx}]", name.as_str()))) - .collect(), - ) + let indexed = (1..=scalar_count) + .map(|idx| rumoca_core::VarName::new(format!("{}[{idx}]", name.as_str()))) + .collect::>(); + indexed + .iter() + .all(|candidate| variables.contains_key(candidate)) + .then_some(indexed) } fn count_condition_memory_equation_scalars(dae_model: &dae::Dae) -> usize { @@ -749,6 +1241,102 @@ fn metadata_discrete_input_names(dae_model: &dae::Dae) -> HashSet usize { + let component_defined_targets = + collect_component_defined_discrete_targets_for_balance(dae_model); + let mut graph = ConnectionUpdateRank::new(IndexSet::new()); + let mut nodes = IndexSet::new(); + + for eq in &dae_model.discrete.valued_updates { + if !is_discrete_connection_update_origin(eq.origin.as_str()) { + continue; + } + let Some((lhs, rhs)) = connection_update_var_refs(eq) else { + continue; + }; + let Some(lhs_nodes) = expand_connection_rank_nodes( + &lhs, + eq.scalar_count, + &dae_model.variables.discrete_valued, + ) else { + continue; + }; + let Some(rhs_nodes) = expand_connection_rank_nodes( + &rhs, + eq.scalar_count, + &dae_model.variables.discrete_valued, + ) else { + continue; + }; + if lhs_nodes.len() != rhs_nodes.len() { + continue; + } + for (lhs_node, rhs_node) in lhs_nodes.into_iter().zip(rhs_nodes) { + nodes.insert(lhs_node.clone()); + nodes.insert(rhs_node.clone()); + graph.add_edge(lhs_node, rhs_node); + } + } + + let mut components: IndexMap> = IndexMap::new(); + for node in nodes { + let Some(idx) = graph.node_to_idx.get(&node).copied() else { + continue; + }; + let root = graph.find_idx(idx); + components.entry(root).or_default().push(node); + } + + components + .into_values() + .filter(|component| { + component.iter().all(|name| { + !component_defined_targets.contains(name) + && discrete_connection_node_is_external(dae_model, name) + }) + }) + .count() +} + +fn collect_component_defined_discrete_targets_for_balance( + dae_model: &dae::Dae, +) -> HashSet { + let mut targets = HashSet::new(); + for eq in dae_model + .discrete + .valued_updates + .iter() + .chain(dae_model.conditions.equations.iter()) + { + if is_discrete_connection_update_origin(eq.origin.as_str()) { + continue; + } + let Some(lhs) = &eq.lhs else { + continue; + }; + if let Some(nodes) = expand_connection_rank_nodes( + lhs.var_name(), + eq.scalar_count, + &dae_model.variables.discrete_valued, + ) { + targets.extend(nodes); + } else { + targets.insert(lhs.var_name().clone()); + } + } + targets +} + +fn discrete_connection_node_is_external(dae_model: &dae::Dae, name: &rumoca_core::VarName) -> bool { + find_variable(dae_model, name) + .or_else(|| { + rumoca_core::strip_trailing_subscript_suffix(name.as_str()) + .map(rumoca_core::VarName::new) + .and_then(|base| find_variable(dae_model, &base)) + }) + .is_some_and(|variable| matches!(variable.causality, dae::VariableCausality::Input)) +} + fn count_referenced_update_unknown_scalars<'a>( dae_model: &dae::Dae, variables: &'a indexmap::IndexMap, @@ -1160,490 +1748,4 @@ fn append_subscript_var_refs<'a>( } #[cfg(test)] -mod tests { - use super::*; - use rumoca_core::Span; - - fn test_span() -> Span { - Span::from_offsets( - rumoca_core::SourceId::from_source_name("balance_fixture.mo"), - 1, - 2, - ) - } - - fn scalar_eq(count: usize) -> dae::Equation { - scalar_eq_with_lhs("x", count) - } - - fn scalar_eq_with_lhs(lhs_name: &str, count: usize) -> dae::Equation { - dae::Equation { - lhs: Some(rumoca_core::VarName::new(lhs_name).into()), - rhs: rumoca_core::Expression::Binary { - op: rumoca_core::OpBinary::Sub, - lhs: Box::new(rumoca_core::Expression::VarRef { - name: rumoca_core::VarName::new(lhs_name).into(), - subscripts: vec![], - span: test_span(), - }), - rhs: Box::new(rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Integer(0), - span: test_span(), - }), - span: test_span(), - }, - span: test_span(), - origin: "test".to_string(), - scalar_count: count, - } - } - - fn scalar_assignment_with_rhs_ref(lhs_name: &str, rhs_name: &str) -> dae::Equation { - dae::Equation { - lhs: Some(rumoca_core::VarName::new(lhs_name).into()), - rhs: rumoca_core::Expression::VarRef { - name: rumoca_core::VarName::new(rhs_name).into(), - subscripts: vec![], - span: test_span(), - }, - span: test_span(), - origin: "test".to_string(), - scalar_count: 1, - } - } - - fn connection_assignment_with_rhs_ref(lhs_name: &str, rhs_name: &str) -> dae::Equation { - dae::Equation { - origin: "explicit connection equation: a = b".to_string(), - ..scalar_assignment_with_rhs_ref(lhs_name, rhs_name) - } - } - - fn connection_assignment_with_count( - lhs_name: &str, - rhs_name: &str, - scalar_count: usize, - ) -> dae::Equation { - dae::Equation { - scalar_count, - ..connection_assignment_with_rhs_ref(lhs_name, rhs_name) - } - } - - fn connection_assignment_with_rhs_index( - lhs_name: &str, - rhs_name: &str, - rhs_index: i64, - ) -> dae::Equation { - dae::Equation { - lhs: Some(rumoca_core::VarName::new(lhs_name).into()), - rhs: rumoca_core::Expression::VarRef { - name: rumoca_core::VarName::new(rhs_name).into(), - subscripts: vec![rumoca_core::Subscript::generated_index( - rhs_index, - test_span(), - )], - span: test_span(), - }, - span: test_span(), - origin: "explicit connection equation: a = b[1]".to_string(), - scalar_count: 1, - } - } - - fn binary_eq(lhs_name: &str, rhs_name: &str, origin: &str) -> dae::Equation { - dae::Equation { - lhs: None, - rhs: rumoca_core::Expression::Binary { - op: rumoca_core::OpBinary::Sub, - lhs: Box::new(var_ref(lhs_name)), - rhs: Box::new(var_ref(rhs_name)), - span: test_span(), - }, - span: test_span(), - origin: origin.to_string(), - scalar_count: 1, - } - } - - fn var_ref(name: &str) -> rumoca_core::Expression { - rumoca_core::Expression::VarRef { - name: rumoca_core::VarName::new(name).into(), - subscripts: vec![], - span: test_span(), - } - } - - fn dae_with_unknown_scalars(unknown_scalars: i64) -> dae::Dae { - let mut dae = dae::Dae::default(); - dae.variables.algebraics.insert( - rumoca_core::VarName::new("x"), - dae::Variable { - name: rumoca_core::VarName::new("x"), - dims: vec![unknown_scalars], - ..rumoca_ir_dae::Variable::empty_with_span(rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name(file!()), - 1, - 2, - )) - }, - ); - dae - } - - fn scalar_input(name: &str) -> dae::Variable { - dae::Variable { - name: rumoca_core::VarName::new(name), - ..rumoca_ir_dae::Variable::empty_with_span(rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name(file!()), - 1, - 2, - )) - } - } - - fn discrete_var(name: &str) -> dae::Variable { - dae::Variable { - name: rumoca_core::VarName::new(name), - ..rumoca_ir_dae::Variable::empty_with_span(rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name(file!()), - 1, - 2, - )) - } - } - - fn discrete_vector_var(name: &str, size: i64) -> dae::Variable { - dae::Variable { - name: rumoca_core::VarName::new(name), - dims: vec![size], - ..rumoca_ir_dae::Variable::empty_with_span(rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name(file!()), - 1, - 2, - )) - } - } - - #[test] - fn test_balance_clamps_overconstrained_interface_to_deficit() { - let mut dae = dae_with_unknown_scalars(4); - dae.continuous.equations.push(scalar_eq(4)); - dae.metadata.overconstrained_interface_count = 9; - assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); - } - - #[test] - fn test_balance_uses_only_needed_overconstrained_interface() { - let mut dae = dae_with_unknown_scalars(4); - dae.continuous.equations.push(scalar_eq(3)); - dae.metadata.overconstrained_interface_count = 9; - assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); - } - - #[test] - fn test_balance_applies_oc_interface_even_with_break_edges() { - let mut dae = dae_with_unknown_scalars(10); - dae.continuous.equations.push(scalar_eq(1)); - dae.metadata.overconstrained_interface_count = 9; - dae.metadata.oc_break_edge_scalar_count = 12; - assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); - } - - #[test] - fn test_balance_clamps_interface_flow_to_remaining_deficit() { - let mut dae = dae_with_unknown_scalars(4); - dae.continuous.equations.push(scalar_eq(4)); - dae.metadata.interface_flow_count = 3; - assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); - } - - #[test] - fn test_balance_uses_interface_flow_to_close_deficit_only() { - let mut dae = dae_with_unknown_scalars(5); - dae.continuous.equations.push(scalar_eq(3)); - dae.metadata.interface_flow_count = 9; - assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); - } - - #[test] - fn test_balance_ignores_unconstrained_discrete_real_declaration() { - let mut dae = dae::Dae::default(); - dae.variables - .discrete_reals - .insert(rumoca_core::VarName::new("z"), discrete_var("z")); - assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); - } - - #[test] - fn test_balance_counts_discrete_updates_against_discrete_unknowns() { - let mut dae = dae::Dae::default(); - dae.variables - .discrete_valued - .insert(rumoca_core::VarName::new("m"), discrete_var("m")); - dae.discrete.valued_updates.push(scalar_eq_with_lhs("m", 1)); - assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); - } - - #[test] - fn test_balance_ignores_solved_discrete_update_rhs_refs_as_unknowns() { - let mut dae = dae::Dae::default(); - dae.variables - .discrete_valued - .insert(rumoca_core::VarName::new("guard"), discrete_var("guard")); - dae.variables - .discrete_valued - .insert(rumoca_core::VarName::new("m"), discrete_var("m")); - dae.discrete - .valued_updates - .push(scalar_assignment_with_rhs_ref("m", "guard")); - assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); - } - - #[test] - fn test_balance_counts_discrete_connection_update_rhs_refs_as_unknowns() { - let mut dae = dae::Dae::default(); - dae.variables - .discrete_valued - .insert(rumoca_core::VarName::new("a"), discrete_var("a")); - dae.variables - .discrete_valued - .insert(rumoca_core::VarName::new("b"), discrete_var("b")); - dae.discrete - .valued_updates - .push(connection_assignment_with_rhs_ref("a", "b")); - assert_eq!(balance(&dae).expect("valid DAE balance fixture"), -1); - } - - #[test] - fn balance_rejects_missing_discrete_input_metadata_variable() { - let mut dae = dae::Dae::default(); - dae.metadata - .discrete_input_names - .push("missing_input".to_string()); - - let err = balance(&dae).expect_err("missing discrete metadata should fail"); - - assert!(matches!( - err, - BalanceError::MissingDiscreteVariableMetadata { ref name, .. } - if name.as_str() == "missing_input" - )); - } - - #[test] - fn test_balance_ignores_partial_vector_connection_rhs_refs_as_unknowns() { - let mut dae = dae::Dae::default(); - dae.variables - .discrete_valued - .insert(rumoca_core::VarName::new("a"), discrete_var("a")); - dae.variables - .discrete_valued - .insert(rumoca_core::VarName::new("b"), discrete_vector_var("b", 2)); - dae.discrete - .valued_updates - .push(connection_assignment_with_rhs_ref("a", "b")); - assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); - } - - #[test] - fn test_balance_counts_discrete_connection_cycles_by_graph_rank() { - let mut dae = dae::Dae::default(); - for name in ["a", "b", "c"] { - dae.variables - .discrete_valued - .insert(rumoca_core::VarName::new(name), discrete_var(name)); - } - dae.discrete - .valued_updates - .push(scalar_assignment_with_rhs_ref("a", "source")); - dae.discrete - .valued_updates - .push(connection_assignment_with_rhs_ref("a", "b")); - dae.discrete - .valued_updates - .push(connection_assignment_with_rhs_ref("b", "c")); - dae.discrete - .valued_updates - .push(connection_assignment_with_rhs_ref("c", "a")); - assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); - } - - #[test] - fn test_balance_skips_discrete_connection_between_component_defined_targets() { - let mut dae = dae::Dae::default(); - for name in ["a", "b"] { - dae.variables - .discrete_valued - .insert(rumoca_core::VarName::new(name), discrete_var(name)); - dae.discrete - .valued_updates - .push(scalar_assignment_with_rhs_ref(name, "source")); - } - dae.discrete - .valued_updates - .push(connection_assignment_with_rhs_ref("a", "b")); - assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); - } - - #[test] - fn test_balance_matches_subscripted_connection_rhs_to_vector_anchor() { - let mut dae = dae::Dae::default(); - dae.variables - .discrete_valued - .insert(rumoca_core::VarName::new("a"), discrete_var("a")); - dae.variables - .discrete_valued - .insert(rumoca_core::VarName::new("b"), discrete_vector_var("b", 2)); - dae.discrete - .valued_updates - .push(scalar_assignment_with_rhs_ref("a", "source")); - dae.discrete.valued_updates.push(dae::Equation { - lhs: Some(rumoca_core::VarName::new("b").into()), - scalar_count: 2, - ..scalar_assignment_with_rhs_ref("b", "source") - }); - dae.discrete - .valued_updates - .push(connection_assignment_with_rhs_index("a", "b", 1)); - assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); - } - - #[test] - fn test_balance_counts_vector_discrete_connection_cycles_by_scalar_rank() { - let mut dae = dae::Dae::default(); - for name in ["a", "b", "c"] { - dae.variables.discrete_valued.insert( - rumoca_core::VarName::new(name), - discrete_vector_var(name, 2), - ); - } - dae.discrete.valued_updates.push(dae::Equation { - scalar_count: 2, - ..scalar_assignment_with_rhs_ref("a", "source") - }); - dae.discrete - .valued_updates - .push(connection_assignment_with_count("a", "b", 2)); - dae.discrete - .valued_updates - .push(connection_assignment_with_count("b", "c", 2)); - dae.discrete - .valued_updates - .push(connection_assignment_with_count("c", "a", 2)); - assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); - } - - #[test] - fn test_balance_counts_discrete_real_updates_against_discrete_real_unknowns() { - let mut dae = dae::Dae::default(); - dae.variables - .discrete_reals - .insert(rumoca_core::VarName::new("z"), discrete_var("z")); - dae.discrete.real_updates.push(scalar_eq_with_lhs("z", 1)); - assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); - } - - #[test] - fn test_balance_counts_residual_discrete_real_updates() { - let mut dae = dae::Dae::default(); - dae.variables - .discrete_reals - .insert(rumoca_core::VarName::new("z"), discrete_var("z")); - let mut eq = scalar_eq_with_lhs("z", 1); - eq.lhs = None; - dae.discrete.real_updates.push(eq); - assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); - } - - #[test] - fn test_balance_ignores_state_reinit_updates_for_static_balance() { - let mut dae = dae::Dae::default(); - dae.variables - .states - .insert(rumoca_core::VarName::new("x"), discrete_var("x")); - dae.continuous.equations.push(scalar_eq(1)); - dae.discrete.real_updates.push(scalar_eq(1)); - assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); - } - - #[test] - fn test_balance_counts_condition_equations_against_condition_unknowns() { - let mut dae = dae::Dae::default(); - dae.variables - .discrete_valued - .insert(rumoca_core::VarName::new("c"), discrete_var("c")); - dae.conditions.equations.push(scalar_eq_with_lhs("c", 1)); - assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); - } - - #[test] - fn test_balance_ignores_condition_rhs_refs_as_unknowns() { - let mut dae = dae::Dae::default(); - dae.variables - .discrete_valued - .insert(rumoca_core::VarName::new("guard"), discrete_var("guard")); - dae.variables - .discrete_valued - .insert(rumoca_core::VarName::new("c"), discrete_var("c")); - dae.conditions - .equations - .push(scalar_assignment_with_rhs_ref("c", "guard")); - assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); - } - - #[test] - fn component_defined_targets_include_second_binary_ref_when_it_is_the_only_unknown() { - let mut dae = dae::Dae::default(); - dae.variables - .algebraics - .insert(rumoca_core::VarName::new("y"), discrete_var("y")); - dae.variables - .inputs - .insert(rumoca_core::VarName::new("u"), scalar_input("u")); - dae.variables.inputs.insert( - rumoca_core::VarName::new("external"), - scalar_input("external"), - ); - - dae.continuous - .equations - .push(binary_eq("u", "y", "component equation")); - dae.continuous - .equations - .push(binary_eq("y", "external", "connect(y, external)")); - - assert_eq!( - balance(&dae).expect("valid DAE balance fixture"), - 0, - "the connection alias should be redundant because y is already constrained" - ); - } - - #[test] - fn component_defined_targets_do_not_treat_two_unknown_residual_as_two_definitions() { - let mut dae = dae::Dae::default(); - for name in ["a", "b"] { - dae.variables - .algebraics - .insert(rumoca_core::VarName::new(name), discrete_var(name)); - } - dae.variables.inputs.insert( - rumoca_core::VarName::new("external"), - scalar_input("external"), - ); - - dae.continuous - .equations - .push(binary_eq("a", "b", "component equation")); - dae.continuous - .equations - .push(binary_eq("b", "external", "connect(b, external)")); - - assert_eq!( - balance(&dae).expect("valid DAE balance fixture"), - 0, - "the connection still supplies the second equation for coupled unknowns" - ); - } -} +mod tests; diff --git a/crates/rumoca-phase-dae/src/balance/tests.rs b/crates/rumoca-phase-dae/src/balance/tests.rs new file mode 100644 index 000000000..ee7ab0334 --- /dev/null +++ b/crates/rumoca-phase-dae/src/balance/tests.rs @@ -0,0 +1,1084 @@ +use super::*; +use rumoca_core::Span; + +fn test_span() -> Span { + Span::from_offsets( + rumoca_core::SourceId::from_source_name("balance_fixture.mo"), + 1, + 2, + ) +} + +fn scalar_eq(count: usize) -> dae::Equation { + scalar_eq_with_lhs("x", count) +} + +fn scalar_eq_with_lhs(lhs_name: &str, count: usize) -> dae::Equation { + dae::Equation { + lhs: Some(rumoca_core::VarName::new(lhs_name).into()), + rhs: rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new(lhs_name).into(), + subscripts: vec![], + span: test_span(), + }), + rhs: Box::new(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(0), + span: test_span(), + }), + span: test_span(), + }, + span: test_span(), + origin: "test".to_string(), + scalar_count: count, + } +} + +fn scalar_assignment_with_rhs_ref(lhs_name: &str, rhs_name: &str) -> dae::Equation { + dae::Equation { + lhs: Some(rumoca_core::VarName::new(lhs_name).into()), + rhs: rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new(rhs_name).into(), + subscripts: vec![], + span: test_span(), + }, + span: test_span(), + origin: "test".to_string(), + scalar_count: 1, + } +} + +fn connection_assignment_with_rhs_ref(lhs_name: &str, rhs_name: &str) -> dae::Equation { + dae::Equation { + origin: "explicit connection equation: a = b".to_string(), + ..scalar_assignment_with_rhs_ref(lhs_name, rhs_name) + } +} + +fn connection_assignment_with_count( + lhs_name: &str, + rhs_name: &str, + scalar_count: usize, +) -> dae::Equation { + dae::Equation { + scalar_count, + ..connection_assignment_with_rhs_ref(lhs_name, rhs_name) + } +} + +fn binary_residual_eq_with_count( + lhs_name: &str, + rhs_name: &str, + origin: &str, + scalar_count: usize, +) -> dae::Equation { + dae::Equation { + lhs: Some(rumoca_core::VarName::new(lhs_name).into()), + rhs: rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new(lhs_name).into(), + subscripts: vec![], + span: test_span(), + }), + rhs: Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new(rhs_name).into(), + subscripts: vec![], + span: test_span(), + }), + span: test_span(), + }, + span: test_span(), + origin: origin.to_string(), + scalar_count, + } +} + +fn der_ref(name: &str) -> rumoca_core::Expression { + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Der, + args: vec![rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new(name).into(), + subscripts: vec![], + span: test_span(), + }], + span: test_span(), + } +} + +fn overconstrained_derivative_alias_eq(lhs_name: &str, rhs_name: &str) -> dae::Equation { + dae::Equation { + lhs: None, + rhs: rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(der_ref(lhs_name)), + rhs: Box::new(der_ref(rhs_name)), + span: test_span(), + }, + span: test_span(), + origin: format!("overconstrained derivative alias: {lhs_name} = {rhs_name}"), + scalar_count: 1, + } +} + +fn structured_reference(name: &str) -> rumoca_core::Reference { + let var_name = rumoca_core::VarName::new(name); + let component_ref = rumoca_core::component_reference_from_flat_name(&var_name, test_span()) + .expect("structured component reference"); + rumoca_core::Reference::with_component_reference(name, component_ref) +} + +fn binary_residual_eq_with_structured_refs( + lhs_name: &str, + rhs_name: &str, + origin: &str, + scalar_count: usize, +) -> dae::Equation { + dae::Equation { + lhs: Some(structured_reference(lhs_name)), + rhs: rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(rumoca_core::Expression::VarRef { + name: structured_reference(lhs_name), + subscripts: vec![], + span: test_span(), + }), + rhs: Box::new(rumoca_core::Expression::VarRef { + name: structured_reference(rhs_name), + subscripts: vec![], + span: test_span(), + }), + span: test_span(), + }, + span: test_span(), + origin: origin.to_string(), + scalar_count, + } +} + +fn connection_assignment_with_rhs_index( + lhs_name: &str, + rhs_name: &str, + rhs_index: i64, +) -> dae::Equation { + dae::Equation { + lhs: Some(rumoca_core::VarName::new(lhs_name).into()), + rhs: rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new(rhs_name).into(), + subscripts: vec![rumoca_core::Subscript::generated_index( + rhs_index, + test_span(), + )], + span: test_span(), + }, + span: test_span(), + origin: "explicit connection equation: a = b[1]".to_string(), + scalar_count: 1, + } +} + +fn dae_with_unknown_scalars(unknown_scalars: i64) -> dae::Dae { + let mut dae = dae::Dae::default(); + dae.variables.algebraics.insert( + rumoca_core::VarName::new("x"), + dae::Variable { + name: rumoca_core::VarName::new("x"), + dims: vec![unknown_scalars], + ..rumoca_ir_dae::Variable::empty_with_span(rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + )) + }, + ); + dae +} + +fn discrete_var(name: &str) -> dae::Variable { + dae::Variable { + name: rumoca_core::VarName::new(name), + ..rumoca_ir_dae::Variable::empty_with_span(rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + )) + } +} + +fn discrete_vector_var(name: &str, size: i64) -> dae::Variable { + dae::Variable { + name: rumoca_core::VarName::new(name), + dims: vec![size], + ..rumoca_ir_dae::Variable::empty_with_span(rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + )) + } +} + +fn algebraic_vector_var(name: &str, size: i64) -> dae::Variable { + dae::Variable { + name: rumoca_core::VarName::new(name), + dims: vec![size], + ..rumoca_ir_dae::Variable::empty_with_span(test_span()) + } +} + +fn algebraic_var(name: &str) -> dae::Variable { + dae::Variable { + name: rumoca_core::VarName::new(name), + ..rumoca_ir_dae::Variable::empty_with_span(test_span()) + } +} + +fn top_level_connector_algebraic_var(name: &str) -> dae::Variable { + let name = rumoca_core::VarName::new(name); + dae::Variable { + component_ref: rumoca_core::component_reference_from_flat_name(&name, test_span()), + name, + ..rumoca_ir_dae::Variable::empty_with_span(test_span()) + } +} + +fn unconnected_flow_anchor(name: &str) -> dae::Equation { + dae::Equation { + origin: format!("unconnected flow: {name} = 0"), + ..scalar_eq_with_lhs(name, 1) + } +} + +fn explicit_scalar_equation(name: &str) -> dae::Equation { + dae::Equation { + origin: format!("equation from {name} = 0"), + ..scalar_eq_with_lhs(name, 1) + } +} + +fn connection_flow_sum_equation(lhs_name: &str, rhs_name: &str) -> dae::Equation { + dae::Equation { + lhs: None, + rhs: rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new(lhs_name).into(), + subscripts: vec![], + span: test_span(), + }), + rhs: Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new(rhs_name).into(), + subscripts: vec![], + span: test_span(), + }), + span: test_span(), + }, + span: test_span(), + origin: format!("connect({lhs_name}, {rhs_name})"), + scalar_count: 1, + } +} + +#[test] +fn test_balance_clamps_overconstrained_interface_to_deficit() { + let mut dae = dae_with_unknown_scalars(4); + dae.continuous.equations.push(scalar_eq(4)); + dae.metadata.overconstrained_interface_count = 9; + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); +} + +#[test] +fn admission_accepts_exact_initial_deficit_closure() { + let mut dae = dae_with_unknown_scalars(3); + dae.continuous.equations.push(scalar_eq(1)); + dae.initialization + .equations + .push(scalar_eq_with_lhs("x", 2)); + + let detail = initial_closure_balance_detail(&dae).expect("valid DAE initial balance fixture"); + + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), -2); + assert_eq!(detail.deficit_before, 2); + assert_eq!(detail.initial_equation_scalars, 2); + assert_eq!(detail.closure_used, 2); + assert_eq!(detail.deficit_after, 0); + assert!(detail.is_admissible()); + assert!(is_balanced_for_admission(&dae).expect("valid DAE balance fixture")); +} + +#[test] +fn admission_rejects_partial_initial_deficit_closure() { + let mut dae = dae_with_unknown_scalars(4); + dae.continuous.equations.push(scalar_eq(1)); + dae.initialization + .equations + .push(scalar_eq_with_lhs("x", 2)); + + let detail = initial_closure_balance_detail(&dae).expect("valid DAE initial balance fixture"); + + assert_eq!(detail.deficit_before, 3); + assert_eq!(detail.closure_used, 2); + assert_eq!(detail.deficit_after, 1); + assert!(!detail.is_admissible()); + assert!(!is_balanced_for_admission(&dae).expect("valid DAE balance fixture")); +} + +#[test] +fn admission_uses_overconstrained_break_edges_as_deficit_closure() { + let mut dae = dae_with_unknown_scalars(4); + dae.continuous.equations.push(scalar_eq(2)); + dae.metadata.oc_break_edge_scalar_count = 2; + let detail = initial_closure_balance_detail(&dae).expect("valid DAE fixture"); + + assert_eq!( + ( + balance(&dae).expect("valid DAE fixture"), + detail.deficit_before, + detail.overconstrained_break_edge_scalars, + detail.closure_used, + detail.deficit_after, + is_balanced_for_admission(&dae).expect("valid DAE fixture"), + ), + (-2, 2, 2, 2, 0, true) + ); +} + +#[test] +fn admission_rejects_missing_initial_deficit_closure() { + let mut dae = dae_with_unknown_scalars(2); + dae.continuous.equations.push(scalar_eq(1)); + + let detail = initial_closure_balance_detail(&dae).expect("valid DAE initial balance fixture"); + + assert_eq!(detail.deficit_before, 1); + assert_eq!(detail.initial_equation_scalars, 0); + assert_eq!(detail.deficit_after, 1); + assert!(!is_balanced_for_admission(&dae).expect("valid DAE balance fixture")); +} + +#[test] +fn admission_does_not_mask_overdetermined_balance_with_initial_equations() { + let mut dae = dae_with_unknown_scalars(1); + dae.continuous.equations.push(scalar_eq(2)); + dae.initialization + .equations + .push(scalar_eq_with_lhs("x", 3)); + + let detail = initial_closure_balance_detail(&dae).expect("valid DAE initial balance fixture"); + + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 1); + assert_eq!(detail.deficit_before, 0); + assert_eq!(detail.closure_used, 0); + assert!(!detail.is_admissible()); + assert!(!is_balanced_for_admission(&dae).expect("valid DAE balance fixture")); +} + +#[test] +fn test_balance_uses_only_needed_overconstrained_interface() { + let mut dae = dae_with_unknown_scalars(4); + dae.continuous.equations.push(scalar_eq(3)); + dae.metadata.overconstrained_interface_count = 9; + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); +} + +#[test] +fn test_balance_applies_oc_interface_even_with_break_edges() { + let mut dae = dae_with_unknown_scalars(10); + dae.continuous.equations.push(scalar_eq(1)); + dae.metadata.overconstrained_interface_count = 9; + dae.metadata.oc_break_edge_scalar_count = 12; + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); +} + +#[test] +fn test_balance_clamps_interface_flow_to_remaining_deficit() { + let mut dae = dae_with_unknown_scalars(4); + dae.continuous.equations.push(scalar_eq(4)); + dae.metadata.interface_flow_count = 3; + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); +} + +#[test] +fn test_balance_uses_interface_flow_to_close_deficit_only() { + let mut dae = dae_with_unknown_scalars(5); + dae.continuous.equations.push(scalar_eq(3)); + dae.metadata.interface_flow_count = 9; + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); +} + +#[test] +fn top_level_flow_anchor_is_not_double_counted_when_single_unknown_equation_closes_it() { + let mut dae = dae::Dae::default(); + dae.variables.algebraics.insert( + rumoca_core::VarName::new("pin.i"), + top_level_connector_algebraic_var("pin.i"), + ); + dae.continuous + .equations + .push(explicit_scalar_equation("pin.i")); + dae.continuous + .equations + .push(unconnected_flow_anchor("pin.i")); + + let detail = balance_detail(&dae).expect("valid DAE balance fixture"); + + assert_eq!(detail.f_x_scalar, 1); + assert_eq!(detail.balance(), 0); +} + +#[test] +fn top_level_flow_anchor_is_counted_when_flow_only_appears_in_connection_sum() { + let mut dae = dae::Dae::default(); + dae.variables.algebraics.insert( + rumoca_core::VarName::new("a.i"), + top_level_connector_algebraic_var("a.i"), + ); + dae.variables.algebraics.insert( + rumoca_core::VarName::new("b.i"), + top_level_connector_algebraic_var("b.i"), + ); + dae.continuous + .equations + .push(connection_flow_sum_equation("a.i", "b.i")); + dae.continuous + .equations + .push(unconnected_flow_anchor("a.i")); + + let detail = balance_detail(&dae).expect("valid DAE balance fixture"); + + assert_eq!(detail.f_x_scalar, 2); + assert_eq!(detail.balance(), 0); +} + +#[test] +fn test_balance_clamps_stream_interface_equations_to_remaining_deficit() { + let mut dae = dae_with_unknown_scalars(4); + dae.continuous.equations.push(scalar_eq(4)); + dae.metadata.stream_interface_equation_count = 3; + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); +} + +#[test] +fn test_balance_uses_stream_interface_equations_to_close_deficit_only() { + let mut dae = dae_with_unknown_scalars(5); + dae.continuous.equations.push(scalar_eq(3)); + dae.metadata.stream_interface_equation_count = 9; + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); +} + +#[test] +fn test_balance_ignores_unconstrained_discrete_real_declaration() { + let mut dae = dae::Dae::default(); + dae.variables + .discrete_reals + .insert(rumoca_core::VarName::new("z"), discrete_var("z")); + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); +} + +#[test] +fn test_balance_counts_discrete_updates_against_discrete_unknowns() { + let mut dae = dae::Dae::default(); + dae.variables + .discrete_valued + .insert(rumoca_core::VarName::new("m"), discrete_var("m")); + dae.discrete.valued_updates.push(scalar_eq_with_lhs("m", 1)); + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); +} + +#[test] +fn test_balance_ignores_solved_discrete_update_rhs_refs_as_unknowns() { + let mut dae = dae::Dae::default(); + dae.variables + .discrete_valued + .insert(rumoca_core::VarName::new("guard"), discrete_var("guard")); + dae.variables + .discrete_valued + .insert(rumoca_core::VarName::new("m"), discrete_var("m")); + dae.discrete + .valued_updates + .push(scalar_assignment_with_rhs_ref("m", "guard")); + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); +} + +#[test] +fn test_balance_counts_discrete_connection_update_rhs_refs_as_unknowns() { + let mut dae = dae::Dae::default(); + dae.variables + .discrete_valued + .insert(rumoca_core::VarName::new("a"), discrete_var("a")); + dae.variables + .discrete_valued + .insert(rumoca_core::VarName::new("b"), discrete_var("b")); + dae.discrete + .valued_updates + .push(connection_assignment_with_rhs_ref("a", "b")); + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), -1); +} + +#[test] +fn balance_rejects_missing_discrete_input_metadata_variable() { + let mut dae = dae::Dae::default(); + dae.metadata + .discrete_input_names + .push("missing_input".to_string()); + + let err = balance(&dae).expect_err("missing discrete metadata should fail"); + + assert!(matches!( + err, + BalanceError::MissingDiscreteVariableMetadata { ref name, .. } + if name.as_str() == "missing_input" + )); +} + +#[test] +fn test_balance_ignores_partial_vector_connection_rhs_refs_as_unknowns() { + let mut dae = dae::Dae::default(); + dae.variables + .discrete_valued + .insert(rumoca_core::VarName::new("a"), discrete_var("a")); + dae.variables + .discrete_valued + .insert(rumoca_core::VarName::new("b"), discrete_vector_var("b", 2)); + dae.discrete + .valued_updates + .push(connection_assignment_with_rhs_ref("a", "b")); + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); +} + +#[test] +fn test_balance_expands_scalarized_discrete_connection_targets() { + let mut dae = dae::Dae::default(); + for name in ["a[1]", "a[2]", "b[1]", "b[2]"] { + dae.variables + .discrete_valued + .insert(rumoca_core::VarName::new(name), discrete_var(name)); + } + dae.discrete + .valued_updates + .push(connection_assignment_with_count("a", "b", 2)); + + assert_eq!( + balance(&dae).expect("scalarized vector connection should expand to children"), + 2 + ); +} + +#[test] +fn test_balance_anchors_size_one_discrete_input_vectors_by_index() { + let mut dae = dae::Dae::default(); + dae.variables.discrete_valued.insert( + rumoca_core::VarName::new("source"), + discrete_vector_var("source", 1), + ); + for name in ["a", "b"] { + dae.variables + .discrete_valued + .insert(rumoca_core::VarName::new(name), discrete_var(name)); + } + dae.metadata.discrete_input_names.push("source".to_string()); + dae.discrete.valued_updates.push(scalar_eq_with_lhs("b", 1)); + dae.discrete + .valued_updates + .push(connection_assignment_with_rhs_index("a", "source", 1)); + dae.discrete + .valued_updates + .push(connection_assignment_with_rhs_ref("a", "b")); + + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); +} + +#[test] +fn test_balance_counts_discrete_connection_cycles_by_graph_rank() { + let mut dae = dae::Dae::default(); + for name in ["a", "b", "c"] { + dae.variables + .discrete_valued + .insert(rumoca_core::VarName::new(name), discrete_var(name)); + } + dae.discrete + .valued_updates + .push(scalar_assignment_with_rhs_ref("a", "source")); + dae.discrete + .valued_updates + .push(connection_assignment_with_rhs_ref("a", "b")); + dae.discrete + .valued_updates + .push(connection_assignment_with_rhs_ref("b", "c")); + dae.discrete + .valued_updates + .push(connection_assignment_with_rhs_ref("c", "a")); + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); +} + +#[test] +fn test_balance_skips_discrete_connection_between_component_defined_targets() { + let mut dae = dae::Dae::default(); + for name in ["a", "b"] { + dae.variables + .discrete_valued + .insert(rumoca_core::VarName::new(name), discrete_var(name)); + dae.discrete + .valued_updates + .push(scalar_assignment_with_rhs_ref(name, "source")); + } + dae.discrete + .valued_updates + .push(connection_assignment_with_rhs_ref("a", "b")); + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); +} + +#[test] +fn test_balance_matches_subscripted_connection_rhs_to_vector_anchor() { + let mut dae = dae::Dae::default(); + dae.variables + .discrete_valued + .insert(rumoca_core::VarName::new("a"), discrete_var("a")); + dae.variables + .discrete_valued + .insert(rumoca_core::VarName::new("b"), discrete_vector_var("b", 2)); + dae.discrete + .valued_updates + .push(scalar_assignment_with_rhs_ref("a", "source")); + dae.discrete.valued_updates.push(dae::Equation { + lhs: Some(rumoca_core::VarName::new("b").into()), + scalar_count: 2, + ..scalar_assignment_with_rhs_ref("b", "source") + }); + dae.discrete + .valued_updates + .push(connection_assignment_with_rhs_index("a", "b", 1)); + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); +} + +#[test] +fn test_balance_counts_vector_discrete_connection_cycles_by_scalar_rank() { + let mut dae = dae::Dae::default(); + for name in ["a", "b", "c"] { + dae.variables.discrete_valued.insert( + rumoca_core::VarName::new(name), + discrete_vector_var(name, 2), + ); + } + dae.discrete.valued_updates.push(dae::Equation { + scalar_count: 2, + ..scalar_assignment_with_rhs_ref("a", "source") + }); + dae.discrete + .valued_updates + .push(connection_assignment_with_count("a", "b", 2)); + dae.discrete + .valued_updates + .push(connection_assignment_with_count("b", "c", 2)); + dae.discrete + .valued_updates + .push(connection_assignment_with_count("c", "a", 2)); + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); +} + +#[test] +fn test_balance_counts_discrete_real_updates_against_discrete_real_unknowns() { + let mut dae = dae::Dae::default(); + dae.variables + .discrete_reals + .insert(rumoca_core::VarName::new("z"), discrete_var("z")); + dae.discrete.real_updates.push(scalar_eq_with_lhs("z", 1)); + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); +} + +#[test] +fn test_balance_counts_residual_discrete_real_updates() { + let mut dae = dae::Dae::default(); + dae.variables + .discrete_reals + .insert(rumoca_core::VarName::new("z"), discrete_var("z")); + let mut eq = scalar_eq_with_lhs("z", 1); + eq.lhs = None; + dae.discrete.real_updates.push(eq); + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); +} + +#[test] +fn test_balance_ignores_state_reinit_updates_for_static_balance() { + let mut dae = dae::Dae::default(); + dae.variables + .states + .insert(rumoca_core::VarName::new("x"), discrete_var("x")); + dae.continuous.equations.push(scalar_eq(1)); + dae.discrete.real_updates.push(scalar_eq(1)); + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); +} + +#[test] +fn test_balance_counts_condition_equations_against_condition_unknowns() { + let mut dae = dae::Dae::default(); + dae.variables + .discrete_valued + .insert(rumoca_core::VarName::new("c"), discrete_var("c")); + dae.conditions.equations.push(scalar_eq_with_lhs("c", 1)); + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); +} + +#[test] +fn balance_spends_component_vector_aliases_only_against_surplus() { + let mut dae = dae::Dae::default(); + for name in ["a", "b", "c"] { + dae.variables.algebraics.insert( + rumoca_core::VarName::new(name), + algebraic_vector_var(name, 3), + ); + } + dae.continuous.equations.push(binary_residual_eq_with_count( + "a", + "source", + "component equation", + 3, + )); + dae.continuous.equations.push(dae::Equation { + scalar_count: 6, + ..scalar_eq_with_lhs("b", 6) + }); + dae.continuous.equations.push(binary_residual_eq_with_count( + "a", + "c", + "connection equation: a = c", + 3, + )); + + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); + + for name in ["d", "e"] { + dae.variables.algebraics.insert( + rumoca_core::VarName::new(name), + algebraic_vector_var(name, 3), + ); + } + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), -3); +} + +#[test] +fn balance_spends_component_vector_forwarding_only_against_surplus() { + let mut dae = dae::Dae::default(); + for name in ["source", "alias"] { + dae.variables.algebraics.insert( + rumoca_core::VarName::new(name), + algebraic_vector_var(name, 3), + ); + } + dae.continuous.equations.push(binary_residual_eq_with_count( + "source", + "source_driver", + "component equation", + 3, + )); + dae.continuous.equations.push(binary_residual_eq_with_count( + "alias", + "alias_driver", + "component equation", + 3, + )); + dae.continuous.equations.push(binary_residual_eq_with_count( + "alias", + "source", + "equation from adaptor", + 3, + )); + + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); + + for name in ["d", "e"] { + dae.variables.algebraics.insert( + rumoca_core::VarName::new(name), + algebraic_vector_var(name, 3), + ); + } + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), -3); +} + +#[test] +fn balance_spends_overconstrained_derivative_aliases_only_against_surplus() { + let mut dae = dae::Dae::default(); + for name in ["branch.port.reference.gamma", "root.port.reference.gamma"] { + dae.variables + .algebraics + .insert(rumoca_core::VarName::new(name), algebraic_var(name)); + } + dae.continuous.equations.push(binary_residual_eq_with_count( + "branch.port.reference.gamma", + "branch_driver", + "component equation", + 1, + )); + dae.continuous.equations.push(binary_residual_eq_with_count( + "root.port.reference.gamma", + "root_driver", + "component equation", + 1, + )); + dae.continuous + .equations + .push(overconstrained_derivative_alias_eq( + "branch.port.reference.gamma", + "root.port.reference.gamma", + )); + + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); + + for name in ["unclosed_a", "unclosed_b"] { + dae.variables + .algebraics + .insert(rumoca_core::VarName::new(name), algebraic_var(name)); + } + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), -1); +} + +#[test] +fn balance_does_not_spend_overconstrained_origin_without_derivative_alias_shape() { + let mut dae = dae::Dae::default(); + for name in ["branch.port.reference.gamma", "root.port.reference.gamma"] { + dae.variables + .algebraics + .insert(rumoca_core::VarName::new(name), algebraic_var(name)); + } + dae.continuous.equations.push(binary_residual_eq_with_count( + "branch.port.reference.gamma", + "branch_driver", + "component equation", + 1, + )); + dae.continuous.equations.push(binary_residual_eq_with_count( + "root.port.reference.gamma", + "root_driver", + "component equation", + 1, + )); + dae.continuous.equations.push(binary_residual_eq_with_count( + "branch.port.reference.gamma", + "root.port.reference.gamma", + "overconstrained derivative alias: branch.port.reference.gamma = root.port.reference.gamma", + 1, + )); + + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 1); +} + +#[test] +fn balance_spends_scalarized_vector_bindings_only_against_surplus() { + let mut dae = dae::Dae::default(); + dae.variables.outputs.insert( + rumoca_core::VarName::new("alias"), + algebraic_vector_var("alias", 3), + ); + for name in ["source", "extra"] { + dae.variables.algebraics.insert( + rumoca_core::VarName::new(name), + algebraic_vector_var(name, 3), + ); + } + dae.continuous.equations.push(binary_residual_eq_with_count( + "source", + "source_driver", + "component equation", + 3, + )); + dae.continuous.equations.push(binary_residual_eq_with_count( + "extra", + "extra_driver", + "component equation", + 6, + )); + for idx in 1..=3 { + dae.continuous.equations.push(binary_residual_eq_with_count( + "alias", + "source", + &format!("binding equation for alias [scalarized {idx}]"), + 1, + )); + } + + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); + + for name in ["d", "e"] { + dae.variables.algebraics.insert( + rumoca_core::VarName::new(name), + algebraic_vector_var(name, 3), + ); + } + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), -3); +} + +#[test] +fn balance_spends_scalar_component_binding_aliases_only_against_surplus() { + let mut dae = dae::Dae::default(); + dae.variables + .outputs + .insert(rumoca_core::VarName::new("alias"), algebraic_var("alias")); + for name in ["source", "extra"] { + dae.variables + .algebraics + .insert(rumoca_core::VarName::new(name), algebraic_var(name)); + } + dae.continuous.equations.push(binary_residual_eq_with_count( + "source", + "source_driver", + "component equation", + 1, + )); + dae.continuous.equations.push(binary_residual_eq_with_count( + "extra", + "extra_driver", + "component equation", + 2, + )); + dae.continuous.equations.push(binary_residual_eq_with_count( + "alias", + "source", + "binding equation for alias", + 1, + )); + + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); + + for name in ["d", "e"] { + dae.variables + .algebraics + .insert(rumoca_core::VarName::new(name), algebraic_var(name)); + } + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), -1); +} + +#[test] +fn balance_spends_simple_continuous_binding_aliases_only_against_surplus() { + let mut dae = dae::Dae::default(); + for name in ["alias", "source", "extra"] { + dae.variables + .algebraics + .insert(rumoca_core::VarName::new(name), algebraic_var(name)); + } + dae.continuous.equations.push(binary_residual_eq_with_count( + "source", + "source_driver", + "component equation", + 1, + )); + dae.continuous.equations.push(binary_residual_eq_with_count( + "extra", + "extra_driver", + "component equation", + 2, + )); + dae.continuous.equations.push(binary_residual_eq_with_count( + "alias", + "source", + "binding equation for alias", + 1, + )); + + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); + + for name in ["d", "e"] { + dae.variables + .algebraics + .insert(rumoca_core::VarName::new(name), algebraic_var(name)); + } + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), -1); +} + +#[test] +fn balance_spends_simple_continuous_component_equation_aliases_only_against_surplus() { + let mut dae = dae::Dae::default(); + for name in ["comp.alias", "comp.source", "comp.extra"] { + dae.variables + .algebraics + .insert(rumoca_core::VarName::new(name), algebraic_var(name)); + } + dae.continuous.equations.push(binary_residual_eq_with_count( + "comp.source", + "source_driver", + "component equation", + 1, + )); + dae.continuous.equations.push(binary_residual_eq_with_count( + "comp.extra", + "extra_driver", + "component equation", + 2, + )); + dae.continuous + .equations + .push(binary_residual_eq_with_structured_refs( + "comp.alias", + "comp.source", + "equation from comp.alias", + 1, + )); + + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); + + for name in ["d", "e"] { + dae.variables + .algebraics + .insert(rumoca_core::VarName::new(name), algebraic_var(name)); + } + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), -1); +} + +#[test] +fn balance_spends_scalarized_vector_forwarding_only_against_surplus() { + let mut dae = dae::Dae::default(); + dae.variables.outputs.insert( + rumoca_core::VarName::new("alias"), + algebraic_vector_var("alias", 3), + ); + for name in ["source", "extra"] { + dae.variables.algebraics.insert( + rumoca_core::VarName::new(name), + algebraic_vector_var(name, 3), + ); + } + dae.continuous.equations.push(binary_residual_eq_with_count( + "source", + "source_driver", + "component equation", + 3, + )); + dae.continuous.equations.push(binary_residual_eq_with_count( + "extra", + "extra_driver", + "component equation", + 6, + )); + for idx in 1..=3 { + dae.continuous.equations.push(binary_residual_eq_with_count( + "alias", + "source", + &format!("equation from adaptor [scalarized {idx}]"), + 1, + )); + } + + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); + + for name in ["d", "e"] { + dae.variables.algebraics.insert( + rumoca_core::VarName::new(name), + algebraic_vector_var(name, 3), + ); + } + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), -3); +} + +#[test] +fn test_balance_ignores_condition_rhs_refs_as_unknowns() { + let mut dae = dae::Dae::default(); + dae.variables + .discrete_valued + .insert(rumoca_core::VarName::new("guard"), discrete_var("guard")); + dae.variables + .discrete_valued + .insert(rumoca_core::VarName::new("c"), discrete_var("c")); + dae.conditions + .equations + .push(scalar_assignment_with_rhs_ref("c", "guard")); + assert_eq!(balance(&dae).expect("valid DAE balance fixture"), 0); +} diff --git a/crates/rumoca-phase-dae/src/balance_alias.rs b/crates/rumoca-phase-dae/src/balance_alias.rs new file mode 100644 index 000000000..551556b8e --- /dev/null +++ b/crates/rumoca-phase-dae/src/balance_alias.rs @@ -0,0 +1,983 @@ +use super::{BalanceSymbolSet, eq_binary_var_refs}; +use rumoca_ir_dae as dae; + +pub(super) fn is_vector_forwarding_alias( + eq: &dae::Equation, + continuous_unknowns: &BalanceSymbolSet, + _output_names: &BalanceSymbolSet, + _component_defined_targets: &BalanceSymbolSet, +) -> bool { + if eq.scalar_count <= 1 { + let Some((lhs, rhs)) = eq_binary_lhs_rhs(&eq.rhs) else { + return false; + }; + return eq.origin.starts_with("equation from ") + && (is_unresolved_reference_field_alias(lhs, rhs, continuous_unknowns) + || is_rooted_two_pin_reference_anchor_alias(eq.origin.as_str(), lhs, rhs)); + } + if !(eq.origin.starts_with("binding equation for") || eq.origin.starts_with("equation from ")) { + return false; + } + let Some((lhs, rhs)) = eq_binary_lhs_rhs(&eq.rhs) else { + return false; + }; + if matches!(lhs, rumoca_core::Expression::FieldAccess { .. }) { + return !lhs_matches_symbols(lhs, continuous_unknowns); + } + if forwarded_value_ref(rhs).is_none() { + return false; + } + if !lhs_is_forwarding_target(lhs) { + return false; + } + !lhs_matches_symbols(lhs, continuous_unknowns) +} + +pub(super) fn is_surplus_component_vector_forwarding_alias( + eq: &dae::Equation, + continuous_unknowns: &BalanceSymbolSet, + _output_names: &BalanceSymbolSet, + _component_defined_targets: &BalanceSymbolSet, +) -> bool { + if (eq.scalar_count <= 1 && !eq.origin.contains(" [scalarized ")) + || !(eq.origin.starts_with("binding equation for") + || eq.origin.starts_with("equation from ")) + { + return false; + } + let Some((lhs, rhs)) = eq_binary_lhs_rhs(&eq.rhs) else { + return false; + }; + if forwarded_value_ref(rhs).is_none() { + return false; + } + lhs_is_forwarding_target(lhs) && lhs_matches_symbols(lhs, continuous_unknowns) +} + +pub(super) fn is_surplus_overconstrained_derivative_alias(eq: &dae::Equation) -> bool { + if eq.scalar_count != 1 || !eq.origin.starts_with("overconstrained derivative alias:") { + return false; + } + let Some((lhs, rhs)) = eq_binary_lhs_rhs(&eq.rhs) else { + return false; + }; + let Some(lhs_name) = derivative_alias_reference_name(lhs) else { + return false; + }; + let Some(rhs_name) = derivative_alias_reference_name(rhs) else { + return false; + }; + lhs_name.as_str().ends_with(".reference.gamma") + && rhs_name.as_str().ends_with(".reference.gamma") +} + +fn derivative_alias_reference_name(expr: &rumoca_core::Expression) -> Option { + let rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Der, + args, + .. + } = expr + else { + return None; + }; + let [arg] = args.as_slice() else { + return None; + }; + expression_var_name(arg) +} + +fn lhs_is_forwarding_target(expr: &rumoca_core::Expression) -> bool { + matches!( + expr, + rumoca_core::Expression::VarRef { .. } | rumoca_core::Expression::FieldAccess { .. } + ) +} + +fn is_unresolved_reference_field_alias( + lhs: &rumoca_core::Expression, + rhs: &rumoca_core::Expression, + continuous_unknowns: &BalanceSymbolSet, +) -> bool { + let Some(lhs_name) = field_access_var_name(lhs) else { + return false; + }; + let Some(rhs_name) = field_access_var_name(rhs) else { + return false; + }; + lhs_name.as_str().ends_with(".reference.gamma") + && rhs_name.as_str().ends_with(".reference.gamma") + && !continuous_unknowns.matches_name(&lhs_name) +} + +fn is_rooted_two_pin_reference_anchor_alias( + origin: &str, + lhs: &rumoca_core::Expression, + rhs: &rumoca_core::Expression, +) -> bool { + let Some(component) = origin.strip_prefix("equation from ") else { + return false; + }; + let Some(lhs_name) = expression_var_name(lhs) else { + return false; + }; + let Some(rhs_name) = expression_var_name(rhs) else { + return false; + }; + let plug_p_reference = format!("{component}.plug_p.reference.gamma"); + let plug_n_reference = format!("{component}.plug_n.reference.gamma"); + let internal_reference = format!("{component}.gamma"); + lhs_name.as_str() == plug_p_reference + && (rhs_name.as_str() == plug_n_reference || rhs_name.as_str() == internal_reference) +} + +fn lhs_matches_symbols(expr: &rumoca_core::Expression, symbols: &BalanceSymbolSet) -> bool { + match expr { + rumoca_core::Expression::VarRef { name, .. } => symbols.matches_reference(name), + rumoca_core::Expression::FieldAccess { base, field, .. } => { + field_access_var_name(expr).is_some_and(|name| symbols.matches_name(&name)) + || field_access_array_members_match_symbols(base, field, symbols) + || matches!( + base.as_ref(), + rumoca_core::Expression::VarRef { name, .. } if symbols.matches_reference(name) + ) + } + _ => false, + } +} + +fn field_access_array_members_match_symbols( + base: &rumoca_core::Expression, + field: &str, + symbols: &BalanceSymbolSet, +) -> bool { + let rumoca_core::Expression::VarRef { name, .. } = base else { + return false; + }; + let array_prefix = format!("{}[", name.var_name().as_str()); + let field_marker = format!("].{field}"); + symbols.names.iter().any(|candidate| { + let candidate = candidate.as_str(); + candidate.starts_with(&array_prefix) && candidate.contains(&field_marker) + }) || symbols.prefixes.iter().any(|candidate| { + let candidate = candidate.as_str(); + candidate.starts_with(&array_prefix) && candidate.contains(&field_marker) + }) +} + +fn field_access_var_name(expr: &rumoca_core::Expression) -> Option { + let rumoca_core::Expression::FieldAccess { base, field, .. } = expr else { + return None; + }; + Some(rumoca_core::VarName::new(format!( + "{}.{field}", + indexed_expression_var_name(base)? + ))) +} + +fn expression_var_name(expr: &rumoca_core::Expression) -> Option { + match expr { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } => { + let suffix = subscript_suffix(subscripts)?; + Some(rumoca_core::VarName::new(format!( + "{}{suffix}", + name.var_name().as_str() + ))) + } + rumoca_core::Expression::FieldAccess { .. } => field_access_var_name(expr), + rumoca_core::Expression::Index { .. } => Some(rumoca_core::VarName::new( + indexed_expression_var_name(expr)?, + )), + _ => None, + } +} + +fn indexed_expression_var_name(expr: &rumoca_core::Expression) -> Option { + match expr { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } => Some(format!( + "{}{}", + name.var_name().as_str(), + subscript_suffix(subscripts)? + )), + rumoca_core::Expression::Index { + base, subscripts, .. + } => Some(format!( + "{}{}", + indexed_expression_var_name(base)?, + subscript_suffix(subscripts)? + )), + rumoca_core::Expression::FieldAccess { base, field, .. } => { + Some(format!("{}.{field}", indexed_expression_var_name(base)?)) + } + _ => None, + } +} + +fn subscript_suffix(subscripts: &[rumoca_core::Subscript]) -> Option { + let mut suffix = String::new(); + for subscript in subscripts { + let rumoca_core::Subscript::Index { value, .. } = subscript else { + return None; + }; + suffix.push('['); + suffix.push_str(&value.to_string()); + suffix.push(']'); + } + Some(suffix) +} + +fn forwarded_value_ref(expr: &rumoca_core::Expression) -> Option<&rumoca_core::Reference> { + match expr { + rumoca_core::Expression::VarRef { name, .. } => Some(name), + rumoca_core::Expression::FieldAccess { base, .. } => { + let rumoca_core::Expression::VarRef { name, .. } = base.as_ref() else { + return None; + }; + Some(name) + } + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Fill, + args, + .. + } => { + let rumoca_core::Expression::VarRef { name, .. } = args.first()? else { + return None; + }; + Some(name) + } + _ => None, + } +} + +pub(super) fn is_absent_lhs_component_alias( + eq: &dae::Equation, + continuous_unknowns: &BalanceSymbolSet, + input_names: &BalanceSymbolSet, +) -> bool { + let Some(lhs) = eq_lhs_var_ref(&eq.rhs) else { + return false; + }; + !continuous_unknowns.names.contains(lhs.var_name()) + && !input_names.names.contains(lhs.var_name()) +} + +pub(super) fn is_non_constraining_binding_alias( + eq: &dae::Equation, + continuous_unknowns: &BalanceSymbolSet, + _output_names: &BalanceSymbolSet, + component_defined_targets: &BalanceSymbolSet, +) -> bool { + let refs = eq_binary_var_refs(&eq.rhs); + let [lhs, rhs] = refs.as_slice() else { + return false; + }; + let lhs_is_continuous_unknown = continuous_unknowns.matches_reference(lhs); + let rhs_is_continuous_unknown = continuous_unknowns.matches_reference(rhs); + let lhs_is_component_defined = component_defined_targets.matches_reference(lhs); + if !lhs_is_continuous_unknown { + return true; + } + if !rhs_is_continuous_unknown { + return lhs_is_component_defined; + } + lhs_is_component_defined && component_defined_targets.matches_reference(rhs) +} + +pub(super) fn is_input_forwarding_connection_alias( + eq: &dae::Equation, + continuous_unknowns: &BalanceSymbolSet, + input_names: &BalanceSymbolSet, +) -> bool { + let refs = eq_binary_var_refs(&eq.rhs); + let [lhs, rhs] = refs.as_slice() else { + return false; + }; + let lhs_is_input = input_names.matches_reference(lhs); + let rhs_is_input = input_names.matches_reference(rhs); + let lhs_is_continuous_unknown = continuous_unknowns.matches_reference(lhs); + let rhs_is_continuous_unknown = continuous_unknowns.matches_reference(rhs); + (lhs_is_input && !rhs_is_continuous_unknown) || (rhs_is_input && !lhs_is_continuous_unknown) +} + +fn eq_lhs_var_ref(expr: &rumoca_core::Expression) -> Option<&rumoca_core::Reference> { + let rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs, + .. + } = expr + else { + return None; + }; + let rumoca_core::Expression::VarRef { name, .. } = lhs.as_ref() else { + return None; + }; + Some(name) +} + +fn eq_binary_lhs_rhs( + expr: &rumoca_core::Expression, +) -> Option<(&rumoca_core::Expression, &rumoca_core::Expression)> { + let rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs, + rhs, + .. + } = expr + else { + return None; + }; + Some((lhs, rhs)) +} + +#[cfg(test)] +mod tests { + use super::super::{BalanceResult, balance}; + use rumoca_core::Span; + use rumoca_ir_dae as dae; + + fn test_span() -> Span { + Span::from_offsets( + rumoca_core::SourceId::from_source_name("balance_alias_fixture.mo"), + 1, + 2, + ) + } + + fn var_ref(name: &str) -> rumoca_core::Expression { + rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new(name).into(), + subscripts: vec![], + span: test_span(), + } + } + + fn binary_eq(lhs_name: &str, rhs_name: &str, origin: &str) -> dae::Equation { + dae::Equation { + lhs: None, + rhs: rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(var_ref(lhs_name)), + rhs: Box::new(var_ref(rhs_name)), + span: test_span(), + }, + span: test_span(), + origin: origin.to_string(), + scalar_count: 1, + } + } + + fn vector_binary_eq( + lhs_name: &str, + rhs_name: &str, + origin: &str, + count: usize, + ) -> dae::Equation { + dae::Equation { + scalar_count: count, + ..binary_eq(lhs_name, rhs_name, origin) + } + } + + fn vector_fill_eq(lhs_name: &str, rhs_name: &str, origin: &str, count: usize) -> dae::Equation { + dae::Equation { + lhs: None, + rhs: rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(var_ref(lhs_name)), + rhs: Box::new(rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Fill, + args: vec![ + var_ref(rhs_name), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(count as i64), + span: test_span(), + }, + ], + span: test_span(), + }), + span: test_span(), + }, + span: test_span(), + origin: origin.to_string(), + scalar_count: count, + } + } + + fn vector_field_access_eq( + lhs_base: &str, + lhs_field: &str, + rhs_name: &str, + origin: &str, + count: usize, + ) -> dae::Equation { + dae::Equation { + lhs: None, + rhs: rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(rumoca_core::Expression::FieldAccess { + base: Box::new(var_ref(lhs_base)), + field: lhs_field.to_string(), + span: test_span(), + }), + rhs: Box::new(var_ref(rhs_name)), + span: test_span(), + }, + span: test_span(), + origin: origin.to_string(), + scalar_count: count, + } + } + + fn field_access_expr( + base: rumoca_core::Expression, + fields: &[&str], + ) -> rumoca_core::Expression { + fields + .iter() + .fold(base, |base, field| rumoca_core::Expression::FieldAccess { + base: Box::new(base), + field: (*field).to_string(), + span: test_span(), + }) + } + + fn indexed_field_access_expr( + base: &str, + index: i64, + fields: &[&str], + ) -> rumoca_core::Expression { + field_access_expr( + rumoca_core::Expression::Index { + base: Box::new(var_ref(base)), + subscripts: vec![rumoca_core::Subscript::Index { + value: index, + span: test_span(), + }], + span: test_span(), + }, + fields, + ) + } + + fn scalar_expr_eq( + lhs: rumoca_core::Expression, + rhs: rumoca_core::Expression, + origin: &str, + ) -> dae::Equation { + dae::Equation { + lhs: None, + rhs: rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span: test_span(), + }, + span: test_span(), + origin: origin.to_string(), + scalar_count: 1, + } + } + + fn scalar_input(name: &str) -> dae::Variable { + dae::Variable { + name: rumoca_core::VarName::new(name), + ..rumoca_ir_dae::Variable::empty_with_span(test_span()) + } + } + + fn algebraic_var(name: &str) -> dae::Variable { + dae::Variable { + name: rumoca_core::VarName::new(name), + ..rumoca_ir_dae::Variable::empty_with_span(test_span()) + } + } + + fn vector_algebraic_var(name: &str, size: i64) -> dae::Variable { + dae::Variable { + dims: vec![size], + ..algebraic_var(name) + } + } + + fn vector_output_var(name: &str, size: i64) -> dae::Variable { + dae::Variable { + dims: vec![size], + ..algebraic_var(name) + } + } + + fn balance_value(dae: &dae::Dae) -> BalanceResult { + balance(dae) + } + + #[test] + fn balance_skips_binding_alias_to_absent_public_alias() { + let mut dae = dae::Dae::default(); + dae.variables + .algebraics + .insert(rumoca_core::VarName::new("x"), algebraic_var("x")); + dae.continuous + .equations + .push(binary_eq("x", "source", "component equation")); + dae.continuous.equations.push(binary_eq( + "publicAlias", + "x", + "binding equation for publicAlias", + )); + + assert_eq!(balance_value(&dae).expect("valid DAE balance fixture"), 0); + } + + #[test] + fn balance_skips_binding_alias_between_component_defined_targets() { + let mut dae = dae::Dae::default(); + for name in ["a", "b"] { + dae.variables + .algebraics + .insert(rumoca_core::VarName::new(name), algebraic_var(name)); + dae.continuous + .equations + .push(binary_eq(name, "source", "component equation")); + } + dae.continuous + .equations + .push(binary_eq("a", "b", "binding equation for a")); + + assert_eq!(balance_value(&dae).expect("valid DAE balance fixture"), 0); + } + + #[test] + fn balance_counts_binding_from_parameter_to_undefined_unknown() { + let mut dae = dae::Dae::default(); + dae.variables + .algebraics + .insert(rumoca_core::VarName::new("secret"), algebraic_var("secret")); + dae.variables + .parameters + .insert(rumoca_core::VarName::new("k"), algebraic_var("k")); + dae.continuous + .equations + .push(binary_eq("secret", "k", "binding equation for secret")); + + assert_eq!(balance_value(&dae).expect("valid DAE balance fixture"), 0); + } + + #[test] + fn balance_skips_binding_alias_from_output_projection_when_lhs_is_defined() { + let mut dae = dae::Dae::default(); + dae.variables.outputs.insert( + rumoca_core::VarName::new("driver"), + vector_output_var("driver", 3), + ); + dae.variables.algebraics.insert( + rumoca_core::VarName::new("alias"), + vector_algebraic_var("alias", 3), + ); + + dae.continuous.equations.push(vector_binary_eq( + "driver", + "source", + "component equation", + 3, + )); + dae.continuous.equations.push(vector_binary_eq( + "alias", + "external", + "component equation", + 3, + )); + dae.continuous.equations.push(vector_binary_eq( + "alias", + "driver", + "binding equation for alias", + 3, + )); + + assert_eq!(balance_value(&dae).expect("valid DAE balance fixture"), 0); + } + + #[test] + fn balance_skips_vector_binding_alias_to_public_output_projection() { + let mut dae = dae::Dae::default(); + dae.variables.outputs.insert( + rumoca_core::VarName::new("public"), + vector_output_var("public", 3), + ); + dae.variables.algebraics.insert( + rumoca_core::VarName::new("internal"), + vector_algebraic_var("internal", 3), + ); + + dae.continuous.equations.push(vector_binary_eq( + "internal", + "source", + "component equation", + 3, + )); + dae.continuous.equations.push(vector_binary_eq( + "public", + "external", + "component equation", + 3, + )); + dae.continuous.equations.push(vector_binary_eq( + "public", + "internal", + "binding equation for public", + 3, + )); + + assert_eq!(balance_value(&dae).expect("valid DAE balance fixture"), 0); + } + + #[test] + fn balance_skips_internal_component_input_forwarding_without_continuous_unknown() { + let mut dae = dae::Dae::default(); + dae.variables.inputs.insert( + rumoca_core::VarName::new("replicator.u"), + scalar_input("replicator.u"), + ); + dae.continuous.equations.push(binary_eq( + "replicator.y", + "replicator.u", + "equation from replicator", + )); + + assert_eq!(balance_value(&dae).expect("valid DAE balance fixture"), 0); + } + + #[test] + fn balance_skips_connection_input_forwarding_to_missing_alias() { + let mut dae = dae::Dae::default(); + dae.variables + .inputs + .insert(rumoca_core::VarName::new("sink.u"), scalar_input("sink.u")); + dae.continuous.equations.push(binary_eq( + "source.y", + "sink.u", + "connection equation: source.y = sink.u", + )); + + assert_eq!(balance_value(&dae).expect("valid DAE balance fixture"), 0); + } + + #[test] + fn balance_skips_vector_forwarding_aliases() { + let mut dae = dae::Dae::default(); + for name in ["alias", "source", "filled"] { + dae.variables.algebraics.insert( + rumoca_core::VarName::new(name), + vector_algebraic_var(name, 3), + ); + } + dae.continuous.equations.push(vector_binary_eq( + "source", + "driver", + "component equation", + 3, + )); + dae.continuous.equations.push(vector_binary_eq( + "alias", + "aliasDriver", + "component equation", + 3, + )); + dae.continuous.equations.push(vector_binary_eq( + "filled", + "filledDriver", + "component equation", + 3, + )); + dae.continuous.equations.push(vector_binary_eq( + "alias", + "source", + "binding equation for alias", + 3, + )); + dae.continuous.equations.push(vector_fill_eq( + "filled", + "source", + "equation from replicator", + 3, + )); + + assert_eq!(balance_value(&dae).expect("valid DAE balance fixture"), 0); + } + + #[test] + fn balance_skips_vector_forwarding_from_output_projection() { + let mut dae = dae::Dae::default(); + dae.variables.outputs.insert( + rumoca_core::VarName::new("driver"), + vector_output_var("driver", 3), + ); + dae.variables.algebraics.insert( + rumoca_core::VarName::new("alias"), + vector_algebraic_var("alias", 3), + ); + + dae.continuous.equations.push(vector_binary_eq( + "driver", + "source", + "component equation", + 3, + )); + dae.continuous.equations.push(vector_binary_eq( + "alias", + "external", + "component equation", + 3, + )); + dae.continuous.equations.push(vector_binary_eq( + "alias", + "driver", + "equation from adaptor", + 3, + )); + + assert_eq!(balance_value(&dae).expect("valid DAE balance fixture"), 0); + } + + #[test] + fn balance_skips_component_defined_output_forwarding_from_input() { + let mut dae = dae::Dae::default(); + dae.variables + .outputs + .insert(rumoca_core::VarName::new("y"), vector_output_var("y", 3)); + dae.variables + .inputs + .insert(rumoca_core::VarName::new("u"), vector_algebraic_var("u", 3)); + + dae.continuous + .equations + .push(vector_binary_eq("y", "source", "component equation", 3)); + dae.continuous + .equations + .push(vector_binary_eq("y", "u", "equation from replicator", 3)); + + assert_eq!(balance_value(&dae).expect("valid DAE balance fixture"), 0); + } + + #[test] + fn balance_skips_component_defined_field_projection_from_input() { + let mut dae = dae::Dae::default(); + dae.variables.algebraics.insert( + rumoca_core::VarName::new("pin.v"), + vector_algebraic_var("pin.v", 3), + ); + dae.variables + .inputs + .insert(rumoca_core::VarName::new("u"), vector_algebraic_var("u", 3)); + + dae.continuous + .equations + .push(vector_binary_eq("pin.v", "source", "component equation", 3)); + dae.continuous.equations.push(vector_field_access_eq( + "pin", + "v", + "u", + "equation from plugToPin", + 3, + )); + + assert_eq!(balance_value(&dae).expect("valid DAE balance fixture"), 0); + } + + #[test] + fn balance_skips_component_defined_array_field_projection_from_input() { + let mut dae = dae::Dae::default(); + for name in ["pin[1].v.re", "pin[1].v.im"] { + dae.variables + .algebraics + .insert(rumoca_core::VarName::new(name), algebraic_var(name)); + } + dae.variables + .inputs + .insert(rumoca_core::VarName::new("u"), vector_algebraic_var("u", 2)); + + dae.continuous + .equations + .push(binary_eq("pin[1].v.re", "source_re", "component equation")); + dae.continuous + .equations + .push(binary_eq("pin[1].v.im", "source_im", "component equation")); + dae.continuous.equations.push(vector_field_access_eq( + "pin", + "v", + "u", + "equation from plugToPin", + 2, + )); + + assert_eq!(balance_value(&dae).expect("valid DAE balance fixture"), 0); + } + + #[test] + fn balance_counts_output_projection_when_lhs_is_not_component_defined() { + let mut dae = dae::Dae::default(); + dae.variables + .outputs + .insert(rumoca_core::VarName::new("y"), vector_output_var("y", 3)); + dae.variables + .algebraics + .insert(rumoca_core::VarName::new("x"), vector_algebraic_var("x", 3)); + + dae.continuous + .equations + .push(vector_binary_eq("x", "source", "component equation", 3)); + dae.continuous + .equations + .push(vector_binary_eq("y", "x", "equation from plant", 3)); + + assert_eq!(balance_value(&dae).expect("valid DAE balance fixture"), 0); + } + + #[test] + fn balance_skips_unresolved_indexed_reference_field_alias() { + let mut dae = dae::Dae::default(); + dae.variables.algebraics.insert( + rumoca_core::VarName::new("sensor.sensor[1].pin_p.reference.gamma"), + algebraic_var("sensor.sensor[1].pin_p.reference.gamma"), + ); + dae.continuous.equations.push(binary_eq( + "sensor.sensor[1].pin_p.reference.gamma", + "driver", + "component equation", + )); + dae.continuous.equations.push(scalar_expr_eq( + indexed_field_access_expr("sensor", 1, &["pin_p", "reference", "gamma"]), + indexed_field_access_expr("sensor", 1, &["pin_n", "reference", "gamma"]), + "equation from sensor.sensor[1]", + )); + + assert_eq!(balance_value(&dae).expect("valid DAE balance fixture"), 0); + } + + #[test] + fn balance_skips_rooted_two_pin_reference_anchor_aliases() { + let mut dae = dae::Dae::default(); + dae.variables.algebraics.insert( + rumoca_core::VarName::new("source.plug_p.reference.gamma"), + algebraic_var("source.plug_p.reference.gamma"), + ); + dae.continuous.equations.push(binary_eq( + "source.plug_p.reference.gamma", + "driver", + "component equation", + )); + dae.continuous.equations.push(binary_eq( + "source.plug_p.reference.gamma", + "source.gamma", + "equation from source", + )); + dae.continuous.equations.push(binary_eq( + "source.plug_p.reference.gamma", + "source.plug_n.reference.gamma", + "equation from source", + )); + + assert_eq!(balance_value(&dae).expect("valid DAE balance fixture"), 0); + } + + #[test] + fn balance_counts_vector_pass_through_from_non_output_driver() { + let mut dae = dae::Dae::default(); + dae.variables + .outputs + .insert(rumoca_core::VarName::new("y"), vector_output_var("y", 3)); + dae.variables + .algebraics + .insert(rumoca_core::VarName::new("u"), vector_algebraic_var("u", 3)); + + dae.continuous + .equations + .push(vector_binary_eq("u", "source", "component equation", 3)); + dae.continuous + .equations + .push(vector_binary_eq("y", "u", "equation from passThrough", 3)); + + assert_eq!(balance_value(&dae).expect("valid DAE balance fixture"), 0); + } + + #[test] + fn balance_skips_connection_alias_between_component_defined_vectors() { + let mut dae = dae::Dae::default(); + for name in ["sensor.v", "adaptor.v"] { + dae.variables.algebraics.insert( + rumoca_core::VarName::new(name), + vector_algebraic_var(name, 3), + ); + } + + dae.continuous.equations.push(vector_binary_eq( + "sensor.v", + "sensorSource", + "component equation", + 3, + )); + dae.continuous.equations.push(vector_binary_eq( + "adaptor.v", + "adaptorSource", + "component equation", + 3, + )); + dae.continuous.equations.push(vector_binary_eq( + "sensor.v", + "adaptor.v", + "connection equation: sensor.v = adaptor.v", + 3, + )); + + assert_eq!(balance_value(&dae).expect("valid DAE balance fixture"), 0); + } + + #[test] + fn component_defined_targets_include_second_binary_ref_when_it_is_the_only_unknown() { + let mut dae = dae::Dae::default(); + dae.variables + .algebraics + .insert(rumoca_core::VarName::new("y"), algebraic_var("y")); + dae.variables + .inputs + .insert(rumoca_core::VarName::new("u"), scalar_input("u")); + dae.variables.inputs.insert( + rumoca_core::VarName::new("external"), + scalar_input("external"), + ); + + dae.continuous + .equations + .push(binary_eq("u", "y", "component equation")); + dae.continuous + .equations + .push(binary_eq("y", "external", "connect(y, external)")); + + assert_eq!(balance_value(&dae).expect("valid DAE balance fixture"), 0); + } + + #[test] + fn component_defined_targets_do_not_treat_two_unknown_residual_as_two_definitions() { + let mut dae = dae::Dae::default(); + for name in ["a", "b"] { + dae.variables + .algebraics + .insert(rumoca_core::VarName::new(name), algebraic_var(name)); + } + dae.variables.inputs.insert( + rumoca_core::VarName::new("external"), + scalar_input("external"), + ); + + dae.continuous + .equations + .push(binary_eq("a", "b", "component equation")); + dae.continuous + .equations + .push(binary_eq("b", "external", "connect(b, external)")); + + assert_eq!(balance_value(&dae).expect("valid DAE balance fixture"), 0); + } +} diff --git a/crates/rumoca-phase-dae/src/balance_initial_closure.rs b/crates/rumoca-phase-dae/src/balance_initial_closure.rs new file mode 100644 index 000000000..e1c57f250 --- /dev/null +++ b/crates/rumoca-phase-dae/src/balance_initial_closure.rs @@ -0,0 +1,33 @@ +/// Balance detail after applying initialization-only deficit closure. +/// +/// Modelica initialization has its own equation system: fixed starts and +/// initial equations constrain otherwise free initial unknowns without making +/// an overdetermined simulation DAE acceptable. This detail is therefore +/// deficit-only: initial equations can close missing scalar equations, but they +/// never mask surplus equations. +#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +pub struct InitialClosureBalanceDetail { + pub scalar_equations: usize, + pub scalar_unknowns: usize, + pub deficit_before: i64, + pub overconstrained_root_gauge_scalars: i64, + pub overconstrained_break_edge_scalars: i64, + pub initial_equation_scalars: i64, + pub initial_algorithm_scalars: i64, + pub closure_used: i64, + pub deficit_after: i64, +} + +impl InitialClosureBalanceDetail { + pub fn scalar_equations_with_closure(&self) -> usize { + self.scalar_equations + self.closure_used as usize + } + + pub fn balance_with_closure(&self) -> i64 { + self.scalar_equations_with_closure() as i64 - self.scalar_unknowns as i64 + } + + pub fn is_admissible(&self) -> bool { + self.balance_with_closure() == 0 + } +} diff --git a/crates/rumoca-phase-dae/src/binding_conversion.rs b/crates/rumoca-phase-dae/src/binding_conversion.rs index ea5f1038e..0d3d1b09e 100644 --- a/crates/rumoca-phase-dae/src/binding_conversion.rs +++ b/crates/rumoca-phase-dae/src/binding_conversion.rs @@ -8,8 +8,8 @@ use rustc_hash::FxHashMap; use std::collections::HashSet; use crate::{ - ToDaeError, classification, flat_to_dae_expression_with_refs, flat_to_dae_var_name, - path_utils::strip_all_subscripts, + ToDaeError, analysis::variable_analysis, classification, flat_to_dae_expression_with_refs, + flat_to_dae_var_name, path_utils::strip_all_subscripts, }; type Dae = dae::Dae; @@ -67,8 +67,8 @@ pub(super) fn convert_bindings_to_equations( .collect(); let unknown_prefix_children = build_unknown_prefix_children(&unknowns)?; let internal_inputs = super::InternalInputIndex::new(flat)?; - let connected_inputs_only_connected_to_inputs = - super::find_connected_inputs_only_connected_to_inputs(flat, &internal_inputs); + let connected_input_binding_anchors = + super::find_connected_input_binding_anchors(flat, &internal_inputs); // Build a map from flat equation LHS to the binding expression it exactly // shadows. This is an origin-selection rule at the conversion boundary, not @@ -83,6 +83,10 @@ pub(super) fn convert_bindings_to_equations( let defined_by_unknown_rhs = collect_vars_with_unknown_rhs(flat, &unknowns); for (name, var) in &flat.variables { + if variable_analysis::is_external_constructor_handle(flat, name, var) { + continue; + } + if !var.is_primitive && prefix_children.contains_key(name.as_str()) { continue; } @@ -103,12 +107,8 @@ pub(super) fn convert_bindings_to_equations( // Connected input-only alias sets (MLS §9.1) must preserve their value anchor // from declaration bindings; otherwise the system can become underdetermined. - let keep_connected_input_binding = should_keep_connected_input_binding( - &kind, - name, - var, - &connected_inputs_only_connected_to_inputs, - ); + let keep_connected_input_binding = + should_keep_connected_input_binding(&kind, name, var, &connected_input_binding_anchors); if should_skip_variable_binding(&kind, name, connected_inputs) && !keep_connected_input_binding @@ -563,6 +563,9 @@ fn collect_unknowns(flat: &Model, state_vars: &IndexSet) -> HashSet { span, } => { let suppressed = self.suppress_events - || matches!(function, rumoca_core::BuiltinFunction::NoEvent); + || matches!( + function, + rumoca_core::BuiltinFunction::NoEvent + | rumoca_core::BuiltinFunction::Smooth + ); let mut arg_rewriter = ConditionRewriter { relations: self.relations, relation_spans: self.relation_spans, @@ -742,8 +746,11 @@ impl ExpressionVisitor for ConditionCandidateCollector<'_> { function: &rumoca_core::BuiltinFunction, args: &[rumoca_core::Expression], ) { - let suppressed = - self.suppress_events || matches!(function, rumoca_core::BuiltinFunction::NoEvent); + let suppressed = self.suppress_events + || matches!( + function, + rumoca_core::BuiltinFunction::NoEvent | rumoca_core::BuiltinFunction::Smooth + ); if !suppressed && matches!( function, diff --git a/crates/rumoca-phase-dae/src/constructor_field_selection.rs b/crates/rumoca-phase-dae/src/constructor_field_selection.rs new file mode 100644 index 000000000..a43fa2fa0 --- /dev/null +++ b/crates/rumoca-phase-dae/src/constructor_field_selection.rs @@ -0,0 +1,77 @@ +pub(crate) fn positional_constructor_arg_for_field<'a>( + args: &'a [rumoca_core::Expression], + field: &str, +) -> Option<&'a rumoca_core::Expression> { + args.iter() + .filter(|arg| !is_named_arg_marker(arg)) + .find(|arg| expression_leaf_name(arg) == Some(field)) +} + +fn is_named_arg_marker(arg: &rumoca_core::Expression) -> bool { + matches!( + arg, + rumoca_core::Expression::FunctionCall { + name, + is_constructor: true, + .. + } if name + .as_str() + .starts_with(rumoca_core::NAMED_FUNCTION_ARG_PREFIX) + ) +} + +fn expression_leaf_name(expr: &rumoca_core::Expression) -> Option<&str> { + match expr { + rumoca_core::Expression::VarRef { name, .. } => name + .component_ref() + .and_then(|component_ref| component_ref.parts.last()) + .map(|part| part.ident.as_str()), + rumoca_core::Expression::Index { base, .. } => expression_leaf_name(base), + rumoca_core::Expression::FieldAccess { field, .. } => Some(field.as_str()), + _ => None, + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn component_ref_expr(parts: &[&str]) -> rumoca_core::Expression { + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference( + parts.join("."), + rumoca_core::ComponentReference { + local: false, + span: rumoca_core::Span::DUMMY, + parts: parts + .iter() + .map(|part| rumoca_core::ComponentRefPart { + ident: (*part).to_string(), + span: rumoca_core::Span::DUMMY, + subs: Vec::new(), + }) + .collect(), + def_id: None, + }, + ), + subscripts: vec![], + span: rumoca_core::Span::DUMMY, + } + } + + #[test] + fn selects_positional_constructor_arg_by_structured_leaf_name() { + let args = vec![ + component_ref_expr(&["pCur1", "V_flow"]), + component_ref_expr(&["pCur1", "dp"]), + ]; + + let Some(rumoca_core::Expression::VarRef { name, .. }) = + positional_constructor_arg_for_field(&args, "V_flow") + else { + panic!("expected V_flow field proxy"); + }; + + assert_eq!(name.as_str(), "pCur1.V_flow"); + } +} diff --git a/crates/rumoca-phase-dae/src/convert.rs b/crates/rumoca-phase-dae/src/convert.rs index 0585e3e81..b87383dc3 100644 --- a/crates/rumoca-phase-dae/src/convert.rs +++ b/crates/rumoca-phase-dae/src/convert.rs @@ -1,3 +1,8 @@ +//! SPEC_0021 file-size exception: Flat-to-DAE conversion still owns reference +//! metadata rehydration, equation conversion, and discrete update rewrites in +//! one module while DAE admission is stabilizing. split plan: move reference +//! metadata resolution and discrete update rewriting into focused submodules. + use crate::errors::ToDaeError; use indexmap::{IndexMap, IndexSet}; use rumoca_core::ExpressionRewriter; @@ -29,6 +34,7 @@ pub fn attach_dae_reference_metadata(dae: &mut dae::Dae) -> Result<(), ToDaeErro let mut rewriter = DaeMetadataReferenceRewriter { scope, local_names: Vec::new(), + active_component_scope: None, error: None, }; rewriter.rewrite_equations(&mut dae.continuous.equations); @@ -587,8 +593,21 @@ impl DaeReferenceRewriter<'_> { field: &str, span: rumoca_core::Span, ) -> Result { + let base = self.rewrite_expression(base)?; + if let rumoca_core::Expression::FunctionCall { + args, + is_constructor: true, + .. + } = &base + && let Some(value) = + crate::constructor_field_selection::positional_constructor_arg_for_field( + args, field, + ) + { + return self.rewrite_expression(&value.clone().with_span(span)); + } Ok(rumoca_core::Expression::FieldAccess { - base: Box::new(self.rewrite_expression(base)?), + base: Box::new(base), field: field.to_owned(), span, }) @@ -765,6 +784,7 @@ struct DaeReferenceMetadata { struct DaeReferenceScope { variables: IndexMap, + variables_by_def_id: IndexMap, aggregate_prefixes: IndexMap, enum_literal_ordinals: IndexMap, } @@ -772,7 +792,10 @@ struct DaeReferenceScope { impl DaeReferenceScope { fn new(dae: &dae::Dae) -> Self { let mut variables = IndexMap::new(); + let mut variables_by_def_id = IndexMap::new(); let mut aggregate_prefixes = IndexMap::new(); + let mut aggregate_prefix_def_ids = IndexMap::new(); + let mut next_aggregate_def_id = next_component_ref_def_id(dae); Self::insert_partition(&mut variables, &dae.variables.states); Self::insert_partition(&mut variables, &dae.variables.algebraics); Self::insert_partition(&mut variables, &dae.variables.inputs); @@ -781,11 +804,28 @@ impl DaeReferenceScope { Self::insert_partition(&mut variables, &dae.variables.constants); Self::insert_partition(&mut variables, &dae.variables.discrete_reals); Self::insert_partition(&mut variables, &dae.variables.discrete_valued); + for (name, metadata) in &variables { + if let Some(def_id) = metadata + .component_ref + .as_ref() + .and_then(|component_ref| component_ref.def_id) + { + variables_by_def_id + .entry(def_id) + .or_insert_with(|| (name.clone(), metadata.clone())); + } + } for metadata in variables.values() { - Self::insert_aggregate_prefixes(&mut aggregate_prefixes, metadata); + Self::insert_aggregate_prefixes( + &mut aggregate_prefixes, + &mut aggregate_prefix_def_ids, + &mut next_aggregate_def_id, + metadata, + ); } Self { variables, + variables_by_def_id, aggregate_prefixes, enum_literal_ordinals: dae.symbols.enum_literal_ordinals.clone(), } @@ -809,6 +849,8 @@ impl DaeReferenceScope { fn insert_aggregate_prefixes( aggregate_prefixes: &mut IndexMap, + aggregate_prefix_def_ids: &mut IndexMap, + next_def_id: &mut u32, metadata: &DaeReferenceMetadata, ) { let Some(component_ref) = metadata.component_ref.as_ref() else { @@ -826,17 +868,94 @@ impl DaeReferenceScope { local: component_ref.local, span: component_ref.span, parts: prefix_parts.to_vec(), - def_id: None, + def_id: Some(Self::aggregate_prefix_def_id( + aggregate_prefix_def_ids, + next_def_id, + component_ref.local, + component_ref.span, + prefix_parts, + )), }; let key = prefix.to_var_name(); aggregate_prefixes .entry(key) .or_insert_with(|| DaeReferenceMetadata { component_ref: Some(prefix), - origin: metadata.origin, + origin: dae::VariableOrigin::Generated, source_span: metadata.source_span, }); + if prefix_parts.iter().any(|part| !part.subs.is_empty()) { + let array_prefix = Self::array_prefix_without_subscripts( + component_ref.local, + component_ref.span, + prefix_parts, + aggregate_prefix_def_ids, + next_def_id, + ); + let key = array_prefix.to_var_name(); + aggregate_prefixes + .entry(key) + .or_insert_with(|| DaeReferenceMetadata { + component_ref: Some(array_prefix), + origin: dae::VariableOrigin::Generated, + source_span: metadata.source_span, + }); + } + } + } + + fn array_prefix_without_subscripts( + local: bool, + span: rumoca_core::Span, + prefix_parts: &[rumoca_core::ComponentRefPart], + aggregate_prefix_def_ids: &mut IndexMap, + next_def_id: &mut u32, + ) -> rumoca_core::ComponentReference { + let mut array_prefix = rumoca_core::ComponentReference { + local, + span, + parts: prefix_parts.to_vec(), + def_id: None, + }; + for part in &mut array_prefix.parts { + part.subs.clear(); + } + let key = array_prefix.to_var_name(); + array_prefix.def_id = Some(Self::aggregate_prefix_def_id_for_key( + aggregate_prefix_def_ids, + next_def_id, + key, + )); + array_prefix + } + + fn aggregate_prefix_def_id( + aggregate_prefix_def_ids: &mut IndexMap, + next_def_id: &mut u32, + local: bool, + span: rumoca_core::Span, + prefix_parts: &[rumoca_core::ComponentRefPart], + ) -> rumoca_core::DefId { + let key = rumoca_core::ComponentReference { + local, + span, + parts: prefix_parts.to_vec(), + def_id: None, } + .to_var_name(); + Self::aggregate_prefix_def_id_for_key(aggregate_prefix_def_ids, next_def_id, key) + } + + fn aggregate_prefix_def_id_for_key( + aggregate_prefix_def_ids: &mut IndexMap, + next_def_id: &mut u32, + key: rumoca_core::VarName, + ) -> rumoca_core::DefId { + *aggregate_prefix_def_ids.entry(key).or_insert_with(|| { + let def_id = rumoca_core::DefId::new(*next_def_id); + *next_def_id = next_def_id.saturating_add(1); + def_id + }) } fn reference_for( @@ -866,6 +985,14 @@ impl DaeReferenceScope { if let Some(enriched) = self.enrich_element_reference(name, span)? { return Ok(enriched); } + if let Some(def_id) = name.target_def_id() + && let Some((target, metadata)) = self.variables_by_def_id.get(&def_id) + { + return self.reference_from_metadata(target, target.as_str(), metadata, span); + } + if let Some(reference) = self.overexpanded_record_reference(name, span)? { + return Ok(reference); + } if name.has_structure() { return Ok(name.clone()); } @@ -878,6 +1005,24 @@ impl DaeReferenceScope { }) } + fn reference_for_component_local( + &self, + component_scope: &str, + name: &rumoca_core::Reference, + span: rumoca_core::Span, + ) -> Option> { + if name.is_generated() + || name.as_str() == "time" + || name.as_str().contains('.') + || self.enum_literal_ordinals.contains_key(name.as_str()) + { + return None; + } + let key = rumoca_core::VarName::new(format!("{component_scope}.{}", name.as_str())); + let metadata = self.variables.get(&key)?; + Some(self.reference_from_metadata(&key, key.as_str(), metadata, span)) + } + /// Enrich an element reference whose base is a declared aggregate /// variable: the base is derived from the structured parts (no name /// parsing), its metadata supplies the def-ids, and the element's own @@ -932,7 +1077,15 @@ impl DaeReferenceScope { span: rumoca_core::Span, ) -> Result { match metadata.origin { - dae::VariableOrigin::Generated => Ok(rumoca_core::Reference::generated(rendered)), + dae::VariableOrigin::Generated => { + if let Some(component_ref) = metadata.component_ref.clone() { + return Ok(rumoca_core::Reference::with_component_reference( + rendered, + component_ref, + )); + } + Ok(rumoca_core::Reference::generated(rendered)) + } dae::VariableOrigin::Source => { let component_ref = metadata.component_ref.clone().ok_or_else(|| { @@ -950,11 +1103,191 @@ impl DaeReferenceScope { } } } + + fn overexpanded_record_reference( + &self, + name: &rumoca_core::Reference, + span: rumoca_core::Span, + ) -> Result, ToDaeError> { + let Some(component_ref) = name.component_ref() else { + return Ok(None); + }; + let mut current = component_ref.clone(); + while current.parts.len() > 1 { + let field = current.parts.last().map(|part| part.ident.as_str()); + let prefix_ends_with_field = current + .parts + .get(current.parts.len() - 2) + .zip(field) + .is_some_and(|(prefix_leaf, field)| prefix_leaf.ident == field); + if !prefix_ends_with_field { + return self.penultimate_field_reference(¤t, span); + } + current.parts.pop(); + current.def_id = None; + if let Some(reference) = self.reference_for_known_component_ref(¤t, span)? { + return Ok(Some(reference)); + } + if let Some(reference) = self.penultimate_field_reference(¤t, span)? { + return Ok(Some(reference)); + } + } + Ok(None) + } + + fn penultimate_field_reference( + &self, + component_ref: &rumoca_core::ComponentReference, + span: rumoca_core::Span, + ) -> Result, ToDaeError> { + if component_ref.parts.len() < 3 { + return Ok(None); + } + let mut candidate = component_ref.clone(); + let penultimate_idx = candidate.parts.len() - 2; + candidate.parts.remove(penultimate_idx); + candidate.def_id = None; + if let Some(reference) = self.reference_for_known_component_ref(&candidate, span)? { + return Ok(Some(reference)); + } + if candidate.parts.len() != component_ref.parts.len() { + let candidate_ref = rumoca_core::Reference::from_component_reference(candidate); + return self.overexpanded_record_reference(&candidate_ref, span); + } + Ok(None) + } + + fn reference_for_known_component_ref( + &self, + component_ref: &rumoca_core::ComponentReference, + span: rumoca_core::Span, + ) -> Result, ToDaeError> { + let target = component_ref.to_var_name(); + if let Some(reference) = self.reference_for_known_rendered_path(target.as_str(), span)? { + return Ok(Some(reference)); + } + self.indexed_field_variant_reference(component_ref, span) + } + + fn reference_for_known_rendered_path( + &self, + rendered: &str, + span: rumoca_core::Span, + ) -> Result, ToDaeError> { + let target = rumoca_core::VarName::new(rendered); + if let Some(metadata) = self.variables.get(&target) { + return self + .reference_from_metadata(&target, rendered, metadata, span) + .map(Some); + } + if let Some(metadata) = self.aggregate_prefixes.get(&target) { + return self + .reference_from_metadata(&target, rendered, metadata, span) + .map(Some); + } + let array_prefix = format!("{rendered}["); + if let Some((_, leaf_metadata)) = self + .variables + .iter() + .find(|(name, _)| name.as_str().starts_with(&array_prefix)) + { + let metadata = DaeReferenceMetadata { + component_ref: Some(rumoca_core::ComponentReference::from_flat_segments( + rendered, + leaf_metadata.source_span, + None, + )), + origin: dae::VariableOrigin::Generated, + source_span: leaf_metadata.source_span, + }; + return self + .reference_from_metadata(&target, rendered, &metadata, span) + .map(Some); + } + Ok(None) + } + + fn indexed_field_variant_reference( + &self, + component_ref: &rumoca_core::ComponentReference, + span: rumoca_core::Span, + ) -> Result, ToDaeError> { + let Some(field) = component_ref.parts.last() else { + return Ok(None); + }; + if component_ref.parts.len() < 2 { + return Ok(None); + } + let mut base_ref = component_ref.clone(); + base_ref.parts.pop(); + base_ref.def_id = None; + let base = base_ref.to_var_name(); + let indexed_base_prefix = format!("{}[", base.as_str()); + let indexed_field_suffix = format!("].{}", field.ident); + if let Some((_, leaf_metadata)) = self.variables.iter().find(|(name, _)| { + let name = name.as_str(); + name.starts_with(&indexed_base_prefix) && name.contains(&indexed_field_suffix) + }) { + let target = component_ref.to_var_name(); + let metadata = DaeReferenceMetadata { + component_ref: Some(component_ref.clone()), + origin: dae::VariableOrigin::Generated, + source_span: leaf_metadata.source_span, + }; + return self + .reference_from_metadata(&target, target.as_str(), &metadata, span) + .map(Some); + } + Ok(None) + } + + fn indexed_record_field_reference( + &self, + base: &rumoca_core::Expression, + field: &str, + span: rumoca_core::Span, + ) -> Option { + let rumoca_core::Expression::Index { + base, subscripts, .. + } = base + else { + return None; + }; + let rumoca_core::Expression::VarRef { + name, + subscripts: base_subscripts, + .. + } = base.as_ref() + else { + return None; + }; + if !base_subscripts.is_empty() { + return None; + } + let mut field_ref = name.component_ref()?.clone(); + field_ref.parts.push(rumoca_core::ComponentRefPart { + ident: field.to_string(), + span, + subs: Vec::new(), + }); + field_ref.def_id = None; + let field_name = field_ref.to_var_name(); + let metadata = self.variables.get(&field_name)?; + let reference = self + .reference_from_metadata(&field_name, field_name.as_str(), metadata, span) + .ok()?; + Some(rumoca_core::Expression::VarRef { + name: reference, + subscripts: subscripts.clone(), + span, + }) + } } struct DaeMetadataReferenceRewriter { scope: DaeReferenceScope, local_names: Vec>, + active_component_scope: Option, error: Option, } @@ -968,11 +1301,15 @@ impl DaeMetadataReferenceRewriter { fn rewrite_equations(&mut self, equations: &mut [dae::Equation]) { for equation in equations { + let previous_scope = self.active_component_scope.clone(); + self.active_component_scope = equation_origin_component_scope(&equation.origin); if let Err(error) = self.rewrite_equation_lhs(equation) { self.error = Some(error); + self.active_component_scope = previous_scope; return; } equation.rhs = self.rewrite_expression(&equation.rhs); + self.active_component_scope = previous_scope; if self.error.is_some() { return; } @@ -1086,6 +1423,27 @@ impl ExpressionRewriter for DaeMetadataReferenceRewriter { span, }; } + if let Some(reference) = self + .active_component_scope + .as_deref() + .and_then(|scope| self.scope.reference_for_component_local(scope, name, span)) + { + return match reference { + Ok(reference) => rumoca_core::Expression::VarRef { + name: reference, + subscripts: self.rewrite_subscripts(subscripts), + span, + }, + Err(error) => { + self.error = Some(error); + rumoca_core::Expression::VarRef { + name: name.clone(), + subscripts: self.rewrite_subscripts(subscripts), + span, + } + } + }; + } match self.scope.reference_for(name, span) { Ok(reference) => rumoca_core::Expression::VarRef { name: reference, @@ -1102,6 +1460,33 @@ impl ExpressionRewriter for DaeMetadataReferenceRewriter { } } } + + fn walk_field_access_expression( + &mut self, + base: &rumoca_core::Expression, + field: &str, + span: rumoca_core::Span, + ) -> rumoca_core::Expression { + let base = self.rewrite_expression(base); + if let Some(rewritten) = self + .scope + .indexed_record_field_reference(&base, field, span) + { + return rewritten; + } + rumoca_core::Expression::FieldAccess { + base: Box::new(base), + field: field.to_owned(), + span, + } + } +} + +fn equation_origin_component_scope(origin: &str) -> Option { + origin + .strip_prefix("equation from ") + .filter(|scope| !scope.is_empty()) + .map(str::to_string) } pub(crate) fn remap_flat_structured_equations( @@ -1567,6 +1952,263 @@ mod tests { assert_eq!(component_ref.parts[1].ident, "modelcard"); } + #[test] + fn dae_metadata_attachment_assigns_def_id_to_aggregate_prefix_reference() { + let field_name = rumoca_core::VarName::new("ductOut.mediums[1].T"); + let span = test_span(68, 75); + let mut field_ref = rumoca_core::component_reference_from_flat_name(&field_name, span) + .expect("test reference should parse static subscripts"); + field_ref.def_id = Some(rumoca_core::DefId::new(140)); + let aggregate_ref = + rumoca_core::ComponentReference::from_flat_segments("ductOut.mediums", span, None); + let mut dae = dae::Dae::new(); + dae.variables.algebraics.insert( + field_name.clone(), + dae::Variable { + name: field_name, + component_ref: Some(field_ref), + origin: dae::VariableOrigin::Source, + ..rumoca_ir_dae::Variable::empty_with_span(span) + }, + ); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: rumoca_core::Expression::Array { + elements: vec![ + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference( + "ductOut.mediums", + aggregate_ref.clone(), + ), + subscripts: Vec::new(), + span, + }, + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference( + "ductOut.mediums", + aggregate_ref, + ), + subscripts: Vec::new(), + span, + }, + ], + is_matrix: false, + span, + }, + span, + origin: "test".to_string(), + scalar_count: 2, + }); + + attach_dae_reference_metadata(&mut dae).expect("aggregate prefix should resolve"); + + let rumoca_core::Expression::Array { elements, .. } = &dae.continuous.equations[0].rhs + else { + panic!("expected array expression"); + }; + let def_ids = elements + .iter() + .map(|element| match element { + rumoca_core::Expression::VarRef { name, .. } => { + name.target_def_id().expect("aggregate prefix needs def-id") + } + _ => panic!("expected aggregate prefix reference"), + }) + .collect::>(); + assert_eq!(def_ids[0], def_ids[1]); + assert_ne!(def_ids[0], rumoca_core::DefId::new(140)); + } + + #[test] + fn dae_metadata_attachment_assigns_def_id_to_unsubscripted_nested_array_prefix() { + let field_name = rumoca_core::VarName::new("ductOut.mediums[1].state.X"); + let span = test_span(69, 76); + let mut field_ref = rumoca_core::component_reference_from_flat_name(&field_name, span) + .expect("test reference should parse static subscripts"); + field_ref.def_id = Some(rumoca_core::DefId::new(141)); + let nested_prefix_ref = rumoca_core::ComponentReference::from_flat_segments( + "ductOut.mediums.state", + span, + None, + ); + let mut dae = dae::Dae::new(); + dae.variables.algebraics.insert( + field_name.clone(), + dae::Variable { + name: field_name, + component_ref: Some(field_ref), + origin: dae::VariableOrigin::Source, + ..rumoca_ir_dae::Variable::empty_with_span(span) + }, + ); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference( + "ductOut.mediums.state", + nested_prefix_ref, + ), + subscripts: Vec::new(), + span, + }, + span, + origin: "test".to_string(), + scalar_count: 1, + }); + + attach_dae_reference_metadata(&mut dae) + .expect("nested unsubscripted array prefix should resolve"); + + let rumoca_core::Expression::VarRef { name, .. } = &dae.continuous.equations[0].rhs else { + panic!("expected variable reference"); + }; + let def_id = name + .target_def_id() + .expect("nested aggregate prefix needs def-id"); + assert_ne!(def_id, rumoca_core::DefId::new(141)); + } + + #[test] + fn dae_metadata_attachment_collapses_overexpanded_record_array_prefix() { + let field_name = rumoca_core::VarName::new("pipe.flowModel.states.phase[1]"); + let span = test_span(70, 77); + let field_ref = rumoca_core::ComponentReference::from_flat_segments( + field_name.as_str(), + span, + Some(rumoca_core::DefId::new(701)), + ); + let mut dae = dae::Dae::new(); + dae.variables.algebraics.insert( + field_name.clone(), + dae::Variable { + name: field_name, + component_ref: Some(field_ref), + origin: dae::VariableOrigin::Source, + ..rumoca_ir_dae::Variable::empty_with_span(span) + }, + ); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::from_component_reference( + rumoca_core::ComponentReference::from_flat_segments( + "pipe.flowModel.states.phase.phase.phase", + span, + None, + ), + ), + subscripts: Vec::new(), + span, + }, + span, + origin: "test".to_string(), + scalar_count: 1, + }); + + attach_dae_reference_metadata(&mut dae) + .expect("overexpanded record prefix should resolve to aggregate metadata"); + + let rumoca_core::Expression::VarRef { name, .. } = &dae.continuous.equations[0].rhs else { + panic!("expected variable reference"); + }; + assert_eq!(name.as_str(), "pipe.flowModel.states.phase"); + assert!(name.has_structure()); + } + + #[test] + fn dae_metadata_attachment_collapses_overexpanded_indexed_record_field_prefix() { + let field_name = rumoca_core::VarName::new("pipe.flowModel.states[1].phase"); + let span = test_span(74, 81); + let field_ref = rumoca_core::ComponentReference::from_flat_segments( + field_name.as_str(), + span, + Some(rumoca_core::DefId::new(702)), + ); + let mut dae = dae::Dae::new(); + dae.variables.algebraics.insert( + field_name.clone(), + dae::Variable { + name: field_name, + component_ref: Some(field_ref), + origin: dae::VariableOrigin::Source, + ..rumoca_ir_dae::Variable::empty_with_span(span) + }, + ); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::from_component_reference( + rumoca_core::ComponentReference::from_flat_segments( + "pipe.flowModel.states.phase.phase.phase", + span, + None, + ), + ), + subscripts: Vec::new(), + span, + }, + span, + origin: "test".to_string(), + scalar_count: 1, + }); + + attach_dae_reference_metadata(&mut dae) + .expect("overexpanded indexed record field prefix should resolve"); + + let rumoca_core::Expression::VarRef { name, .. } = &dae.continuous.equations[0].rhs else { + panic!("expected variable reference"); + }; + assert_eq!(name.as_str(), "pipe.flowModel.states.phase"); + assert!(name.has_structure()); + } + + #[test] + fn dae_metadata_attachment_collapses_overexpanded_record_sibling_field_prefix() { + let field_name = rumoca_core::VarName::new("pipe.flowModel.states[1].h"); + let span = test_span(76, 83); + let field_ref = rumoca_core::ComponentReference::from_flat_segments( + field_name.as_str(), + span, + Some(rumoca_core::DefId::new(703)), + ); + let mut dae = dae::Dae::new(); + dae.variables.algebraics.insert( + field_name.clone(), + dae::Variable { + name: field_name, + component_ref: Some(field_ref), + origin: dae::VariableOrigin::Source, + ..rumoca_ir_dae::Variable::empty_with_span(span) + }, + ); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::from_component_reference( + rumoca_core::ComponentReference::from_flat_segments( + "pipe.flowModel.states.phase.phase.h", + span, + None, + ), + ), + subscripts: Vec::new(), + span, + }, + span, + origin: "test".to_string(), + scalar_count: 1, + }); + + attach_dae_reference_metadata(&mut dae) + .expect("overexpanded record sibling field should resolve"); + + let rumoca_core::Expression::VarRef { name, .. } = &dae.continuous.equations[0].rhs else { + panic!("expected variable reference"); + }; + assert_eq!(name.as_str(), "pipe.flowModel.states.h"); + assert!(name.has_structure()); + } + #[test] fn dae_metadata_attachment_recovers_subscripted_array_prefix() { let field_name = rumoca_core::VarName::new("J[1,1].value"); @@ -1641,6 +2283,85 @@ mod tests { assert_eq!(component_ref.parts[0].subs.len(), 2); } + #[test] + fn dae_metadata_attachment_projects_indexed_record_field_to_leaf_variable() { + let field_name = rumoca_core::VarName::new("plant.aw.re"); + let span = test_span(90, 97); + let field_def = rumoca_core::DefId::new(123); + let field_ref = rumoca_core::ComponentReference { + local: false, + span, + parts: vec![ + rumoca_core::ComponentRefPart { + ident: "plant".to_string(), + span, + subs: Vec::new(), + }, + rumoca_core::ComponentRefPart { + ident: "aw".to_string(), + span, + subs: Vec::new(), + }, + rumoca_core::ComponentRefPart { + ident: "re".to_string(), + span, + subs: Vec::new(), + }, + ], + def_id: Some(field_def), + }; + let mut dae = dae::Dae::new(); + dae.variables.algebraics.insert( + field_name.clone(), + dae::Variable { + name: field_name, + component_ref: Some(field_ref), + origin: dae::VariableOrigin::Source, + dims: vec![3], + ..rumoca_ir_dae::Variable::empty_with_span(rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + )) + }, + ); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::Index { + base: Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("plant.aw"), + subscripts: Vec::new(), + span, + }), + subscripts: vec![rumoca_core::Subscript::Index { value: 1, span }], + span, + }), + field: "re".to_string(), + span, + }, + span, + origin: "record-array field projection".to_string(), + scalar_count: 1, + }); + + attach_dae_reference_metadata(&mut dae) + .expect("indexed record field should resolve through its leaf variable"); + + let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = &dae.continuous.equations[0].rhs + else { + panic!("expected projected leaf variable reference"); + }; + assert_eq!(name.as_str(), "plant.aw.re"); + assert_eq!(name.target_def_id(), Some(field_def)); + assert!(matches!( + subscripts.as_slice(), + [rumoca_core::Subscript::Index { value: 1, .. }] + )); + } + #[test] fn dae_metadata_attachment_does_not_invent_bare_component_prefixes() { let field_name = rumoca_core::VarName::new("mp.modelcard.VTO"); diff --git a/crates/rumoca-phase-dae/src/dae_lowering.rs b/crates/rumoca-phase-dae/src/dae_lowering.rs index f09012414..00dfcddd5 100644 --- a/crates/rumoca-phase-dae/src/dae_lowering.rs +++ b/crates/rumoca-phase-dae/src/dae_lowering.rs @@ -1,24 +1,42 @@ -//! DAE-level lowering passes for code generation. -//! -//! This module contains record function parameter decomposition, array size -//! argument insertion, parameter dependency sorting, and vector equation -//! scalarization that operate on the DAE IR before code generation. - +//! DAE pre-codegen lowering. SPEC_0021 file-size exception: split plan: move pass families into focused `dae_lowering` submodules. use crate::ToDaeError; use crate::scalar_size::compute_var_size; use indexmap::{IndexMap, IndexSet}; -use rumoca_core::{ExpressionRewriter, ExpressionVisitor, StatementRewriter}; +use rumoca_core::{ + BuiltinFunction as Builtin, Expression as Expr, ExpressionRewriter, ExpressionVisitor, + Function, OpBinary, Reference, Span, StatementRewriter, Subscript, VarName, +}; use rumoca_ir_dae as dae; use rumoca_ir_dae::DaeExpressionRewriter; use std::collections::{BTreeMap, HashMap, HashSet}; type Dae = dae::Dae; -type RecordArgMap = HashMap)>>; +type RecordArgMap = HashMap>; +type RecordArrayFieldMap = HashMap; + +#[derive(Debug, Clone)] +struct RecordArrayFieldVariants { + variants: Vec, + field_dims: Vec, +} + +struct RecordArgDecomposition { + original_index: usize, + param_name: String, + fields: Vec, +} + +mod record_field_inference; +use record_field_inference::{FieldUseMap, infer_record_fields_by_function}; +mod colon_slice_dot; +use colon_slice_dot::{ + DotOperand, classify_dot_operand, is_colon_slice, lower_colon_slice_dot_products, +}; #[derive(Default)] struct ArrayParamMap { - by_def_id: HashMap>, - by_name: HashMap>, + by_def_id: HashMap>, + by_name: HashMap>, } /// DAE value prepared for code generators that need DAE-level convenience @@ -48,6 +66,7 @@ impl CodegenDae { pub fn prepare_dae_for_codegen(dae: &dae::Dae) -> Result { let mut prepared = dae.clone(); lower_record_function_params_dae(&mut prepared)?; + unwrap_block_constructor_value_wrappers(&mut prepared); insert_array_size_args_dae(&mut prepared)?; Ok(CodegenDae { dae: prepared }) } @@ -64,6 +83,65 @@ pub fn prepare_dae_for_fmi_model_description(dae: &dae::Dae) -> Result { + functions: &'a IndexMap, +} + +impl ExpressionRewriter for BlockConstructorValueUnwrapper<'_> { + fn rewrite_expression(&mut self, expr: &rumoca_core::Expression) -> rumoca_core::Expression { + if let rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor: true, + span, + } = expr + { + let args = self.rewrite_expressions(args); + if args.len() == 1 + && self + .functions + .get(name.var_name()) + .is_some_and(is_block_constructor_value_wrapper) + { + return args.into_iter().next().expect("checked len"); + } + return rumoca_core::Expression::FunctionCall { + name: name.clone(), + args, + is_constructor: true, + span: *span, + }; + } + self.walk_expression(expr) + } +} + +impl StatementRewriter for BlockConstructorValueUnwrapper<'_> {} +impl DaeExpressionRewriter for BlockConstructorValueUnwrapper<'_> {} + +fn is_block_constructor_value_wrapper(function: &rumoca_core::Function) -> bool { + function.is_constructor + && !function.inputs.is_empty() + && function + .inputs + .iter() + .all(|input| input.type_class == Some(rumoca_core::ClassType::Connector)) +} + // ============================================================================= // Record function parameter decomposition (DAE level) // ============================================================================= @@ -111,46 +189,34 @@ pub fn lower_record_function_params_dae(dae: &mut Dae) -> Result<(), ToDaeError> }) .collect::>(); + let inferred_fields_by_function = + infer_record_fields_by_function(&dae.symbols.functions, &record_fields_by_type); + // Identify functions with record params and rewrite their signatures. - let mut decomp_map: HashMap)>> = HashMap::new(); + let mut decomp_map = RecordArgMap::new(); for (func_name, func) in dae.symbols.functions.iter_mut() { - let mut decomposed: Vec<(usize, String, Vec)> = Vec::new(); - for (idx, input) in func.inputs.iter().enumerate() { - if let Some(fields) = record_fields_by_type.get(&input.type_name) { - decomposed.push((idx, input.name.clone(), fields.clone())); - } - } - if decomposed.is_empty() { - continue; - } + let decomposed = record_inputs_to_decompose( + func_name.as_str(), + func, + &record_fields_by_type, + &inferred_fields_by_function, + ); + let inferred_function_fields = inferred_fields_by_function.get(func_name.as_str()); // Replace record inputs with scalar field inputs - let old_inputs = std::mem::take(&mut func.inputs); - for (idx, input) in old_inputs.into_iter().enumerate() { - let Some((_, param_name, fields)) = decomposed.iter().find(|(i, _, _)| *i == idx) - else { - func.inputs.push(input); - continue; - }; - for field in fields { - func.inputs.push(rumoca_core::FunctionParam { - def_id: None, - name: format!("{param_name}_{field}"), - span: input.span, - type_name: "Real".to_string(), - type_class: None, - dims: vec![], - shape_expr: Vec::new(), - default: None, - description: None, - }); - } + rewrite_record_function_inputs(func, &decomposed, inferred_function_fields); + if decomposed.is_empty() { + continue; } - let entry: Vec<(usize, Vec)> = decomposed + let entry: Vec = decomposed .iter() - .map(|(idx, _, fields)| (*idx, fields.clone())) + .map(|(idx, param_name, fields)| RecordArgDecomposition { + original_index: *idx, + param_name: param_name.clone(), + fields: fields.clone(), + }) .collect(); decomp_map.insert(func_name.as_str().to_string(), entry); } @@ -175,6 +241,153 @@ pub fn lower_record_function_params_dae(dae: &mut Dae) -> Result<(), ToDaeError> Ok(()) } +fn record_inputs_to_decompose( + func_name: &str, + func: &rumoca_core::Function, + record_fields_by_type: &HashMap>, + inferred_fields_by_function: &HashMap, +) -> Vec<(usize, String, Vec)> { + func.inputs + .iter() + .enumerate() + .filter_map(|(idx, input)| { + let fields = record_fields_for_input( + func_name, + input, + record_fields_by_type, + inferred_fields_by_function, + ); + (!fields.is_empty()).then(|| (idx, input.name.clone(), fields)) + }) + .collect() +} + +fn record_fields_for_input( + func_name: &str, + input: &rumoca_core::FunctionParam, + record_fields_by_type: &HashMap>, + inferred_fields_by_function: &HashMap, +) -> Vec { + if input.type_class == Some(rumoca_core::ClassType::Class) { + return Vec::new(); + } + record_fields_by_type + .get(&input.type_name) + .into_iter() + .flatten() + .cloned() + .chain( + inferred_fields_by_function + .get(func_name) + .and_then(|fields| fields.get(&input.name)) + .into_iter() + .flat_map(|fields| fields.keys().cloned()), + ) + .collect::>() + .into_iter() + .collect() +} + +fn rewrite_record_function_inputs( + func: &mut rumoca_core::Function, + decomposed: &[(usize, String, Vec)], + inferred_fields: Option<&FieldUseMap>, +) { + let old_inputs = std::mem::take(&mut func.inputs); + let mut seen_inputs = HashSet::::new(); + let mut flattened_prefixes = HashSet::::new(); + for (idx, input) in old_inputs.into_iter().enumerate() { + rewrite_one_record_function_input( + func, + input, + idx, + decomposed, + inferred_fields, + &mut seen_inputs, + &mut flattened_prefixes, + ); + } + append_inferred_flattened_inputs(func, inferred_fields, flattened_prefixes, &mut seen_inputs); +} + +fn rewrite_one_record_function_input( + func: &mut rumoca_core::Function, + input: rumoca_core::FunctionParam, + idx: usize, + decomposed: &[(usize, String, Vec)], + inferred_fields: Option<&FieldUseMap>, + seen_inputs: &mut HashSet, + flattened_prefixes: &mut HashSet, +) { + let Some((_, param_name, fields)) = decomposed.iter().find(|(i, _, _)| *i == idx) else { + if let Some((prefix, _)) = input.name.split_once('_') { + flattened_prefixes.insert(prefix.to_string()); + } + seen_inputs.insert(input.name.clone()); + func.inputs.push(input); + return; + }; + for field in fields { + let dims = inferred_fields + .and_then(|by_prefix| by_prefix.get(param_name)) + .and_then(|by_field| by_field.get(field)) + .cloned() + .unwrap_or_default(); + push_flat_record_input(func, format!("{param_name}_{field}"), input.span, dims); + seen_inputs.insert(format!("{param_name}_{field}")); + } +} + +fn append_inferred_flattened_inputs( + func: &mut rumoca_core::Function, + inferred_fields: Option<&FieldUseMap>, + flattened_prefixes: HashSet, + seen_inputs: &mut HashSet, +) { + let Some(inferred) = inferred_fields else { + return; + }; + for prefix in flattened_prefixes { + append_inferred_flattened_prefix(func, inferred, &prefix, seen_inputs); + } +} + +fn append_inferred_flattened_prefix( + func: &mut rumoca_core::Function, + inferred: &FieldUseMap, + prefix: &str, + seen_inputs: &mut HashSet, +) { + let Some(fields) = inferred.get(prefix) else { + return; + }; + for (field, dims) in fields { + let name = format!("{prefix}_{field}"); + if seen_inputs.insert(name.clone()) { + push_flat_record_input(func, name, func.span, dims.clone()); + } + } +} + +fn push_flat_record_input( + func: &mut rumoca_core::Function, + name: String, + span: rumoca_core::Span, + dims: Vec, +) { + func.inputs.push(rumoca_core::FunctionParam { + def_id: None, + name, + span, + type_name: "Real".to_string(), + type_class: None, + dims, + shape_expr: Vec::new(), + default: None, + description: None, + }); +} + pub(crate) fn lower_enum_literal_refs_to_ordinals(dae: &mut Dae) { if dae.symbols.enum_literal_ordinals.is_empty() { return; @@ -184,6 +397,14 @@ pub(crate) fn lower_enum_literal_refs_to_ordinals(dae: &mut Dae) { ordinals: &ordinals, } .rewrite_dae(dae); + lower_enum_literal_refs_in_structured_families( + &mut dae.continuous.structured_equations, + &ordinals, + ); + lower_enum_literal_refs_in_structured_families( + &mut dae.initialization.structured_equations, + &ordinals, + ); for function in dae.symbols.functions.values_mut() { function.body = EnumLiteralOrdinalLowerer { ordinals: &ordinals, @@ -192,6 +413,19 @@ pub(crate) fn lower_enum_literal_refs_to_ordinals(dae: &mut Dae) { } } +fn lower_enum_literal_refs_in_structured_families( + families: &mut [dae::StructuredEquationFamily], + ordinals: &IndexMap, +) { + for family in families { + let Some(template) = &mut family.template else { + continue; + }; + let mut lowerer = EnumLiteralOrdinalLowerer { ordinals }; + template.body = lowerer.rewrite_expressions(&template.body); + } +} + struct EnumLiteralOrdinalLowerer<'a> { ordinals: &'a IndexMap, } @@ -289,28 +523,71 @@ impl DaeExpressionRewriter for DaeRecordArgDecomposer<'_> {} fn decompose_dae_record_args( old_args: &[rumoca_core::Expression], - decomposed: &[(usize, Vec)], + decomposed: &[RecordArgDecomposition], call_span: rumoca_core::Span, ) -> Result, ToDaeError> { let mut args = Vec::new(); - let mut old_idx = 0; - for (param_idx, fields) in decomposed { - while old_idx < *param_idx && old_idx < old_args.len() { - args.push(old_args[old_idx].clone()); - old_idx += 1; + let positional_args = old_args + .iter() + .filter(|arg| named_function_arg_value(arg).is_none()) + .collect::>(); + let mut consumed_named_args = HashSet::new(); + let mut positional_idx = 0; + for entry in decomposed { + while positional_idx < entry.original_index && positional_idx < positional_args.len() { + args.push(positional_args[positional_idx].clone()); + positional_idx += 1; } - if old_idx < old_args.len() { - expand_dae_record_arg(&old_args[old_idx], fields, &mut args, call_span)?; - old_idx += 1; + if let Some(named_arg) = named_record_arg_value(old_args, &entry.param_name) { + consumed_named_args.insert(entry.param_name.as_str()); + expand_dae_record_arg(named_arg, &entry.fields, &mut args, call_span)?; + } else if positional_idx < positional_args.len() { + expand_dae_record_arg( + positional_args[positional_idx], + &entry.fields, + &mut args, + call_span, + )?; + positional_idx += 1; } } - while old_idx < old_args.len() { - args.push(old_args[old_idx].clone()); - old_idx += 1; + while positional_idx < positional_args.len() { + args.push(positional_args[positional_idx].clone()); + positional_idx += 1; } + args.extend(old_args.iter().filter_map(|arg| { + let (name, _) = named_function_arg_value(arg)?; + (!consumed_named_args.contains(name)).then(|| arg.clone()) + })); Ok(args) } +fn named_record_arg_value<'a>( + args: &'a [rumoca_core::Expression], + param_name: &str, +) -> Option<&'a rumoca_core::Expression> { + args.iter().find_map(|arg| { + let (name, value) = named_function_arg_value(arg)?; + (name == param_name).then_some(value) + }) +} + +fn named_function_arg_value( + arg: &rumoca_core::Expression, +) -> Option<(&str, &rumoca_core::Expression)> { + let rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor: true, + .. + } = arg + else { + return None; + }; + let param_name = name.as_str().strip_prefix("__rumoca_named_arg__.")?; + Some((param_name, args.first()?)) +} + fn named_constructor_arg_dae<'a>( ctor_args: &'a [rumoca_core::Expression], field: &str, @@ -370,9 +647,11 @@ fn expand_dae_record_arg( // Variable reference → emit field VarRefs if let rumoca_core::Expression::VarRef { name, .. } = arg { + let base_name = + record_arg_projection_base_name(name, fields).unwrap_or_else(|| name.clone()); for field in fields { out.push(rumoca_core::Expression::VarRef { - name: name.with_appended_field(field), + name: base_name.with_appended_field(field), subscripts: vec![], span: owner_span, }); @@ -380,21 +659,6 @@ fn expand_dae_record_arg( return Ok(()); } - // Check for FieldAccess on the record variable (e.g. `c.re` passed directly) - if let rumoca_core::Expression::FieldAccess { .. } = arg { - // Single field access on a record — just push the base.field expression - // This handles cases like passing `c.re` where `c` was the original record param. - out.push(arg.clone()); - // Pad remaining fields with 0.0 - for _ in 1..fields.len() { - out.push(rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(0.0), - span: owner_span, - }); - } - return Ok(()); - } - // Scalar expression passed to a record-typed parameter (e.g. Real expr passed // where Complex is expected) — treat as the first field, zero-fill the rest. // This avoids generating invalid C like `(expr).re`. @@ -420,6 +684,24 @@ fn expand_dae_record_arg( Ok(()) } +fn record_arg_projection_base_name( + name: &rumoca_core::Reference, + fields: &[String], +) -> Option { + let component_ref = name.component_ref()?; + let terminal = component_ref.last_ident()?; + if !fields.iter().any(|field| field == terminal) { + return None; + } + if component_ref.parts.len() > 1 { + let mut base_ref = component_ref.clone(); + base_ref.parts.pop(); + base_ref.def_id = None; + return Some(rumoca_core::Reference::from_component_reference(base_ref)); + } + None +} + /// Returns true if the expression is obviously a scalar (not a record type). fn is_obviously_scalar(expr: &rumoca_core::Expression) -> bool { match expr { @@ -470,12 +752,12 @@ fn is_obviously_scalar(expr: &rumoca_core::Expression) -> bool { pub fn insert_array_size_args_dae(dae: &mut Dae) -> Result<(), ToDaeError> { let mut array_param_map = ArrayParamMap::default(); for (name, func) in &dae.symbols.functions { - let indices: Vec = func + let indices: Vec<(usize, String)> = func .inputs .iter() .enumerate() .filter(|(_, p)| !p.dims.is_empty()) - .map(|(i, _)| i) + .map(|(index, parameter)| (index, parameter.name.clone())) .collect(); if indices.is_empty() { continue; @@ -575,7 +857,7 @@ impl DaeExpressionRewriter for DaeSizeArgInserter<'_> {} fn array_param_indices_for_call<'a>( map: &'a ArrayParamMap, call_name: &rumoca_core::Reference, -) -> Option<&'a Vec> { +) -> Option<&'a Vec<(usize, String)>> { call_name .component_ref() .and_then(|reference| reference.def_id) @@ -585,22 +867,40 @@ fn array_param_indices_for_call<'a>( fn insert_dae_size_args( args: &mut Vec, - array_indices: &[usize], + array_params: &[(usize, String)], call_span: rumoca_core::Span, ) -> Result<(), ToDaeError> { - for ¶m_idx in array_indices.iter().rev() { - if param_idx >= args.len() { + for (param_idx, param_name) in array_params.iter().rev() { + let Some(arg_idx) = argument_index_for_param(args, *param_idx, param_name) else { continue; - } + }; // If the argument is an Array literal with known element count, use the // literal count directly instead of size(), which some backends cannot // render for compound literals. - let size_expr = array_size_expr_for_arg(&args[param_idx], call_span)?; - args.insert(param_idx + 1, size_expr); + let size_expr = + array_size_expr_for_arg(argument_value_for_size(&args[arg_idx]), call_span)?; + args.insert(arg_idx + 1, size_expr); } Ok(()) } +fn argument_index_for_param( + args: &[rumoca_core::Expression], + param_idx: usize, + param_name: &str, +) -> Option { + args.iter() + .position(|arg| named_function_arg_value(arg).is_some_and(|(name, _)| name == param_name)) + .or((param_idx < args.len()).then_some(param_idx)) +} + +fn argument_value_for_size(arg: &rumoca_core::Expression) -> &rumoca_core::Expression { + if let Some((_, value)) = named_function_arg_value(arg) { + return value; + } + arg +} + fn array_size_expr_for_arg( arg: &rumoca_core::Expression, call_span: rumoca_core::Span, @@ -759,33 +1059,41 @@ impl ExpressionVisitor for VarRefListCollector { // Vector equation scalarization // ============================================================================= -/// Scalarize vector equations that reference "phantom" base names. -/// -/// In Modelica, connector arrays like `plug_p.pin[3]` produce scalarized -/// variables (`sineVoltage.plug_p.pin[1].v`, `…pin[2].v`, `…pin[3].v`) -/// but some component-level equations reference the unsubscripted base name -/// (`sineVoltage.plug_p.pin.v`) as a vector. These phantom base names do -/// not appear in any DAE variable map, so backends that render equations -/// directly (CasADi, SymPy, JAX) produce undefined identifiers. -/// -/// This pass detects equations with `scalar_count > 1` whose expressions -/// contain such phantom VarRefs, and expands each into `scalar_count` -/// scalar equations — one per element — with every phantom VarRef replaced -/// by its indexed variant and every declared-array VarRef subscripted. +/// Expand phantom connector-array equations, indexing phantom bases and declared-array arguments. pub fn scalarize_phantom_vector_equations(dae: &mut Dae) -> Result<(), ToDaeError> { let known_names = build_known_var_name_set(dae); let phantom_map = build_phantom_expansion_map(dae, &known_names); - let array_dims = build_array_dims_map(dae); - - canonicalize_embedded_subscript_equation_list(&mut dae.continuous.equations, &array_dims)?; - canonicalize_embedded_subscript_equation_list(&mut dae.initialization.equations, &array_dims)?; - canonicalize_embedded_subscript_equation_list(&mut dae.discrete.real_updates, &array_dims)?; - canonicalize_embedded_subscript_equation_list(&mut dae.discrete.valued_updates, &array_dims)?; - canonicalize_embedded_subscript_equation_list(&mut dae.conditions.equations, &array_dims)?; + let var_dims = build_dae_var_dims_map(dae); + let mut array_dims = var_dims.clone(); + array_dims.retain(|_, dims| !dims.is_empty()); + let record_array_projection_aliases = build_record_array_projection_alias_map(dae)?; + let record_array_fields = build_record_array_field_map(dae); - if phantom_map.is_empty() { - return Ok(()); - } + canonicalize_embedded_subscript_equation_list( + &mut dae.continuous.equations, + &array_dims, + &record_array_projection_aliases, + )?; + canonicalize_embedded_subscript_equation_list( + &mut dae.initialization.equations, + &array_dims, + &record_array_projection_aliases, + )?; + canonicalize_embedded_subscript_equation_list( + &mut dae.discrete.real_updates, + &array_dims, + &record_array_projection_aliases, + )?; + canonicalize_embedded_subscript_equation_list( + &mut dae.discrete.valued_updates, + &array_dims, + &record_array_projection_aliases, + )?; + canonicalize_embedded_subscript_equation_list( + &mut dae.conditions.equations, + &array_dims, + &record_array_projection_aliases, + )?; // Expanding an array equation into scalar rows shifts every later row, so the // partitions that carry structured families (continuous, initialization) must @@ -795,48 +1103,254 @@ pub fn scalarize_phantom_vector_equations(dae: &mut Dae) -> Result<(), ToDaeErro &mut dae.continuous.equations, &phantom_map, &array_dims, + &var_dims, + &record_array_fields, &dae.symbols.functions, + false, )?; rumoca_ir_dae::remap_structured_families_after_expansion( &mut dae.continuous.structured_equations, &continuous_spans, ); + if dae.initialization.equation_provenance.len() != dae.initialization.equations.len() { + return Err(ToDaeError::runtime_metadata_violation(format!( + "initial equation provenance cardinality {} does not match equation cardinality {} before phantom scalarization", + dae.initialization.equation_provenance.len(), + dae.initialization.equations.len() + ))); + } + let initialization_provenance = dae.initialization.equation_provenance.clone(); let initialization_spans = scalarize_equation_list( &mut dae.initialization.equations, &phantom_map, &array_dims, + &var_dims, + &record_array_fields, &dae.symbols.functions, + false, )?; rumoca_ir_dae::remap_structured_families_after_expansion( &mut dae.initialization.structured_equations, &initialization_spans, ); + dae.initialization.equation_provenance = initialization_spans + .iter() + .zip(initialization_provenance) + .flat_map(|((_, new_len), provenance)| std::iter::repeat_n(provenance, *new_len)) + .collect(); // The discrete and condition partitions carry no structured families. scalarize_equation_list( &mut dae.discrete.real_updates, &phantom_map, &array_dims, + &var_dims, + &record_array_fields, &dae.symbols.functions, + true, )?; scalarize_equation_list( &mut dae.discrete.valued_updates, &phantom_map, &array_dims, + &var_dims, + &record_array_fields, &dae.symbols.functions, + true, )?; scalarize_equation_list( &mut dae.conditions.equations, &phantom_map, &array_dims, + &var_dims, + &record_array_fields, &dae.symbols.functions, + false, )?; Ok(()) } -/// Build the set of all variable names known to the DAE. -fn build_known_var_name_set(dae: &Dae) -> HashSet { - let mut names = HashSet::new(); - for map in [ +pub(crate) fn sync_materialized_structured_equation_templates( + dae: &mut Dae, +) -> Result<(), ToDaeError> { + let var_dims = build_dae_var_dims_map(dae); + sync_structured_partition_templates( + &dae.continuous.structured_equations, + &mut dae.continuous.equations, + &var_dims, + )?; + sync_structured_partition_templates( + &dae.initialization.structured_equations, + &mut dae.initialization.equations, + &var_dims, + )?; + Ok(()) +} + +pub(crate) fn repair_external_table_event_handles(dae: &mut Dae) { + let table_ids = dae + .metadata + .nonnumeric_variable_names + .iter() + .filter(|name| name.ends_with(".tableID")) + .cloned() + .collect::>(); + repair_external_table_event_handles_in_equations(&mut dae.discrete.real_updates, &table_ids); + repair_external_table_event_handles_in_equations(&mut dae.discrete.valued_updates, &table_ids); + repair_external_table_event_handles_in_equations(&mut dae.conditions.equations, &table_ids); +} + +fn repair_external_table_event_handles_in_equations( + equations: &mut [dae::Equation], + table_ids: &HashSet, +) { + for equation in equations { + let Some(prefix) = equation + .lhs + .as_ref() + .and_then(|lhs| external_table_event_prefix(lhs.as_str())) + else { + continue; + }; + let table_id_name = format!("{prefix}.tableID"); + if !table_ids.contains(&table_id_name) { + continue; + } + equation.rhs = ExternalTableEventHandleRepair { + table_id_name: &table_id_name, + span: equation.span, + } + .rewrite_expression(&equation.rhs); + } +} + +fn external_table_event_prefix(name: &str) -> Option<&str> { + name.strip_suffix(".nextTimeEventScaled") + .or_else(|| name.strip_suffix(".nextTimeEvent")) +} + +struct ExternalTableEventHandleRepair<'a> { + table_id_name: &'a str, + span: rumoca_core::Span, +} + +impl ExpressionRewriter for ExternalTableEventHandleRepair<'_> { + fn walk_function_call_expression( + &mut self, + name: &rumoca_core::Reference, + args: &[rumoca_core::Expression], + is_constructor: bool, + span: rumoca_core::Span, + ) -> rumoca_core::Expression { + let mut args = self.rewrite_expressions(args); + if name.last_segment() == "getNextTimeEvent" + && args + .first() + .is_some_and(external_table_constructor_call_dae) + { + let arg_span = args + .first() + .and_then(rumoca_core::Expression::span) + .unwrap_or(self.span); + args[0] = rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::generated(self.table_id_name.to_string()), + subscripts: Vec::new(), + span: arg_span, + }; + } + rumoca_core::Expression::FunctionCall { + name: name.clone(), + args, + is_constructor, + span, + } + } +} + +fn external_table_constructor_call_dae(expr: &rumoca_core::Expression) -> bool { + matches!( + expr, + rumoca_core::Expression::FunctionCall { + name, + is_constructor: true, + .. + } if matches!( + name.last_segment(), + "ExternalCombiTimeTable" | "ExternalCombiTable1D" + ) + ) +} + +fn sync_structured_partition_templates( + families: &[dae::StructuredEquationFamily], + equations: &mut [dae::Equation], + var_dims: &HashMap>, +) -> Result<(), ToDaeError> { + for family in families { + if !family.interiors_materialized { + continue; + } + let Some(template) = family.template.as_ref() else { + continue; + }; + if template.body.is_empty() + || !family + .equation_counts + .iter() + .all(|count| *count == template.body.len()) + { + continue; + } + sync_structured_template_family(family, template, equations, var_dims)?; + } + Ok(()) +} + +fn sync_structured_template_family( + family: &dae::StructuredEquationFamily, + template: &rumoca_core::ComprehensionTemplate, + equations: &mut [dae::Equation], + var_dims: &HashMap>, +) -> Result<(), ToDaeError> { + let tuples = family.domain.index_tuples().map_err(|err| { + ToDaeError::runtime_metadata_violation_at( + format!("invalid structured equation domain: {err}"), + family.span, + ) + })?; + for (iteration, tuple) in tuples.iter().enumerate() { + let binder_values: HashMap<_, _> = family + .domain + .binders + .iter() + .zip(tuple.iter().copied()) + .map(|(binder, value)| (binder.display_name.clone(), value)) + .collect(); + let row_base = structured_template_row_base(family, template, iteration)?; + for (position, body) in template.body.iter().enumerate() { + let Some(equation) = equations.get_mut(row_base + position) else { + return Err(ToDaeError::runtime_metadata_violation_at( + "structured equation template points past materialized equations".to_string(), + family.span, + )); + }; + equation.rhs = StructuredBinderSubstitution { + values: &binder_values, + span: family.span, + } + .rewrite_expression(body); + if let Some(scalar_count) = + structured_template_residual_scalar_count(&equation.rhs, var_dims) + { + equation.scalar_count = scalar_count; + } + } + } + Ok(()) +} + +fn build_dae_var_dims_map(dae: &Dae) -> HashMap> { + let mut dims = HashMap::new(); + for partition in [ &dae.variables.states, &dae.variables.algebraics, &dae.variables.inputs, @@ -846,7 +1360,198 @@ fn build_known_var_name_set(dae: &Dae) -> HashSet { &dae.variables.discrete_reals, &dae.variables.discrete_valued, ] { - for name in map.keys() { + for (name, variable) in partition { + dims.insert(name.as_str().to_string(), variable.dims.clone()); + } + } + dims +} + +fn structured_template_residual_scalar_count( + expr: &rumoca_core::Expression, + var_dims: &HashMap>, +) -> Option { + let rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs, + .. + } = expr + else { + return None; + }; + expression_selected_scalar_count(lhs, var_dims) +} + +fn expression_selected_scalar_count( + expr: &rumoca_core::Expression, + var_dims: &HashMap>, +) -> Option { + match expr { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } => var_ref_selected_scalar_count(name.as_str(), subscripts, var_dims), + rumoca_core::Expression::Index { + base, subscripts, .. + } => { + let rumoca_core::Expression::VarRef { + name, + subscripts: base_subscripts, + .. + } = base.as_ref() + else { + return None; + }; + if !base_subscripts.is_empty() { + return None; + } + var_ref_selected_scalar_count(name.as_str(), subscripts, var_dims) + } + rumoca_core::Expression::BuiltinCall { function, args, .. } + if matches!(function, rumoca_core::BuiltinFunction::Der) && args.len() == 1 => + { + expression_selected_scalar_count(&args[0], var_dims) + } + _ => None, + } +} + +fn var_ref_selected_scalar_count( + name: &str, + subscripts: &[rumoca_core::Subscript], + var_dims: &HashMap>, +) -> Option { + let dims = var_dims.get(name)?; + if subscripts.is_empty() { + return Some(compute_var_size(dims)); + } + projected_dims_for_subscripts(dims, subscripts).map(|dims| compute_var_size(&dims)) +} + +fn projected_dims_for_subscripts( + dims: &[i64], + subscripts: &[rumoca_core::Subscript], +) -> Option> { + let mut remaining = Vec::new(); + let mut dim_idx = 0usize; + for subscript in subscripts { + if dim_idx >= dims.len() { + break; + } + match subscript { + rumoca_core::Subscript::Index { .. } => { + dim_idx += 1; + } + rumoca_core::Subscript::Expr { expr, .. } => { + if subscript_expr_selects_vector(expr)? { + return None; + } + dim_idx += 1; + } + rumoca_core::Subscript::Colon { .. } => { + remaining.push(dims[dim_idx]); + dim_idx += 1; + } + } + } + remaining.extend_from_slice(&dims[dim_idx..]); + Some(remaining) +} + +fn subscript_expr_selects_vector(expr: &rumoca_core::Expression) -> Option { + match expr { + rumoca_core::Expression::Literal { + value: + rumoca_core::Literal::Integer(_) + | rumoca_core::Literal::Real(_) + | rumoca_core::Literal::Boolean(_), + .. + } => Some(false), + rumoca_core::Expression::Range { .. } | rumoca_core::Expression::Array { .. } => Some(true), + _ => None, + } +} + +fn structured_template_row_base( + family: &dae::StructuredEquationFamily, + template: &rumoca_core::ComprehensionTemplate, + iteration: usize, +) -> Result { + family + .first_equation_index + .checked_add(iteration.checked_mul(template.body.len()).ok_or_else(|| { + ToDaeError::runtime_metadata_violation_at( + "structured equation row index overflows".to_string(), + family.span, + ) + })?) + .ok_or_else(|| { + ToDaeError::runtime_metadata_violation_at( + "structured equation row index overflows".to_string(), + family.span, + ) + }) +} + +struct StructuredBinderSubstitution<'a> { + values: &'a HashMap, + span: rumoca_core::Span, +} + +impl ExpressionRewriter for StructuredBinderSubstitution<'_> { + fn walk_var_ref_expression( + &mut self, + name: &rumoca_core::Reference, + subscripts: &[rumoca_core::Subscript], + span: rumoca_core::Span, + ) -> rumoca_core::Expression { + if subscripts.is_empty() + && let Some(value) = self + .values + .get(name.as_str()) + .or_else(|| self.values.get(name.last_segment())) + { + return rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(*value), + span: if span.is_dummy() { self.span } else { span }, + }; + } + let subscripts = self.rewrite_subscripts(subscripts); + if subscripts + .iter() + .any(|subscript| matches!(subscript, rumoca_core::Subscript::Colon { .. })) + { + return rumoca_core::Expression::Index { + base: Box::new(rumoca_core::Expression::VarRef { + name: name.clone(), + subscripts: Vec::new(), + span, + }), + subscripts, + span, + }; + } + rumoca_core::Expression::VarRef { + name: name.clone(), + subscripts, + span, + } + } +} + +/// Build the set of all variable names known to the DAE. +fn build_known_var_name_set(dae: &Dae) -> HashSet { + let mut names = HashSet::new(); + for map in [ + &dae.variables.states, + &dae.variables.algebraics, + &dae.variables.inputs, + &dae.variables.outputs, + &dae.variables.parameters, + &dae.variables.constants, + &dae.variables.discrete_reals, + &dae.variables.discrete_valued, + ] { + for name in map.keys() { names.insert(name.as_str().to_string()); } } @@ -918,9 +1623,208 @@ fn phantom_variant_reference(dae: &Dae, variant: String) -> rumoca_core::Referen } } -/// Build a map from variable name → dims for declared array variables. -fn build_array_dims_map(dae: &Dae) -> HashMap> { - let mut dims_map = HashMap::new(); +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +struct StructuredProjectionPart { + ident: String, + indices: Vec, +} + +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +struct StructuredProjectionPath { + local: bool, + parts: Vec, +} + +impl StructuredProjectionPath { + fn from_component_ref(component_ref: &rumoca_core::ComponentReference) -> Option { + let parts = component_ref + .parts + .iter() + .map(|part| { + let indices = part + .subs + .iter() + .map(|subscript| match subscript { + rumoca_core::Subscript::Index { value, .. } => Some(*value), + rumoca_core::Subscript::Colon { .. } + | rumoca_core::Subscript::Expr { .. } => None, + }) + .collect::>>()?; + Some(StructuredProjectionPart { + ident: part.ident.clone(), + indices, + }) + }) + .collect::>>()?; + Some(Self { + local: component_ref.local, + parts, + }) + } +} + +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +struct StructuredProjectionIdentity { + path: StructuredProjectionPath, + declaration: rumoca_core::DefId, +} + +#[derive(Debug, Default)] +struct RecordArrayProjectionAliases { + aliases: HashMap, + unique_identity_by_path: HashMap, +} + +impl RecordArrayProjectionAliases { + fn resolve( + &self, + component_ref: &rumoca_core::ComponentReference, + ) -> Option { + let path = StructuredProjectionPath::from_component_ref(component_ref)?; + let identity = match component_ref.def_id { + Some(declaration) => StructuredProjectionIdentity { path, declaration }, + None => self.unique_identity_by_path.get(&path)?.clone(), + }; + self.aliases.get(&identity).cloned() + } +} + +#[derive(Clone, Debug, PartialEq)] +struct DirectProjectionTarget { + identity: Option, + reference: rumoca_core::Reference, +} + +fn dae_variable_partitions(dae: &Dae) -> [&IndexMap; 8] { + [ + &dae.variables.states, + &dae.variables.algebraics, + &dae.variables.inputs, + &dae.variables.outputs, + &dae.variables.parameters, + &dae.variables.constants, + &dae.variables.discrete_reals, + &dae.variables.discrete_valued, + ] +} + +fn build_record_array_projection_alias_map( + dae: &Dae, +) -> Result { + let mut direct_paths = HashMap::new(); + for partition in dae_variable_partitions(dae) { + for (name, variable) in partition { + let Some(component_ref) = variable.component_ref.as_ref() else { + continue; + }; + let Some(path) = StructuredProjectionPath::from_component_ref(component_ref) else { + continue; + }; + let target = DirectProjectionTarget { + identity: component_ref + .def_id + .map(|declaration| StructuredProjectionIdentity { + path: path.clone(), + declaration, + }), + reference: rumoca_core::Reference::with_component_reference( + name.as_str(), + component_ref.clone(), + ), + }; + if let Some(existing) = direct_paths.insert(path, target.clone()) + && existing != target + { + return Err(ToDaeError::runtime_metadata_violation_at( + format!("conflicting directly declared projection target for `{name}`"), + variable.source_span, + )); + } + } + } + + let mut aliases = RecordArrayProjectionAliases::default(); + for partition in dae_variable_partitions(dae) { + for (name, variable) in partition { + append_record_array_projection_aliases(&mut aliases, &direct_paths, name, variable)?; + } + } + Ok(aliases) +} + +fn append_record_array_projection_aliases( + aliases: &mut RecordArrayProjectionAliases, + direct_paths: &HashMap, + name: &rumoca_core::VarName, + variable: &dae::Variable, +) -> Result<(), ToDaeError> { + let Some(component_ref) = variable.component_ref.as_ref() else { + return Ok(()); + }; + let Some(declaration) = component_ref.def_id else { + return Ok(()); + }; + for index in 0..component_ref.parts.len().saturating_sub(1) { + if component_ref.parts[index].subs.is_empty() { + continue; + } + if StructuredProjectionPath::from_component_ref(component_ref).is_none() { + continue; + } + let mut projection = component_ref.clone(); + let subscripts = std::mem::take(&mut projection.parts[index].subs); + projection + .parts + .last_mut() + .expect("component reference with an indexed part has a leaf") + .subs + .extend(subscripts); + let Some(path) = StructuredProjectionPath::from_component_ref(&projection) else { + continue; + }; + let identity = StructuredProjectionIdentity { + path: path.clone(), + declaration, + }; + let reference = + rumoca_core::Reference::with_component_reference(name.as_str(), component_ref.clone()); + if let Some(direct) = direct_paths.get(&path) { + if direct.identity.as_ref() == Some(&identity) && direct.reference == reference { + continue; + } + return Err(ToDaeError::runtime_metadata_violation_at( + format!( + "record-array projection for `{name}` collides with a directly declared variable" + ), + variable.source_span, + )); + } + if let Some(existing) = aliases.aliases.get(&identity) { + if existing == &reference { + continue; + } + return Err(ToDaeError::runtime_metadata_violation_at( + format!("conflicting record-array projection alias for `{name}`"), + variable.source_span, + )); + } + if let Some(existing_identity) = aliases.unique_identity_by_path.get(&path) + && existing_identity != &identity + { + return Err(ToDaeError::runtime_metadata_violation_at( + format!("duplicate record-array projection alias for `{name}`"), + variable.source_span, + )); + } + aliases.aliases.insert(identity.clone(), reference); + aliases.unique_identity_by_path.insert(path, identity); + } + Ok(()) +} + +fn build_record_array_field_map(dae: &Dae) -> RecordArrayFieldMap { + let mut fields: HashMap, Vec)> = + HashMap::new(); for map in [ &dae.variables.states, &dae.variables.algebraics, @@ -932,20 +1836,94 @@ fn build_array_dims_map(dae: &Dae) -> HashMap> { &dae.variables.discrete_valued, ] { for (name, var) in map { - if !var.dims.is_empty() { - dims_map.insert(name.as_str().to_string(), var.dims.clone()); + let Some(component_ref) = var.component_ref.as_ref() else { + continue; + }; + let Some((key, index)) = record_array_field_key(component_ref) else { + continue; + }; + let (indexed, field_dims) = fields.entry(key).or_default(); + if field_dims.is_empty() { + *field_dims = var.dims.clone(); } + indexed.insert( + index, + rumoca_core::Reference::with_component_reference( + name.as_str(), + component_ref.clone(), + ), + ); } } - dims_map + fields + .into_iter() + .filter_map(|(key, (indexed, field_dims))| { + let mut variants = Vec::with_capacity(indexed.len()); + for expected in 1..=indexed.len() { + let value = indexed.get(&expected)?.clone(); + variants.push(value); + } + Some(( + key, + RecordArrayFieldVariants { + variants, + field_dims, + }, + )) + }) + .collect() +} + +fn record_array_field_key( + component_ref: &rumoca_core::ComponentReference, +) -> Option<(String, usize)> { + if component_ref.parts.len() < 2 { + return None; + } + let container_index = component_ref + .parts + .iter() + .enumerate() + .take(component_ref.parts.len() - 1) + .rev() + .find_map(|(index, part)| { + single_positive_index_subscript(&part.subs).map(|sub| (index, sub)) + })?; + let (container_index, _) = container_index; + let suffix = component_ref.parts.get(container_index + 1..)?; + if suffix.is_empty() { + return None; + } + let suffix = suffix + .iter() + .map(|part| part.ident.as_str()) + .collect::>() + .join("."); + let container = component_ref.parts.get(container_index)?; + let index = single_positive_index_subscript(&container.subs)?; + let mut base_ref = component_ref.clone(); + base_ref.parts.truncate(container_index + 1); + base_ref.parts.last_mut()?.subs.clear(); + base_ref.def_id = None; + let base = rumoca_core::ComponentPath::from_component_reference(&base_ref).to_flat_string(); + Some((format!("{base}.{suffix}"), index)) +} + +fn single_positive_index_subscript(subscripts: &[rumoca_core::Subscript]) -> Option { + let [rumoca_core::Subscript::Index { value, .. }] = subscripts else { + return None; + }; + usize::try_from(*value).ok().filter(|value| *value > 0) } fn canonicalize_embedded_subscript_equation_list( equations: &mut [dae::Equation], array_dims: &HashMap>, + record_array_projection_aliases: &RecordArrayProjectionAliases, ) -> Result<(), ToDaeError> { let mut canonicalizer = EmbeddedSubscriptCanonicalizer { array_dims, + record_array_projection_aliases, error: None, }; for equation in equations { @@ -962,6 +1940,7 @@ fn canonicalize_embedded_subscript_equation_list( struct EmbeddedSubscriptCanonicalizer<'a> { array_dims: &'a HashMap>, + record_array_projection_aliases: &'a RecordArrayProjectionAliases, error: Option, } @@ -979,6 +1958,13 @@ impl ExpressionRewriter for EmbeddedSubscriptCanonicalizer<'_> { span, }; } + if let Some(reference) = self.record_array_projection_alias(name, subscripts) { + return rumoca_core::Expression::VarRef { + name: reference, + subscripts: Vec::new(), + span, + }; + } if subscripts.is_empty() && let Some(scalar_name) = rumoca_core::parse_scalar_name(name.as_str()) && self.array_dims.contains_key(scalar_name.base) @@ -1020,15 +2006,46 @@ impl ExpressionRewriter for EmbeddedSubscriptCanonicalizer<'_> { span, }; } + let subscripts = self.rewrite_subscripts(subscripts); + if subscripts + .iter() + .any(|subscript| matches!(subscript, rumoca_core::Subscript::Colon { .. })) + { + return rumoca_core::Expression::Index { + base: Box::new(rumoca_core::Expression::VarRef { + name: name.clone(), + subscripts: Vec::new(), + span, + }), + subscripts, + span, + }; + } rumoca_core::Expression::VarRef { name: name.clone(), - subscripts: self.rewrite_subscripts(subscripts), + subscripts, span, } } } impl EmbeddedSubscriptCanonicalizer<'_> { + fn record_array_projection_alias( + &self, + name: &rumoca_core::Reference, + subscripts: &[rumoca_core::Subscript], + ) -> Option { + let mut projection = name.component_ref()?.clone(); + if !subscripts.is_empty() { + projection + .parts + .last_mut()? + .subs + .extend_from_slice(subscripts); + } + self.record_array_projection_aliases.resolve(&projection) + } + fn record_error( &mut self, error: ToDaeError, @@ -1097,6 +2114,7 @@ fn merge_phantom_widths(left: Option, right: Option) -> Option bool { match expr { rumoca_core::Expression::ArrayComprehension { .. } => true, @@ -1141,32 +2159,214 @@ fn expr_has_array_comprehension(expr: &rumoca_core::Expression) -> bool { } } -fn subscript_has_array_comprehension(subscript: &rumoca_core::Subscript) -> bool { - match subscript { - rumoca_core::Subscript::Expr { expr, .. } => expr_has_array_comprehension(expr), - rumoca_core::Subscript::Index { .. } | rumoca_core::Subscript::Colon { .. } => false, +fn expr_has_record_array_member_slice(expr: &rumoca_core::Expression) -> bool { + match expr { + rumoca_core::Expression::FieldAccess { base, .. } => { + matches!( + base.as_ref(), + rumoca_core::Expression::Index { + subscripts, + .. + } if subscripts.iter().any(subscript_is_record_array_member_slice) + ) || expr_has_record_array_member_slice(base) + } + rumoca_core::Expression::Binary { lhs, rhs, .. } => { + expr_has_record_array_member_slice(lhs) || expr_has_record_array_member_slice(rhs) + } + rumoca_core::Expression::Unary { rhs, .. } => expr_has_record_array_member_slice(rhs), + rumoca_core::Expression::BuiltinCall { args, .. } + | rumoca_core::Expression::FunctionCall { args, .. } => { + args.iter().any(expr_has_record_array_member_slice) + } + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + branches.iter().any(|(condition, value)| { + expr_has_record_array_member_slice(condition) + || expr_has_record_array_member_slice(value) + }) || expr_has_record_array_member_slice(else_branch) + } + rumoca_core::Expression::Array { elements, .. } + | rumoca_core::Expression::Tuple { elements, .. } => { + elements.iter().any(expr_has_record_array_member_slice) + } + rumoca_core::Expression::Range { + start, step, end, .. + } => { + expr_has_record_array_member_slice(start) + || step + .as_ref() + .is_some_and(|step| expr_has_record_array_member_slice(step)) + || expr_has_record_array_member_slice(end) + } + rumoca_core::Expression::Index { + base, subscripts, .. + } => { + expr_has_record_array_member_slice(base) + || subscripts + .iter() + .any(subscript_has_record_array_member_slice) + } + _ => false, } } -/// Scalarize an expression at index `k` (0-based). -/// -/// - Phantom VarRefs are replaced by the k-th indexed variant from `phantom_map`. -/// - Declared array VarRefs (with no subscripts) get subscript `[k+1]` (1-based) -/// only while expanding an equation that already contains a phantom reference. -/// MLS §10.6: ordinary declared-array equations must remain array equations so -/// later matrix-aware scalarization can preserve linear algebra semantics. -/// - All other expressions are recursively processed. +fn expr_has_colon_slice(expr: &rumoca_core::Expression) -> bool { + match expr { + rumoca_core::Expression::Index { + base, subscripts, .. + } => { + subscripts + .iter() + .any(|subscript| matches!(subscript, rumoca_core::Subscript::Colon { .. })) + || expr_has_colon_slice(base) + || subscripts.iter().any(subscript_has_colon_slice) + } + rumoca_core::Expression::FieldAccess { base, .. } => expr_has_colon_slice(base), + rumoca_core::Expression::Binary { lhs, rhs, .. } => { + expr_has_colon_slice(lhs) || expr_has_colon_slice(rhs) + } + rumoca_core::Expression::Unary { rhs, .. } => expr_has_colon_slice(rhs), + rumoca_core::Expression::BuiltinCall { args, .. } + | rumoca_core::Expression::FunctionCall { args, .. } => { + args.iter().any(expr_has_colon_slice) + } + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + branches.iter().any(|(condition, value)| { + expr_has_colon_slice(condition) || expr_has_colon_slice(value) + }) || expr_has_colon_slice(else_branch) + } + rumoca_core::Expression::Array { elements, .. } + | rumoca_core::Expression::Tuple { elements, .. } => { + elements.iter().any(expr_has_colon_slice) + } + rumoca_core::Expression::Range { + start, step, end, .. + } => { + expr_has_colon_slice(start) + || step.as_ref().is_some_and(|step| expr_has_colon_slice(step)) + || expr_has_colon_slice(end) + } + _ => false, + } +} + +#[cfg(test)] +fn subscript_has_array_comprehension(subscript: &rumoca_core::Subscript) -> bool { + match subscript { + rumoca_core::Subscript::Expr { expr, .. } => expr_has_array_comprehension(expr), + rumoca_core::Subscript::Index { .. } | rumoca_core::Subscript::Colon { .. } => false, + } +} + +fn expr_has_vectorized_scalar_function_call( + expr: &rumoca_core::Expression, + array_dims: &HashMap>, + functions: &IndexMap, +) -> bool { + match expr { + rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor: false, + .. + } => { + let function = functions.get(name.var_name()); + if scalar_output_function(function) { + let mut positional_idx = 0usize; + if args.iter().any(|arg| { + let formal_rank = + scalarized_function_arg_formal_rank(function, arg, &mut positional_idx); + expr_has_vectorized_scalar_actual(arg, formal_rank, array_dims) + }) { + return true; + } + } + args.iter() + .any(|arg| expr_has_vectorized_scalar_function_call(arg, array_dims, functions)) + } + rumoca_core::Expression::Binary { lhs, rhs, .. } => { + expr_has_vectorized_scalar_function_call(lhs, array_dims, functions) + || expr_has_vectorized_scalar_function_call(rhs, array_dims, functions) + } + rumoca_core::Expression::Unary { rhs, .. } => { + expr_has_vectorized_scalar_function_call(rhs, array_dims, functions) + } + rumoca_core::Expression::BuiltinCall { args, .. } + | rumoca_core::Expression::Tuple { elements: args, .. } + | rumoca_core::Expression::Array { elements: args, .. } => args + .iter() + .any(|arg| expr_has_vectorized_scalar_function_call(arg, array_dims, functions)), + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + branches.iter().any(|(condition, value)| { + expr_has_vectorized_scalar_function_call(condition, array_dims, functions) + || expr_has_vectorized_scalar_function_call(value, array_dims, functions) + }) || expr_has_vectorized_scalar_function_call(else_branch, array_dims, functions) + } + rumoca_core::Expression::Index { base, .. } + | rumoca_core::Expression::FieldAccess { base, .. } => { + expr_has_vectorized_scalar_function_call(base, array_dims, functions) + } + _ => false, + } +} +fn expr_has_vectorized_scalar_actual( + expr: &rumoca_core::Expression, + formal_rank: usize, + array_dims: &HashMap>, +) -> bool { + if let Some((_, value)) = named_function_arg_value(expr) { + return expr_has_vectorized_scalar_actual(value, formal_rank, array_dims); + } + matches!( + expr, + rumoca_core::Expression::VarRef { + name, + subscripts, + .. + } if subscripts.is_empty() + && array_dims + .get(name.as_str()) + .is_some_and(|dims| dims.len() == formal_rank + 1) + ) +} +fn subscript_has_record_array_member_slice(subscript: &rumoca_core::Subscript) -> bool { + match subscript { + rumoca_core::Subscript::Expr { expr, .. } => expr_has_record_array_member_slice(expr), + rumoca_core::Subscript::Index { .. } | rumoca_core::Subscript::Colon { .. } => false, + } +} +fn subscript_has_colon_slice(subscript: &rumoca_core::Subscript) -> bool { + match subscript { + rumoca_core::Subscript::Expr { expr, .. } => expr_has_colon_slice(expr), + rumoca_core::Subscript::Index { .. } | rumoca_core::Subscript::Colon { .. } => false, + } +} fn scalarize_expr_at( expr: &rumoca_core::Expression, k: usize, phantom_map: &HashMap>, array_dims: &HashMap>, + var_dims: &HashMap>, + record_array_fields: &RecordArrayFieldMap, functions: &IndexMap, ) -> Result { let ctx = ScalarizeExprContext { k, phantom_map, array_dims, + var_dims, + record_array_fields, functions, }; scalarize_expr_with_context(expr, &ctx) @@ -1176,6 +2376,8 @@ struct ScalarizeExprContext<'a> { k: usize, phantom_map: &'a HashMap>, array_dims: &'a HashMap>, + var_dims: &'a HashMap>, + record_array_fields: &'a RecordArrayFieldMap, functions: &'a IndexMap, } @@ -1195,15 +2397,11 @@ fn scalarize_expr_with_context( ctx.k, ctx.phantom_map, ctx.array_dims, + ctx.record_array_fields, ) .and_then(|projected| projected.map_or_else(|| Ok(expr.clone()), Ok)), rumoca_core::Expression::Binary { op, lhs, rhs, span } => { - Ok(rumoca_core::Expression::Binary { - op: op.clone(), - lhs: Box::new(scalarize_expr_with_context(lhs, ctx)?), - rhs: Box::new(scalarize_expr_with_context(rhs, ctx)?), - span: *span, - }) + scalarize_binary_expr_at(op, lhs, rhs, *span, ctx) } rumoca_core::Expression::Unary { op, rhs, span } => Ok(rumoca_core::Expression::Unary { op: op.clone(), @@ -1215,15 +2413,11 @@ fn scalarize_expr_with_context( args, span, } => { - if let Some(expr) = scalarize_builtin_array_constructor_at( - *function, - args, - *span, - ctx.k, - ctx.phantom_map, - ctx.array_dims, - ctx.functions, - )? { + if let Some(expr) = scalarize_builtin_vector_output_at(*function, args, *span, ctx)? { + return Ok(expr); + } + if let Some(expr) = scalarize_builtin_array_constructor_at(*function, args, *span, ctx)? + { return Ok(expr); } Ok(rumoca_core::Expression::BuiltinCall { @@ -1241,15 +2435,32 @@ fn scalarize_expr_with_context( is_constructor, span, } => scalarize_function_call_at(name, args, *is_constructor, *span, ctx), + rumoca_core::Expression::Index { + base, + subscripts, + span, + } => scalarize_index_expr_at(base, subscripts, *span, ctx), + rumoca_core::Expression::FieldAccess { base, field, span } => { + scalarize_field_access_at(base, field, *span, ctx) + } rumoca_core::Expression::If { branches, else_branch, span, } => scalarize_if_expr_at(branches, else_branch, *span, ctx), rumoca_core::Expression::Array { elements, .. } => { - // An array literal in a vector equation context: extract element k - if ctx.k < elements.len() { - scalarize_expr_with_context(&elements[ctx.k], ctx) + if let Some((element_index, element_lane)) = + scalarized_array_literal_lane(elements, ctx.k, ctx.array_dims) + { + let element_ctx = ScalarizeExprContext { + k: element_lane, + phantom_map: ctx.phantom_map, + array_dims: ctx.array_dims, + var_dims: ctx.var_dims, + record_array_fields: ctx.record_array_fields, + functions: ctx.functions, + }; + scalarize_expr_with_context(&elements[element_index], &element_ctx) } else { Ok(expr.clone()) } @@ -1264,601 +2475,2495 @@ fn scalarize_expr_with_context( } } -fn scalarize_array_comprehension_at( - original: &rumoca_core::Expression, - inner: &rumoca_core::Expression, - indices: &[rumoca_core::ComprehensionIndex], - filter: Option<&rumoca_core::Expression>, - span: rumoca_core::Span, +fn scalarize_binary_expr_at( + op: &OpBinary, + lhs: &Expr, + rhs: &Expr, + span: Span, ctx: &ScalarizeExprContext<'_>, -) -> Result { - if filter.is_some() || indices.len() != 1 { - return Ok(original.clone()); - } - let Some(value) = scalarized_comprehension_index_value(&indices[0].range, ctx.k) else { - return Ok(original.clone()); - }; - let mut substitution = ComprehensionIndexSubstitution { - name: indices[0].name.clone(), - value, - span: indices[0].range.span().unwrap_or(span), - }; - let selected = substitution.rewrite_expression(inner); - scalarize_expr_with_context(&selected, ctx) -} - -fn scalarized_comprehension_index_value(range: &rumoca_core::Expression, k: usize) -> Option { - let rumoca_core::Expression::Range { - start, - step, - end: _, - .. - } = range - else { - return None; - }; - let start = integer_literal_value(start)?; - let step = match step.as_deref() { - Some(step) => integer_literal_value(step)?, - None => 1, - }; - Some(start + (k as i64) * step) -} - -fn integer_literal_value(expr: &rumoca_core::Expression) -> Option { - let rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Integer(value), - .. - } = expr - else { - return None; - }; - Some(*value) -} - -struct ComprehensionIndexSubstitution { - name: String, - value: i64, - span: rumoca_core::Span, -} - -impl ExpressionRewriter for ComprehensionIndexSubstitution { - fn walk_var_ref_expression( - &mut self, - name: &rumoca_core::Reference, - subscripts: &[rumoca_core::Subscript], - span: rumoca_core::Span, - ) -> rumoca_core::Expression { - if name.as_str() == self.name && subscripts.is_empty() { - return rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Integer(self.value), - span: self.span, - }; - } - rumoca_core::Expression::VarRef { - name: name.clone(), - subscripts: self.rewrite_subscripts(subscripts), - span, - } - } - - fn walk_array_comprehension_expression( - &mut self, - expr: &rumoca_core::Expression, - indices: &[rumoca_core::ComprehensionIndex], - filter: Option<&rumoca_core::Expression>, - span: rumoca_core::Span, - ) -> rumoca_core::Expression { - if indices.iter().any(|index| index.name == self.name) { - return rumoca_core::Expression::ArrayComprehension { - expr: Box::new(expr.clone()), - indices: indices.to_vec(), - filter: filter.cloned().map(Box::new), - span, - }; +) -> Result { + if matches!(op, OpBinary::Sub) { + let projector = Projector(ctx.var_dims, ctx.functions); + if projector.has_descendant_matrix_product_candidate(rhs, false)? { + let dims = projector.dims(lhs, span)?; + if dims.is_none() + && matches!(projector.dims(rhs, span)?, Some(dims) if !dims.is_empty()) + { + return Err(projection_error("unknown target shape", span)); + } + if let Some(dims) = dims + && let Some(rhs) = projector.project(rhs, ctx.k, &dims)? + { + let indices = projection_lane_indices(ctx.k, &dims) + .ok_or_else(|| projection_error("result shape mismatch", span))?; + let lhs = projector.element(lhs, &indices, span)?; + return Ok(vectorized_binary_expr(op.clone(), lhs, rhs, span)); + } } - ExpressionRewriter::walk_array_comprehension_expression(self, expr, indices, filter, span) } + Ok(vectorized_binary_expr( + op.clone(), + scalarize_expr_with_context(lhs, ctx)?, + scalarize_expr_with_context(rhs, ctx)?, + span, + )) } - -fn scalarize_var_ref_at( - name: &rumoca_core::Reference, - subscripts: &[rumoca_core::Subscript], +fn scalarize_builtin_vector_output_at( + function: rumoca_core::BuiltinFunction, + args: &[rumoca_core::Expression], span: rumoca_core::Span, - k: usize, - phantom_map: &HashMap>, - array_dims: &HashMap>, + ctx: &ScalarizeExprContext<'_>, ) -> Result, ToDaeError> { - let n = name.as_str(); - if !subscripts.is_empty() { + let Some(output_width) = builtin_fixed_vector_output_width(function) else { return Ok(None); - } - if let Some(variants) = phantom_map.get(n) - && k < variants.len() - { - return Ok(Some(rumoca_core::Expression::VarRef { - name: variants[k].clone(), - subscripts: vec![], - span, - })); - } - if !array_dims.contains_key(n) { + }; + if ctx.k >= output_width { return Ok(None); } - let index = one_based_scalar_index(k, span, "DAE phantom scalarized variable subscript")?; - Ok(Some(rumoca_core::Expression::VarRef { - name: name.clone(), + let index = one_based_scalar_index(ctx.k, span, "DAE scalarized builtin output subscript")?; + Ok(Some(rumoca_core::Expression::Index { + base: Box::new(rumoca_core::Expression::BuiltinCall { + function, + args: args + .iter() + .map(|arg| vectorize_builtin_vector_arg(arg, ctx)) + .collect::, _>>()?, + span, + }), subscripts: vec![generated_index_subscript( index, span, - "DAE phantom scalarized variable subscript", + "DAE scalarized builtin output subscript", )?], span, })) } +fn vectorize_builtin_vector_arg( + arg: &rumoca_core::Expression, + ctx: &ScalarizeExprContext<'_>, +) -> Result { + if let Some(vector) = try_project_colon_slice_var_ref(arg, ctx.array_dims)? { + return Ok(vector); + } + Ok(vectorize_phantom_expr(arg, ctx.phantom_map)) +} -fn reference_or_wrapper_span( - reference: &rumoca_core::Reference, - span: rumoca_core::Span, -) -> rumoca_core::Span { - if span.is_dummy() { - reference.span().unwrap_or(span) - } else { - span +fn builtin_fixed_vector_output_width(function: rumoca_core::BuiltinFunction) -> Option { + match function { + rumoca_core::BuiltinFunction::Cross => Some(3), + _ => None, } } -fn scalarize_function_call_at( - name: &rumoca_core::Reference, - args: &[rumoca_core::Expression], - is_constructor: bool, +fn scalarize_index_expr_at( + base: &rumoca_core::Expression, + subscripts: &[rumoca_core::Subscript], span: rumoca_core::Span, ctx: &ScalarizeExprContext<'_>, ) -> Result { - let function = ctx.functions.get(&rumoca_core::VarName::new(name.as_str())); - let first_output_size = - first_function_output_size(name.as_str(), ctx.functions).ok_or_else(|| { - ToDaeError::runtime_contract_violation_at( - format!("missing function output metadata for `{name}`"), + if let rumoca_core::Expression::VarRef { + name, + subscripts: base_subscripts, + .. + } = base + && base_subscripts.is_empty() + { + if subscripts + .iter() + .all(|subscript| matches!(subscript, rumoca_core::Subscript::Colon { .. })) + && let Some(expr) = scalarize_var_ref_at( + name, + base_subscripts, span, - ) - })?; - if first_output_size > 1 { - let index = - one_based_scalar_index(ctx.k, span, "DAE scalarized function output subscript")?; - return Ok(rumoca_core::Expression::Index { - base: Box::new(rumoca_core::Expression::FunctionCall { + ctx.k, + ctx.phantom_map, + ctx.array_dims, + ctx.record_array_fields, + )? + { + return Ok(expr); + } + if let Some(dims) = ctx.array_dims.get(name.as_str()) + && let Some(projected_subscripts) = + project_slice_subscripts_for_lane(dims, subscripts, ctx.k, span)? + { + return Ok(rumoca_core::Expression::VarRef { name: name.clone(), - args: args - .iter() - .map(|arg| vectorize_phantom_expr(arg, ctx.phantom_map)) - .collect(), - is_constructor, - span, - }), - subscripts: vec![generated_index_subscript( - index, + subscripts: projected_subscripts, span, - "DAE scalarized function output subscript", - )?], - span, + }); + } + } + Ok(rumoca_core::Expression::Index { + base: Box::new(scalarize_expr_with_context(base, ctx)?), + subscripts: subscripts.to_vec(), + span, + }) +} + +fn scalarize_field_access_at( + base: &rumoca_core::Expression, + field: &str, + span: rumoca_core::Span, + ctx: &ScalarizeExprContext<'_>, +) -> Result { + if let Some(projected) = + scalarize_record_array_member_slice_at(base, field, span, ctx.k, ctx.record_array_fields)? + { + return Ok(projected); + } + if let Some(projected) = scalarize_record_array_field_at(base, field, span, ctx)? { + return Ok(projected); + } + let base = scalarize_expr_with_context(base, ctx)?; + if let rumoca_core::Expression::VarRef { name, span, .. } = &base + && let rumoca_core::Expression::VarRef { subscripts, .. } = &base + && subscripts.is_empty() + { + let base_name = record_arg_projection_base_name(name, &[field.to_string()]) + .unwrap_or_else(|| name.clone()); + return Ok(rumoca_core::Expression::VarRef { + name: base_name.with_appended_field(field), + subscripts: Vec::new(), + span: *span, }); } - Ok(rumoca_core::Expression::FunctionCall { - name: name.clone(), - args: scalarize_function_call_args(args, function, ctx)?, - is_constructor, + Ok(rumoca_core::Expression::FieldAccess { + base: Box::new(base), + field: field.to_string(), span, }) } -fn scalarize_function_call_args( - args: &[rumoca_core::Expression], - function: Option<&rumoca_core::Function>, +fn scalarize_record_array_field_at( + base: &rumoca_core::Expression, + field: &str, + span: rumoca_core::Span, ctx: &ScalarizeExprContext<'_>, -) -> Result, ToDaeError> { - args.iter() - .enumerate() - .map(|(index, arg)| { - if function_input_expects_array(function, arg, index) { - return Ok(vectorize_phantom_expr(arg, ctx.phantom_map)); - } - scalarize_expr_with_context(arg, ctx) - }) - .collect() +) -> Result, ToDaeError> { + project_record_array_field_rhs_at(base, field, span, ctx.k, ctx.record_array_fields) } -fn function_input_expects_array( - function: Option<&rumoca_core::Function>, - arg: &rumoca_core::Expression, - index: usize, -) -> bool { - let Some(function) = function else { - return false; +fn project_record_array_field_rhs_at( + base: &rumoca_core::Expression, + field: &str, + span: rumoca_core::Span, + k: usize, + record_array_fields: &RecordArrayFieldMap, +) -> Result, ToDaeError> { + let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = base + else { + return Ok(None); }; - let input_name = named_argument_input_name(arg); - let param = input_name - .and_then(|name| function.inputs.iter().find(|param| param.name == name)) - .or_else(|| function.inputs.get(index)); - param.is_some_and(|param| !param.dims.is_empty() || !param.shape_expr.is_empty()) + if !subscripts.is_empty() { + return Ok(None); + } + let key = format!("{}.{}", name.as_str(), field); + project_record_array_field_entry_at(record_array_fields.get(&key), k, span) } -fn named_argument_input_name(arg: &rumoca_core::Expression) -> Option<&str> { - let rumoca_core::Expression::FunctionCall { +fn project_record_array_field_entry_at( + entry: Option<&RecordArrayFieldVariants>, + k: usize, + span: rumoca_core::Span, +) -> Result, ToDaeError> { + let Some(entry) = entry else { + return Ok(None); + }; + let field_width = record_array_field_width(&entry.field_dims); + if field_width == 0 { + return Ok(None); + } + let record_index = k / field_width; + let field_index = k % field_width; + let Some(projected) = entry.variants.get(record_index).cloned() else { + return Ok(None); + }; + let subscripts = if field_width == 1 { + Vec::new() + } else { + let index = one_based_scalar_index( + field_index, + span, + "DAE record-array field scalarized subscript", + )?; + vec![generated_index_subscript( + index, + span, + "DAE record-array field scalarized subscript", + )?] + }; + Ok(Some(rumoca_core::Expression::VarRef { + name: projected, + subscripts, + span, + })) +} + +fn record_array_field_width(dims: &[i64]) -> usize { + if dims.is_empty() { + return 1; + } + dims.iter() + .copied() + .try_fold(1usize, |acc, dim| { + let dim = usize::try_from(dim).ok()?; + (dim > 0).then_some(acc.saturating_mul(dim)) + }) + .unwrap_or(0) +} + +fn record_array_field_scalar_count(entry: &RecordArrayFieldVariants) -> usize { + entry + .variants + .len() + .saturating_mul(record_array_field_width(&entry.field_dims)) +} + +fn scalarize_record_array_member_slice_at( + base: &rumoca_core::Expression, + field: &str, + span: rumoca_core::Span, + k: usize, + record_array_fields: &RecordArrayFieldMap, +) -> Result, ToDaeError> { + if let Some((name, subscript, suffix)) = record_array_member_slice_parts(base, field) + && let Some(component_ref) = name.component_ref() + && let Some(key) = record_array_member_slice_field_key(component_ref, &suffix) + && let Some(entry) = record_array_fields.get(&key) + && let Some(projected) = + project_record_array_field_slice_entry_at(entry, subscript, k, span)? + { + return Ok(Some(projected)); + } + + let rumoca_core::Expression::Index { + base: inner, + subscripts, + .. + } = base + else { + return Ok(None); + }; + let [subscript] = subscripts.as_slice() else { + return Ok(None); + }; + let rumoca_core::Expression::VarRef { name, - is_constructor: true, + subscripts: ref_subscripts, .. - } = arg + } = inner.as_ref() else { - return None; + return Ok(None); }; - name.as_str().strip_prefix("__rumoca_named_arg__.") + if !ref_subscripts.is_empty() { + return Ok(None); + } + let Some(component_ref) = name.component_ref() else { + return Ok(None); + }; + let mut element_ref = component_ref.clone(); + let Some(part) = element_ref.parts.last_mut() else { + return Ok(None); + }; + let index = record_array_member_slice_index(subscript, k, span)?; + part.subs = vec![generated_index_subscript( + index, + span, + "DAE record-array member slice subscript", + )?]; + element_ref.parts.push(rumoca_core::ComponentRefPart { + ident: field.to_string(), + span, + subs: Vec::new(), + }); + Ok(Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::from_component_reference(element_ref), + subscripts: Vec::new(), + span, + })) } -fn vectorize_phantom_array_formal_args( +fn record_array_member_slice_parts<'a>( + base: &'a rumoca_core::Expression, + field: &str, +) -> Option<( + &'a rumoca_core::Reference, + &'a rumoca_core::Subscript, + Vec, +)> { + let (name, subscript, mut suffix) = collect_record_array_member_slice_base(base)?; + suffix.push(field.to_string()); + Some((name, subscript, suffix)) +} + +fn collect_record_array_member_slice_base( expr: &rumoca_core::Expression, - phantom_map: &HashMap>, - functions: &IndexMap, -) -> rumoca_core::Expression { - PhantomArrayFormalArgVectorizer { - phantom_map, - functions, +) -> Option<( + &rumoca_core::Reference, + &rumoca_core::Subscript, + Vec, +)> { + match expr { + rumoca_core::Expression::Index { + base, subscripts, .. + } => { + let [subscript] = subscripts.as_slice() else { + return None; + }; + let rumoca_core::Expression::VarRef { + name, + subscripts: ref_subscripts, + .. + } = base.as_ref() + else { + return None; + }; + ref_subscripts + .is_empty() + .then_some((name, subscript, Vec::new())) + } + rumoca_core::Expression::FieldAccess { base, field, .. } => { + let (name, subscript, mut suffix) = collect_record_array_member_slice_base(base)?; + suffix.push(field.clone()); + Some((name, subscript, suffix)) + } + _ => None, } - .rewrite_expression(expr) } -struct PhantomArrayFormalArgVectorizer<'a> { - phantom_map: &'a HashMap>, - functions: &'a IndexMap, +fn record_array_member_slice_field_key( + component_ref: &rumoca_core::ComponentReference, + suffix: &[String], +) -> Option { + if suffix.is_empty() { + return None; + } + let mut base_ref = component_ref.clone(); + base_ref.parts.last_mut()?.subs.clear(); + base_ref.def_id = None; + let base = rumoca_core::ComponentPath::from_component_reference(&base_ref).to_flat_string(); + Some(format!("{base}.{}", suffix.join("."))) } -impl ExpressionRewriter for PhantomArrayFormalArgVectorizer<'_> { - fn walk_function_call_expression( - &mut self, - name: &rumoca_core::Reference, - args: &[rumoca_core::Expression], - is_constructor: bool, - span: rumoca_core::Span, - ) -> rumoca_core::Expression { - let function = self - .functions - .get(&rumoca_core::VarName::new(name.as_str())); - rumoca_core::Expression::FunctionCall { - name: name.clone(), - args: args - .iter() - .enumerate() - .map(|(index, arg)| { - if function_input_expects_array(function, arg, index) { - vectorize_phantom_expr(arg, self.phantom_map) - } else { - self.rewrite_expression(arg) - } - }) - .collect(), - is_constructor, - span, - } +fn project_record_array_field_slice_entry_at( + entry: &RecordArrayFieldVariants, + subscript: &rumoca_core::Subscript, + k: usize, + span: rumoca_core::Span, +) -> Result, ToDaeError> { + let field_width = record_array_field_width(&entry.field_dims); + if field_width == 0 { + return Ok(None); } + let record_lane = k / field_width; + let field_index = k % field_width; + let element_index = record_array_member_slice_index(subscript, record_lane, span)?; + let variant_index = usize::try_from(element_index.checked_sub(1).ok_or_else(|| { + ToDaeError::runtime_metadata_violation( + "record-array member slice produced a non-positive index".to_string(), + ) + })?) + .map_err(|_| { + ToDaeError::runtime_metadata_violation( + "record-array member slice index is outside supported range".to_string(), + ) + })?; + let Some(projected) = entry.variants.get(variant_index).cloned() else { + return Ok(None); + }; + let subscripts = if field_width == 1 { + Vec::new() + } else { + let index = one_based_scalar_index( + field_index, + span, + "DAE record-array member field scalarized subscript", + )?; + vec![generated_index_subscript( + index, + span, + "DAE record-array member field scalarized subscript", + )?] + }; + Ok(Some(rumoca_core::Expression::VarRef { + name: projected, + subscripts, + span, + })) } -fn one_based_scalar_index( - zero_based: usize, - span: rumoca_core::Span, - context: &'static str, -) -> Result { - zero_based - .checked_add(1) - .and_then(|index| i64::try_from(index).ok()) - .ok_or_else(|| { - ToDaeError::runtime_contract_violation_at( - format!("{context} {zero_based} exceeds i64 range"), - span, - ) - }) +fn subscript_is_record_array_member_slice(subscript: &rumoca_core::Subscript) -> bool { + matches!(subscript, rumoca_core::Subscript::Colon { .. }) + || matches!( + subscript, + rumoca_core::Subscript::Expr { expr, .. } + if matches!(expr.as_ref(), rumoca_core::Expression::Range { .. }) + ) } -fn generated_index_subscript( - index: i64, +fn record_array_member_slice_index( + subscript: &rumoca_core::Subscript, + k: usize, span: rumoca_core::Span, - context: &'static str, -) -> Result { - rumoca_core::Subscript::try_generated_index(index, span, context).map_err(|err| { - if span.is_dummy() { - ToDaeError::runtime_metadata_violation(err.to_string()) - } else { - ToDaeError::runtime_metadata_violation_at(err.to_string(), span) +) -> Result { + match subscript { + rumoca_core::Subscript::Colon { .. } => { + one_based_scalar_index(k, span, "DAE record-array member slice subscript") } - }) + rumoca_core::Subscript::Expr { expr, .. } => { + scalarized_record_array_subscript_index(expr, k).ok_or_else(|| { + ToDaeError::runtime_metadata_violation( + "record-array member slice requires a compile-time integer or range subscript" + .to_string(), + ) + }) + } + rumoca_core::Subscript::Index { value, .. } => Ok(*value), + } } -fn scalarize_if_expr_at( - branches: &[(rumoca_core::Expression, rumoca_core::Expression)], - else_branch: &rumoca_core::Expression, +fn scalarize_array_comprehension_at( + original: &rumoca_core::Expression, + inner: &rumoca_core::Expression, + indices: &[rumoca_core::ComprehensionIndex], + filter: Option<&rumoca_core::Expression>, span: rumoca_core::Span, ctx: &ScalarizeExprContext<'_>, ) -> Result { - Ok(rumoca_core::Expression::If { - branches: branches - .iter() - .map(|(condition, value)| { - Ok(( - scalarize_expr_with_context(condition, ctx)?, - scalarize_expr_with_context(value, ctx)?, - )) - }) - .collect::, ToDaeError>>()?, - else_branch: Box::new(scalarize_expr_with_context(else_branch, ctx)?), - span, - }) + if filter.is_some() || indices.len() != 1 { + return Ok(original.clone()); + } + let Some(value) = scalarized_comprehension_index_value(&indices[0].range, ctx.k) else { + return Ok(original.clone()); + }; + let mut substitution = ComprehensionIndexSubstitution { + name: indices[0].name.clone(), + value, + span: indices[0].range.span().unwrap_or(span), + }; + let selected = substitution.rewrite_expression(inner); + scalarize_expr_with_context(&selected, ctx) } -fn scalarize_builtin_array_constructor_at( - function: rumoca_core::BuiltinFunction, - args: &[rumoca_core::Expression], - span: rumoca_core::Span, +fn scalarized_comprehension_index_value(range: &rumoca_core::Expression, k: usize) -> Option { + let rumoca_core::Expression::Range { + start, + step, + end: _, + .. + } = range + else { + return None; + }; + let start = integer_constant_value(start)?; + let step = match step.as_deref() { + Some(step) => integer_constant_value(step)?, + None => 1, + }; + Some(start + (k as i64) * step) +} + +fn scalarized_record_array_subscript_index( + expr: &rumoca_core::Expression, k: usize, - phantom_map: &HashMap>, - array_dims: &HashMap>, - functions: &IndexMap, -) -> Result, ToDaeError> { - match function { - rumoca_core::BuiltinFunction::Zeros => Ok(Some(real_literal(0.0, span))), - rumoca_core::BuiltinFunction::Ones => Ok(Some(real_literal(1.0, span))), - rumoca_core::BuiltinFunction::Fill => { - let Some(value) = args.first() else { - return Ok(None); - }; - scalarize_expr_at(value, k, phantom_map, array_dims, functions).map(Some) - } - rumoca_core::BuiltinFunction::Identity => { - let Some(n) = args.first().and_then(literal_positive_usize) else { - return Ok(None); - }; - let row = k / n; - let col = k % n; - Ok(Some(real_literal(if row == col { 1.0 } else { 0.0 }, span))) - } - _ => Ok(None), - } +) -> Option { + integer_constant_value(expr).or_else(|| scalarized_comprehension_index_value(expr, k)) } -fn literal_positive_usize(expr: &rumoca_core::Expression) -> Option { +fn integer_constant_value(expr: &rumoca_core::Expression) -> Option { match expr { rumoca_core::Expression::Literal { value: rumoca_core::Literal::Integer(value), .. - } => usize::try_from(*value).ok().filter(|value| *value > 0), - rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(value), - .. - } if *value > 0.0 && value.fract() == 0.0 => Some(*value as usize), + } => Some(*value), + rumoca_core::Expression::Unary { op, rhs, .. } => { + let value = integer_constant_value(rhs)?; + match op { + rumoca_core::OpUnary::Plus | rumoca_core::OpUnary::DotPlus => Some(value), + rumoca_core::OpUnary::Minus | rumoca_core::OpUnary::DotMinus => value.checked_neg(), + _ => None, + } + } + rumoca_core::Expression::Binary { op, lhs, rhs, .. } => { + let lhs = integer_constant_value(lhs)?; + let rhs = integer_constant_value(rhs)?; + rumoca_core::eval_ast_integer_binary(op, lhs, rhs) + } _ => None, } } -fn real_literal(value: f64, span: rumoca_core::Span) -> rumoca_core::Expression { - rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(value), - span, - } -} - -fn first_function_output_size( - name: &str, - functions: &IndexMap, -) -> Option { - let lookup_name = rumoca_core::VarName::new(name); - let function = functions.get(&lookup_name)?; - let output = function.outputs.first()?; - Some(compute_var_size(&output.dims)) -} - -fn vectorize_phantom_expr( - expr: &rumoca_core::Expression, - phantom_map: &HashMap>, -) -> rumoca_core::Expression { - PhantomVectorizer { phantom_map }.rewrite_expression(expr) +fn integer_literal_value(expr: &rumoca_core::Expression) -> Option { + let rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + .. + } = expr + else { + return None; + }; + Some(*value) } -struct PhantomVectorizer<'a> { - phantom_map: &'a HashMap>, +struct ComprehensionIndexSubstitution { + name: String, + value: i64, + span: rumoca_core::Span, } -impl ExpressionRewriter for PhantomVectorizer<'_> { - fn rewrite_var_ref_expression( +impl ExpressionRewriter for ComprehensionIndexSubstitution { + fn walk_var_ref_expression( &mut self, name: &rumoca_core::Reference, subscripts: &[rumoca_core::Subscript], span: rumoca_core::Span, ) -> rumoca_core::Expression { - if subscripts.is_empty() - && let Some(variants) = self.phantom_map.get(name.as_str()) - { - return rumoca_core::Expression::Array { - elements: variants - .iter() - .map(|variant| rumoca_core::Expression::VarRef { - name: variant.clone(), - subscripts: Vec::new(), - span, - }) - .collect(), - is_matrix: false, - span, + if name.as_str() == self.name && subscripts.is_empty() { + return rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(self.value), + span: self.span, }; } - self.walk_var_ref_expression(name, subscripts, span) - } - - fn walk_binary_expression( - &mut self, - op: &rumoca_core::OpBinary, - lhs: &rumoca_core::Expression, - rhs: &rumoca_core::Expression, - span: rumoca_core::Span, - ) -> rumoca_core::Expression { - let lhs = self.rewrite_expression(lhs); - let rhs = self.rewrite_expression(rhs); - vectorized_binary_expr(op.clone(), lhs, rhs, span) + rumoca_core::Expression::VarRef { + name: name.clone(), + subscripts: self.rewrite_subscripts(subscripts), + span, + } } - fn walk_unary_expression( + fn walk_array_comprehension_expression( &mut self, - op: &rumoca_core::OpUnary, - rhs: &rumoca_core::Expression, + expr: &rumoca_core::Expression, + indices: &[rumoca_core::ComprehensionIndex], + filter: Option<&rumoca_core::Expression>, span: rumoca_core::Span, ) -> rumoca_core::Expression { - let rhs = self.rewrite_expression(rhs); - if let rumoca_core::Expression::Array { elements, .. } = rhs { - return rumoca_core::Expression::Array { - elements: elements - .into_iter() - .map(|element| rumoca_core::Expression::Unary { - op: op.clone(), - rhs: Box::new(element), - span, - }) - .collect(), - is_matrix: false, + if indices.iter().any(|index| index.name == self.name) { + return rumoca_core::Expression::ArrayComprehension { + expr: Box::new(expr.clone()), + indices: indices.to_vec(), + filter: filter.cloned().map(Box::new), span, }; } - rumoca_core::Expression::Unary { - op: op.clone(), - rhs: Box::new(rhs), - span, - } + ExpressionRewriter::walk_array_comprehension_expression(self, expr, indices, filter, span) } } -fn vectorized_binary_expr( - op: rumoca_core::OpBinary, - lhs: rumoca_core::Expression, - rhs: rumoca_core::Expression, +fn scalarize_var_ref_at( + name: &rumoca_core::Reference, + subscripts: &[rumoca_core::Subscript], span: rumoca_core::Span, -) -> rumoca_core::Expression { - match (lhs, rhs) { - ( - rumoca_core::Expression::Array { - elements: lhs_values, - .. - }, - rumoca_core::Expression::Array { - elements: rhs_values, - .. - }, - ) if lhs_values.len() == rhs_values.len() => { - array_from_binary_elements(op, lhs_values.into_iter().zip(rhs_values).collect(), span) - } - ( - rumoca_core::Expression::Array { - elements: lhs_values, - .. - }, - rhs, - ) => array_from_binary_elements( - op, - lhs_values - .into_iter() - .map(|lhs| (lhs, rhs.clone())) - .collect(), - span, - ), - ( - lhs, - rumoca_core::Expression::Array { - elements: rhs_values, - .. - }, - ) => array_from_binary_elements( - op, - rhs_values - .into_iter() - .map(|rhs| (lhs.clone(), rhs)) - .collect(), - span, - ), - (lhs, rhs) => rumoca_core::Expression::Binary { - op, - lhs: Box::new(lhs), - rhs: Box::new(rhs), + k: usize, + phantom_map: &HashMap>, + array_dims: &HashMap>, + record_array_fields: &RecordArrayFieldMap, +) -> Result, ToDaeError> { + let n = name.as_str(); + if !subscripts.is_empty() { + return Ok(None); + } + if let Some(variants) = phantom_map.get(n) + && k < variants.len() + { + return Ok(Some(rumoca_core::Expression::VarRef { + name: variants[k].clone(), + subscripts: vec![], span, - }, + })); } -} - -fn array_from_binary_elements( - op: rumoca_core::OpBinary, - pairs: Vec<(rumoca_core::Expression, rumoca_core::Expression)>, - span: rumoca_core::Span, -) -> rumoca_core::Expression { - rumoca_core::Expression::Array { + if let Some(projected) = + project_record_array_field_entry_at(record_array_fields.get(n), k, span)? + { + return Ok(Some(projected)); + } + if !array_dims.contains_key(n) { + return Ok(None); + } + let index = one_based_scalar_index(k, span, "DAE phantom scalarized variable subscript")?; + Ok(Some(rumoca_core::Expression::VarRef { + name: name.clone(), + subscripts: vec![generated_index_subscript( + index, + span, + "DAE phantom scalarized variable subscript", + )?], + span, + })) +} + +fn reference_or_wrapper_span( + reference: &rumoca_core::Reference, + span: rumoca_core::Span, +) -> rumoca_core::Span { + if span.is_dummy() { + reference.span().unwrap_or(span) + } else { + span + } +} + +fn scalarize_function_call_at( + name: &rumoca_core::Reference, + args: &[rumoca_core::Expression], + is_constructor: bool, + span: rumoca_core::Span, + ctx: &ScalarizeExprContext<'_>, +) -> Result { + let function = ctx.functions.get(&rumoca_core::VarName::new(name.as_str())); + let first_output_size = + first_function_output_size(name.as_str(), ctx.functions).ok_or_else(|| { + ToDaeError::runtime_contract_violation_at( + format!("missing function output metadata for `{name}`"), + span, + ) + })?; + if first_output_size > 1 { + let index = + one_based_scalar_index(ctx.k, span, "DAE scalarized function output subscript")?; + return Ok(rumoca_core::Expression::Index { + base: Box::new(rumoca_core::Expression::FunctionCall { + name: name.clone(), + args: args + .iter() + .map(|arg| vectorize_phantom_expr(arg, ctx.phantom_map)) + .collect(), + is_constructor, + span, + }), + subscripts: vec![generated_index_subscript( + index, + span, + "DAE scalarized function output subscript", + )?], + span, + }); + } + Ok(rumoca_core::Expression::FunctionCall { + name: name.clone(), + args: scalarize_function_call_args(args, function, ctx)?, + is_constructor, + span, + }) +} + +fn scalarize_function_call_args( + args: &[rumoca_core::Expression], + function: Option<&rumoca_core::Function>, + ctx: &ScalarizeExprContext<'_>, +) -> Result, ToDaeError> { + let scalar_output = scalar_output_function(function); + let mut positional_idx = 0usize; + args.iter() + .enumerate() + .map(|(index, arg)| { + let formal_rank = + scalarized_function_arg_formal_rank(function, arg, &mut positional_idx); + if scalar_output + && function_input_expects_array(function, arg, index) + && expr_phantom_ref_width(arg, ctx.phantom_map).is_some() + { + return Ok(vectorize_phantom_expr(arg, ctx.phantom_map)); + } + if scalar_output { + return project_scalarized_function_arg_at( + arg, + formal_rank, + ctx.k, + ctx.array_dims, + ctx.record_array_fields, + ctx.functions, + ); + } + if function_input_expects_array(function, arg, index) { + return Ok(vectorize_phantom_expr(arg, ctx.phantom_map)); + } + scalarize_expr_with_context(arg, ctx) + }) + .collect() +} + +fn function_input_expects_array( + function: Option<&rumoca_core::Function>, + arg: &rumoca_core::Expression, + index: usize, +) -> bool { + let Some(function) = function else { + return false; + }; + let input_name = named_argument_input_name(arg); + let param = input_name + .and_then(|name| function.inputs.iter().find(|param| param.name == name)) + .or_else(|| function.inputs.get(index)); + param.is_some_and(|param| !param.dims.is_empty() || !param.shape_expr.is_empty()) +} + +fn named_argument_input_name(arg: &rumoca_core::Expression) -> Option<&str> { + let rumoca_core::Expression::FunctionCall { + name, + is_constructor: true, + .. + } = arg + else { + return None; + }; + name.as_str().strip_prefix("__rumoca_named_arg__.") +} + +fn vectorize_phantom_array_formal_args( + expr: &rumoca_core::Expression, + phantom_map: &HashMap>, + functions: &IndexMap, +) -> rumoca_core::Expression { + PhantomArrayFormalArgVectorizer { + phantom_map, + functions, + } + .rewrite_expression(expr) +} + +struct PhantomArrayFormalArgVectorizer<'a> { + phantom_map: &'a HashMap>, + functions: &'a IndexMap, +} + +impl ExpressionRewriter for PhantomArrayFormalArgVectorizer<'_> { + fn walk_function_call_expression( + &mut self, + name: &rumoca_core::Reference, + args: &[rumoca_core::Expression], + is_constructor: bool, + span: rumoca_core::Span, + ) -> rumoca_core::Expression { + let function = self + .functions + .get(&rumoca_core::VarName::new(name.as_str())); + rumoca_core::Expression::FunctionCall { + name: name.clone(), + args: args + .iter() + .enumerate() + .map(|(index, arg)| { + if function_input_expects_array(function, arg, index) { + vectorize_phantom_expr(arg, self.phantom_map) + } else { + self.rewrite_expression(arg) + } + }) + .collect(), + is_constructor, + span, + } + } +} + +fn one_based_scalar_index( + zero_based: usize, + span: rumoca_core::Span, + context: &'static str, +) -> Result { + zero_based + .checked_add(1) + .and_then(|index| i64::try_from(index).ok()) + .ok_or_else(|| { + ToDaeError::runtime_contract_violation_at( + format!("{context} {zero_based} exceeds i64 range"), + span, + ) + }) +} + +fn generated_index_subscript( + index: i64, + span: rumoca_core::Span, + context: &'static str, +) -> Result { + rumoca_core::Subscript::try_generated_index(index, span, context).map_err(|err| { + if span.is_dummy() { + ToDaeError::runtime_metadata_violation(err.to_string()) + } else { + ToDaeError::runtime_metadata_violation_at(err.to_string(), span) + } + }) +} + +fn generated_colon_subscript( + span: rumoca_core::Span, + context: &'static str, +) -> Result { + rumoca_core::Subscript::try_generated_colon(span, context).map_err(|err| { + if span.is_dummy() { + ToDaeError::runtime_metadata_violation(err.to_string()) + } else { + ToDaeError::runtime_metadata_violation_at(err.to_string(), span) + } + }) +} + +fn scalarize_if_expr_at( + branches: &[(rumoca_core::Expression, rumoca_core::Expression)], + else_branch: &rumoca_core::Expression, + span: rumoca_core::Span, + ctx: &ScalarizeExprContext<'_>, +) -> Result { + Ok(rumoca_core::Expression::If { + branches: branches + .iter() + .map(|(condition, value)| { + Ok(( + scalarize_expr_with_context(condition, ctx)?, + scalarize_expr_with_context(value, ctx)?, + )) + }) + .collect::, ToDaeError>>()?, + else_branch: Box::new(scalarize_expr_with_context(else_branch, ctx)?), + span, + }) +} +fn scalarize_builtin_array_constructor_at( + function: rumoca_core::BuiltinFunction, + args: &[rumoca_core::Expression], + span: rumoca_core::Span, + ctx: &ScalarizeExprContext<'_>, +) -> Result, ToDaeError> { + match function { + rumoca_core::BuiltinFunction::Zeros => Ok(Some(real_literal(0.0, span))), + rumoca_core::BuiltinFunction::Ones => Ok(Some(real_literal(1.0, span))), + rumoca_core::BuiltinFunction::Fill => { + let Some(value) = args.first() else { + return Ok(None); + }; + let value_lane = scalarized_fill_value_lane( + value, + ctx.k, + ctx.phantom_map, + ctx.array_dims, + ctx.record_array_fields, + ); + scalarize_expr_at( + value, + value_lane, + ctx.phantom_map, + ctx.array_dims, + ctx.var_dims, + ctx.record_array_fields, + ctx.functions, + ) + .map(Some) + } + rumoca_core::BuiltinFunction::Identity => { + let Some(n) = args.first().and_then(literal_positive_usize) else { + return Ok(None); + }; + let row = ctx.k / n; + let col = ctx.k % n; + Ok(Some(real_literal(if row == col { 1.0 } else { 0.0 }, span))) + } + _ => Ok(None), + } +} +fn scalarized_fill_value_lane( + value: &rumoca_core::Expression, + k: usize, + phantom_map: &HashMap>, + array_dims: &HashMap>, + record_array_fields: &RecordArrayFieldMap, +) -> usize { + let Some(width) = + scalarized_fill_value_width(value, phantom_map, array_dims, record_array_fields) + else { + return k; + }; + if width <= 1 { 0 } else { k % width } +} +fn scalarized_fill_value_width( + value: &rumoca_core::Expression, + phantom_map: &HashMap>, + array_dims: &HashMap>, + record_array_fields: &RecordArrayFieldMap, +) -> Option { + match value { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } if subscripts.is_empty() => phantom_map + .get(name.as_str()) + .map(Vec::len) + .or_else(|| { + record_array_fields + .get(name.as_str()) + .map(record_array_field_scalar_count) + }) + .or_else(|| { + array_dims + .get(name.as_str()) + .map(|dims| compute_var_size(dims)) + }), + rumoca_core::Expression::Array { elements, .. } => Some(elements.len()), + rumoca_core::Expression::Index { + base, subscripts, .. + } => { + let rumoca_core::Expression::VarRef { + name, + subscripts: base_subscripts, + .. + } = base.as_ref() + else { + return None; + }; + if !base_subscripts.is_empty() { + return None; + } + array_dims + .get(name.as_str()) + .and_then(|dims| projected_dims_for_subscripts(dims, subscripts)) + .map(|dims| compute_var_size(&dims)) + } + _ => None, + } +} +fn literal_positive_usize(expr: &rumoca_core::Expression) -> Option { + match expr { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + .. + } => usize::try_from(*value).ok().filter(|value| *value > 0), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(value), + .. + } if *value > 0.0 && value.fract() == 0.0 => Some(*value as usize), + _ => None, + } +} +fn real_literal(value: f64, span: rumoca_core::Span) -> rumoca_core::Expression { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(value), + span, + } +} +fn first_function_output_size( + name: &str, + functions: &IndexMap, +) -> Option { + let lookup_name = rumoca_core::VarName::new(name); + let function = functions.get(&lookup_name)?; + let output = function.outputs.first()?; + Some(compute_var_size(&output.dims)) +} +fn vectorize_phantom_expr( + expr: &rumoca_core::Expression, + phantom_map: &HashMap>, +) -> rumoca_core::Expression { + PhantomVectorizer { phantom_map }.rewrite_expression(expr) +} +struct PhantomVectorizer<'a> { + phantom_map: &'a HashMap>, +} +impl ExpressionRewriter for PhantomVectorizer<'_> { + fn rewrite_var_ref_expression( + &mut self, + name: &rumoca_core::Reference, + subscripts: &[rumoca_core::Subscript], + span: rumoca_core::Span, + ) -> rumoca_core::Expression { + if subscripts.is_empty() + && let Some(variants) = self.phantom_map.get(name.as_str()) + { + return rumoca_core::Expression::Array { + elements: variants + .iter() + .map(|variant| rumoca_core::Expression::VarRef { + name: variant.clone(), + subscripts: Vec::new(), + span, + }) + .collect(), + is_matrix: false, + span, + }; + } + self.walk_var_ref_expression(name, subscripts, span) + } + + fn walk_binary_expression( + &mut self, + op: &rumoca_core::OpBinary, + lhs: &rumoca_core::Expression, + rhs: &rumoca_core::Expression, + span: rumoca_core::Span, + ) -> rumoca_core::Expression { + let lhs = self.rewrite_expression(lhs); + let rhs = self.rewrite_expression(rhs); + vectorized_binary_expr(op.clone(), lhs, rhs, span) + } + + fn walk_unary_expression( + &mut self, + op: &rumoca_core::OpUnary, + rhs: &rumoca_core::Expression, + span: rumoca_core::Span, + ) -> rumoca_core::Expression { + let rhs = self.rewrite_expression(rhs); + if let rumoca_core::Expression::Array { elements, .. } = rhs { + return rumoca_core::Expression::Array { + elements: elements + .into_iter() + .map(|element| rumoca_core::Expression::Unary { + op: op.clone(), + rhs: Box::new(element), + span, + }) + .collect(), + is_matrix: false, + span, + }; + } + rumoca_core::Expression::Unary { + op: op.clone(), + rhs: Box::new(rhs), + span, + } + } +} + +fn vectorized_binary_expr( + op: rumoca_core::OpBinary, + lhs: rumoca_core::Expression, + rhs: rumoca_core::Expression, + span: rumoca_core::Span, +) -> rumoca_core::Expression { + match (lhs, rhs) { + ( + rumoca_core::Expression::Array { + elements: lhs_values, + .. + }, + rumoca_core::Expression::Array { + elements: rhs_values, + .. + }, + ) if lhs_values.len() == rhs_values.len() => { + array_from_binary_elements(op, lhs_values.into_iter().zip(rhs_values).collect(), span) + } + ( + rumoca_core::Expression::Array { + elements: lhs_values, + .. + }, + rhs, + ) => array_from_binary_elements( + op, + lhs_values + .into_iter() + .map(|lhs| (lhs, rhs.clone())) + .collect(), + span, + ), + ( + lhs, + rumoca_core::Expression::Array { + elements: rhs_values, + .. + }, + ) => array_from_binary_elements( + op, + rhs_values + .into_iter() + .map(|rhs| (lhs.clone(), rhs)) + .collect(), + span, + ), + (lhs, rhs) => rumoca_core::Expression::Binary { + op, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span, + }, + } +} + +fn array_from_binary_elements( + op: rumoca_core::OpBinary, + pairs: Vec<(rumoca_core::Expression, rumoca_core::Expression)>, + span: rumoca_core::Span, +) -> rumoca_core::Expression { + rumoca_core::Expression::Array { elements: pairs .into_iter() .map(|(lhs, rhs)| rumoca_core::Expression::Binary { op: op.clone(), - lhs: Box::new(lhs), - rhs: Box::new(rhs), + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span, + }) + .collect(), + is_matrix: false, + span, + } +} + +fn scalarize_equation_list( + equations: &mut Vec, + phantom_map: &HashMap>, + array_dims: &HashMap>, + var_dims: &HashMap>, + record_array_fields: &RecordArrayFieldMap, + functions: &IndexMap, + recover_discrete_assignments: bool, +) -> Result, ToDaeError> { + let mut new_equations = Vec::with_capacity(equations.len()); + let mut spans = Vec::with_capacity(equations.len()); + for eq in equations.drain(..) { + let new_start = new_equations.len(); + let phantom_width = expr_phantom_ref_width(&eq.rhs, phantom_map); + let effective_scalar_count = if eq.scalar_count > 1 { + eq.scalar_count + } else { + structured_template_residual_scalar_count(&eq.rhs, array_dims) + .unwrap_or(eq.scalar_count) + }; + if effective_scalar_count > 1 + && (phantom_width.is_some() + || expr_has_record_array_member_slice(&eq.rhs) + || expr_has_colon_slice(&eq.rhs) + || expr_has_vectorized_scalar_function_call(&eq.rhs, array_dims, functions)) + { + // Expand into scalar_count individual equations + for k in 0..effective_scalar_count { + let scalar_rhs = scalarize_expr_at( + &eq.rhs, + k, + phantom_map, + array_dims, + var_dims, + record_array_fields, + functions, + )?; + let origin = format!("{} [scalarized {}]", eq.origin, k + 1); + new_equations.push(scalarized_equation_at( + &eq, + scalar_rhs, + k, + origin, + phantom_map, + array_dims, + recover_discrete_assignments, + )?); + } + } else if phantom_width == Some(1) { + let scalar_rhs = scalarize_expr_at( + &eq.rhs, + 0, + phantom_map, + array_dims, + var_dims, + record_array_fields, + functions, + )?; + new_equations.push(dae::Equation { + rhs: scalar_rhs, + ..eq + }); + } else if phantom_width.is_some() { + let rhs = vectorize_phantom_array_formal_args(&eq.rhs, phantom_map, functions); + new_equations.push(dae::Equation { rhs, ..eq }); + } else { + new_equations.push(project_scalarized_residual_rhs( + eq, + array_dims, + var_dims, + record_array_fields, + functions, + )?); + } + spans.push((new_start, new_equations.len() - new_start)); + } + *equations = new_equations; + Ok(spans) +} +fn project_scalarized_residual_rhs( + eq: dae::Equation, + array_dims: &HashMap>, + var_dims: &HashMap>, + record_array_fields: &RecordArrayFieldMap, + functions: &IndexMap, +) -> Result { + if eq.scalar_count != 1 { + return Ok(eq); + } + let rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs, + rhs, + span, + } = &eq.rhs + else { + let rhs = lower_colon_slice_dot_products(&eq.rhs, var_dims)?; + return Ok(dae::Equation { rhs, ..eq }); + }; + let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = lhs.as_ref() + else { + let rhs = lower_colon_slice_dot_products(&eq.rhs, var_dims)?; + return Ok(dae::Equation { rhs, ..eq }); + }; + let Some(target_dims) = Projector(var_dims, functions).dims(lhs, *span)? else { + let rhs = lower_colon_slice_dot_products(&eq.rhs, var_dims)?; + return Ok(dae::Equation { rhs, ..eq }); + }; + let k = if target_dims.is_empty() { + scalarized_lhs_zero_based_index_or_singleton(name, subscripts, array_dims).unwrap_or(0) + } else { + let Some(k) = scalarized_lhs_zero_based_index_or_singleton(name, subscripts, array_dims) + else { + let rhs = lower_colon_slice_dot_products(&eq.rhs, var_dims)?; + return Ok(dae::Equation { rhs, ..eq }); + }; + k + }; + let scalar_lhs = if target_dims.is_empty() { + lhs.as_ref().clone() + } else { + project_scalarized_rhs_expr_at(lhs, k, array_dims, record_array_fields, functions)? + }; + let projector = Projector(var_dims, functions); + let projected = projector.project(rhs, k, &target_dims)?; + let scalar_rhs = match projected { + Some(projected) => projected, + None if target_dims.is_empty() + && !matches!(rhs.as_ref(), Expr::ArrayComprehension { .. }) + && !expr_has_vectorized_scalar_function_call(rhs, array_dims, functions) => + { + rhs.as_ref().clone() + } + None => project_scalarized_rhs_expr_at(rhs, k, array_dims, record_array_fields, functions)?, + }; + let scalar_lhs = lower_colon_slice_dot_products(&scalar_lhs, var_dims)?; + let scalar_rhs = lower_colon_slice_dot_products(&scalar_rhs, var_dims)?; + Ok(dae::Equation { + rhs: rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(scalar_lhs), + rhs: Box::new(scalar_rhs), + span: *span, + }, + ..eq + }) +} + +fn project_scalarized_rhs_expr_at( + expr: &rumoca_core::Expression, + k: usize, + array_dims: &HashMap>, + record_array_fields: &RecordArrayFieldMap, + functions: &IndexMap, +) -> Result { + RhsProjectionCtx { + k, + array_dims, + record_array_fields, + functions, + } + .project(expr) +} + +fn try_project_colon_slice_var_ref( + expr: &rumoca_core::Expression, + array_dims: &HashMap>, +) -> Result, ToDaeError> { + let rumoca_core::Expression::Index { + base, + subscripts, + span, + } = expr + else { + return Ok(None); + }; + let rumoca_core::Expression::VarRef { + name, + subscripts: base_subscripts, + .. + } = base.as_ref() + else { + return Ok(None); + }; + if !base_subscripts.is_empty() || !subscripts_have_colon(subscripts) { + return Ok(None); + } + let Some(dims) = array_dims.get(name.as_str()) else { + return Ok(None); + }; + let Some(projected_dims) = projected_dims_for_subscripts(dims, subscripts) else { + return Ok(None); + }; + let scalar_count = compute_var_size(&projected_dims); + if scalar_count <= 1 { + return Ok(None); + } + let Some(elements) = project_colon_slice_elements(name, dims, subscripts, scalar_count, *span)? + else { + return Ok(None); + }; + Ok(Some(rumoca_core::Expression::Array { + elements, + is_matrix: false, + span: *span, + })) +} + +fn subscripts_have_colon(subscripts: &[rumoca_core::Subscript]) -> bool { + subscripts + .iter() + .any(|subscript| matches!(subscript, rumoca_core::Subscript::Colon { .. })) +} + +fn project_colon_slice_elements( + name: &rumoca_core::Reference, + dims: &[i64], + subscripts: &[rumoca_core::Subscript], + scalar_count: usize, + span: rumoca_core::Span, +) -> Result>, ToDaeError> { + let mut elements = Vec::with_capacity(scalar_count); + for idx in 0..scalar_count { + let Some(projected_subscripts) = + project_slice_subscripts_for_lane(dims, subscripts, idx, span)? + else { + return Ok(None); + }; + elements.push(rumoca_core::Expression::VarRef { + name: name.clone(), + subscripts: projected_subscripts, + span, + }); + } + Ok(Some(elements)) +} + +fn lower_colon_slice_binary_expr( + op: &rumoca_core::OpBinary, + lhs: &rumoca_core::Expression, + rhs: &rumoca_core::Expression, + span: rumoca_core::Span, + array_dims: &HashMap>, +) -> Result { + if !matches!(op, rumoca_core::OpBinary::Mul) { + return Ok(rumoca_core::Expression::Binary { + op: op.clone(), + lhs: Box::new(lower_colon_slice_dot_products(lhs, array_dims)?), + rhs: Box::new(lower_colon_slice_dot_products(rhs, array_dims)?), + span, + }); + } + if is_colon_slice(lhs) || is_colon_slice(rhs) { + let preserve = match ( + classify_dot_operand(lhs, array_dims)?, + classify_dot_operand(rhs, array_dims)?, + ) { + (DotOperand::Vector(projected_lhs), DotOperand::Vector(projected_rhs)) => { + if let ( + Expr::Array { + elements: lhs_values, + .. + }, + Expr::Array { + elements: rhs_values, + .. + }, + ) = (&projected_lhs, &projected_rhs) + && lhs_values.len() == rhs_values.len() + && let Some(dot) = dot_product_expr(lhs_values, rhs_values, span) + { + return Ok(dot); + } + true + } + (DotOperand::Unsafe, _) | (_, DotOperand::Unsafe) => true, + _ => false, + }; + if preserve { + return Ok(rumoca_core::Expression::Binary { + op: op.clone(), + lhs: Box::new(lower_colon_slice_dot_products(lhs, array_dims)?), + rhs: Box::new(lower_colon_slice_dot_products(rhs, array_dims)?), + span, + }); + } + } + let lhs = lower_colon_slice_dot_operand(lhs, array_dims)?; + let rhs = lower_colon_slice_dot_operand(rhs, array_dims)?; + if let ( + Expr::Array { + elements: lhs_values, + .. + }, + Expr::Array { + elements: rhs_values, + .. + }, + ) = (&lhs, &rhs) + && lhs_values.len() == rhs_values.len() + && let Some(dot) = dot_product_expr(lhs_values, rhs_values, span) + { + return Ok(dot); + } + Ok(vectorized_binary_expr(op.clone(), lhs, rhs, span)) +} + +fn lower_colon_slice_dot_operand( + expr: &rumoca_core::Expression, + array_dims: &HashMap>, +) -> Result { + if let Some(projected) = try_project_colon_slice_var_ref(expr, array_dims)? { + return Ok(projected); + } + lower_colon_slice_dot_products(expr, array_dims) +} +fn dot_product_expr(lhs_values: &[Expr], rhs_values: &[Expr], span: Span) -> Option { + lhs_values + .iter() + .cloned() + .zip(rhs_values.iter().cloned()) + .map(|(lhs, rhs)| vectorized_binary_expr(OpBinary::Mul, lhs, rhs, span)) + .reduce(|lhs, rhs| vectorized_binary_expr(OpBinary::Add, lhs, rhs, span)) +} +struct Projector<'a>( + &'a HashMap>, + &'a IndexMap, +); +impl Projector<'_> { + fn dims(&self, expr: &Expr, span: Span) -> Result>, ToDaeError> { + if let Some((name, subscripts)) = matrix_var_slice(expr) { + return match self.0.get(name.as_str()) { + Some(dims) if dims.iter().any(|dim| *dim < 0) => { + Err(projection_error("negative dimension", span)) + } + Some(dims) => Ok(proven_projected_dims(dims, subscripts)), + None => Ok(None), + }; + } + match expr { + Expr::BuiltinCall { + function: Builtin::Fill, + args, + .. + } => literal_fill_dims(args, span), + Expr::BuiltinCall { function, args, .. } + if matches!(function, Builtin::Transpose | Builtin::Der) => + { + let [arg] = args.as_slice() else { + return Err(projection_error("unknown operand shape", span)); + }; + match (function, self.dims(arg, span)?) { + (_, None) => Ok(None), + (Builtin::Der, Some(dims)) => Ok(Some(dims)), + (Builtin::Transpose, Some(dims)) if dims.len() == 2 => { + Ok(Some(vec![dims[1], dims[0]])) + } + _ => Err(projection_error("unsupported rank", span)), + } + } + Expr::BuiltinCall { function, .. } => Ok(builtin_fixed_vector_output_width(*function) + .and_then(|width| i64::try_from(width).ok().map(|width| vec![width]))), + Expr::FunctionCall { name, .. } => Ok(self + .1 + .get(name.var_name()) + .and_then(|f| f.outputs.first()) + .map(|v| v.dims.clone())), + Expr::Literal { .. } => Ok(Some(Vec::new())), + Expr::Binary { op, lhs, rhs, .. } => { + let (Some(lhs_dims), Some(rhs_dims)) = + (self.dims(lhs, span)?, self.dims(rhs, span)?) + else { + return Ok(None); + }; + if matches!(op, OpBinary::Mul | OpBinary::MulElem) { + product_dims(op, &lhs_dims, &rhs_dims, span).map(Some) + } else if (lhs_dims.is_empty() && rhs_dims.is_empty()) + || (matches!(op, OpBinary::Add | OpBinary::Sub) && lhs_dims == rhs_dims) + { + Ok(Some(lhs_dims)) + } else { + Err(projection_error("unknown operand shape", span)) + } + } + _ => Ok(None), + } + } + fn project( + &self, + expr: &Expr, + k: usize, + target_dims: &[i64], + ) -> Result, ToDaeError> { + let Expr::Binary { op, lhs, rhs, span } = expr else { + return Ok(None); + }; + if matches!(op, OpBinary::Add | OpBinary::Sub) { + let has_array_syntax = has_array_slice_syntax(expr); + let result_dims = self.dims(expr, *span)?; + if !has_array_syntax + && (result_dims.is_none() + || target_dims.is_empty() && result_dims.as_ref().is_some_and(Vec::is_empty)) + { + return Ok(None); + } + let result_dims = + result_dims.ok_or_else(|| projection_error("unknown operand shape", *span))?; + if result_dims != target_dims { + return Err(projection_error("result shape mismatch", *span)); + } + let indices = projection_lane_indices(k, target_dims) + .ok_or_else(|| projection_error("result shape mismatch", *span))?; + return self.element(expr, &indices, *span).map(Some); + } + if !matches!(op, OpBinary::Mul | OpBinary::MulElem) { + return Ok(None); + } + if target_dims.iter().any(|dim| *dim < 0) { + return Err(projection_error("negative dimension", *span)); + } + let lhs_dims = self.dims(lhs, *span)?; + let rhs_dims = self.dims(rhs, *span)?; + if scalar_times_unknown_non_product(self, lhs, &lhs_dims, rhs, &rhs_dims)? { + return Ok(None); + } + let has_array = lhs_dims.as_ref().is_some_and(|dims| !dims.is_empty()) + || rhs_dims.as_ref().is_some_and(|dims| !dims.is_empty()) + || has_array_slice_syntax(lhs) + || has_array_slice_syntax(rhs); + if target_dims.is_empty() && !has_array && (lhs_dims.is_none() || rhs_dims.is_none()) { + return Ok(None); + } + let lhs_dims = lhs_dims.ok_or_else(|| projection_error("unknown operand shape", *span))?; + let rhs_dims = rhs_dims.ok_or_else(|| projection_error("unknown operand shape", *span))?; + if (!lhs_dims.is_empty() && matches!(lhs.as_ref(), Expr::FunctionCall { .. })) + || (!rhs_dims.is_empty() && matches!(rhs.as_ref(), Expr::FunctionCall { .. })) + { + return Err(projection_error( + "array-valued function output cannot be projected", + *span, + )); + } + let result_dims = product_dims(op, &lhs_dims, &rhs_dims, *span)?; + if result_dims != target_dims { + let detail = if target_dims.is_empty() { + "non-scalar result in scalar context" + } else { + "result shape mismatch" + }; + return Err(projection_error(detail, *span)); + } + if result_dims.is_empty() && lhs_dims.is_empty() && rhs_dims.is_empty() { + return Ok(None); + } + let indices = projection_lane_indices(k, target_dims) + .ok_or_else(|| projection_error("result shape mismatch", *span))?; + if matches!(op, OpBinary::MulElem) { + return Ok(Some(vectorized_binary_expr( + op.clone(), + self.element(lhs, &indices, *span)?, + self.element(rhs, &indices, *span)?, + *span, + ))); + } + if lhs_dims.is_empty() || rhs_dims.is_empty() { + let lhs_indices: &[i64] = if lhs_dims.is_empty() { &[] } else { &indices }; + let rhs_indices: &[i64] = if rhs_dims.is_empty() { &[] } else { &indices }; + let lhs = self.element(lhs, lhs_indices, *span)?; + let rhs = self.element(rhs, rhs_indices, *span)?; + return Ok(Some(vectorized_binary_expr(op.clone(), lhs, rhs, *span))); + } + let inner = *lhs_dims.last().unwrap(); + if inner < 0 { + return Err(projection_error("negative dimension", *span)); + } + if inner == 0 { + return Ok(Some(real_literal(0.0, *span))); + } + let (mut lhs_values, mut rhs_values) = (Vec::new(), Vec::new()); + let row = (lhs_dims.len() == 2).then(|| indices[0]); + let column = (rhs_dims.len() == 2).then(|| *indices.last().unwrap()); + for inner_index in 1..=inner { + let lhs_indices = row.into_iter().chain([inner_index]).collect::>(); + let rhs_indices = [inner_index].into_iter().chain(column).collect::>(); + lhs_values.push(self.element(lhs, &lhs_indices, *span)?); + rhs_values.push(self.element(rhs, &rhs_indices, *span)?); + } + dot_product_expr(&lhs_values, &rhs_values, *span) + .map(Some) + .ok_or_else(|| projection_error("unsupported inner dimension", *span)) + } + fn element(&self, expr: &Expr, indices: &[i64], span: Span) -> Result { + if !matches!(expr, Expr::Binary { .. }) + && matches!(self.dims(expr, span)?, Some(dims) if dims.is_empty()) + { + return Ok(expr.clone()); + } + if let Some((name, subscripts)) = matrix_var_slice(expr) { + let expr_span = expr.span().unwrap_or(span); + let base_dims = self + .0 + .get(name.as_str()) + .ok_or_else(|| projection_error("unknown operand shape", span))?; + let lane = proven_projected_dims(base_dims, subscripts) + .and_then(|dims| linear_lane_for_indices(indices, &dims)) + .ok_or_else(|| projection_error("result shape mismatch", span))?; + let projected = + project_slice_subscripts_for_lane(base_dims, subscripts, lane, expr_span)? + .ok_or_else(|| projection_error("unknown operand shape", span))?; + return Ok(Expr::VarRef { + name: name.clone(), + subscripts: projected, + span: expr_span, + }); + } + match expr { + Expr::BuiltinCall { + function: Builtin::Der, + args, + span, + } if args.len() == 1 => Ok(Expr::BuiltinCall { + function: Builtin::Der, + args: vec![self.element(&args[0], indices, *span)?], + span: *span, + }), + Expr::BuiltinCall { + function: Builtin::Transpose, + args, + span, + } if indices.len() == 2 && args.len() == 1 => { + self.element(&args[0], &[indices[1], indices[0]], *span) + } + Expr::BuiltinCall { + function: Builtin::Fill, + args, + span, + } => self.fill_element(args, indices, *span), + Expr::BuiltinCall { function, span, .. } + if builtin_fixed_vector_output_width(*function).is_some() && indices.len() == 1 => + { + Ok(Expr::Index { + base: Box::new(expr.clone()), + subscripts: vec![generated_index_subscript( + indices[0], + *span, + "DAE matrix-product builtin projection", + )?], + span: *span, + }) + } + Expr::Binary { op, lhs, rhs, span } if matches!(op, OpBinary::Add | OpBinary::Sub) => { + Ok(vectorized_binary_expr( + op.clone(), + self.element(lhs, indices, *span)?, + self.element(rhs, indices, *span)?, + *span, + )) + } + Expr::Binary { .. } => { + let dims = self + .dims(expr, span)? + .ok_or_else(|| projection_error("unknown operand shape", span))?; + let lane = linear_lane_for_indices(indices, &dims) + .ok_or_else(|| projection_error("result shape mismatch", span))?; + self.project(expr, lane, &dims)? + .or_else(|| dims.is_empty().then(|| expr.clone())) + .ok_or_else(|| projection_error("array operand cannot be projected", span)) + } + Expr::FunctionCall { span, .. } => { + let dims = self + .dims(expr, *span)? + .ok_or_else(|| projection_error("unknown operand shape", *span))?; + let lane = linear_lane_for_indices(indices, &dims) + .ok_or_else(|| projection_error("result shape mismatch", *span))?; + let index = one_based_scalar_index( + lane, + *span, + "DAE matrix-product function output projection", + )?; + Ok(Expr::Index { + base: Box::new(expr.clone()), + subscripts: vec![generated_index_subscript( + index, + *span, + "DAE matrix-product function output projection", + )?], + span: *span, + }) + } + _ => Err(projection_error("unknown operand shape", span)), + } + } + + fn fill_element(&self, args: &[Expr], indices: &[i64], span: Span) -> Result { + let (value, _) = args + .split_first() + .ok_or_else(|| projection_error("unknown operand shape", span))?; + let dims = literal_fill_dims(args, span)? + .ok_or_else(|| projection_error("unknown operand shape", span))?; + linear_lane_for_indices(indices, &dims) + .ok_or_else(|| projection_error("result shape mismatch", span))?; + self.element(value, &[], span) + } + + fn has_descendant_matrix_product_candidate( + &self, + expr: &Expr, + allow_non_scalar_evidence: bool, + ) -> Result { + match expr { + Expr::Binary { + op: OpBinary::Mul | OpBinary::MulElem, + lhs, + rhs, + span, + } => { + let operands = [lhs.as_ref(), rhs.as_ref()]; + let operand_dims = operands + .iter() + .map(|operand| match operand { + Expr::Binary { op, .. } + if !matches!( + op, + OpBinary::Add | OpBinary::Sub | OpBinary::Mul | OpBinary::MulElem + ) => + { + Ok(None) + } + _ => self.dims(operand, *span), + }) + .collect::, _>>()?; + if operand_dims.iter().flatten().any(Vec::is_empty) { + return operands + .into_iter() + .map(|operand| self.has_descendant_matrix_product_candidate(operand, false)) + .collect::, _>>() + .map(|candidates| candidates.into_iter().any(std::convert::identity)); + } + Ok(operand_dims.iter().flatten().any(|dims| !dims.is_empty()) + || operands + .into_iter() + .zip(operand_dims) + .filter(|(_, dims)| dims.is_none()) + .map(|(operand, _)| { + Ok::( + has_array_slice_syntax(operand) + || self + .has_descendant_matrix_product_candidate(operand, true)?, + ) + }) + .collect::, _>>()? + .into_iter() + .any(std::convert::identity)) + } + Expr::Binary { lhs, rhs, .. } => Ok(self + .has_descendant_matrix_product_candidate(lhs, allow_non_scalar_evidence)? + || self.has_descendant_matrix_product_candidate(rhs, allow_non_scalar_evidence)?), + Expr::Unary { rhs, .. } => { + self.has_descendant_matrix_product_candidate(rhs, allow_non_scalar_evidence) + } + Expr::BuiltinCall { function, args, .. } + if !builtin_is_scalar_boundary(function, args.len()) => + { + args.iter() + .map(|arg| { + self.has_descendant_matrix_product_candidate(arg, allow_non_scalar_evidence) + }) + .collect::, _>>() + .map(|candidates| candidates.into_iter().any(std::convert::identity)) + } + Expr::If { + branches, + else_branch, + .. + } => { + let branch_candidates = branches + .iter() + .map(|(_, value)| { + self.has_descendant_matrix_product_candidate( + value, + allow_non_scalar_evidence, + ) + }) + .collect::, _>>()?; + Ok(branch_candidates.into_iter().any(std::convert::identity) + || self.has_descendant_matrix_product_candidate( + else_branch, + allow_non_scalar_evidence, + )?) + } + _ if allow_non_scalar_evidence => { + let Some(span) = expr.span() else { + return Ok(false); + }; + Ok(has_array_slice_syntax(expr) + || matches!(self.dims(expr, span)?, Some(dims) if !dims.is_empty())) + } + _ => Ok(false), + } + } +} +fn matrix_var_slice(expr: &Expr) -> Option<(&Reference, &[Subscript])> { + match expr { + Expr::VarRef { + name, subscripts, .. + } => Some((name, subscripts)), + Expr::Index { + base, subscripts, .. + } => matrix_var_slice(base) + .filter(|(_, base_subscripts)| base_subscripts.is_empty()) + .map(|(name, _)| (name, subscripts.as_slice())), + _ => None, + } +} +fn has_array_slice_syntax(expr: &Expr) -> bool { + match expr { + Expr::VarRef { .. } | Expr::Index { .. } => { + matrix_var_slice(expr).is_some_and(|(_, subscripts)| { + subscripts_have_colon(subscripts) + || subscripts.iter().any( + |subscript| matches!(subscript, Subscript::Expr { expr, .. } if subscript_expr_selects_vector(expr) == Some(true)), + ) + }) + } + Expr::Unary { rhs, .. } => has_array_slice_syntax(rhs), + Expr::BuiltinCall { function, args, .. } + if !builtin_is_scalar_boundary(function, args.len()) => + { + args.iter().any(has_array_slice_syntax) + } + Expr::Binary { lhs, rhs, .. } => { + has_array_slice_syntax(lhs) || has_array_slice_syntax(rhs) + } + Expr::If { + branches, + else_branch, + .. + } => branches + .iter() + .map(|(_, value)| value) + .chain([else_branch.as_ref()]) + .any(has_array_slice_syntax), + _ => false, + } +} +fn builtin_is_scalar_boundary(function: &Builtin, arity: usize) -> bool { + matches!( + function, + Builtin::Sum | Builtin::Product | Builtin::Scalar | Builtin::Ndims | Builtin::Size + ) || arity == 1 && matches!(function, Builtin::Min | Builtin::Max) +} +fn literal_fill_dims(args: &[Expr], span: Span) -> Result>, ToDaeError> { + let Some((_, dimension_args)) = args.split_first().filter(|(_, dims)| !dims.is_empty()) else { + return Err(projection_error("unknown operand shape", span)); + }; + let Some(dims) = dimension_args + .iter() + .map(integer_literal_value) + .collect::>>() + else { + return Ok(None); + }; + if dims.iter().any(|dim| *dim < 0) { + return Err(projection_error("negative dimension", span)); + } + Ok(Some(dims)) +} +fn scalar_times_unknown_non_product( + projector: &Projector<'_>, + lhs: &Expr, + lhs_dims: &Option>, + rhs: &Expr, + rhs_dims: &Option>, +) -> Result { + let lhs_unknown_non_product = + lhs_dims.is_none() && !projector.has_descendant_matrix_product_candidate(lhs, false)?; + let rhs_unknown_non_product = + rhs_dims.is_none() && !projector.has_descendant_matrix_product_candidate(rhs, false)?; + let result = lhs_dims.as_ref().is_some_and(Vec::is_empty) && rhs_unknown_non_product + || rhs_dims.as_ref().is_some_and(Vec::is_empty) && lhs_unknown_non_product; + Ok(result) +} +fn product_dims(op: &OpBinary, lhs: &[i64], rhs: &[i64], at: Span) -> Result, ToDaeError> { + if matches!(op, OpBinary::MulElem) { + return (lhs == rhs) + .then(|| lhs.to_vec()) + .ok_or_else(|| projection_error("elementwise shape mismatch", at)); + } + match (lhs.len(), rhs.len()) { + (0, _) => Ok(rhs.to_vec()), + (_, 0) => Ok(lhs.to_vec()), + (l, r) if l > 2 || r > 2 => Err(projection_error("unsupported rank", at)), + _ if lhs.last() != rhs.first() => Err(projection_error("inner dimension mismatch", at)), + (1, 1) => Ok(Vec::new()), + (2, 1) => Ok(vec![lhs[0]]), + (1, 2) => Ok(vec![rhs[1]]), + (2, 2) => Ok(vec![lhs[0], rhs[1]]), + _ => unreachable!(), + } +} +fn proven_projected_dims(dims: &[i64], subscripts: &[Subscript]) -> Option> { + (subscripts.len() <= dims.len()).then_some(())?; + projected_dims_for_subscripts(dims, subscripts) +} +fn linear_lane_for_indices(indices: &[i64], dims: &[i64]) -> Option { + (indices.len() == dims.len()).then_some(())?; + indices + .iter() + .zip(dims) + .try_fold(0usize, |lane, (index, dim)| { + let dim = usize::try_from(*dim).ok()?; + let index = usize::try_from(index.checked_sub(1)?).ok()?; + (index < dim).then_some(lane.checked_mul(dim)?.checked_add(index)?) + }) +} +fn projection_error(why: &str, span: Span) -> ToDaeError { + ToDaeError::runtime_contract_violation_at(format!("DAE matrix-product projection: {why}"), span) +} +fn project_slice_subscripts_for_lane( + dims: &[i64], + subscripts: &[rumoca_core::Subscript], + k: usize, + span: rumoca_core::Span, +) -> Result>, ToDaeError> { + let Some(selected_dims) = + projected_dims_for_subscripts(dims, subscripts).filter(|dims| !dims.is_empty()) + else { + return Ok(None); + }; + let Some(lane_indices) = lane_indices_for_dims(k, &selected_dims) else { + return Ok(None); + }; + let mut lane_iter = lane_indices.into_iter(); + let mut projected = Vec::with_capacity(dims.len()); + for subscript in subscripts { + match subscript { + rumoca_core::Subscript::Colon { .. } => { + let Some(index) = lane_iter.next() else { + return Ok(None); + }; + projected.push(generated_index_subscript( + index, + span, + "DAE scalarized colon slice projection", + )?); + } + rumoca_core::Subscript::Index { .. } | rumoca_core::Subscript::Expr { .. } => { + projected.push(subscript.clone()); + } + } + } + for _ in subscripts.len()..dims.len() { + let Some(index) = lane_iter.next() else { + return Ok(None); + }; + projected.push(generated_index_subscript( + index, + span, + "DAE scalarized trailing array projection", + )?); + } + Ok(Some(projected)) +} +fn lane_indices_for_dims(k: usize, dims: &[i64]) -> Option> { + let mut remaining = k; + let mut indices = Vec::with_capacity(dims.len()); + for dim in dims.iter().rev() { + let dim = usize::try_from(*dim).ok()?; + if dim == 0 { + return None; + } + indices.push(i64::try_from(remaining % dim + 1).ok()?); + remaining /= dim; + } + indices.reverse(); + (remaining == 0).then_some(indices) +} +fn projection_lane_indices(k: usize, dims: &[i64]) -> Option> { + if dims.is_empty() { + return Some(Vec::new()); + } + lane_indices_for_dims(k, dims) +} +struct RhsProjectionCtx<'a> { + k: usize, + array_dims: &'a HashMap>, + record_array_fields: &'a RecordArrayFieldMap, + functions: &'a IndexMap, +} +impl RhsProjectionCtx<'_> { + fn project( + &self, + expr: &rumoca_core::Expression, + ) -> Result { + match expr { + rumoca_core::Expression::VarRef { .. } => self.project_var_ref(expr), + rumoca_core::Expression::Index { + base, + subscripts, + span, + } => self.project_index(base, subscripts, *span, expr), + rumoca_core::Expression::Array { elements, .. } => self.project_array(elements, expr), + rumoca_core::Expression::ArrayComprehension { + expr: inner, + indices, + filter, + span, + } => self.project_array_comprehension(inner, indices, filter.as_deref(), *span, expr), + rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor, + span, + } => self.project_function_call(name, args, *is_constructor, *span), + rumoca_core::Expression::FieldAccess { base, field, span } => { + self.project_field_access(base, field, *span) + } + rumoca_core::Expression::Binary { op, lhs, rhs, span } => { + self.project_binary(op, lhs, rhs, *span) + } + rumoca_core::Expression::Unary { op, rhs, span } => self.project_unary(op, rhs, *span), + rumoca_core::Expression::If { + branches, + else_branch, + span, + } => self.project_if(branches, else_branch, *span), + _ => Ok(expr.clone()), + } + } + fn project_var_ref( + &self, + expr: &rumoca_core::Expression, + ) -> Result { + let rumoca_core::Expression::VarRef { + name, + subscripts, + span, + } = expr + else { + return Ok(expr.clone()); + }; + if subscripts.is_empty() && self.array_dims.contains_key(name.as_str()) { + let index = one_based_scalar_index(self.k, *span, "DAE scalar lhs RHS projection")?; + return Ok(rumoca_core::Expression::VarRef { + name: name.clone(), + subscripts: vec![generated_index_subscript( + index, + *span, + "DAE scalar lhs RHS projection", + )?], + span: *span, + }); + } + if subscripts.is_empty() + && let Some(projected) = project_record_array_field_entry_at( + self.record_array_fields.get(name.as_str()), + self.k, + *span, + )? + { + return Ok(projected); + } + Ok(expr.clone()) + } + fn project_index( + &self, + base: &rumoca_core::Expression, + subscripts: &[rumoca_core::Subscript], + span: rumoca_core::Span, + expr: &rumoca_core::Expression, + ) -> Result { + let rumoca_core::Expression::VarRef { + name, + subscripts: base_subscripts, + .. + } = base + else { + return Ok(expr.clone()); + }; + if !base_subscripts.is_empty() { + return Ok(expr.clone()); + } + let Some(dims) = self.array_dims.get(name.as_str()) else { + return Ok(expr.clone()); + }; + let Some(projected_subscripts) = + project_slice_subscripts_for_lane(dims, subscripts, self.k, span)? + else { + return Ok(expr.clone()); + }; + Ok(rumoca_core::Expression::VarRef { + name: name.clone(), + subscripts: projected_subscripts, + span, + }) + } + fn project_array( + &self, + elements: &[rumoca_core::Expression], + expr: &rumoca_core::Expression, + ) -> Result { + if let Some((element_index, element_lane)) = + scalarized_array_literal_lane(elements, self.k, self.array_dims) + { + return RhsProjectionCtx { + k: element_lane, + ..*self + } + .project(&elements[element_index]); + } + Ok(expr.clone()) + } + fn project_array_comprehension( + &self, + inner: &rumoca_core::Expression, + indices: &[rumoca_core::ComprehensionIndex], + filter: Option<&rumoca_core::Expression>, + span: rumoca_core::Span, + expr: &rumoca_core::Expression, + ) -> Result { + if filter.is_some() || indices.len() != 1 { + return Ok(expr.clone()); + } + let Some(value) = scalarized_comprehension_index_value(&indices[0].range, self.k) else { + return Ok(expr.clone()); + }; + let mut substitution = ComprehensionIndexSubstitution { + name: indices[0].name.clone(), + value, + span: indices[0].range.span().unwrap_or(span), + }; + let selected = substitution.rewrite_expression(inner); + self.project(&selected) + } + fn project_function_call( + &self, + name: &rumoca_core::Reference, + args: &[rumoca_core::Expression], + is_constructor: bool, + span: rumoca_core::Span, + ) -> Result { + let function = self.functions.get(name.var_name()); + if scalar_output_function(function) { + return self.project_function_call_with_formals(name, args, false, span, function); + } + if let Some(function) = function { + return self.project_function_call_with_formals( + name, + args, + is_constructor, span, + Some(function), + ); + } + let projected_args = args + .iter() + .map(|arg| self.project(arg)) + .collect::, ToDaeError>>()?; + Ok(rumoca_core::Expression::FunctionCall { + name: name.clone(), + args: projected_args, + is_constructor, + span, + }) + } + fn project_function_call_with_formals( + &self, + name: &rumoca_core::Reference, + args: &[rumoca_core::Expression], + is_constructor: bool, + span: rumoca_core::Span, + function: Option<&rumoca_core::Function>, + ) -> Result { + let mut positional_idx = 0usize; + let projected_args = args + .iter() + .map(|arg| { + let formal_rank = + scalarized_function_arg_formal_rank(function, arg, &mut positional_idx); + project_scalarized_function_arg_at( + arg, + formal_rank, + self.k, + self.array_dims, + self.record_array_fields, + self.functions, + ) }) - .collect(), - is_matrix: false, - span, + .collect::, ToDaeError>>()?; + Ok(rumoca_core::Expression::FunctionCall { + name: name.clone(), + args: projected_args, + is_constructor, + span, + }) + } + fn project_field_access( + &self, + base: &rumoca_core::Expression, + field: &str, + span: rumoca_core::Span, + ) -> Result { + if let Some(projected) = scalarize_record_array_member_slice_at( + base, + field, + span, + self.k, + self.record_array_fields, + )? { + return Ok(projected); + } + if let Some(projected) = + project_record_array_field_rhs_at(base, field, span, self.k, self.record_array_fields)? + { + return Ok(projected); + } + Ok(rumoca_core::Expression::FieldAccess { + base: Box::new(self.project(base)?), + field: field.to_string(), + span, + }) + } + fn project_binary( + &self, + op: &rumoca_core::OpBinary, + lhs: &rumoca_core::Expression, + rhs: &rumoca_core::Expression, + span: rumoca_core::Span, + ) -> Result { + Ok(rumoca_core::Expression::Binary { + op: op.clone(), + lhs: Box::new(self.project(lhs)?), + rhs: Box::new(self.project(rhs)?), + span, + }) + } + fn project_unary( + &self, + op: &rumoca_core::OpUnary, + rhs: &rumoca_core::Expression, + span: rumoca_core::Span, + ) -> Result { + Ok(rumoca_core::Expression::Unary { + op: op.clone(), + rhs: Box::new(self.project(rhs)?), + span, + }) + } + fn project_if( + &self, + branches: &[(rumoca_core::Expression, rumoca_core::Expression)], + else_branch: &rumoca_core::Expression, + span: rumoca_core::Span, + ) -> Result { + let branches = branches + .iter() + .map(|(condition, value)| Ok((condition.clone(), self.project(value)?))) + .collect::, ToDaeError>>()?; + Ok(rumoca_core::Expression::If { + branches, + else_branch: Box::new(self.project(else_branch)?), + span, + }) } } - -/// Process an equation list, expanding vector equations with phantom refs. -/// Scalarize phantom/comprehension array equations in place, returning one -/// `(new_start, new_len)` span per input equation (indexed by pre-expansion -/// position). An array equation that expands into `scalar_count` rows reports -/// `new_len == scalar_count`; every other equation reports `new_len == 1`. The -/// spans feed [`rumoca_ir_dae::remap_structured_families_after_expansion`] so -/// structured families stay pointed at their post-expansion row blocks. -fn scalarize_equation_list( - equations: &mut Vec, - phantom_map: &HashMap>, +fn scalarized_array_literal_lane( + elements: &[rumoca_core::Expression], + k: usize, + array_dims: &HashMap>, +) -> Option<(usize, usize)> { + if elements.is_empty() { + return None; + } + let widths = elements + .iter() + .map(|element| scalarized_array_literal_element_width(element, array_dims)) + .collect::>>()?; + let width = *widths.first()?; + if width == 0 || widths.iter().any(|candidate| *candidate != width) { + return None; + } + let element_index = k / width; + if element_index >= elements.len() { + return None; + } + Some((element_index, k % width)) +} +fn scalarized_array_literal_element_width( + element: &rumoca_core::Expression, + array_dims: &HashMap>, +) -> Option { + match element { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } if subscripts.is_empty() => array_dims + .get(name.as_str()) + .map(|dims| compute_var_size(dims)) + .or(Some(1)), + _ => Some(1), + } +} +fn scalar_output_function(function: Option<&rumoca_core::Function>) -> bool { + matches!(function.and_then(|function| function.outputs.first()), Some(output) if output.dims.is_empty()) +} +fn scalarized_function_arg_formal_rank( + function: Option<&rumoca_core::Function>, + arg: &rumoca_core::Expression, + positional_idx: &mut usize, +) -> usize { + let Some(function) = function else { + return 0; + }; + if let Some(input_name) = named_argument_input_name(arg) { + return function + .inputs + .iter() + .find(|input| input.name == input_name) + .map_or(0, |input| input.dims.len()); + } + let rank = function + .inputs + .get(*positional_idx) + .map_or(0, |input| input.dims.len()); + *positional_idx += 1; + rank +} +fn project_scalarized_function_arg_at( + arg: &rumoca_core::Expression, + formal_rank: usize, + k: usize, array_dims: &HashMap>, + record_array_fields: &RecordArrayFieldMap, functions: &IndexMap, -) -> Result, ToDaeError> { - let mut new_equations = Vec::with_capacity(equations.len()); - let mut spans = Vec::with_capacity(equations.len()); - for eq in equations.drain(..) { - let new_start = new_equations.len(); - let phantom_width = expr_phantom_ref_width(&eq.rhs, phantom_map); - if eq.scalar_count > 1 && (phantom_width.is_some() || expr_has_array_comprehension(&eq.rhs)) - { - // Expand into scalar_count individual equations - for k in 0..eq.scalar_count { - let scalar_rhs = scalarize_expr_at(&eq.rhs, k, phantom_map, array_dims, functions)?; - let origin = format!("{} [scalarized {}]", eq.origin, k + 1); - new_equations.push(scalarized_equation_at( - &eq, - scalar_rhs, - k, - origin, - phantom_map, - array_dims, - )?); - } - } else if phantom_width == Some(1) { - let scalar_rhs = scalarize_expr_at(&eq.rhs, 0, phantom_map, array_dims, functions)?; - new_equations.push(dae::Equation { - rhs: scalar_rhs, - ..eq - }); - } else if phantom_width.is_some() { - let rhs = vectorize_phantom_array_formal_args(&eq.rhs, phantom_map, functions); - new_equations.push(dae::Equation { rhs, ..eq }); - } else { - new_equations.push(eq); - } - spans.push((new_start, new_equations.len() - new_start)); +) -> Result { + if let Some((_, value)) = named_function_arg_value(arg) + && let rumoca_core::Expression::FunctionCall { + name, + is_constructor, + span, + .. + } = arg + { + return Ok(rumoca_core::Expression::FunctionCall { + name: name.clone(), + args: vec![project_scalarized_function_arg_at( + value, + formal_rank, + k, + array_dims, + record_array_fields, + functions, + )?], + is_constructor: *is_constructor, + span: *span, + }); } - *equations = new_equations; - Ok(spans) + if let rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor: false, + span, + } = arg + && is_stream_passthrough_intrinsic_dae(name.as_str()) + { + let projected_args = match args.as_slice() { + [inner] => vec![project_scalarized_function_arg_at( + inner, + formal_rank, + k, + array_dims, + record_array_fields, + functions, + )?], + _ => args.clone(), + }; + return Ok(rumoca_core::Expression::FunctionCall { + name: name.clone(), + args: projected_args, + is_constructor: false, + span: *span, + }); + } + let rumoca_core::Expression::VarRef { + name, + subscripts, + span, + } = arg + else { + return project_scalarized_rhs_expr_at(arg, k, array_dims, record_array_fields, functions); + }; + if !subscripts.is_empty() { + return Ok(arg.clone()); + } + if let Some(projected) = + project_record_array_field_entry_at(record_array_fields.get(name.as_str()), k, *span)? + { + return Ok(projected); + } + let Some(dims) = array_dims.get(name.as_str()) else { + return Ok(arg.clone()); + }; + if dims.len() == formal_rank { + return Ok(arg.clone()); + } + if dims.len() != formal_rank + 1 { + return project_scalarized_rhs_expr_at(arg, k, array_dims, record_array_fields, functions); + } + let index = one_based_scalar_index(k, *span, "DAE scalarized function argument projection")?; + let index_subscript = + generated_index_subscript(index, *span, "DAE scalarized function argument projection")?; + if formal_rank == 0 { + return Ok(rumoca_core::Expression::VarRef { + name: name.clone(), + subscripts: vec![index_subscript], + span: *span, + }); + } + let mut projected_subscripts = Vec::with_capacity(dims.len()); + projected_subscripts.push(index_subscript); + for _ in 0..formal_rank { + projected_subscripts.push(generated_colon_subscript( + *span, + "DAE scalarized function argument projection", + )?); + } + Ok(rumoca_core::Expression::Index { + base: Box::new(rumoca_core::Expression::VarRef { + name: name.clone(), + subscripts: Vec::new(), + span: *span, + }), + subscripts: projected_subscripts, + span: *span, + }) +} +fn is_stream_passthrough_intrinsic_dae(name: &str) -> bool { + rumoca_core::qualified_type_name_matches(name, "actualStream") + || rumoca_core::qualified_type_name_matches(name, "inStream") +} +fn scalarized_lhs_zero_based_index_or_singleton( + name: &rumoca_core::Reference, + subscripts: &[rumoca_core::Subscript], + array_dims: &HashMap>, +) -> Option { + let dims = array_dims.get(name.as_str())?; + if subscripts.is_empty() { + return (dims.iter().copied().product::() == 1).then_some(0); + } + (subscripts.len() == dims.len()).then_some(())?; + let indices = subscripts + .iter() + .map(|subscript| match subscript { + rumoca_core::Subscript::Index { value, .. } => Some(*value), + rumoca_core::Subscript::Expr { expr, .. } => integer_literal_value(expr), + rumoca_core::Subscript::Colon { .. } => None, + }) + .collect::>>()?; + linear_lane_for_indices(&indices, dims) } - fn scalarized_equation_at( eq: &dae::Equation, scalar_rhs: rumoca_core::Expression, @@ -1866,8 +4971,14 @@ fn scalarized_equation_at( origin: String, phantom_map: &HashMap>, array_dims: &HashMap>, + recover_discrete_assignments: bool, ) -> Result { let Some(lhs) = &eq.lhs else { + if recover_discrete_assignments + && let Some((lhs, rhs)) = residual_assignment_parts(scalar_rhs.clone())? + { + return Ok(dae::Equation::explicit(lhs, rhs, eq.span, origin)); + } return Ok(dae::Equation::residual(scalar_rhs, eq.span, origin)); }; Ok(dae::Equation::explicit( @@ -1877,6 +4988,97 @@ fn scalarized_equation_at( origin, )) } +fn residual_assignment_parts( + expr: rumoca_core::Expression, +) -> Result, ToDaeError> { + let rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs, + rhs, + .. + } = expr + else { + return Ok(None); + }; + let Some(lhs) = assignment_target_reference(*lhs)? else { + return Ok(None); + }; + Ok(Some((lhs, *rhs))) +} +fn assignment_target_reference( + expr: rumoca_core::Expression, +) -> Result, ToDaeError> { + match expr { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } => { + if subscripts.is_empty() { + return Ok(Some(name)); + } + if let Some(mut component_ref) = name.component_ref().cloned() { + let Some(part) = component_ref.parts.last_mut() else { + return Ok(None); + }; + part.subs.extend(subscripts); + return Ok(Some(rumoca_core::Reference::from_component_reference( + component_ref, + ))); + } + let Some(rendered) = render_literal_subscripted_reference(name.as_str(), &subscripts) + else { + return Ok(None); + }; + if name.is_generated() { + Ok(Some(rumoca_core::Reference::generated(rendered))) + } else { + Ok(Some(rumoca_core::Reference::new(rendered))) + } + } + rumoca_core::Expression::FieldAccess { base, field, .. } => { + let Some(base) = assignment_target_reference(*base)? else { + return Ok(None); + }; + Ok(Some(base.with_appended_field(&field))) + } + rumoca_core::Expression::Index { + base, subscripts, .. + } => { + let Some(mut base) = assignment_target_reference(*base)? else { + return Ok(None); + }; + for subscript in subscripts { + let rumoca_core::Subscript::Index { value, span } = subscript else { + return Ok(None); + }; + base = base.with_appended_index( + value, + span.require_provenance("DAE assignment target index") + .map_err(|error| { + ToDaeError::runtime_metadata_violation(error.to_string()) + })?, + ); + } + Ok(Some(base)) + } + _ => Ok(None), + } +} + +fn render_literal_subscripted_reference( + base: &str, + subscripts: &[rumoca_core::Subscript], +) -> Option { + let mut rendered = base.to_string(); + for subscript in subscripts { + let rumoca_core::Subscript::Index { value, .. } = subscript else { + return None; + }; + rendered.push('['); + rendered.push_str(&value.to_string()); + rendered.push(']'); + } + Some(rendered) +} fn scalarize_lhs_name_at( name: &rumoca_core::VarName, diff --git a/crates/rumoca-phase-dae/src/dae_lowering/colon_slice_dot.rs b/crates/rumoca-phase-dae/src/dae_lowering/colon_slice_dot.rs new file mode 100644 index 000000000..ecf545bd0 --- /dev/null +++ b/crates/rumoca-phase-dae/src/dae_lowering/colon_slice_dot.rs @@ -0,0 +1,147 @@ +use super::*; +use rumoca_core::FallibleExpressionRewriter; + +pub(super) enum DotOperand { + Scalar, + Vector(rumoca_core::Expression), + Unsafe, +} + +pub(super) fn is_colon_slice(expr: &rumoca_core::Expression) -> bool { + matches!(expr, rumoca_core::Expression::Index { subscripts, .. } if subscripts_have_colon(subscripts)) +} + +pub(super) fn classify_dot_operand( + expr: &rumoca_core::Expression, + array_dims: &HashMap>, +) -> Result { + let (name, subscripts, span) = match expr { + rumoca_core::Expression::VarRef { + name, + subscripts, + span, + } => (name, subscripts.as_slice(), *span), + rumoca_core::Expression::Index { + base, + subscripts, + span, + } => match base.as_ref() { + rumoca_core::Expression::VarRef { + name, + subscripts: base_subscripts, + .. + } if base_subscripts.is_empty() => (name, subscripts.as_slice(), *span), + _ => return Ok(DotOperand::Unsafe), + }, + rumoca_core::Expression::Array { + elements, + is_matrix: false, + .. + } => { + for element in elements { + if !matches!( + classify_dot_operand(element, array_dims)?, + DotOperand::Scalar + ) { + return Ok(DotOperand::Unsafe); + } + } + return Ok(DotOperand::Vector(expr.clone())); + } + rumoca_core::Expression::Array { .. } => return Ok(DotOperand::Unsafe), + rumoca_core::Expression::Literal { .. } => return Ok(DotOperand::Scalar), + rumoca_core::Expression::Unary { rhs, .. } => { + return Ok(scalar_dot_operand(classify_dot_operand(rhs, array_dims)?)); + } + rumoca_core::Expression::Binary { lhs, rhs, .. } => { + return Ok(scalar_dot_operand_pair( + classify_dot_operand(lhs, array_dims)?, + classify_dot_operand(rhs, array_dims)?, + )); + } + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + for (condition, value) in branches { + if !matches!( + classify_dot_operand(condition, array_dims)?, + DotOperand::Scalar + ) || !matches!(classify_dot_operand(value, array_dims)?, DotOperand::Scalar) + { + return Ok(DotOperand::Unsafe); + } + } + return Ok(scalar_dot_operand(classify_dot_operand( + else_branch, + array_dims, + )?)); + } + _ => return Ok(DotOperand::Unsafe), + }; + let Some(dims) = array_dims.get(name.as_str()) else { + return Ok(DotOperand::Unsafe); + }; + let Some(projected_dims) = projected_dims_for_subscripts(dims, subscripts) else { + return Ok(DotOperand::Unsafe); + }; + match projected_dims.as_slice() { + [] => Ok(DotOperand::Scalar), + [_] => { + let Some(elements) = project_colon_slice_elements( + name, + dims, + subscripts, + compute_var_size(&projected_dims), + span, + )? + else { + return Ok(DotOperand::Unsafe); + }; + Ok(DotOperand::Vector(rumoca_core::Expression::Array { + elements, + is_matrix: false, + span, + })) + } + _ => Ok(DotOperand::Unsafe), + } +} + +fn scalar_dot_operand(operand: DotOperand) -> DotOperand { + scalar_dot_operand_pair(operand, DotOperand::Scalar) +} + +fn scalar_dot_operand_pair(lhs: DotOperand, rhs: DotOperand) -> DotOperand { + if matches!(lhs, DotOperand::Scalar) && matches!(rhs, DotOperand::Scalar) { + DotOperand::Scalar + } else { + DotOperand::Unsafe + } +} + +pub(super) fn lower_colon_slice_dot_products( + expr: &rumoca_core::Expression, + array_dims: &HashMap>, +) -> Result { + ColonSliceDotLowerer { array_dims }.rewrite_expression(expr) +} + +struct ColonSliceDotLowerer<'a> { + array_dims: &'a HashMap>, +} + +impl FallibleExpressionRewriter for ColonSliceDotLowerer<'_> { + type Error = ToDaeError; + + fn walk_binary_expression( + &mut self, + op: &rumoca_core::OpBinary, + lhs: &rumoca_core::Expression, + rhs: &rumoca_core::Expression, + span: rumoca_core::Span, + ) -> Result { + lower_colon_slice_binary_expr(op, lhs, rhs, span, self.array_dims) + } +} diff --git a/crates/rumoca-phase-dae/src/dae_lowering/record_field_inference.rs b/crates/rumoca-phase-dae/src/dae_lowering/record_field_inference.rs new file mode 100644 index 000000000..ea9fd2958 --- /dev/null +++ b/crates/rumoca-phase-dae/src/dae_lowering/record_field_inference.rs @@ -0,0 +1,383 @@ +use indexmap::IndexMap; +use std::collections::{BTreeMap, HashMap, HashSet}; + +use rumoca_core::ExpressionVisitor; + +pub(super) type FieldUseMap = BTreeMap>>; + +pub(super) fn infer_record_fields_by_function( + functions: &IndexMap, + record_fields_by_type: &HashMap>, +) -> HashMap { + let mut fields = HashMap::new(); + for (name, function) in functions { + fields.insert( + name.as_str().to_string(), + collect_local_record_field_uses(function, record_fields_by_type), + ); + } + + let mut changed = true; + while changed { + changed = false; + for (name, function) in functions { + let propagated = collect_callee_record_field_uses(function, functions, &fields); + let entry = fields.entry(name.as_str().to_string()).or_default(); + if merge_field_use_map(entry, propagated) { + changed = true; + } + } + } + fields +} + +fn collect_local_record_field_uses( + function: &rumoca_core::Function, + record_fields_by_type: &HashMap>, +) -> FieldUseMap { + let prefixes = record_field_prefixes(function, record_fields_by_type); + let mut collector = RecordFieldUseCollector { + prefixes: &prefixes, + fields: BTreeMap::new(), + }; + for statement in &function.body { + visit_statement_expressions(statement, &mut collector); + } + collector.fields +} + +fn record_field_prefixes( + function: &rumoca_core::Function, + record_fields_by_type: &HashMap>, +) -> HashSet { + let mut prefixes = HashSet::new(); + for input in &function.inputs { + if input.type_class == Some(rumoca_core::ClassType::Record) + || record_fields_by_type.contains_key(&input.type_name) + { + prefixes.insert(input.name.clone()); + } + if let Some((prefix, _)) = input.name.split_once('_') { + prefixes.insert(prefix.to_string()); + } + } + prefixes +} + +struct RecordFieldUseCollector<'a> { + prefixes: &'a HashSet, + fields: FieldUseMap, +} + +impl RecordFieldUseCollector<'_> { + fn record_var_ref(&mut self, name: &rumoca_core::Reference, dims: Vec) { + let Some((prefix, field)) = split_record_field_name(name.as_str()) else { + return; + }; + if !self.prefixes.contains(prefix) { + return; + } + merge_field_dims( + self.fields + .entry(prefix.to_string()) + .or_default() + .entry(field.to_string()) + .or_default(), + dims, + ); + } +} + +impl ExpressionVisitor for RecordFieldUseCollector<'_> { + fn visit_var_ref( + &mut self, + name: &rumoca_core::Reference, + subscripts: &[rumoca_core::Subscript], + ) { + self.record_var_ref(name, dims_from_subscripts(subscripts)); + for subscript in subscripts { + self.visit_subscript(subscript); + } + } + + fn visit_expression(&mut self, expr: &rumoca_core::Expression) { + if let rumoca_core::Expression::Index { + base, subscripts, .. + } = expr + && let rumoca_core::Expression::VarRef { + name, + subscripts: base_subscripts, + .. + } = base.as_ref() + { + let dims = dims_from_subscripts(base_subscripts) + .into_iter() + .chain(dims_from_subscripts(subscripts)) + .collect(); + self.record_var_ref(name, dims); + } + ExpressionVisitor::walk_expression(self, expr); + } +} + +fn split_record_field_name(name: &str) -> Option<(&str, &str)> { + let (prefix, field) = name.split_once('_')?; + (!prefix.is_empty() && !field.is_empty()).then_some((prefix, field)) +} + +fn dims_from_subscripts(subscripts: &[rumoca_core::Subscript]) -> Vec { + let max_index = subscripts + .iter() + .filter_map(const_subscript_index) + .max() + .unwrap_or_default(); + if max_index > 0 { + vec![max_index] + } else { + Vec::new() + } +} + +fn const_subscript_index(subscript: &rumoca_core::Subscript) -> Option { + match subscript { + rumoca_core::Subscript::Index { value, .. } => Some(*value), + rumoca_core::Subscript::Expr { expr, .. } => { + let rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + .. + } = expr.as_ref() + else { + return None; + }; + Some(*value) + } + rumoca_core::Subscript::Colon { .. } => None, + } +} + +fn collect_callee_record_field_uses( + function: &rumoca_core::Function, + functions: &IndexMap, + fields_by_function: &HashMap, +) -> FieldUseMap { + let mut collector = CalleeFieldUseCollector { + functions, + fields_by_function, + fields: BTreeMap::new(), + }; + for statement in &function.body { + visit_statement_expressions(statement, &mut collector); + } + collector.fields +} + +struct CalleeFieldUseCollector<'a> { + functions: &'a IndexMap, + fields_by_function: &'a HashMap, + fields: FieldUseMap, +} + +impl ExpressionVisitor for CalleeFieldUseCollector<'_> { + fn visit_expression(&mut self, expr: &rumoca_core::Expression) { + if let rumoca_core::Expression::FunctionCall { name, args, .. } = expr + && let Some(callee_fields) = self.fields_by_function.get(name.as_str()) + && let Some(callee) = self.functions.get(name.var_name()) + { + merge_callee_field_uses(args, callee, callee_fields, &mut self.fields); + } + ExpressionVisitor::walk_expression(self, expr); + } +} + +fn merge_callee_field_uses( + args: &[rumoca_core::Expression], + callee: &rumoca_core::Function, + callee_fields: &FieldUseMap, + target: &mut FieldUseMap, +) { + for (idx, input) in callee.inputs.iter().enumerate() { + let Some(fields) = callee_fields.get(&input.name) else { + continue; + }; + let Some(rumoca_core::Expression::VarRef { name: arg_name, .. }) = args.get(idx) else { + continue; + }; + let entry = target.entry(arg_name.as_str().to_string()).or_default(); + for (field, dims) in fields { + merge_field_dims(entry.entry(field.clone()).or_default(), dims.clone()); + } + } +} + +fn merge_field_use_map(target: &mut FieldUseMap, source: FieldUseMap) -> bool { + let mut changed = false; + for (prefix, fields) in source { + let target_fields = target.entry(prefix).or_default(); + for (field, dims) in fields { + let target_dims = target_fields.entry(field).or_default(); + let before = target_dims.clone(); + merge_field_dims(target_dims, dims); + changed |= *target_dims != before; + } + } + changed +} + +fn visit_statement_expressions( + statement: &rumoca_core::Statement, + visitor: &mut impl ExpressionVisitor, +) { + match statement { + rumoca_core::Statement::Assignment { value, .. } + | rumoca_core::Statement::Reinit { value, .. } => visitor.visit_expression(value), + rumoca_core::Statement::For { + indices, equations, .. + } => { + for index in indices { + visitor.visit_expression(&index.range); + } + for statement in equations { + visit_statement_expressions(statement, visitor); + } + } + rumoca_core::Statement::While { block, .. } => { + visitor.visit_expression(&block.cond); + for statement in &block.stmts { + visit_statement_expressions(statement, visitor); + } + } + rumoca_core::Statement::If { + cond_blocks, + else_block, + .. + } => { + for block in cond_blocks { + visitor.visit_expression(&block.cond); + for statement in &block.stmts { + visit_statement_expressions(statement, visitor); + } + } + if let Some(else_block) = else_block { + for statement in else_block { + visit_statement_expressions(statement, visitor); + } + } + } + rumoca_core::Statement::When { blocks, .. } => { + for block in blocks { + visitor.visit_expression(&block.cond); + for statement in &block.stmts { + visit_statement_expressions(statement, visitor); + } + } + } + rumoca_core::Statement::FunctionCall { args, .. } => { + for arg in args { + visitor.visit_expression(arg); + } + } + rumoca_core::Statement::Assert { + condition, + message, + level, + .. + } => { + visitor.visit_expression(condition); + visitor.visit_expression(message); + if let Some(level) = level { + visitor.visit_expression(level); + } + } + rumoca_core::Statement::Empty { .. } + | rumoca_core::Statement::Return { .. } + | rumoca_core::Statement::Break { .. } => {} + } +} + +fn merge_field_dims(target: &mut Vec, source: Vec) { + if source.len() > target.len() { + target.resize(source.len(), 0); + } + for (idx, value) in source.into_iter().enumerate() { + target[idx] = target[idx].max(value); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use rumoca_core::{ClassType, Expression, Function, FunctionParam, Span, Statement, VarName}; + + fn span(start: usize) -> Span { + Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + start, + start + 1, + ) + } + + fn var_ref(name: &str) -> Expression { + Expression::VarRef { + name: VarName::new(name).into(), + subscripts: vec![], + span: span(1), + } + } + + fn assignment(value: Expression) -> Statement { + Statement::Assignment { + comp: rumoca_core::ComponentReference { + local: false, + span: span(1), + parts: vec![rumoca_core::ComponentRefPart { + ident: "y".to_string(), + span: span(1), + subs: vec![], + }], + def_id: None, + }, + value, + span: span(1), + } + } + + #[test] + fn callee_field_uses_follow_signature_order_not_btree_order() { + let mut callee = Function::new("Pkg.callee", span(1)); + callee.add_input( + FunctionParam::new("port", "Pkg.Port", span(1)).with_type_class(ClassType::Connector), + ); + callee.add_input( + FunctionParam::new("state", "Pkg.State", span(1)).with_type_class(ClassType::Record), + ); + callee.body.push(assignment(var_ref("state_phase"))); + + let mut caller = Function::new("Pkg.caller", span(1)); + caller.body.push(assignment(Expression::FunctionCall { + name: VarName::new("Pkg.callee").into(), + args: vec![var_ref("port_a"), var_ref("state_a")], + is_constructor: false, + span: span(1), + })); + + let mut functions = IndexMap::new(); + functions.insert(callee.name.clone(), callee); + functions.insert(caller.name.clone(), caller); + + let record_fields_by_type = + HashMap::from([("Pkg.State".to_string(), vec!["phase".to_string()])]); + let inferred = infer_record_fields_by_function(&functions, &record_fields_by_type); + let caller_fields = inferred.get("Pkg.caller").expect("caller fields"); + + assert!( + !caller_fields.contains_key("port_a"), + "connector argument must not inherit the callee record field" + ); + assert_eq!( + caller_fields + .get("state_a") + .and_then(|fields| fields.get("phase")), + Some(&Vec::::new()) + ); + } +} diff --git a/crates/rumoca-phase-dae/src/dae_lowering/record_lowering_tests.rs b/crates/rumoca-phase-dae/src/dae_lowering/record_lowering_tests.rs index 8e43bb83e..8098e3ad8 100644 --- a/crates/rumoca-phase-dae/src/dae_lowering/record_lowering_tests.rs +++ b/crates/rumoca-phase-dae/src/dae_lowering/record_lowering_tests.rs @@ -17,6 +17,16 @@ fn var_ref(name: &str, span: Span) -> rumoca_core::Expression { } } +fn structured_var_ref(name: &str, span: Span) -> rumoca_core::Expression { + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::from_component_reference( + rumoca_core::ComponentReference::from_flat_segments(name, span, None), + ), + subscripts: vec![], + span, + } +} + fn record_constructor() -> rumoca_core::Function { let mut constructor = rumoca_core::Function::new("Pkg.Record", test_span(1)); constructor.is_constructor = true; @@ -43,6 +53,101 @@ fn function_with_array_input() -> rumoca_core::Function { function } +fn block_like_constructor() -> rumoca_core::Function { + let mut constructor = rumoca_core::Function::new("Pkg.Divide", test_span(1)); + constructor.is_constructor = true; + constructor.add_input( + rumoca_core::FunctionParam::new("u1", "RealInput", test_span(1)) + .with_type_class(ClassType::Connector), + ); + constructor.add_input( + rumoca_core::FunctionParam::new("u2", "RealInput", test_span(1)) + .with_type_class(ClassType::Connector), + ); + constructor +} + +fn assignment_to(name: &str, value: rumoca_core::Expression) -> rumoca_core::Statement { + rumoca_core::Statement::Assignment { + comp: rumoca_core::ComponentReference { + local: false, + span: test_span(1), + parts: vec![rumoca_core::ComponentRefPart { + ident: name.to_string(), + span: test_span(1), + subs: vec![], + }], + def_id: None, + }, + value, + span: test_span(1), + } +} + +#[test] +fn prepare_dae_for_codegen_unwraps_block_constructor_value_wrapper() { + let mut dae = Dae::default(); + dae.symbols + .functions + .insert(VarName::new("Pkg.Divide"), block_like_constructor()); + dae.continuous.equations.push(rumoca_ir_dae::Equation { + lhs: Some(VarName::new("x").into()), + rhs: rumoca_core::Expression::FunctionCall { + name: VarName::new("Pkg.f").into(), + args: vec![rumoca_core::Expression::FunctionCall { + name: VarName::new("Pkg.Divide").into(), + args: vec![var_ref("u", test_span(1))], + is_constructor: true, + span: test_span(1), + }], + is_constructor: false, + span: test_span(1), + }, + span: test_span(1), + origin: "test".to_string(), + scalar_count: 1, + }); + + let prepared = prepare_dae_for_codegen(&dae).expect("codegen preparation"); + + let rumoca_core::Expression::FunctionCall { args, .. } = + &prepared.as_dae().continuous.equations[0].rhs + else { + panic!("expected outer function call"); + }; + assert!(matches!( + &args[0], + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "u" + )); +} + +#[test] +fn prepare_dae_for_codegen_keeps_record_constructor_value() { + let mut dae = Dae::default(); + dae.symbols + .functions + .insert(VarName::new("Pkg.Record"), record_constructor()); + dae.continuous.equations.push(rumoca_ir_dae::Equation { + lhs: Some(VarName::new("x").into()), + rhs: rumoca_core::Expression::FunctionCall { + name: VarName::new("Pkg.Record").into(), + args: vec![var_ref("u", test_span(1))], + is_constructor: true, + span: test_span(1), + }, + span: test_span(1), + origin: "test".to_string(), + scalar_count: 1, + }); + + let prepared = prepare_dae_for_codegen(&dae).expect("codegen preparation"); + + assert!(matches!( + &prepared.as_dae().continuous.equations[0].rhs, + rumoca_core::Expression::FunctionCall { name, .. } if name.as_str() == "Pkg.Record" + )); +} + #[test] fn dae_record_param_lowering_uses_constructor_signature_metadata() { let span = test_span(1); @@ -306,3 +411,348 @@ fn dae_record_param_lowering_leaves_unknown_record_metadata_unexpanded() { }; assert_eq!(args.len(), 1); } + +#[test] +fn dae_record_param_lowering_keeps_external_object_inputs_opaque() { + let span = test_span(1); + let mut dae = Dae::default(); + + let mut external_constructor = rumoca_core::Function::new("Pkg.ExternalTable", span); + external_constructor.is_constructor = true; + external_constructor.add_input(rumoca_core::FunctionParam::new("table", "Real", span)); + external_constructor.add_input(rumoca_core::FunctionParam::new("fileName", "String", span)); + dae.symbols + .functions + .insert(VarName::new("Pkg.ExternalTable"), external_constructor); + + let mut record_typed_user = rumoca_core::Function::new("Pkg.recordUser", span); + record_typed_user.add_input( + rumoca_core::FunctionParam::new("recordish", "Pkg.ExternalTable", span) + .with_type_class(ClassType::Record), + ); + record_typed_user.add_output(rumoca_core::FunctionParam::new("y", "Real", span)); + dae.symbols + .functions + .insert(VarName::new("Pkg.recordUser"), record_typed_user); + + let mut external_user = rumoca_core::Function::new("Pkg.getTableValue", span); + external_user.add_input( + rumoca_core::FunctionParam::new("tableID", "Pkg.ExternalTable", span) + .with_type_class(ClassType::Class), + ); + external_user.add_input(rumoca_core::FunctionParam::new("column", "Integer", span)); + external_user.add_output(rumoca_core::FunctionParam::new("y", "Real", span)); + dae.symbols + .functions + .insert(VarName::new("Pkg.getTableValue"), external_user); + + dae.continuous.equations.push(rumoca_ir_dae::Equation { + lhs: Some(VarName::new("x").into()), + rhs: rumoca_core::Expression::FunctionCall { + name: VarName::new("Pkg.getTableValue").into(), + args: vec![ + var_ref("integerTable.combiTimeTable.tableID", span), + rumoca_core::Expression::Literal { + value: Literal::Integer(1), + span, + }, + ], + is_constructor: false, + span, + }, + span, + origin: "test".to_string(), + scalar_count: 1, + }); + + lower_record_function_params_dae(&mut dae).expect("record params lower"); + + let external_user = dae + .symbols + .functions + .get(&VarName::new("Pkg.getTableValue")) + .expect("external-object function remains"); + let input_names = external_user + .inputs + .iter() + .map(|input| input.name.as_str()) + .collect::>(); + assert_eq!(input_names, vec!["tableID", "column"]); + + let rumoca_core::Expression::FunctionCall { args, .. } = &dae.continuous.equations[0].rhs + else { + panic!("expected external-object function call"); + }; + assert_eq!(args.len(), 2); + assert!(matches!( + &args[0], + rumoca_core::Expression::VarRef { name, span: arg_span, .. } + if name.as_str() == "integerTable.combiTimeTable.tableID" && *arg_span == span + )); +} + +#[test] +fn dae_record_param_lowering_infers_fields_from_already_lowered_body() { + let mut dae = Dae::default(); + let mut function = function_with_record_input(); + function.body.push(assignment_to( + "y", + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: Box::new(var_ref("r_T", test_span(1))), + rhs: Box::new(rumoca_core::Expression::Index { + base: Box::new(var_ref("r_X", test_span(1))), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(rumoca_core::Expression::Literal { + value: Literal::Integer(1), + span: test_span(1), + }), + span: test_span(1), + }], + span: test_span(1), + }), + span: test_span(1), + }, + )); + dae.symbols + .functions + .insert(VarName::new("Pkg.f"), function); + dae.continuous.equations.push(rumoca_ir_dae::Equation { + lhs: Some(VarName::new("x").into()), + rhs: rumoca_core::Expression::FunctionCall { + name: VarName::new("Pkg.f").into(), + args: vec![var_ref("rec", test_span(1))], + is_constructor: false, + span: test_span(1), + }, + span: test_span(1), + origin: "test".to_string(), + scalar_count: 1, + }); + + lower_record_function_params_dae(&mut dae).expect("record params lower"); + + let function = dae + .symbols + .functions + .get(&VarName::new("Pkg.f")) + .expect("function remains"); + let input_names = function + .inputs + .iter() + .map(|input| input.name.as_str()) + .collect::>(); + assert_eq!(input_names, vec!["r_T", "r_X"]); + let rumoca_core::Expression::FunctionCall { args, .. } = &dae.continuous.equations[0].rhs + else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 2); + assert!(matches!( + &args[0], + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "rec.T" + )); + assert!(matches!( + &args[1], + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "rec.X" + )); +} + +#[test] +fn dae_record_arg_expansion_projects_overexpanded_record_field_actual_to_base() { + let mut out = Vec::new(); + expand_dae_record_arg( + &structured_var_ref("pipe.flowModel.states.phase", test_span(1)), + &[ + "phase".to_string(), + "h".to_string(), + "d".to_string(), + "T".to_string(), + "p".to_string(), + ], + &mut out, + test_span(1), + ) + .expect("record arg expansion should project to record base"); + + let names = out + .iter() + .map(|expr| match expr { + rumoca_core::Expression::VarRef { name, .. } => name.as_str().to_string(), + other => panic!("expected VarRef, got {other:?}"), + }) + .collect::>(); + assert_eq!( + names, + vec![ + "pipe.flowModel.states.phase", + "pipe.flowModel.states.h", + "pipe.flowModel.states.d", + "pipe.flowModel.states.T", + "pipe.flowModel.states.p", + ] + ); +} + +#[test] +fn dae_record_param_lowering_merges_metadata_and_body_inferred_fields() { + let mut dae = Dae::default(); + let mut constructor = rumoca_core::Function::new("Pkg.Record", test_span(1)); + constructor.is_constructor = true; + constructor.add_input(rumoca_core::FunctionParam::new("p", "Real", test_span(1))); + constructor.add_input(rumoca_core::FunctionParam::new("T", "Real", test_span(1))); + dae.symbols + .functions + .insert(VarName::new("Pkg.Record"), constructor); + + let mut function = function_with_record_input(); + function.body.push(assignment_to( + "y", + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: Box::new(var_ref("r_T", test_span(1))), + rhs: Box::new(var_ref("r_X", test_span(1))), + span: test_span(1), + }, + )); + dae.symbols + .functions + .insert(VarName::new("Pkg.f"), function); + + lower_record_function_params_dae(&mut dae).expect("record params lower"); + + let function = dae + .symbols + .functions + .get(&VarName::new("Pkg.f")) + .expect("function remains"); + let input_names = function + .inputs + .iter() + .map(|input| input.name.as_str()) + .collect::>(); + assert_eq!(input_names, vec!["r_p", "r_T", "r_X"]); +} + +#[test] +fn dae_record_param_lowering_propagates_nested_callee_field_requirements() { + let mut dae = Dae::default(); + let mut constructor = rumoca_core::Function::new("Pkg.Record", test_span(1)); + constructor.is_constructor = true; + constructor.add_input(rumoca_core::FunctionParam::new("p", "Real", test_span(1))); + constructor.add_input(rumoca_core::FunctionParam::new("T", "Real", test_span(1))); + dae.symbols + .functions + .insert(VarName::new("Pkg.Record"), constructor); + + let mut callee = rumoca_core::Function::new("Pkg.g", test_span(1)); + callee.add_input( + rumoca_core::FunctionParam::new("state", "Pkg.Record", test_span(1)) + .with_type_class(ClassType::Record), + ); + callee.add_output(rumoca_core::FunctionParam::new("y", "Real", test_span(1))); + callee + .body + .push(assignment_to("y", var_ref("state_X", test_span(1)))); + dae.symbols.functions.insert(VarName::new("Pkg.g"), callee); + + let mut caller = rumoca_core::Function::new("Pkg.f", test_span(1)); + caller.add_input( + rumoca_core::FunctionParam::new("state", "Pkg.Record", test_span(1)) + .with_type_class(ClassType::Record), + ); + caller.add_output(rumoca_core::FunctionParam::new("y", "Real", test_span(1))); + caller.body.push(assignment_to( + "y", + rumoca_core::Expression::FunctionCall { + name: VarName::new("Pkg.g").into(), + args: vec![var_ref("state", test_span(1))], + is_constructor: false, + span: test_span(1), + }, + )); + dae.symbols.functions.insert(VarName::new("Pkg.f"), caller); + + lower_record_function_params_dae(&mut dae).expect("record params lower"); + + let caller = dae + .symbols + .functions + .get(&VarName::new("Pkg.f")) + .expect("caller remains"); + let input_names = caller + .inputs + .iter() + .map(|input| input.name.as_str()) + .collect::>(); + assert_eq!(input_names, vec!["state_p", "state_T", "state_X"]); + let rumoca_core::Statement::Assignment { value, .. } = &caller.body[0] else { + panic!("expected assignment"); + }; + let rumoca_core::Expression::FunctionCall { args, .. } = value else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 3); + assert!( + matches!( + &args[2], + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "state.X" + ), + "unexpected propagated callee args: {args:#?}" + ); +} + +#[test] +fn dae_record_param_lowering_infers_air_state_x_width_from_indexed_body_use() { + let mut dae = Dae::default(); + + let mut function = + rumoca_core::Function::new("Buildings.Media.Air.specificHeatCapacityCp", test_span(1)); + function.add_input(rumoca_core::FunctionParam::new( + "state_p", + "AbsolutePressure", + test_span(1), + )); + function.add_input(rumoca_core::FunctionParam::new( + "state_T", + "Temperature", + test_span(1), + )); + function.add_output(rumoca_core::FunctionParam::new( + "cp", + "SpecificHeatCapacity", + test_span(1), + )); + function.body.push(assignment_to( + "cp", + rumoca_core::Expression::Index { + base: Box::new(var_ref("state_X", test_span(1))), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(rumoca_core::Expression::Literal { + value: Literal::Integer(1), + span: test_span(1), + }), + span: test_span(1), + }], + span: test_span(1), + }, + )); + dae.symbols.functions.insert( + VarName::new("Buildings.Media.Air.specificHeatCapacityCp"), + function, + ); + + lower_record_function_params_dae(&mut dae).expect("record params lower"); + + let function = dae + .symbols + .functions + .get(&VarName::new("Buildings.Media.Air.specificHeatCapacityCp")) + .expect("function remains"); + let state_x = function + .inputs + .iter() + .find(|input| input.name == "state_X") + .expect("state_X input should be synthesized"); + assert_eq!(state_x.dims, vec![1]); +} diff --git a/crates/rumoca-phase-dae/src/dae_lowering/tests.rs b/crates/rumoca-phase-dae/src/dae_lowering/tests.rs index 34543ee74..1baf53486 100644 --- a/crates/rumoca-phase-dae/src/dae_lowering/tests.rs +++ b/crates/rumoca-phase-dae/src/dae_lowering/tests.rs @@ -1,6 +1,12 @@ use super::*; use rumoca_core::Span; +mod colon_slice_dot; +mod initialization_provenance; +mod matrix_product_projection; +mod record_array_member; +mod record_array_projection_alias; + fn test_span() -> Span { Span::from_offsets( rumoca_core::SourceId::from_source_name("dae_lowering_fixture.mo"), @@ -18,6 +24,29 @@ fn var_ref(name: &str) -> rumoca_core::Expression { } } +fn component_ref(parts: &[(&str, Option)]) -> rumoca_core::ComponentReference { + rumoca_core::ComponentReference { + local: false, + span: test_span(), + parts: parts + .iter() + .map(|(ident, index)| rumoca_core::ComponentRefPart { + ident: ident.to_string(), + span: test_span(), + subs: index + .map(|value| { + vec![rumoca_core::Subscript::Index { + value, + span: test_span(), + }] + }) + .unwrap_or_default(), + }) + .collect(), + def_id: None, + } +} + fn var_ref_with_expr_subscript( name: &str, subscript: rumoca_core::Expression, @@ -152,6 +181,20 @@ fn all_var_names(expr: &rumoca_core::Expression) -> Vec { names } +fn var_ref_subscript_indices(expr: &rumoca_core::Expression) -> Vec { + let rumoca_core::Expression::VarRef { subscripts, .. } = expr else { + panic!("expected VarRef, got {expr:?}"); + }; + subscripts.clone() +} + +fn index_expr_subscripts(expr: &rumoca_core::Expression) -> Vec { + let rumoca_core::Expression::Index { subscripts, .. } = expr else { + panic!("expected Index, got {expr:?}"); + }; + subscripts.clone() +} + fn collect_var_names_rec(expr: &rumoca_core::Expression, names: &mut Vec) { match expr { rumoca_core::Expression::VarRef { @@ -230,6 +273,92 @@ fn assert_any_true_array_arg(expr: &rumoca_core::Expression, expected_names: &[& } } +fn contains_literal_index_ref(expr: &rumoca_core::Expression, needle: &str, index: i64) -> bool { + match expr { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } => { + name.as_str() == needle + && matches!( + subscripts.as_slice(), + [rumoca_core::Subscript::Expr { + expr, + .. + }] if matches!( + expr.as_ref(), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + .. + } if *value == index + ) + ) + } + rumoca_core::Expression::Binary { lhs, rhs, .. } => { + contains_literal_index_ref(lhs, needle, index) + || contains_literal_index_ref(rhs, needle, index) + } + rumoca_core::Expression::Unary { rhs, .. } => { + contains_literal_index_ref(rhs, needle, index) + } + rumoca_core::Expression::BuiltinCall { args, .. } + | rumoca_core::Expression::FunctionCall { args, .. } => args + .iter() + .any(|arg| contains_literal_index_ref(arg, needle, index)), + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + branches.iter().any(|(condition, value)| { + contains_literal_index_ref(condition, needle, index) + || contains_literal_index_ref(value, needle, index) + }) || contains_literal_index_ref(else_branch, needle, index) + } + rumoca_core::Expression::Array { elements, .. } + | rumoca_core::Expression::Tuple { elements, .. } => elements + .iter() + .any(|element| contains_literal_index_ref(element, needle, index)), + rumoca_core::Expression::Range { + start, step, end, .. + } => { + contains_literal_index_ref(start, needle, index) + || step + .as_ref() + .is_some_and(|step| contains_literal_index_ref(step, needle, index)) + || contains_literal_index_ref(end, needle, index) + } + rumoca_core::Expression::ArrayComprehension { + expr, + indices, + filter, + .. + } => { + contains_literal_index_ref(expr, needle, index) + || indices + .iter() + .any(|index_def| contains_literal_index_ref(&index_def.range, needle, index)) + || filter + .as_ref() + .is_some_and(|filter| contains_literal_index_ref(filter, needle, index)) + } + rumoca_core::Expression::Index { + base, subscripts, .. + } => { + contains_literal_index_ref(base, needle, index) + || subscripts.iter().any(|subscript| match subscript { + rumoca_core::Subscript::Expr { expr, .. } => { + contains_literal_index_ref(expr, needle, index) + } + _ => false, + }) + } + rumoca_core::Expression::FieldAccess { base, .. } => { + contains_literal_index_ref(base, needle, index) + } + rumoca_core::Expression::Literal { .. } | rumoca_core::Expression::Empty { .. } => false, + } +} + #[test] fn test_parameter_start_dependency_sort_deduplicates_repeated_refs() { let mut dae = dae::Dae::default(); @@ -477,6 +606,43 @@ fn test_scalarize_phantom_vector_equations_preserves_event_assignment_lhs() { } } +#[test] +fn test_scalarize_discrete_residual_vector_equations_recovers_explicit_assignments() { + let mut dae = Dae::new(); + let mut target = dae::Variable::new(rumoca_core::VarName::new("target"), test_span()); + target.dims = vec![2]; + dae.variables + .discrete_valued + .insert(rumoca_core::VarName::new("target"), target); + + for k in 1..=2 { + let name = format!("source[{k}]"); + dae.variables.discrete_valued.insert( + rumoca_core::VarName::new(&name), + dae::Variable::new(rumoca_core::VarName::new(&name), test_span()), + ); + } + + dae.discrete.valued_updates.push(dae::Equation::residual( + sub(var_ref("target"), var_ref("source")), + test_span(), + "discrete valued residual assignment", + )); + dae.discrete.valued_updates[0].scalar_count = 2; + + scalarize_phantom_vector_equations(&mut dae).unwrap(); + + assert_eq!(dae.discrete.valued_updates.len(), 2); + for (k, eq) in dae.discrete.valued_updates.iter().enumerate() { + assert_eq!( + eq.lhs.as_ref().map(|lhs| lhs.as_str()), + Some(format!("target[{}]", k + 1).as_str()) + ); + let names = all_var_names(&eq.rhs); + assert_eq!(names, vec![format!("source[{}]", k + 1)]); + } +} + #[test] fn scalarize_phantom_vector_equations_reports_missing_function_shape_with_call_span() { let mut dae = Dae::new(); @@ -747,6 +913,57 @@ fn embedded_scalar_reference_canonicalization_requires_provenance() { ); } +#[test] +fn test_scalarize_preserves_plain_matrix_literal_assignment_for_structural_scalarizer() { + let mut dae = Dae::new(); + + let mut skew = dae::Variable::new( + rumoca_core::VarName::new("skew"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ); + skew.dims = vec![3, 3]; + dae.variables + .algebraics + .insert(rumoca_core::VarName::new("skew"), skew); + + let row = |values: [f64; 3]| rumoca_core::Expression::Array { + elements: values + .into_iter() + .map(|value| rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(value), + span: test_span(), + }) + .collect(), + is_matrix: false, + span: test_span(), + }; + let matrix = rumoca_core::Expression::Array { + elements: vec![ + row([0.0, -1.0, 0.0]), + row([1.0, 0.0, 0.0]), + row([0.0, 0.0, 0.0]), + ], + is_matrix: false, + span: test_span(), + }; + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: sub(var_ref("skew"), matrix.clone()), + span: test_span(), + origin: "plain matrix literal residual".to_string(), + scalar_count: 1, + }); + + scalarize_phantom_vector_equations(&mut dae).unwrap(); + + assert_eq!(dae.continuous.equations.len(), 1); + assert_eq!(dae.continuous.equations[0].scalar_count, 1); + assert_eq!( + dae.continuous.equations[0].rhs, + sub(var_ref("skew"), matrix) + ); +} + #[test] fn test_scalarize_vector_binding_preserves_array_comprehension_without_phantom_refs() { let mut dae = Dae::new(); @@ -805,6 +1022,58 @@ fn test_scalarize_vector_binding_preserves_array_comprehension_without_phantom_r ); } +#[test] +fn test_scalarize_scalar_lhs_projects_array_comprehension_rhs() { + let mut dae = Dae::new(); + let mut vs = dae::Variable::new( + rumoca_core::VarName::new("pipe.vs"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ); + vs.dims = vec![2]; + dae.variables + .algebraics + .insert(rumoca_core::VarName::new("pipe.vs"), vs); + + let rhs = sub( + var_ref_with_expr_subscript("pipe.vs", int_lit(1)), + rumoca_core::Expression::ArrayComprehension { + expr: Box::new(mul( + var_ref_with_expr_subscript("pipe.crossAreas", var_ref("i")), + var_ref_with_expr_subscript("pipe.lengths", var_ref("i")), + )), + indices: vec![rumoca_core::ComprehensionIndex { + name: "i".to_string(), + range: rumoca_core::Expression::Range { + start: Box::new(int_lit(1)), + step: None, + end: Box::new(int_lit(2)), + span: test_span(), + }, + }], + filter: None, + span: test_span(), + }, + ); + dae.continuous.equations.push(dae::Equation::residual( + rhs, + test_span(), + "binding equation for pipe.vs[1]", + )); + + scalarize_phantom_vector_equations(&mut dae).unwrap(); + + let rewritten = &dae.continuous.equations[0].rhs; + assert!( + !expr_has_array_comprehension(rewritten), + "scalar lhs equation should project aggregate RHS: {rewritten:?}" + ); + assert!( + contains_literal_index_ref(rewritten, "pipe.crossAreas", 1) + && contains_literal_index_ref(rewritten, "pipe.lengths", 1), + "scalar lhs equation should select matching RHS element: {rewritten:?}" + ); +} + #[test] fn test_scalarize_phantom_vector_equations_selects_zeros() { let mut dae = Dae::new(); @@ -881,6 +1150,61 @@ fn test_scalarize_phantom_vector_equations_selects_zeros() { } } +#[test] +fn test_scalarize_fill_repeats_vector_value_across_outer_dimension() { + let mut dae = Dae::new(); + let span = test_span(); + + for port in 1..=2 { + let name = format!("freshAir.ports[{port}].C_outflow"); + let mut var = dae::Variable::new(rumoca_core::VarName::new(&name), span); + var.dims = vec![1]; + dae.variables + .algebraics + .insert(rumoca_core::VarName::new(&name), var); + } + + let mut input = dae::Variable::new(rumoca_core::VarName::new("freshAir.C_in_internal"), span); + input.dims = vec![1]; + dae.variables + .algebraics + .insert(rumoca_core::VarName::new("freshAir.C_in_internal"), input); + + let residual = sub( + var_ref("freshAir.ports.C_outflow"), + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Fill, + args: vec![var_ref("freshAir.C_in_internal"), int_lit(2)], + span, + }, + ); + dae.continuous.equations.push(dae::Equation::residual_array( + residual, + span, + "equation from freshAir", + 2, + )); + + scalarize_phantom_vector_equations(&mut dae).unwrap(); + + assert_eq!(dae.continuous.equations.len(), 2); + for (idx, eq) in dae.continuous.equations.iter().enumerate() { + let names = all_var_names(&eq.rhs); + assert!( + names.contains(&format!("freshAir.ports[{}].C_outflow", idx + 1)), + "equation {idx} should select the matching port: {names:?}" + ); + assert!( + names.contains(&"freshAir.C_in_internal[1]".to_string()), + "fill should repeat the singleton concentration vector for each port: {names:?}" + ); + assert!( + !names.contains(&"freshAir.C_in_internal[2]".to_string()), + "outer fill index must not be used as the concentration component lane: {names:?}" + ); + } +} + #[test] fn test_scalarize_preserves_vector_function_arguments_for_array_output() { let mut dae = Dae::new(); @@ -976,6 +1300,397 @@ fn test_scalarize_preserves_vector_function_arguments_for_array_output() { } } +#[test] +fn test_scalarize_projects_vectorized_scalar_function_arguments() { + let mut dae = Dae::new(); + let span = test_span(); + + for (name, dims) in [ + ("rhos", vec![4]), + ("states.p", vec![4]), + ("states.T", vec![4]), + ("states.X", vec![4, 2]), + ] { + let mut var = dae::Variable::new(rumoca_core::VarName::new(name), span); + var.dims = dims; + dae.variables + .algebraics + .insert(rumoca_core::VarName::new(name), var); + } + + let mut density = rumoca_core::Function::new("Medium.density", span); + density + .inputs + .push(rumoca_core::FunctionParam::new("state_p", "Real", span)); + density + .inputs + .push(rumoca_core::FunctionParam::new("state_T", "Real", span)); + density + .inputs + .push(rumoca_core::FunctionParam::new("state_X", "Real", span).with_dims(vec![0])); + density + .outputs + .push(rumoca_core::FunctionParam::new("d", "Real", span)); + dae.symbols.functions.insert(density.name.clone(), density); + + let eq_rhs = sub( + var_ref("rhos"), + function_call( + "Medium.density", + vec![ + var_ref("states.p"), + var_ref("states.T"), + var_ref("states.X"), + ], + ), + ); + dae.continuous.equations.push(dae::Equation::residual_array( + eq_rhs, + span, + "density vector equation", + 4, + )); + + scalarize_phantom_vector_equations(&mut dae).unwrap(); + + assert_eq!(dae.continuous.equations.len(), 4); + let rumoca_core::Expression::Binary { rhs, .. } = &dae.continuous.equations[2].rhs else { + panic!("expected scalarized subtraction residual"); + }; + let rumoca_core::Expression::FunctionCall { args, .. } = rhs.as_ref() else { + panic!("expected scalarized RHS to remain a scalar density call"); + }; + assert_eq!(all_var_names(&args[0]), vec!["states.p[3]".to_string()]); + assert_eq!( + var_ref_subscript_indices(&args[0]), + vec![rumoca_core::Subscript::generated_index(3, span)] + ); + assert_eq!( + index_expr_subscripts(&args[2]), + vec![ + rumoca_core::Subscript::generated_index(3, span), + rumoca_core::Subscript::generated_colon(span), + ] + ); +} + +#[test] +fn test_scalarize_scalar_lhs_preserves_stream_wrapped_array_formal_argument() { + let mut dae = Dae::new(); + let span = test_span(); + + for (name, dims) in [ + ("volume.portInDensities", vec![4]), + ("volume.vessel_ps_static", vec![4]), + ("volume.ports[2].Xi_outflow", vec![1]), + ] { + let mut var = dae::Variable::new(rumoca_core::VarName::new(name), span); + var.dims = dims; + dae.variables + .algebraics + .insert(rumoca_core::VarName::new(name), var); + } + + let mut set_state = rumoca_core::Function::new("Medium.setState_phX", span); + set_state + .inputs + .push(rumoca_core::FunctionParam::new("p", "Real", span)); + set_state + .inputs + .push(rumoca_core::FunctionParam::new("X", "Real", span).with_dims(vec![0])); + set_state.outputs.push(rumoca_core::FunctionParam::new( + "state", + "Medium.ThermodynamicState", + span, + )); + dae.symbols + .functions + .insert(set_state.name.clone(), set_state); + + let lhs = rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new("volume.portInDensities").into(), + subscripts: vec![rumoca_core::Subscript::generated_index(2, span)], + span, + }; + let eq_rhs = sub( + lhs, + function_call( + "Medium.setState_phX", + vec![ + var_ref("volume.vessel_ps_static"), + function_call("inStream", vec![var_ref("volume.ports[2].Xi_outflow")]), + ], + ), + ); + dae.continuous + .equations + .push(dae::Equation::residual(eq_rhs, span, "volume port density")); + + scalarize_phantom_vector_equations(&mut dae).unwrap(); + + let rumoca_core::Expression::Binary { rhs, .. } = &dae.continuous.equations[0].rhs else { + panic!("expected residual subtraction"); + }; + let rumoca_core::Expression::FunctionCall { args, .. } = rhs.as_ref() else { + panic!("expected setState call"); + }; + assert_eq!( + var_ref_subscript_indices(&args[0]), + vec![rumoca_core::Subscript::generated_index(2, span)] + ); + let rumoca_core::Expression::FunctionCall { + args: stream_args, .. + } = &args[1] + else { + panic!("expected stream passthrough argument"); + }; + assert_eq!( + all_var_names(&stream_args[0]), + vec!["volume.ports[2].Xi_outflow"] + ); +} + +#[test] +fn test_scalarize_canonicalizes_var_ref_colon_slice_to_index_expr() { + let mut dae = Dae::new(); + let span = test_span(); + let mut a = dae::Variable::new(rumoca_core::VarName::new("a"), span); + a.dims = vec![2, 3]; + dae.variables + .algebraics + .insert(rumoca_core::VarName::new("a"), a); + dae.continuous.equations.push(dae::Equation::residual( + rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new("a").into(), + subscripts: vec![ + rumoca_core::Subscript::generated_index(1, span), + rumoca_core::Subscript::generated_colon(span), + ], + span, + }, + span, + "colon slice", + )); + + scalarize_phantom_vector_equations(&mut dae).unwrap(); + + let rumoca_core::Expression::Index { + base, subscripts, .. + } = &dae.continuous.equations[0].rhs + else { + panic!( + "colon VarRef should be normalized to Index, got {:?}", + dae.continuous.equations[0].rhs + ); + }; + assert!(matches!( + base.as_ref(), + rumoca_core::Expression::VarRef { + name, + subscripts, + .. + } if name.as_str() == "a" && subscripts.is_empty() + )); + assert_eq!( + subscripts, + &[ + rumoca_core::Subscript::generated_index(1, span), + rumoca_core::Subscript::generated_colon(span), + ] + ); +} + +#[test] +fn sync_structured_templates_replaces_stale_materialized_external_table_arg() { + let mut dae = Dae::new(); + let span = test_span(); + let stale_constructor = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Modelica.Blocks.Types.ExternalCombiTimeTable"), + args: vec![rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Fill, + args: vec![int_lit(0), int_lit(0), int_lit(2)], + span, + }], + is_constructor: true, + span, + }; + let stale_residual = sub( + var_ref("integerTable.combiTimeTable.y"), + function_call( + "Modelica.Blocks.Tables.Internal.getTimeTableValueNoDer", + vec![stale_constructor, int_lit(1), var_ref("time")], + ), + ); + dae.continuous.equations.push(dae::Equation::residual( + stale_residual, + span, + "materialized", + )); + + let canonical_residual = sub( + var_ref("integerTable.combiTimeTable.y"), + function_call( + "Modelica.Blocks.Tables.Internal.getTimeTableValueNoDer", + vec![ + var_ref("integerTable.combiTimeTable.tableID"), + var_ref("i"), + var_ref("time"), + ], + ), + ); + dae.continuous + .structured_equations + .push(dae::StructuredEquationFamily { + domain: rumoca_core::StructuredIndexDomain { + binders: vec![rumoca_core::StructuredIndexBinder { + id: 0, + display_name: "i".to_string(), + lower: 1, + upper: 1, + step: 1, + }], + }, + first_equation_index: 0, + equation_counts: vec![1], + span, + origin: "equation from integerTable.combiTimeTable".to_string(), + regular: None, + template: Some(rumoca_core::ComprehensionTemplate { + body: vec![canonical_residual], + }), + interiors_materialized: true, + }); + + sync_materialized_structured_equation_templates(&mut dae) + .expect("structured template sync should succeed"); + + let rumoca_core::Expression::Binary { rhs, .. } = &dae.continuous.equations[0].rhs else { + panic!("expected residual subtraction"); + }; + let rumoca_core::Expression::FunctionCall { args, .. } = rhs.as_ref() else { + panic!("expected table lookup call"); + }; + assert!(matches!( + &args[0], + rumoca_core::Expression::VarRef { name, .. } + if name.as_str() == "integerTable.combiTimeTable.tableID" + )); + assert!(matches!( + &args[1], + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(1), + .. + } + )); +} + +#[test] +fn sync_structured_templates_preserves_vector_equation_scalar_count() { + let mut dae = Dae::new(); + let span = test_span(); + dae.continuous.equations.push(dae::Equation::residual_array( + sub( + rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new("leg_v_b").into(), + subscripts: vec![ + rumoca_core::Subscript::generated_colon(span), + rumoca_core::Subscript::generated_index(1, span), + ], + span, + }, + var_ref("v_b"), + ), + span, + "materialized vector column equation", + 3, + )); + + let template_residual = sub( + rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new("leg_v_b").into(), + subscripts: vec![ + rumoca_core::Subscript::generated_colon(span), + rumoca_core::Subscript::Expr { + expr: Box::new(var_ref("i")), + span, + }, + ], + span, + }, + var_ref("v_b"), + ); + dae.continuous + .structured_equations + .push(dae::StructuredEquationFamily { + domain: rumoca_core::StructuredIndexDomain { + binders: vec![rumoca_core::StructuredIndexBinder { + id: 0, + display_name: "i".to_string(), + lower: 1, + upper: 1, + step: 1, + }], + }, + first_equation_index: 0, + equation_counts: vec![1], + span, + origin: "equation from vector column loop".to_string(), + regular: None, + template: Some(rumoca_core::ComprehensionTemplate { + body: vec![template_residual], + }), + interiors_materialized: true, + }); + + sync_materialized_structured_equation_templates(&mut dae) + .expect("structured template sync should preserve vector equation metadata"); + + assert_eq!( + dae.continuous.equations[0].scalar_count, 3, + "template sync must not collapse vector-valued loop equations to scalar rows" + ); +} + +#[test] +fn repair_external_table_events_uses_component_table_id_handle() { + let mut dae = Dae::new(); + let span = test_span(); + dae.metadata + .nonnumeric_variable_names + .push("integerTable.combiTimeTable.tableID".to_string()); + let stale_constructor = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Modelica.Blocks.Types.ExternalCombiTimeTable"), + args: vec![rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Fill, + args: vec![int_lit(0), int_lit(0), int_lit(2)], + span, + }], + is_constructor: true, + span, + }; + dae.discrete.real_updates.push(dae::Equation::explicit( + rumoca_core::Reference::new("integerTable.combiTimeTable.nextTimeEventScaled"), + function_call( + "Modelica.Blocks.Tables.Internal.getNextTimeEvent", + vec![stale_constructor, var_ref("time")], + ), + span, + "guarded when equation assignment", + )); + + repair_external_table_event_handles(&mut dae); + + let rumoca_core::Expression::FunctionCall { args, .. } = &dae.discrete.real_updates[0].rhs + else { + panic!("expected getNextTimeEvent call"); + }; + assert!(matches!( + &args[0], + rumoca_core::Expression::VarRef { name, .. } + if name.as_str() == "integerTable.combiTimeTable.tableID" + )); +} + #[test] fn test_scalarize_leaves_scalar_equations_unchanged() { let mut dae = Dae::new(); @@ -1041,6 +1756,156 @@ fn test_scalarize_ignores_vector_equations_without_phantom_refs() { assert_eq!(dae.continuous.equations[0].scalar_count, 3); } +#[test] +fn test_scalarize_vector_equation_projects_indexed_record_field_slice() { + let mut dae = Dae::new(); + + for name in ["T1[1]", "T1[2]", "ele[1].vol1.T", "ele[2].vol1.T"] { + dae.variables.algebraics.insert( + rumoca_core::VarName::new(name), + dae::Variable::new(name.into(), test_span()), + ); + } + + let mut ele = dae::Variable::new(rumoca_core::VarName::new("ele"), test_span()); + ele.dims = vec![2]; + dae.variables + .algebraics + .insert(rumoca_core::VarName::new("ele"), ele); + + let ele_slice = rumoca_core::Expression::Index { + base: Box::new(var_ref("ele")), + subscripts: vec![rumoca_core::Subscript::Colon { span: test_span() }], + span: test_span(), + }; + let rhs = rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FieldAccess { + base: Box::new(ele_slice), + field: "vol1".to_string(), + span: test_span(), + }), + field: "T".to_string(), + span: test_span(), + }; + dae.continuous.equations.push(dae::Equation::residual_array( + sub(var_ref("T1"), rhs), + test_span(), + "record field slice vector equation", + 2, + )); + + scalarize_phantom_vector_equations(&mut dae).unwrap(); + + assert_eq!(dae.continuous.equations.len(), 2); + assert_eq!( + all_var_names(&dae.continuous.equations[0].rhs), + vec!["T1[1]", "ele[1]"] + ); + assert_eq!( + all_var_names(&dae.continuous.equations[1].rhs), + vec!["T1[2]", "ele[2]"] + ); +} + +#[test] +fn test_scalarize_record_array_field_uses_field_lane_before_record_lane() { + let mut dae = Dae::new(); + + for record_index in [1, 2] { + let name = format!("ductOut.mediums[{record_index}].state.X"); + let mut var = dae::Variable::new(rumoca_core::VarName::new(&name), test_span()); + var.dims = vec![2]; + var.component_ref = Some(component_ref(&[ + ("ductOut", None), + ("mediums", Some(record_index)), + ("state", None), + ("X", None), + ])); + dae.variables.algebraics.insert(var.name.clone(), var); + } + for record_index in [1, 2, 3] { + let name = format!("ductOut.statesFM[{record_index}].X"); + let mut var = dae::Variable::new(rumoca_core::VarName::new(&name), test_span()); + var.dims = vec![2]; + var.component_ref = Some(component_ref(&[ + ("ductOut", None), + ("statesFM", Some(record_index)), + ("X", None), + ])); + dae.variables.algebraics.insert(var.name.clone(), var); + } + + let record_array_fields = build_record_array_field_map(&dae); + let mut array_dims = build_dae_var_dims_map(&dae); + array_dims.retain(|_, dims| !dims.is_empty()); + let expr = rumoca_core::Expression::FieldAccess { + base: Box::new(var_ref("ductOut.mediums.state")), + field: "X".to_string(), + span: test_span(), + }; + + let projected = (0..4) + .map(|k| { + project_scalarized_rhs_expr_at( + &expr, + k, + &HashSet::new(), + &array_dims, + &record_array_fields, + &IndexMap::new(), + ExpressionForm::Other, + ) + .unwrap() + }) + .collect::>(); + + assert_eq!( + projected.iter().flat_map(all_var_names).collect::>(), + vec![ + "ductOut.mediums[1].state.X[1]", + "ductOut.mediums[1].state.X[2]", + "ductOut.mediums[2].state.X[1]", + "ductOut.mediums[2].state.X[2]", + ] + ); + + let lhs = rumoca_core::Expression::Array { + elements: vec![ + var_ref("ductOut.statesFM[2].X"), + var_ref("ductOut.statesFM[3].X"), + ], + is_matrix: false, + span: test_span(), + }; + let projected_lhs = (0..4) + .map(|k| { + project_scalarized_rhs_expr_at( + &lhs, + k, + &HashSet::new(), + &array_dims, + &record_array_fields, + &IndexMap::new(), + ExpressionForm::Other, + ) + .unwrap() + }) + .collect::>(); + + assert_eq!( + projected_lhs + .iter() + .flat_map(all_var_names) + .collect::>(), + vec![ + "ductOut.statesFM[2].X[1]", + "ductOut.statesFM[2].X[2]", + "ductOut.statesFM[3].X[1]", + "ductOut.statesFM[3].X[2]", + ] + ); +} + #[test] fn test_scalarize_preserves_declared_matrix_vector_equations() { let mut dae = Dae::new(); @@ -1090,3 +1955,47 @@ fn test_scalarize_preserves_declared_matrix_vector_equations() { "declared matrix/vector equation should remain symbolic for structural scalarization" ); } + +fn record_lane_variable(name: &str, declaration: u32) -> dae::Variable { + let mut variable = dae::Variable::new(rumoca_core::VarName::new(name), test_span()); + let mut reference = component_ref(&[("bus", None), ("cells", Some(1)), ("x", None)]); + reference.def_id = Some(rumoca_core::DefId::new(declaration)); + variable.component_ref = Some(reference); + variable +} + +#[test] +fn test_record_array_projection_alias_rejects_direct_name_collision() { + let mut dae = Dae::new(); + let lane = record_lane_variable("bus.cells[1].x", 90_101); + dae.variables.algebraics.insert(lane.name.clone(), lane); + + let mut direct = dae::Variable::new(rumoca_core::VarName::new("bus.cells.x[1]"), test_span()); + let mut direct_ref = component_ref(&[("bus", None), ("cells", None), ("x", Some(1))]); + direct_ref.def_id = Some(rumoca_core::DefId::new(90_102)); + direct.component_ref = Some(direct_ref); + dae.variables.algebraics.insert(direct.name.clone(), direct); + + let error = build_record_array_projection_alias_map(&dae).unwrap_err(); + assert!( + error + .to_string() + .contains("collides with a directly declared variable") + ); +} + +#[test] +fn test_record_array_projection_alias_rejects_duplicate_structured_path() { + let mut dae = Dae::new(); + let first = record_lane_variable("first_lane", 90_111); + let second = record_lane_variable("second_lane", 90_112); + dae.variables.algebraics.insert(first.name.clone(), first); + dae.variables.algebraics.insert(second.name.clone(), second); + + let error = build_record_array_projection_alias_map(&dae).unwrap_err(); + assert!( + error + .to_string() + .contains("conflicting directly declared projection target") + ); +} diff --git a/crates/rumoca-phase-dae/src/dae_lowering/tests/colon_slice_dot.rs b/crates/rumoca-phase-dae/src/dae_lowering/tests/colon_slice_dot.rs new file mode 100644 index 000000000..6f274f861 --- /dev/null +++ b/crates/rumoca-phase-dae/src/dae_lowering/tests/colon_slice_dot.rs @@ -0,0 +1,418 @@ +use super::*; + +fn column_slice(name: &str, column: i64) -> rumoca_core::Expression { + let span = test_span(); + rumoca_core::Expression::Index { + base: Box::new(var_ref(name)), + subscripts: vec![ + rumoca_core::Subscript::generated_colon(span), + rumoca_core::Subscript::generated_index(column, span), + ], + span, + } +} + +fn rotation_column_slice() -> rumoca_core::Expression { + column_slice("rotation", 3) +} + +fn booster_dot_names() -> Vec { + (1..=3) + .flat_map(|lane| { + [ + format!("rotation[{lane},3]"), + format!("deck_normal_w[{lane}]"), + ] + }) + .collect() +} + +fn scalar_if(condition: rumoca_core::Expression) -> rumoca_core::Expression { + let span = test_span(); + rumoca_core::Expression::If { + branches: vec![(condition, var_ref("gain"))], + else_branch: Box::new(rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: Box::new(var_ref("gain")), + rhs: Box::new(int_lit(1)), + span, + }), + span, + } +} + +fn call_count(expr: &rumoca_core::Expression) -> usize { + struct Counter(usize); + + impl rumoca_core::ExpressionVisitor for Counter { + fn visit_builtin_call( + &mut self, + function: &rumoca_core::BuiltinFunction, + args: &[rumoca_core::Expression], + ) { + self.0 += 1; + self.walk_builtin_call(function, args); + } + + fn visit_function_call( + &mut self, + name: &rumoca_core::Reference, + args: &[rumoca_core::Expression], + is_constructor: bool, + ) { + self.0 += 1; + self.walk_function_call(name, args, is_constructor); + } + } + + let mut counter = Counter(0); + rumoca_core::ExpressionVisitor::visit_expression(&mut counter, expr); + counter.0 +} + +fn assert_unknown_if_condition_remains_unprojected(product: rumoca_core::Expression) { + let array_dims = HashMap::from([ + ("rotation".to_string(), vec![3, 3]), + ("gain".to_string(), vec![]), + ]); + let lowered = lower_colon_slice_dot_products(&product, &array_dims) + .expect("unknown condition shape should remain representable"); + + assert!(matches!(lowered, rumoca_core::Expression::Binary { .. })); + assert_eq!(call_count(&lowered), 1); +} + +#[test] +fn colon_slice_times_plain_vector_lowers_to_scalar_dot_product() { + let array_dims = HashMap::from([ + ("rotation".to_string(), vec![3, 3]), + ("deck_normal_w".to_string(), vec![3]), + ("normal_basis".to_string(), vec![3, 2]), + ]); + + let lowered = lower_colon_slice_dot_products( + &mul(rotation_column_slice(), var_ref("deck_normal_w")), + &array_dims, + ) + .expect("proven vector product should lower"); + assert!(!matches!(lowered, rumoca_core::Expression::Array { .. })); + assert_eq!(all_var_names(&lowered), booster_dot_names()); + + let symmetric = lower_colon_slice_dot_products( + &mul(var_ref("deck_normal_w"), rotation_column_slice()), + &array_dims, + ) + .expect("symmetric proven vector product should lower"); + let mut symmetric_names = all_var_names(&symmetric); + let mut lowered_names = all_var_names(&lowered); + symmetric_names.sort(); + lowered_names.sort(); + assert_eq!(symmetric_names, lowered_names); + + let slice_dot = lower_colon_slice_dot_products( + &mul(rotation_column_slice(), column_slice("normal_basis", 2)), + &array_dims, + ) + .expect("two proven slices should lower"); + assert_eq!( + all_var_names(&slice_dot), + vec![ + "rotation[1,3]", + "normal_basis[1,2]", + "rotation[2,3]", + "normal_basis[2,2]", + "rotation[3,3]", + "normal_basis[3,2]", + ] + ); +} + +#[test] +fn colon_slice_dot_product_requires_two_proven_equal_rank_one_vectors() { + let span = test_span(); + let slice = |name: &str, subscripts| rumoca_core::Expression::Index { + base: Box::new(var_ref(name)), + subscripts, + span, + }; + let array_dims = HashMap::from([ + ("rotation".to_string(), vec![3, 3]), + ("short".to_string(), vec![2]), + ("matrix".to_string(), vec![3, 3]), + ("cube".to_string(), vec![2, 2, 2]), + ("vector4".to_string(), vec![4]), + ("short_basis".to_string(), vec![2, 2]), + ("gain".to_string(), vec![]), + ]); + + for invalid in [ + mul(rotation_column_slice(), var_ref("short")), + mul(rotation_column_slice(), var_ref("unknown")), + mul(rotation_column_slice(), var_ref("matrix")), + mul( + rotation_column_slice(), + function_call("unknownVector", vec![]), + ), + mul(rotation_column_slice(), column_slice("short_basis", 1)), + mul( + slice( + "cube", + vec![ + rumoca_core::Subscript::generated_colon(span), + rumoca_core::Subscript::generated_colon(span), + rumoca_core::Subscript::generated_index(1, span), + ], + ), + var_ref("vector4"), + ), + mul( + rotation_column_slice(), + rumoca_core::Expression::Array { + elements: vec![ + rumoca_core::Expression::Array { + elements: vec![int_lit(1), int_lit(2)], + is_matrix: false, + span, + }, + rumoca_core::Expression::Array { + elements: vec![int_lit(3), int_lit(4)], + is_matrix: false, + span, + }, + ], + is_matrix: true, + span, + }, + ), + mul( + rotation_column_slice(), + rumoca_core::Expression::Array { + elements: vec![rumoca_core::Expression::Array { + elements: vec![int_lit(1), int_lit(2), int_lit(3)], + is_matrix: false, + span, + }], + is_matrix: false, + span, + }, + ), + ] { + let lowered = lower_colon_slice_dot_products(&invalid, &array_dims) + .expect("unsupported dot-product shape should remain representable"); + assert!( + matches!(lowered, rumoca_core::Expression::Binary { .. }), + "unsupported product lowered to {lowered:?}" + ); + } + + for scalar in [int_lit(2), var_ref("gain")] { + let scaled = + lower_colon_slice_dot_products(&mul(rotation_column_slice(), scalar), &array_dims) + .expect("slice scaling should remain elementwise"); + assert!(matches!( + scaled, + rumoca_core::Expression::Array { ref elements, .. } if elements.len() == 3 + )); + } + + let elementwise = rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::MulElem, + lhs: Box::new(rotation_column_slice()), + rhs: Box::new(var_ref("short")), + span, + }; + let lowered = lower_colon_slice_dot_products(&elementwise, &array_dims) + .expect("MulElem must not become a dot product"); + assert!(matches!( + lowered, + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::MulElem, + .. + } + )); +} + +#[test] +fn colon_slice_product_with_builtin_call_remains_unprojected() { + let span = test_span(); + let array_dims = HashMap::from([ + ("rotation".to_string(), vec![3, 3]), + ("gain".to_string(), vec![]), + ]); + let product = mul( + rotation_column_slice(), + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Abs, + args: vec![var_ref("gain")], + span, + }, + ); + + let lowered = lower_colon_slice_dot_products(&product, &array_dims) + .expect("builtin-call shape should remain representable"); + assert!(matches!(lowered, rumoca_core::Expression::Binary { .. })); +} + +#[test] +fn scalar_composites_keep_colon_slice_scaling_elementwise_in_both_orders() { + let span = test_span(); + let array_dims = HashMap::from([ + ("rotation".to_string(), vec![3, 3]), + ("gain".to_string(), vec![]), + ]); + let scalars = vec![ + rumoca_core::Expression::Unary { + op: rumoca_core::OpUnary::Minus, + rhs: Box::new(var_ref("gain")), + span, + }, + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: Box::new(var_ref("gain")), + rhs: Box::new(int_lit(1)), + span, + }, + scalar_if(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Boolean(true), + span, + }), + scalar_if(var_ref("gain")), + ]; + + for scalar in scalars { + for product in [ + mul(rotation_column_slice(), scalar.clone()), + mul(scalar, rotation_column_slice()), + ] { + let lowered = lower_colon_slice_dot_products(&product, &array_dims) + .expect("proven scalar composite must preserve vector scaling"); + assert!(matches!( + lowered, + rumoca_core::Expression::Array { ref elements, .. } if elements.len() == 3 + )); + } + } +} + +#[test] +fn slice_left_of_if_with_unknown_condition_remains_single_call() { + let scalar = scalar_if(function_call("unknownCondition", vec![])); + assert_unknown_if_condition_remains_unprojected(mul(rotation_column_slice(), scalar)); +} + +#[test] +fn slice_right_of_if_with_unknown_condition_remains_single_call() { + let scalar = scalar_if(function_call("unknownCondition", vec![])); + assert_unknown_if_condition_remains_unprojected(mul(scalar, rotation_column_slice())); +} + +#[test] +fn slice_if_with_builtin_condition_remains_single_call_in_both_orders() { + let span = test_span(); + let scalar = scalar_if(rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Abs, + args: vec![var_ref("gain")], + span, + }); + + for product in [ + mul(rotation_column_slice(), scalar.clone()), + mul(scalar, rotation_column_slice()), + ] { + assert_unknown_if_condition_remains_unprojected(product); + } +} + +#[test] +fn user_function_argument_rewrites_nested_colon_slice_dot_product() { + let mut dae = Dae::new(); + let span = test_span(); + for (name, dims) in [("rotation", vec![3, 3]), ("deck_normal_w", vec![3])] { + let mut variable = dae::Variable::new(rumoca_core::VarName::new(name), span); + variable.dims = dims; + dae.variables + .algebraics + .insert(rumoca_core::VarName::new(name), variable); + } + dae.continuous.equations.push(dae::Equation::residual( + function_call( + "Pkg.consume", + vec![mul(rotation_column_slice(), var_ref("deck_normal_w"))], + ), + span, + "result = Pkg.consume(rotation[:, 3] * deck_normal_w)", + )); + + scalarize_phantom_vector_equations(&mut dae).expect("scalarize function argument"); + + let rumoca_core::Expression::FunctionCall { args, .. } = &dae.continuous.equations[0].rhs + else { + panic!("expected user function call"); + }; + assert!(matches!( + args.as_slice(), + [rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + .. + }] + )); + assert_eq!(all_var_names(&args[0]), booster_dot_names()); +} + +#[test] +fn booster_nested_min_max_colon_slice_vector_product_is_scalar() { + let mut dae = Dae::new(); + let span = test_span(); + for (name, dims) in [("rotation", vec![3, 3]), ("deck_normal_w", vec![3])] { + let mut variable = dae::Variable::new(rumoca_core::VarName::new(name), span); + variable.dims = dims; + dae.variables + .algebraics + .insert(rumoca_core::VarName::new(name), variable); + } + let clamp = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Min, + args: vec![ + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Max, + args: vec![ + mul(rotation_column_slice(), var_ref("deck_normal_w")), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(-1.0), + span, + }, + ], + span, + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(1.0), + span, + }, + ], + span, + }; + dae.continuous.equations.push(dae::Equation::residual( + clamp, + span, + "body_up_surface_cosine = min(max(rotation[:, 3] * deck_normal_w, -1.0), 1.0)", + )); + + scalarize_phantom_vector_equations(&mut dae).unwrap(); + + let rumoca_core::Expression::BuiltinCall { args: min_args, .. } = + &dae.continuous.equations[0].rhs + else { + panic!("expected outer min"); + }; + let rumoca_core::Expression::BuiltinCall { args: max_args, .. } = &min_args[0] else { + panic!("expected nested max"); + }; + assert!(matches!( + max_args[0], + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + .. + } + )); + assert_eq!(all_var_names(&max_args[0]), booster_dot_names()); +} diff --git a/crates/rumoca-phase-dae/src/dae_lowering/tests/initialization_provenance.rs b/crates/rumoca-phase-dae/src/dae_lowering/tests/initialization_provenance.rs new file mode 100644 index 000000000..5495ab30f --- /dev/null +++ b/crates/rumoca-phase-dae/src/dae_lowering/tests/initialization_provenance.rs @@ -0,0 +1,37 @@ +use super::*; + +#[test] +fn scalarization_repeats_typed_initialization_provenance() { + let mut dae = Dae::new(); + let mut target = dae::Variable::new(rumoca_core::VarName::new("target"), test_span()); + target.dims = vec![3]; + dae.variables + .algebraics + .insert(rumoca_core::VarName::new("target"), target); + for k in 1..=3 { + let name = format!("connector.pin[{k}].v"); + dae.variables.algebraics.insert( + rumoca_core::VarName::new(&name), + dae::Variable::new(rumoca_core::VarName::new(&name), test_span()), + ); + } + dae.initialization + .equations + .push(dae::Equation::residual_array( + sub(var_ref("target"), var_ref("connector.pin.v")), + test_span(), + "phantom initial equation", + 3, + )); + dae.initialization + .equation_provenance + .push(dae::InitializationEquationProvenance::FixedStart); + + scalarize_phantom_vector_equations(&mut dae).unwrap(); + + assert_eq!(dae.initialization.equations.len(), 3); + assert_eq!( + dae.initialization.equation_provenance, + vec![dae::InitializationEquationProvenance::FixedStart; 3] + ); +} diff --git a/crates/rumoca-phase-dae/src/dae_lowering/tests/matrix_product_projection.rs b/crates/rumoca-phase-dae/src/dae_lowering/tests/matrix_product_projection.rs new file mode 100644 index 000000000..a13ce3acb --- /dev/null +++ b/crates/rumoca-phase-dae/src/dae_lowering/tests/matrix_product_projection.rs @@ -0,0 +1,116 @@ +use super::*; + +#[test] +fn bare_negative_dimensions_propagate_from_projection_entry() { + let dimensions = HashMap::from([("A".to_string(), vec![-1]), ("B".to_string(), vec![-1])]); + let expression = mul(var_ref("A"), var_ref("B")); + let error = Projector(&dimensions, &IndexMap::new()) + .project(&expression, 0, &[]) + .expect_err("invalid bare operands must fail closed during projection entry"); + + assert!(error.to_string().contains("negative dimension")); + assert_eq!(error.source_span(), Some(test_span())); +} + +#[test] +fn scalar_product_sum_ignores_surrounding_equation_lane() { + let dimensions = HashMap::from([ + ("p".to_string(), vec![3]), + ("R".to_string(), vec![3, 3]), + ("leg_r_b".to_string(), vec![3, 4]), + ]); + let indexed = |name, subscripts| rumoca_core::Expression::Index { + base: Box::new(var_ref(name)), + subscripts, + span: test_span(), + }; + let index = |value| rumoca_core::Subscript::Index { + value, + span: test_span(), + }; + let colon = || rumoca_core::Subscript::Colon { span: test_span() }; + let expression = rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: Box::new(indexed("p", vec![index(3)])), + rhs: Box::new(mul( + indexed("R", vec![index(3), colon()]), + indexed("leg_r_b", vec![colon(), index(2)]), + )), + span: test_span(), + }; + + let projected = Projector(&dimensions, &IndexMap::new()) + .project(&expression, 1, &[]) + .expect("a scalar product result must not consume the surrounding equation lane") + .expect("the nested vector product must be projected"); + + let rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + rhs, + .. + } = projected + else { + panic!("expected scalar sum projection"); + }; + assert_eq!( + all_var_names(&rhs), + [ + "R[3,1]", + "leg_r_b[1,2]", + "R[3,2]", + "leg_r_b[2,2]", + "R[3,3]", + "leg_r_b[3,2]", + ] + ); +} + +#[test] +fn scalar_residual_keeps_colon_slice_dot_product_lane_local() { + let mut dae = Dae::new(); + let span = test_span(); + for (name, dims) in [("p", vec![3]), ("R", vec![3, 3]), ("leg_r_b", vec![3, 4])] { + let mut variable = dae::Variable::new(rumoca_core::VarName::new(name), span); + variable.dims = dims; + dae.variables + .algebraics + .insert(rumoca_core::VarName::new(name), variable); + } + let index = |value| rumoca_core::Subscript::Index { value, span }; + let colon = || rumoca_core::Subscript::Colon { span }; + let reference = |name, subscripts| rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new(name).into(), + subscripts, + span, + }; + let residual = rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(reference("p", vec![index(3)])), + rhs: Box::new(mul( + reference("R", vec![index(3), colon()]), + reference("leg_r_b", vec![colon(), index(2)]), + )), + span, + }; + dae.continuous.equations.push(dae::Equation::residual( + residual, + span, + "p[3] = R[3, :] * leg_r_b[:, 2]", + )); + + scalarize_phantom_vector_equations(&mut dae).expect("lower scalar matrix product residual"); + + assert_eq!(dae.continuous.equations.len(), 1); + assert_eq!( + all_var_names(&dae.continuous.equations[0].rhs), + [ + "p[3]", + "R[3,1]", + "leg_r_b[1,2]", + "R[3,2]", + "leg_r_b[2,2]", + "R[3,3]", + "leg_r_b[3,2]", + ] + ); +} diff --git a/crates/rumoca-phase-dae/src/dae_lowering/tests/record_array_member.rs b/crates/rumoca-phase-dae/src/dae_lowering/tests/record_array_member.rs new file mode 100644 index 000000000..67a92560e --- /dev/null +++ b/crates/rumoca-phase-dae/src/dae_lowering/tests/record_array_member.rs @@ -0,0 +1,81 @@ +use super::*; + +#[test] +fn test_scalarize_record_array_member_index_projects_record_field() { + let mut dae = Dae::new(); + + for record_index in [1, 2, 3] { + let name = format!("ductOut.statesFM[{record_index}].X"); + let mut var = dae::Variable::new(rumoca_core::VarName::new(&name), test_span()); + var.dims = vec![2]; + var.component_ref = Some(component_ref(&[ + ("ductOut", None), + ("statesFM", Some(record_index)), + ("X", None), + ])); + dae.variables.algebraics.insert(var.name.clone(), var); + } + + let record_array_fields = build_record_array_field_map(&dae); + assert_eq!( + record_array_fields + .get("ductOut.statesFM.X") + .map(|entry| entry.field_dims.as_slice()), + Some(&[2][..]) + ); + let mut array_dims = build_dae_var_dims_map(&dae); + array_dims.retain(|_, dims| !dims.is_empty()); + let subscript_cases = [ + rumoca_core::Subscript::Index { + value: 3, + span: test_span(), + }, + rumoca_core::Subscript::Expr { + expr: Box::new(int_lit(3)), + span: test_span(), + }, + rumoca_core::Subscript::Expr { + expr: Box::new(sub(int_lit(4), int_lit(1))), + span: test_span(), + }, + ]; + + for subscript in subscript_cases { + let expr = rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::Index { + base: Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference( + "ductOut.statesFM", + component_ref(&[("ductOut", None), ("statesFM", None)]), + ), + subscripts: Vec::new(), + span: test_span(), + }), + subscripts: vec![subscript], + span: test_span(), + }), + field: "X".to_string(), + span: test_span(), + }; + + let projected = (0..2) + .map(|k| { + project_scalarized_rhs_expr_at( + &expr, + k, + &HashSet::new(), + &array_dims, + &record_array_fields, + &IndexMap::new(), + ExpressionForm::Other, + ) + .unwrap() + }) + .collect::>(); + + assert_eq!( + projected.iter().flat_map(all_var_names).collect::>(), + vec!["ductOut.statesFM[3].X[1]", "ductOut.statesFM[3].X[2]"] + ); + } +} diff --git a/crates/rumoca-phase-dae/src/dae_lowering/tests/record_array_projection_alias.rs b/crates/rumoca-phase-dae/src/dae_lowering/tests/record_array_projection_alias.rs new file mode 100644 index 000000000..76e8216cf --- /dev/null +++ b/crates/rumoca-phase-dae/src/dae_lowering/tests/record_array_projection_alias.rs @@ -0,0 +1,91 @@ +use super::*; + +#[test] +fn test_record_array_projection_alias_only_resolves_concrete_indices() { + let mut dae = Dae::new(); + let lane = record_lane_variable("bus.cells[1].x", 90_121); + dae.variables.algebraics.insert(lane.name.clone(), lane); + let aliases = build_record_array_projection_alias_map(&dae).unwrap(); + + let mut concrete = component_ref(&[("bus", None), ("cells", None), ("x", Some(1))]); + concrete.def_id = None; + assert_eq!( + aliases.resolve(&concrete).unwrap().as_str(), + "bus.cells[1].x" + ); + concrete.def_id = Some(rumoca_core::DefId::new(99_999)); + assert!( + aliases.resolve(&concrete).is_none(), + "a resolved but mismatched declaration identity must not fall back to path-only lookup" + ); + + let mut colon = component_ref(&[("bus", None), ("cells", None), ("x", None)]); + colon.parts.last_mut().unwrap().subs = + vec![rumoca_core::Subscript::Colon { span: test_span() }]; + assert!(aliases.resolve(&colon).is_none()); + + let mut expression = component_ref(&[("bus", None), ("cells", None), ("x", None)]); + expression.parts.last_mut().unwrap().subs = vec![rumoca_core::Subscript::Expr { + expr: Box::new(var_ref("i")), + span: test_span(), + }]; + assert!(aliases.resolve(&expression).is_none()); +} + +fn record_lane_alias_identity( + variable: &dae::Variable, +) -> ( + StructuredProjectionPath, + StructuredProjectionIdentity, + rumoca_core::Reference, +) { + let component_ref = variable.component_ref.as_ref().unwrap(); + let mut projection = component_ref.clone(); + let subscripts = std::mem::take(&mut projection.parts[1].subs); + projection.parts.last_mut().unwrap().subs.extend(subscripts); + let path = StructuredProjectionPath::from_component_ref(&projection).unwrap(); + let identity = StructuredProjectionIdentity { + path: path.clone(), + declaration: component_ref.def_id.unwrap(), + }; + let reference = rumoca_core::Reference::with_component_reference( + variable.name.as_str(), + component_ref.clone(), + ); + (path, identity, reference) +} + +#[test] +fn test_record_array_projection_alias_duplicate_same_identity_and_target_is_idempotent() { + let variable = record_lane_variable("bus.cells[1].x", 90_131); + let mut aliases = RecordArrayProjectionAliases::default(); + let direct = HashMap::new(); + + append_record_array_projection_aliases(&mut aliases, &direct, &variable.name, &variable) + .unwrap(); + append_record_array_projection_aliases(&mut aliases, &direct, &variable.name, &variable) + .unwrap(); + + assert_eq!(aliases.aliases.len(), 1); + assert_eq!(aliases.unique_identity_by_path.len(), 1); +} + +#[test] +fn test_record_array_projection_alias_same_direct_identity_and_target_is_noop() { + let variable = record_lane_variable("bus.cells[1].x", 90_141); + let (path, identity, reference) = record_lane_alias_identity(&variable); + let direct = HashMap::from([( + path, + DirectProjectionTarget { + identity: Some(identity), + reference, + }, + )]); + let mut aliases = RecordArrayProjectionAliases::default(); + + append_record_array_projection_aliases(&mut aliases, &direct, &variable.name, &variable) + .unwrap(); + + assert!(aliases.aliases.is_empty()); + assert!(aliases.unique_identity_by_path.is_empty()); +} diff --git a/crates/rumoca-phase-dae/src/equation_conversion.rs b/crates/rumoca-phase-dae/src/equation_conversion.rs index e7ae2ec0c..abea55399 100644 --- a/crates/rumoca-phase-dae/src/equation_conversion.rs +++ b/crates/rumoca-phase-dae/src/equation_conversion.rs @@ -1,4 +1,9 @@ //! Continuous equation filtering/conversion for ToDAE. +//! +//! SPEC_0021 file-size exception: this module still owns the current Flat +//! equation-to-DAE conversion path. split plan: move connector/input boundary +//! filtering, structured-equation expansion, and expression conversion context +//! into separate sibling modules as each concern receives focused tests. use std::collections::{HashMap, HashSet}; @@ -45,6 +50,85 @@ pub(super) fn is_input_input_connection(eq: &flat::Equation, dae: &dae::Dae) -> } } +fn is_flat_stream_name(flat: &flat::Model, name: &rumoca_core::VarName) -> bool { + name_resolution::resolve_var_name_with_subscript_fallback(name, |candidate| { + flat.variables + .get(candidate) + .is_some_and(|variable| variable.stream) + }) + .is_some() +} + +fn connection_sides(eq: &flat::Equation) -> Option<(rumoca_core::VarName, rumoca_core::VarName)> { + if !eq.origin.is_connection() { + return None; + } + + let rumoca_core::Expression::Binary { op, lhs, rhs, .. } = &eq.residual else { + return None; + }; + if !matches!(op, rumoca_core::OpBinary::Sub) { + return None; + } + + Some(( + name_resolution::extract_varref_name(lhs)?, + name_resolution::extract_varref_name(rhs)?, + )) +} + +pub(super) fn is_stream_stream_connection(eq: &flat::Equation, flat: &flat::Model) -> bool { + let Some((lhs_name, rhs_name)) = connection_sides(eq) else { + return false; + }; + is_flat_stream_name(flat, &lhs_name) && is_flat_stream_name(flat, &rhs_name) +} + +fn should_skip_stream_stream_connection(eq: &flat::Equation, ctx: &EqFilterContext<'_>) -> bool { + let Some((lhs_name, rhs_name)) = connection_sides(eq) else { + return false; + }; + if !is_stream_stream_connection(eq, ctx.flat) { + return false; + } + + let lhs_consumed = + name_resolution::resolve_name_against_set(&lhs_name, ctx.non_connection_rhs_var_refs) + .is_some(); + let rhs_consumed = + name_resolution::resolve_name_against_set(&rhs_name, ctx.non_connection_rhs_var_refs) + .is_some(); + !(lhs_consumed && rhs_consumed) +} + +fn is_identity_equation(eq: &flat::Equation) -> bool { + let rumoca_core::Expression::Binary { op, lhs, rhs, .. } = &eq.residual else { + return false; + }; + if !matches!(op, rumoca_core::OpBinary::Sub) { + return false; + } + let rumoca_core::Expression::VarRef { + name: lhs_name, + subscripts: lhs_subscripts, + .. + } = lhs.as_ref() + else { + return false; + }; + let rumoca_core::Expression::VarRef { + name: rhs_name, + subscripts: rhs_subscripts, + .. + } = rhs.as_ref() + else { + return false; + }; + + lhs_name.var_name() == rhs_name.var_name() + && subscripts_match_semantically(lhs_subscripts, rhs_subscripts) +} + /// Check if an equation defines an input variable with a constant/parameter value. /// /// Equations of the form `input_var = literal` where `input_var` is an input @@ -315,6 +399,8 @@ struct EqFilterStats { kept_other: usize, skipped_top_level_oc: usize, skipped_input_input: usize, + skipped_stream_stream: usize, + skipped_identity: usize, skipped_output_alias: usize, skipped_input_default: usize, skipped_explicit_zero: usize, @@ -333,13 +419,15 @@ impl EqFilterStats { fn log(&self) { crate::log_equation_filter_debug(format!( - "eq-filter: kept(connection={}, flow_sum={}, unconnected_flow={}, other={}) skipped(top_level_oc={}, input_input={}, output_alias={}, input_default={}, explicit_zero={}, inferred_zero={})", + "eq-filter: kept(connection={}, flow_sum={}, unconnected_flow={}, other={}) skipped(top_level_oc={}, input_input={}, stream_stream={}, identity={}, output_alias={}, input_default={}, explicit_zero={}, inferred_zero={})", self.kept_connection, self.kept_flow_sum, self.kept_unconnected_flow, self.kept_other, self.skipped_top_level_oc, self.skipped_input_input, + self.skipped_stream_stream, + self.skipped_identity, self.skipped_output_alias, self.skipped_input_default, self.skipped_explicit_zero, @@ -423,6 +511,18 @@ fn skip_equation_pre_classification( return true; } + if should_skip_stream_stream_connection(eq, ctx) { + stats.skipped_stream_stream += 1; + log_skip(ctx.debug_eq_filter, "stream_stream", eq); + return true; + } + + if is_identity_equation(eq) { + stats.skipped_identity += 1; + log_skip(ctx.debug_eq_filter, "identity", eq); + return true; + } + if let Some(output_name) = output_alias_skip_reason(eq, ctx, dae) { stats.skipped_output_alias += 1; if ctx.debug_eq_filter { @@ -454,7 +554,7 @@ fn skip_equation_pre_classification( fn compute_scalar_count( eq: &flat::Equation, flat: &flat::Model, - prefix_counts: &FxHashMap, + prefix_counts: &super::ScalarInferenceMetadata, linearized_embedded_lhs_bases: &HashSet, stats: &mut EqFilterStats, debug_eq_filter: bool, @@ -513,6 +613,7 @@ fn record_field_specs_for_call( #[derive(Debug, Clone)] struct RecordFieldSpec { param: rumoca_core::FunctionParam, + match_by_name: bool, } impl RecordFieldSpec { @@ -527,7 +628,10 @@ impl RecordFieldSpec { (!params.is_empty()).then(|| { params .into_iter() - .map(|param| Self { param }) + .map(|param| Self { + param, + match_by_name: false, + }) .collect::>() }) } @@ -552,13 +656,23 @@ impl RecordFieldSpec { } } + fn is_lhs_derived(&self) -> bool { + self.match_by_name + } + fn matches_component_ref( &self, field_ref: &rumoca_core::ComponentReference, symbol_ancestry: &IndexMap>, ) -> bool { self.param.def_id.is_some_and(|expected| { - field_ref.def_id == Some(expected) + let name_matches = self.match_by_name + && field_ref + .parts + .last() + .is_some_and(|part| part.ident == self.param.name); + name_matches + || field_ref.def_id == Some(expected) || field_ref.def_id.is_some_and(|actual| { symbol_ancestry .get(&actual) @@ -642,6 +756,12 @@ fn rhs_field_expression( return arg; } + if let Some(projected) = + project_complex_field_expression(rhs, field.name(), flat, equation_span) + { + return projected; + } + let span = match rhs.span() { Some(span) => span, None => equation_span, @@ -649,6 +769,317 @@ fn rhs_field_expression( field.field_access(rhs.clone(), span) } +fn project_complex_field_expression( + expr: &rumoca_core::Expression, + field: &str, + flat: &flat::Model, + context_span: rumoca_core::Span, +) -> Option { + let (re, im) = complex_parts(expr, flat, context_span)?; + match field { + "re" => Some(re), + "im" => Some(im), + _ => None, + } +} + +fn complex_parts( + expr: &rumoca_core::Expression, + flat: &flat::Model, + context_span: rumoca_core::Span, +) -> Option<(rumoca_core::Expression, rumoca_core::Expression)> { + match expr { + rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor, + .. + } => complex_function_call_parts(name, args, *is_constructor, flat, expr, context_span), + rumoca_core::Expression::VarRef { + name, + subscripts, + span, + } => complex_var_ref_parts(expr, name, subscripts, *span, flat), + rumoca_core::Expression::Literal { span, .. } => Some((expr.clone(), zero_literal(*span))), + rumoca_core::Expression::Unary { op, rhs, span } => { + complex_unary_parts(op, rhs, *span, flat) + } + rumoca_core::Expression::Binary { op, lhs, rhs, span } => { + complex_binary_parts(op, lhs, rhs, *span, flat) + } + rumoca_core::Expression::FieldAccess { span, .. } => { + Some((expr.clone(), zero_literal(*span))) + } + _ => None, + } +} + +fn complex_function_call_parts( + name: &rumoca_core::Reference, + args: &[rumoca_core::Expression], + is_constructor: bool, + flat: &flat::Model, + expr: &rumoca_core::Expression, + context_span: rumoca_core::Span, +) -> Option<(rumoca_core::Expression, rumoca_core::Expression)> { + if !flat + .functions + .get(name.var_name()) + .is_some_and(|function| is_constructor || function.is_constructor) + { + return None; + } + let fields = record_field_specs_for_call(name, is_constructor, flat)?; + let re = constructor_arg_by_field_name(args, &fields, "re") + .unwrap_or_else(|| zero_literal(expr_or_context_span(expr, context_span))); + let im = constructor_arg_by_field_name(args, &fields, "im") + .unwrap_or_else(|| zero_literal(expr_or_context_span(expr, context_span))); + Some((re, im)) +} + +fn complex_var_ref_parts( + expr: &rumoca_core::Expression, + name: &rumoca_core::Reference, + subscripts: &[rumoca_core::Subscript], + span: rumoca_core::Span, + flat: &flat::Model, +) -> Option<(rumoca_core::Expression, rumoca_core::Expression)> { + if !subscripts.is_empty() { + return None; + } + if let Some(parts) = aggregate_var_complex_parts(name.var_name(), flat, span) { + return Some(parts); + } + flat.variables.get(name.var_name()).and_then(|variable| { + (variable.is_primitive && super::compute_var_size(&variable.dims) == 1) + .then(|| (expr.clone(), zero_literal(span))) + }) +} + +fn complex_unary_parts( + op: &rumoca_core::OpUnary, + rhs: &rumoca_core::Expression, + span: rumoca_core::Span, + flat: &flat::Model, +) -> Option<(rumoca_core::Expression, rumoca_core::Expression)> { + let (re, im) = complex_parts(rhs, flat, span)?; + match op { + rumoca_core::OpUnary::Plus | rumoca_core::OpUnary::DotPlus => Some((re, im)), + rumoca_core::OpUnary::Minus | rumoca_core::OpUnary::DotMinus => { + Some((unary_minus(re, span), unary_minus(im, span))) + } + _ => None, + } +} + +fn complex_binary_parts( + op: &rumoca_core::OpBinary, + lhs: &rumoca_core::Expression, + rhs: &rumoca_core::Expression, + span: rumoca_core::Span, + flat: &flat::Model, +) -> Option<(rumoca_core::Expression, rumoca_core::Expression)> { + let lhs_parts = complex_parts(lhs, flat, span)?; + let rhs_parts = complex_parts(rhs, flat, span)?; + match op { + rumoca_core::OpBinary::Add | rumoca_core::OpBinary::AddElem => Some(complex_add_sub_parts( + op.clone(), + lhs_parts, + rhs_parts, + span, + )), + rumoca_core::OpBinary::Sub | rumoca_core::OpBinary::SubElem => Some(complex_add_sub_parts( + op.clone(), + lhs_parts, + rhs_parts, + span, + )), + rumoca_core::OpBinary::Mul | rumoca_core::OpBinary::MulElem => { + Some(complex_product_parts(lhs_parts, rhs_parts, span)) + } + rumoca_core::OpBinary::Div | rumoca_core::OpBinary::DivElem => { + Some(complex_quotient_parts(lhs_parts, rhs_parts, span)) + } + _ => None, + } +} + +fn complex_add_sub_parts( + op: rumoca_core::OpBinary, + lhs: (rumoca_core::Expression, rumoca_core::Expression), + rhs: (rumoca_core::Expression, rumoca_core::Expression), + span: rumoca_core::Span, +) -> (rumoca_core::Expression, rumoca_core::Expression) { + ( + binary_expr(op.clone(), lhs.0, rhs.0, span), + binary_expr(op, lhs.1, rhs.1, span), + ) +} + +fn complex_product_parts( + lhs: (rumoca_core::Expression, rumoca_core::Expression), + rhs: (rumoca_core::Expression, rumoca_core::Expression), + span: rumoca_core::Span, +) -> (rumoca_core::Expression, rumoca_core::Expression) { + let re = binary_expr( + rumoca_core::OpBinary::Sub, + binary_expr( + rumoca_core::OpBinary::Mul, + lhs.0.clone(), + rhs.0.clone(), + span, + ), + binary_expr( + rumoca_core::OpBinary::Mul, + lhs.1.clone(), + rhs.1.clone(), + span, + ), + span, + ); + let im = binary_expr( + rumoca_core::OpBinary::Add, + binary_expr(rumoca_core::OpBinary::Mul, lhs.0, rhs.1, span), + binary_expr(rumoca_core::OpBinary::Mul, lhs.1, rhs.0, span), + span, + ); + (re, im) +} + +fn complex_quotient_parts( + lhs: (rumoca_core::Expression, rumoca_core::Expression), + rhs: (rumoca_core::Expression, rumoca_core::Expression), + span: rumoca_core::Span, +) -> (rumoca_core::Expression, rumoca_core::Expression) { + let denominator = complex_quotient_denominator(&rhs, span); + let re = binary_expr( + rumoca_core::OpBinary::Div, + binary_expr( + rumoca_core::OpBinary::Add, + binary_expr( + rumoca_core::OpBinary::Mul, + lhs.0.clone(), + rhs.0.clone(), + span, + ), + binary_expr( + rumoca_core::OpBinary::Mul, + lhs.1.clone(), + rhs.1.clone(), + span, + ), + span, + ), + denominator.clone(), + span, + ); + let im = binary_expr( + rumoca_core::OpBinary::Div, + binary_expr( + rumoca_core::OpBinary::Sub, + binary_expr(rumoca_core::OpBinary::Mul, lhs.1, rhs.0, span), + binary_expr(rumoca_core::OpBinary::Mul, lhs.0, rhs.1, span), + span, + ), + denominator, + span, + ); + (re, im) +} + +fn complex_quotient_denominator( + rhs: &(rumoca_core::Expression, rumoca_core::Expression), + span: rumoca_core::Span, +) -> rumoca_core::Expression { + binary_expr( + rumoca_core::OpBinary::Add, + binary_expr( + rumoca_core::OpBinary::Mul, + rhs.0.clone(), + rhs.0.clone(), + span, + ), + binary_expr( + rumoca_core::OpBinary::Mul, + rhs.1.clone(), + rhs.1.clone(), + span, + ), + span, + ) +} + +fn constructor_arg_by_field_name( + args: &[rumoca_core::Expression], + fields: &[RecordFieldSpec], + name: &str, +) -> Option { + fields + .iter() + .enumerate() + .find(|(_, field)| field.name() == name) + .and_then(|(index, field)| constructor_field_arg(args, field, index)) +} + +fn aggregate_var_complex_parts( + name: &rumoca_core::VarName, + flat: &flat::Model, + span: rumoca_core::Span, +) -> Option<(rumoca_core::Expression, rumoca_core::Expression)> { + let re_name = rumoca_core::VarName::new(format!("{}.re", name.as_str())); + let im_name = rumoca_core::VarName::new(format!("{}.im", name.as_str())); + let re_var = flat.variables.get(&re_name)?; + let im_var = flat.variables.get(&im_name)?; + Some(( + rumoca_core::Expression::VarRef { + name: reference_for_variable(re_var), + subscripts: Vec::new(), + span, + }, + rumoca_core::Expression::VarRef { + name: reference_for_variable(im_var), + subscripts: Vec::new(), + span, + }, + )) +} + +fn expr_or_context_span( + expr: &rumoca_core::Expression, + context_span: rumoca_core::Span, +) -> rumoca_core::Span { + expr.span().unwrap_or(context_span) +} + +fn zero_literal(span: rumoca_core::Span) -> rumoca_core::Expression { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(0.0), + span, + } +} + +fn unary_minus(expr: rumoca_core::Expression, span: rumoca_core::Span) -> rumoca_core::Expression { + rumoca_core::Expression::Unary { + op: rumoca_core::OpUnary::Minus, + rhs: Box::new(expr), + span, + } +} + +fn binary_expr( + op: rumoca_core::OpBinary, + lhs: rumoca_core::Expression, + rhs: rumoca_core::Expression, + span: rumoca_core::Span, +) -> rumoca_core::Expression { + rumoca_core::Expression::Binary { + op, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span, + } +} + fn reference_for_variable(field_var: &flat::Variable) -> rumoca_core::Reference { if let Some(component_ref) = field_var.component_ref.clone() { return rumoca_core::Reference::with_component_reference( @@ -682,10 +1113,21 @@ fn field_var_ref(field_var: &flat::Variable) -> rumoca_core::Expression { fn field_lhs_expression( field_vars: &[&flat::Variable], + selection_subscripts: Option<&[rumoca_core::Subscript]>, equation_span: rumoca_core::Span, ) -> rumoca_core::Expression { if let [field_var] = field_vars { - return field_var_ref(field_var); + let mut expr = field_var_ref(field_var); + if let Some(subscripts) = selection_subscripts + && !subscripts.is_empty() + { + expr = rumoca_core::Expression::Index { + base: Box::new(expr), + subscripts: subscripts.to_vec(), + span: equation_span, + }; + } + return expr; } rumoca_core::Expression::Array { elements: field_vars @@ -705,11 +1147,28 @@ fn field_scalar_count(field_vars: &[&flat::Variable]) -> usize { .max(1) } +fn selected_field_scalar_count( + field_vars: &[&flat::Variable], + selection_subscripts: Option<&[rumoca_core::Subscript]>, + flat: &flat::Model, +) -> usize { + if let ([field_var], Some(subscripts)) = (field_vars, selection_subscripts) + && !subscripts.is_empty() + && let Some(size) = + super::compute_subscripted_size_with_context(&field_var.dims, subscripts, flat) + { + return size.max(1); + } + field_scalar_count(field_vars) +} + fn component_ref_matches_record_field( lhs_ref: &rumoca_core::ComponentReference, + field_var: &flat::Variable, field_ref: &rumoca_core::ComponentReference, field: &RecordFieldSpec, symbol_ancestry: &IndexMap>, + flat: &flat::Model, ) -> bool { let lhs_parts = lhs_ref.parts.as_slice(); let field_parts = field_ref.parts.as_slice(); @@ -734,7 +1193,42 @@ fn component_ref_matches_record_field( } continue; } - if record_owner_subscripts_match(lhs_part, field_part, field_leaf) { + if record_owner_subscripts_match(lhs_part, field_part, field_leaf, field_var, flat) { + continue; + } + return false; + } + true +} + +fn component_ref_is_record_field_child( + lhs_ref: &rumoca_core::ComponentReference, + field_ref: &rumoca_core::ComponentReference, + flat: &flat::Model, +) -> bool { + let lhs_parts = lhs_ref.parts.as_slice(); + let field_parts = field_ref.parts.as_slice(); + if field_parts.len() != lhs_parts.len() + 1 { + return false; + } + + let Some(lhs_leaf_index) = lhs_parts.len().checked_sub(1) else { + return false; + }; + let field_leaf = &field_parts[field_parts.len() - 1]; + for (lhs_index, lhs_part) in lhs_parts.iter().enumerate() { + let field_part = &field_parts[lhs_index]; + if lhs_part.ident != field_part.ident { + return false; + } + if lhs_index != lhs_leaf_index { + if !subscripts_match_semantically(&lhs_part.subs, &field_part.subs) { + return false; + } + continue; + } + if record_owner_subscripts_match_without_field_dims(lhs_part, field_part, field_leaf, flat) + { continue; } return false; @@ -746,6 +1240,34 @@ fn record_owner_subscripts_match( lhs_part: &rumoca_core::ComponentRefPart, field_owner_part: &rumoca_core::ComponentRefPart, field_leaf: &rumoca_core::ComponentRefPart, + field_var: &flat::Variable, + flat: &flat::Model, +) -> bool { + if lhs_part.subs.is_empty() { + return true; + } + if subscripts_match_semantically(&lhs_part.subs, &field_owner_part.subs) { + return true; + } + if subscripts_select_field_owner(&lhs_part.subs, &field_owner_part.subs, flat) { + return true; + } + if field_owner_part.subs.is_empty() + && subscript_prefix_matches(&field_leaf.subs, &lhs_part.subs) + { + return true; + } + field_owner_part.subs.is_empty() + && field_leaf.subs.is_empty() + && !field_var.dims.is_empty() + && lhs_part.subs.len() <= field_var.dims.len() +} + +fn record_owner_subscripts_match_without_field_dims( + lhs_part: &rumoca_core::ComponentRefPart, + field_owner_part: &rumoca_core::ComponentRefPart, + field_leaf: &rumoca_core::ComponentRefPart, + flat: &flat::Model, ) -> bool { if lhs_part.subs.is_empty() { return true; @@ -753,9 +1275,189 @@ fn record_owner_subscripts_match( if subscripts_match_semantically(&lhs_part.subs, &field_owner_part.subs) { return true; } + if subscripts_select_field_owner(&lhs_part.subs, &field_owner_part.subs, flat) { + return true; + } field_owner_part.subs.is_empty() && subscript_prefix_matches(&field_leaf.subs, &lhs_part.subs) } +fn subscripts_select_field_owner( + lhs: &[rumoca_core::Subscript], + field_owner: &[rumoca_core::Subscript], + flat: &flat::Model, +) -> bool { + lhs.len() == field_owner.len() + && lhs + .iter() + .zip(field_owner.iter()) + .all(|(selection, concrete)| subscript_selects_index(selection, concrete, flat)) +} + +fn subscript_selects_index( + selection: &rumoca_core::Subscript, + concrete: &rumoca_core::Subscript, + flat: &flat::Model, +) -> bool { + let rumoca_core::Subscript::Index { value, .. } = concrete else { + return subscript_matches_semantically(selection, concrete); + }; + match selection { + rumoca_core::Subscript::Index { + value: selected, .. + } => selected == value, + rumoca_core::Subscript::Expr { expr, .. } => { + expression_selects_index(expr, *value, flat).unwrap_or(false) + } + rumoca_core::Subscript::Colon { .. } => true, + } +} + +fn expression_selects_index( + expr: &rumoca_core::Expression, + index: i64, + flat: &flat::Model, +) -> Option { + match expr { + rumoca_core::Expression::Range { + start, step, end, .. + } => { + let start = eval_integer_expression(start, flat)?; + let end = eval_integer_expression(end, flat)?; + let step = step + .as_deref() + .map(|step| eval_integer_expression(step, flat)) + .unwrap_or(Some(1))?; + if step == 0 { + return Some(false); + } + Some(if step > 0 { + index >= start && index <= end && (index - start) % step == 0 + } else { + index <= start && index >= end && (start - index) % (-step) == 0 + }) + } + _ => eval_integer_expression(expr, flat).map(|selected| selected == index), + } +} + +fn eval_integer_expression(expr: &rumoca_core::Expression, flat: &flat::Model) -> Option { + match expr { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + .. + } => Some(*value), + rumoca_core::Expression::VarRef { + name, subscripts, .. + } => { + let binding = &flat.variables.get(name.var_name())?.binding; + let Some(binding) = binding else { + return None; + }; + if let Some(index) = eval_single_array_subscript(subscripts, flat) { + let values = eval_integer_array_expression(binding, flat, 0)?; + let selected = usize::try_from(index.checked_sub(1)?).ok()?; + return values.get(selected).copied(); + } + eval_integer_expression(binding, flat) + } + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs, + rhs, + .. + } => Some(eval_integer_expression(lhs, flat)? + eval_integer_expression(rhs, flat)?), + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs, + rhs, + .. + } => Some(eval_integer_expression(lhs, flat)? - eval_integer_expression(rhs, flat)?), + _ => None, + } +} + +fn eval_single_array_subscript( + subscripts: &[rumoca_core::Subscript], + flat: &flat::Model, +) -> Option { + let [subscript] = subscripts else { + return None; + }; + match subscript { + rumoca_core::Subscript::Index { value, .. } => Some(*value), + rumoca_core::Subscript::Expr { expr, .. } => eval_integer_expression(expr, flat), + rumoca_core::Subscript::Colon { .. } => None, + } +} + +fn eval_integer_array_expression( + expr: &rumoca_core::Expression, + flat: &flat::Model, + depth: u8, +) -> Option> { + if depth > 8 { + return None; + } + match expr { + rumoca_core::Expression::Array { + elements, + is_matrix: false, + .. + } => elements + .iter() + .map(|element| eval_integer_expression(element, flat)) + .collect(), + rumoca_core::Expression::VarRef { name, .. } => { + let binding = flat.variables.get(name.var_name())?.binding.as_ref()?; + eval_integer_array_expression(binding, flat, depth + 1) + } + rumoca_core::Expression::FunctionCall { name, args, .. } + if is_polyphase_index_non_positive_sequence(name.as_str()) && args.len() == 1 => + { + let m = eval_integer_expression(&args[0], flat)?; + index_non_positive_sequence(m) + } + _ => None, + } +} + +fn is_polyphase_index_non_positive_sequence(name: &str) -> bool { + matches!( + name, + "indexNonPositiveSequence" + | "Modelica.Electrical.Polyphase.Functions.indexNonPositiveSequence" + ) +} + +fn number_of_symmetric_base_systems(m: i64) -> Option { + if m <= 0 { + return None; + } + if m % 2 != 0 || m == 2 { + return Some(1); + } + Some(2 * number_of_symmetric_base_systems(m / 2)?) +} + +fn index_non_positive_sequence(m: i64) -> Option> { + let n_base = number_of_symmetric_base_systems(m)?; + let m_base = m.checked_div(n_base)?; + if m_base == 1 { + return Some(Vec::new()); + } + if m_base == 2 { + return Some((1..=n_base).map(|k| 2 + 2 * (k - 1)).collect()); + } + + let mut values = Vec::new(); + for k in 1..=n_base { + for value in 2..=m_base { + values.push(value + m_base * (k - 1)); + } + } + Some(values) +} + fn subscripts_match_semantically( lhs: &[rumoca_core::Subscript], rhs: &[rumoca_core::Subscript], @@ -846,9 +1548,11 @@ fn record_field_variables<'a>( let field_ref = field_var.component_ref.as_ref()?; component_ref_matches_record_field( lhs_component_ref, + field_var, field_ref, field, &flat.symbol_ancestry, + flat, ) .then_some((field_var, field_ref)) }) @@ -861,6 +1565,31 @@ fn record_field_variables<'a>( .collect()) } +fn indexed_record_field_selection_subscripts<'a>( + lhs_name: &'a rumoca_core::Reference, + field_vars: &[&flat::Variable], +) -> Option<&'a [rumoca_core::Subscript]> { + let [field_var] = field_vars else { + return None; + }; + if field_var.dims.is_empty() { + return None; + } + let lhs_ref = lhs_name.component_ref()?; + let field_ref = field_var.component_ref.as_ref()?; + if field_ref.parts.len() != lhs_ref.parts.len() + 1 { + return None; + } + let leaf_index = lhs_ref.parts.len().checked_sub(1)?; + let lhs_leaf = &lhs_ref.parts[leaf_index]; + let field_owner = &field_ref.parts[leaf_index]; + let field_leaf = field_ref.parts.last()?; + if lhs_leaf.subs.is_empty() || !field_owner.subs.is_empty() || !field_leaf.subs.is_empty() { + return None; + } + Some(lhs_leaf.subs.as_slice()) +} + fn record_field_expansion_error( lhs_name: &rumoca_core::Reference, field: &RecordFieldSpec, @@ -876,6 +1605,93 @@ fn record_field_expansion_error( ) } +fn record_field_specs_for_lhs( + lhs_ref: &rumoca_core::Reference, + flat: &flat::Model, + span: rumoca_core::Span, +) -> Result>, ToDaeError> { + let Some(lhs_component_ref) = lhs_ref.component_ref() else { + return Ok(None); + }; + + let mut fields = IndexMap::::new(); + for field_var in flat.variables.values() { + if !field_var.is_primitive { + continue; + } + let Some(field_ref) = field_var.component_ref.as_ref() else { + continue; + }; + if !component_ref_is_record_field_child(lhs_component_ref, field_ref, flat) { + continue; + } + let Some(field_leaf) = field_ref.parts.last() else { + continue; + }; + let Some(def_id) = field_ref.def_id else { + return Err(ToDaeError::runtime_contract_violation_at( + format!( + "record equation for `{}` has field `{}` without DefId", + lhs_ref.as_str(), + field_leaf.ident + ), + span, + )); + }; + fields.entry(field_leaf.ident.clone()).or_insert(def_id); + } + + let specs = fields + .into_iter() + .map(|(name, def_id)| RecordFieldSpec { + param: rumoca_core::FunctionParam::new(name, "Real", span).with_def_id(def_id), + match_by_name: true, + }) + .collect::>(); + Ok((!specs.is_empty()).then_some(specs)) +} + +fn lhs_record_reference( + lhs: &rumoca_core::Expression, +) -> Option<(rumoca_core::Reference, rumoca_core::Span, bool)> { + match lhs { + rumoca_core::Expression::VarRef { + name, + subscripts, + span, + } if subscripts.is_empty() => Some((name.clone(), *span, false)), + rumoca_core::Expression::Index { + base, + subscripts, + span, + } => { + let rumoca_core::Expression::VarRef { + name, + subscripts: base_subscripts, + .. + } = base.as_ref() + else { + return None; + }; + if !base_subscripts.is_empty() { + return None; + } + let mut component_ref = name.component_ref()?.clone(); + component_ref + .parts + .last_mut()? + .subs + .extend(subscripts.clone()); + Some(( + rumoca_core::Reference::from_component_reference(component_ref), + *span, + true, + )) + } + _ => None, + } +} + pub(crate) fn expand_record_field_equation( eq: &flat::Equation, flat: &flat::Model, @@ -889,36 +1705,40 @@ pub(crate) fn expand_record_field_equation( else { return Ok(None); }; - let rumoca_core::Expression::VarRef { - name: lhs_name, - subscripts, - span: lhs_span, - } = lhs.as_ref() + let Some((lhs_name, lhs_span, lhs_is_indexed_selection)) = lhs_record_reference(lhs.as_ref()) else { return Ok(None); }; - if !subscripts.is_empty() - || flat - .variables - .get(lhs_name.var_name()) - .is_some_and(|var| var.is_primitive) + if flat + .variables + .get(lhs_name.var_name()) + .is_some_and(|var| var.is_primitive) { return Ok(None); } - let Some(field_specs) = record_field_specs_for_rhs(rhs, flat) else { - return Ok(None); + let field_specs = match record_field_specs_for_rhs(rhs, flat) { + Some(field_specs) => field_specs, + None => match record_field_specs_for_lhs(&lhs_name, flat, lhs_span)? { + Some(field_specs) => field_specs, + None => return Ok(None), + }, }; let mut equations = Vec::new(); for (index, field) in field_specs.iter().enumerate() { - let field_vars = record_field_variables(lhs_name, field, flat, *lhs_span)?; + let field_vars = record_field_variables(&lhs_name, field, flat, lhs_span)?; if field_vars.is_empty() { - return Err(record_field_expansion_error(lhs_name, field, eq.span)); + if field.is_lhs_derived() || lhs_is_indexed_selection { + continue; + } + return Err(record_field_expansion_error(&lhs_name, field, eq.span)); } - let scalar_count = field_scalar_count(&field_vars); + let selection_subscripts = + indexed_record_field_selection_subscripts(&lhs_name, &field_vars); + let scalar_count = selected_field_scalar_count(&field_vars, selection_subscripts, flat); equations.push(flat::Equation::new_array( field_residual( - field_lhs_expression(&field_vars, eq.span), + field_lhs_expression(&field_vars, selection_subscripts, eq.span), rhs_field_expression(rhs, field, index, flat, eq.span), eq.span, ), @@ -937,6 +1757,7 @@ fn route_classified_equation( eq: &flat::Equation, dae_eq: dae::Equation, discrete_valued_lhs_counts: &HashMap, + discrete_valued_binding_targets: &HashSet, ) -> Result<(), ToDaeError> { let discrete_bucket = classify_residual_discrete_bucket(dae, &eq.residual); @@ -957,6 +1778,7 @@ fn route_classified_equation( &split_eq, split_dae_eq, discrete_valued_lhs_counts, + discrete_valued_binding_targets, )?; } return Ok(()); @@ -973,6 +1795,7 @@ fn route_classified_equation( &dae_eq, true, discrete_valued_lhs_counts, + discrete_valued_binding_targets, )? { dae.discrete.valued_updates.push(dae_eq); } @@ -995,6 +1818,7 @@ fn route_classified_equation( &dae_eq, true, discrete_valued_lhs_counts, + discrete_valued_binding_targets, )? { dae.discrete.valued_updates.push(dae_eq); } @@ -1370,6 +2194,22 @@ fn collect_explicit_discrete_assignments( dae: &dae::Dae, discrete_valued_lhs_counts: &HashMap, equation_span: rumoca_core::Span, +) -> Result>, ToDaeError> { + collect_explicit_discrete_assignments_with_binding_targets( + expr, + dae, + discrete_valued_lhs_counts, + equation_span, + &HashSet::new(), + ) +} + +fn collect_explicit_discrete_assignments_with_binding_targets( + expr: &rumoca_core::Expression, + dae: &dae::Dae, + discrete_valued_lhs_counts: &HashMap, + equation_span: rumoca_core::Span, + binding_targets: &HashSet, ) -> Result>, ToDaeError> { match expr { rumoca_core::Expression::Binary { @@ -1383,6 +2223,7 @@ fn collect_explicit_discrete_assignments( dae, discrete_valued_lhs_counts, equation_span, + binding_targets, ), rumoca_core::Expression::If { branches, @@ -1399,11 +2240,12 @@ fn collect_explicit_discrete_assignments( op: rumoca_core::OpUnary::Minus, rhs, .. - } => collect_explicit_discrete_assignments( + } => collect_explicit_discrete_assignments_with_binding_targets( rhs, dae, discrete_valued_lhs_counts, equation_span, + binding_targets, ), _ => Ok(None), } @@ -1415,6 +2257,7 @@ fn collect_binary_explicit_discrete_assignments( dae: &dae::Dae, discrete_valued_lhs_counts: &HashMap, equation_span: rumoca_core::Span, + binding_targets: &HashSet, ) -> Result>, ToDaeError> { if let Some(assignments) = collect_oriented_discrete_alias_assignment( lhs, @@ -1422,6 +2265,7 @@ fn collect_binary_explicit_discrete_assignments( dae, discrete_valued_lhs_counts, equation_span, + binding_targets, )? { return Ok(Some(assignments)); } @@ -1605,6 +2449,7 @@ fn collect_oriented_discrete_alias_assignment( dae: &dae::Dae, discrete_valued_lhs_counts: &HashMap, equation_span: rumoca_core::Span, + binding_targets: &HashSet, ) -> Result>, ToDaeError> { let Some(lhs_target) = explicit_assignment_target(lhs) else { return Ok(None); @@ -1624,6 +2469,17 @@ fn collect_oriented_discrete_alias_assignment( return Ok(None); } + if binding_targets.contains(&lhs_target.name) && !binding_targets.contains(&rhs_target.name) { + let mut result = HashMap::new(); + result.insert(rhs_target.name, lhs.clone()); + return Ok(Some(result)); + } + if binding_targets.contains(&rhs_target.name) && !binding_targets.contains(&lhs_target.name) { + let mut result = HashMap::new(); + result.insert(lhs_target.name, rhs.clone()); + return Ok(Some(result)); + } + let lhs_definitions = required_discrete_lhs_count( &lhs_target.name, discrete_valued_lhs_counts, @@ -1712,18 +2568,32 @@ fn push_explicit_discrete_assignments( equation: &dae::Equation, discrete_valued: bool, discrete_valued_lhs_counts: &HashMap, + discrete_valued_binding_targets: &HashSet, ) -> Result { let rhs = crate::dae_to_flat_expression(&equation.rhs); - let Some(assignments) = collect_explicit_discrete_assignments( + let Some(assignments) = collect_explicit_discrete_assignments_with_binding_targets( &rhs, dae, discrete_valued_lhs_counts, equation.span, + discrete_valued_binding_targets, )? else { return Ok(false); }; + // Aggregate array targets can be classified from their scalarized DAE + // lanes before the residual itself has been scalarized. Keep that residual + // in f_m here; DAE scalarization will recover one explicit assignment per + // concrete lane with the required component-reference metadata. + if equation.scalar_count > 1 + && assignments + .keys() + .any(|target| !is_discrete_valued_target(dae, target)) + { + return Ok(false); + } + let mut ordered: Vec<_> = assignments.into_iter().collect(); ordered.sort_unstable_by(|(lhs, _), (rhs, _)| lhs.as_str().cmp(rhs.as_str())); @@ -1753,13 +2623,14 @@ fn push_explicit_discrete_assignments( pub(super) fn classify_equations( dae: &mut dae::Dae, flat: &flat::Model, - prefix_counts: &FxHashMap, + prefix_counts: &super::ScalarInferenceMetadata, ) -> Result<(), ToDaeError> { let outputs_with_component_eqs = collect_vars_with_component_equations(flat); let non_connection_rhs_var_refs = collect_non_connection_rhs_var_refs(flat); let top_level_oc_connectors = collect_top_level_overconstrained_connectors(flat); let linearized_embedded_lhs_bases = super::collect_linearized_embedded_lhs_bases(flat); let discrete_valued_lhs_counts = collect_discrete_valued_lhs_target_counts(dae, flat); + let discrete_valued_binding_targets = collect_discrete_valued_binding_targets(dae, flat); let debug_eq_filter = crate::equation_filter_debug_enabled(); let mut stats = EqFilterStats::default(); let filter_ctx = EqFilterContext { @@ -1802,7 +2673,14 @@ pub(super) fn classify_equations( expanded_eq.origin.to_string(), scalar_count, ); - route_classified_equation(dae, flat, expanded_eq, dae_eq, &discrete_valued_lhs_counts)?; + route_classified_equation( + dae, + flat, + expanded_eq, + dae_eq, + &discrete_valued_lhs_counts, + &discrete_valued_binding_targets, + )?; } if dae.continuous.equations.len() == fx_index_before + 1 { flat_to_fx_index.insert(flat_idx, fx_index_before); @@ -1833,13 +2711,46 @@ fn collect_discrete_valued_lhs_target_counts( continue; } for target in crate::discrete_partition::residual_lhs_targets(&equation.residual) { - if is_discrete_valued_target(dae, &target) { - *counts.entry(target).or_insert(0) += 1; + let concrete_targets = discrete_lhs_count_targets(dae, equation, target); + for concrete_target in concrete_targets { + *counts.entry(concrete_target).or_insert(0) += 1; } } } counts } +fn discrete_lhs_count_targets( + dae: &dae::Dae, + equation: &flat::Equation, + target: rumoca_core::VarName, +) -> Vec { + if is_discrete_valued_target(dae, &target) { + return vec![target]; + } + let Some(reference) = + crate::discrete_partition::residual_target_component_reference(&equation.residual, &target) + else { + return Vec::new(); + }; + crate::discrete_partition::scalarized_discrete_targets_for_reference( + &dae.variables.discrete_valued, + &reference, + ) +} + +fn collect_discrete_valued_binding_targets( + dae: &dae::Dae, + flat: &flat::Model, +) -> HashSet { + flat.variables + .iter() + .filter(|(name, variable)| { + variable.binding.is_some() && is_discrete_valued_target(dae, name) + }) + .map(|(name, _)| name.clone()) + .collect() +} + #[cfg(test)] mod tests; diff --git a/crates/rumoca-phase-dae/src/equation_conversion/tests.rs b/crates/rumoca-phase-dae/src/equation_conversion/tests.rs index e15b32f0d..d43c52f47 100644 --- a/crates/rumoca-phase-dae/src/equation_conversion/tests.rs +++ b/crates/rumoca-phase-dae/src/equation_conversion/tests.rs @@ -6,9 +6,11 @@ use rumoca_ir_dae as dae; use rumoca_ir_flat as flat; use super::{ - EqFilterContext, classify_equations, collect_discrete_valued_lhs_target_counts, - collect_explicit_discrete_assignments, expand_record_field_equation, - explicit_lhs_reference_from_target, output_alias_skip_reason, output_has_component_equation, + EqFilterContext, classify_equations, collect_discrete_valued_binding_targets, + collect_discrete_valued_lhs_target_counts, collect_explicit_discrete_assignments, + collect_explicit_discrete_assignments_with_binding_targets, expand_record_field_equation, + explicit_lhs_reference_from_target, is_identity_equation, is_stream_stream_connection, + output_alias_skip_reason, output_has_component_equation, should_skip_stream_stream_connection, }; use crate::ToDaeError; @@ -42,6 +44,48 @@ fn call(name: &str) -> rumoca_core::Expression { } } +fn binary( + op: rumoca_core::OpBinary, + lhs: rumoca_core::Expression, + rhs: rumoca_core::Expression, +) -> rumoca_core::Expression { + rumoca_core::Expression::Binary { + op, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span: fixture_span(), + } +} + +#[test] +fn identity_equation_detects_same_variable_residual() { + let eq = flat::Equation { + residual: residual(var_ref("medium.h"), var_ref("medium.h")), + span: fixture_span(), + origin: flat::EquationOrigin::ComponentEquation { + component: "medium".to_string(), + }, + scalar_count: 1, + }; + + assert!(is_identity_equation(&eq)); +} + +#[test] +fn identity_equation_rejects_distinct_alias_residual() { + let eq = flat::Equation { + residual: residual(var_ref("port_a.p"), var_ref("port_b.p")), + span: fixture_span(), + origin: flat::EquationOrigin::Connection { + lhs: "port_a.p".to_string(), + rhs: "port_b.p".to_string(), + }, + scalar_count: 1, + }; + + assert!(!is_identity_equation(&eq)); +} + fn component_ref_with_def_id( parts: Vec<(&str, Vec)>, def_id: Option, @@ -80,6 +124,13 @@ fn var_ref_with_parts(name: &str, parts: Vec<(&str, Vec)>) -> rumoca_core:: } } +fn integer_literal(value: i64) -> rumoca_core::Expression { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + span: fixture_span(), + } +} + fn primitive_variable_with_dims_and_parts( name: &str, dims: Vec, @@ -132,6 +183,96 @@ fn add_pair_record_function(flat_model: &mut flat::Model, name: &str) { flat_model.functions.insert(function.name.clone(), function); } +fn add_complex_constructor(flat_model: &mut flat::Model) { + let mut constructor = rumoca_core::Function::new("Complex", fixture_span()); + constructor.is_constructor = true; + constructor.add_input( + rumoca_core::FunctionParam::new("re", "Real", fixture_span()) + .with_def_id(rumoca_core::DefId::new(201)), + ); + constructor.add_input( + rumoca_core::FunctionParam::new("im", "Real", fixture_span()) + .with_def_id(rumoca_core::DefId::new(202)), + ); + flat_model + .functions + .insert(constructor.name.clone(), constructor); +} + +fn add_converter_symmetrical_component_fields(flat_model: &mut flat::Model) { + for (name, parts, def_id) in [ + ( + "converter.iSymmetricalComponent[2].re", + vec![ + ("converter", vec![]), + ("iSymmetricalComponent", vec![2]), + ("re", vec![]), + ], + rumoca_core::DefId::new(201), + ), + ( + "converter.iSymmetricalComponent[2].im", + vec![ + ("converter", vec![]), + ("iSymmetricalComponent", vec![2]), + ("im", vec![]), + ], + rumoca_core::DefId::new(202), + ), + ( + "converter.iSymmetricalComponent[3].re", + vec![ + ("converter", vec![]), + ("iSymmetricalComponent", vec![3]), + ("re", vec![]), + ], + rumoca_core::DefId::new(201), + ), + ( + "converter.iSymmetricalComponent[3].im", + vec![ + ("converter", vec![]), + ("iSymmetricalComponent", vec![3]), + ("im", vec![]), + ], + rumoca_core::DefId::new(202), + ), + ] { + let var = primitive_variable_with_parts(name, parts, def_id); + flat_model.variables.insert(var.name.clone(), var); + } +} + +fn add_complex_scalar_fields(flat_model: &mut flat::Model, base: &str) { + for (field, def_id) in [ + ("re", rumoca_core::DefId::new(201)), + ("im", rumoca_core::DefId::new(202)), + ] { + let name = format!("{base}.{field}"); + let var = + primitive_variable_with_parts(&name, vec![(base, vec![]), (field, vec![])], def_id); + flat_model.variables.insert(var.name.clone(), var); + } +} + +fn add_index_non_pos_parameter(flat_model: &mut flat::Model) { + flat_model.variables.insert( + rumoca_core::VarName::new("converter.indexNonPos"), + flat::Variable { + name: rumoca_core::VarName::new("converter.indexNonPos"), + dims: vec![2], + variability: rumoca_core::Variability::Parameter(Default::default()), + binding: Some(rumoca_core::Expression::Array { + elements: vec![integer_literal(2), integer_literal(3)], + is_matrix: false, + span: fixture_span(), + }), + is_primitive: true, + ..flat::Variable::empty_with_span(fixture_span()) + }, + ); +} + #[test] fn test_record_function_equation_expands_to_declared_fields() { let mut flat_model = flat::Model::new(); @@ -312,6 +453,137 @@ fn test_record_function_equation_matches_subscripts_semantically() { assert!(format!("{:?}", expanded[0].residual).contains("controller.y[1].alpha")); } +#[test] +fn test_record_field_equation_expands_parameter_array_selected_lhs() { + let mut flat_model = flat::Model::new(); + add_converter_symmetrical_component_fields(&mut flat_model); + add_index_non_pos_parameter(&mut flat_model); + add_complex_constructor(&mut flat_model); + + let selected_lhs = rumoca_core::Expression::Index { + base: Box::new(var_ref_with_parts( + "converter.iSymmetricalComponent", + vec![("converter", vec![]), ("iSymmetricalComponent", vec![])], + )), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new("converter.indexNonPos").into(), + subscripts: vec![rumoca_core::Subscript::generated_index(1, fixture_span())], + span: fixture_span(), + }), + span: fixture_span(), + }], + span: fixture_span(), + }; + let equation = flat::Equation::new( + residual( + selected_lhs, + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("Complex").into(), + args: vec![integer_literal(0), integer_literal(0)], + is_constructor: true, + span: fixture_span(), + }, + ), + fixture_span(), + flat::EquationOrigin::ComponentEquation { + component: "converter".to_string(), + }, + ); + + let expanded = expand_record_field_equation(&equation, &flat_model) + .unwrap() + .expect("parameter-selected record equation should expand"); + assert_eq!(expanded.len(), 2); + assert!(format!("{:?}", expanded[0].residual).contains("iSymmetricalComponent[2].re")); + assert!(format!("{:?}", expanded[1].residual).contains("iSymmetricalComponent[2].im")); +} + +#[test] +fn test_record_field_equation_projects_complex_expression_fields() { + let mut flat_model = flat::Model::new(); + add_complex_constructor(&mut flat_model); + add_complex_scalar_fields(&mut flat_model, "out"); + add_complex_scalar_fields(&mut flat_model, "u"); + flat_model.variables.insert( + rumoca_core::VarName::new("scale"), + flat::Variable { + name: rumoca_core::VarName::new("scale"), + is_primitive: true, + ..flat::Variable::empty_with_span(fixture_span()) + }, + ); + + let complex_bias = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("Complex").into(), + args: vec![integer_literal(1), integer_literal(2)], + is_constructor: true, + span: fixture_span(), + }; + let rhs = binary( + rumoca_core::OpBinary::Add, + binary(rumoca_core::OpBinary::Mul, var_ref("scale"), var_ref("u")), + complex_bias, + ); + let equation = flat::Equation::new( + residual(var_ref_with_parts("out", vec![("out", vec![])]), rhs), + fixture_span(), + flat::EquationOrigin::ComponentEquation { + component: "out".to_string(), + }, + ); + + let expanded = expand_record_field_equation(&equation, &flat_model) + .unwrap() + .expect("complex record equation should expand into real and imaginary fields"); + + assert_eq!(expanded.len(), 2); + let re_residual = format!("{:?}", expanded[0].residual); + let im_residual = format!("{:?}", expanded[1].residual); + assert!(re_residual.contains("out.re")); + assert!(re_residual.contains("u.re")); + assert!(!re_residual.contains("FieldAccess")); + assert!(im_residual.contains("out.im")); + assert!(im_residual.contains("u.im")); + assert!(!im_residual.contains("FieldAccess")); +} + +#[test] +fn test_record_field_equation_projects_complex_division_fields() { + let mut flat_model = flat::Model::new(); + add_complex_constructor(&mut flat_model); + for base in ["out", "u", "v"] { + add_complex_scalar_fields(&mut flat_model, base); + } + + let equation = flat::Equation::new( + residual( + var_ref_with_parts("out", vec![("out", vec![])]), + binary(rumoca_core::OpBinary::Div, var_ref("u"), var_ref("v")), + ), + fixture_span(), + flat::EquationOrigin::ComponentEquation { + component: "out".to_string(), + }, + ); + + let expanded = expand_record_field_equation(&equation, &flat_model) + .unwrap() + .expect("complex division equation should expand into real and imaginary fields"); + + assert_eq!(expanded.len(), 2); + let re_residual = format!("{:?}", expanded[0].residual); + let im_residual = format!("{:?}", expanded[1].residual); + assert!(re_residual.contains("u.re")); + assert!(re_residual.contains("v.re")); + assert!(re_residual.contains("v.im")); + assert!(!re_residual.contains("FieldAccess")); + assert!(im_residual.contains("u.im")); + assert!(im_residual.contains("v.re")); + assert!(im_residual.contains("v.im")); + assert!(!im_residual.contains("FieldAccess")); +} + #[test] fn test_record_function_equation_matches_record_array_index_on_field_leaf() { let mut flat_model = flat::Model::new(); @@ -373,6 +645,58 @@ fn test_record_function_equation_matches_record_array_index_on_field_leaf() { assert!(!second.contains("controller.y.beta[2]")); } +#[test] +fn test_record_function_equation_matches_record_array_index_on_field_dims() { + let mut flat_model = flat::Model::new(); + for (name, field, def_id) in [ + ("controller.y.alpha", "alpha", rumoca_core::DefId::new(101)), + ("controller.y.beta", "beta", rumoca_core::DefId::new(102)), + ] { + let var = primitive_variable_with_dims_and_parts( + name, + vec![2], + vec![("controller", vec![]), ("y", vec![]), (field, vec![])], + def_id, + ); + flat_model.variables.insert(var.name.clone(), var); + } + add_pair_constructor(&mut flat_model); + add_pair_record_function(&mut flat_model, "Records.makePair"); + + let equation = flat::Equation::new_array( + residual( + rumoca_core::Expression::Index { + base: Box::new(var_ref_with_parts( + "controller.y", + vec![("controller", vec![]), ("y", vec![])], + )), + subscripts: vec![rumoca_core::Subscript::generated_index(1, fixture_span())], + span: fixture_span(), + }, + call("Records.makePair"), + ), + fixture_span(), + flat::EquationOrigin::ComponentEquation { + component: "controller".to_string(), + }, + 1, + ); + + let expanded = expand_record_field_equation(&equation, &flat_model) + .unwrap() + .expect("record array element equation should match field variables with array dims"); + assert_eq!(expanded.len(), 2); + assert_eq!(expanded[0].scalar_count, 1); + assert_eq!(expanded[1].scalar_count, 1); + + let first = format!("{:?}", expanded[0].residual); + assert!(first.contains("controller.y.alpha")); + assert!(first.contains("Index")); + let second = format!("{:?}", expanded[1].residual); + assert!(second.contains("controller.y.beta")); + assert!(second.contains("Index")); +} + #[test] fn test_record_function_equation_matches_field_def_id_not_spelling() { let mut flat_model = flat::Model::new(); @@ -532,6 +856,152 @@ fn test_discrete_alias_assignment_orients_to_unowned_rhs() { assert!(!assignments.contains_key(&rumoca_core::VarName::new("phase"))); } +#[test] +fn test_scalarized_boolean_array_definitions_count_each_concrete_lane() { + let declaration = rumoca_core::DefId::new(301); + let mut dae_model = dae::Dae::new(); + for (name, parts) in [ + ( + "parallel.split[1].set", + vec![("parallel", vec![]), ("split", vec![1]), ("set", vec![])], + ), + ( + "parallel.split[2].set", + vec![("parallel", vec![]), ("split", vec![2]), ("set", vec![])], + ), + ( + "step.inPort[1].set", + vec![("step", vec![]), ("inPort", vec![1]), ("set", vec![])], + ), + ("trigger", vec![("trigger", vec![])]), + ] { + let name = rumoca_core::VarName::new(name); + let mut variable = dae::Variable::new(name.clone(), fixture_span()); + variable.component_ref = Some(component_ref_with_def_id(parts, Some(declaration))); + dae_model.variables.discrete_valued.insert(name, variable); + } + + let mut flat_model = flat::Model::new(); + // Scalarized Flat models need not retain an aggregate declaration map + // entry; the source equation's structured LHS remains authoritative. + flat_model.equations.push(flat::Equation::new( + residual( + var_ref_with_parts( + "parallel.split.set", + vec![("parallel", vec![]), ("split", vec![]), ("set", vec![])], + ), + var_ref_with_parts("trigger", vec![("trigger", vec![])]), + ), + fixture_span(), + flat::EquationOrigin::ComponentEquation { + component: "parallel".to_string(), + }, + )); + flat_model.equations.push(flat::Equation::new( + residual( + var_ref_with_parts( + "parallel.split[2].set", + vec![("parallel", vec![]), ("split", vec![2]), ("set", vec![])], + ), + var_ref_with_parts( + "step.inPort[1].set", + vec![("step", vec![]), ("inPort", vec![1]), ("set", vec![])], + ), + ), + fixture_span(), + flat::EquationOrigin::Connection { + lhs: "parallel.split[2].set".to_string(), + rhs: "step.inPort[1].set".to_string(), + }, + )); + + let counts = collect_discrete_valued_lhs_target_counts(&dae_model, &flat_model); + assert_eq!( + counts[&rumoca_core::VarName::new("parallel.split[1].set")], + 1 + ); + assert_eq!( + counts[&rumoca_core::VarName::new("parallel.split[2].set")], + 2 + ); + + let assignments = collect_explicit_discrete_assignments( + &flat_model.equations[1].residual, + &dae_model, + &counts, + flat_model.equations[1].span, + ) + .unwrap() + .expect("scalarized alias assignment"); + assert!(assignments.contains_key(&rumoca_core::VarName::new("step.inPort[1].set"))); + assert!(!assignments.contains_key(&rumoca_core::VarName::new("parallel.split[2].set"))); +} + +#[test] +fn test_discrete_alias_assignment_orients_away_from_binding_owned_target() { + let mut dae_model = dae::Dae::new(); + for name in [ + "stateGraphRoot.suspend", + "stateGraphRoot.subgraphStatePort.suspend", + ] { + dae_model.variables.discrete_valued.insert( + rumoca_core::VarName::new(name), + dae::Variable::new(rumoca_core::VarName::new(name), fixture_span()), + ); + } + + let mut flat_model = flat::Model::new(); + flat_model.add_variable( + rumoca_core::VarName::new("stateGraphRoot.suspend"), + flat::Variable { + name: rumoca_core::VarName::new("stateGraphRoot.suspend"), + is_primitive: true, + is_discrete_type: true, + binding: Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Boolean(false), + span: fixture_span(), + }), + ..flat::Variable::empty_with_span(fixture_span()) + }, + ); + flat_model.add_variable( + rumoca_core::VarName::new("stateGraphRoot.subgraphStatePort.suspend"), + flat::Variable { + name: rumoca_core::VarName::new("stateGraphRoot.subgraphStatePort.suspend"), + is_primitive: true, + is_discrete_type: true, + ..flat::Variable::empty_with_span(fixture_span()) + }, + ); + flat_model.equations.push(flat::Equation::new( + residual( + var_ref("stateGraphRoot.suspend"), + var_ref("stateGraphRoot.subgraphStatePort.suspend"), + ), + fixture_span(), + flat::EquationOrigin::ComponentEquation { + component: "stateGraphRoot".to_string(), + }, + )); + + let counts = collect_discrete_valued_lhs_target_counts(&dae_model, &flat_model); + let binding_targets = collect_discrete_valued_binding_targets(&dae_model, &flat_model); + let assignments = collect_explicit_discrete_assignments_with_binding_targets( + &flat_model.equations[0].residual, + &dae_model, + &counts, + flat_model.equations[0].span, + &binding_targets, + ) + .unwrap() + .expect("binding-owned alias assignment"); + + assert!(assignments.contains_key(&rumoca_core::VarName::new( + "stateGraphRoot.subgraphStatePort.suspend" + ))); + assert!(!assignments.contains_key(&rumoca_core::VarName::new("stateGraphRoot.suspend"))); +} + #[test] fn test_discrete_alias_assignment_preserves_indexed_lhs_target() { let mut dae_model = dae::Dae::new(); @@ -854,6 +1324,149 @@ fn test_output_alias_skip_applies_when_both_sides_are_component_defined() { ); } +#[test] +fn test_stream_stream_connection_is_not_continuous_dae_residual() { + let mut flat_model = flat::Model::new(); + for name in ["pipe.port_b.h_outflow", "sink.port.h_outflow"] { + let var_name = rumoca_core::VarName::new(name); + flat_model.add_variable( + var_name.clone(), + flat::Variable { + name: var_name, + is_primitive: true, + stream: true, + ..rumoca_ir_flat::Variable::empty_with_span(fixture_span()) + }, + ); + } + flat_model.add_equation(flat::Equation::new( + residual( + var_ref("pipe.port_b.h_outflow"), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(1.0), + span: fixture_span(), + }, + ), + fixture_span(), + flat::EquationOrigin::ComponentEquation { + component: "pipe".to_string(), + }, + )); + flat_model.add_equation(flat::Equation::new( + residual( + var_ref("sink.port.h_outflow"), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(2.0), + span: fixture_span(), + }, + ), + fixture_span(), + flat::EquationOrigin::ComponentEquation { + component: "sink".to_string(), + }, + )); + let connection_eq = flat::Equation::new( + residual( + var_ref("pipe.port_b.h_outflow"), + var_ref("sink.port.h_outflow"), + ), + fixture_span(), + flat::EquationOrigin::Connection { + lhs: "pipe.port_b.h_outflow".to_string(), + rhs: "sink.port.h_outflow".to_string(), + }, + ); + assert!(is_stream_stream_connection(&connection_eq, &flat_model)); + flat_model.add_equation(connection_eq); + + let mut dae_model = dae::Dae::new(); + for name in ["pipe.port_b.h_outflow", "sink.port.h_outflow"] { + let var_name = rumoca_core::VarName::new(name); + dae_model.variables.algebraics.insert( + var_name.clone(), + dae::Variable::new(var_name, fixture_span()), + ); + } + + let scalar_metadata = crate::build_prefix_counts(&flat_model); + classify_equations(&mut dae_model, &flat_model, &scalar_metadata).unwrap(); + + assert_eq!( + dae_model.continuous.equations.len(), + 2, + "stream aliases support stream operators and must not overconstrain f_x" + ); + assert!(dae_model.continuous.equations.iter().all(|equation| { + !equation + .origin + .contains("pipe.port_b.h_outflow = sink.port.h_outflow") + })); +} + +#[test] +fn test_stream_stream_connection_is_kept_when_both_sides_are_consumed() { + let mut flat_model = flat::Model::new(); + for name in [ + "tank.ports[1].h_outflow", + "radiator.port_b.h_outflow", + "y1", + "y2", + ] { + let var_name = rumoca_core::VarName::new(name); + flat_model.add_variable( + var_name.clone(), + flat::Variable { + name: var_name, + is_primitive: true, + stream: name.ends_with("h_outflow"), + ..rumoca_ir_flat::Variable::empty_with_span(fixture_span()) + }, + ); + } + flat_model.add_equation(flat::Equation::new( + residual(var_ref("y1"), var_ref("tank.ports[1].h_outflow")), + fixture_span(), + flat::EquationOrigin::ComponentEquation { + component: "tank".to_string(), + }, + )); + flat_model.add_equation(flat::Equation::new( + residual(var_ref("y2"), var_ref("radiator.port_b.h_outflow")), + fixture_span(), + flat::EquationOrigin::ComponentEquation { + component: "radiator".to_string(), + }, + )); + let connection_eq = flat::Equation::new( + residual( + var_ref("tank.ports[1].h_outflow"), + var_ref("radiator.port_b.h_outflow"), + ), + fixture_span(), + flat::EquationOrigin::Connection { + lhs: "tank.ports[1].h_outflow".to_string(), + rhs: "radiator.port_b.h_outflow".to_string(), + }, + ); + + let outputs_with_component_eqs = HashSet::default(); + let non_connection_rhs_var_refs = super::collect_non_connection_rhs_var_refs(&flat_model); + let top_level_oc_connectors: IndexSet = IndexSet::new(); + let ctx = EqFilterContext { + flat: &flat_model, + outputs_with_component_eqs: &outputs_with_component_eqs, + non_connection_rhs_var_refs: &non_connection_rhs_var_refs, + top_level_oc_connectors: &top_level_oc_connectors, + debug_eq_filter: false, + }; + + assert!(is_stream_stream_connection(&connection_eq, &flat_model)); + assert!( + !should_skip_stream_stream_connection(&connection_eq, &ctx), + "stream aliases consumed on both sides carry a structural constraint" + ); +} + #[test] fn test_classify_equations_preserves_repeated_residuals_for_validation() { let residual = rumoca_core::Expression::Binary { @@ -894,12 +1507,8 @@ fn test_classify_equations_preserves_repeated_residuals_for_validation() { } let mut dae_model = dae::Dae::new(); - classify_equations( - &mut dae_model, - &flat_model, - &rustc_hash::FxHashMap::default(), - ) - .unwrap(); + let scalar_metadata = crate::build_prefix_counts(&flat_model); + classify_equations(&mut dae_model, &flat_model, &scalar_metadata).unwrap(); assert_eq!( dae_model.continuous.equations.len(), diff --git a/crates/rumoca-phase-dae/src/fold_start_values.rs b/crates/rumoca-phase-dae/src/fold_start_values.rs index a71497866..31f9001cd 100644 --- a/crates/rumoca-phase-dae/src/fold_start_values.rs +++ b/crates/rumoca-phase-dae/src/fold_start_values.rs @@ -7,10 +7,12 @@ //! with literal values. use crate::errors::ToDaeError; -use rumoca_core::ExpressionVisitor; -use rumoca_core::{Expression, Span, VarName}; -use rumoca_eval_dae::constant::{ConstValue, eval_const_expr_with}; -use rumoca_ir_dae::{Dae, DaeVariableMutVisitor, DaeVariablePartition, DaeVisitor, Variable}; +use rumoca_core::{Expression, ExpressionRewriter, ExpressionVisitor, Literal, Reference, Span}; +use rumoca_core::{StatementRewriter, Subscript, VarName}; +use rumoca_eval_dae::constant::{ConstValue, eval_const_expr_with_shape}; +use rumoca_ir_dae::{ + Dae, DaeVariableMutVisitor, DaeVariablePartition, DaeVariables, DaeVisitor, Variable, +}; use std::collections::HashMap; /// Evaluate all parameter/state/constant start expressions to typed literals @@ -18,7 +20,43 @@ use std::collections::HashMap; pub(crate) fn fold_start_values_to_literals(dae: &mut Dae) -> Result<(), ToDaeError> { // Phase 1: build a name→value map from constants, enum ordinals, and // parameter start expressions (fixed-point iteration). - let values = collect_foldable_start_values(dae); + let mut values: HashMap = HashMap::new(); + let dims = rumoca_eval_dae::collect_var_dims(dae); + + // Seed with enum literal ordinals + for (name, ordinal) in &dae.symbols.enum_literal_ordinals { + values.insert(name.clone(), ConstValue::Real(*ordinal as f64)); + } + seed_modelica_standard_constants(&mut values); + + // Collect all named start bindings (constants, parameters, inputs, states, + // discrete reals, discrete valued, algebraics, outputs) + let mut bindings = Vec::new(); + StartBindingCollector { + bindings: &mut bindings, + } + .visit_dae(dae); + let declared_start_names = declared_start_names(&bindings, dae); + + // Fixed-point iteration: resolve chains like A = B, B = 3.14 + let max_passes = bindings.len().max(1) * 2; + for _ in 0..max_passes { + let mut changed = false; + for (name, expr) in &bindings { + if values.contains_key(name.as_str()) { + continue; + } + if let Some(value) = eval_start_const_expr(expr, &values, &dims) + && value.is_finite() + { + values.insert(name.to_string(), value); + changed = true; + } + } + if !changed { + break; + } + } // Set of parameter names: a parameter start expression that references // another parameter stays symbolic (see below), matching the bare-alias @@ -51,7 +89,7 @@ pub(crate) fn fold_start_values_to_literals(dae: &mut Dae) -> Result<(), ToDaeEr name, subscripts, .. } = start && subscripts.is_empty() - && name.as_str() == var.name.as_str() + && is_self_start_reference(name, var) { var.start = None; return; @@ -65,8 +103,11 @@ pub(crate) fn fold_start_values_to_literals(dae: &mut Dae) -> Result<(), ToDaeEr // flow through to dependents instead of being locked at // compile time. Topo-sorting (sort_parameters_by_start_deps) // already orders the chain for forward-eval templates. - if let Expression::VarRef { subscripts, .. } = start + if let Expression::VarRef { + name, subscripts, .. + } = start && subscripts.is_empty() + && is_declared_start_reference(name, &declared_start_names) { return; } @@ -83,19 +124,22 @@ pub(crate) fn fold_start_values_to_literals(dae: &mut Dae) -> Result<(), ToDaeEr if is_parameter && start_references_parameter(start, ¶m_names) && !expr_depends_on_string_value(start, &values) + && !expr_is_shape_only(start) { return; } let Some(val) = values.get(var.name.as_str()).cloned() else { return; }; - let mut refs: Vec = Vec::new(); - start.collect_var_refs(&mut refs); - consumed_params.extend( - refs.iter() - .map(|name| name.as_str().to_string()) - .filter(|name| param_names.contains(name)), - ); + if !expr_is_shape_only(start) { + let mut refs: Vec = Vec::new(); + start.collect_var_refs(&mut refs); + consumed_params.extend( + refs.iter() + .map(|name| name.as_str().to_string()) + .filter(|name| param_names.contains(name)), + ); + } let span = match folded_start_span(var, start) { Ok(span) => span, Err(err) => { @@ -123,38 +167,140 @@ pub(crate) fn fold_start_values_to_literals(dae: &mut Dae) -> Result<(), ToDaeEr Ok(()) } -fn collect_foldable_start_values(dae: &Dae) -> HashMap { - let mut values: HashMap = HashMap::new(); +/// Replace Modelica Standard Library package constants with literals anywhere +/// they survive lowering as references. These constants are translation-time +/// package declarations, not runtime variables, so unresolved VarRefs to them +/// should not reach DAE reference validation. +pub(crate) fn fold_known_package_constants_to_literals(dae: &mut Dae) { + let mut rewriter = KnownPackageConstantRewriter; + rewrite_variable_attributes(&mut dae.variables, &mut rewriter); + for equation in dae + .continuous + .equations + .iter_mut() + .chain(dae.discrete.real_updates.iter_mut()) + .chain(dae.discrete.valued_updates.iter_mut()) + .chain(dae.conditions.equations.iter_mut()) + .chain(dae.initialization.equations.iter_mut()) + { + equation.rhs = rewriter.rewrite_expression(&equation.rhs); + } + rewrite_expressions(&mut dae.conditions.relations, &mut rewriter); + rewrite_expressions(&mut dae.events.synthetic_root_conditions, &mut rewriter); + for action in &mut dae.events.event_actions { + action.condition = rewriter.rewrite_expression(&action.condition); + let (rumoca_ir_dae::DaeEventActionKind::Assert { message } + | rumoca_ir_dae::DaeEventActionKind::Terminate { message }) = &mut action.kind; + *message = rewriter.rewrite_expression(message); + } + rewrite_expressions(&mut dae.clocks.constructor_exprs, &mut rewriter); + rewrite_expressions(&mut dae.clocks.triggered_conditions, &mut rewriter); + for function in dae.symbols.functions.values_mut() { + for param in function + .inputs + .iter_mut() + .chain(function.outputs.iter_mut()) + .chain(function.locals.iter_mut()) + { + if let Some(default) = &mut param.default { + *default = rewriter.rewrite_expression(default); + } + } + function.body = rewriter.rewrite_statements(&function.body); + } +} - for (name, ordinal) in &dae.symbols.enum_literal_ordinals { - values.insert(name.clone(), ConstValue::Real(*ordinal as f64)); +fn rewrite_variable_attributes( + variables: &mut DaeVariables, + rewriter: &mut KnownPackageConstantRewriter, +) { + for variable in variables + .states + .values_mut() + .chain(variables.algebraics.values_mut()) + .chain(variables.inputs.values_mut()) + .chain(variables.outputs.values_mut()) + .chain(variables.parameters.values_mut()) + .chain(variables.constants.values_mut()) + .chain(variables.discrete_reals.values_mut()) + .chain(variables.discrete_valued.values_mut()) + { + rewrite_optional_expression(&mut variable.start, rewriter); + rewrite_optional_expression(&mut variable.min, rewriter); + rewrite_optional_expression(&mut variable.max, rewriter); + rewrite_optional_expression(&mut variable.nominal, rewriter); } +} - let mut bindings = Vec::new(); - StartBindingCollector { - bindings: &mut bindings, +fn rewrite_expressions( + expressions: &mut [Expression], + rewriter: &mut KnownPackageConstantRewriter, +) { + for expression in expressions { + *expression = rewriter.rewrite_expression(expression); } - .visit_dae(dae); +} - let max_passes = bindings.len().max(1) * 2; - for _ in 0..max_passes { - let mut changed = false; - for (name, expr) in &bindings { - if values.contains_key(name.as_str()) { - continue; - } - if let Some(value) = eval_start_const_expr(expr, &values) - && value.is_finite() - { - values.insert(name.to_string(), value); - changed = true; - } - } - if !changed { - break; +fn rewrite_optional_expression( + expression: &mut Option, + rewriter: &mut KnownPackageConstantRewriter, +) { + if let Some(expression) = expression { + *expression = rewriter.rewrite_expression(expression); + } +} + +struct KnownPackageConstantRewriter; + +impl ExpressionRewriter for KnownPackageConstantRewriter { + fn rewrite_var_ref_expression( + &mut self, + name: &Reference, + subscripts: &[Subscript], + span: Span, + ) -> Expression { + if subscripts.is_empty() + && let Some(value) = modelica_standard_constant_value(name.as_str()) + { + return Expression::Literal { + value: Literal::Real(value), + span, + }; } + self.walk_var_ref_expression(name, subscripts, span) + } +} + +impl StatementRewriter for KnownPackageConstantRewriter {} + +fn seed_modelica_standard_constants(values: &mut HashMap) { + for (name, value) in MODELICA_STANDARD_CONSTANTS { + values + .entry(name.to_string()) + .or_insert(ConstValue::Real(*value)); } - values +} + +const MODELICA_STANDARD_CONSTANTS: &[(&str, f64)] = &[ + ("Modelica.Constants.pi", std::f64::consts::PI), + ("Modelica.Constants.e", std::f64::consts::E), + ("Modelica.Constants.small", 1.0e-60), + ("Modelica.Constants.eps", f64::EPSILON), +]; + +fn modelica_standard_constant_value(name: &str) -> Option { + MODELICA_STANDARD_CONSTANTS + .iter() + .find_map(|(candidate, value)| (*candidate == name).then_some(*value)) +} + +fn is_declared_start_reference( + name: &rumoca_core::Reference, + declared_start_names: &std::collections::HashSet, +) -> bool { + reference_lookup_names(name) + .iter() + .any(|candidate| declared_start_names.contains(candidate)) } fn folded_start_span(var: &Variable, start: &Expression) -> Result { @@ -212,10 +358,70 @@ fn start_references_parameter( refs.iter().any(|name| param_names.contains(name.as_str())) } +fn expr_is_shape_only(expr: &Expression) -> bool { + match expr { + Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + args, + .. + } => args.len() >= 2, + Expression::FunctionCall { name, args, .. } if name.last_segment() == "size" => { + args.len() >= 2 + } + Expression::Unary { rhs, .. } => expr_is_shape_only(rhs), + Expression::Binary { lhs, rhs, .. } => expr_is_shape_only(lhs) && expr_is_shape_only(rhs), + Expression::If { + branches, + else_branch, + .. + } => { + branches + .iter() + .all(|(cond, branch)| expr_is_shape_only(cond) && expr_is_shape_only(branch)) + && expr_is_shape_only(else_branch) + } + _ => false, + } +} + +fn is_self_start_reference(name: &rumoca_core::Reference, var: &Variable) -> bool { + match (name.component_ref(), var.component_ref.as_ref()) { + (Some(name_ref), Some(var_ref)) => component_refs_same_target(name_ref, var_ref), + (Some(_), None) => false, + _ => name.as_str() == var.name.as_str(), + } +} + +fn component_refs_same_target( + lhs: &rumoca_core::ComponentReference, + rhs: &rumoca_core::ComponentReference, +) -> bool { + if let (Some(lhs_def), Some(rhs_def)) = (lhs.def_id, rhs.def_id) { + return lhs_def == rhs_def; + } + lhs.parts.len() == rhs.parts.len() + && lhs + .parts + .iter() + .zip(&rhs.parts) + .all(|(lhs, rhs)| lhs.ident == rhs.ident && lhs.subs == rhs.subs) +} + struct StartBindingCollector<'a> { bindings: &'a mut Vec<(VarName, Expression)>, } +fn declared_start_names( + bindings: &[(VarName, Expression)], + dae: &Dae, +) -> std::collections::HashSet { + bindings + .iter() + .map(|(name, _)| name.as_str().to_string()) + .chain(dae.symbols.enum_literal_ordinals.keys().cloned()) + .collect() +} + impl DaeVisitor for StartBindingCollector<'_> { fn visit_variable( &mut self, @@ -257,16 +463,46 @@ where fn eval_start_const_expr( expr: &Expression, env: &HashMap, + dims: &indexmap::IndexMap>, ) -> Option { - eval_const_expr_with(expr, &|name, subscripts| { - if !subscripts.is_empty() { - return None; - } - env.get(name.as_str()).cloned().or_else(|| { - rumoca_ir_dae::component_base_name(name.as_str()) - .and_then(|base| env.get(&base).cloned()) - }) - }) + eval_const_expr_with_shape( + expr, + &|name, subscripts| { + if !subscripts.is_empty() { + return None; + } + for key in reference_lookup_names(name) { + if let Some(value) = env.get(key.as_str()).cloned().or_else(|| { + rumoca_ir_dae::component_base_name(key.as_str()) + .and_then(|base| env.get(&base).cloned()) + }) { + return Some(value); + } + } + None + }, + &|name| dims.get(name).cloned(), + ) +} + +fn reference_lookup_names(name: &rumoca_core::Reference) -> Vec { + let mut names = Vec::new(); + if let Some(component_ref) = name.component_ref() { + names.push(component_ref_flat_name(component_ref)); + } + names.push(name.as_str().to_string()); + names.sort(); + names.dedup(); + names +} + +fn component_ref_flat_name(component_ref: &rumoca_core::ComponentReference) -> String { + component_ref + .parts + .iter() + .map(|part| part.ident.as_str()) + .collect::>() + .join(".") } // --------------------------------------------------------------------------- @@ -355,17 +591,15 @@ pub(crate) fn sort_algebraics_by_equation_deps(dae: &mut Dae) -> Result<(), ToDa for eq in &dae.continuous.equations { let refs = collect_param_refs(&eq.rhs, &alg_names); // This equation may define one of our algebraic vars. - // Try to identify which variable this equation defines - // by checking if it matches the pattern `0 = var - expr` or additive form. - for alg_name in &alg_names { - if equation_defines_var(&eq.rhs, alg_name) { - let deps: Vec = refs - .iter() - .filter(|r| r.as_str() != alg_name.as_str()) - .cloned() - .collect(); - eq_deps.insert(alg_name.clone(), deps); - } + // Identify candidate definitions once per equation instead of scanning + // every algebraic/output name against the same expression tree. + for alg_name in collect_defined_var_refs(&eq.rhs, &alg_names) { + let deps: Vec = refs + .iter() + .filter(|r| r.as_str() != alg_name.as_str()) + .cloned() + .collect(); + eq_deps.insert(alg_name, deps); } } @@ -376,39 +610,20 @@ pub(crate) fn sort_algebraics_by_equation_deps(dae: &mut Dae) -> Result<(), ToDa Ok(()) } -/// Check if an equation's RHS defines a given variable (appears as LHS of subtraction -/// or as a term in an additive equation). -fn equation_defines_var(rhs: &Expression, var_name: &str) -> bool { - match rhs { - Expression::Binary { - op, - lhs, - rhs: rhs_inner, - .. - } => { - if matches!(op, rumoca_core::OpBinary::Sub) { - // 0 = var - expr or 0 = expr - var - if is_var_ref_named(lhs, var_name) || is_var_ref_named(rhs_inner, var_name) { - return true; - } - } - if matches!(op, rumoca_core::OpBinary::Add) { - // Check additive terms - let terms = collect_additive_var_refs(rhs); - if terms.iter().any(|t| t == var_name) { - return true; - } - } - false - } - Expression::Unary { - op: rumoca_core::OpUnary::Minus, - rhs: inner, - .. - } => equation_defines_var(inner, var_name), - Expression::VarRef { name, .. } => name.as_str() == var_name, - _ => false, - } +#[cfg(test)] +fn is_var_ref_named(expr: &Expression, name: &str) -> bool { + matches!(expr, Expression::VarRef { name: n, .. } if n.as_str() == name || var_base_name(n.as_str()) == name) +} + +fn collect_defined_var_refs( + rhs: &Expression, + alg_names: &std::collections::HashSet, +) -> Vec { + let mut seen = std::collections::HashSet::new(); + collect_additive_var_refs(rhs) + .into_iter() + .filter(|name| alg_names.contains(name) && seen.insert(name.clone())) + .collect() } /// Extract the base name from a possibly-subscripted variable name. @@ -417,10 +632,6 @@ fn var_base_name(name: &str) -> &str { rumoca_core::strip_scalar_name_subscripts(name).unwrap_or(name) } -fn is_var_ref_named(expr: &Expression, name: &str) -> bool { - matches!(expr, Expression::VarRef { name: n, .. } if n.as_str() == name || var_base_name(n.as_str()) == name) -} - /// Collect all VarRef names from an additive expression tree. fn collect_additive_var_refs(expr: &Expression) -> Vec { match expr { @@ -559,6 +770,21 @@ mod tests { } } + fn structured_var_ref(rendered: &str, target: &str, def_id: u32) -> Expression { + Expression::VarRef { + name: rumoca_core::Reference::with_component_reference( + rendered, + rumoca_core::ComponentReference::from_flat_segments( + target, + test_span(1, 2), + Some(rumoca_core::DefId::new(def_id)), + ), + ), + subscripts: vec![], + span: test_span(1, 2), + } + } + fn real(value: f64) -> Expression { Expression::Literal { value: Literal::Real(value), @@ -590,6 +816,16 @@ mod tests { var } + fn source_constant(name: &str, def_id: u32, start: Expression) -> Variable { + let mut var = parameter(name, start); + var.component_ref = Some(rumoca_core::ComponentReference::from_flat_segments( + name, + test_span(1, 2), + Some(rumoca_core::DefId::new(def_id)), + )); + var + } + #[test] fn var_base_name_strips_only_trailing_subscripts() { assert_eq!(var_base_name("e[1]"), "e"); @@ -600,6 +836,89 @@ mod tests { assert_eq!(var_base_name("record[index.re]"), "record[index.re]"); } + #[test] + fn fold_start_values_keeps_structured_alias_with_rewritten_display_name() { + let mut dae = Dae::new(); + dae.variables.constants.insert( + VarName::new("sineVoltage.pi"), + source_constant( + "sineVoltage.pi", + 39973, + structured_var_ref("sineVoltage.pi", "Modelica.Constants.pi", 86), + ), + ); + dae.variables.constants.insert( + VarName::new("Modelica.Constants.pi"), + source_constant("Modelica.Constants.pi", 86, real(std::f64::consts::PI)), + ); + + fold_start_values_to_literals(&mut dae).expect("constant alias starts should be valid"); + + let Some(Expression::VarRef { name, .. }) = + &dae.variables.constants[&VarName::new("sineVoltage.pi")].start + else { + panic!("structured alias start must not be treated as a self-reference"); + }; + assert_eq!(name.as_str(), "sineVoltage.pi"); + assert_eq!( + name.component_ref().map(component_ref_flat_name).as_deref(), + Some("Modelica.Constants.pi") + ); + } + + #[test] + fn fold_start_values_folds_unmaterialized_modelica_package_constants() { + let mut dae = Dae::new(); + dae.variables.constants.insert( + VarName::new("phaseShift"), + source_constant("phaseShift", 39974, var("Modelica.Constants.pi")), + ); + + fold_start_values_to_literals(&mut dae) + .expect("well-known package constants should fold before reference validation"); + + let Some(Expression::Literal { + value: Literal::Real(value), + .. + }) = &dae.variables.constants[&VarName::new("phaseShift")].start + else { + panic!("package constant start should fold to a real literal"); + }; + assert!((*value - std::f64::consts::PI).abs() < f64::EPSILON); + } + + #[test] + fn fold_known_package_constants_rewrites_equation_rhs() { + let mut dae = Dae::new(); + dae.continuous.equations.push(rumoca_ir_dae::Equation { + lhs: Some(VarName::new("y").into()), + rhs: Expression::Binary { + op: OpBinary::Mul, + lhs: Box::new(var("Modelica.Constants.pi")), + rhs: Box::new(var("gain")), + span: test_span(9, 10), + }, + span: test_span(9, 10), + origin: "package constant equation".to_string(), + scalar_count: 1, + }); + + fold_known_package_constants_to_literals(&mut dae); + + let Expression::Binary { lhs, rhs, .. } = &dae.continuous.equations[0].rhs else { + panic!("equation rhs should remain a product"); + }; + let Expression::Literal { + value: Literal::Real(value), + .. + } = lhs.as_ref() + else { + panic!("Modelica.Constants.pi should fold to a real literal"); + }; + assert!((*value - std::f64::consts::PI).abs() < f64::EPSILON); + assert!(is_var_ref_named(rhs, "gain")); + } + #[test] fn var_ref_dependency_match_uses_trailing_scalar_subscript_base() { assert!(is_var_ref_named(&var("e[1]"), "e")); @@ -772,6 +1091,38 @@ mod tests { ); } + #[test] + fn folds_parameter_start_size_from_dae_variable_dims() { + let mut dae = Dae::new(); + let mut table = parameter("table", real(0.0)); + table.dims = vec![3, 2]; + dae.variables + .parameters + .insert(VarName::new("table"), table); + dae.variables.parameters.insert( + VarName::new("nout"), + parameter( + "nout", + Expression::BuiltinCall { + function: BuiltinFunction::Size, + args: vec![var("table"), integer(1)], + span: test_span(30, 45), + }, + ), + ); + + fold_start_values_to_literals(&mut dae) + .unwrap_or_else(|err| panic!("start folding should succeed: {err}")); + + assert!(matches!( + dae.variables.parameters[&VarName::new("nout")].start, + Some(Expression::Literal { + value: Literal::Real(value), + .. + }) if (value - 3.0).abs() <= 1.0e-12 + )); + } + #[test] fn folded_start_literal_preserves_expression_span_without_attribute_span() { let mut dae = Dae::new(); diff --git a/crates/rumoca-phase-dae/src/initial.rs b/crates/rumoca-phase-dae/src/initial.rs index 5b0fd9848..2697f79b5 100644 --- a/crates/rumoca-phase-dae/src/initial.rs +++ b/crates/rumoca-phase-dae/src/initial.rs @@ -1,12 +1,13 @@ //! Initial-equation conversion for ToDAE. +use crate::{ + ScalarInferenceMetadata, ToDaeError, flat_to_dae_expression_with_refs, + remap_flat_structured_equations, +}; use indexmap::IndexMap; use rumoca_core::{Expression, Literal, ProvenanceSpan, Reference, Span, VarName}; use rumoca_ir_dae as dae; use rumoca_ir_flat as flat; -use rustc_hash::FxHashMap; - -use crate::{ToDaeError, flat_to_dae_expression_with_refs, remap_flat_structured_equations}; /// Determine scalar count for one initial equation. /// @@ -30,11 +31,11 @@ fn initial_equation_scalar_count( pub(crate) fn convert_initial_equations( dae: &mut dae::Dae, flat: &flat::Model, - prefix_counts: &FxHashMap, + prefix_counts: &ScalarInferenceMetadata, infer_scalar_count: F, ) -> Result<(), ToDaeError> where - F: Fn(&rumoca_core::Expression, &flat::Model, &FxHashMap) -> usize, + F: Fn(&rumoca_core::Expression, &flat::Model, &ScalarInferenceMetadata) -> usize, { let mut flat_to_dae_index: IndexMap = IndexMap::new(); @@ -58,6 +59,9 @@ where scalar_count, ); dae.initialization.equations.push(dae_eq); + dae.initialization + .equation_provenance + .push(dae::InitializationEquationProvenance::User); } if dae.initialization.equations.len() == dae_index_before + 1 { flat_to_dae_index.insert(flat_idx, dae_index_before); @@ -79,7 +83,14 @@ pub(crate) fn add_fixed_start_initial_equations(dae: &mut dae::Dae) -> Result<() collect_fixed_start_equations(&dae.variables.states, &mut equations)?; collect_fixed_start_equations(&dae.variables.algebraics, &mut equations)?; collect_fixed_start_equations(&dae.variables.outputs, &mut equations)?; + let generated = equations.len(); dae.initialization.equations.extend(equations); + dae.initialization + .equation_provenance + .extend(std::iter::repeat_n( + dae::InitializationEquationProvenance::FixedStart, + generated, + )); Ok(()) } diff --git a/crates/rumoca-phase-dae/src/lib.rs b/crates/rumoca-phase-dae/src/lib.rs index 27fc4d386..20997e39d 100644 --- a/crates/rumoca-phase-dae/src/lib.rs +++ b/crates/rumoca-phase-dae/src/lib.rs @@ -29,6 +29,7 @@ mod binding_conversion; mod condition_activation; mod condition_lowering; mod connector_input_analysis; +mod constructor_field_selection; mod convert; mod dae_lowering; mod equation_conversion; @@ -68,6 +69,7 @@ use dae_lowering::sort_parameters_by_start_dependency; use indexmap::{IndexMap, IndexSet}; use path_utils::subscript_fallback_chain; use reference_validation::{validate_dae_constructor_field_selections, validate_dae_references}; +use rumoca_core::ExpressionVisitor; #[cfg(test)] use rumoca_core::strip_subscript; use rumoca_core::timing::{maybe_elapsed_seconds, maybe_start_timer_if}; @@ -80,18 +82,21 @@ use rumoca_ir_dae::{Dae, Variable}; use rumoca_ir_flat as flat; use rumoca_ir_flat::Model; use runtime_precompute::populate_runtime_precompute; -use rustc_hash::FxHashMap; +use rustc_hash::{FxHashMap, FxHashSet}; use scalar_inference::*; use std::collections::{HashMap, HashSet}; use variable_analysis::{ InternalInputIndex, count_interface_flows, count_overconstrained_interface, filter_state_variables, find_connected_inputs, find_discrete_connected_internal_inputs, - find_equation_defined_inputs, find_when_only_vars, is_when_only_var, - validate_flat_function_calls, + find_equation_defined_inputs, find_overconstrained_derivative_alias_roots, find_when_only_vars, + is_when_only_var, validate_flat_function_calls, }; use when_conversion::convert_when_clause; -pub use balance::{BalanceError, balance, balance_detail, equations_unknowns, is_balanced}; +pub use balance::{ + BalanceError, InitialClosureBalanceDetail, balance, balance_detail, equations_unknowns, + initial_closure_balance_detail, is_balanced, is_balanced_for_admission, +}; pub use dae_lowering::{ CodegenDae, insert_array_size_args_dae, lower_record_function_params_dae, prepare_dae_for_codegen, prepare_dae_for_fmi_model_description, @@ -100,9 +105,9 @@ pub use dae_lowering::{ pub use errors::{ToDaeError, ToDaeResult}; // Re-export moved functions so sibling modules can still use `super::`. pub(crate) use variable_analysis::{ - collect_continuous_equation_lhs, find_connected_inputs_only_connected_to_inputs, - infer_record_subscript_size_from_prefix_chain, is_continuous_unknown, is_internal_input, - record_subscript_scalar_size, resolve_flat_function, + collect_continuous_equation_lhs, find_connected_input_binding_anchors, + find_connected_inputs_only_connected_to_inputs, infer_record_subscript_size_from_prefix_chain, + is_continuous_unknown, is_internal_input, record_subscript_scalar_size, resolve_flat_function, }; #[cfg(test)] @@ -245,11 +250,7 @@ pub fn to_dae_with_options( })?; } - // MLS §4.7: Propagate partial status and class type for balance checking - dae.metadata.is_partial = flat.is_partial; - dae.metadata.class_type = flat.class_type.clone(); - dae.metadata.model_description = flat.model_description.clone(); - dae.metadata.symbol_ancestry = flat.symbol_ancestry.clone(); + initialize_dae_metadata(&mut dae, flat); let classification_indexes = build_variable_classification_indexes(flat)?; let prefix_children = &classification_indexes.prefix_children; @@ -355,6 +356,17 @@ pub fn to_dae_with_options( run_todae_phase(todae_subphase_timing, "pre_lowering", || { pre_lowering::lower_pre_operator(&mut dae) })?; + run_todae_phase( + todae_subphase_timing, + "overconstrained_derivative_alias_rewrite", + || { + rewrite_overconstrained_derivative_alias_refs( + &mut dae, + flat, + &classification_indexes.overconstrained_derivative_alias_roots, + ) + }, + )?; dae.symbols.functions = flat_to_dae_function_map(&flat.functions); finalize_lowered_dae(&mut dae, flat, state_vars, todae_subphase_timing, options)?; @@ -372,6 +384,28 @@ pub fn fold_hidden_component_outputs_for_projection(dae: &mut dae::Dae) { inline_hidden_component_algebraics::inline_hidden_component_algebraics(dae); } +fn initialize_dae_metadata(dae: &mut dae::Dae, flat: &flat::Model) { + // MLS §4.7: Propagate partial status and class type for balance checking. + dae.metadata.is_partial = flat.is_partial; + dae.metadata.class_type = flat.class_type.clone(); + dae.metadata.model_description = flat.model_description.clone(); + dae.metadata.symbol_ancestry = flat.symbol_ancestry.clone(); + dae.metadata.nonnumeric_variable_names = nonnumeric_variable_names(flat); +} + +fn nonnumeric_variable_names(flat: &flat::Model) -> Vec { + flat.variable_type_names + .iter() + .filter(|(name, type_name)| { + rumoca_core::qualified_type_name_matches(type_name, "String") + || flat.variables.get(*name).is_some_and(|var| { + variable_analysis::is_external_constructor_handle(flat, name, var) + }) + }) + .map(|(name, _)| name.as_str().to_string()) + .collect() +} + fn finalize_lowered_dae( dae: &mut dae::Dae, flat: &flat::Model, @@ -395,12 +429,29 @@ fn finalize_lowered_dae( dae_lowering::scalarize_phantom_vector_equations(dae) })?; + run_todae_phase(todae_subphase_timing, "sync_structured_templates", || { + dae_lowering::sync_materialized_structured_equation_templates(dae) + })?; + + run_todae_phase( + todae_subphase_timing, + "repair_external_table_events", + || { + dae_lowering::repair_external_table_event_handles(dae); + Ok::<(), ToDaeError>(()) + }, + )?; + // Fold symbolic start-value expressions to literal constants where // possible. Safe as an always-on pass: start values are init-time // metadata, not user-observable at runtime. run_todae_phase(todae_subphase_timing, "fold_start_values", || { fold_start_values::fold_start_values_to_literals(dae) })?; + run_todae_phase(todae_subphase_timing, "fold_package_constants", || { + fold_start_values::fold_known_package_constants_to_literals(dae); + Ok::<(), ToDaeError>(()) + })?; // Reorder algebraics so any algebraic used in another's defining // equation appears first. Pure reorder, no information loss; lets @@ -421,6 +472,9 @@ fn finalize_lowered_dae( run_todae_phase(todae_subphase_timing, "reference_metadata", || { attach_dae_reference_metadata(dae) })?; + run_todae_phase(todae_subphase_timing, "discrete_input_metadata", || { + refresh_external_discrete_input_metadata(dae, flat); + }); run_todae_phase(todae_subphase_timing, "appendix_b_validation", || { appendix_b_validation::validate_appendix_b_invariants(dae) })?; @@ -428,8 +482,11 @@ fn finalize_lowered_dae( run_todae_phase(todae_subphase_timing, "metadata_counts", || { // MLS §4.7 / §4.8 / §9.4: propagate interface counts from flatten. dae.metadata.interface_flow_count = count_interface_flows(flat); + dae.metadata.stream_interface_equation_count = flat.stream_interface_equation_count; dae.metadata.oc_break_edge_scalar_count = flat.oc_break_edge_scalar_count; overconstrained_interface::validate_connection_graph(flat)?; + dae.metadata.overconstrained_root_gauge_count = + overconstrained_interface::count_overconstrained_root_gauge(flat, state_vars)?; let oc_correction = count_overconstrained_interface(flat, state_vars)?; if oc_correction >= 0 { dae.metadata.overconstrained_interface_count = oc_correction; @@ -453,6 +510,14 @@ fn finalize_lowered_dae( run_todae_phase(todae_subphase_timing, "reference_validation", || { validate_dae_references(dae, &known_flat_var_names) })?; + run_todae_phase( + todae_subphase_timing, + "prune_unreferenced_algebraics", + || { + prune_unreferenced_local_algebraics(dae); + Ok::<(), ToDaeError>(()) + }, + )?; if ir_boundary_validation_enabled() { dae.validate_shape_contract().map_err(|err| { ToDaeError::runtime_contract_violation_at( @@ -462,7 +527,10 @@ fn finalize_lowered_dae( })?; } - if options.error_on_unbalanced && !dae.metadata.is_partial && balance::balance(dae)? != 0 { + if options.error_on_unbalanced + && !dae.metadata.is_partial + && !balance::is_balanced_for_admission(dae)? + { let (equations, unknowns) = balance::equations_unknowns(dae)?; return Err(ToDaeError::unbalanced(equations, unknowns)); } @@ -470,6 +538,105 @@ fn finalize_lowered_dae( Ok(()) } +fn prune_unreferenced_local_algebraics(dae: &mut dae::Dae) { + let referenced = collect_dae_referenced_var_names(dae); + dae.variables.algebraics.retain(|name, variable| { + referenced.contains(name) + || variable.causality != dae::VariableCausality::Local + || variable.origin != dae::VariableOrigin::Source + || !is_structured_prune_candidate(variable) + }); +} + +fn is_structured_prune_candidate(variable: &dae::Variable) -> bool { + !variable.dims.is_empty() + || variable + .component_ref + .as_ref() + .is_some_and(|component_ref| { + component_ref.parts.len() > 1 + || component_ref.parts.iter().any(|part| !part.subs.is_empty()) + }) +} + +fn collect_dae_referenced_var_names(dae: &dae::Dae) -> HashSet { + let mut collector = DaeVarRefCollector { + names: HashSet::new(), + }; + collector.collect_equations(&dae.continuous.equations); + collector.collect_equations(&dae.initialization.equations); + collector.collect_equations(&dae.discrete.real_updates); + collector.collect_equations(&dae.discrete.valued_updates); + collector.collect_equations(&dae.conditions.equations); + collector.collect_expressions(&dae.conditions.relations); + collector.collect_expressions(&dae.events.synthetic_root_conditions); + for action in &dae.events.event_actions { + collector.visit_expression(&action.condition); + } + collector.collect_expressions(&dae.clocks.constructor_exprs); + collector.collect_expressions(&dae.clocks.triggered_conditions); + collector.collect_variable_attributes(&dae.variables); + collector.names +} + +struct DaeVarRefCollector { + names: HashSet, +} + +impl DaeVarRefCollector { + fn collect_equations(&mut self, equations: &[dae::Equation]) { + for equation in equations { + if let Some(lhs) = &equation.lhs { + self.names.insert(lhs.var_name().clone()); + } + self.visit_expression(&equation.rhs); + } + } + + fn collect_expressions(&mut self, expressions: &[Expression]) { + for expression in expressions { + self.visit_expression(expression); + } + } + + fn collect_variable_attributes(&mut self, variables: &dae::DaeVariables) { + self.collect_variable_partition_attributes(&variables.states); + self.collect_variable_partition_attributes(&variables.algebraics); + self.collect_variable_partition_attributes(&variables.inputs); + self.collect_variable_partition_attributes(&variables.outputs); + self.collect_variable_partition_attributes(&variables.parameters); + self.collect_variable_partition_attributes(&variables.constants); + self.collect_variable_partition_attributes(&variables.discrete_reals); + self.collect_variable_partition_attributes(&variables.discrete_valued); + } + + fn collect_variable_partition_attributes( + &mut self, + variables: &IndexMap, + ) { + for variable in variables.values() { + for expression in [ + variable.start.as_ref(), + variable.min.as_ref(), + variable.max.as_ref(), + variable.nominal.as_ref(), + ] + .into_iter() + .flatten() + { + self.visit_expression(expression); + } + } + } +} + +impl ExpressionVisitor for DaeVarRefCollector { + fn visit_var_ref(&mut self, name: &rumoca_core::Reference, subscripts: &[Subscript]) { + self.names.insert(name.var_name().clone()); + self.walk_var_ref(name, subscripts); + } +} + fn ir_boundary_validation_enabled() -> bool { cfg!(any( debug_assertions, @@ -478,6 +645,57 @@ fn ir_boundary_validation_enabled() -> bool { )) } +fn refresh_external_discrete_input_metadata(dae: &mut dae::Dae, flat: &flat::Model) { + let targeted = targeted_discrete_variable_names(dae); + dae.metadata.discrete_input_names.clear(); + let mut names = Vec::new(); + names.extend(external_discrete_input_names( + &dae.variables.discrete_valued, + &targeted, + flat, + )); + names.sort(); + names.dedup(); + dae.metadata.discrete_input_names = names; +} + +fn targeted_discrete_variable_names(dae: &dae::Dae) -> HashSet { + dae.discrete + .real_updates + .iter() + .chain(dae.discrete.valued_updates.iter()) + .chain(dae.conditions.equations.iter()) + .filter_map(|eq| eq.lhs.as_ref().map(|lhs| lhs.var_name().clone())) + .collect() +} + +fn external_discrete_input_names( + variables: &IndexMap, + targeted: &HashSet, + flat: &flat::Model, +) -> Vec { + variables + .iter() + .filter(|(name, variable)| { + matches!( + variable.causality, + dae::VariableCausality::Input | dae::VariableCausality::Output + ) && !targeted.contains(*name) + && flat_variable_is_connected(flat, name) + }) + .map(|(name, _)| name.as_str().to_string()) + .collect() +} + +fn flat_variable_is_connected(flat: &flat::Model, name: &rumoca_core::VarName) -> bool { + flat.variables + .get(name) + .is_some_and(|variable| variable.connected) + || subscript_fallback_chain(name.as_str()) + .into_iter() + .any(|candidate| flat.variables.get(&candidate).is_some_and(|v| v.connected)) +} + /// Determine if an algebraic variable should be stored as discrete or regular algebraic. /// Returns Some(map_name) if not a regular algebraic, None if it's a regular algebraic. enum AlgebraicCategory { @@ -500,6 +718,183 @@ fn categorize_algebraic( } } +#[derive(Debug)] +struct ExpandableProjectionLane { + declaration_span: Span, + instance: rumoca_core::DefId, + indices: Vec, + field_dims: Vec, + projection: ComponentReference, +} + +fn concrete_component_indices(subscripts: &[Subscript]) -> Option> { + if subscripts.is_empty() { + return None; + } + subscripts + .iter() + .map(|subscript| match subscript { + Subscript::Index { value, .. } => Some(*value), + Subscript::Colon { .. } | Subscript::Expr { .. } => None, + }) + .collect() +} + +fn same_component_projection(lhs: &ComponentReference, rhs: &ComponentReference) -> bool { + lhs.local == rhs.local + && lhs.parts.len() == rhs.parts.len() + && lhs.parts.iter().zip(&rhs.parts).all(|(lhs, rhs)| { + lhs.ident == rhs.ident + && lhs.subs.len() == rhs.subs.len() + && lhs.subs.iter().zip(&rhs.subs).all(|(lhs, rhs)| { + matches!( + (lhs, rhs), + ( + Subscript::Index { value: lhs, .. }, + Subscript::Index { value: rhs, .. } + ) if lhs == rhs + ) + }) + }) +} + +fn aggregate_has_only_complete_projection_connections( + flat: &flat::Model, + name: &VarName, + expected_scalar_count: usize, +) -> bool { + let mut seen = false; + for equation in &flat.equations { + if !collect_var_refs(&equation.residual).contains(name) { + continue; + } + seen = true; + if !equation.origin.is_connection() || equation.scalar_count != expected_scalar_count { + return false; + } + } + seen +} + +fn is_complete_expandable_projection( + flat: &flat::Model, + name: &VarName, + lanes: &[ExpandableProjectionLane], + directly_defined: &HashSet, +) -> Option<()> { + let aggregate = flat.variables.get(name)?; + let aggregate_ref = aggregate.component_ref.as_ref()?; + let first = lanes.first()?; + let rank = first.indices.len(); + let domain_dims = aggregate.dims.get(..rank)?; + let field_dims = aggregate.dims.get(rank..)?; + + let aggregate_is_only_projection = aggregate.from_expandable_connector + && aggregate.is_primitive + && aggregate.binding.is_none() + && matches!(aggregate.causality, rumoca_core::Causality::Empty) + && !directly_defined.contains(name) + && !domain_dims.is_empty() + && domain_dims.iter().all(|dim| *dim > 0) + && field_dims == first.field_dims + && aggregate_ref.def_id.is_none() + && lanes.iter().all(|lane| { + lane.declaration_span == first.declaration_span + && lane.indices.len() == rank + && lane.field_dims == first.field_dims + && same_component_projection(&lane.projection, aggregate_ref) + && lane + .indices + .iter() + .zip(domain_dims) + .all(|(index, dim)| (1..=*dim).contains(index)) + }); + aggregate_is_only_projection.then_some(())?; + + let scalar_width = |dims: &[i64]| { + dims.iter().try_fold(1_usize, |count, dim| { + usize::try_from(*dim) + .ok() + .and_then(|dim| count.checked_mul(dim)) + }) + }; + let expected = scalar_width(domain_dims)?; + let expected_scalar_count = expected.checked_mul(scalar_width(field_dims)?)?; + if aggregate.connected + && !aggregate_has_only_complete_projection_connections(flat, name, expected_scalar_count) + { + return None; + } + let concrete_domain = lanes + .iter() + .map(|lane| lane.indices.clone()) + .collect::>(); + let concrete_instances = lanes + .iter() + .map(|lane| lane.instance) + .collect::>(); + (concrete_domain.len() == expected + && concrete_instances.len() == expected + && lanes.len() == expected) + .then_some(()) +} + +fn expandable_aggregate_projection_names(flat: &flat::Model) -> HashSet { + let mut lanes: HashMap> = HashMap::new(); + for candidate in flat.variables.values() { + if !candidate.from_expandable_connector || !candidate.is_primitive { + continue; + } + let Some(component_ref) = candidate.component_ref.as_ref() else { + continue; + }; + if component_ref.span.is_dummy() { + continue; + } + let Some(instance) = component_ref.def_id else { + continue; + }; + let mut candidate_projections = Vec::new(); + for (index, part) in component_ref.parts.iter().enumerate() { + let Some(indices) = concrete_component_indices(&part.subs) else { + continue; + }; + let mut projection = component_ref.clone(); + projection.parts[index].subs.clear(); + let projection_name = projection.to_var_name(); + if flat + .variables + .get(&projection_name) + .is_some_and(|aggregate| aggregate.from_expandable_connector) + { + candidate_projections.push((projection_name, projection, indices)); + } + } + // More than one indexed container would make the projection identity ambiguous. + let [(projection_name, projection, indices)] = candidate_projections.as_slice() else { + continue; + }; + lanes + .entry(projection_name.clone()) + .or_default() + .push(ExpandableProjectionLane { + declaration_span: component_ref.span, + instance, + indices: indices.clone(), + field_dims: candidate.dims.clone(), + projection: projection.clone(), + }); + } + + let (directly_defined, _) = collect_continuous_equation_lhs(flat); + lanes + .into_iter() + .filter_map(|(name, lanes)| { + is_complete_expandable_projection(flat, &name, &lanes, &directly_defined).map(|()| name) + }) + .collect() +} + fn has_clocked_binding(var: &flat::Variable) -> bool { var.binding .as_ref() @@ -510,6 +905,7 @@ fn has_clocked_binding(var: &flat::Variable) -> bool { struct VariableClassificationIndexes { prefix_children: FxHashMap>, state_vars: IndexSet, + overconstrained_derivative_alias_roots: FxHashMap, connected_inputs: IndexSet, discrete_connected_inputs: IndexSet, input_only_connected_inputs: IndexSet, @@ -539,6 +935,8 @@ fn build_variable_classification_indexes( let prefix_children = build_prefix_children(flat); let internal_inputs = InternalInputIndex::new(flat)?; let der_vars = classification::find_state_variables(flat); + let overconstrained_derivative_alias_roots = + find_overconstrained_derivative_alias_roots(&der_vars, flat); let state_vars = filter_state_variables(der_vars, flat, &internal_inputs); let mut connected_input_set = find_connected_inputs(flat, &internal_inputs); connected_input_set.extend(find_equation_defined_inputs(flat, &internal_inputs)); @@ -560,6 +958,7 @@ fn build_variable_classification_indexes( Ok(VariableClassificationIndexes { prefix_children, state_vars, + overconstrained_derivative_alias_roots, connected_inputs, discrete_connected_inputs, input_only_connected_inputs, @@ -606,8 +1005,22 @@ fn classify_variables( .keys() .map(|name| name.as_str().to_string()) .collect(); + let expandable_aggregate_projections = expandable_aggregate_projection_names(flat); for (name, var) in &flat.variables { + if variable_analysis::is_external_constructor_handle(flat, name, var) { + continue; + } + + // A nested array inside an expandable connector can be represented by + // both an aggregate field projection (`bus.cells.x`) and its concrete + // structured lanes (`bus.cells[1].x`, ...). The aggregate is only a + // view used to expand connection equations, not a second runtime + // variable with discrete semantics. + if expandable_aggregate_projections.contains(name) { + continue; + } + // Skip non-primitive aggregate variables whose primitive fields are // represented separately (MLS §4.8). Keep non-primitive leaves so // connector-typed scalar aliases remain available in the DAE. @@ -725,6 +1138,274 @@ fn classify_variables( Ok(()) } +fn rewrite_overconstrained_derivative_alias_refs( + dae: &mut dae::Dae, + flat: &flat::Model, + alias_roots: &FxHashMap, +) -> Result<(), ToDaeError> { + if alias_roots.is_empty() { + return Ok(()); + } + + let mut rewritten_aliases = FxHashSet::default(); + for equation in dae + .continuous + .equations + .iter_mut() + .chain(dae.initialization.equations.iter_mut()) + .chain(dae.discrete.real_updates.iter_mut()) + .chain(dae.discrete.valued_updates.iter_mut()) + .chain(dae.conditions.equations.iter_mut()) + { + rewrite_overconstrained_derivative_alias_expr( + &mut equation.rhs, + alias_roots, + &mut rewritten_aliases, + ); + } + + for expr in dae + .conditions + .relations + .iter_mut() + .chain(dae.events.synthetic_root_conditions.iter_mut()) + .chain(dae.clocks.constructor_exprs.iter_mut()) + .chain(dae.clocks.triggered_conditions.iter_mut()) + { + rewrite_overconstrained_derivative_alias_expr(expr, alias_roots, &mut rewritten_aliases); + } + + add_overconstrained_derivative_alias_equations(dae, flat, alias_roots) +} + +fn add_overconstrained_derivative_alias_equations( + dae: &mut dae::Dae, + flat: &flat::Model, + alias_roots: &FxHashMap, +) -> Result<(), ToDaeError> { + // Rewriting der(alias) to der(root) preserves the dynamic derivative use, + // but the alias variable remains an algebraic unknown after state filtering. + // Connection-graph break-edge accounting cannot stand in for this local + // alias closure, so emit one residual row for every removed derivative state. + for (alias, root) in alias_roots { + let Some(alias_var) = flat.variables.get(alias) else { + continue; + }; + if !flat.variables.contains_key(root) { + continue; + } + let residual = derivative_alias_residual(alias, root, alias_var.source_span); + let rhs = flat_to_dae_expression_with_refs(&residual, flat)?; + dae.continuous.equations.push(dae::Equation::residual_array( + rhs, + alias_var.source_span, + format!( + "overconstrained derivative alias: {} = {}", + alias.as_str(), + root.as_str() + ), + flat_variable_scalar_count(alias_var), + )); + } + Ok(()) +} + +fn rewrite_overconstrained_derivative_alias_expr( + expr: &mut rumoca_core::Expression, + alias_roots: &FxHashMap, + rewritten_aliases: &mut FxHashSet, +) { + match expr { + rumoca_core::Expression::BuiltinCall { + function, + args, + span, + } if *function == BuiltinFunction::Der => { + if rewrite_derivative_alias_call(args, *span, alias_roots, rewritten_aliases) { + return; + } + rewrite_overconstrained_derivative_alias_exprs(args, alias_roots, rewritten_aliases); + } + rumoca_core::Expression::Binary { lhs, rhs, .. } => { + rewrite_overconstrained_derivative_alias_expr(lhs, alias_roots, rewritten_aliases); + rewrite_overconstrained_derivative_alias_expr(rhs, alias_roots, rewritten_aliases); + } + rumoca_core::Expression::Unary { rhs, .. } => { + rewrite_overconstrained_derivative_alias_expr(rhs, alias_roots, rewritten_aliases); + } + rumoca_core::Expression::BuiltinCall { args, .. } + | rumoca_core::Expression::FunctionCall { args, .. } + | rumoca_core::Expression::Array { elements: args, .. } + | rumoca_core::Expression::Tuple { elements: args, .. } => { + rewrite_overconstrained_derivative_alias_exprs(args, alias_roots, rewritten_aliases); + } + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + for (condition, value) in branches { + rewrite_overconstrained_derivative_alias_expr( + condition, + alias_roots, + rewritten_aliases, + ); + rewrite_overconstrained_derivative_alias_expr( + value, + alias_roots, + rewritten_aliases, + ); + } + rewrite_overconstrained_derivative_alias_expr( + else_branch, + alias_roots, + rewritten_aliases, + ); + } + rumoca_core::Expression::Range { + start, step, end, .. + } => { + rewrite_overconstrained_derivative_alias_expr(start, alias_roots, rewritten_aliases); + if let Some(step) = step { + rewrite_overconstrained_derivative_alias_expr(step, alias_roots, rewritten_aliases); + } + rewrite_overconstrained_derivative_alias_expr(end, alias_roots, rewritten_aliases); + } + rumoca_core::Expression::ArrayComprehension { + expr, + indices, + filter, + .. + } => { + rewrite_overconstrained_derivative_alias_expr(expr, alias_roots, rewritten_aliases); + for index in indices { + rewrite_overconstrained_derivative_alias_expr( + &mut index.range, + alias_roots, + rewritten_aliases, + ); + } + if let Some(filter) = filter { + rewrite_overconstrained_derivative_alias_expr( + filter, + alias_roots, + rewritten_aliases, + ); + } + } + rumoca_core::Expression::Index { + base, subscripts, .. + } => { + rewrite_overconstrained_derivative_alias_expr(base, alias_roots, rewritten_aliases); + rewrite_overconstrained_derivative_alias_subscripts( + subscripts, + alias_roots, + rewritten_aliases, + ); + } + rumoca_core::Expression::FieldAccess { base, .. } => { + rewrite_overconstrained_derivative_alias_expr(base, alias_roots, rewritten_aliases); + } + rumoca_core::Expression::VarRef { subscripts, .. } => { + rewrite_overconstrained_derivative_alias_subscripts( + subscripts, + alias_roots, + rewritten_aliases, + ); + } + rumoca_core::Expression::Literal { .. } | rumoca_core::Expression::Empty { .. } => {} + } +} + +fn rewrite_derivative_alias_call( + args: &mut Vec, + span: Span, + alias_roots: &FxHashMap, + rewritten_aliases: &mut FxHashSet, +) -> bool { + let Some(rumoca_core::Expression::VarRef { + name, subscripts, .. + }) = args.first_mut() + else { + return false; + }; + if !subscripts.is_empty() { + return false; + } + let Some(root) = alias_roots.get(name.var_name()) else { + return false; + }; + rewritten_aliases.insert(name.var_name().clone()); + *args = vec![derivative_alias_var_ref(root, span)]; + true +} + +fn rewrite_overconstrained_derivative_alias_exprs( + exprs: &mut [rumoca_core::Expression], + alias_roots: &FxHashMap, + rewritten_aliases: &mut FxHashSet, +) { + for expr in exprs { + rewrite_overconstrained_derivative_alias_expr(expr, alias_roots, rewritten_aliases); + } +} + +fn rewrite_overconstrained_derivative_alias_subscripts( + subscripts: &mut [rumoca_core::Subscript], + alias_roots: &FxHashMap, + rewritten_aliases: &mut FxHashSet, +) { + for subscript in subscripts { + if let rumoca_core::Subscript::Expr { expr, .. } = subscript { + rewrite_overconstrained_derivative_alias_expr(expr, alias_roots, rewritten_aliases); + } + } +} + +fn derivative_alias_var_ref(name: &rumoca_core::VarName, span: Span) -> rumoca_core::Expression { + rumoca_core::Expression::VarRef { + name: name.clone().into(), + subscripts: Vec::new(), + span, + } +} + +fn derivative_alias_residual( + alias: &rumoca_core::VarName, + root: &rumoca_core::VarName, + span: Span, +) -> rumoca_core::Expression { + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(derivative_var_ref(alias, span)), + rhs: Box::new(derivative_var_ref(root, span)), + span, + } +} + +fn derivative_var_ref(name: &rumoca_core::VarName, span: Span) -> rumoca_core::Expression { + rumoca_core::Expression::BuiltinCall { + function: BuiltinFunction::Der, + args: vec![rumoca_core::Expression::VarRef { + name: name.clone().into(), + subscripts: Vec::new(), + span, + }], + span, + } +} + +fn flat_variable_scalar_count(var: &flat::Variable) -> usize { + if var.dims.is_empty() { + return 1; + } + var.dims + .iter() + .map(|dimension| usize::try_from((*dimension).max(1)).unwrap_or(1)) + .product::() + .max(1) +} + fn record_variable_start_metadata( dae: &mut dae::Dae, name: &rumoca_core::VarName, @@ -989,7 +1670,7 @@ fn get_output_in_input_output_connection( fn classify_equations( dae: &mut dae::Dae, flat: &flat::Model, - prefix_counts: &FxHashMap, + prefix_counts: &ScalarInferenceMetadata, ) -> Result<(), ToDaeError> { equation_conversion::classify_equations(dae, flat, prefix_counts) } diff --git a/crates/rumoca-phase-dae/src/overconstrained_interface.rs b/crates/rumoca-phase-dae/src/overconstrained_interface.rs index 5d897c624..ae011bc26 100644 --- a/crates/rumoca-phase-dae/src/overconstrained_interface.rs +++ b/crates/rumoca-phase-dae/src/overconstrained_interface.rs @@ -157,6 +157,124 @@ pub(crate) fn count_overconstrained_interface( Ok(correction) } +/// Count rooted overconstrained component gauge freedoms. +/// +/// A VCG component with a definite or selected potential root has one free +/// reference coordinate for the overconstrained record. This is not an ordinary +/// missing equation; admission consumes it only as deficit-only closure. +pub(crate) fn count_overconstrained_root_gauge( + flat: &flat::Model, + state_vars: &IndexSet, +) -> Result { + let record_groups = collect_oc_record_groups(flat, state_vars)?; + if record_groups.is_empty() { + return Ok(0); + } + + let record_paths: Vec<&str> = record_groups.keys().map(|s| s.as_str()).collect(); + let (comp_of, n_comps) = + build_record_components(&record_paths, &flat.branches, &flat.optional_edges); + let rooted_components = rooted_component_flags(flat, &comp_of, n_comps); + + let mut gauge_by_component = vec![0usize; n_comps]; + for (rec_path, group) in &record_groups { + let Some(&comp_id) = comp_of.get(rec_path.as_str()) else { + continue; + }; + if !rooted_components[comp_id] { + continue; + } + let gauge = group + .total_scalar_size + .saturating_sub(group.eq_constraint_size); + gauge_by_component[comp_id] = gauge_by_component[comp_id].max(gauge); + } + + Ok(gauge_by_component.into_iter().sum()) +} + +fn collect_oc_record_groups( + flat: &flat::Model, + state_vars: &IndexSet, +) -> Result, ToDaeError> { + let defined_lhs_paths = collect_nonconnection_lhs_paths(flat); + let mut record_groups: FxHashMap = FxHashMap::default(); + for (name, var) in &flat.variables { + if !var.is_overconstrained { + continue; + } + let Some(ref rec_path) = var.oc_record_path else { + continue; + }; + if !crate::is_continuous_unknown(flat, state_vars, name) { + continue; + } + let group = match record_groups.entry(rec_path.clone()) { + std::collections::hash_map::Entry::Occupied(entry) => entry.into_mut(), + std::collections::hash_map::Entry::Vacant(entry) => { + let eq_constraint_size = var.oc_eq_constraint_size.ok_or_else(|| { + ToDaeError::runtime_contract_violation_at( + format!( + "overconstrained variable `{name}` is missing equalityConstraint size" + ), + var.source_span, + ) + })?; + entry.insert(OcRecordGroup { + total_scalar_size: 0, + eq_constraint_size, + has_internal_definition: defined_lhs_paths + .iter() + .any(|lhs| is_same_or_child(lhs, rec_path)), + }) + } + }; + group.total_scalar_size += flat_variable_scalar_size(var); + } + Ok(record_groups) +} + +fn rooted_component_flags( + flat: &flat::Model, + comp_of: &FxHashMap<&str, usize>, + n_comps: usize, +) -> Vec { + let mut has_root = vec![false; n_comps]; + for (rec_path, &comp_id) in comp_of { + if flat.definite_roots.contains(*rec_path) { + has_root[comp_id] = true; + } + for root in &flat.definite_roots { + if rec_path.starts_with(root.as_str()) + && rec_path.as_bytes().get(root.len()) == Some(&b'.') + { + has_root[comp_id] = true; + } + } + } + + for (pot_root, _priority) in &flat.potential_roots { + for (rec_path, &comp_id) in comp_of { + let matches = rec_path.starts_with(pot_root.as_str()) + || pot_root.starts_with(*rec_path) + || crate::path_utils::first_rendered_segment(rec_path) + .is_some_and(|prefix| pot_root.starts_with(prefix)); + if matches && !has_root[comp_id] { + has_root[comp_id] = true; + } + } + } + has_root +} + +fn flat_variable_scalar_size(var: &flat::Variable) -> usize { + if var.dims.is_empty() { + 1 + } else { + var.dims.iter().map(|&d| d.max(1) as usize).product() + } +} + /// Build connected components from record paths using VCG branches. /// Returns (path -> component_id, number_of_components). pub(crate) fn build_record_components<'a>( diff --git a/crates/rumoca-phase-dae/src/path_utils.rs b/crates/rumoca-phase-dae/src/path_utils.rs index c6057ad88..5ef681060 100644 --- a/crates/rumoca-phase-dae/src/path_utils.rs +++ b/crates/rumoca-phase-dae/src/path_utils.rs @@ -35,6 +35,43 @@ pub(crate) fn resolve_known_path_suffix( .find(|candidate| known_names.contains(candidate)) } +pub(crate) fn repeated_indexed_component_path_candidates(path: &str) -> Vec { + let parts = split_top_level_path_parts(path); + let mut candidates = Vec::new(); + for index in 0..parts.len() { + let Some(bracket) = parts[index].find('[') else { + continue; + }; + let repeated = &parts[index][..bracket]; + if repeated.is_empty() { + continue; + } + let mut candidate = parts.clone(); + candidate.insert(index, repeated.to_string()); + candidates.push(candidate.join(".")); + } + candidates +} + +fn split_top_level_path_parts(path: &str) -> Vec { + let mut parts = Vec::new(); + let mut depth = 0i32; + let mut start = 0usize; + for (index, ch) in path.char_indices() { + match ch { + '[' => depth += 1, + ']' => depth -= 1, + '.' if depth == 0 => { + parts.push(path[start..index].to_string()); + start = index + 1; + } + _ => {} + } + } + parts.push(path[start..].to_string()); + parts +} + #[cfg(test)] mod tests { use super::*; @@ -135,4 +172,16 @@ mod tests { Some("aimcData.statorCoreParameters.wRef") ); } + + #[test] + fn test_repeated_indexed_component_path_candidates_insert_repeated_segment() { + assert_eq!( + repeated_indexed_component_path_candidates("triac[1].thyristor1.off"), + vec!["triac.triac[1].thyristor1.off"] + ); + assert_eq!( + repeated_indexed_component_path_candidates("imc.stator.resistor[1].v"), + vec!["imc.stator.resistor.resistor[1].v"] + ); + } } diff --git a/crates/rumoca-phase-dae/src/pre_lowering.rs b/crates/rumoca-phase-dae/src/pre_lowering.rs index 70ad7b58f..6449df6c9 100644 --- a/crates/rumoca-phase-dae/src/pre_lowering.rs +++ b/crates/rumoca-phase-dae/src/pre_lowering.rs @@ -421,6 +421,20 @@ fn resolve_pre_targets( target.source_name = target_name.clone(); continue; } + if let Some((source_name, partition)) = + repeated_indexed_component_pre_target(dae, target_name) + { + if target.require_discrete { + validate_pre_target_partition( + partition, + &source_name, + target.span, + target.allow_continuous_target, + )?; + } + target.source_name = source_name; + continue; + } let index = scalarized_index.get_or_insert_with(|| build_scalarized_field_index(dae)); if let Some(source_name) = singleton_scalarized_field_name(index, target_name, target.require_discrete) @@ -439,6 +453,21 @@ fn resolve_pre_targets( Ok(()) } +fn repeated_indexed_component_pre_target( + dae: &dae::Dae, + target_name: &rumoca_core::VarName, +) -> Option<(rumoca_core::VarName, dae::DaeVariablePartition)> { + for candidate in + crate::path_utils::repeated_indexed_component_path_candidates(target_name.as_str()) + { + let candidate = rumoca_core::VarName::new(candidate); + if let Some((partition, _)) = find_variable_partition(dae, &candidate) { + return Some((candidate, partition)); + } + } + None +} + type ScalarizedFieldIndex = std::collections::HashMap<(String, String), Vec<(rumoca_core::VarName, bool)>>; diff --git a/crates/rumoca-phase-dae/src/pre_lowering/tests.rs b/crates/rumoca-phase-dae/src/pre_lowering/tests.rs index 0d4a5f86e..baa91ab3c 100644 --- a/crates/rumoca-phase-dae/src/pre_lowering/tests.rs +++ b/crates/rumoca-phase-dae/src/pre_lowering/tests.rs @@ -69,6 +69,23 @@ fn indexed_field_access(base_name: &str, index: i64, field: &str) -> rumoca_core } } +fn repeated_indexed_component_field_pre_call( + base_name: &str, + index: i64, + component: &str, + field: &str, +) -> rumoca_core::Expression { + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Pre, + args: vec![rumoca_core::Expression::FieldAccess { + base: Box::new(indexed_field_access(base_name, index, component)), + field: field.to_string(), + span: test_span(45, 60 + component.len() + field.len()), + }], + span: test_span(20, 65 + base_name.len() + component.len() + field.len()), + } +} + fn integer_literal(value: i64) -> rumoca_core::Expression { rumoca_core::Expression::Literal { value: rumoca_core::Literal::Integer(value), @@ -751,6 +768,38 @@ fn test_lower_pre_rejects_missing_target_variable() -> Result<(), ToDaeError> { Ok(()) } +#[test] +fn test_lower_pre_resolves_repeated_indexed_component_target() -> Result<(), ToDaeError> { + let mut dae = dae::Dae::new(); + dae.variables.discrete_valued.insert( + rumoca_core::VarName::new("triac.triac[1].thyristor1.off"), + discrete_valued_var("triac.triac[1].thyristor1.off"), + ); + dae.discrete.valued_updates.push(dae::Equation::residual( + repeated_indexed_component_field_pre_call("triac", 1, "thyristor1", "off"), + test_span(1, 2), + "repeated indexed component pre target".to_string(), + )); + + lower_pre_operator(&mut dae)?; + + assert!( + dae.variables + .parameters + .contains_key(&rumoca_core::VarName::new( + "__pre__.triac.triac[1].thyristor1.off" + )), + "pre parameter should use the declared repeated component target" + ); + match &dae.discrete.valued_updates[0].rhs { + rumoca_core::Expression::VarRef { name, .. } => { + assert_eq!(name.as_str(), "__pre__.triac.triac[1].thyristor1.off"); + } + other => panic!("Expected VarRef, got {:?}", other), + } + Ok(()) +} + #[test] fn test_lower_pre_ignores_enum_literals_in_edge_relations() -> Result<(), ToDaeError> { let mut dae = dae::Dae::new(); diff --git a/crates/rumoca-phase-dae/src/reference_validation.rs b/crates/rumoca-phase-dae/src/reference_validation.rs index 734a63294..208a2744a 100644 --- a/crates/rumoca-phase-dae/src/reference_validation.rs +++ b/crates/rumoca-phase-dae/src/reference_validation.rs @@ -156,7 +156,7 @@ fn validate_constructor_field_selection( .collect(); short_matches.sort(); crate::log_todae_debug(format!( - "DEBUG TODAE missing constructor selection={} args_len={} short_matches={short_matches:?} total_functions={}", + "TODAE missing constructor selection={} args_len={} short_matches={short_matches:?} total_functions={}", selected_name, args.len(), candidates.len() @@ -170,7 +170,10 @@ fn validate_constructor_field_selection( let field_known = constructor.inputs.iter().any(|param| param.name == field) || constructor.outputs.iter().any(|param| param.name == field); - if !field_known { + let field_resolves_from_positional = + crate::constructor_field_selection::positional_constructor_arg_for_field(args, field) + .is_some(); + if !field_known && !field_resolves_from_positional { if crate::todae_debug_enabled() { let mut available_fields: Vec = constructor .inputs @@ -185,7 +188,7 @@ fn validate_constructor_field_selection( .collect(); available_fields.sort(); crate::log_todae_debug(format!( - "DEBUG TODAE constructor field missing selection={} available={available_fields:?}", + "TODAE constructor field missing selection={} available={available_fields:?}", selected_name )); } diff --git a/crates/rumoca-phase-dae/src/runtime_precompute/clock.rs b/crates/rumoca-phase-dae/src/runtime_precompute/clock.rs index 9e74a4cb4..cd456653f 100644 --- a/crates/rumoca-phase-dae/src/runtime_precompute/clock.rs +++ b/crates/rumoca-phase-dae/src/runtime_precompute/clock.rs @@ -421,7 +421,10 @@ impl ExpressionVisitor for SyntheticRootConditionCollector<'_> { let mut else_suppressed = self.suppress_events; for (cond, value) in branches { let condition_activation = condition_activation::runtime_activation(cond); - let cond_suppressed = else_suppressed || condition_activation.is_some(); + let event_suppressed_condition = is_event_suppressed_wrapper(cond); + let cond_suppressed = else_suppressed + || condition_activation.is_some() + || event_suppressed_condition; push_relation_root_if_event_condition( cond, cond_suppressed, @@ -431,7 +434,9 @@ impl ExpressionVisitor for SyntheticRootConditionCollector<'_> { self.visit_expr_with_suppression(cond, cond_suppressed); self.visit_expr_with_suppression( value, - else_suppressed || matches!(condition_activation, Some(false)), + else_suppressed + || matches!(condition_activation, Some(false)) + || event_suppressed_condition, ); else_suppressed |= matches!(condition_activation, Some(true)); } @@ -447,7 +452,8 @@ impl ExpressionVisitor for SyntheticRootConditionCollector<'_> { self.walk_expression(expr); } rumoca_core::Expression::BuiltinCall { - function: rumoca_core::BuiltinFunction::NoEvent, + function: + rumoca_core::BuiltinFunction::NoEvent | rumoca_core::BuiltinFunction::Smooth, args, .. } => { @@ -472,6 +478,16 @@ impl ExpressionVisitor for SyntheticRootConditionCollector<'_> { } } +fn is_event_suppressed_wrapper(expr: &rumoca_core::Expression) -> bool { + matches!( + expr, + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::NoEvent, + .. + } + ) +} + fn push_relation_root_if_event_condition( expr: &rumoca_core::Expression, suppress_events: bool, @@ -714,6 +730,52 @@ fn eval_clock_scalar_from_var_ref( inferred } +fn eval_clock_static_subscript( + subscript: &rumoca_core::Subscript, + constants: &HashMap, + sources: &SourceMap<'_>, + remaining_depth: usize, + visiting: &mut HashSet, +) -> Option { + let raw = match subscript { + rumoca_core::Subscript::Index { value, .. } => *value as f64, + rumoca_core::Subscript::Expr { expr, .. } => { + eval_clock_scalar_child(expr, constants, sources, remaining_depth, visiting)? + } + rumoca_core::Subscript::Colon { .. } => return None, + }; + let rounded = raw.round(); + if !rounded.is_finite() || rounded < 1.0 || (raw - rounded).abs() > 1.0e-12 { + return None; + } + Some(rounded as usize) +} + +fn eval_clock_array_index( + base: &rumoca_core::Expression, + subscripts: &[rumoca_core::Subscript], + constants: &HashMap, + sources: &SourceMap<'_>, + remaining_depth: usize, + visiting: &mut HashSet, +) -> Option { + let [subscript] = subscripts else { + return None; + }; + let index = + eval_clock_static_subscript(subscript, constants, sources, remaining_depth, visiting)?; + let rumoca_core::Expression::Array { + elements, + is_matrix: false, + .. + } = base + else { + return None; + }; + let element = elements.get(index.checked_sub(1)?)?; + eval_clock_scalar_child(element, constants, sources, remaining_depth, visiting) +} + fn eval_clock_scalar_with_sources( expr: &rumoca_core::Expression, constants: &HashMap, @@ -787,6 +849,22 @@ fn eval_clock_scalar_with_sources( eval_clock_scalar_child(lhs, constants, sources, remaining_depth, visiting)?; Some(numerator / denominator) } + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Integer, + args, + .. + } => eval_clock_scalar_child(args.first()?, constants, sources, remaining_depth, visiting) + .map(f64::floor), + rumoca_core::Expression::Index { + base, subscripts, .. + } => eval_clock_array_index( + base, + subscripts, + constants, + sources, + remaining_depth, + visiting, + ), _ => None, } } diff --git a/crates/rumoca-phase-dae/src/runtime_precompute/mod.rs b/crates/rumoca-phase-dae/src/runtime_precompute/mod.rs index b143d9182..e310e8a5c 100644 --- a/crates/rumoca-phase-dae/src/runtime_precompute/mod.rs +++ b/crates/rumoca-phase-dae/src/runtime_precompute/mod.rs @@ -95,8 +95,9 @@ pub(crate) fn populate_runtime_precompute(dae_model: &mut dae::Dae) -> Result<() &mut scheduled_time_events, ); } - scheduled_time_events.sort_by(f64::total_cmp); - scheduled_time_events.dedup_by(|a, b| (*a - *b).abs() <= 1e-12 * (1.0 + a.abs().max(b.abs()))); + scheduled_time_events.sort_by(|a, b| a.time.total_cmp(&b.time)); + scheduled_time_events + .dedup_by(|a, b| (a.time - b.time).abs() <= 1e-12 * (1.0 + a.time.abs().max(b.time.abs()))); log_runtime_precompute_profile("scheduled_time_events", time_event_start); let clock_metadata_start = maybe_start_timer_if(profile); @@ -250,9 +251,34 @@ fn prune_unreferenced_condition_memory( dae_model.conditions.relations = kept_relations; dae_model.conditions.equations = kept_conditions; rewrite_condition_memory_references(dae_model, &condition_name, &replacements)?; + resize_condition_memory_variables(dae_model, &condition_name); Ok(()) } +fn resize_condition_memory_variables(dae_model: &mut dae::Dae, condition_name: &str) { + let condition_len = dae_model.conditions.relations.len() as i64; + let condition_key = rumoca_core::VarName::new(condition_name); + let pre_condition_key = rumoca_core::VarName::new(format!("__pre__.{condition_name}")); + resize_condition_memory_variable( + dae_model.variables.discrete_valued.get_mut(&condition_key), + condition_len, + ); + resize_condition_memory_variable( + dae_model.variables.parameters.get_mut(&pre_condition_key), + condition_len, + ); +} + +fn resize_condition_memory_variable(variable: Option<&mut dae::Variable>, condition_len: i64) { + let Some(variable) = variable else { + return; + }; + variable.dims = vec![condition_len]; + if let Some(rumoca_core::Expression::Array { elements, .. }) = &mut variable.start { + elements.truncate(condition_len.max(0) as usize); + } +} + fn condition_memory_can_be_direct( relation: &rumoca_core::Expression, constants: &HashMap, @@ -595,12 +621,13 @@ fn collect_array_elements_scalar_entries( name: &rumoca_core::VarName, elements: &[rumoca_core::Expression], values: &HashMap, + dims: &indexmap::IndexMap>, ) -> Vec<(String, f64)> { elements .iter() .enumerate() .filter_map(|(index, element)| { - eval_scalar_const_expr(element, values).map(|value| { + eval_scalar_const_expr_with_dims(element, values, dims).map(|value| { ( dae::format_subscript_key(name.as_str(), &[index + 1]), value, @@ -615,16 +642,17 @@ fn collect_range_scalar_entries( step: Option<&rumoca_core::Expression>, end: &rumoca_core::Expression, values: &HashMap, + dims: &indexmap::IndexMap>, ) -> Vec<(String, f64)> { - let Some(mut current) = eval_scalar_const_expr(start, values) else { + let Some(mut current) = eval_scalar_const_expr_with_dims(start, values, dims) else { return Vec::new(); }; - let Some(end_value) = eval_scalar_const_expr(end, values) else { + let Some(end_value) = eval_scalar_const_expr_with_dims(end, values, dims) else { return Vec::new(); }; let step_value = match step { Some(expr) => { - let Some(value) = eval_scalar_const_expr(expr, values) else { + let Some(value) = eval_scalar_const_expr_with_dims(expr, values, dims) else { return Vec::new(); }; value @@ -656,20 +684,22 @@ fn collect_array_scalar_entries( name: &rumoca_core::VarName, start: &rumoca_core::Expression, values: &HashMap, + dims: &indexmap::IndexMap>, ) -> Vec<(String, f64)> { match start { rumoca_core::Expression::Array { elements, .. } | rumoca_core::Expression::Tuple { elements, .. } => { - collect_array_elements_scalar_entries(name, elements, values) + collect_array_elements_scalar_entries(name, elements, values, dims) } rumoca_core::Expression::Range { start, step, end, .. - } => collect_range_scalar_entries(name, start, step.as_deref(), end, values), + } => collect_range_scalar_entries(name, start, step.as_deref(), end, values, dims), _ => Vec::new(), } } fn collect_compile_time_scalars(dae_model: &dae::Dae) -> HashMap { let mut values = HashMap::new(); + let dims = rumoca_eval_dae::collect_var_dims(dae_model); for (literal, ordinal) in &dae_model.symbols.enum_literal_ordinals { values.insert(literal.clone(), *ordinal as f64); } @@ -679,6 +709,7 @@ fn collect_compile_time_scalars(dae_model: &dae::Dae) -> HashMap { .iter() .chain(dae_model.variables.constants.iter()) .chain(dae_model.variables.inputs.iter()) + .chain(dae_model.variables.discrete_valued.iter()) .filter_map(|(name, variable)| variable.start.as_ref().map(|start| (name, start))) .collect(); @@ -686,10 +717,11 @@ fn collect_compile_time_scalars(dae_model: &dae::Dae) -> HashMap { for _ in 0..max_passes { let mut changed = false; for (name, start) in &bindings { - if let Some(value) = eval_scalar_const_expr(start, &values) { + if let Some(value) = eval_scalar_const_expr_with_dims(start, &values, &dims) { changed |= insert_compile_time_scalar(&mut values, name.as_str(), value); } - for (indexed_name, indexed_value) in collect_array_scalar_entries(name, start, &values) + for (indexed_name, indexed_value) in + collect_array_scalar_entries(name, start, &values, &dims) { changed |= insert_compile_time_scalar(&mut values, &indexed_name, indexed_value); } @@ -809,6 +841,28 @@ fn eval_scalar_const_expr( .map(rumoca_eval_dae::constant::ConstValue::Real) }) } + +fn eval_scalar_const_expr_with_dims( + expr: &rumoca_core::Expression, + constants: &HashMap, + dims: &indexmap::IndexMap>, +) -> Option { + rumoca_eval_dae::constant::eval_scalar_const_expr_with_shape( + expr, + &|name, subscripts| { + if subscripts.is_empty() { + return constants + .get(name.as_str()) + .copied() + .map(rumoca_eval_dae::constant::ConstValue::Real); + } + clock::canonical_var_ref_key(name, subscripts, constants) + .and_then(|key| constants.get(&key).copied()) + .map(rumoca_eval_dae::constant::ConstValue::Real) + }, + &|name| dims.get(name).cloned(), + ) +} fn extract_time_event_instant( cond: &rumoca_core::Expression, constants: &HashMap, @@ -918,7 +972,7 @@ fn maybe_push_time_event_condition( suppress_events: bool, constants: &HashMap, seen: &mut HashSet, - out: &mut Vec, + out: &mut Vec, ) { if suppress_events { return; @@ -928,7 +982,10 @@ fn maybe_push_time_event_condition( }; let key = format!("time::{event_time:.15e}"); if seen.insert(key) { - out.push(event_time); + out.push(dae::DaeScheduledTimeEvent { + time: event_time, + source_span: expr.span(), + }); } } fn collect_time_discontinuity_events_expr( @@ -936,7 +993,7 @@ fn collect_time_discontinuity_events_expr( suppress_events: bool, constants: &HashMap, seen: &mut HashSet, - out: &mut Vec, + out: &mut Vec, ) { let mut collector = TimeDiscontinuityEventCollector { suppress_events, @@ -951,7 +1008,7 @@ struct TimeDiscontinuityEventCollector<'a> { suppress_events: bool, constants: &'a HashMap, seen: &'a mut HashSet, - out: &'a mut Vec, + out: &'a mut Vec, } impl TimeDiscontinuityEventCollector<'_> { diff --git a/crates/rumoca-phase-dae/src/runtime_precompute/tests/clock_schedule_tests.rs b/crates/rumoca-phase-dae/src/runtime_precompute/tests/clock_schedule_tests.rs new file mode 100644 index 000000000..9d8407075 --- /dev/null +++ b/crates/rumoca-phase-dae/src/runtime_precompute/tests/clock_schedule_tests.rs @@ -0,0 +1,1077 @@ +use super::*; + +#[test] +fn test_runtime_precompute_extracts_affine_time_event() { + let cond = rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Le, + lhs: Box::new(rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: Box::new(var("time")), + rhs: Box::new(var("delay")), + span: test_span(1, 2), + }), + rhs: Box::new(var("switch_at")), + span: test_span(1, 2), + }; + let mut dae_model = dae_with_if_condition(cond); + let mut delay = dae::Variable::new( + rumoca_core::VarName::new("delay"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ); + delay.start = Some(lit(0.25)); + dae_model + .variables + .parameters + .insert(rumoca_core::VarName::new("delay"), delay); + let mut switch_at = dae::Variable::new( + rumoca_core::VarName::new("switch_at"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ); + switch_at.start = Some(lit(1.5)); + dae_model + .variables + .parameters + .insert(rumoca_core::VarName::new("switch_at"), switch_at); + + populate_conditions(&mut dae_model); + populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); + assert_eq!(dae_model.events.scheduled_time_events.len(), 1); + assert!((dae_model.events.scheduled_time_events[0].time - 1.25).abs() <= 1e-12); +} + +#[test] +fn test_runtime_precompute_extracts_discrete_partition_events() { + let mut dae_model = dae::Dae::default(); + dae_model.variables.discrete_valued.insert( + rumoca_core::VarName::new("c"), + dae::Variable::new( + rumoca_core::VarName::new("c"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ), + ); + dae_model + .discrete + .valued_updates + .push(dae::Equation::residual( + sub( + var("c"), + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Gt, + lhs: Box::new(var("time")), + rhs: Box::new(lit(0.5)), + span: test_span(1, 2), + }, + ), + test_span(1, 2), + "test_discrete_partition", + )); + + populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); + assert_eq!(scheduled_times(&dae_model), vec![0.5]); +} + +#[test] +fn test_runtime_precompute_collects_clock_constructor_exprs() { + let mut dae_model = dae::Dae::default(); + dae_model.variables.discrete_reals.insert( + rumoca_core::VarName::new("s"), + dae::Variable::new( + rumoca_core::VarName::new("s"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ), + ); + dae_model + .discrete + .real_updates + .push(dae::Equation::residual( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new("s").into(), + subscripts: vec![], + span: test_span(1, 2), + }), + rhs: Box::new(rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Sample, + args: vec![ + rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new("u").into(), + subscripts: vec![], + span: test_span(100, 101), + }, + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("Clock").into(), + args: vec![rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(0.1), + span: test_span(108, 111), + }], + is_constructor: false, + span: test_span(102, 112), + }, + ], + span: test_span(95, 113), + }), + span: test_span(1, 2), + }, + test_span(1, 2), + "test_clock_constructor", + )); + + populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); + + assert_eq!(dae_model.clocks.constructor_exprs.len(), 1); + assert_eq!(dae_model.clocks.schedules.len(), 1); + assert!((dae_model.clocks.schedules[0].period_seconds - 0.1).abs() <= 1e-12); + assert!(dae_model.clocks.schedules[0].phase_seconds.abs() <= 1e-12); + assert!((dae_model.clocks.intervals["s"] - 0.1).abs() <= 1e-12); +} + +#[test] +fn test_runtime_precompute_rejects_static_clock_constructor_without_source_provenance() { + let mut dae_model = dae::Dae::default(); + dae_model.variables.discrete_reals.insert( + rumoca_core::VarName::new("s"), + dae::Variable::new( + rumoca_core::VarName::new("s"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ), + ); + dae_model + .discrete + .real_updates + .push(dae::Equation::residual( + sub( + var("s"), + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("Clock").into(), + args: vec![rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(0.1), + span: rumoca_core::Span::DUMMY, + }], + is_constructor: false, + span: rumoca_core::Span::DUMMY, + }, + ), + Span::DUMMY, + "unspanned_static_clock_constructor", + )); + + let err = populate_runtime_precompute(&mut dae_model) + .expect_err("source-free static clock constructors must fail fast"); + assert!(matches!( + err, + crate::ToDaeError::RuntimeMetadataViolation { detail } + if detail.contains("source provenance") + )); +} + +#[test] +fn test_runtime_precompute_collects_sample_start_interval_schedule() { + let mut dae_model = dae::Dae::default(); + dae_model.variables.discrete_reals.insert( + rumoca_core::VarName::new("s"), + dae::Variable::new( + rumoca_core::VarName::new("s"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ), + ); + dae_model + .discrete + .real_updates + .push(dae::Equation::residual( + sub( + var("s"), + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Sample, + args: vec![ + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(0.2), + span: test_span(120, 123), + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(0.1), + span: test_span(125, 128), + }, + ], + span: test_span(113, 129), + }, + ), + test_span(1, 2), + // MLS §16.5.1: sample(start, interval) defines a periodic event. + "periodic_sample_start_interval", + )); + + populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); + + assert_eq!(dae_model.clocks.schedules.len(), 1); + assert!((dae_model.clocks.schedules[0].period_seconds - 0.1).abs() <= 1e-12); + assert!((dae_model.clocks.schedules[0].phase_seconds - 0.2).abs() <= 1e-12); +} + +#[test] +fn test_runtime_precompute_marks_schedule_backed_sample_root_condition() { + let mut dae_model = dae::Dae::default(); + let relation = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new(rumoca_core::INTERNAL_SAMPLE_FUNCTION_NAME).into(), + args: vec![lit(42.0), lit(0.05), lit(0.1)], + is_constructor: false, + span: test_span(120, 149), + }; + dae_model.conditions.relations.push(relation.clone()); + dae_model.conditions.equations.push(dae::Equation::explicit( + condition_lhs("c", 1), + relation, + test_span(1, 2), + "scheduled sample condition memory", + )); + + populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); + + assert_eq!(dae_model.events.scheduled_root_conditions.len(), 1); + let root = &dae_model.events.scheduled_root_conditions[0]; + assert_eq!(root.root_index, 0); + assert!((root.period_seconds - 0.1).abs() <= 1e-12); + assert!((root.phase_seconds - 0.05).abs() <= 1e-12); +} + +#[test] +fn test_runtime_precompute_collects_sample_schedule_from_initial_time_parameter() { + let mut dae_model = dae::Dae::default(); + let mut frequency = dae::Variable::new( + rumoca_core::VarName::new("mean.f"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ); + frequency.start = Some(lit(150.0)); + dae_model + .variables + .parameters + .insert(rumoca_core::VarName::new("mean.f"), frequency); + dae_model.variables.parameters.insert( + rumoca_core::VarName::new("mean.t0"), + dae::Variable::new( + rumoca_core::VarName::new("mean.t0"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ), + ); + dae_model.variables.discrete_reals.insert( + rumoca_core::VarName::new("s"), + dae::Variable::new( + rumoca_core::VarName::new("s"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ), + ); + dae_model + .initialization + .equations + .push(dae::Equation::residual( + sub(var("mean.t0"), var("time")), + test_span(1, 2), + "initial t0 = time", + )); + + let interval = rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Div, + lhs: Box::new(lit(1.0)), + rhs: Box::new(var("mean.f")), + span: test_span(140, 148), + }; + dae_model + .discrete + .real_updates + .push(dae::Equation::residual( + sub( + var("s"), + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new(rumoca_core::INTERNAL_SAMPLE_FUNCTION_NAME) + .into(), + args: vec![ + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: Box::new(var("mean.t0")), + rhs: Box::new(interval.clone()), + span: test_span(130, 148), + }, + interval, + ], + is_constructor: false, + span: test_span(120, 149), + }, + ), + test_span(1, 2), + "periodic sample with initial-time origin", + )); + + populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); + + assert_eq!(dae_model.clocks.schedules.len(), 1); + assert!((dae_model.clocks.schedules[0].period_seconds - 1.0 / 150.0).abs() <= 1e-12); + assert!((dae_model.clocks.schedules[0].phase_seconds - 1.0 / 150.0).abs() <= 1e-12); +} + +#[test] +fn test_runtime_precompute_assigns_implicit_sample_interval_from_unique_schedule() { + let mut dae_model = dae::Dae::default(); + dae_model.variables.discrete_reals.insert( + rumoca_core::VarName::new("simTime"), + dae::Variable::new( + rumoca_core::VarName::new("simTime"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ), + ); + dae_model.variables.discrete_reals.insert( + rumoca_core::VarName::new("clockY"), + dae::Variable::new( + rumoca_core::VarName::new("clockY"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ), + ); + + // simTime = sample(time) (implicit clock sample form) + dae_model + .discrete + .real_updates + .push(dae::Equation::residual( + sub( + var("simTime"), + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Sample, + args: vec![var("time")], + span: test_span(1, 2), + }, + ), + test_span(1, 2), + "implicit_clocked_sample", + )); + + dae_model + .discrete + .real_updates + .push(dae::Equation::residual( + sub( + var("clockY"), + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("Clock").into(), + args: vec![rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(0.1), + span: test_span(140, 143), + }], + is_constructor: false, + span: test_span(134, 144), + }, + ), + test_span(1, 2), + "periodic_clock_constructor", + )); + + populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); + assert_eq!(dae_model.clocks.schedules.len(), 1); + assert!((dae_model.clocks.intervals["simTime"] - 0.1).abs() <= 1e-12); +} + +#[test] +fn test_runtime_precompute_propagates_no_argument_clock_guard_timing() { + let mut dae_model = dae::Dae::default(); + for name in ["u", "dummy", "b"] { + dae_model.variables.discrete_valued.insert( + rumoca_core::VarName::new(name), + dae::Variable::new( + rumoca_core::VarName::new(name), + rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + ), + ), + ); + } + let clock_span = test_span(1_000, 1_005); + dae_model.discrete.valued_updates.push(dae::Equation { + lhs: Some(rumoca_core::VarName::new("u").into()), + rhs: clock_call(0.02), + span: test_span(1, 2), + origin: "u = Clock(0.02)".to_string(), + scalar_count: 1, + }); + dae_model.discrete.valued_updates.push(dae::Equation { + lhs: Some(rumoca_core::VarName::new("dummy").into()), + rhs: if_then_else( + no_argument_clock_call(clock_span), + var("u"), + var("__pre__.dummy"), + ), + span: test_span(1, 2), + origin: "when Clock() then dummy = u".to_string(), + scalar_count: 1, + }); + dae_model.discrete.valued_updates.push(dae::Equation { + lhs: Some(rumoca_core::VarName::new("b").into()), + rhs: if_then_else( + no_argument_clock_call(clock_span), + rumoca_core::Expression::Unary { + op: rumoca_core::OpUnary::Not, + rhs: Box::new(var("__pre__.__pre__.b")), + span: test_span(1, 2), + }, + var("__pre__.b"), + ), + span: test_span(1, 2), + origin: "when Clock() then b = not previous(b)".to_string(), + scalar_count: 1, + }); + + populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); + + assert!((dae_model.clocks.intervals["dummy"] - 0.02).abs() <= 1e-12); + assert!((dae_model.clocks.intervals["b"] - 0.02).abs() <= 1e-12); +} + +#[test] +fn test_runtime_precompute_assigns_clock_interval_to_algebraic_alias_chain() { + let mut dae_model = dae::Dae::default(); + dae_model.variables.inputs.insert( + rumoca_core::VarName::new("u"), + dae::Variable::new( + rumoca_core::VarName::new("u"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ), + ); + dae_model.variables.algebraics.insert( + rumoca_core::VarName::new("feedback.y"), + dae::Variable::new( + rumoca_core::VarName::new("feedback.y"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ), + ); + dae_model.variables.algebraics.insert( + rumoca_core::VarName::new("PI.u"), + dae::Variable::new( + rumoca_core::VarName::new("PI.u"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ), + ); + dae_model.variables.discrete_reals.insert( + rumoca_core::VarName::new("sample2.y"), + dae::Variable::new( + rumoca_core::VarName::new("sample2.y"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ), + ); + dae_model.variables.discrete_valued.insert( + rumoca_core::VarName::new("sample2.clock"), + dae::Variable::new( + rumoca_core::VarName::new("sample2.clock"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ), + ); + + dae_model + .discrete + .real_updates + .push(dae::Equation::explicit( + rumoca_core::VarName::new("sample2.clock"), + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("Clock").into(), + args: vec![rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(0.1), + span: test_span(150, 153), + }], + is_constructor: false, + span: test_span(144, 154), + }, + test_span(1, 2), + "explicit_clock_alias", + )); + dae_model + .discrete + .real_updates + .push(dae::Equation::explicit( + rumoca_core::VarName::new("sample2.y"), + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Sample, + args: vec![var("u"), var("sample2.clock")], + span: test_span(1, 2), + }, + test_span(1, 2), + "explicit_sample_value", + )); + dae_model.continuous.equations.push(dae::Equation::residual( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(var("sample2.y")), + rhs: Box::new(var("feedback.y")), + span: test_span(1, 2), + }, + test_span(1, 2), + "sample_alias", + )); + dae_model.continuous.equations.push(dae::Equation::residual( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(var("feedback.y")), + rhs: Box::new(var("PI.u")), + span: test_span(1, 2), + }, + test_span(1, 2), + "controller_input_alias", + )); + populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); + assert!((dae_model.clocks.intervals["sample2.clock"] - 0.1).abs() <= 1e-12); + assert!((dae_model.clocks.intervals["sample2.y"] - 0.1).abs() <= 1e-12); + assert!((dae_model.clocks.intervals["feedback.y"] - 0.1).abs() <= 1e-12); + assert!((dae_model.clocks.intervals["PI.u"] - 0.1).abs() <= 1e-12); +} + +#[test] +fn test_runtime_precompute_propagates_clock_across_indexed_vector_and_previous() { + let mut dae_model = dae::Dae::default(); + let mut sampled = dae::Variable::new(rumoca_core::VarName::new("sampled"), test_span(1, 2)); + sampled.dims = vec![2]; + dae_model + .variables + .discrete_reals + .insert(sampled.name.clone(), sampled); + for name in ["delay1.u", "delay1.y", "delay2.u", "delay2.y"] { + dae_model.variables.discrete_reals.insert( + rumoca_core::VarName::new(name), + dae::Variable::new(rumoca_core::VarName::new(name), test_span(1, 2)), + ); + } + dae_model.variables.discrete_valued.insert( + rumoca_core::VarName::new("clock"), + dae::Variable::new(rumoca_core::VarName::new("clock"), test_span(1, 2)), + ); + + dae_model + .discrete + .valued_updates + .push(dae::Equation::explicit( + rumoca_core::VarName::new("clock"), + clock_call(0.02), + test_span(1, 2), + "clock_source", + )); + dae_model + .discrete + .real_updates + .push(dae::Equation::explicit( + rumoca_core::VarName::new("sampled"), + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Sample, + args: vec![var("u"), var("clock")], + span: test_span(1, 2), + }, + test_span(1, 2), + "sampled_vector", + )); + for (delay, index) in [("delay1", 1), ("delay2", 2)] { + dae_model + .discrete + .real_updates + .push(dae::Equation::explicit( + rumoca_core::VarName::new(format!("{delay}.u")), + condition_memory_ref("sampled", index), + test_span(1, 2), + "indexed_vector_alias", + )); + dae_model + .discrete + .real_updates + .push(dae::Equation::explicit( + rumoca_core::VarName::new(format!("{delay}.y")), + var(&format!("__pre__.{delay}.u")), + test_span(1, 2), + "clocked_previous_value", + )); + } + + populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); + + for name in ["sampled", "delay1.u", "delay1.y", "delay2.u", "delay2.y"] { + assert!((dae_model.clocks.intervals[name] - 0.02).abs() <= 1e-12); + } +} + +#[test] +fn test_runtime_precompute_propagates_uniform_clock_through_vector_alias_projection() { + let mut dae_model = dae::Dae::default(); + for name in ["sampled", "alias"] { + let mut variable = dae::Variable::new(rumoca_core::VarName::new(name), test_span(1, 2)); + variable.dims = vec![2]; + dae_model + .variables + .discrete_reals + .insert(variable.name.clone(), variable); + } + for name in ["delay.u", "delay.y"] { + dae_model.variables.discrete_reals.insert( + rumoca_core::VarName::new(name), + dae::Variable::new(rumoca_core::VarName::new(name), test_span(1, 2)), + ); + } + + dae_model + .discrete + .real_updates + .push(dae::Equation::explicit( + rumoca_core::VarName::new("sampled"), + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Sample, + args: vec![var("u"), clock_call(0.02)], + span: test_span(1, 2), + }, + test_span(1, 2), + "direct_clocked_vector", + )); + dae_model.continuous.equations.push(dae::Equation::explicit( + rumoca_core::VarName::new("alias"), + var("sampled"), + test_span(1, 2), + "untimed_vector_alias", + )); + dae_model + .discrete + .real_updates + .push(dae::Equation::explicit( + rumoca_core::VarName::new("delay.u"), + condition_memory_ref("alias", 1), + test_span(1, 2), + "indexed_alias_consumer", + )); + dae_model + .discrete + .real_updates + .push(dae::Equation::explicit( + rumoca_core::VarName::new("delay.y"), + var("__pre__.delay.u"), + test_span(1, 2), + "previous_consumer", + )); + + populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); + + for name in ["sampled", "alias", "delay.u", "delay.y"] { + assert!((dae_model.clocks.intervals[name] - 0.02).abs() <= 1e-12); + } +} + +#[test] +fn test_runtime_precompute_keeps_distinct_clocks_for_array_elements() { + let mut dae_model = dae::Dae::default(); + let mut sampled = dae::Variable::new(rumoca_core::VarName::new("sampled"), test_span(1, 2)); + sampled.dims = vec![2]; + dae_model + .variables + .discrete_reals + .insert(sampled.name.clone(), sampled); + for name in ["clock1", "clock2"] { + dae_model.variables.discrete_valued.insert( + rumoca_core::VarName::new(name), + dae::Variable::new(rumoca_core::VarName::new(name), test_span(1, 2)), + ); + } + for (name, period) in [("clock1", 0.1), ("clock2", 0.2)] { + dae_model + .discrete + .valued_updates + .push(dae::Equation::explicit( + rumoca_core::VarName::new(name), + clock_call(period), + test_span(1, 2), + "independent_clock_source", + )); + } + for (index, clock) in [(1, "clock1"), (2, "clock2")] { + dae_model.discrete.real_updates.push(dae::Equation { + lhs: Some(condition_lhs("sampled", index)), + rhs: rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Sample, + args: vec![var("u"), var(clock)], + span: test_span(1, 2), + }, + span: test_span(1, 2), + origin: "independently_clocked_array_element".to_string(), + scalar_count: 1, + }); + } + + populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); + + assert!(!dae_model.clocks.intervals.contains_key("sampled")); + assert!((dae_model.clocks.intervals["sampled[1]"] - 0.1).abs() <= 1e-12); + assert!((dae_model.clocks.intervals["sampled[2]"] - 0.2).abs() <= 1e-12); +} + +#[test] +fn test_runtime_precompute_promotes_equal_element_clocks_to_uniform_array() { + let mut dae_model = dae::Dae::default(); + for name in ["sampled", "alias"] { + let mut variable = dae::Variable::new(rumoca_core::VarName::new(name), test_span(1, 2)); + variable.dims = vec![2]; + dae_model + .variables + .discrete_reals + .insert(variable.name.clone(), variable); + } + for name in ["clock1", "clock2"] { + dae_model.variables.discrete_valued.insert( + rumoca_core::VarName::new(name), + dae::Variable::new(rumoca_core::VarName::new(name), test_span(1, 2)), + ); + dae_model + .discrete + .valued_updates + .push(dae::Equation::explicit( + rumoca_core::VarName::new(name), + clock_call(0.1), + test_span(1, 2), + "uniform_clock_source", + )); + } + for (index, clock) in [(1, "clock1"), (2, "clock2")] { + dae_model.discrete.real_updates.push(dae::Equation { + lhs: Some(condition_lhs("sampled", index)), + rhs: rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Sample, + args: vec![var("u"), var(clock)], + span: test_span(1, 2), + }, + span: test_span(1, 2), + origin: "uniformly_clocked_array_element".to_string(), + scalar_count: 1, + }); + } + dae_model.continuous.equations.push(dae::Equation::explicit( + rumoca_core::VarName::new("alias"), + var("sampled"), + test_span(1, 2), + "whole_array_alias", + )); + + populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); + + for name in ["sampled", "alias"] { + let interval = dae_model.clocks.intervals.get(name).unwrap_or_else(|| { + panic!( + "missing {name}; intervals={:?}", + dae_model.clocks.intervals.keys().collect::>() + ) + }); + assert!((*interval - 0.1).abs() <= 1e-12); + } +} + +#[test] +fn test_runtime_precompute_does_not_assign_fallback_interval_for_non_sample_clock_ops() { + let mut dae_model = dae::Dae::default(); + dae_model.variables.discrete_valued.insert( + rumoca_core::VarName::new("b"), + dae::Variable::new( + rumoca_core::VarName::new("b"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ), + ); + dae_model.variables.discrete_reals.insert( + rumoca_core::VarName::new("clockY"), + dae::Variable::new( + rumoca_core::VarName::new("clockY"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ), + ); + + // b = pre(b) is discrete/event logic, not an implicit sample(..) form. + dae_model + .discrete + .valued_updates + .push(dae::Equation::residual( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(var("b")), + rhs: Box::new(rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Pre, + args: vec![var("b")], + span: test_span(1, 2), + }), + span: test_span(1, 2), + }, + test_span(1, 2), + "pre_based_discrete_update", + )); + + // Add one static periodic schedule in the model. + dae_model + .discrete + .real_updates + .push(dae::Equation::residual( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(var("clockY")), + rhs: Box::new(rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("Clock").into(), + args: vec![rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(0.1), + span: test_span(160, 163), + }], + is_constructor: false, + span: test_span(154, 164), + }), + span: test_span(1, 2), + }, + test_span(1, 2), + "periodic_clock_constructor", + )); + + populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); + assert_eq!(dae_model.clocks.schedules.len(), 1); + assert!( + !dae_model.clocks.intervals.contains_key("b"), + "fallback interval must only apply to implicit sample(..) sources", + ); +} + +#[test] +fn test_runtime_precompute_extracts_shifted_clock_schedule() { + let mut dae_model = dae::Dae::default(); + dae_model + .discrete + .valued_updates + .push(dae::Equation::residual( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new("b").into(), + subscripts: vec![], + span: test_span(1, 2), + }), + rhs: Box::new(rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("shiftSample").into(), + args: vec![ + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("Clock").into(), + args: vec![rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(0.2), + span: test_span(170, 173), + }], + is_constructor: false, + span: test_span(164, 174), + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(1.0), + span: test_span(176, 179), + }, + ], + is_constructor: false, + span: test_span(152, 180), + }), + span: test_span(1, 2), + }, + test_span(1, 2), + "test_shifted_clock_constructor", + )); + + populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); + assert_eq!(dae_model.clocks.constructor_exprs.len(), 2); + assert_eq!(dae_model.clocks.schedules.len(), 2); + + let has_base = dae_model.clocks.schedules.iter().any(|sched| { + (sched.period_seconds - 0.2).abs() <= 1e-12 && sched.phase_seconds.abs() <= 1e-12 + }); + let has_shifted = dae_model.clocks.schedules.iter().any(|sched| { + (sched.period_seconds - 0.2).abs() <= 1e-12 && (sched.phase_seconds - 0.2).abs() <= 1e-12 + }); + assert!(has_base); + assert!(has_shifted); +} + +#[test] +fn test_runtime_precompute_extracts_fractional_shift_sample_schedule() { + let mut dae_model = dae::Dae::default(); + dae_model + .discrete + .valued_updates + .push(dae::Equation::residual( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(var("b")), + rhs: Box::new(rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("shiftSample").into(), + args: vec![ + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("Clock").into(), + args: vec![rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(0.2), + span: test_span(190, 193), + }], + is_constructor: false, + span: test_span(184, 194), + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(1.0), + span: test_span(196, 199), + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(5.0), + span: test_span(201, 204), + }, + ], + is_constructor: false, + span: test_span(180, 205), + }), + span: test_span(1, 2), + }, + test_span(1, 2), + "test_fractional_shifted_clock_constructor", + )); + + populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); + assert!( + dae_model.clocks.schedules.iter().any(|sched| { + (sched.period_seconds - 0.2).abs() <= 1e-12 + && (sched.phase_seconds - 0.04).abs() <= 1e-12 + }), + "expected shiftSample(Clock(0.2), 1, 5) to shift by 1/5 of the base period" + ); +} + +#[test] +fn test_runtime_precompute_extracts_fractional_back_sample_schedule() { + let mut dae_model = dae::Dae::default(); + dae_model + .discrete + .valued_updates + .push(dae::Equation::residual( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(var("b")), + rhs: Box::new(rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("backSample").into(), + args: vec![ + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("shiftSample").into(), + args: vec![ + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("Clock").into(), + args: vec![rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(0.2), + span: test_span(210, 213), + }], + is_constructor: false, + span: test_span(204, 214), + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(2.0), + span: test_span(216, 219), + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(5.0), + span: test_span(221, 224), + }, + ], + is_constructor: false, + span: test_span(192, 225), + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(1.0), + span: test_span(227, 230), + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(5.0), + span: test_span(232, 235), + }, + ], + is_constructor: false, + span: test_span(180, 236), + }), + span: test_span(1, 2), + }, + test_span(1, 2), + "test_fractional_back_sample_clock_constructor", + )); + + populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); + assert!( + dae_model.clocks.schedules.iter().any(|sched| { + (sched.period_seconds - 0.2).abs() <= 1e-12 + && (sched.phase_seconds - 0.04).abs() <= 1e-12 + }), + "expected backSample(shiftSample(Clock(0.2), 2, 5), 1, 5) to land at phase 0.04" + ); +} + +#[test] +fn test_runtime_precompute_records_per_variable_clock_phase() { + let mut dae_model = dae::Dae::default(); + dae_model + .variables + .discrete_valued + .insert(rumoca_core::VarName::new("u"), { + let mut source = dae::Variable::new( + rumoca_core::VarName::new("u"), + rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + ), + ); + source.start = Some(var("u_start")); + source + }); + let mut start = dae::Variable::new( + rumoca_core::VarName::new("u_start"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ); + start.start = Some(lit(1.0)); + dae_model + .variables + .parameters + .insert(rumoca_core::VarName::new("u_start"), start); + dae_model.variables.discrete_valued.insert( + rumoca_core::VarName::new("y"), + dae::Variable::new( + rumoca_core::VarName::new("y"), + rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), + ), + ); + dae_model.discrete.valued_updates.push(dae::Equation { + lhs: Some(rumoca_core::VarName::new("u").into()), + rhs: rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("shiftSample").into(), + args: vec![clock_call(0.02), lit(4.0), lit(3.0)], + is_constructor: false, + span: test_span(1, 2), + }, + span: test_span(1, 2), + origin: "u = shiftSample(Clock(0.02), 4, 3)".to_string(), + scalar_count: 1, + }); + dae_model.discrete.valued_updates.push(dae::Equation { + lhs: Some(rumoca_core::VarName::new("y").into()), + rhs: rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("backSample").into(), + args: vec![var("u"), lit(4.0), lit(3.0)], + is_constructor: false, + span: test_span(1, 2), + }, + span: test_span(1, 2), + origin: "y = backSample(u, 4, 3)".to_string(), + scalar_count: 1, + }); + + populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); + + let u = dae_model + .clocks + .timings + .get("u") + .expect("shifted source timing should be recorded"); + assert!((u.period_seconds - 0.02).abs() <= 1e-12); + assert!((u.phase_seconds - ((4.0 / 3.0) * 0.02)).abs() <= 1e-12); + + let y = dae_model + .clocks + .timings + .get("y") + .expect("back-sampled target timing should be recorded"); + assert!((y.period_seconds - 0.02).abs() <= 1e-12); + assert!(y.phase_seconds.abs() <= 1e-12); + assert!((dae_model.clocks.intervals["y"] - 0.02).abs() <= 1e-12); +} diff --git a/crates/rumoca-phase-dae/src/runtime_precompute/tests/condition_memory_resize.rs b/crates/rumoca-phase-dae/src/runtime_precompute/tests/condition_memory_resize.rs new file mode 100644 index 000000000..7167cf397 --- /dev/null +++ b/crates/rumoca-phase-dae/src/runtime_precompute/tests/condition_memory_resize.rs @@ -0,0 +1,83 @@ +use super::*; + +#[test] +fn test_runtime_precompute_resizes_condition_memory_variables_after_prune() { + let time_only = time_gt(2.5); + let root_cond = rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Gt, + lhs: Box::new(var("x")), + rhs: Box::new(lit(0.0)), + span: test_span(1, 2), + }; + let mut dae_model = dae_with_if_condition(root_cond.clone()); + dae_model.conditions.relations = vec![time_only.clone(), root_cond]; + dae_model.conditions.equations = vec![ + dae::Equation::explicit( + condition_lhs("c", 1), + time_only, + test_span(60, 65), + "condition equation from test", + ), + dae::Equation::explicit( + condition_lhs("c", 2), + dae_model.conditions.relations[1].clone(), + test_span(60, 65), + "condition equation from test", + ), + ]; + let mut condition = dae::Variable::new(rumoca_core::VarName::new("c"), test_span(60, 65)); + condition.dims = vec![2]; + condition.start = Some(rumoca_core::Expression::Array { + elements: vec![ + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Boolean(false), + span: test_span(60, 65), + }; + 2 + ], + is_matrix: false, + span: test_span(60, 65), + }); + let mut pre_condition = + dae::Variable::new(rumoca_core::VarName::new("__pre__.c"), test_span(60, 65)); + pre_condition.dims = vec![2]; + pre_condition.start = condition.start.clone(); + dae_model + .variables + .discrete_valued + .insert(rumoca_core::VarName::new("c"), condition); + dae_model + .variables + .parameters + .insert(rumoca_core::VarName::new("__pre__.c"), pre_condition); + + populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); + + assert_eq!( + dae_model + .variables + .discrete_valued + .get(&rumoca_core::VarName::new("c")) + .map(|variable| variable.dims.as_slice()), + Some(&[1][..]) + ); + assert_eq!( + dae_model + .variables + .parameters + .get(&rumoca_core::VarName::new("__pre__.c")) + .map(|variable| variable.dims.as_slice()), + Some(&[1][..]) + ); + assert!( + matches!( + &dae_model + .variables + .discrete_valued + .get(&rumoca_core::VarName::new("c")) + .and_then(|variable| variable.start.as_ref()), + Some(rumoca_core::Expression::Array { elements, .. }) if elements.len() == 1 + ), + "condition memory start vector must match the pruned relation count" + ); +} diff --git a/crates/rumoca-phase-dae/src/runtime_precompute/tests/mod.rs b/crates/rumoca-phase-dae/src/runtime_precompute/tests/mod.rs index a37e8e5ed..443e76108 100644 --- a/crates/rumoca-phase-dae/src/runtime_precompute/tests/mod.rs +++ b/crates/rumoca-phase-dae/src/runtime_precompute/tests/mod.rs @@ -4,6 +4,8 @@ use super::*; mod clock_alias_resolution_tests; mod clock_alias_tests; +mod clock_schedule_tests; +mod condition_memory_resize; mod dynamic_clock_tests; fn populate_conditions(dae_model: &mut dae::Dae) { @@ -15,6 +17,15 @@ fn test_span(start: usize, end: usize) -> Span { Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), start, end) } +fn scheduled_times(dae_model: &dae::Dae) -> Vec { + dae_model + .events + .scheduled_time_events + .iter() + .map(|event| event.time) + .collect() +} + fn time_gt(value: f64) -> rumoca_core::Expression { rumoca_core::Expression::Binary { op: rumoca_core::OpBinary::Gt, @@ -286,7 +297,7 @@ fn test_runtime_precompute_suppresses_initial_only_time_events() { populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); assert_eq!( - dae_model.events.scheduled_time_events, + scheduled_times(&dae_model), vec![0.5], "time events reachable only during initialization must not be scheduled for simulation" ); @@ -400,7 +411,7 @@ fn test_runtime_precompute_collects_event_without_synthetic_root_for_time_if_con .events .scheduled_time_events .iter() - .any(|time| (*time - 5.0).abs() <= 1.0e-12), + .any(|event| (event.time - 5.0).abs() <= 1.0e-12), "expected precompute to capture scheduled event at t=5" ); } @@ -524,7 +535,7 @@ fn test_runtime_precompute_interns_and_orders_root_and_time_event_metadata() { "Appendix B relations must not be duplicated as synthetic roots" ); assert_eq!( - dae_model.events.scheduled_time_events, + scheduled_times(&dae_model), vec![0.5, 1.5], "scheduled time events should be canonicalized and sorted" ); @@ -618,6 +629,99 @@ fn test_runtime_precompute_skips_noevent_wrapped_conditions_for_events() { ); } +#[test] +fn test_runtime_precompute_skips_smooth_wrapped_conditions_for_synthetic_roots() { + let mut dae_model = dae::Dae::default(); + let cond = rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Ge, + lhs: Box::new(var("w")), + rhs: Box::new(lit(0.0)), + span: test_span(10, 18), + }; + dae_model.continuous.equations.push(dae::Equation::residual( + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Smooth, + args: vec![ + lit(1.0), + rumoca_core::Expression::If { + branches: vec![( + cond, + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs: Box::new(var("w")), + rhs: Box::new(var("w")), + span: test_span(20, 25), + }, + )], + else_branch: Box::new(rumoca_core::Expression::Unary { + op: rumoca_core::OpUnary::Minus, + rhs: Box::new(rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs: Box::new(var("w")), + rhs: Box::new(var("w")), + span: test_span(30, 35), + }), + span: test_span(29, 35), + }), + span: test_span(10, 35), + }, + ], + span: test_span(1, 36), + }, + test_span(1, 36), + "smooth directional loss", + )); + + populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); + + assert!( + dae_model.events.synthetic_root_conditions.is_empty(), + "smooth-wrapped relations should not become synthetic zero-crossing roots" + ); +} + +#[test] +fn test_runtime_precompute_suppresses_branch_roots_guarded_by_noevent_condition() { + let mut dae_model = dae::Dae::default(); + let guard = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::NoEvent, + args: vec![rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Gt, + lhs: Box::new(rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Abs, + args: vec![var("w")], + span: test_span(10, 16), + }), + rhs: Box::new(var("wLinear")), + span: test_span(10, 26), + }], + span: test_span(2, 27), + }; + dae_model.continuous.equations.push(dae::Equation::residual( + rumoca_core::Expression::If { + branches: vec![( + guard, + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Sign, + args: vec![var("w")], + span: test_span(30, 37), + }, + )], + else_branch: Box::new(var("w")), + span: test_span(1, 40), + }, + test_span(1, 40), + "noevent guarded sign", + )); + + populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); + + assert!( + dae_model.events.synthetic_root_conditions.is_empty(), + "roots inside branches guarded by noEvent conditions should not fire at unreachable surfaces" + ); +} + #[test] fn test_runtime_precompute_skips_time_vs_parameter_synthetic_roots() { let cond = rumoca_core::Expression::Binary { @@ -654,7 +758,7 @@ fn test_runtime_precompute_skips_time_vs_parameter_synthetic_roots() { .all(|expr| format!("{expr:?}") != format!("{cond:?}")), "time-vs-parameter branch conditions should be scheduled time events, not synthetic roots" ); - assert_eq!(dae_model.events.scheduled_time_events, vec![2.5]); + assert_eq!(scheduled_times(&dae_model), vec![2.5]); assert!( dae_model.conditions.relations.is_empty(), "time-vs-parameter conditions should be represented as scheduled events, not solver roots" @@ -917,1079 +1021,3 @@ fn test_runtime_precompute_keeps_time_vs_state_synthetic_roots() { "time-vs-state branch conditions should not be scheduled as static events" ); } - -#[test] -fn test_runtime_precompute_extracts_affine_time_event() { - let cond = rumoca_core::Expression::Binary { - op: rumoca_core::OpBinary::Le, - lhs: Box::new(rumoca_core::Expression::Binary { - op: rumoca_core::OpBinary::Add, - lhs: Box::new(var("time")), - rhs: Box::new(var("delay")), - span: test_span(1, 2), - }), - rhs: Box::new(var("switch_at")), - span: test_span(1, 2), - }; - let mut dae_model = dae_with_if_condition(cond); - let mut delay = dae::Variable::new( - rumoca_core::VarName::new("delay"), - rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), - ); - delay.start = Some(lit(0.25)); - dae_model - .variables - .parameters - .insert(rumoca_core::VarName::new("delay"), delay); - let mut switch_at = dae::Variable::new( - rumoca_core::VarName::new("switch_at"), - rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), - ); - switch_at.start = Some(lit(1.5)); - dae_model - .variables - .parameters - .insert(rumoca_core::VarName::new("switch_at"), switch_at); - - populate_conditions(&mut dae_model); - populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); - assert_eq!(dae_model.events.scheduled_time_events.len(), 1); - assert!((dae_model.events.scheduled_time_events[0] - 1.25).abs() <= 1e-12); -} - -#[test] -fn test_runtime_precompute_extracts_discrete_partition_events() { - let mut dae_model = dae::Dae::default(); - dae_model.variables.discrete_valued.insert( - rumoca_core::VarName::new("c"), - dae::Variable::new( - rumoca_core::VarName::new("c"), - rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), - ), - ); - dae_model - .discrete - .valued_updates - .push(dae::Equation::residual( - sub( - var("c"), - rumoca_core::Expression::Binary { - op: rumoca_core::OpBinary::Gt, - lhs: Box::new(var("time")), - rhs: Box::new(lit(0.5)), - span: test_span(1, 2), - }, - ), - test_span(1, 2), - "test_discrete_partition", - )); - - populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); - assert_eq!(dae_model.events.scheduled_time_events, vec![0.5]); -} - -#[test] -fn test_runtime_precompute_collects_clock_constructor_exprs() { - let mut dae_model = dae::Dae::default(); - dae_model.variables.discrete_reals.insert( - rumoca_core::VarName::new("s"), - dae::Variable::new( - rumoca_core::VarName::new("s"), - rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), - ), - ); - dae_model - .discrete - .real_updates - .push(dae::Equation::residual( - rumoca_core::Expression::Binary { - op: rumoca_core::OpBinary::Sub, - lhs: Box::new(rumoca_core::Expression::VarRef { - name: rumoca_core::VarName::new("s").into(), - subscripts: vec![], - span: test_span(1, 2), - }), - rhs: Box::new(rumoca_core::Expression::BuiltinCall { - function: rumoca_core::BuiltinFunction::Sample, - args: vec![ - rumoca_core::Expression::VarRef { - name: rumoca_core::VarName::new("u").into(), - subscripts: vec![], - span: test_span(100, 101), - }, - rumoca_core::Expression::FunctionCall { - name: rumoca_core::VarName::new("Clock").into(), - args: vec![rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(0.1), - span: test_span(108, 111), - }], - is_constructor: false, - span: test_span(102, 112), - }, - ], - span: test_span(95, 113), - }), - span: test_span(1, 2), - }, - test_span(1, 2), - "test_clock_constructor", - )); - - populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); - - assert_eq!(dae_model.clocks.constructor_exprs.len(), 1); - assert_eq!(dae_model.clocks.schedules.len(), 1); - assert!((dae_model.clocks.schedules[0].period_seconds - 0.1).abs() <= 1e-12); - assert!(dae_model.clocks.schedules[0].phase_seconds.abs() <= 1e-12); - assert!((dae_model.clocks.intervals["s"] - 0.1).abs() <= 1e-12); -} - -#[test] -fn test_runtime_precompute_rejects_static_clock_constructor_without_source_provenance() { - let mut dae_model = dae::Dae::default(); - dae_model.variables.discrete_reals.insert( - rumoca_core::VarName::new("s"), - dae::Variable::new( - rumoca_core::VarName::new("s"), - rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), - ), - ); - dae_model - .discrete - .real_updates - .push(dae::Equation::residual( - sub( - var("s"), - rumoca_core::Expression::FunctionCall { - name: rumoca_core::VarName::new("Clock").into(), - args: vec![rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(0.1), - span: rumoca_core::Span::DUMMY, - }], - is_constructor: false, - span: rumoca_core::Span::DUMMY, - }, - ), - Span::DUMMY, - "unspanned_static_clock_constructor", - )); - - let err = populate_runtime_precompute(&mut dae_model) - .expect_err("source-free static clock constructors must fail fast"); - assert!(matches!( - err, - crate::ToDaeError::RuntimeMetadataViolation { detail } - if detail.contains("source provenance") - )); -} - -#[test] -fn test_runtime_precompute_collects_sample_start_interval_schedule() { - let mut dae_model = dae::Dae::default(); - dae_model.variables.discrete_reals.insert( - rumoca_core::VarName::new("s"), - dae::Variable::new( - rumoca_core::VarName::new("s"), - rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), - ), - ); - dae_model - .discrete - .real_updates - .push(dae::Equation::residual( - sub( - var("s"), - rumoca_core::Expression::BuiltinCall { - function: rumoca_core::BuiltinFunction::Sample, - args: vec![ - rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(0.2), - span: test_span(120, 123), - }, - rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(0.1), - span: test_span(125, 128), - }, - ], - span: test_span(113, 129), - }, - ), - test_span(1, 2), - // MLS §16.5.1: sample(start, interval) defines a periodic event. - "periodic_sample_start_interval", - )); - - populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); - - assert_eq!(dae_model.clocks.schedules.len(), 1); - assert!((dae_model.clocks.schedules[0].period_seconds - 0.1).abs() <= 1e-12); - assert!((dae_model.clocks.schedules[0].phase_seconds - 0.2).abs() <= 1e-12); -} - -#[test] -fn test_runtime_precompute_marks_schedule_backed_sample_root_condition() { - let mut dae_model = dae::Dae::default(); - let relation = rumoca_core::Expression::FunctionCall { - name: rumoca_core::VarName::new(rumoca_core::INTERNAL_SAMPLE_FUNCTION_NAME).into(), - args: vec![lit(42.0), lit(0.05), lit(0.1)], - is_constructor: false, - span: test_span(120, 149), - }; - dae_model.conditions.relations.push(relation.clone()); - dae_model.conditions.equations.push(dae::Equation::explicit( - condition_lhs("c", 1), - relation, - test_span(1, 2), - "scheduled sample condition memory", - )); - - populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); - - assert_eq!(dae_model.events.scheduled_root_conditions.len(), 1); - let root = &dae_model.events.scheduled_root_conditions[0]; - assert_eq!(root.root_index, 0); - assert!((root.period_seconds - 0.1).abs() <= 1e-12); - assert!((root.phase_seconds - 0.05).abs() <= 1e-12); -} - -#[test] -fn test_runtime_precompute_collects_sample_schedule_from_initial_time_parameter() { - let mut dae_model = dae::Dae::default(); - let mut frequency = dae::Variable::new( - rumoca_core::VarName::new("mean.f"), - rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), - ); - frequency.start = Some(lit(150.0)); - dae_model - .variables - .parameters - .insert(rumoca_core::VarName::new("mean.f"), frequency); - dae_model.variables.parameters.insert( - rumoca_core::VarName::new("mean.t0"), - dae::Variable::new( - rumoca_core::VarName::new("mean.t0"), - rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), - ), - ); - dae_model.variables.discrete_reals.insert( - rumoca_core::VarName::new("s"), - dae::Variable::new( - rumoca_core::VarName::new("s"), - rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), - ), - ); - dae_model - .initialization - .equations - .push(dae::Equation::residual( - sub(var("mean.t0"), var("time")), - test_span(1, 2), - "initial t0 = time", - )); - - let interval = rumoca_core::Expression::Binary { - op: rumoca_core::OpBinary::Div, - lhs: Box::new(lit(1.0)), - rhs: Box::new(var("mean.f")), - span: test_span(140, 148), - }; - dae_model - .discrete - .real_updates - .push(dae::Equation::residual( - sub( - var("s"), - rumoca_core::Expression::FunctionCall { - name: rumoca_core::VarName::new(rumoca_core::INTERNAL_SAMPLE_FUNCTION_NAME) - .into(), - args: vec![ - rumoca_core::Expression::Binary { - op: rumoca_core::OpBinary::Add, - lhs: Box::new(var("mean.t0")), - rhs: Box::new(interval.clone()), - span: test_span(130, 148), - }, - interval, - ], - is_constructor: false, - span: test_span(120, 149), - }, - ), - test_span(1, 2), - "periodic sample with initial-time origin", - )); - - populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); - - assert_eq!(dae_model.clocks.schedules.len(), 1); - assert!((dae_model.clocks.schedules[0].period_seconds - 1.0 / 150.0).abs() <= 1e-12); - assert!((dae_model.clocks.schedules[0].phase_seconds - 1.0 / 150.0).abs() <= 1e-12); -} - -#[test] -fn test_runtime_precompute_assigns_implicit_sample_interval_from_unique_schedule() { - let mut dae_model = dae::Dae::default(); - dae_model.variables.discrete_reals.insert( - rumoca_core::VarName::new("simTime"), - dae::Variable::new( - rumoca_core::VarName::new("simTime"), - rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), - ), - ); - dae_model.variables.discrete_reals.insert( - rumoca_core::VarName::new("clockY"), - dae::Variable::new( - rumoca_core::VarName::new("clockY"), - rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), - ), - ); - - // simTime = sample(time) (implicit clock sample form) - dae_model - .discrete - .real_updates - .push(dae::Equation::residual( - sub( - var("simTime"), - rumoca_core::Expression::BuiltinCall { - function: rumoca_core::BuiltinFunction::Sample, - args: vec![var("time")], - span: test_span(1, 2), - }, - ), - test_span(1, 2), - "implicit_clocked_sample", - )); - - dae_model - .discrete - .real_updates - .push(dae::Equation::residual( - sub( - var("clockY"), - rumoca_core::Expression::FunctionCall { - name: rumoca_core::VarName::new("Clock").into(), - args: vec![rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(0.1), - span: test_span(140, 143), - }], - is_constructor: false, - span: test_span(134, 144), - }, - ), - test_span(1, 2), - "periodic_clock_constructor", - )); - - populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); - assert_eq!(dae_model.clocks.schedules.len(), 1); - assert!((dae_model.clocks.intervals["simTime"] - 0.1).abs() <= 1e-12); -} - -#[test] -fn test_runtime_precompute_propagates_no_argument_clock_guard_timing() { - let mut dae_model = dae::Dae::default(); - for name in ["u", "dummy", "b"] { - dae_model.variables.discrete_valued.insert( - rumoca_core::VarName::new(name), - dae::Variable::new( - rumoca_core::VarName::new(name), - rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name(file!()), - 1, - 2, - ), - ), - ); - } - let clock_span = test_span(1_000, 1_005); - dae_model.discrete.valued_updates.push(dae::Equation { - lhs: Some(rumoca_core::VarName::new("u").into()), - rhs: clock_call(0.02), - span: test_span(1, 2), - origin: "u = Clock(0.02)".to_string(), - scalar_count: 1, - }); - dae_model.discrete.valued_updates.push(dae::Equation { - lhs: Some(rumoca_core::VarName::new("dummy").into()), - rhs: if_then_else( - no_argument_clock_call(clock_span), - var("u"), - var("__pre__.dummy"), - ), - span: test_span(1, 2), - origin: "when Clock() then dummy = u".to_string(), - scalar_count: 1, - }); - dae_model.discrete.valued_updates.push(dae::Equation { - lhs: Some(rumoca_core::VarName::new("b").into()), - rhs: if_then_else( - no_argument_clock_call(clock_span), - rumoca_core::Expression::Unary { - op: rumoca_core::OpUnary::Not, - rhs: Box::new(var("__pre__.__pre__.b")), - span: test_span(1, 2), - }, - var("__pre__.b"), - ), - span: test_span(1, 2), - origin: "when Clock() then b = not previous(b)".to_string(), - scalar_count: 1, - }); - - populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); - - assert!((dae_model.clocks.intervals["dummy"] - 0.02).abs() <= 1e-12); - assert!((dae_model.clocks.intervals["b"] - 0.02).abs() <= 1e-12); -} - -#[test] -fn test_runtime_precompute_assigns_clock_interval_to_algebraic_alias_chain() { - let mut dae_model = dae::Dae::default(); - dae_model.variables.inputs.insert( - rumoca_core::VarName::new("u"), - dae::Variable::new( - rumoca_core::VarName::new("u"), - rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), - ), - ); - dae_model.variables.algebraics.insert( - rumoca_core::VarName::new("feedback.y"), - dae::Variable::new( - rumoca_core::VarName::new("feedback.y"), - rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), - ), - ); - dae_model.variables.algebraics.insert( - rumoca_core::VarName::new("PI.u"), - dae::Variable::new( - rumoca_core::VarName::new("PI.u"), - rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), - ), - ); - dae_model.variables.discrete_reals.insert( - rumoca_core::VarName::new("sample2.y"), - dae::Variable::new( - rumoca_core::VarName::new("sample2.y"), - rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), - ), - ); - dae_model.variables.discrete_valued.insert( - rumoca_core::VarName::new("sample2.clock"), - dae::Variable::new( - rumoca_core::VarName::new("sample2.clock"), - rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), - ), - ); - - dae_model - .discrete - .real_updates - .push(dae::Equation::explicit( - rumoca_core::VarName::new("sample2.clock"), - rumoca_core::Expression::FunctionCall { - name: rumoca_core::VarName::new("Clock").into(), - args: vec![rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(0.1), - span: test_span(150, 153), - }], - is_constructor: false, - span: test_span(144, 154), - }, - test_span(1, 2), - "explicit_clock_alias", - )); - dae_model - .discrete - .real_updates - .push(dae::Equation::explicit( - rumoca_core::VarName::new("sample2.y"), - rumoca_core::Expression::BuiltinCall { - function: rumoca_core::BuiltinFunction::Sample, - args: vec![var("u"), var("sample2.clock")], - span: test_span(1, 2), - }, - test_span(1, 2), - "explicit_sample_value", - )); - dae_model.continuous.equations.push(dae::Equation::residual( - rumoca_core::Expression::Binary { - op: rumoca_core::OpBinary::Sub, - lhs: Box::new(var("sample2.y")), - rhs: Box::new(var("feedback.y")), - span: test_span(1, 2), - }, - test_span(1, 2), - "sample_alias", - )); - dae_model.continuous.equations.push(dae::Equation::residual( - rumoca_core::Expression::Binary { - op: rumoca_core::OpBinary::Sub, - lhs: Box::new(var("feedback.y")), - rhs: Box::new(var("PI.u")), - span: test_span(1, 2), - }, - test_span(1, 2), - "controller_input_alias", - )); - populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); - assert!((dae_model.clocks.intervals["sample2.clock"] - 0.1).abs() <= 1e-12); - assert!((dae_model.clocks.intervals["sample2.y"] - 0.1).abs() <= 1e-12); - assert!((dae_model.clocks.intervals["feedback.y"] - 0.1).abs() <= 1e-12); - assert!((dae_model.clocks.intervals["PI.u"] - 0.1).abs() <= 1e-12); -} - -#[test] -fn test_runtime_precompute_propagates_clock_across_indexed_vector_and_previous() { - let mut dae_model = dae::Dae::default(); - let mut sampled = dae::Variable::new(rumoca_core::VarName::new("sampled"), test_span(1, 2)); - sampled.dims = vec![2]; - dae_model - .variables - .discrete_reals - .insert(sampled.name.clone(), sampled); - for name in ["delay1.u", "delay1.y", "delay2.u", "delay2.y"] { - dae_model.variables.discrete_reals.insert( - rumoca_core::VarName::new(name), - dae::Variable::new(rumoca_core::VarName::new(name), test_span(1, 2)), - ); - } - dae_model.variables.discrete_valued.insert( - rumoca_core::VarName::new("clock"), - dae::Variable::new(rumoca_core::VarName::new("clock"), test_span(1, 2)), - ); - - dae_model - .discrete - .valued_updates - .push(dae::Equation::explicit( - rumoca_core::VarName::new("clock"), - clock_call(0.02), - test_span(1, 2), - "clock_source", - )); - dae_model - .discrete - .real_updates - .push(dae::Equation::explicit( - rumoca_core::VarName::new("sampled"), - rumoca_core::Expression::BuiltinCall { - function: rumoca_core::BuiltinFunction::Sample, - args: vec![var("u"), var("clock")], - span: test_span(1, 2), - }, - test_span(1, 2), - "sampled_vector", - )); - for (delay, index) in [("delay1", 1), ("delay2", 2)] { - dae_model - .discrete - .real_updates - .push(dae::Equation::explicit( - rumoca_core::VarName::new(format!("{delay}.u")), - condition_memory_ref("sampled", index), - test_span(1, 2), - "indexed_vector_alias", - )); - dae_model - .discrete - .real_updates - .push(dae::Equation::explicit( - rumoca_core::VarName::new(format!("{delay}.y")), - var(&format!("__pre__.{delay}.u")), - test_span(1, 2), - "clocked_previous_value", - )); - } - - populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); - - for name in ["sampled", "delay1.u", "delay1.y", "delay2.u", "delay2.y"] { - assert!((dae_model.clocks.intervals[name] - 0.02).abs() <= 1e-12); - } -} - -#[test] -fn test_runtime_precompute_propagates_uniform_clock_through_vector_alias_projection() { - let mut dae_model = dae::Dae::default(); - for name in ["sampled", "alias"] { - let mut variable = dae::Variable::new(rumoca_core::VarName::new(name), test_span(1, 2)); - variable.dims = vec![2]; - dae_model - .variables - .discrete_reals - .insert(variable.name.clone(), variable); - } - for name in ["delay.u", "delay.y"] { - dae_model.variables.discrete_reals.insert( - rumoca_core::VarName::new(name), - dae::Variable::new(rumoca_core::VarName::new(name), test_span(1, 2)), - ); - } - - dae_model - .discrete - .real_updates - .push(dae::Equation::explicit( - rumoca_core::VarName::new("sampled"), - rumoca_core::Expression::BuiltinCall { - function: rumoca_core::BuiltinFunction::Sample, - args: vec![var("u"), clock_call(0.02)], - span: test_span(1, 2), - }, - test_span(1, 2), - "direct_clocked_vector", - )); - dae_model.continuous.equations.push(dae::Equation::explicit( - rumoca_core::VarName::new("alias"), - var("sampled"), - test_span(1, 2), - "untimed_vector_alias", - )); - dae_model - .discrete - .real_updates - .push(dae::Equation::explicit( - rumoca_core::VarName::new("delay.u"), - condition_memory_ref("alias", 1), - test_span(1, 2), - "indexed_alias_consumer", - )); - dae_model - .discrete - .real_updates - .push(dae::Equation::explicit( - rumoca_core::VarName::new("delay.y"), - var("__pre__.delay.u"), - test_span(1, 2), - "previous_consumer", - )); - - populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); - - for name in ["sampled", "alias", "delay.u", "delay.y"] { - assert!((dae_model.clocks.intervals[name] - 0.02).abs() <= 1e-12); - } -} - -#[test] -fn test_runtime_precompute_keeps_distinct_clocks_for_array_elements() { - let mut dae_model = dae::Dae::default(); - let mut sampled = dae::Variable::new(rumoca_core::VarName::new("sampled"), test_span(1, 2)); - sampled.dims = vec![2]; - dae_model - .variables - .discrete_reals - .insert(sampled.name.clone(), sampled); - for name in ["clock1", "clock2"] { - dae_model.variables.discrete_valued.insert( - rumoca_core::VarName::new(name), - dae::Variable::new(rumoca_core::VarName::new(name), test_span(1, 2)), - ); - } - for (name, period) in [("clock1", 0.1), ("clock2", 0.2)] { - dae_model - .discrete - .valued_updates - .push(dae::Equation::explicit( - rumoca_core::VarName::new(name), - clock_call(period), - test_span(1, 2), - "independent_clock_source", - )); - } - for (index, clock) in [(1, "clock1"), (2, "clock2")] { - dae_model.discrete.real_updates.push(dae::Equation { - lhs: Some(condition_lhs("sampled", index)), - rhs: rumoca_core::Expression::BuiltinCall { - function: rumoca_core::BuiltinFunction::Sample, - args: vec![var("u"), var(clock)], - span: test_span(1, 2), - }, - span: test_span(1, 2), - origin: "independently_clocked_array_element".to_string(), - scalar_count: 1, - }); - } - - populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); - - assert!(!dae_model.clocks.intervals.contains_key("sampled")); - assert!((dae_model.clocks.intervals["sampled[1]"] - 0.1).abs() <= 1e-12); - assert!((dae_model.clocks.intervals["sampled[2]"] - 0.2).abs() <= 1e-12); -} - -#[test] -fn test_runtime_precompute_promotes_equal_element_clocks_to_uniform_array() { - let mut dae_model = dae::Dae::default(); - for name in ["sampled", "alias"] { - let mut variable = dae::Variable::new(rumoca_core::VarName::new(name), test_span(1, 2)); - variable.dims = vec![2]; - dae_model - .variables - .discrete_reals - .insert(variable.name.clone(), variable); - } - for name in ["clock1", "clock2"] { - dae_model.variables.discrete_valued.insert( - rumoca_core::VarName::new(name), - dae::Variable::new(rumoca_core::VarName::new(name), test_span(1, 2)), - ); - dae_model - .discrete - .valued_updates - .push(dae::Equation::explicit( - rumoca_core::VarName::new(name), - clock_call(0.1), - test_span(1, 2), - "uniform_clock_source", - )); - } - for (index, clock) in [(1, "clock1"), (2, "clock2")] { - dae_model.discrete.real_updates.push(dae::Equation { - lhs: Some(condition_lhs("sampled", index)), - rhs: rumoca_core::Expression::BuiltinCall { - function: rumoca_core::BuiltinFunction::Sample, - args: vec![var("u"), var(clock)], - span: test_span(1, 2), - }, - span: test_span(1, 2), - origin: "uniformly_clocked_array_element".to_string(), - scalar_count: 1, - }); - } - dae_model.continuous.equations.push(dae::Equation::explicit( - rumoca_core::VarName::new("alias"), - var("sampled"), - test_span(1, 2), - "whole_array_alias", - )); - - populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); - - for name in ["sampled", "alias"] { - let interval = dae_model.clocks.intervals.get(name).unwrap_or_else(|| { - panic!( - "missing {name}; intervals={:?}", - dae_model.clocks.intervals.keys().collect::>() - ) - }); - assert!((*interval - 0.1).abs() <= 1e-12); - } -} - -#[test] -fn test_runtime_precompute_does_not_assign_fallback_interval_for_non_sample_clock_ops() { - let mut dae_model = dae::Dae::default(); - dae_model.variables.discrete_valued.insert( - rumoca_core::VarName::new("b"), - dae::Variable::new( - rumoca_core::VarName::new("b"), - rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), - ), - ); - dae_model.variables.discrete_reals.insert( - rumoca_core::VarName::new("clockY"), - dae::Variable::new( - rumoca_core::VarName::new("clockY"), - rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), - ), - ); - - // b = pre(b) is discrete/event logic, not an implicit sample(..) form. - dae_model - .discrete - .valued_updates - .push(dae::Equation::residual( - rumoca_core::Expression::Binary { - op: rumoca_core::OpBinary::Sub, - lhs: Box::new(var("b")), - rhs: Box::new(rumoca_core::Expression::BuiltinCall { - function: rumoca_core::BuiltinFunction::Pre, - args: vec![var("b")], - span: test_span(1, 2), - }), - span: test_span(1, 2), - }, - test_span(1, 2), - "pre_based_discrete_update", - )); - - // Add one static periodic schedule in the model. - dae_model - .discrete - .real_updates - .push(dae::Equation::residual( - rumoca_core::Expression::Binary { - op: rumoca_core::OpBinary::Sub, - lhs: Box::new(var("clockY")), - rhs: Box::new(rumoca_core::Expression::FunctionCall { - name: rumoca_core::VarName::new("Clock").into(), - args: vec![rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(0.1), - span: test_span(160, 163), - }], - is_constructor: false, - span: test_span(154, 164), - }), - span: test_span(1, 2), - }, - test_span(1, 2), - "periodic_clock_constructor", - )); - - populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); - assert_eq!(dae_model.clocks.schedules.len(), 1); - assert!( - !dae_model.clocks.intervals.contains_key("b"), - "fallback interval must only apply to implicit sample(..) sources", - ); -} - -#[test] -fn test_runtime_precompute_extracts_shifted_clock_schedule() { - let mut dae_model = dae::Dae::default(); - dae_model - .discrete - .valued_updates - .push(dae::Equation::residual( - rumoca_core::Expression::Binary { - op: rumoca_core::OpBinary::Sub, - lhs: Box::new(rumoca_core::Expression::VarRef { - name: rumoca_core::VarName::new("b").into(), - subscripts: vec![], - span: test_span(1, 2), - }), - rhs: Box::new(rumoca_core::Expression::FunctionCall { - name: rumoca_core::VarName::new("shiftSample").into(), - args: vec![ - rumoca_core::Expression::FunctionCall { - name: rumoca_core::VarName::new("Clock").into(), - args: vec![rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(0.2), - span: test_span(170, 173), - }], - is_constructor: false, - span: test_span(164, 174), - }, - rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(1.0), - span: test_span(176, 179), - }, - ], - is_constructor: false, - span: test_span(152, 180), - }), - span: test_span(1, 2), - }, - test_span(1, 2), - "test_shifted_clock_constructor", - )); - - populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); - assert_eq!(dae_model.clocks.constructor_exprs.len(), 2); - assert_eq!(dae_model.clocks.schedules.len(), 2); - - let has_base = dae_model.clocks.schedules.iter().any(|sched| { - (sched.period_seconds - 0.2).abs() <= 1e-12 && sched.phase_seconds.abs() <= 1e-12 - }); - let has_shifted = dae_model.clocks.schedules.iter().any(|sched| { - (sched.period_seconds - 0.2).abs() <= 1e-12 && (sched.phase_seconds - 0.2).abs() <= 1e-12 - }); - assert!(has_base); - assert!(has_shifted); -} - -#[test] -fn test_runtime_precompute_extracts_fractional_shift_sample_schedule() { - let mut dae_model = dae::Dae::default(); - dae_model - .discrete - .valued_updates - .push(dae::Equation::residual( - rumoca_core::Expression::Binary { - op: rumoca_core::OpBinary::Sub, - lhs: Box::new(var("b")), - rhs: Box::new(rumoca_core::Expression::FunctionCall { - name: rumoca_core::VarName::new("shiftSample").into(), - args: vec![ - rumoca_core::Expression::FunctionCall { - name: rumoca_core::VarName::new("Clock").into(), - args: vec![rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(0.2), - span: test_span(190, 193), - }], - is_constructor: false, - span: test_span(184, 194), - }, - rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(1.0), - span: test_span(196, 199), - }, - rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(5.0), - span: test_span(201, 204), - }, - ], - is_constructor: false, - span: test_span(180, 205), - }), - span: test_span(1, 2), - }, - test_span(1, 2), - "test_fractional_shifted_clock_constructor", - )); - - populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); - assert!( - dae_model.clocks.schedules.iter().any(|sched| { - (sched.period_seconds - 0.2).abs() <= 1e-12 - && (sched.phase_seconds - 0.04).abs() <= 1e-12 - }), - "expected shiftSample(Clock(0.2), 1, 5) to shift by 1/5 of the base period" - ); -} - -#[test] -fn test_runtime_precompute_extracts_fractional_back_sample_schedule() { - let mut dae_model = dae::Dae::default(); - dae_model - .discrete - .valued_updates - .push(dae::Equation::residual( - rumoca_core::Expression::Binary { - op: rumoca_core::OpBinary::Sub, - lhs: Box::new(var("b")), - rhs: Box::new(rumoca_core::Expression::FunctionCall { - name: rumoca_core::VarName::new("backSample").into(), - args: vec![ - rumoca_core::Expression::FunctionCall { - name: rumoca_core::VarName::new("shiftSample").into(), - args: vec![ - rumoca_core::Expression::FunctionCall { - name: rumoca_core::VarName::new("Clock").into(), - args: vec![rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(0.2), - span: test_span(210, 213), - }], - is_constructor: false, - span: test_span(204, 214), - }, - rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(2.0), - span: test_span(216, 219), - }, - rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(5.0), - span: test_span(221, 224), - }, - ], - is_constructor: false, - span: test_span(192, 225), - }, - rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(1.0), - span: test_span(227, 230), - }, - rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(5.0), - span: test_span(232, 235), - }, - ], - is_constructor: false, - span: test_span(180, 236), - }), - span: test_span(1, 2), - }, - test_span(1, 2), - "test_fractional_back_sample_clock_constructor", - )); - - populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); - assert!( - dae_model.clocks.schedules.iter().any(|sched| { - (sched.period_seconds - 0.2).abs() <= 1e-12 - && (sched.phase_seconds - 0.04).abs() <= 1e-12 - }), - "expected backSample(shiftSample(Clock(0.2), 2, 5), 1, 5) to land at phase 0.04" - ); -} - -#[test] -fn test_runtime_precompute_records_per_variable_clock_phase() { - let mut dae_model = dae::Dae::default(); - dae_model - .variables - .discrete_valued - .insert(rumoca_core::VarName::new("u"), { - let mut source = dae::Variable::new( - rumoca_core::VarName::new("u"), - rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name(file!()), - 1, - 2, - ), - ); - source.start = Some(var("u_start")); - source - }); - let mut start = dae::Variable::new( - rumoca_core::VarName::new("u_start"), - rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), - ); - start.start = Some(lit(1.0)); - dae_model - .variables - .parameters - .insert(rumoca_core::VarName::new("u_start"), start); - dae_model.variables.discrete_valued.insert( - rumoca_core::VarName::new("y"), - dae::Variable::new( - rumoca_core::VarName::new("y"), - rumoca_core::Span::from_offsets(rumoca_core::SourceId::from_source_name(file!()), 1, 2), - ), - ); - dae_model.discrete.valued_updates.push(dae::Equation { - lhs: Some(rumoca_core::VarName::new("u").into()), - rhs: rumoca_core::Expression::FunctionCall { - name: rumoca_core::VarName::new("shiftSample").into(), - args: vec![clock_call(0.02), lit(4.0), lit(3.0)], - is_constructor: false, - span: test_span(1, 2), - }, - span: test_span(1, 2), - origin: "u = shiftSample(Clock(0.02), 4, 3)".to_string(), - scalar_count: 1, - }); - dae_model.discrete.valued_updates.push(dae::Equation { - lhs: Some(rumoca_core::VarName::new("y").into()), - rhs: rumoca_core::Expression::FunctionCall { - name: rumoca_core::VarName::new("backSample").into(), - args: vec![var("u"), lit(4.0), lit(3.0)], - is_constructor: false, - span: test_span(1, 2), - }, - span: test_span(1, 2), - origin: "y = backSample(u, 4, 3)".to_string(), - scalar_count: 1, - }); - - populate_runtime_precompute(&mut dae_model).expect("runtime precompute should succeed"); - - let u = dae_model - .clocks - .timings - .get("u") - .expect("shifted source timing should be recorded"); - assert!((u.period_seconds - 0.02).abs() <= 1e-12); - assert!((u.phase_seconds - ((4.0 / 3.0) * 0.02)).abs() <= 1e-12); - - let y = dae_model - .clocks - .timings - .get("y") - .expect("back-sampled target timing should be recorded"); - assert!((y.period_seconds - 0.02).abs() <= 1e-12); - assert!(y.phase_seconds.abs() <= 1e-12); - assert!((dae_model.clocks.intervals["y"] - 0.02).abs() <= 1e-12); -} diff --git a/crates/rumoca-phase-dae/src/scalar_inference/inference_and_bindings.rs b/crates/rumoca-phase-dae/src/scalar_inference/inference_and_bindings.rs index 4633a4e9f..708ca9ef9 100644 --- a/crates/rumoca-phase-dae/src/scalar_inference/inference_and_bindings.rs +++ b/crates/rumoca-phase-dae/src/scalar_inference/inference_and_bindings.rs @@ -16,15 +16,15 @@ pub(crate) fn flat_function_output_dims(func: &rumoca_core::Function) -> Option< /// Compute total scalar size of a function's outputs. pub(crate) fn flat_function_output_scalar_size(func: &rumoca_core::Function) -> usize { + if func.is_constructor && !func.inputs.is_empty() { + return func + .inputs + .iter() + .map(|input| compute_var_size(&input.dims).max(1)) + .sum::() + .max(1); + } if func.outputs.is_empty() { - if func.is_constructor && !func.inputs.is_empty() { - return func - .inputs - .iter() - .map(|input| compute_var_size(&input.dims).max(1)) - .sum::() - .max(1); - } return 1; } func.outputs @@ -76,7 +76,50 @@ pub(crate) fn extract_integer_from_flat_expr(expr: &Expression) -> Option /// For record types like `Orientation = {T[3,3], w[3]}`, the prefix "rev.R_rel" maps to 12 /// (9 for T + 3 for w), not 2 (number of child entries). This ensures record-level equations /// like `rev.R_rel = Frames.planarRotation(...)` get the correct scalar count. -pub(crate) fn build_prefix_counts(flat: &Model) -> FxHashMap { +#[derive(Debug, Clone, Hash, PartialEq, Eq)] +struct ProjectedResultFieldKey { + function: rumoca_core::DefId, + field: String, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +enum ProjectedResultFieldShape { + Known(Vec), + Unknown, +} + +pub(crate) struct ScalarInferenceMetadata { + prefix_counts: FxHashMap, + projected_result_fields: FxHashMap, +} + +impl std::ops::Deref for ScalarInferenceMetadata { + type Target = FxHashMap; + + fn deref(&self) -> &Self::Target { + &self.prefix_counts + } +} + +impl ScalarInferenceMetadata { + pub(crate) fn projected_result_field_dims( + &self, + function: &rumoca_core::Reference, + field: &str, + ) -> Option<&[i64]> { + let function = function.component_ref()?.def_id?; + let key = ProjectedResultFieldKey { + function, + field: field.to_string(), + }; + match self.projected_result_fields.get(&key)? { + ProjectedResultFieldShape::Known(dims) => Some(dims), + ProjectedResultFieldShape::Unknown => None, + } + } +} + +pub(crate) fn build_prefix_counts(flat: &Model) -> ScalarInferenceMetadata { fn normalize_embedded_subscripts(name: &str) -> String { let mut out = String::with_capacity(name.len()); let mut depth = 0usize; @@ -114,7 +157,121 @@ pub(crate) fn build_prefix_counts(flat: &Model) -> FxHashMap { } } } - counts + ScalarInferenceMetadata { + prefix_counts: counts, + projected_result_fields: build_projected_result_field_shapes(flat), + } +} + +fn build_projected_result_field_shapes( + flat: &Model, +) -> FxHashMap { + let mut shapes = FxHashMap::default(); + let known_function_ids = flat + .functions + .values() + .filter_map(|function| function.def_id) + .collect::>(); + collect_constructor_result_field_shapes(flat, &mut shapes); + for variable in flat.variables.values() { + let Some(Expression::FieldAccess { base, field, .. }) = variable.binding.as_ref() else { + continue; + }; + let Expression::FunctionCall { name, .. } = base.as_ref() else { + continue; + }; + let Some(function) = resolved_function_def_id(name, &known_function_ids) else { + continue; + }; + merge_projected_result_field_shape( + &mut shapes, + ProjectedResultFieldKey { + function, + field: field.clone(), + }, + concrete_result_field_dims(&variable.dims), + ); + } + shapes +} + +fn collect_constructor_result_field_shapes( + flat: &Model, + shapes: &mut FxHashMap, +) { + for function in flat.functions.values() { + let (Some(function_def_id), [output]) = (function.def_id, function.outputs.as_slice()) + else { + continue; + }; + if output.type_class != Some(rumoca_core::ClassType::Record) { + continue; + } + let Some(fields) = crate::dae_lowering::record_constructor_fields_from_metadata( + flat.functions.iter(), + &output.type_name, + ) else { + continue; + }; + for field in fields { + let dims = concrete_constructor_result_field_dims(&field); + merge_projected_result_field_shape( + shapes, + ProjectedResultFieldKey { + function: function_def_id, + field: field.name, + }, + dims, + ); + } + } +} + +fn resolved_function_def_id( + reference: &rumoca_core::Reference, + known_function_ids: &FxHashSet, +) -> Option { + let referenced = reference + .component_ref() + .and_then(|component_ref| component_ref.def_id)?; + known_function_ids + .contains(&referenced) + .then_some(referenced) +} + +fn concrete_result_field_dims(dims: &[i64]) -> Option> { + dims.iter() + .all(|dimension| *dimension >= 0) + .then(|| dims.to_vec()) +} + +fn concrete_constructor_result_field_dims(field: &rumoca_core::FunctionParam) -> Option> { + field + .shape_expr + .is_empty() + .then(|| concrete_result_field_dims(&field.dims)) + .flatten() +} + +fn merge_projected_result_field_shape( + shapes: &mut FxHashMap, + key: ProjectedResultFieldKey, + dims: Option>, +) { + let candidate = dims.map_or( + ProjectedResultFieldShape::Unknown, + ProjectedResultFieldShape::Known, + ); + match shapes.entry(key) { + std::collections::hash_map::Entry::Vacant(entry) => { + entry.insert(candidate); + } + std::collections::hash_map::Entry::Occupied(mut entry) => { + if entry.get() != &candidate { + entry.insert(ProjectedResultFieldShape::Unknown); + } + } + } } /// Build a prefix-to-children index: maps each dotted prefix to all descendant variable names. @@ -205,6 +362,15 @@ fn extract_lhs_var_size_from_var_name( flat: &Model, prefix_counts: &FxHashMap, ) -> Option { + if let Expression::VarRef { + name, subscripts, .. + } = lhs + && !subscripts.is_empty() + && let Some(var) = flat.variables.get(name.var_name()) + { + return compute_subscripted_size_with_context(&var.dims, subscripts, flat); + } + // Try to extract the variable name from the LHS (no subscripts) let var_name = extract_var_from_lhs(lhs)?; @@ -254,6 +420,15 @@ fn extract_lhs_var_size_from_var_name( return Some(count); } + if let Some(size) = resolve_singleton_indexed_lhs_path_size(&var_name, flat, prefix_counts) { + return Some(size); + } + + if let Some(size) = resolve_repeated_indexed_component_path_size(&var_name, flat, prefix_counts) + { + return Some(size); + } + // Also try progressively stripping subscripts: // "port_a[1].T[1]" -> "port_a[1].T" -> "port_a.T" for base in subscript_fallback_chain(var_name.as_str()) { @@ -357,6 +532,22 @@ pub(crate) fn extract_lhs_var_size_with_linearized_bases( return Some(size); } + if matches!(lhs.as_ref(), Expression::FieldAccess { .. }) + && let Some(var_name) = render_lhs_path(lhs) + && let Some(span) = lhs.span() + && let Some(size) = extract_lhs_var_size_from_var_name( + &Expression::VarRef { + name: var_name.into(), + subscripts: vec![], + span, + }, + flat, + prefix_counts, + ) + { + return Some(size); + } + if let Expression::Index { base, subscripts, .. } = lhs.as_ref() @@ -409,11 +600,16 @@ fn indexed_lhs_scalar_size( } let total = *prefix_counts.get(base_name.as_str())?; - let full_name = format!( - "{}{}", - base_name.as_str(), - render_subscript_suffix(subscripts)? - ); + let full_name = if let Some(suffix) = render_subscript_suffix(subscripts) { + format!("{}{}", base_name.as_str(), suffix) + } else if subscripts.is_empty() { + return None; + } else { + // Dynamic indexing into a scalarized record array still selects one + // array element. The exact index is runtime-valued, but every element + // has the same scalar record width in flat variables. + format!("{}[1]", base_name.as_str()) + }; Some(record_subscript_scalar_size( &full_name, base_name.as_str(), @@ -443,6 +639,30 @@ fn render_subscript_suffix(subscripts: &[Subscript]) -> Option { Some(out) } +fn render_lhs_path(lhs: &Expression) -> Option { + match lhs { + Expression::VarRef { + name, subscripts, .. + } => { + let mut out = name.as_str().to_string(); + out.push_str(&render_subscript_suffix(subscripts)?); + Some(VarName::new(out)) + } + Expression::Index { + base, subscripts, .. + } => { + let mut out = render_lhs_path(base)?.as_str().to_string(); + out.push_str(&render_subscript_suffix(subscripts)?); + Some(VarName::new(out)) + } + Expression::FieldAccess { base, field, .. } => { + let base = render_lhs_path(base)?; + Some(VarName::new(format!("{}.{}", base.as_str(), field))) + } + _ => None, + } +} + /// Extract a variable name from an LHS expression. /// /// Handles: @@ -497,6 +717,75 @@ pub(crate) fn extract_var_from_lhs(lhs: &Expression) -> Option { } } +fn resolve_singleton_indexed_lhs_path_size( + var_name: &VarName, + flat: &Model, + prefix_counts: &FxHashMap, +) -> Option { + for candidate in singleton_indexed_path_candidates(var_name.as_str()) { + let candidate_name = VarName::new(candidate.clone()); + if let Some(var) = flat.variables.get(&candidate_name) { + return Some(compute_var_size(&var.dims)); + } + if let Some(&count) = prefix_counts.get(candidate.as_str()) { + return Some(count); + } + } + None +} + +fn resolve_repeated_indexed_component_path_size( + var_name: &VarName, + flat: &Model, + prefix_counts: &FxHashMap, +) -> Option { + for candidate in + crate::path_utils::repeated_indexed_component_path_candidates(var_name.as_str()) + { + let candidate_name = VarName::new(candidate.clone()); + if let Some(var) = flat.variables.get(&candidate_name) { + return Some(compute_var_size(&var.dims)); + } + if let Some(&count) = prefix_counts.get(candidate.as_str()) { + return Some(count); + } + } + None +} + +fn singleton_indexed_path_candidates(path: &str) -> Vec { + let parts = split_lhs_path_parts(path); + let mut candidates = Vec::new(); + for index in 0..parts.len() { + if parts[index].contains('[') { + continue; + } + let mut candidate = parts.clone(); + candidate[index] = format!("{}[1]", candidate[index]); + candidates.push(candidate.join(".")); + } + candidates +} + +fn split_lhs_path_parts(path: &str) -> Vec { + let mut parts = Vec::new(); + let mut depth = 0i32; + let mut start = 0usize; + for (index, ch) in path.char_indices() { + match ch { + '[' => depth += 1, + ']' => depth -= 1, + '.' if depth == 0 => { + parts.push(path[start..index].to_string()); + start = index + 1; + } + _ => {} + } + } + parts.push(path[start..].to_string()); + parts +} + /// Create a Variable from a flat::Variable. pub(crate) fn resolve_missing_start_ref( name: &VarName, @@ -528,12 +817,21 @@ impl ExpressionRewriter for StartRefRewriter<'_> { subscripts: &[rumoca_core::Subscript], span: rumoca_core::Span, ) -> Expression { - let resolved_name = if let Some(resolved) = + let resolved_name = if name.has_structure() { + if !self.known_var_names.contains(name.as_str()) + && let Some(resolved) = crate::path_utils::resolve_known_path_suffix( + name.as_str(), + self.known_var_names, + ) + { + VarName::new(resolved).into() + } else { + name.clone() + } + } else if let Some(resolved) = crate::path_utils::resolve_known_path_suffix(name.as_str(), self.known_var_names) { VarName::new(resolved).into() - } else if name.has_structure() { - name.clone() } else { resolve_missing_start_ref(name.var_name(), self.known_var_names).into() }; @@ -654,21 +952,23 @@ fn select_scalar_start_record_alias( let lhs_path = rumoca_core::ComponentPath::from_flat_path(lhs_name.as_str()); let leaf_field = lhs_path.parts().last(); let Some(lhs_base) = flat::component_base_name(lhs_name.as_str()) else { - return select_leaf_start_record_alias(expr, leaf_field, known_var_names, owner_span) - .unwrap_or_else(|| expr.clone()); + return select_leaf_or_original(expr, leaf_field, known_var_names, owner_span); }; if lhs_base == lhs_name.as_str() { - return select_leaf_start_record_alias(expr, leaf_field, known_var_names, owner_span) - .unwrap_or_else(|| expr.clone()); + return select_leaf_or_original(expr, leaf_field, known_var_names, owner_span); } let Some(field_suffix) = lhs_name.as_str().strip_prefix(lhs_base.as_str()) else { - return select_leaf_start_record_alias(expr, leaf_field, known_var_names, owner_span) - .unwrap_or_else(|| expr.clone()); + return select_leaf_or_original(expr, leaf_field, known_var_names, owner_span); }; if !field_suffix.starts_with('.') { - return select_leaf_start_record_alias(expr, leaf_field, known_var_names, owner_span) - .unwrap_or_else(|| expr.clone()); + return select_leaf_or_original(expr, leaf_field, known_var_names, owner_span); } + let selector = StartRecordAliasSelector { + field_suffix, + leaf_field, + known_var_names, + owner_span, + }; match expr { Expression::VarRef { @@ -676,36 +976,7 @@ fn select_scalar_start_record_alias( subscripts, span, } if subscripts.is_empty() => { - // A binding that already names a known scalar variable is a - // complete value; record-field selection would graft the LHS - // field onto an unrelated variable (e.g. `resistor.m = - // multiStarResistance.mBasic` must not become `m`). - if known_var_names.contains(rhs_name.as_str()) { - return expr.clone(); - } - let selected = format!("{}{}", rhs_name.as_str(), field_suffix); - if let Some(selected) = - crate::path_utils::resolve_known_path_suffix(&selected, known_var_names) - { - return Expression::VarRef { - name: VarName::new(selected).into(), - subscripts: vec![], - span: real_or_owner_span(*span, owner_span), - }; - } - if let Some(field) = leaf_field { - let selected = format!("{}.{}", rhs_name.as_str(), field); - if let Some(selected) = - crate::path_utils::resolve_known_path_suffix(&selected, known_var_names) - { - return Expression::VarRef { - name: VarName::new(selected).into(), - subscripts: vec![], - span: real_or_owner_span(*span, owner_span), - }; - } - } - expr.clone() + select_varref_start_record_alias(expr, rhs_name.var_name(), *span, &selector) } Expression::FieldAccess { base, field, span } => { let Expression::VarRef { @@ -719,37 +990,157 @@ fn select_scalar_start_record_alias( if !subscripts.is_empty() { return expr.clone(); } - if known_var_names.contains(&format!("{}.{}", rhs_name.as_str(), field)) { - return expr.clone(); - } - let selected = format!("{}.{}{}", rhs_name.as_str(), field, field_suffix); - if let Some(selected) = - crate::path_utils::resolve_known_path_suffix(&selected, known_var_names) - { - return Expression::VarRef { - name: VarName::new(selected).into(), - subscripts: vec![], - span: real_or_owner_span(*span, owner_span), - }; - } - if let Some(lhs_leaf) = leaf_field { - let selected = format!("{}.{}.{}", rhs_name.as_str(), field, lhs_leaf); - if let Some(selected) = - crate::path_utils::resolve_known_path_suffix(&selected, known_var_names) - { - return Expression::VarRef { - name: VarName::new(selected).into(), - subscripts: vec![], - span: real_or_owner_span(*span, owner_span), - }; - } - } - expr.clone() + select_field_start_record_alias(expr, rhs_name.var_name(), field, *span, &selector) } + Expression::FunctionCall { + name, + is_constructor: false, + .. + } if is_state_constructor_function(name) => select_function_record_field_start( + expr, + field_suffix, + leaf_field.map(String::as_str), + owner_span, + ) + .unwrap_or_else(|| expr.clone()), _ => expr.clone(), } } +struct StartRecordAliasSelector<'a> { + field_suffix: &'a str, + leaf_field: Option<&'a String>, + known_var_names: &'a HashSet, + owner_span: rumoca_core::Span, +} + +fn select_leaf_or_original( + expr: &Expression, + leaf_field: Option<&String>, + known_var_names: &HashSet, + owner_span: rumoca_core::Span, +) -> Expression { + select_leaf_start_record_alias(expr, leaf_field, known_var_names, owner_span) + .unwrap_or_else(|| expr.clone()) +} + +fn select_varref_start_record_alias( + expr: &Expression, + rhs_name: &VarName, + span: rumoca_core::Span, + selector: &StartRecordAliasSelector<'_>, +) -> Expression { + // A binding that already names a known scalar variable is a complete value; + // record-field selection would graft the LHS field onto an unrelated + // variable (e.g. `resistor.m = multiStarResistance.mBasic` must not become + // `m`). + if selector.known_var_names.contains(rhs_name.as_str()) { + return expr.clone(); + } + resolve_selected_start_var( + &format!("{}{}", rhs_name.as_str(), selector.field_suffix), + selector.known_var_names, + span, + selector.owner_span, + ) + .or_else(|| { + selector.leaf_field.and_then(|field| { + resolve_selected_start_var( + &format!("{}.{}", rhs_name.as_str(), field), + selector.known_var_names, + span, + selector.owner_span, + ) + }) + }) + .unwrap_or_else(|| expr.clone()) +} + +fn select_field_start_record_alias( + expr: &Expression, + rhs_name: &VarName, + field: &str, + span: rumoca_core::Span, + selector: &StartRecordAliasSelector<'_>, +) -> Expression { + if selector + .known_var_names + .contains(&format!("{}.{}", rhs_name.as_str(), field)) + { + return expr.clone(); + } + resolve_selected_start_var( + &format!("{}.{}{}", rhs_name.as_str(), field, selector.field_suffix), + selector.known_var_names, + span, + selector.owner_span, + ) + .or_else(|| { + selector.leaf_field.and_then(|lhs_leaf| { + resolve_selected_start_var( + &format!("{}.{}.{}", rhs_name.as_str(), field, lhs_leaf), + selector.known_var_names, + span, + selector.owner_span, + ) + }) + }) + .unwrap_or_else(|| expr.clone()) +} + +fn resolve_selected_start_var( + candidate: &str, + known_var_names: &HashSet, + span: rumoca_core::Span, + owner_span: rumoca_core::Span, +) -> Option { + crate::path_utils::resolve_known_path_suffix(candidate, known_var_names).map(|selected| { + Expression::VarRef { + name: VarName::new(selected).into(), + subscripts: vec![], + span: real_or_owner_span(span, owner_span), + } + }) +} + +fn is_state_constructor_function(name: &rumoca_core::Reference) -> bool { + matches!( + name.var_name().last_segment(), + "setState_pTX" + | "setState_pT" + | "setState_dTX" + | "setState_phX" + | "setState_ph" + | "setState_psX" + | "setState_ps" + | "setSmoothState" + ) +} + +fn select_function_record_field_start( + expr: &Expression, + field_suffix: &str, + leaf_field: Option<&str>, + owner_span: rumoca_core::Span, +) -> Option { + let suffix = field_suffix.strip_prefix('.')?; + let mut segments = rumoca_core::ComponentPath::from_flat_path(suffix).into_parts(); + if segments.is_empty() + && let Some(field) = leaf_field + { + segments.push(field.to_string()); + } + let mut selected = expr.clone(); + for field in segments { + selected = Expression::FieldAccess { + base: Box::new(selected), + field, + span: owner_span, + }; + } + Some(selected) +} + fn select_leaf_start_record_alias( expr: &Expression, leaf_field: Option<&String>, @@ -796,6 +1187,15 @@ fn select_leaf_start_record_alias( }, ) } + Expression::FunctionCall { + name, + is_constructor: false, + .. + } if is_state_constructor_function(name) => Some(Expression::FieldAccess { + base: Box::new(expr.clone()), + field: field.clone(), + span: owner_span, + }), _ => None, } } @@ -854,6 +1254,208 @@ mod tests { ); } + #[test] + fn record_field_start_from_function_call_preserves_field_access() { + let span = test_span(); + let expr = Expression::FunctionCall { + name: rumoca_core::Reference::new("Buildings.Media.Air.setState_pTX"), + args: vec![], + is_constructor: false, + span, + }; + let selected = select_scalar_start_record_alias( + &VarName::new("state_default.X"), + &expr, + &HashSet::from(["state_default.X".to_string()]), + span, + ); + + let Expression::FieldAccess { base, field, .. } = selected else { + panic!("expected function field access start, got {selected:?}"); + }; + assert_eq!(field, "X"); + assert!(matches!( + base.as_ref(), + Expression::FunctionCall { name, .. } + if name.as_str() == "Buildings.Media.Air.setState_pTX" + )); + } + + #[test] + fn record_field_start_from_constructor_call_is_not_rewritten_to_field_access() { + let span = test_span(); + let expr = Expression::FunctionCall { + name: rumoca_core::Reference::new("Modelica.Blocks.Types.ExternalCombiTable1D"), + args: vec![], + is_constructor: true, + span, + }; + let selected = select_scalar_start_record_alias( + &VarName::new("tableID"), + &expr, + &HashSet::from(["tableID".to_string()]), + span, + ); + + assert!(matches!( + selected, + Expression::FunctionCall { + is_constructor: true, + .. + } + )); + } + + #[test] + fn scalar_start_from_function_call_is_not_rewritten_to_field_access() { + let span = test_span(); + let expr = Expression::FunctionCall { + name: rumoca_core::Reference::new("Modelica.Units.Conversions.from_degC"), + args: vec![], + is_constructor: false, + span, + }; + let selected = select_scalar_start_record_alias( + &VarName::new("component.T_start"), + &expr, + &HashSet::from(["component.T_start".to_string()]), + span, + ); + + assert!(matches!( + selected, + Expression::FunctionCall { + is_constructor: false, + .. + } + )); + } + + #[test] + fn indexed_lhs_scalar_size_uses_record_element_width_for_dynamic_subscript() { + let mut flat = Model::new(); + for name in ["sym[1].re", "sym[1].im", "sym[2].re", "sym[2].im", "idx[1]"] { + flat.add_variable( + VarName::new(name), + flat::Variable { + name: VarName::new(name), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + + let residual = Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(Expression::Index { + base: Box::new(Expression::VarRef { + name: VarName::new("sym").into(), + subscripts: vec![], + span: test_span(), + }), + subscripts: vec![Subscript::Expr { + expr: Box::new(Expression::VarRef { + name: VarName::new("idx").into(), + subscripts: vec![Subscript::Index { + value: 1, + span: test_span(), + }], + span: test_span(), + }), + span: test_span(), + }], + span: test_span(), + }), + rhs: Box::new(Expression::FunctionCall { + name: VarName::new("Complex").into(), + args: vec![ + Expression::Literal { + value: Literal::Integer(0), + span: test_span(), + }, + Expression::Literal { + value: Literal::Integer(0), + span: test_span(), + }, + ], + is_constructor: true, + span: test_span(), + }), + span: test_span(), + }; + + let prefix_counts = build_prefix_counts(&flat); + assert_eq!( + infer_equation_scalar_count(&residual, &flat, &prefix_counts), + 2 + ); + } + + #[test] + fn record_constructor_residual_uses_scalarized_unknown_width() { + let mut flat = Model::new(); + let mut complex = rumoca_core::Function::new("Complex", test_span()); + complex.is_constructor = true; + complex.add_input(rumoca_core::FunctionParam::new("re", "Real", test_span())); + complex.add_input(rumoca_core::FunctionParam::new("im", "Real", test_span())); + complex.add_output(rumoca_core::FunctionParam::new( + "result", + "Complex", + test_span(), + )); + flat.add_function(complex); + for name in ["currentSensor.i.re", "currentSensor.i.im"] { + flat.add_variable( + VarName::new(name), + flat::Variable { + name: VarName::new(name), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + + let residual = Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(Expression::FunctionCall { + name: VarName::new("Complex").into(), + args: vec![ + Expression::Literal { + value: Literal::Integer(0), + span: test_span(), + }, + Expression::Literal { + value: Literal::Integer(0), + span: test_span(), + }, + ], + is_constructor: true, + span: test_span(), + }), + rhs: Box::new(Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: Box::new(Expression::VarRef { + name: VarName::new("currentSensor.i").into(), + subscripts: vec![], + span: test_span(), + }), + rhs: Box::new(Expression::VarRef { + name: VarName::new("currentSensor.i").into(), + subscripts: vec![], + span: test_span(), + }), + span: test_span(), + }), + span: test_span(), + }; + + let prefix_counts = build_prefix_counts(&flat); + assert_eq!( + infer_equation_scalar_count(&residual, &flat, &prefix_counts), + 2 + ); + } + #[test] fn create_dae_variable_preserves_structured_binding_references() { let record_def_id = rumoca_core::DefId::new(42); @@ -1065,4 +1667,5 @@ mod tests { assert!(matches!(err, ToDaeError::RuntimeMetadataViolation { .. })); } + } diff --git a/crates/rumoca-phase-dae/src/scalar_inference/mod.rs b/crates/rumoca-phase-dae/src/scalar_inference/mod.rs index 375b5b190..37ab64b53 100644 --- a/crates/rumoca-phase-dae/src/scalar_inference/mod.rs +++ b/crates/rumoca-phase-dae/src/scalar_inference/mod.rs @@ -280,7 +280,7 @@ pub(crate) fn infer_binary_expression_form( lhs: &Expression, rhs: &Expression, flat: &Model, - prefix_counts: &FxHashMap, + prefix_counts: &ScalarInferenceMetadata, ) -> ExpressionForm { let lhs_form = infer_expression_form(lhs, flat, prefix_counts); let rhs_form = infer_expression_form(rhs, flat, prefix_counts); @@ -324,7 +324,7 @@ pub(crate) fn infer_builtin_expression_form( function: &BuiltinFunction, args: &[Expression], flat: &Model, - prefix_counts: &FxHashMap, + prefix_counts: &ScalarInferenceMetadata, ) -> ExpressionForm { if is_reduction_builtin(function) { return ExpressionForm::Scalar; @@ -346,7 +346,7 @@ pub(crate) fn infer_array_expression_form( elements: &[Expression], is_matrix: bool, flat: &Model, - prefix_counts: &FxHashMap, + prefix_counts: &ScalarInferenceMetadata, ) -> ExpressionForm { if is_matrix { return ExpressionForm::Other; @@ -367,7 +367,7 @@ pub(crate) fn infer_if_expression_form( branches: &[(Expression, Expression)], else_branch: &Expression, flat: &Model, - prefix_counts: &FxHashMap, + prefix_counts: &ScalarInferenceMetadata, ) -> ExpressionForm { let else_form = infer_expression_form(else_branch, flat, prefix_counts); if branches @@ -383,7 +383,7 @@ pub(crate) fn infer_if_expression_form( pub(crate) fn infer_expression_form( expr: &Expression, flat: &Model, - prefix_counts: &FxHashMap, + prefix_counts: &ScalarInferenceMetadata, ) -> ExpressionForm { match expr { Expression::Literal { value: _, .. } => ExpressionForm::Scalar, @@ -422,18 +422,39 @@ pub(crate) fn infer_expression_form( _ => ExpressionForm::Other, } } + Expression::FieldAccess { base, field, .. } => { + infer_projected_result_field_form(base, field, prefix_counts) + } Expression::Tuple { .. } | Expression::Range { .. } - | Expression::FieldAccess { .. } | Expression::ArrayComprehension { .. } | Expression::Empty { .. } => ExpressionForm::Other, } } +fn infer_projected_result_field_form( + base: &Expression, + field: &str, + metadata: &ScalarInferenceMetadata, +) -> ExpressionForm { + let Expression::FunctionCall { name, .. } = base else { + return ExpressionForm::Other; + }; + let Some(dims) = metadata.projected_result_field_dims(name, field) else { + return ExpressionForm::Other; + }; + match dims { + [] => ExpressionForm::Scalar, + [dimension] => ExpressionForm::Vector(*dimension as usize), + [rows, columns] => ExpressionForm::Matrix(*rows as usize, *columns as usize), + _ => ExpressionForm::Vector(compute_var_size(dims)), + } +} + pub(crate) fn infer_equation_scalar_count_from_forms( residual: &Expression, flat: &Model, - prefix_counts: &FxHashMap, + prefix_counts: &ScalarInferenceMetadata, ) -> Option { let Expression::Binary { op, lhs, rhs, .. } = residual else { return None; @@ -470,7 +491,7 @@ pub(crate) fn infer_equation_scalar_count_from_forms( pub(crate) fn infer_equation_scalar_count( residual: &Expression, flat: &Model, - prefix_counts: &FxHashMap, + prefix_counts: &ScalarInferenceMetadata, ) -> usize { // First try extracting size from a simple LHS pattern (var - expr = 0) if let Some(size) = extract_lhs_var_size(residual, flat, prefix_counts) { @@ -499,6 +520,11 @@ pub(crate) fn infer_equation_scalar_count( let form_size = infer_equation_scalar_count_from_forms(residual, flat, prefix_counts); let varref_size = infer_scalar_count_from_varrefs(residual, flat, prefix_counts); match (form_size, varref_size) { + (Some(1), Some(varref)) + if varref > 1 && has_top_level_record_constructor(residual, flat) => + { + varref + } (Some(a), Some(b)) => a.min(b), (Some(a), None) => a, (None, Some(b)) => b, @@ -506,6 +532,35 @@ pub(crate) fn infer_equation_scalar_count( } } +fn has_top_level_record_constructor(residual: &Expression, flat: &Model) -> bool { + let Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs, + rhs, + .. + } = residual + else { + return false; + }; + expression_is_record_constructor(lhs, flat) || expression_is_record_constructor(rhs, flat) +} + +fn expression_is_record_constructor(expr: &Expression, flat: &Model) -> bool { + let Expression::FunctionCall { + name, + is_constructor, + .. + } = expr + else { + return false; + }; + *is_constructor + || flat + .functions + .get(name.var_name()) + .is_some_and(|function| function.is_constructor && function.inputs.len() > 1) +} + /// Infer scalar count by finding VarRefs in an expression and checking for record expansion. /// /// Per MLS §10.2, equations involving record types represent multiple scalar equations, @@ -520,7 +575,7 @@ pub(crate) fn infer_scalar_count_from_varrefs( prefix_counts: &FxHashMap, ) -> Option { let mut var_refs = Vec::new(); - collect_var_refs_skip_reductions(expr, &mut var_refs); + collect_var_refs_for_cardinality(expr, &mut var_refs); infer_scalar_count_from_collected_varrefs(&var_refs, flat, prefix_counts) } @@ -558,7 +613,7 @@ pub(crate) fn infer_flow_sum_scalar_count( prefix_counts: &FxHashMap, ) -> Option { let mut var_refs = Vec::new(); - collect_var_refs_skip_reductions(residual, &mut var_refs); + collect_var_refs_for_cardinality(residual, &mut var_refs); if var_refs.is_empty() { return None; } diff --git a/crates/rumoca-phase-dae/src/scalar_inference/parts.rs b/crates/rumoca-phase-dae/src/scalar_inference/parts.rs index 74c37afb9..2b043dfcc 100644 --- a/crates/rumoca-phase-dae/src/scalar_inference/parts.rs +++ b/crates/rumoca-phase-dae/src/scalar_inference/parts.rs @@ -71,6 +71,14 @@ pub(crate) fn infer_scalar_count_from_collected_varrefs( continue; } + if var_ref.subscripts.is_empty() + && let Some(size) = repeated_indexed_component_size(var_name, flat, prefix_counts) + { + any_found = true; + max_size = max_size.max(size); + continue; + } + // Try progressively stripping embedded subscripts: // "a[1].b[2]" -> "a[1].b" -> "a.b" let fallback_chain = subscript_fallback_chain(var_name.as_str()); @@ -115,17 +123,45 @@ pub(crate) fn infer_scalar_count_from_collected_varrefs( } } +fn repeated_indexed_component_size( + var_name: &VarName, + flat: &Model, + prefix_counts: &FxHashMap, +) -> Option { + crate::path_utils::repeated_indexed_component_path_candidates(var_name.as_str()) + .into_iter() + .find_map(|candidate| { + repeated_indexed_component_candidate_size(&candidate, flat, prefix_counts) + }) +} + +fn repeated_indexed_component_candidate_size( + candidate: &str, + flat: &Model, + prefix_counts: &FxHashMap, +) -> Option { + let candidate_name = VarName::new(candidate.to_string()); + flat.variables + .get(&candidate_name) + .map(|var| compute_var_size(&var.dims)) + .or_else(|| prefix_counts.get(candidate).copied()) +} + struct VarRefCollectionVisitor<'a> { vars: &'a mut Vec, - skip_function_args: bool, + mode: VarRefCollectionMode, +} + +#[derive(Clone, Copy, PartialEq, Eq)] +enum VarRefCollectionMode { + Dependencies, + Cardinality, + SkipFunctionArgs, } impl<'a> VarRefCollectionVisitor<'a> { - fn new(vars: &'a mut Vec, skip_function_args: bool) -> Self { - Self { - vars, - skip_function_args, - } + fn new(vars: &'a mut Vec, mode: VarRefCollectionMode) -> Self { + Self { vars, mode } } } @@ -152,12 +188,131 @@ impl rumoca_core::ExpressionVisitor for VarRefCollectionVisitor<'_> { args: &[Expression], is_constructor: bool, ) { - if self.skip_function_args { + if self.mode == VarRefCollectionMode::SkipFunctionArgs { // Function arguments are not shaped like function output. return; } self.walk_function_call(name, args, is_constructor); } + + fn visit_index(&mut self, base: &Expression, subscripts: &[Subscript]) { + if let Some(name) = render_index_path(base, subscripts) { + self.vars.push(CollectedVarRef { + name, + subscripts: Vec::new(), + }); + return; + } + if let Expression::VarRef { name, .. } = base { + self.vars.push(CollectedVarRef { + name: name.var_name().clone(), + subscripts: subscripts.to_vec(), + }); + return; + } + rumoca_core::ExpressionVisitor::visit_expression(self, base); + for subscript in subscripts { + self.visit_subscript(subscript); + } + } + + fn visit_field_access(&mut self, base: &Expression, field: &str) { + if let Some(name) = render_field_path(base, field) { + self.vars.push(CollectedVarRef { + name, + subscripts: Vec::new(), + }); + return; + } + if self.mode == VarRefCollectionMode::Cardinality + && matches!(base, Expression::FunctionCall { .. }) + { + // The selected result field owns expression shape. Function-call + // arguments remain dependencies, but their record widths do not + // determine the result field's equation cardinality. + return; + } + rumoca_core::ExpressionVisitor::visit_expression(self, base); + } +} + +fn render_index_path(base: &Expression, subscripts: &[Subscript]) -> Option { + let mut out = render_expr_path(base)?.as_str().to_string(); + out.push_str(&render_static_subscript_suffix(subscripts)?); + Some(VarName::new(out)) +} + +fn render_field_path(base: &Expression, field: &str) -> Option { + let base = render_expr_path(base)?; + Some(VarName::new(format!("{}.{}", base.as_str(), field))) +} + +enum Tail { + Field(String), + Subscripts(String), +} + +fn render_expr_path(expr: &Expression) -> Option { + let mut current = expr; + let mut tail = Vec::new(); + loop { + match current { + Expression::VarRef { + name, subscripts, .. + } => { + let mut out = name.as_str().to_string(); + out.push_str(&render_static_subscript_suffix(subscripts)?); + for part in tail.iter().rev() { + append_tail_part(&mut out, part); + } + return Some(VarName::new(out)); + } + Expression::Index { + base, subscripts, .. + } => { + tail.push(Tail::Subscripts(render_static_subscript_suffix( + subscripts, + )?)); + current = base; + } + Expression::FieldAccess { base, field, .. } => { + tail.push(Tail::Field(field.clone())); + current = base; + } + _ => return None, + } + } +} + +fn append_tail_part(out: &mut String, part: &Tail) { + match part { + Tail::Field(field) => { + out.push('.'); + out.push_str(field); + } + Tail::Subscripts(subscripts) => out.push_str(subscripts), + } +} + +fn render_static_subscript_suffix(subscripts: &[Subscript]) -> Option { + let mut out = String::new(); + for subscript in subscripts { + match subscript { + Subscript::Index { value, .. } => out.push_str(&format!("[{value}]")), + Subscript::Expr { expr, .. } => { + let Expression::Literal { + value: Literal::Integer(value), + .. + } = expr.as_ref() + else { + return None; + }; + out.push_str(&format!("[{value}]")); + } + Subscript::Colon { .. } => out.push_str("[:]"), + } + } + Some(out) } /// Collect VarRefs from an expression, skipping arguments to reduction builtins. @@ -165,7 +320,17 @@ impl rumoca_core::ExpressionVisitor for VarRefCollectionVisitor<'_> { /// Reduction builtins like sum(), product() reduce arrays to scalars (MLS §10.3.4), /// so their array arguments should not inflate the equation's scalar count. pub(crate) fn collect_var_refs_skip_reductions(expr: &Expression, vars: &mut Vec) { - let mut collector = VarRefCollectionVisitor::new(vars, false); + let mut collector = VarRefCollectionVisitor::new(vars, VarRefCollectionMode::Dependencies); + rumoca_core::ExpressionVisitor::visit_expression(&mut collector, expr); +} + +/// Collect VarRefs that can own expression cardinality. +/// +/// A selected function-result field takes its shape from result-field metadata, +/// so its call arguments are excluded here. Semantic dependency collectors use +/// [`collect_var_refs_skip_reductions`] and still traverse those arguments. +pub(crate) fn collect_var_refs_for_cardinality(expr: &Expression, vars: &mut Vec) { + let mut collector = VarRefCollectionVisitor::new(vars, VarRefCollectionMode::Cardinality); rumoca_core::ExpressionVisitor::visit_expression(&mut collector, expr); } @@ -178,7 +343,7 @@ pub(crate) fn collect_var_refs_skip_reductions_and_function_args( expr: &Expression, vars: &mut Vec, ) { - let mut collector = VarRefCollectionVisitor::new(vars, true); + let mut collector = VarRefCollectionVisitor::new(vars, VarRefCollectionMode::SkipFunctionArgs); rumoca_core::ExpressionVisitor::visit_expression(&mut collector, expr); } diff --git a/crates/rumoca-phase-dae/src/tests/conditions.rs b/crates/rumoca-phase-dae/src/tests/conditions.rs index 66df58a33..8070214ec 100644 --- a/crates/rumoca-phase-dae/src/tests/conditions.rs +++ b/crates/rumoca-phase-dae/src/tests/conditions.rs @@ -552,7 +552,7 @@ fn test_todae_suppresses_not_initial_else_relations_from_runtime_roots() { } #[test] -fn test_todae_keeps_smooth_if_conditions_event_generating() { +fn test_todae_suppresses_smooth_if_conditions_from_runtime_roots() { let mut flat = Model::new(); flat.add_variable(VarName::new("x"), scalar_var("x")); flat.add_variable(VarName::new("u"), input_var("u")); @@ -566,15 +566,11 @@ fn test_todae_keeps_smooth_if_conditions_event_generating() { let dae_model = to_dae_unbalanced_ok(&flat); - assert_eq!( - dae_model.conditions.relations.len(), - 1, - "smooth is not noEvent; relations inside smooth may still generate events" - ); assert!( - rumoca_core::expressions_semantically_equal(&dae_model.conditions.relations[0], &condition), - "relation should match the source condition" + dae_model.conditions.relations.is_empty(), + "relations inside smooth() may be evaluated without zero-crossing events" ); + assert!(dae_model.conditions.equations.is_empty()); assert!( matches!( &dae_model.continuous.equations[0].rhs, @@ -587,17 +583,12 @@ fn test_todae_keeps_smooth_if_conditions_event_generating() { rumoca_core::Expression::If { branches, .. } if matches!( &branches[0].0, - rumoca_core::Expression::VarRef { name, subscripts, .. } - if name.as_str() == "c" - && matches!( - subscripts.as_slice(), - [rumoca_core::Subscript::Index { value: 1, .. }] - ) + rumoca_core::Expression::Binary { .. } ) ) ) ), - "relations inside smooth should be rewritten through Appendix B condition memory" + "relations inside smooth should remain direct expressions" ); } diff --git a/crates/rumoca-phase-dae/src/tests/mod.rs b/crates/rumoca-phase-dae/src/tests/mod.rs index 94c02c429..3a4d02408 100644 --- a/crates/rumoca-phase-dae/src/tests/mod.rs +++ b/crates/rumoca-phase-dae/src/tests/mod.rs @@ -4,3 +4,119 @@ mod algorithm_lowering; mod conditions; mod initialization; mod root; + +fn var_ref(name: &str) -> Expression { + Expression::VarRef { + name: VarName::new(name).into(), + subscripts: vec![], + span: Span::DUMMY, + } +} + +fn der_ref(name: &str) -> Expression { + Expression::BuiltinCall { + function: BuiltinFunction::Der, + args: vec![var_ref(name)], + span: Span::DUMMY, + } +} + +#[test] +fn prune_unreferenced_local_algebraics_keeps_referenced_and_public_vars() { + let mut dae = dae::Dae::new(); + for name in ["used", "unused_scalar", "unused_structured", "public_out"] { + let mut variable = dae::Variable::new(VarName::new(name), Span::DUMMY); + variable.origin = dae::VariableOrigin::Source; + dae.variables + .algebraics + .insert(VarName::new(name), variable); + } + dae.variables + .algebraics + .get_mut(&VarName::new("unused_structured")) + .expect("fixture variable exists") + .component_ref = rumoca_core::component_reference_from_flat_name( + &VarName::new("unused_structured.field"), + Span::DUMMY, + ); + dae.variables + .algebraics + .get_mut(&VarName::new("public_out")) + .expect("fixture variable exists") + .causality = dae::VariableCausality::Output; + dae.continuous.equations.push(dae::Equation::explicit( + VarName::new("used"), + var_ref("source"), + Span::DUMMY, + "test equation", + )); + + prune_unreferenced_local_algebraics(&mut dae); + + assert!(dae.variables.algebraics.contains_key(&VarName::new("used"))); + assert!( + dae.variables + .algebraics + .contains_key(&VarName::new("public_out")) + ); + assert!( + !dae.variables + .algebraics + .contains_key(&VarName::new("unused_structured")) + ); + assert!( + dae.variables + .algebraics + .contains_key(&VarName::new("unused_scalar")) + ); +} + +#[test] +fn overconstrained_derivative_alias_rewrite_targets_root_state() { + let mut dae = dae::Dae::new(); + dae.continuous.equations.push(dae::Equation::residual( + Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(der_ref("branch.port.reference.gamma")), + rhs: Box::new(var_ref("branch.omega")), + span: Span::DUMMY, + }, + Span::DUMMY, + "branch omega", + )); + let alias_roots = FxHashMap::from_iter([( + VarName::new("branch.port.reference.gamma"), + VarName::new("root.port.reference.gamma"), + )]); + + let mut flat = flat::Model::new(); + flat.oc_break_edge_scalar_count = 1; + for name in ["branch.port.reference.gamma", "root.port.reference.gamma"] { + flat.add_variable( + VarName::new(name), + flat::Variable { + name: VarName::new(name), + is_primitive: true, + ..flat::Variable::empty_with_span(Span::DUMMY) + }, + ); + } + rewrite_overconstrained_derivative_alias_refs(&mut dae, &flat, &alias_roots) + .expect("alias rewrite should succeed"); + + assert_eq!(dae.continuous.equations.len(), 2); + let Expression::Binary { lhs, .. } = &dae.continuous.equations[0].rhs else { + panic!("expected residual binary expression"); + }; + let Expression::BuiltinCall { args, .. } = lhs.as_ref() else { + panic!("expected der() call on lhs"); + }; + let Expression::VarRef { name, .. } = &args[0] else { + panic!("expected der() argument var ref"); + }; + assert_eq!(name.var_name(), &VarName::new("root.port.reference.gamma")); + assert_eq!( + dae.continuous.equations[1].origin, + "overconstrained derivative alias: branch.port.reference.gamma = root.port.reference.gamma" + ); +} diff --git a/crates/rumoca-phase-dae/src/tests/root/input_binding_tests.rs b/crates/rumoca-phase-dae/src/tests/root/input_binding_tests.rs index fd252fca4..fd8048a02 100644 --- a/crates/rumoca-phase-dae/src/tests/root/input_binding_tests.rs +++ b/crates/rumoca-phase-dae/src/tests/root/input_binding_tests.rs @@ -605,6 +605,65 @@ fn test_is_input_default_equation_false_for_unknown_rhs() { assert!(!is_input_default_equation(&eq, &flat, &dae)); } +#[test] +fn test_input_default_dependency_keeps_projected_function_call_arguments() { + let mut flat = Model::new(); + flat.top_level_input_components + .insert("top_input".to_string()); + + let mut dae = Dae::new(); + dae.variables.inputs.insert( + rumoca_core::VarName::new("top_input"), + Variable::new( + rumoca_core::VarName::new("top_input"), + crate::test_support::test_span(), + ), + ); + dae.variables.algebraics.insert( + rumoca_core::VarName::new("continuous_x"), + Variable::new( + rumoca_core::VarName::new("continuous_x"), + crate::test_support::test_span(), + ), + ); + + let eq = rumoca_ir_flat::Equation { + residual: rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(rumoca_core::Expression::VarRef { + name: VarName::new("top_input").into(), + subscripts: vec![], + span: crate::test_support::test_span(), + }), + rhs: Box::new(rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FunctionCall { + name: VarName::new("calc").into(), + args: vec![rumoca_core::Expression::VarRef { + name: VarName::new("continuous_x").into(), + subscripts: vec![], + span: crate::test_support::test_span(), + }], + is_constructor: false, + span: crate::test_support::test_span(), + }), + field: "field".to_string(), + span: crate::test_support::test_span(), + }), + span: crate::test_support::test_span(), + }, + span: crate::test_support::test_span(), + origin: rumoca_ir_flat::EquationOrigin::ComponentEquation { + component: "model".to_string(), + }, + scalar_count: 1, + }; + + assert!( + !is_input_default_equation(&eq, &flat, &dae), + "continuous function arguments must keep the top-input equation as a constraint" + ); +} + #[test] fn test_is_input_default_equation_false_for_rhs_input_alias() { let mut flat = Model::new(); @@ -767,6 +826,57 @@ fn test_connected_input_binding_kept_for_input_only_connection_alias() { ); } +#[test] +fn test_connected_input_alias_keeps_single_binding_anchor() { + let mut flat = Model::new(); + for (name, value) in [("inner.p", 1.0), ("inner.q", 2.0)] { + flat.add_variable( + VarName::new(name), + crate::test_support::with_component_ref(flat::Variable { + name: VarName::new(name), + causality: rumoca_core::Causality::Input(rumoca_core::Token::default()), + variability: rumoca_core::Variability::Empty, + is_primitive: true, + binding: Some(rumoca_core::Expression::Literal { + value: Literal::Real(value), + span: crate::test_support::test_span(), + }), + ..rumoca_ir_flat::Variable::empty_with_span(rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + )) + }), + ); + } + + add_connection_equation(&mut flat, "inner.q", "inner.p"); + + let dae = to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("to_dae should succeed for connected internal input aliases"); + + let binding_equations = dae + .continuous + .equations + .iter() + .filter(|eq| eq.origin.starts_with("binding equation for inner.")) + .count(); + assert_eq!( + binding_equations, 1, + "one input-only alias component should keep exactly one binding anchor" + ); + assert_eq!( + crate::balance::balance(&dae).expect("valid DAE balance fixture"), + 0, + "multiple default bindings in one input-only alias set must not over-constrain balance" + ); +} + #[test] fn test_connected_input_alias_with_multilayer_subscripts_promotes_internal_inputs() { let mut flat = Model::new(); diff --git a/crates/rumoca-phase-dae/src/tests/root/mod.rs b/crates/rumoca-phase-dae/src/tests/root/mod.rs index 5019ea6bc..d62d2a090 100644 --- a/crates/rumoca-phase-dae/src/tests/root/mod.rs +++ b/crates/rumoca-phase-dae/src/tests/root/mod.rs @@ -521,7 +521,7 @@ fn test_todae_ignores_unreachable_function_without_body() { } #[test] -fn test_todae_rejects_member_style_function_call_without_resolved_name() { +fn test_todae_accepts_unique_member_style_function_call_by_leaf() { let mut flat = Model::new(); add_primitive_real(&mut flat, "x"); @@ -536,13 +536,40 @@ fn test_todae_rejects_member_style_function_call_without_resolved_name() { add_scalar_ode_with_rhs_call(&mut flat, "x", "world.gravityAcceleration"); + to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("unique leaf match should resolve member-style call for codegen export"); +} + +#[test] +fn test_todae_rejects_ambiguous_member_style_function_call_by_leaf() { + let mut flat = Model::new(); + add_primitive_real(&mut flat, "x"); + + for name in [ + "Modelica.Mechanics.MultiBody.World.gravityAcceleration", + "Modelica.Mechanics.MultiBody.Examples.gravityAcceleration", + ] { + let mut fn_def = rumoca_core::Function::new(name, crate::test_support::test_span()); + fn_def.body.push(rumoca_core::Statement::Return { + span: crate::test_support::test_span(), + }); + flat.add_function(fn_def); + } + + add_scalar_ode_with_rhs_call(&mut flat, "x", "world.gravityAcceleration"); + let err = to_dae_with_options( &flat, ToDaeOptions { error_on_unbalanced: false, }, ) - .expect_err("member-style call should fail without prior name resolution"); + .expect_err("ambiguous leaf match should fail closed"); assert!( matches!( @@ -550,7 +577,7 @@ fn test_todae_rejects_member_style_function_call_without_resolved_name() { ToDaeError::UnresolvedFunctionCall { ref name, .. } if name == "world.gravityAcceleration" ), - "expected unresolved function diagnostic for member-style call, got {err:?}" + "expected unresolved function diagnostic for ambiguous member-style call, got {err:?}" ); } @@ -639,6 +666,63 @@ fn test_todae_accepts_record_constructor_calls_for_known_type_names() { .expect("record constructor calls should be accepted for known type names"); } +#[test] +fn test_todae_keeps_external_object_constructor_out_of_fx() { + let mut flat = Model::new(); + let mut constructor = + rumoca_core::Function::new("Pkg.SpawnExternalObject", crate::test_support::test_span()); + constructor.is_constructor = true; + constructor.external = Some(rumoca_core::ExternalFunction { + language: "C".to_string(), + function_name: Some("allocate_object".to_string()), + output_name: Some("adapter".to_string()), + ..Default::default() + }); + flat.add_function(constructor); + flat.add_variable( + VarName::new("obj"), + crate::test_support::with_component_ref(flat::Variable { + name: VarName::new("obj"), + is_primitive: true, + binding: Some(rumoca_core::Expression::FunctionCall { + name: VarName::new("Pkg.SpawnExternalObject").into(), + args: vec![], + is_constructor: true, + span: crate::test_support::test_span(), + }), + ..rumoca_ir_flat::Variable::empty_with_span(rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + )) + }), + ); + flat.variable_type_names + .insert(VarName::new("obj"), "Pkg.SpawnExternalObject".to_string()); + + let dae = to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("external object handles are metadata, not continuous residuals"); + + assert!( + dae.metadata + .nonnumeric_variable_names + .contains(&"obj".to_string()) + ); + assert!(!dae.variables.algebraics.contains_key(&VarName::new("obj"))); + assert!(!dae.variables.outputs.contains_key(&VarName::new("obj"))); + assert!( + dae.continuous + .equations + .iter() + .all(|eq| !eq.origin.contains("binding equation for obj")) + ); +} + #[test] fn test_todae_rejects_constructor_field_selection_without_signature() { let mut flat = Model::new(); diff --git a/crates/rumoca-phase-dae/src/tests/root/tests_regressions.rs b/crates/rumoca-phase-dae/src/tests/root/tests_regressions.rs index d0d93fb7b..92ffe8cba 100644 --- a/crates/rumoca-phase-dae/src/tests/root/tests_regressions.rs +++ b/crates/rumoca-phase-dae/src/tests/root/tests_regressions.rs @@ -2,6 +2,9 @@ use super::*; mod assertion_actions_tests; mod clocked_tuple_tests; +mod matrix_product_compound_operands; +mod matrix_product_projection; +mod parameter_binding_tests; mod regression_more_tests; mod when_inactive_tests; mod when_lowering_tests; @@ -1480,6 +1483,117 @@ fn test_discrete_input_alias_chain_to_local_output_counts_as_local_unknown() { ); } +#[test] +fn test_discrete_input_metadata_records_only_untargeted_dae_inputs() { + let mut flat = Model::new(); + for name in ["untargeted.u", "targeted.u", "untargeted.y", "targeted.y"] { + flat.add_variable( + VarName::new(name), + flat::Variable { + name: VarName::new(name), + connected: true, + ..rumoca_ir_flat::Variable::empty_with_span(rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + )) + }, + ); + } + flat.add_variable( + VarName::new("unconnected.u"), + flat::Variable { + name: VarName::new("unconnected.u"), + connected: false, + ..rumoca_ir_flat::Variable::empty_with_span(rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + )) + }, + ); + + let mut dae = rumoca_ir_dae::Dae::default(); + dae.variables.discrete_valued.insert( + VarName::new("untargeted.u"), + dae_discrete_port_var("untargeted.u", rumoca_ir_dae::VariableCausality::Input), + ); + dae.variables.discrete_valued.insert( + VarName::new("targeted.u"), + dae_discrete_port_var("targeted.u", rumoca_ir_dae::VariableCausality::Input), + ); + dae.variables.discrete_valued.insert( + VarName::new("untargeted.y"), + dae_discrete_port_var("untargeted.y", rumoca_ir_dae::VariableCausality::Output), + ); + dae.variables.discrete_valued.insert( + VarName::new("targeted.y"), + dae_discrete_port_var("targeted.y", rumoca_ir_dae::VariableCausality::Output), + ); + dae.variables.discrete_valued.insert( + VarName::new("unconnected.u"), + dae_discrete_port_var("unconnected.u", rumoca_ir_dae::VariableCausality::Input), + ); + dae.variables.discrete_reals.insert( + VarName::new("real_port.u"), + dae_discrete_port_var("real_port.u", rumoca_ir_dae::VariableCausality::Input), + ); + flat.add_variable( + VarName::new("real_port.u"), + flat::Variable { + name: VarName::new("real_port.u"), + connected: true, + ..rumoca_ir_flat::Variable::empty_with_span(rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + )) + }, + ); + dae.discrete + .valued_updates + .push(test_discrete_assignment("targeted.u")); + dae.discrete + .valued_updates + .push(test_discrete_assignment("targeted.y")); + + crate::refresh_external_discrete_input_metadata(&mut dae, &flat); + + assert_eq!( + dae.metadata.discrete_input_names, + vec!["untargeted.u".to_string(), "untargeted.y".to_string()], + "only connected discrete-valued input/output variables without a DAE update target are external" + ); +} + +fn dae_discrete_port_var( + name: &str, + causality: rumoca_ir_dae::VariableCausality, +) -> rumoca_ir_dae::Variable { + rumoca_ir_dae::Variable { + name: VarName::new(name), + causality, + ..rumoca_ir_dae::Variable::empty_with_span(rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + )) + } +} + +fn test_discrete_assignment(lhs_name: &str) -> rumoca_ir_dae::Equation { + rumoca_ir_dae::Equation { + lhs: Some(VarName::new(lhs_name).into()), + rhs: Expression::Literal { + value: Literal::Boolean(true), + span: crate::test_support::test_span(), + }, + span: crate::test_support::test_span(), + origin: "test discrete assignment".to_string(), + scalar_count: 1, + } +} + #[test] fn test_connected_real_input_propagates_discrete_partition_from_peer() { let mut flat = Model::new(); @@ -1637,11 +1751,19 @@ fn test_when_clause_guard_for_clock_condition_uses_clock_tick_directly() { }), ); + let clock_span = crate::test_support::test_span(); let clock_call = Expression::FunctionCall { - name: VarName::new("Clock").into(), + name: rumoca_core::Reference::with_component_reference( + "Clock", + rumoca_core::ComponentReference::from_flat_segments( + "Clock", + clock_span, + Some(rumoca_core::BuiltinTypeIdentity::Clock.def_id()), + ), + ), args: vec![], is_constructor: false, - span: crate::test_support::test_span(), + span: clock_span, }; let previous_y2 = Expression::FunctionCall { name: VarName::new("previous").into(), @@ -1831,82 +1953,3 @@ fn collect_edge_guard_names(expr: &rumoca_core::Expression, names: &mut Vec {} } } - -/// A scalar parameter binding that references a known variable must keep -/// that reference; the record-field start alias selection used to graft the -/// LHS leaf onto the RHS and suffix-resolve to an unrelated variable -/// (`resistor.m = multiStar.mBasic` became the top-level `m`, so MSL -/// PowerConverters models evaluated `fill(300.15, m)` with the wrong phase -/// count). -#[test] -fn test_scalar_binding_to_known_variable_keeps_reference() { - let mut flat = Model::new(); - for (name, value) in [("m", 3), ("multiStar.mBasic", 1)] { - let var_name = VarName::new(name); - flat.add_variable( - var_name.clone(), - crate::test_support::with_component_ref(flat::Variable { - name: var_name, - variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), - binding: Some(Expression::Literal { - value: rumoca_core::Literal::Integer(value), - span: crate::test_support::test_span(), - }), - is_primitive: true, - ..rumoca_ir_flat::Variable::empty_with_span(rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name(file!()), - 1, - 2, - )) - }), - ); - } - let target = VarName::new("multiStar.resistor.m"); - flat.add_variable( - target.clone(), - crate::test_support::with_component_ref(flat::Variable { - name: target.clone(), - variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), - binding: Some(Expression::VarRef { - name: rumoca_core::Reference::from_component_reference( - rumoca_core::component_reference_from_flat_name( - &VarName::new("multiStar.mBasic"), - crate::test_support::test_span(), - ) - .expect("fixture name must form a component reference"), - ), - subscripts: vec![], - span: crate::test_support::test_span(), - }), - is_primitive: true, - ..rumoca_ir_flat::Variable::empty_with_span(rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name(file!()), - 1, - 2, - )) - }), - ); - - let dae = to_dae_with_options( - &flat, - ToDaeOptions { - error_on_unbalanced: false, - }, - ) - .expect("scalar parameter bindings should convert"); - - let var = dae - .variables - .parameters - .get(&target) - .expect("target parameter should exist in DAE"); - let start = var.start.as_ref().expect("binding becomes parameter start"); - let Expression::VarRef { name, .. } = start else { - panic!("expected a variable reference start, got {start:?}"); - }; - assert_eq!( - name.var_name().as_str(), - "multiStar.mBasic", - "binding reference must not be grafted onto an unrelated variable" - ); -} diff --git a/crates/rumoca-phase-dae/src/tests/root/tests_regressions/matrix_product_compound_operands.rs b/crates/rumoca-phase-dae/src/tests/root/tests_regressions/matrix_product_compound_operands.rs new file mode 100644 index 000000000..1385b02fc --- /dev/null +++ b/crates/rumoca-phase-dae/src/tests/root/tests_regressions/matrix_product_compound_operands.rs @@ -0,0 +1,254 @@ +use super::matrix_product_projection::{binary, declare_array, literal_subscripts}; +use super::*; + +fn assert_compound_dot_term(expr: &Expression, row: i64, inner: i64) { + let Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs, + rhs, + .. + } = expr + else { + panic!("expected compound dot term, got {expr:?}"); + }; + let Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: lhs_a, + rhs: lhs_b, + .. + } = lhs.as_ref() + else { + panic!("expected matrix inner sum, got {lhs:?}"); + }; + assert_eq!(literal_subscripts(lhs_a), Some(("A", vec![row, inner]))); + assert_eq!(literal_subscripts(lhs_b), Some(("B", vec![row, inner]))); + let Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: rhs_c, + rhs: rhs_d, + .. + } = rhs.as_ref() + else { + panic!("expected vector inner sum, got {rhs:?}"); + }; + assert_eq!(literal_subscripts(rhs_c), Some(("C", vec![inner]))); + assert_eq!(literal_subscripts(rhs_d), Some(("D", vec![inner]))); +} + +#[test] +fn test_todae_projects_bare_compound_array_operands_as_complete_dots() { + let mut flat = Model::new(); + for (name, dims) in [ + ("A", [2, 2].as_slice()), + ("B", [2, 2].as_slice()), + ("C", [2].as_slice()), + ("D", [2].as_slice()), + ("Y", [2].as_slice()), + ] { + declare_array(&mut flat, name, dims); + } + let add = |lhs, rhs| binary(rumoca_core::OpBinary::Add, lhs, rhs); + let product = binary( + rumoca_core::OpBinary::Mul, + add(make_structured_var_ref("A"), make_structured_var_ref("B")), + add(make_structured_var_ref("C"), make_structured_var_ref("D")), + ); + flat.add_equation(flat::Equation { + residual: binary( + rumoca_core::OpBinary::Sub, + Expression::Index { + base: Box::new(make_structured_var_ref("Y")), + subscripts: vec![rumoca_core::Subscript::Colon { + span: crate::test_support::test_span(), + }], + span: crate::test_support::test_span(), + }, + product, + ), + span: crate::test_support::test_span(), + origin: flat::EquationOrigin::ComponentEquation { + component: "CompoundArrayProduct".to_string(), + }, + scalar_count: 2, + }); + + let dae = to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("bare compound array operands must project as a matrix-vector product"); + + assert_eq!(dae.continuous.equations.len(), 2); + for (lane, equation) in dae.continuous.equations.iter().enumerate() { + let row = i64::try_from(lane + 1).expect("two lanes fit i64"); + let Expression::Binary { lhs, rhs, .. } = &equation.rhs else { + panic!("expected scalar residual"); + }; + assert_eq!(literal_subscripts(lhs), Some(("Y", vec![row]))); + let Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: first, + rhs: second, + .. + } = rhs.as_ref() + else { + panic!("lane {row} must contain the complete two-term dot, got {rhs:?}"); + }; + assert_compound_dot_term(first, row, 1); + assert_compound_dot_term(second, row, 2); + } +} + +#[test] +fn test_todae_preserves_scalar_scaling_of_vector_scalar_division() { + for scalar_on_left in [true, false] { + let mut flat = Model::new(); + declare_array(&mut flat, "V", &[2]); + declare_array(&mut flat, "s", &[]); + declare_array(&mut flat, "Y", &[2]); + let vector = Expression::Index { + base: Box::new(make_structured_var_ref("V")), + subscripts: vec![rumoca_core::Subscript::Colon { + span: crate::test_support::test_span(), + }], + span: crate::test_support::test_span(), + }; + let quotient = binary( + rumoca_core::OpBinary::Div, + vector, + make_structured_var_ref("s"), + ); + let scalar = Expression::Literal { + value: Literal::Real(2.0), + span: crate::test_support::test_span(), + }; + let product = if scalar_on_left { + binary(rumoca_core::OpBinary::Mul, scalar, quotient) + } else { + binary(rumoca_core::OpBinary::Mul, quotient, scalar) + }; + flat.add_equation(flat::Equation { + residual: binary( + rumoca_core::OpBinary::Sub, + Expression::Index { + base: Box::new(make_structured_var_ref("Y")), + subscripts: vec![rumoca_core::Subscript::Colon { + span: crate::test_support::test_span(), + }], + span: crate::test_support::test_span(), + }, + product, + ), + span: crate::test_support::test_span(), + origin: flat::EquationOrigin::ComponentEquation { + component: "ArrayScalarDivision".to_string(), + }, + scalar_count: 2, + }); + + let dae = to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("scalar scaling of vector/scalar division is legal ARR-030 input"); + + assert_eq!(dae.continuous.equations.len(), 2); + for (lane, equation) in dae.continuous.equations.iter().enumerate() { + let index = i64::try_from(lane + 1).expect("two lanes fit i64"); + let Expression::Binary { lhs, rhs, .. } = &equation.rhs else { + panic!("expected scalar residual"); + }; + assert_eq!(literal_subscripts(lhs), Some(("Y", vec![index]))); + let Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs, + rhs, + .. + } = rhs.as_ref() + else { + panic!("expected preserved scalar product, got {rhs:?}"); + }; + let (scalar, quotient) = if scalar_on_left { + (lhs.as_ref(), rhs.as_ref()) + } else { + (rhs.as_ref(), lhs.as_ref()) + }; + assert!(matches!( + scalar, + Expression::Literal { + value: Literal::Real(2.0), + .. + } + )); + let Expression::Binary { + op: rumoca_core::OpBinary::Div, + lhs: dividend, + rhs: divisor, + .. + } = quotient + else { + panic!("expected indexed vector/scalar quotient, got {quotient:?}"); + }; + assert_eq!(literal_subscripts(dividend), Some(("V", vec![index]))); + assert_eq!(literal_subscripts(divisor), Some(("s", vec![]))); + } + } +} + +#[test] +fn test_todae_rejects_bare_divided_array_operands_in_matrix_product() { + let mut flat = Model::new(); + for (name, dims) in [ + ("A", [2, 2].as_slice()), + ("x", [2].as_slice()), + ("s", [].as_slice()), + ("t", [].as_slice()), + ("Y", [2].as_slice()), + ] { + declare_array(&mut flat, name, dims); + } + let divided = |array, scalar| { + binary( + rumoca_core::OpBinary::Div, + make_structured_var_ref(array), + make_structured_var_ref(scalar), + ) + }; + flat.add_equation(flat::Equation { + residual: binary( + rumoca_core::OpBinary::Sub, + Expression::Index { + base: Box::new(make_structured_var_ref("Y")), + subscripts: vec![rumoca_core::Subscript::Colon { + span: crate::test_support::test_span(), + }], + span: crate::test_support::test_span(), + }, + binary( + rumoca_core::OpBinary::Mul, + divided("A", "s"), + divided("x", "t"), + ), + ), + span: crate::test_support::test_span(), + origin: flat::EquationOrigin::ComponentEquation { + component: "BareDividedArrayProduct".to_string(), + }, + scalar_count: 2, + }); + + let error = to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect_err("bare divided array operands must not be scalarized lane-wise"); + + assert!(error.to_string().contains("unknown operand shape")); + assert_eq!(error.source_span(), Some(crate::test_support::test_span())); +} diff --git a/crates/rumoca-phase-dae/src/tests/root/tests_regressions/matrix_product_projection.rs b/crates/rumoca-phase-dae/src/tests/root/tests_regressions/matrix_product_projection.rs new file mode 100644 index 000000000..80c39be3a --- /dev/null +++ b/crates/rumoca-phase-dae/src/tests/root/tests_regressions/matrix_product_projection.rs @@ -0,0 +1,1925 @@ +use super::*; + +pub(super) fn declare_array(flat: &mut Model, name: &str, dims: &[i64]) { + flat.add_variable( + VarName::new(name), + crate::test_support::with_component_ref(flat::Variable { + name: VarName::new(name), + dims: dims.to_vec(), + is_primitive: true, + ..flat::Variable::empty_with_span(crate::test_support::test_span()) + }), + ); +} + +fn colon_vector(name: &str) -> Expression { + colon_array(name, 1) +} + +fn colon_array(name: &str, rank: usize) -> Expression { + Expression::Index { + base: Box::new(make_structured_var_ref(name)), + subscripts: (0..rank) + .map(|_| rumoca_core::Subscript::Colon { + span: crate::test_support::test_span(), + }) + .collect(), + span: crate::test_support::test_span(), + } +} + +fn row_slice(name: &str, row: i64) -> Expression { + Expression::Index { + base: Box::new(make_structured_var_ref(name)), + subscripts: vec![ + rumoca_core::Subscript::Index { + value: row, + span: crate::test_support::test_span(), + }, + rumoca_core::Subscript::Colon { + span: crate::test_support::test_span(), + }, + ], + span: crate::test_support::test_span(), + } +} + +fn multiply(lhs: Expression, rhs: Expression) -> Expression { + binary(rumoca_core::OpBinary::Mul, lhs, rhs) +} + +pub(super) fn binary(op: rumoca_core::OpBinary, lhs: Expression, rhs: Expression) -> Expression { + Expression::Binary { + op, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span: crate::test_support::test_span(), + } +} + +fn real(value: f64) -> Expression { + Expression::Literal { + value: Literal::Real(value), + span: crate::test_support::test_span(), + } +} + +fn builtin(function: rumoca_core::BuiltinFunction, args: Vec) -> Expression { + Expression::BuiltinCall { + function, + args, + span: crate::test_support::test_span(), + } +} + +fn add_equation(flat: &mut Model, lhs: Expression, rhs: Expression, scalar_count: usize) { + flat.add_equation(flat::Equation { + residual: Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span: crate::test_support::test_span(), + }, + span: crate::test_support::test_span(), + origin: flat::EquationOrigin::ComponentEquation { + component: "MatrixProductProjection".to_string(), + }, + scalar_count, + }); +} + +pub(super) fn literal_subscripts(expr: &Expression) -> Option<(&str, Vec)> { + let Expression::VarRef { + name, subscripts, .. + } = expr + else { + return None; + }; + let indices = subscripts + .iter() + .map(|subscript| match subscript { + rumoca_core::Subscript::Index { value, .. } => Some(*value), + rumoca_core::Subscript::Expr { expr, .. } => match expr.as_ref() { + Expression::Literal { + value: Literal::Integer(value), + .. + } => Some(*value), + _ => None, + }, + rumoca_core::Subscript::Colon { .. } => None, + }) + .collect::>>()?; + Some((name.as_str(), indices)) +} + +type ProductTerm<'a> = ((&'a str, Vec), (&'a str, Vec)); + +fn flatten_dot_terms<'a>(expr: &'a Expression, terms: &mut Vec>) -> bool { + match expr { + Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs, + rhs, + .. + } => flatten_dot_terms(lhs, terms) && flatten_dot_terms(rhs, terms), + Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs, + rhs, + .. + } => { + let (Some(lhs), Some(rhs)) = (literal_subscripts(lhs), literal_subscripts(rhs)) else { + return false; + }; + terms.push((lhs, rhs)); + true + } + _ => false, + } +} + +type LiteralProductTerm<'a> = ((&'a str, Vec), f64); + +fn flatten_literal_dot_terms<'a>( + expr: &'a Expression, + terms: &mut Vec>, +) -> bool { + match expr { + Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs, + rhs, + .. + } => flatten_literal_dot_terms(lhs, terms) && flatten_literal_dot_terms(rhs, terms), + Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs, + rhs, + .. + } => { + let Some(lhs) = literal_subscripts(lhs) else { + return false; + }; + let Expression::Literal { + value: Literal::Real(rhs), + .. + } = rhs.as_ref() + else { + return false; + }; + terms.push((lhs, *rhs)); + true + } + _ => false, + } +} + +fn flatten_builtin_dot_terms(expr: &Expression, terms: &mut Vec<(Vec, i64)>) -> bool { + match expr { + Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs, + rhs, + .. + } => flatten_builtin_dot_terms(lhs, terms) && flatten_builtin_dot_terms(rhs, terms), + Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs, + rhs, + .. + } => { + let Some(("rotation", lhs_indices)) = literal_subscripts(lhs) else { + return false; + }; + let Expression::Index { + base, subscripts, .. + } = rhs.as_ref() + else { + return false; + }; + if !matches!( + base.as_ref(), + Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Cross, + .. + } + ) { + return false; + } + let [rumoca_core::Subscript::Index { value, .. }] = subscripts.as_slice() else { + return false; + }; + terms.push((lhs_indices, *value)); + true + } + _ => false, + } +} + +fn residual_rhs(equation: &rumoca_ir_dae::Equation) -> &Expression { + let Expression::Binary { + op: rumoca_core::OpBinary::Sub, + rhs, + .. + } = &equation.rhs + else { + panic!("expected scalar residual, got {:?}", equation.rhs); + }; + rhs +} + +fn assert_projection_error(flat: &Model, expected: &str) { + let error = to_dae_with_options( + flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect_err("invalid matrix-product projection must fail closed"); + assert!( + error.to_string().contains(expected), + "expected `{expected}` in {error}" + ); + assert_eq!(error.source_span(), Some(crate::test_support::test_span())); +} + +fn expression_row_slice(name: &str, selector: Expression) -> Expression { + Expression::Index { + base: Box::new(make_structured_var_ref(name)), + subscripts: vec![ + rumoca_core::Subscript::Expr { + expr: Box::new(selector), + span: crate::test_support::test_span(), + }, + rumoca_core::Subscript::Colon { + span: crate::test_support::test_span(), + }, + ], + span: crate::test_support::test_span(), + } +} + +fn declare_dae_array(dae: &mut rumoca_ir_dae::Dae, name: &str, dims: &[i64]) { + let mut variable = + rumoca_ir_dae::Variable::new(VarName::new(name), crate::test_support::test_span()); + variable.dims = dims.to_vec(); + dae.variables + .algebraics + .insert(VarName::new(name), variable); +} + +fn dae_scaling_with_missing_base(variants: &[&str]) -> rumoca_ir_dae::Dae { + let mut dae = rumoca_ir_dae::Dae::new(); + for (name, dims) in [("x", vec![2]), ("y", vec![2])] + .into_iter() + .chain(variants.iter().map(|name| (*name, vec![]))) + { + declare_dae_array(&mut dae, name, &dims); + } + dae.continuous + .equations + .push(rumoca_ir_dae::Equation::residual_array( + binary( + rumoca_core::OpBinary::Sub, + colon_vector("y"), + multiply(make_structured_var_ref("gain"), colon_vector("x")), + ), + crate::test_support::test_span(), + "missing-base scaling", + 2, + )); + dae +} + +#[test] +fn test_todae_preserves_ordinary_scalar_product_operands() { + let cases = [ + builtin( + rumoca_core::BuiltinFunction::Sin, + vec![make_structured_var_ref("x")], + ), + Expression::Unary { + op: rumoca_core::OpUnary::Minus, + rhs: Box::new(make_structured_var_ref("x")), + span: crate::test_support::test_span(), + }, + Expression::If { + branches: vec![( + Expression::Literal { + value: Literal::Boolean(true), + span: crate::test_support::test_span(), + }, + make_structured_var_ref("x"), + )], + else_branch: Box::new(real(1.0)), + span: crate::test_support::test_span(), + }, + ]; + for operand in cases { + let mut flat = Model::new(); + declare_array(&mut flat, "x", &[]); + declare_array(&mut flat, "y", &[]); + add_equation( + &mut flat, + make_structured_var_ref("y"), + binary( + rumoca_core::OpBinary::Add, + real(1.0), + multiply(operand, make_structured_var_ref("x")), + ), + 1, + ); + let dae = to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("ordinary scalar products must retain the existing lowering path"); + assert!(matches!( + residual_rhs(&dae.continuous.equations[0]), + Expression::Binary { + op: rumoca_core::OpBinary::Add, + .. + } + )); + } +} + +#[test] +fn test_todae_preserves_scalar_product_for_selected_array_element_target() { + let mut flat = Model::new(); + declare_array(&mut flat, "x", &[]); + declare_array(&mut flat, "y", &[2]); + let lhs = Expression::VarRef { + name: VarName::new("y").into(), + subscripts: vec![rumoca_core::Subscript::Index { + value: 1, + span: crate::test_support::test_span(), + }], + span: crate::test_support::test_span(), + }; + add_equation( + &mut flat, + lhs, + multiply( + builtin( + rumoca_core::BuiltinFunction::Sin, + vec![make_structured_var_ref("x")], + ), + make_structured_var_ref("x"), + ), + 1, + ); + + let dae = to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("a selected array element is a scalar projection target"); + let Expression::Binary { lhs, rhs, .. } = residual_rhs(&dae.continuous.equations[0]) else { + panic!("expected scalar multiplication"); + }; + let Expression::BuiltinCall { args, .. } = lhs.as_ref() else { + panic!("expected scalar sin operand"); + }; + assert_eq!(literal_subscripts(&args[0]), Some(("x", vec![]))); + assert_eq!(literal_subscripts(rhs), Some(("x", vec![]))); +} + +#[test] +fn test_todae_uses_row_major_lane_for_multidimensional_selected_target() { + let mut flat = Model::new(); + declare_array(&mut flat, "C", &[2, 2]); + declare_array(&mut flat, "x", &[4]); + let lhs = Expression::VarRef { + name: VarName::new("C").into(), + subscripts: vec![ + rumoca_core::Subscript::Index { + value: 2, + span: crate::test_support::test_span(), + }, + rumoca_core::Subscript::Index { + value: 1, + span: crate::test_support::test_span(), + }, + ], + span: crate::test_support::test_span(), + }; + let rhs = Expression::ArrayComprehension { + expr: Box::new(Expression::VarRef { + name: VarName::new("x").into(), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(make_structured_var_ref("i")), + span: crate::test_support::test_span(), + }], + span: crate::test_support::test_span(), + }), + indices: vec![rumoca_core::ComprehensionIndex { + name: "i".to_string(), + range: Expression::Range { + start: Box::new(Expression::Literal { + value: Literal::Integer(1), + span: crate::test_support::test_span(), + }), + step: None, + end: Box::new(Expression::Literal { + value: Literal::Integer(4), + span: crate::test_support::test_span(), + }), + span: crate::test_support::test_span(), + }, + }], + filter: None, + span: crate::test_support::test_span(), + }; + add_equation(&mut flat, lhs, rhs, 1); + + let dae = to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("a multidimensional selected target must retain its row-major lane"); + assert_eq!( + literal_subscripts(residual_rhs(&dae.continuous.equations[0])), + Some(("x", vec![3])) + ); +} + +#[test] +fn test_todae_preserves_derivative_vector_scalar_scaling() { + let mut flat = Model::new(); + declare_array(&mut flat, "x", &[2]); + add_equation( + &mut flat, + builtin(rumoca_core::BuiltinFunction::Der, vec![colon_vector("x")]), + multiply(real(2.0), colon_vector("x")), + 2, + ); + + let dae = to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("derivative vector scaling must retain lane scalarization"); + + assert_eq!(dae.continuous.equations.len(), 2); + for (lane, equation) in dae.continuous.equations.iter().enumerate() { + let Expression::Binary { lhs, .. } = &equation.rhs else { + panic!("expected residual"); + }; + let Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Der, + args, + .. + } = lhs.as_ref() + else { + panic!("expected derivative lhs, got {lhs:?}"); + }; + let index = i64::try_from(lane + 1).expect("two lanes fit i64"); + assert_eq!(literal_subscripts(&args[0]), Some(("x", vec![index]))); + let Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs, + rhs, + .. + } = residual_rhs(equation) + else { + panic!("expected scalar scaling rhs"); + }; + assert!(matches!( + lhs.as_ref(), + Expression::Literal { + value: Literal::Real(2.0), + .. + } + )); + assert_eq!(literal_subscripts(rhs), Some(("x", vec![index]))); + } +} + +#[test] +fn test_todae_preserves_compound_derivative_vector_target() { + let mut flat = Model::new(); + declare_array(&mut flat, "x", &[2]); + add_equation( + &mut flat, + binary( + rumoca_core::OpBinary::Add, + builtin(rumoca_core::BuiltinFunction::Der, vec![colon_vector("x")]), + colon_vector("x"), + ), + multiply(real(2.0), colon_vector("x")), + 2, + ); + + let dae = to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("shape-preserving compound derivative lhs must scalarize by lane"); + + assert_eq!(dae.continuous.equations.len(), 2); + for (lane, equation) in dae.continuous.equations.iter().enumerate() { + let index = i64::try_from(lane + 1).expect("two lanes fit i64"); + let Expression::Binary { lhs, .. } = &equation.rhs else { + panic!("expected residual"); + }; + let Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: derivative, + rhs: current, + .. + } = lhs.as_ref() + else { + panic!("expected compound derivative lhs, got {lhs:?}"); + }; + let Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Der, + args, + .. + } = derivative.as_ref() + else { + panic!("expected derivative term"); + }; + assert_eq!(literal_subscripts(&args[0]), Some(("x", vec![index]))); + assert_eq!(literal_subscripts(current), Some(("x", vec![index]))); + } +} + +#[test] +fn test_todae_projects_matrix_product_for_derivative_vector_target() { + let mut dae = rumoca_ir_dae::Dae::new(); + declare_dae_array(&mut dae, "A", &[3, 3]); + declare_dae_array(&mut dae, "x", &[3]); + declare_dae_array(&mut dae, "y", &[3]); + dae.continuous + .equations + .push(rumoca_ir_dae::Equation::residual_array( + binary( + rumoca_core::OpBinary::Sub, + builtin(rumoca_core::BuiltinFunction::Der, vec![colon_vector("y")]), + multiply( + builtin( + rumoca_core::BuiltinFunction::Transpose, + vec![make_structured_var_ref("A")], + ), + colon_vector("x"), + ), + ), + crate::test_support::test_span(), + "derivative matrix product", + 3, + )); + + scalarize_phantom_vector_equations(&mut dae) + .expect("derivative target shape must support matrix-product projection"); + + assert_eq!(dae.continuous.equations.len(), 3); + for (lane, equation) in dae.continuous.equations.iter().enumerate() { + let output_index = i64::try_from(lane + 1).expect("three lanes fit i64"); + let Expression::Binary { lhs, rhs, .. } = &equation.rhs else { + panic!("expected scalar residual"); + }; + let Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Der, + args, + .. + } = lhs.as_ref() + else { + panic!("expected derivative lhs, got {lhs:?}"); + }; + assert_eq!( + literal_subscripts(&args[0]), + Some(("y", vec![output_index])) + ); + let mut terms = Vec::new(); + assert!( + flatten_dot_terms(rhs, &mut terms), + "expected complete dot: {rhs:?}" + ); + assert_eq!( + terms, + (1_i64..=3) + .map(|row| (("A", vec![row, output_index]), ("x", vec![row]))) + .collect::>() + ); + } +} + +#[test] +fn test_todae_projects_matrix_product_nested_in_vector_addition() { + let mut dae = rumoca_ir_dae::Dae::new(); + declare_dae_array(&mut dae, "position", &[3]); + declare_dae_array(&mut dae, "rotation", &[3, 3]); + declare_dae_array(&mut dae, "offset", &[3]); + declare_dae_array(&mut dae, "target", &[3]); + dae.continuous + .equations + .push(rumoca_ir_dae::Equation::residual_array( + binary( + rumoca_core::OpBinary::Sub, + colon_vector("target"), + binary( + rumoca_core::OpBinary::Add, + colon_vector("position"), + multiply(make_structured_var_ref("rotation"), colon_vector("offset")), + ), + ), + crate::test_support::test_span(), + "compound matrix product", + 3, + )); + + scalarize_phantom_vector_equations(&mut dae) + .expect("shape-preserving addition must project its nested matrix product"); + + for (lane, equation) in dae.continuous.equations.iter().enumerate() { + let index = i64::try_from(lane + 1).expect("three lanes fit i64"); + let Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs, + rhs, + .. + } = residual_rhs(equation) + else { + panic!("expected scalar addition"); + }; + assert_eq!(literal_subscripts(lhs), Some(("position", vec![index]))); + let mut terms = Vec::new(); + assert!( + flatten_dot_terms(rhs, &mut terms), + "expected complete dot: {rhs:?}" + ); + assert_eq!( + terms, + (1_i64..=3) + .map(|inner| { (("rotation", vec![index, inner]), ("offset", vec![inner]),) }) + .collect::>() + ); + } +} + +#[test] +fn test_todae_projects_fixed_vector_builtin_inside_matrix_product() { + let mut dae = rumoca_ir_dae::Dae::new(); + for (name, dims) in [ + ("velocity", &[3][..]), + ("rotation", &[3, 3][..]), + ("omega", &[3][..]), + ("offset", &[3][..]), + ("target", &[3][..]), + ] { + declare_dae_array(&mut dae, name, dims); + } + let cross = builtin( + rumoca_core::BuiltinFunction::Cross, + vec![colon_vector("omega"), colon_vector("offset")], + ); + dae.continuous + .equations + .push(rumoca_ir_dae::Equation::residual_array( + binary( + rumoca_core::OpBinary::Sub, + colon_vector("target"), + binary( + rumoca_core::OpBinary::Add, + colon_vector("velocity"), + multiply(make_structured_var_ref("rotation"), cross), + ), + ), + crate::test_support::test_span(), + "builtin matrix product", + 3, + )); + + scalarize_phantom_vector_equations(&mut dae) + .expect("fixed-vector builtins have sufficient shape evidence for matrix projection"); + + for (lane, equation) in dae.continuous.equations.iter().enumerate() { + let row = i64::try_from(lane + 1).expect("three lanes fit i64"); + let Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs, + rhs, + .. + } = residual_rhs(equation) + else { + panic!("expected scalar addition"); + }; + assert_eq!(literal_subscripts(lhs), Some(("velocity", vec![row]))); + let mut terms = Vec::new(); + assert!(flatten_builtin_dot_terms(rhs, &mut terms)); + assert_eq!( + terms, + (1_i64..=3) + .map(|inner| (vec![row, inner], inner)) + .collect::>() + ); + } +} + +#[test] +fn test_todae_projects_function_sibling_and_nested_matrix_product() { + let mut flat = Model::new(); + declare_array(&mut flat, "A", &[3, 3]); + declare_array(&mut flat, "x", &[3]); + declare_array(&mut flat, "y", &[3]); + let mut function = + rumoca_core::Function::new("arrayFunction", crate::test_support::test_span()); + function.add_output( + rumoca_core::FunctionParam::new("result", "Real", crate::test_support::test_span()) + .with_dims(vec![3]), + ); + function.external = Some(rumoca_core::ExternalFunction { + language: "C".to_string(), + function_name: Some("array_function".to_string()), + output_name: Some("result".to_string()), + ..Default::default() + }); + flat.add_function(function); + let call = Expression::FunctionCall { + name: VarName::new("arrayFunction").into(), + args: Vec::new(), + is_constructor: false, + span: crate::test_support::test_span(), + }; + add_equation( + &mut flat, + colon_vector("y"), + binary( + rumoca_core::OpBinary::Add, + call, + multiply(make_structured_var_ref("A"), colon_vector("x")), + ), + 3, + ); + + let dae = to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("function sibling must not disable nested matrix projection"); + for (lane, equation) in dae.continuous.equations.iter().enumerate() { + let index = i64::try_from(lane + 1).expect("three lanes fit i64"); + let Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs, + rhs, + .. + } = residual_rhs(equation) + else { + panic!("expected projected addition"); + }; + let Expression::Index { subscripts, .. } = lhs.as_ref() else { + panic!("expected indexed function output"); + }; + assert!( + matches!(subscripts.as_slice(), [rumoca_core::Subscript::Index { value, .. }] if *value == index) + ); + let mut terms = Vec::new(); + assert!(flatten_dot_terms(rhs, &mut terms)); + assert_eq!(terms.len(), 3); + } +} + +#[test] +fn test_todae_rejects_compound_vector_rhs_for_scalar_target() { + for explicit_slice in [true, false] { + let mut flat = Model::new(); + declare_array(&mut flat, "A", &[3, 3]); + declare_array(&mut flat, "x", &[3]); + declare_array(&mut flat, "position", &[3]); + declare_array(&mut flat, "y", &[2]); + let lhs = Expression::VarRef { + name: VarName::new("y").into(), + subscripts: vec![rumoca_core::Subscript::Index { + value: 2, + span: crate::test_support::test_span(), + }], + span: crate::test_support::test_span(), + }; + let vector = |name| { + if explicit_slice { + colon_vector(name) + } else { + make_structured_var_ref(name) + } + }; + add_equation( + &mut flat, + lhs, + binary( + rumoca_core::OpBinary::Add, + vector("position"), + multiply(make_structured_var_ref("A"), vector("x")), + ), + 1, + ); + assert_projection_error(&flat, "result shape mismatch"); + } +} + +#[test] +fn test_todae_projects_matrix_product_for_derivative_matrix_target() { + let mut dae = rumoca_ir_dae::Dae::new(); + declare_dae_array(&mut dae, "A", &[2, 3]); + declare_dae_array(&mut dae, "B", &[3, 2]); + declare_dae_array(&mut dae, "C", &[2, 2]); + dae.continuous + .equations + .push(rumoca_ir_dae::Equation::residual_array( + binary( + rumoca_core::OpBinary::Sub, + builtin(rumoca_core::BuiltinFunction::Der, vec![colon_array("C", 2)]), + multiply(make_structured_var_ref("A"), make_structured_var_ref("B")), + ), + crate::test_support::test_span(), + "derivative matrix product", + 4, + )); + + scalarize_phantom_vector_equations(&mut dae) + .expect("derivative matrix target must project every result cell"); + + assert_eq!(dae.continuous.equations.len(), 4); + for (lane, equation) in dae.continuous.equations.iter().enumerate() { + let row = i64::try_from(lane / 2 + 1).expect("row fits i64"); + let column = i64::try_from(lane % 2 + 1).expect("column fits i64"); + let Expression::Binary { lhs, rhs, .. } = &equation.rhs else { + panic!("expected scalar residual"); + }; + let Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Der, + args, + .. + } = lhs.as_ref() + else { + panic!("expected derivative lhs, got {lhs:?}"); + }; + assert_eq!(literal_subscripts(&args[0]), Some(("C", vec![row, column]))); + let mut terms = Vec::new(); + assert!( + flatten_dot_terms(rhs, &mut terms), + "expected complete dot: {rhs:?}" + ); + assert_eq!( + terms, + (1_i64..=3) + .map(|inner| (("A", vec![row, inner]), ("B", vec![inner, column]))) + .collect::>() + ); + } +} + +#[test] +fn test_scalarizer_rejects_unknown_target_for_proven_array_product() { + let mut dae = rumoca_ir_dae::Dae::new(); + declare_dae_array(&mut dae, "A", &[2, 2]); + declare_dae_array(&mut dae, "x", &[2]); + dae.continuous + .equations + .push(rumoca_ir_dae::Equation::residual_array( + binary( + rumoca_core::OpBinary::Sub, + make_structured_var_ref("missing_target"), + multiply(colon_array("A", 2), colon_vector("x")), + ), + crate::test_support::test_span(), + "unknown matrix-product target", + 2, + )); + + let error = scalarize_phantom_vector_equations(&mut dae) + .expect_err("proven array product must not use a same-lane unknown-target fallback"); + assert!(error.to_string().contains("unknown target shape")); + assert_eq!(error.source_span(), Some(crate::test_support::test_span())); +} + +#[test] +fn test_scalarizer_rejects_unknown_target_for_compound_rhs_with_matrix_product() { + let mut dae = rumoca_ir_dae::Dae::new(); + declare_dae_array(&mut dae, "position", &[2]); + declare_dae_array(&mut dae, "A", &[2, 2]); + declare_dae_array(&mut dae, "x", &[2]); + dae.continuous + .equations + .push(rumoca_ir_dae::Equation::residual_array( + binary( + rumoca_core::OpBinary::Sub, + make_structured_var_ref("missing_target"), + binary( + rumoca_core::OpBinary::Add, + colon_vector("position"), + multiply(colon_array("A", 2), colon_vector("x")), + ), + ), + crate::test_support::test_span(), + "unknown compound matrix-product target", + 2, + )); + + let error = scalarize_phantom_vector_equations(&mut dae) + .expect_err("compound matrix-product RHS must reject an unknown target shape"); + assert!(error.to_string().contains("unknown target shape")); + assert_eq!(error.source_span(), Some(crate::test_support::test_span())); +} + +#[test] +fn test_todae_projects_transposed_matrix_vector_rows_as_three_term_dots() { + let mut flat = Model::new(); + declare_array(&mut flat, "A", &[3, 3]); + declare_array(&mut flat, "x", &[3]); + declare_array(&mut flat, "y", &[3]); + + let transpose_a = Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Transpose, + args: vec![make_structured_var_ref("A")], + span: crate::test_support::test_span(), + }; + add_equation( + &mut flat, + colon_vector("y"), + multiply(transpose_a, colon_vector("x")), + 3, + ); + + let dae = to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("shape-correct matrix-vector product should reach finalized DAE"); + + assert_eq!(dae.continuous.equations.len(), 3); + for (lane, equation) in dae.continuous.equations.iter().enumerate() { + let Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs, + rhs, + .. + } = &equation.rhs + else { + panic!("expected scalar residual, got {:?}", equation.rhs); + }; + let output_index = i64::try_from(lane + 1).expect("three lanes fit i64"); + assert_eq!(literal_subscripts(lhs), Some(("y", vec![output_index]))); + + let mut terms = Vec::new(); + assert!( + flatten_dot_terms(rhs, &mut terms), + "DAE lane {} must be a complete dot product, got {rhs:?}", + lane + 1 + ); + let expected = (1_i64..=3) + .map(|row| (("A", vec![row, output_index]), ("x", vec![row]))) + .collect::>(); + assert_eq!( + terms, + expected, + "DAE lane {} must contain every inner-dimension term", + lane + 1 + ); + } +} + +#[test] +fn test_todae_projects_indexed_vector_matrix_columns_as_three_term_dots() { + let mut flat = Model::new(); + declare_array(&mut flat, "source", &[2, 3]); + declare_array(&mut flat, "B", &[3, 2]); + declare_array(&mut flat, "y", &[2]); + add_equation( + &mut flat, + colon_vector("y"), + multiply(row_slice("source", 1), make_structured_var_ref("B")), + 2, + ); + + let dae = to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("indexed vector-matrix product should lower"); + + for (lane, equation) in dae.continuous.equations.iter().enumerate() { + let column = i64::try_from(lane + 1).expect("two lanes fit i64"); + let Expression::Binary { lhs, rhs, .. } = &equation.rhs else { + panic!("expected residual"); + }; + assert_eq!(literal_subscripts(lhs), Some(("y", vec![column]))); + let mut terms = Vec::new(); + assert!(flatten_dot_terms(rhs, &mut terms), "got {rhs:?}"); + assert_eq!( + terms, + (1_i64..=3) + .map(|inner| (("source", vec![1, inner]), ("B", vec![inner, column]))) + .collect::>() + ); + } +} + +#[test] +fn test_todae_projects_matrix_matrix_cells_as_three_term_dots() { + let mut flat = Model::new(); + declare_array(&mut flat, "A", &[2, 3]); + declare_array(&mut flat, "B", &[3, 2]); + declare_array(&mut flat, "C", &[2, 2]); + add_equation( + &mut flat, + colon_array("C", 2), + multiply(colon_array("A", 2), colon_array("B", 2)), + 4, + ); + + let dae = to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("matrix-matrix product should lower"); + + for (lane, equation) in dae.continuous.equations.iter().enumerate() { + let row = i64::try_from(lane / 2 + 1).expect("row fits i64"); + let column = i64::try_from(lane % 2 + 1).expect("column fits i64"); + let Expression::Binary { lhs, rhs, .. } = &equation.rhs else { + panic!("expected residual"); + }; + assert_eq!(literal_subscripts(lhs), Some(("C", vec![row, column]))); + let mut terms = Vec::new(); + assert!(flatten_dot_terms(rhs, &mut terms), "got {rhs:?}"); + assert_eq!( + terms, + (1_i64..=3) + .map(|inner| (("A", vec![row, inner]), ("B", vec![inner, column]))) + .collect::>() + ); + } +} + +#[test] +fn test_todae_projects_proven_scalar_scaling_forms() { + let mut flat = Model::new(); + declare_array(&mut flat, "x", &[3]); + declare_array(&mut flat, "gain", &[]); + for name in ["literal", "declared", "compound", "function", "right"] { + declare_array(&mut flat, name, &[3]); + } + let mut function = + rumoca_core::Function::new("scalarFunction", crate::test_support::test_span()); + function.add_input(rumoca_core::FunctionParam::new( + "u", + "Real", + crate::test_support::test_span(), + )); + function.add_output(rumoca_core::FunctionParam::new( + "y", + "Real", + crate::test_support::test_span(), + )); + function.external = Some(rumoca_core::ExternalFunction { + language: "C".to_string(), + function_name: Some("scalar_function".to_string()), + output_name: Some("y".to_string()), + ..Default::default() + }); + flat.add_function(function); + + let gain = make_structured_var_ref("gain"); + let cases = [ + ("literal", multiply(real(2.0), colon_vector("x"))), + ("declared", multiply(gain.clone(), colon_vector("x"))), + ( + "compound", + multiply( + binary(rumoca_core::OpBinary::Add, gain.clone(), real(1.0)), + colon_vector("x"), + ), + ), + ( + "function", + multiply( + Expression::FunctionCall { + name: VarName::new("scalarFunction").into(), + args: vec![gain.clone()], + is_constructor: false, + span: crate::test_support::test_span(), + }, + colon_vector("x"), + ), + ), + ("right", multiply(colon_vector("x"), gain)), + ]; + for (name, rhs) in cases { + add_equation(&mut flat, colon_vector(name), rhs, 3); + } + + let dae = to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("proven scalar factors should remain scalar"); + + assert_eq!(dae.continuous.equations.len(), 15); + for (index, equation) in dae.continuous.equations.iter().enumerate() { + let lane = i64::try_from(index % 3 + 1).expect("lane fits i64"); + let Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs, + rhs, + .. + } = residual_rhs(equation) + else { + panic!("expected scaling Mul"); + }; + let case = index / 3; + let (factor, vector) = if case == 4 { (rhs, lhs) } else { (lhs, rhs) }; + assert_eq!(literal_subscripts(vector), Some(("x", vec![lane]))); + match case { + 0 => assert!(matches!( + factor.as_ref(), + Expression::Literal { + value: Literal::Real(2.0), + .. + } + )), + 1 | 4 => assert_eq!(literal_subscripts(factor), Some(("gain", vec![]))), + 2 => assert!(matches!( + factor.as_ref(), + Expression::Binary { + op: rumoca_core::OpBinary::Add, + .. + } + )), + 3 => assert!( + matches!(factor.as_ref(), Expression::FunctionCall { name, .. } if name.as_str() == "scalarFunction") + ), + _ => unreachable!(), + } + } +} + +#[test] +fn test_todae_keeps_same_shape_mulelem_on_the_same_lane() { + let mut flat = Model::new(); + for name in ["a", "b", "y"] { + declare_array(&mut flat, name, &[3]); + } + add_equation( + &mut flat, + colon_vector("y"), + binary( + rumoca_core::OpBinary::MulElem, + colon_vector("a"), + colon_vector("b"), + ), + 3, + ); + + let dae = to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("same-shape MulElem should lower elementwise"); + for (lane, equation) in dae.continuous.equations.iter().enumerate() { + let Expression::Binary { + op: rumoca_core::OpBinary::MulElem, + lhs, + rhs, + .. + } = residual_rhs(equation) + else { + panic!("expected MulElem, got {:?}", equation.rhs); + }; + let lane = i64::try_from(lane + 1).expect("lane fits i64"); + assert_eq!(literal_subscripts(lhs), Some(("a", vec![lane]))); + assert_eq!(literal_subscripts(rhs), Some(("b", vec![lane]))); + } +} + +#[test] +fn test_todae_lowers_vector_mul_to_dot_only_for_scalar_targets() { + let mut flat = Model::new(); + declare_array(&mut flat, "a", &[3]); + declare_array(&mut flat, "b", &[3]); + declare_array(&mut flat, "colonDot", &[]); + declare_array(&mut flat, "bareDot", &[]); + add_equation( + &mut flat, + make_structured_var_ref("colonDot"), + multiply(colon_vector("a"), colon_vector("b")), + 1, + ); + add_equation( + &mut flat, + make_structured_var_ref("bareDot"), + multiply(make_structured_var_ref("a"), make_structured_var_ref("b")), + 1, + ); + + let dae = to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("scalar-target vector products should become dots"); + for equation in &dae.continuous.equations { + let mut terms = Vec::new(); + assert!(flatten_dot_terms(residual_rhs(equation), &mut terms)); + assert_eq!(terms.len(), 3); + } +} + +#[test] +fn test_todae_projects_nested_matrix_vector_scaling_in_both_orders() { + let mut flat = Model::new(); + declare_array(&mut flat, "A", &[2, 2]); + declare_array(&mut flat, "x", &[2]); + declare_array(&mut flat, "left", &[2]); + declare_array(&mut flat, "right", &[2]); + let product = || multiply(make_structured_var_ref("A"), colon_vector("x")); + add_equation( + &mut flat, + colon_vector("left"), + multiply(real(2.0), product()), + 2, + ); + add_equation( + &mut flat, + colon_vector("right"), + multiply(product(), real(2.0)), + 2, + ); + + let dae = to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("nested scaling should project only the array side"); + for (index, equation) in dae.continuous.equations.iter().enumerate() { + let Expression::Binary { lhs, rhs, .. } = residual_rhs(equation) else { + panic!("expected outer multiplication"); + }; + let (literal, product) = if index < 2 { (lhs, rhs) } else { (rhs, lhs) }; + assert!(matches!( + literal.as_ref(), + Expression::Literal { + value: Literal::Real(2.0), + .. + } + )); + let mut terms = Vec::new(); + assert!(flatten_dot_terms(product, &mut terms), "got {product:?}"); + let row = i64::try_from(index % 2 + 1).expect("row fits i64"); + assert_eq!( + terms, + (1_i64..=2) + .map(|inner| (("A", vec![row, inner]), ("x", vec![inner]))) + .collect::>() + ); + } +} + +#[test] +fn test_todae_rejects_inner_and_target_shape_mismatches() { + let mut inner = Model::new(); + declare_array(&mut inner, "A", &[2, 3]); + declare_array(&mut inner, "x", &[2]); + declare_array(&mut inner, "y", &[2]); + add_equation( + &mut inner, + colon_vector("y"), + multiply(make_structured_var_ref("A"), colon_vector("x")), + 2, + ); + assert_projection_error(&inner, "inner dimension mismatch"); + + let mut target = Model::new(); + declare_array(&mut target, "A", &[3, 3]); + declare_array(&mut target, "x", &[3]); + declare_array(&mut target, "y", &[2]); + add_equation( + &mut target, + colon_vector("y"), + multiply(make_structured_var_ref("A"), colon_vector("x")), + 2, + ); + assert_projection_error(&target, "result shape mismatch"); +} + +#[test] +fn test_todae_rejects_matrix_result_in_scalar_context_and_rank_three() { + let mut scalar = Model::new(); + declare_array(&mut scalar, "A", &[2, 2]); + declare_array(&mut scalar, "x", &[2]); + declare_array(&mut scalar, "s", &[]); + add_equation( + &mut scalar, + make_structured_var_ref("s"), + multiply(make_structured_var_ref("A"), make_structured_var_ref("x")), + 1, + ); + assert_projection_error(&scalar, "non-scalar result in scalar context"); + + let mut rank_three = Model::new(); + declare_array(&mut rank_three, "T", &[2, 2, 2]); + declare_array(&mut rank_three, "x", &[2]); + declare_array(&mut rank_three, "y", &[2]); + add_equation( + &mut rank_three, + colon_vector("y"), + multiply(make_structured_var_ref("T"), colon_vector("x")), + 2, + ); + assert_projection_error(&rank_three, "unsupported rank"); +} + +#[test] +fn test_todae_rejects_mulelem_shape_mismatch() { + let mut flat = Model::new(); + declare_array(&mut flat, "a", &[3]); + declare_array(&mut flat, "b", &[2]); + declare_array(&mut flat, "y", &[3]); + add_equation( + &mut flat, + colon_vector("y"), + binary( + rumoca_core::OpBinary::MulElem, + colon_vector("a"), + colon_vector("b"), + ), + 3, + ); + assert_projection_error(&flat, "elementwise shape mismatch"); +} + +#[test] +fn test_todae_rejects_dynamic_range_and_unknown_product_operands() { + let range = Expression::Range { + start: Box::new(Expression::Literal { + value: Literal::Integer(1), + span: crate::test_support::test_span(), + }), + step: None, + end: Box::new(Expression::Literal { + value: Literal::Integer(2), + span: crate::test_support::test_span(), + }), + span: crate::test_support::test_span(), + }; + let operands = [ + expression_row_slice("A", make_structured_var_ref("i")), + expression_row_slice("A", range), + Expression::VarRef { + name: VarName::new("unknown").into(), + subscripts: vec![rumoca_core::Subscript::Index { + value: 1, + span: crate::test_support::test_span(), + }], + span: crate::test_support::test_span(), + }, + ]; + for operand in operands { + let mut flat = Model::new(); + declare_array(&mut flat, "A", &[2, 3]); + declare_array(&mut flat, "x", &[3]); + declare_array(&mut flat, "z", &[]); + declare_array(&mut flat, "i", &[]); + add_equation( + &mut flat, + make_structured_var_ref("z"), + multiply(operand, colon_vector("x")), + 1, + ); + assert_projection_error(&flat, "unknown operand shape"); + } +} + +#[test] +fn test_todae_rejects_scalar_dot_with_two_dynamic_row_slices() { + let mut flat = Model::new(); + declare_array(&mut flat, "A", &[2, 3]); + declare_array(&mut flat, "B", &[2, 3]); + declare_array(&mut flat, "i", &[]); + declare_array(&mut flat, "j", &[]); + declare_array(&mut flat, "z", &[]); + add_equation( + &mut flat, + make_structured_var_ref("z"), + multiply( + expression_row_slice("A", make_structured_var_ref("i")), + expression_row_slice("B", make_structured_var_ref("j")), + ), + 1, + ); + assert_projection_error(&flat, "unknown operand shape"); +} + +#[test] +fn test_todae_rejects_scalar_dot_with_unary_wrapped_dynamic_row_slices() { + let mut flat = Model::new(); + declare_array(&mut flat, "A", &[2, 3]); + declare_array(&mut flat, "B", &[2, 3]); + declare_array(&mut flat, "i", &[]); + declare_array(&mut flat, "j", &[]); + declare_array(&mut flat, "z", &[]); + let negate = |expr| Expression::Unary { + op: rumoca_core::OpUnary::Minus, + rhs: Box::new(expr), + span: crate::test_support::test_span(), + }; + add_equation( + &mut flat, + make_structured_var_ref("z"), + multiply( + negate(expression_row_slice("A", make_structured_var_ref("i"))), + negate(expression_row_slice("B", make_structured_var_ref("j"))), + ), + 1, + ); + assert_projection_error(&flat, "unknown operand shape"); +} + +#[test] +fn test_todae_rejects_scalar_dot_with_scaled_dynamic_row_slices() { + let mut flat = Model::new(); + declare_array(&mut flat, "A", &[2, 3]); + declare_array(&mut flat, "B", &[2, 3]); + declare_array(&mut flat, "i", &[]); + declare_array(&mut flat, "j", &[]); + declare_array(&mut flat, "z", &[]); + add_equation( + &mut flat, + make_structured_var_ref("z"), + multiply( + multiply( + real(2.0), + expression_row_slice("A", make_structured_var_ref("i")), + ), + multiply( + real(3.0), + expression_row_slice("B", make_structured_var_ref("j")), + ), + ), + 1, + ); + assert_projection_error(&flat, "unknown operand shape"); +} + +#[test] +fn test_todae_preserves_dynamic_row_slice_reductions_in_scalar_products() { + for reduction in [ + rumoca_core::BuiltinFunction::Sum, + rumoca_core::BuiltinFunction::Product, + rumoca_core::BuiltinFunction::Min, + rumoca_core::BuiltinFunction::Max, + ] { + let mut flat = Model::new(); + declare_array(&mut flat, "A", &[2, 3]); + declare_array(&mut flat, "i", &[]); + declare_array(&mut flat, "z", &[]); + add_equation( + &mut flat, + make_structured_var_ref("z"), + multiply( + real(2.0), + builtin( + reduction, + vec![expression_row_slice("A", make_structured_var_ref("i"))], + ), + ), + 1, + ); + to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("scalar reductions must stay outside matrix-product projection"); + } +} + +#[test] +fn test_todae_preserves_scalar_scaled_sum_of_vector_slices() { + let mut flat = Model::new(); + declare_array(&mut flat, "diameters", &[1]); + declare_array(&mut flat, "dimensions", &[2]); + let range_slice = |start, end| Expression::VarRef { + name: VarName::new("dimensions").into(), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(Expression::Range { + start: Box::new(Expression::Literal { + value: Literal::Integer(start), + span: crate::test_support::test_span(), + }), + step: None, + end: Box::new(Expression::Literal { + value: Literal::Integer(end), + span: crate::test_support::test_span(), + }), + span: crate::test_support::test_span(), + }), + span: crate::test_support::test_span(), + }], + span: crate::test_support::test_span(), + }; + add_equation( + &mut flat, + make_structured_var_ref("diameters"), + multiply( + real(0.5), + binary( + rumoca_core::OpBinary::Add, + range_slice(1, 1), + range_slice(2, 2), + ), + ), + 1, + ); + + let dae = to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("scalar scaling of a vector sum is not matrix-product projection"); + + assert_eq!(dae.continuous.equations.len(), 1); + let Expression::Binary { + op: rumoca_core::OpBinary::Mul, + rhs, + .. + } = residual_rhs(&dae.continuous.equations[0]) + else { + panic!("expected preserved scalar scaling"); + }; + assert!(matches!( + rhs.as_ref(), + Expression::Binary { + op: rumoca_core::OpBinary::Add, + .. + } + )); +} + +#[test] +fn test_scalarizer_preserves_unknown_sum_with_scalar_only_descendant_product() { + let mut dae = rumoca_ir_dae::Dae::new(); + declare_dae_array(&mut dae, "gain", &[]); + declare_dae_array(&mut dae, "u", &[2]); + declare_dae_array(&mut dae, "y", &[2]); + let dynamic_range = Expression::Range { + start: Box::new(make_structured_var_ref("i")), + step: None, + end: Box::new(make_structured_var_ref("j")), + span: crate::test_support::test_span(), + }; + let dynamic_slice = Expression::VarRef { + name: VarName::new("u").into(), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(dynamic_range), + span: crate::test_support::test_span(), + }], + span: crate::test_support::test_span(), + }; + dae.continuous + .equations + .push(rumoca_ir_dae::Equation::residual_array( + binary( + rumoca_core::OpBinary::Sub, + colon_vector("y"), + multiply( + real(2.0), + binary( + rumoca_core::OpBinary::Add, + multiply( + make_structured_var_ref("gain"), + make_structured_var_ref("unknown_scalar"), + ), + dynamic_slice, + ), + ), + ), + crate::test_support::test_span(), + "scalar-only descendant product", + 2, + )); + + scalarize_phantom_vector_equations(&mut dae) + .expect("a scalar-only descendant multiplication is not a matrix-product candidate"); + assert_eq!(dae.continuous.equations.len(), 2); +} + +#[test] +fn test_scalarizer_rejects_scalar_scaled_unknown_sum_containing_matrix_product() { + let mut dae = rumoca_ir_dae::Dae::new(); + declare_dae_array(&mut dae, "A", &[2, 2]); + declare_dae_array(&mut dae, "x", &[2]); + declare_dae_array(&mut dae, "u", &[2]); + declare_dae_array(&mut dae, "y", &[2]); + let dynamic_range = Expression::Range { + start: Box::new(make_structured_var_ref("i")), + step: None, + end: Box::new(make_structured_var_ref("j")), + span: crate::test_support::test_span(), + }; + let dynamic_slice = Expression::VarRef { + name: VarName::new("u").into(), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(dynamic_range), + span: crate::test_support::test_span(), + }], + span: crate::test_support::test_span(), + }; + dae.continuous + .equations + .push(rumoca_ir_dae::Equation::residual_array( + binary( + rumoca_core::OpBinary::Sub, + colon_vector("y"), + multiply( + real(2.0), + binary( + rumoca_core::OpBinary::Add, + multiply(colon_array("A", 2), colon_vector("x")), + dynamic_slice, + ), + ), + ), + crate::test_support::test_span(), + "scalar-scaled unknown sum containing matrix product", + 2, + )); + + let error = scalarize_phantom_vector_equations(&mut dae) + .expect_err("unknown sum containing a matrix product must fail closed"); + assert!(error.to_string().contains("unknown operand shape")); + assert_eq!(error.source_span(), Some(crate::test_support::test_span())); +} + +#[test] +fn test_todae_projects_fill_vector_dot_inside_vector_scaling() { + let mut flat = Model::new(); + declare_array(&mut flat, "velocity", &[2]); + declare_array(&mut flat, "work", &[2]); + let fill = || { + builtin( + rumoca_core::BuiltinFunction::Fill, + vec![ + real(0.5), + Expression::Literal { + value: Literal::Integer(2), + span: crate::test_support::test_span(), + }, + ], + ) + }; + add_equation( + &mut flat, + colon_vector("work"), + multiply(fill(), multiply(colon_vector("velocity"), fill())), + 2, + ); + + let dae = to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("literal fill dimensions must prove the nested vector dot shape"); + + assert_eq!(dae.continuous.equations.len(), 2); + for equation in &dae.continuous.equations { + let Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs, + rhs, + .. + } = residual_rhs(equation) + else { + panic!( + "expected scaled dot product, got {:?}", + residual_rhs(equation) + ); + }; + assert!(matches!( + lhs.as_ref(), + Expression::Literal { + value: Literal::Real(0.5), + .. + } + )); + let mut terms = Vec::new(); + assert!( + flatten_literal_dot_terms(rhs, &mut terms), + "expected a complete fill-vector dot product, got {rhs:?}" + ); + assert_eq!( + terms, + vec![(("velocity", vec![1]), 0.5), (("velocity", vec![2]), 0.5)] + ); + } +} + +#[test] +fn test_todae_rejects_dynamic_row_slices_hidden_by_division_or_unknown_factor() { + for op in [rumoca_core::OpBinary::Div, rumoca_core::OpBinary::DivElem] { + let mut flat = Model::new(); + declare_array(&mut flat, "A", &[2, 3]); + declare_array(&mut flat, "B", &[2, 3]); + declare_array(&mut flat, "i", &[]); + declare_array(&mut flat, "j", &[]); + declare_array(&mut flat, "z", &[]); + add_equation( + &mut flat, + make_structured_var_ref("z"), + multiply( + binary( + op.clone(), + expression_row_slice("A", make_structured_var_ref("i")), + real(2.0), + ), + binary( + op, + expression_row_slice("B", make_structured_var_ref("j")), + real(3.0), + ), + ), + 1, + ); + assert_projection_error(&flat, "unknown operand shape"); + } + + let mut flat = Model::new(); + declare_array(&mut flat, "A", &[2, 3]); + declare_array(&mut flat, "i", &[]); + declare_array(&mut flat, "z", &[]); + add_equation( + &mut flat, + make_structured_var_ref("z"), + multiply( + multiply( + make_structured_var_ref("unknown"), + expression_row_slice("A", make_structured_var_ref("i")), + ), + real(2.0), + ), + 1, + ); + assert_projection_error(&flat, "unknown operand shape"); +} + +#[test] +fn test_todae_rejects_vectorized_builtins_hiding_dynamic_row_slices() { + let mut flat = Model::new(); + declare_array(&mut flat, "A", &[2, 3]); + declare_array(&mut flat, "B", &[2, 3]); + declare_array(&mut flat, "i", &[]); + declare_array(&mut flat, "j", &[]); + declare_array(&mut flat, "z", &[]); + add_equation( + &mut flat, + make_structured_var_ref("z"), + multiply( + builtin( + rumoca_core::BuiltinFunction::Sin, + vec![expression_row_slice("A", make_structured_var_ref("i"))], + ), + builtin( + rumoca_core::BuiltinFunction::Sin, + vec![expression_row_slice("B", make_structured_var_ref("j"))], + ), + ), + 1, + ); + assert_projection_error(&flat, "unknown operand shape"); +} + +#[test] +fn test_todae_rejects_array_valued_function_product_operand() { + let mut flat = Model::new(); + declare_array(&mut flat, "x", &[3]); + declare_array(&mut flat, "z", &[]); + let mut function = + rumoca_core::Function::new("arrayFunction", crate::test_support::test_span()); + function.add_output( + rumoca_core::FunctionParam::new("y", "Real", crate::test_support::test_span()) + .with_dims(vec![3]), + ); + function.external = Some(rumoca_core::ExternalFunction { + language: "C".to_string(), + function_name: Some("array_function".to_string()), + output_name: Some("y".to_string()), + ..Default::default() + }); + flat.add_function(function); + add_equation( + &mut flat, + make_structured_var_ref("z"), + multiply( + Expression::FunctionCall { + name: VarName::new("arrayFunction").into(), + args: Vec::new(), + is_constructor: false, + span: crate::test_support::test_span(), + }, + colon_vector("x"), + ), + 1, + ); + assert_projection_error(&flat, "array-valued function output cannot be projected"); +} + +#[test] +fn test_todae_rejects_bare_unknown_scaling_in_array_projection() { + let mut dae = dae_scaling_with_missing_base(&[]); + let error = scalarize_phantom_vector_equations(&mut dae) + .expect_err("missing declared scalar must fail during array projection"); + assert!( + error.to_string().contains("unknown operand shape"), + "{error}" + ); + assert_eq!(error.source_span(), Some(crate::test_support::test_span())); +} + +#[test] +fn test_scalarizer_rejects_phantom_base_as_scaling_scalar() { + let mut dae = dae_scaling_with_missing_base(&["gain[1]", "gain[2]"]); + let error = scalarize_phantom_vector_equations(&mut dae) + .expect_err("phantom base must not be guessed scalar during array projection"); + assert!( + error.to_string().contains("unknown operand shape"), + "{error}" + ); + assert_eq!(error.source_span(), Some(crate::test_support::test_span())); +} + +#[test] +fn test_todae_projects_zero_inner_matrix_product_to_zero() { + let mut flat = Model::new(); + declare_array(&mut flat, "A", &[2, 0]); + declare_array(&mut flat, "B", &[0, 2]); + declare_array(&mut flat, "C", &[2, 2]); + add_equation( + &mut flat, + colon_array("C", 2), + multiply(colon_array("A", 2), colon_array("B", 2)), + 4, + ); + let dae = to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("zero inner dimensions have the empty-sum value zero"); + assert_eq!(dae.continuous.equations.len(), 4); + for equation in &dae.continuous.equations { + assert!(matches!( + residual_rhs(equation), + Expression::Literal { + value: Literal::Real(0.0), + .. + } + )); + } +} + +#[test] +fn test_todae_rejects_negative_matrix_dimensions() { + let mut flat = Model::new(); + declare_array(&mut flat, "A", &[2, -1]); + declare_array(&mut flat, "x", &[-1]); + declare_array(&mut flat, "y", &[2]); + add_equation( + &mut flat, + colon_vector("y"), + multiply(make_structured_var_ref("A"), colon_vector("x")), + 2, + ); + assert_projection_error(&flat, "NegativeDimension"); +} diff --git a/crates/rumoca-phase-dae/src/tests/root/tests_regressions/parameter_binding_tests.rs b/crates/rumoca-phase-dae/src/tests/root/tests_regressions/parameter_binding_tests.rs new file mode 100644 index 000000000..8298c8492 --- /dev/null +++ b/crates/rumoca-phase-dae/src/tests/root/tests_regressions/parameter_binding_tests.rs @@ -0,0 +1,74 @@ +use super::*; + +/// A scalar parameter binding that references a known variable must keep that +/// reference instead of suffix-resolving to an unrelated variable. +#[test] +fn test_scalar_binding_to_known_variable_keeps_reference() { + let mut flat = Model::new(); + for (name, value) in [("m", 3), ("multiStar.mBasic", 1)] { + let var_name = VarName::new(name); + flat.add_variable( + var_name.clone(), + crate::test_support::with_component_ref(flat::Variable { + name: var_name, + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(Expression::Literal { + value: rumoca_core::Literal::Integer(value), + span: crate::test_support::test_span(), + }), + is_primitive: true, + ..rumoca_ir_flat::Variable::empty_with_span(rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + )) + }), + ); + } + + let target = VarName::new("multiStar.resistor.m"); + flat.add_variable( + target.clone(), + crate::test_support::with_component_ref(flat::Variable { + name: target.clone(), + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(Expression::VarRef { + name: rumoca_core::Reference::from_component_reference( + rumoca_core::component_reference_from_flat_name( + &VarName::new("multiStar.mBasic"), + crate::test_support::test_span(), + ) + .expect("fixture name must form a component reference"), + ), + subscripts: vec![], + span: crate::test_support::test_span(), + }), + is_primitive: true, + ..rumoca_ir_flat::Variable::empty_with_span(rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + )) + }), + ); + + let dae = to_dae_with_options( + &flat, + ToDaeOptions { + error_on_unbalanced: false, + }, + ) + .expect("scalar parameter bindings should convert"); + let start = dae + .variables + .parameters + .get(&target) + .expect("target parameter should exist in DAE") + .start + .as_ref() + .expect("binding becomes parameter start"); + let Expression::VarRef { name, .. } = start else { + panic!("expected a variable reference start, got {start:?}"); + }; + assert_eq!(name.var_name().as_str(), "multiStar.mBasic"); +} diff --git a/crates/rumoca-phase-dae/src/tests/root/tests_regressions/regression_more_tests.rs b/crates/rumoca-phase-dae/src/tests/root/tests_regressions/regression_more_tests.rs index e10060661..914094e6a 100644 --- a/crates/rumoca-phase-dae/src/tests/root/tests_regressions/regression_more_tests.rs +++ b/crates/rumoca-phase-dae/src/tests/root/tests_regressions/regression_more_tests.rs @@ -1,5 +1,133 @@ use super::*; +fn expandable_projection_fixture(lanes: &[i64]) -> Model { + let span = crate::test_support::test_span(); + let mut flat = Model::new(); + let mut aggregate = crate::test_support::with_component_ref(flat::Variable { + name: VarName::new("bus.cells.x"), + dims: vec![2], + is_primitive: true, + from_expandable_connector: true, + ..flat::Variable::empty_with_span(span) + }); + aggregate.component_ref.as_mut().unwrap().def_id = None; + flat.add_variable(aggregate.name.clone(), aggregate); + + for index in lanes { + let name = VarName::new(format!("bus.cells[{index}].x")); + let mut lane = crate::test_support::with_component_ref(flat::Variable { + name: name.clone(), + is_primitive: true, + from_expandable_connector: true, + ..flat::Variable::empty_with_span(span) + }); + lane.component_ref.as_mut().unwrap().def_id = Some(rumoca_core::DefId::new( + 90_000 + u32::try_from(*index).unwrap(), + )); + flat.add_variable(name, lane); + } + flat +} + +#[test] +fn test_expandable_aggregate_projection_requires_complete_lane_domain() { + let complete = expandable_projection_fixture(&[1, 2]); + assert!( + expandable_aggregate_projection_names(&complete).contains(&VarName::new("bus.cells.x")) + ); + + let partial = expandable_projection_fixture(&[1]); + assert!( + !expandable_aggregate_projection_names(&partial).contains(&VarName::new("bus.cells.x")), + "a partial lane set cannot prove that the aggregate is only a projection" + ); +} + +#[test] +fn test_expandable_aggregate_projection_preserves_bound_or_defined_aggregate() { + let span = crate::test_support::test_span(); + let mut bound = expandable_projection_fixture(&[1, 2]); + bound + .variables + .get_mut(&VarName::new("bus.cells.x")) + .unwrap() + .binding = Some(Expression::Literal { + value: Literal::Real(1.0), + span, + }); + assert!(!expandable_aggregate_projection_names(&bound).contains(&VarName::new("bus.cells.x"))); + + let mut defined = expandable_projection_fixture(&[1, 2]); + defined.add_equation(flat::Equation { + residual: Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(Expression::VarRef { + name: VarName::new("bus.cells.x").into(), + subscripts: vec![], + span, + }), + rhs: Box::new(Expression::Literal { + value: Literal::Real(0.0), + span, + }), + span, + }, + span, + origin: flat::EquationOrigin::ComponentEquation { + component: "owner".to_string(), + }, + scalar_count: 2, + }); + assert!( + !expandable_aggregate_projection_names(&defined).contains(&VarName::new("bus.cells.x")), + "a non-connection LHS is an independent aggregate definition" + ); +} + +#[test] +fn test_expandable_aggregate_projection_rejects_distinct_declaration_identity() { + let mut same_file = expandable_projection_fixture(&[1, 2]); + same_file + .variables + .get_mut(&VarName::new("bus.cells[2].x")) + .unwrap() + .component_ref + .as_mut() + .unwrap() + .span = Span::from_offsets( + rumoca_core::SourceId::from_source_name("phase_dae_fixture.mo"), + 3, + 4, + ); + assert!( + !expandable_aggregate_projection_names(&same_file).contains(&VarName::new("bus.cells.x")), + "identical rendered paths from distinct declarations in one source must not be grouped" + ); + + let mut cross_file = expandable_projection_fixture(&[1, 2]); + cross_file + .variables + .get_mut(&VarName::new("bus.cells[2].x")) + .unwrap() + .component_ref + .as_mut() + .unwrap() + .span = Span::from_offsets( + rumoca_core::SourceId::from_source_name("other_declaration.mo"), + 1, + 2, + ); + assert!( + !expandable_aggregate_projection_names(&cross_file).contains(&VarName::new("bus.cells.x")), + "identical rendered paths from different source declarations must not be grouped" + ); +} + +// SPEC_0021 file-size exception: this root regression bucket temporarily +// groups DAE lowering regressions that share root-level fixtures. split plan: +// move member-call, scalar-count, and discrete-alias regressions into focused +// modules under tests_regressions/. + #[test] fn test_anchored_expandable_member_via_input_alias_is_not_interface_input() { // Reproduces Electrical.Cell bus pattern: @@ -1591,6 +1719,207 @@ fn test_infer_equation_scalar_count_connector_field_array_alias() { ); } +#[test] +fn test_infer_equation_scalar_count_indexed_component_field_lhs_uses_field_width() { + let mut flat = Model::new(); + for i in 1..=4 { + for field in ["C_outflow", "m_flow", "p", "h_outflow"] { + let name = format!("src[{i}].ports[1].{field}"); + flat.add_variable( + VarName::new(name.clone()), + flat::Variable { + name: VarName::new(name), + is_primitive: true, + ..flat::Variable::empty_with_span(crate::test_support::test_span()) + }, + ); + } + } + + let lhs = Expression::FieldAccess { + base: Box::new(Expression::FieldAccess { + base: Box::new(Expression::Index { + base: Box::new(Expression::VarRef { + name: VarName::new("src").into(), + subscripts: vec![], + span: crate::test_support::test_span(), + }), + subscripts: vec![rumoca_core::Subscript::Index { + value: 1, + span: crate::test_support::test_span(), + }], + span: crate::test_support::test_span(), + }), + field: "ports".to_string(), + span: crate::test_support::test_span(), + }), + field: "C_outflow".to_string(), + span: crate::test_support::test_span(), + }; + let residual = Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(lhs), + rhs: Box::new(Expression::Literal { + value: rumoca_core::Literal::Integer(0), + span: crate::test_support::test_span(), + }), + span: crate::test_support::test_span(), + }; + + let prefix_counts = build_prefix_counts(&flat); + let scalar_count = infer_equation_scalar_count(&residual, &flat, &prefix_counts); + assert_eq!( + scalar_count, 1, + "indexed component field LHS must use the selected field width, not the whole component prefix" + ); +} + +#[test] +fn test_infer_equation_scalar_count_repeated_indexed_component_leaf_is_scalar() { + let mut flat = Model::new(); + for i in 1..=3 { + for field in ["Goff", "Ron", "s", "unitCurrent", "unitVoltage", "v"] { + let name = format!("triac.triac[{i}].thyristor1.{field}"); + flat.add_variable( + VarName::new(name.clone()), + flat::Variable { + name: VarName::new(name), + is_primitive: true, + ..flat::Variable::empty_with_span(crate::test_support::test_span()) + }, + ); + } + } + + let indexed_component = Expression::FieldAccess { + base: Box::new(Expression::Index { + base: Box::new(Expression::VarRef { + name: VarName::new("triac").into(), + subscripts: vec![], + span: crate::test_support::test_span(), + }), + subscripts: vec![rumoca_core::Subscript::Index { + value: 1, + span: crate::test_support::test_span(), + }], + span: crate::test_support::test_span(), + }), + field: "thyristor1".to_string(), + span: crate::test_support::test_span(), + }; + let lhs = Expression::FieldAccess { + base: Box::new(indexed_component.clone()), + field: "v".to_string(), + span: crate::test_support::test_span(), + }; + let rhs = Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs: Box::new(Expression::FieldAccess { + base: Box::new(indexed_component), + field: "s".to_string(), + span: crate::test_support::test_span(), + }), + rhs: Box::new(Expression::FieldAccess { + base: Box::new(Expression::FieldAccess { + base: Box::new(Expression::Index { + base: Box::new(Expression::VarRef { + name: VarName::new("triac").into(), + subscripts: vec![], + span: crate::test_support::test_span(), + }), + subscripts: vec![rumoca_core::Subscript::Index { + value: 1, + span: crate::test_support::test_span(), + }], + span: crate::test_support::test_span(), + }), + field: "thyristor1".to_string(), + span: crate::test_support::test_span(), + }), + field: "unitCurrent".to_string(), + span: crate::test_support::test_span(), + }), + span: crate::test_support::test_span(), + }; + let residual = Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span: crate::test_support::test_span(), + }; + + let prefix_counts = build_prefix_counts(&flat); + let scalar_count = infer_equation_scalar_count(&residual, &flat, &prefix_counts); + assert_eq!( + scalar_count, 1, + "scalar leaf equation on a repeated indexed component must not inherit the full component aggregate width" + ); +} + +#[test] +fn test_infer_equation_scalar_count_repeated_indexed_connection_residual_is_scalar() { + let mut flat = Model::new(); + for i in 1..=3 { + for field in ["p.i", "n.i"] { + let name = format!("triac.triac[{i}].thyristor1.{field}"); + flat.add_variable( + VarName::new(name.clone()), + flat::Variable { + name: VarName::new(name), + is_primitive: true, + ..flat::Variable::empty_with_span(crate::test_support::test_span()) + }, + ); + } + } + + let pin_current = |pin: &str| Expression::FieldAccess { + base: Box::new(Expression::FieldAccess { + base: Box::new(Expression::FieldAccess { + base: Box::new(Expression::Index { + base: Box::new(Expression::VarRef { + name: VarName::new("triac").into(), + subscripts: vec![], + span: crate::test_support::test_span(), + }), + subscripts: vec![rumoca_core::Subscript::Index { + value: 1, + span: crate::test_support::test_span(), + }], + span: crate::test_support::test_span(), + }), + field: "thyristor1".to_string(), + span: crate::test_support::test_span(), + }), + field: pin.to_string(), + span: crate::test_support::test_span(), + }), + field: "i".to_string(), + span: crate::test_support::test_span(), + }; + let residual = Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(Expression::Literal { + value: rumoca_core::Literal::Integer(0), + span: crate::test_support::test_span(), + }), + rhs: Box::new(Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: Box::new(pin_current("p")), + rhs: Box::new(pin_current("n")), + span: crate::test_support::test_span(), + }), + span: crate::test_support::test_span(), + }; + + let prefix_counts = build_prefix_counts(&flat); + let scalar_count = infer_equation_scalar_count(&residual, &flat, &prefix_counts); + assert_eq!( + scalar_count, 1, + "connection residual over scalar leaves must not inherit repeated component aggregate width" + ); +} + #[test] fn test_infer_equation_scalar_count_record_prefix_uses_scalarized_children() { let mut flat = Model::new(); @@ -1647,6 +1976,57 @@ fn test_infer_equation_scalar_count_record_prefix_uses_scalarized_children() { ); } +#[test] +fn test_infer_equation_scalar_count_record_array_element_uses_field_width() { + let mut flat = Model::new(); + + for name in ["tf.z.re", "tf.z.im"] { + flat.add_variable( + VarName::new(name), + crate::test_support::with_component_ref(flat::Variable { + name: VarName::new(name), + dims: vec![2], + is_primitive: true, + ..rumoca_ir_flat::Variable::empty_with_span(rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + )) + }), + ); + } + + let residual = Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(Expression::Index { + base: Box::new(Expression::VarRef { + name: VarName::new("tf.z").into(), + subscripts: vec![], + span: crate::test_support::test_span(), + }), + subscripts: vec![rumoca_core::Subscript::Index { + value: 1, + span: crate::test_support::test_span(), + }], + span: crate::test_support::test_span(), + }), + rhs: Box::new(Expression::FunctionCall { + name: VarName::new("Cpx").into(), + args: vec![], + is_constructor: true, + span: crate::test_support::test_span(), + }), + span: crate::test_support::test_span(), + }; + + let prefix_counts = build_prefix_counts(&flat); + let scalar_count = infer_equation_scalar_count(&residual, &flat, &prefix_counts); + assert_eq!( + scalar_count, 2, + "a single record-array element assignment should count one scalar per primitive field" + ); +} + #[test] fn test_infer_equation_scalar_count_record_array_range_lhs_uses_full_slice_size() { let mut flat = Model::new(); @@ -1807,6 +2187,145 @@ fn test_infer_equation_scalar_count_structured_range_subscript_uses_slice_size() ); } +#[test] +fn test_infer_equation_scalar_count_matrix_column_slice_uses_preserved_dimension() { + let mut flat = Model::new(); + for (name, dims) in [("vehicle.leg_v_b", vec![3, 4]), ("vehicle.v_b", vec![3])] { + flat.add_variable( + VarName::new(name), + crate::test_support::with_component_ref(flat::Variable { + name: VarName::new(name), + dims, + is_primitive: true, + ..rumoca_ir_flat::Variable::empty_with_span(rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + )) + }), + ); + } + + let residual = Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(Expression::Index { + base: Box::new(Expression::VarRef { + name: VarName::new("vehicle.leg_v_b").into(), + subscripts: vec![], + span: crate::test_support::test_span(), + }), + subscripts: vec![ + rumoca_core::Subscript::Colon { + span: crate::test_support::test_span(), + }, + rumoca_core::Subscript::Expr { + expr: Box::new(Expression::Literal { + value: Literal::Integer(1), + span: crate::test_support::test_span(), + }), + span: crate::test_support::test_span(), + }, + ], + span: crate::test_support::test_span(), + }), + rhs: Box::new(Expression::VarRef { + name: VarName::new("vehicle.v_b").into(), + subscripts: vec![], + span: crate::test_support::test_span(), + }), + span: crate::test_support::test_span(), + }; + + let prefix_counts = build_prefix_counts(&flat); + let scalar_count = infer_equation_scalar_count(&residual, &flat, &prefix_counts); + assert_eq!( + scalar_count, 3, + "matrix column slices such as M[:, 1] should preserve the column height" + ); +} + +#[test] +fn test_infer_equation_scalar_count_qualified_parameter_range_subscript() { + let mut flat = Model::new(); + flat.add_variable( + VarName::new("pipe.n"), + crate::test_support::with_component_ref(flat::Variable { + name: VarName::new("pipe.n"), + is_primitive: true, + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(Expression::Literal { + value: Literal::Integer(2), + span: crate::test_support::test_span(), + }), + ..rumoca_ir_flat::Variable::empty_with_span(rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + )) + }), + ); + for name in ["pipe.m_flows", "pipe.flowModel.m_flows"] { + flat.add_variable( + VarName::new(name), + crate::test_support::with_component_ref(flat::Variable { + name: VarName::new(name), + dims: vec![3], + is_primitive: true, + ..rumoca_ir_flat::Variable::empty_with_span(rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + )) + }), + ); + } + + let range = || { + rumoca_core::Subscript::expr( + Box::new(Expression::Range { + start: Box::new(Expression::Literal { + value: Literal::Integer(1), + span: crate::test_support::test_span(), + }), + step: None, + end: Box::new(Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: Box::new(Expression::VarRef { + name: VarName::new("pipe.n").into(), + subscripts: vec![], + span: crate::test_support::test_span(), + }), + rhs: Box::new(Expression::Literal { + value: Literal::Integer(1), + span: crate::test_support::test_span(), + }), + span: crate::test_support::test_span(), + }), + span: crate::test_support::test_span(), + }), + crate::test_support::test_span(), + ) + }; + let residual = Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(Expression::VarRef { + name: VarName::new("pipe.m_flows").into(), + subscripts: vec![range()], + span: crate::test_support::test_span(), + }), + rhs: Box::new(Expression::VarRef { + name: VarName::new("pipe.flowModel.m_flows").into(), + subscripts: vec![range()], + span: crate::test_support::test_span(), + }), + span: crate::test_support::test_span(), + }; + + let prefix_counts = build_prefix_counts(&flat); + let scalar_count = infer_equation_scalar_count(&residual, &flat, &prefix_counts); + assert_eq!(scalar_count, 3); +} + #[test] fn test_infer_equation_scalar_count_record_array_range_uses_parameter_start_fallback() { let mut flat = Model::new(); @@ -1979,3 +2498,231 @@ fn test_infer_equation_scalar_count_record_array_range_with_scalarized_field_ind "record-array range LHS should infer array length from indexed scalarized fields" ); } + +fn projected_result_field_fixture(field: &str, dims: Vec) -> Model { + let span = crate::test_support::test_span(); + let function_def_id = rumoca_core::DefId::new(91_001); + let mut calc = rumoca_core::Function::new("Pkg.calc", span); + calc.def_id = Some(function_def_id); + calc.add_input( + rumoca_core::FunctionParam::new("params", "Pkg.WideParams", span) + .with_type_class(rumoca_core::ClassType::Record), + ); + calc.add_output( + rumoca_core::FunctionParam::new("result", "Pkg.CalcResult", span) + .with_type_class(rumoca_core::ClassType::Record), + ); + + let mut flat = Model::new(); + flat.add_function(calc); + for index in 0..48 { + let name = format!("params.field{index}"); + flat.add_variable( + VarName::new(name.clone()), + crate::test_support::with_component_ref(flat::Variable { + name: VarName::new(name), + is_primitive: true, + ..flat::Variable::empty_with_span(span) + }), + ); + } + for name in ["ird", "D", "Dinternal"] { + flat.add_variable( + VarName::new(name), + crate::test_support::with_component_ref(flat::Variable { + name: VarName::new(name), + is_primitive: true, + ..flat::Variable::empty_with_span(span) + }), + ); + } + let binding = projected_result_field_call(field, function_def_id, span); + let shape_name = format!("shape_cache.{field}"); + flat.add_variable( + VarName::new(shape_name.clone()), + crate::test_support::with_component_ref(flat::Variable { + name: VarName::new(shape_name), + dims, + binding: Some(binding), + binding_from_modification: true, + is_primitive: true, + ..flat::Variable::empty_with_span(span) + }), + ); + flat +} + +fn projected_result_field_call( + field: &str, + function_def_id: rumoca_core::DefId, + span: rumoca_core::Span, +) -> Expression { + let function_ref = rumoca_core::Reference::with_component_reference( + "Pkg.calc", + rumoca_core::ComponentReference { + local: false, + span, + parts: vec![ + rumoca_core::ComponentRefPart { + ident: "Pkg".to_string(), + span, + subs: vec![], + }, + rumoca_core::ComponentRefPart { + ident: "calc".to_string(), + span, + subs: vec![], + }, + ], + def_id: Some(function_def_id), + }, + ); + Expression::FieldAccess { + base: Box::new(Expression::FunctionCall { + name: function_ref, + args: vec![Expression::VarRef { + name: VarName::new("params").into(), + subscripts: vec![], + span, + }], + is_constructor: false, + span, + }), + field: field.to_string(), + span, + } +} + +fn subtraction(lhs: Expression, rhs: Expression, span: rumoca_core::Span) -> Expression { + Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span, + } +} + +#[test] +fn test_infer_scalar_count_scalar_field_function_call_ignores_record_argument_width() { + let span = crate::test_support::test_span(); + let function_def_id = rumoca_core::DefId::new(91_001); + let flat = projected_result_field_fixture("resistance", vec![]); + let projected = projected_result_field_call("resistance", function_def_id, span); + let residual = subtraction( + Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs: Box::new(Expression::VarRef { + name: VarName::new("ird").into(), + subscripts: vec![], + span, + }), + rhs: Box::new(projected), + span, + }, + subtraction( + Expression::VarRef { + name: VarName::new("D").into(), + subscripts: vec![], + span, + }, + Expression::VarRef { + name: VarName::new("Dinternal").into(), + subscripts: vec![], + span, + }, + span, + ), + span, + ); + + let prefix_counts = build_prefix_counts(&flat); + assert_eq!( + infer_equation_scalar_count(&residual, &flat, &prefix_counts), + 1 + ); +} + +#[test] +fn test_infer_scalar_count_array_field_function_call_preserves_field_width() { + let span = crate::test_support::test_span(); + let function_def_id = rumoca_core::DefId::new(91_001); + let flat = projected_result_field_fixture("curve", vec![3]); + let residual = subtraction( + projected_result_field_call("curve", function_def_id, span), + Expression::Literal { + value: Literal::Real(0.0), + span, + }, + span, + ); + + let prefix_counts = build_prefix_counts(&flat); + assert_eq!( + infer_equation_scalar_count(&residual, &flat, &prefix_counts), + 3 + ); +} + +#[test] +fn test_infer_scalar_count_indexed_function_result_field_is_scalar() { + let span = crate::test_support::test_span(); + let function_def_id = rumoca_core::DefId::new(91_001); + let flat = projected_result_field_fixture("curve", vec![3]); + let residual = subtraction( + Expression::Index { + base: Box::new(projected_result_field_call("curve", function_def_id, span)), + subscripts: vec![Subscript::generated_index(1, span)], + span, + }, + Expression::Literal { + value: Literal::Real(0.0), + span, + }, + span, + ); + + let prefix_counts = build_prefix_counts(&flat); + assert_eq!( + infer_equation_scalar_count(&residual, &flat, &prefix_counts), + 1 + ); +} + +#[test] +fn test_infer_function_result_field_conflicting_instance_shapes_remain_unknown() { + let span = crate::test_support::test_span(); + let function_def_id = rumoca_core::DefId::new(91_001); + let mut flat = projected_result_field_fixture("curve", vec![3]); + flat.add_variable( + VarName::new("shape_cache.conflicting_curve"), + crate::test_support::with_component_ref(flat::Variable { + name: VarName::new("shape_cache.conflicting_curve"), + dims: vec![4], + binding: Some(projected_result_field_call("curve", function_def_id, span)), + binding_from_modification: true, + is_primitive: true, + ..flat::Variable::empty_with_span(span) + }), + ); + + let metadata = build_prefix_counts(&flat); + let projected = projected_result_field_call("curve", function_def_id, span); + assert_eq!( + infer_expression_form(&projected, &flat, &metadata), + ExpressionForm::Other + ); +} + +#[test] +fn test_infer_function_result_field_dynamic_instance_shape_remains_unknown() { + let span = crate::test_support::test_span(); + let function_def_id = rumoca_core::DefId::new(91_001); + let flat = projected_result_field_fixture("curve", vec![-1]); + let metadata = build_prefix_counts(&flat); + let projected = projected_result_field_call("curve", function_def_id, span); + + assert_eq!( + infer_expression_form(&projected, &flat, &metadata), + ExpressionForm::Other + ); +} diff --git a/crates/rumoca-phase-dae/src/when_guard.rs b/crates/rumoca-phase-dae/src/when_guard.rs index b2dc8914e..b2b530ad6 100644 --- a/crates/rumoca-phase-dae/src/when_guard.rs +++ b/crates/rumoca-phase-dae/src/when_guard.rs @@ -66,7 +66,8 @@ impl BooleanAliasUnfolder<'_> { } let result = find_condition_definition_rhs(self.dae, name).and_then(|rhs| { let unfolded = self.rewrite_child(&rhs); - if unfolded.contains_relational_operator() { + if unfolded.contains_relational_operator() || is_clock_constructor_condition(&unfolded) + { Some(unfolded) } else { None @@ -150,7 +151,9 @@ fn is_direct_initial_condition(expr: &Expression) -> bool { fn is_clock_constructor_condition(expr: &Expression) -> bool { matches!( expr, - Expression::FunctionCall { name, .. } if name.last_segment() == "Clock" + Expression::FunctionCall { name, .. } + if name.target_def_id() + == Some(rumoca_core::BuiltinTypeIdentity::Clock.def_id()) ) } @@ -281,6 +284,51 @@ mod tests { } } + fn function_call_with_target( + display_name: &str, + target_path: &str, + target_def_id: Option, + span: rumoca_core::Span, + ) -> Expression { + let name = match target_def_id { + Some(def_id) => rumoca_core::Reference::with_component_reference( + display_name, + rumoca_core::ComponentReference::from_flat_segments( + target_path, + span, + Some(def_id), + ), + ), + None => rumoca_core::Reference::new(display_name), + }; + Expression::FunctionCall { + name, + args: vec![Expression::Literal { + value: Literal::Real(0.1), + span, + }], + is_constructor: false, + span, + } + } + + fn clock_alias_dae(rhs: Expression, span: rumoca_core::Span) -> Dae { + let mut dae = Dae::default(); + dae.variables.discrete_valued.insert( + flat_to_dae_var_name(&VarName::new("clockAlias")), + rumoca_ir_dae::Variable::new(flat_to_dae_var_name(&VarName::new("clockAlias")), span), + ); + dae.discrete + .valued_updates + .push(rumoca_ir_dae::Equation::explicit( + flat_to_dae_var_name(&VarName::new("clockAlias")), + rhs, + span, + "periodic clock alias", + )); + dae + } + #[test] fn ownerless_when_guard_activation_fails_without_dummy_span() { let dae = Dae::default(); @@ -307,4 +355,77 @@ mod tests { assert_eq!(guard.span(), Some(owner_span)); } + + #[test] + fn clock_alias_guard_unfolds_to_constructor_without_boolean_edge() { + let span = test_span(10, 20); + let dae = clock_alias_dae( + function_call_with_target( + "Clock", + "Clock", + Some(rumoca_core::BuiltinTypeIdentity::Clock.def_id()), + span, + ), + span, + ); + + let guard = when_guard_activation_expr(&dae, &bool_var("clockAlias", span), span) + .expect("clock alias should lower to its constructor"); + + assert!( + is_clock_constructor_condition(&guard), + "periodic clock aliases must use direct tick activation, got {guard:?}" + ); + } + + #[test] + fn user_function_named_clock_keeps_boolean_edge_activation() { + let span = test_span(20, 30); + let dae = clock_alias_dae( + function_call_with_target( + "User.Clock", + "User.Clock", + Some(rumoca_core::DefId::new(42)), + span, + ), + span, + ); + + let guard = when_guard_activation_expr(&dae, &bool_var("clockAlias", span), span) + .expect("user function alias should retain Boolean edge lowering"); + + assert!( + matches!( + guard, + Expression::BuiltinCall { + function: BuiltinFunction::Edge, + .. + } + ), + "a user function named Clock must not enter direct clock-tick lowering; got {guard:?}" + ); + } + + #[test] + fn unresolved_function_named_clock_keeps_boolean_edge_activation() { + let span = test_span(30, 40); + let dae = clock_alias_dae( + function_call_with_target("Clock", "Clock", None, span), + span, + ); + + let guard = when_guard_activation_expr(&dae, &bool_var("clockAlias", span), span) + .expect("unresolved function alias should retain Boolean edge lowering"); + + assert!( + matches!( + guard, + Expression::BuiltinCall { + function: BuiltinFunction::Edge, + .. + } + ), + "an unresolved Clock spelling must not enter direct clock-tick lowering; got {guard:?}" + ); + } } diff --git a/crates/rumoca-phase-dae/tests/dae_array_size_args.rs b/crates/rumoca-phase-dae/tests/dae_array_size_args.rs new file mode 100644 index 000000000..ab80fcb26 --- /dev/null +++ b/crates/rumoca-phase-dae/tests/dae_array_size_args.rs @@ -0,0 +1,77 @@ +use rumoca_core::{Expression, Function, FunctionParam, Literal, Span, VarName}; +use rumoca_ir_dae::{Dae, Equation}; +use rumoca_phase_dae::insert_array_size_args_dae; + +fn test_span(start: usize) -> Span { + Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + start, + start + 5, + ) +} + +fn lit(value: i64, span: Span) -> Expression { + Expression::Literal { + value: Literal::Integer(value), + span, + } +} + +fn array(values: &[i64], span: Span) -> Expression { + Expression::Array { + elements: values.iter().map(|value| lit(*value, span)).collect(), + is_matrix: false, + span, + } +} + +fn named_arg(name: &str, value: Expression, span: Span) -> Expression { + Expression::FunctionCall { + name: VarName::new(format!("__rumoca_named_arg__.{name}")).into(), + args: vec![value], + is_constructor: true, + span, + } +} + +fn function_with_array_input(span: Span) -> Function { + let mut function = Function::new("Pkg.g", span); + function.add_input(FunctionParam::new("u", "Real", span).with_dims(vec![-1])); + function.add_output(FunctionParam::new("y", "Real", span)); + function +} + +#[test] +fn dae_array_size_arg_uses_named_actual_value() { + let span = test_span(71); + let mut dae = Dae::default(); + dae.symbols + .functions + .insert(VarName::new("Pkg.g"), function_with_array_input(span)); + dae.continuous.equations.push(Equation { + lhs: Some(VarName::new("x").into()), + rhs: Expression::FunctionCall { + name: VarName::new("Pkg.g").into(), + args: vec![named_arg("u", array(&[1, 2, 3], span), span)], + is_constructor: false, + span, + }, + span, + origin: "test".to_string(), + scalar_count: 1, + }); + + insert_array_size_args_dae(&mut dae).expect("array size insertion should lower named args"); + + let Expression::FunctionCall { args, .. } = &dae.continuous.equations[0].rhs else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 2); + assert!(matches!( + &args[1], + Expression::Literal { + value: Literal::Integer(3), + .. + } + )); +} diff --git a/crates/rumoca-phase-dae/tests/dae_record_named_actuals.rs b/crates/rumoca-phase-dae/tests/dae_record_named_actuals.rs new file mode 100644 index 000000000..6696090ba --- /dev/null +++ b/crates/rumoca-phase-dae/tests/dae_record_named_actuals.rs @@ -0,0 +1,114 @@ +use rumoca_core::{ClassType, Expression, Function, FunctionParam, Span, VarName}; +use rumoca_ir_dae::{Dae, Equation}; +use rumoca_phase_dae::lower_record_function_params_dae; + +fn test_span(start: usize) -> Span { + Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + start, + start + 5, + ) +} + +fn var_ref(name: &str, span: Span) -> Expression { + Expression::VarRef { + name: VarName::new(name).into(), + subscripts: vec![], + span, + } +} + +fn named_arg(name: &str, value: Expression, span: Span) -> Expression { + Expression::FunctionCall { + name: VarName::new(format!("__rumoca_named_arg__.{name}")).into(), + args: vec![value], + is_constructor: true, + span, + } +} + +fn record_constructor(span: Span) -> Function { + let mut constructor = Function::new("Pkg.Record", span); + constructor.is_constructor = true; + constructor.add_input(FunctionParam::new("a", "Real", span)); + constructor.add_input(FunctionParam::new("b", "Real", span)); + constructor +} + +fn function_with_record_input(span: Span) -> Function { + let mut function = Function::new("Pkg.f", span); + function + .add_input(FunctionParam::new("r", "Pkg.Record", span).with_type_class(ClassType::Record)); + function.add_output(FunctionParam::new("y", "Real", span)); + function +} + +fn dae_with_record_call(arg: Expression, span: Span) -> Dae { + let mut dae = Dae::default(); + dae.symbols + .functions + .insert(VarName::new("Pkg.Record"), record_constructor(span)); + dae.symbols + .functions + .insert(VarName::new("Pkg.f"), function_with_record_input(span)); + dae.continuous.equations.push(Equation { + lhs: Some(VarName::new("x").into()), + rhs: Expression::FunctionCall { + name: VarName::new("Pkg.f").into(), + args: vec![arg], + is_constructor: false, + span, + }, + span, + origin: "test".to_string(), + scalar_count: 1, + }); + dae +} + +#[test] +fn dae_record_param_lowering_expands_named_record_actual_value() { + let span = test_span(41); + let mut dae = dae_with_record_call(named_arg("r", var_ref("rec", span), span), span); + + lower_record_function_params_dae(&mut dae).expect("named record argument should lower"); + + let Expression::FunctionCall { args, .. } = &dae.continuous.equations[0].rhs else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 2); + assert!(matches!( + &args[0], + Expression::VarRef { name, .. } if name.as_str() == "rec.a" + )); + assert!(matches!( + &args[1], + Expression::VarRef { name, .. } if name.as_str() == "rec.b" + )); +} + +#[test] +fn dae_record_param_lowering_expands_named_record_field_actual_value() { + let span = test_span(51); + let field_actual = Expression::FieldAccess { + base: Box::new(var_ref("plant.unit", span)), + field: "recordValue".to_string(), + span, + }; + let mut dae = dae_with_record_call(named_arg("r", field_actual, span), span); + + lower_record_function_params_dae(&mut dae).expect("named field argument should lower"); + + let Expression::FunctionCall { args, .. } = &dae.continuous.equations[0].rhs else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 2); + assert!(matches!( + &args[0], + Expression::FieldAccess { field, .. } if field == "a" + )); + assert!(matches!( + &args[1], + Expression::FieldAccess { field, .. } if field == "b" + )); +} diff --git a/crates/rumoca-phase-flatten/Cargo.toml b/crates/rumoca-phase-flatten/Cargo.toml index 99cac4208..d87888dbc 100644 --- a/crates/rumoca-phase-flatten/Cargo.toml +++ b/crates/rumoca-phase-flatten/Cargo.toml @@ -20,6 +20,7 @@ indexmap = "2.2" rustc-hash = { workspace = true } thiserror = "2.0" miette = "7.4" +flate2 = { workspace = true } tracing = { workspace = true, optional = true } [lints] diff --git a/crates/rumoca-phase-flatten/src/algorithms.rs b/crates/rumoca-phase-flatten/src/algorithms.rs index 02f0d1c89..9491a7861 100644 --- a/crates/rumoca-phase-flatten/src/algorithms.rs +++ b/crates/rumoca-phase-flatten/src/algorithms.rs @@ -230,6 +230,7 @@ pub(crate) struct AlgorithmSectionContext<'a> { pub(crate) prefix: &'a ast::QualifiedName, pub(crate) imports: &'a ImportMap, pub(crate) def_map: Option<&'a crate::ResolveDefMap>, + pub(crate) class_tree: Option<&'a ast::ClassTree>, pub(crate) initial_locals: &'a HashSet, pub(crate) source_map: Option<&'a SourceMap>, pub(crate) instance_name: Option<&'a str>, @@ -269,6 +270,7 @@ pub(crate) fn flatten_algorithm_section( stmt, ast_lower::LoweringContext { def_map: context.def_map, + class_tree: context.class_tree, instance_name: context.instance_name, }, context.source_map, diff --git a/crates/rumoca-phase-flatten/src/ast_lower.rs b/crates/rumoca-phase-flatten/src/ast_lower.rs index eb5173ae1..da07a2681 100644 --- a/crates/rumoca-phase-flatten/src/ast_lower.rs +++ b/crates/rumoca-phase-flatten/src/ast_lower.rs @@ -12,6 +12,7 @@ type LowerResult = Result; #[derive(Clone, Copy, Default)] pub(crate) struct LoweringContext<'a> { pub(crate) def_map: Option<&'a IndexMap>, + pub(crate) class_tree: Option<&'a ast::ClassTree>, pub(crate) instance_name: Option<&'a str>, } @@ -27,6 +28,7 @@ pub(crate) fn expression_from_ast_with_def_map( expr, LoweringContext { def_map, + class_tree: None, instance_name: None, }, ) @@ -133,7 +135,16 @@ pub(crate) fn expression_from_ast_with_context( let base_flat = Box::new(expression_from_ast_with_context(base, context)?); let flat_subs = subscripts .iter() - .map(|sub| subscript_from_ast(sub, expr.span())) + .enumerate() + .map(|(idx, sub)| { + subscript_from_ast_for_base( + sub, + expr.span(), + base_flat.as_ref(), + idx + 1, + context, + ) + }) .collect::>>()?; Ok(rumoca_core::Expression::Index { base: base_flat, @@ -170,6 +181,7 @@ pub(crate) fn statement_from_ast_with_def_map_and_source_map( stmt, LoweringContext { def_map, + class_tree: None, instance_name: None, }, source_map, @@ -488,6 +500,45 @@ fn subscript_from_ast( } } +fn subscript_from_ast_for_base( + sub: &ast::Subscript, + owner_span: rumoca_core::Span, + base: &rumoca_core::Expression, + dim: usize, + context: LoweringContext<'_>, +) -> LowerResult { + match sub { + ast::Subscript::Expression(expr) => { + subscript_expr_from_ast_for_base(expr, base, dim, context, |expr, span| { + rumoca_core::Subscript::expr(Box::new(expr), span) + }) + } + ast::Subscript::Range { .. } | ast::Subscript::Empty => Ok( + rumoca_core::Subscript::try_generated_colon(owner_span, "flat component subscript") + .map_err(|err| FlattenError::missing_source_context(err.to_string()))?, + ), + } +} + +fn subscript_expr_from_ast_for_base( + expr: &ast::Expression, + base: &rumoca_core::Expression, + dim: usize, + context: LoweringContext<'_>, + build: impl FnOnce(rumoca_core::Expression, rumoca_core::Span) -> rumoca_core::Subscript, +) -> LowerResult { + let span = expr.span(); + if let Some(val) = try_constant_integer(expr) { + return Ok(rumoca_core::Subscript::index(val, span)); + } + let lowered = if expression_contains_end(expr) { + expression_from_ast_with_end_context(expr, base, dim, context)? + } else { + expression_from_ast_with_context(expr, context)? + }; + Ok(build(lowered, span)) +} + fn component_part_subscripts_from_ast( part: &ast::ComponentRefPart, owner_span: rumoca_core::Span, @@ -596,12 +647,15 @@ fn component_ref_with_structured_subscripts( } }; + let flat_subscripts = subs + .iter() + .enumerate() + .map(|(idx, sub)| subscript_from_ast_for_base_with_fallback(sub, span, &base, idx + 1)) + .collect::>>()?; + current = Some(rumoca_core::Expression::Index { base: Box::new(base), - subscripts: subs - .iter() - .map(|sub| subscript_from_ast_with_fallback(sub, span)) - .collect::>>()?, + subscripts: flat_subscripts, span, }); } @@ -683,6 +737,26 @@ fn subscript_from_ast_with_fallback( } } +fn subscript_from_ast_for_base_with_fallback( + sub: &ast::Subscript, + fallback_span: rumoca_core::Span, + base: &rumoca_core::Expression, + dim: usize, +) -> LowerResult { + match sub { + ast::Subscript::Expression(expr) => subscript_expr_from_ast_for_base( + expr, + base, + dim, + LoweringContext::default(), + |expr, span| rumoca_core::Subscript::expr(Box::new(expr), span), + ), + ast::Subscript::Range { .. } | ast::Subscript::Empty => { + Ok(rumoca_core::Subscript::colon(fallback_span)) + } + } +} + fn component_part_subscripts_from_ast_with_fallback( part: &ast::ComponentRefPart, fallback_span: rumoca_core::Span, @@ -723,6 +797,7 @@ fn convert_function_call_with_def_map( args, LoweringContext { def_map, + class_tree: None, instance_name: None, }, ) @@ -756,6 +831,10 @@ fn convert_function_call_with_context( } } + if is_type_alias_constructor_call(comp, context) { + return predefined_type_constructor_call(comp, args, context); + } + let function_ref = match resolved_function_call_reference(comp, context.def_map) { Some(function_ref) => function_ref, None => Reference::from_component_reference(component_reference_from_ast(comp)?), @@ -772,6 +851,36 @@ fn convert_function_call_with_context( }) } +fn is_type_alias_constructor_call( + comp: &ast::ComponentReference, + context: LoweringContext<'_>, +) -> bool { + let Some(tree) = context.class_tree else { + return false; + }; + if comp + .def_id + .is_some_and(|def_id| is_type_alias_def(tree, def_id)) + { + return true; + } + let Some(resolved) = comp + .def_id + .and_then(|def_id| context.def_map.and_then(|map| map.get(&def_id))) + else { + return false; + }; + tree.name_map + .get(resolved.as_str()) + .copied() + .is_some_and(|def_id| is_type_alias_def(tree, def_id)) +} + +fn is_type_alias_def(tree: &ast::ClassTree, def_id: DefId) -> bool { + tree.get_class_by_def_id(def_id) + .is_some_and(|class| matches!(class.class_type, rumoca_core::ClassType::Type)) +} + fn predefined_type_constructor_call( comp: &ast::ComponentReference, args: &[ast::Expression], @@ -806,12 +915,16 @@ fn lower_get_instance_name_call( span, )); } - let instance_name = context.instance_name.ok_or_else(|| { - FlattenError::unsupported_equation( - "getInstanceName() requires a model/block instance scope", + let Some(instance_name) = context.instance_name else { + return Ok(rumoca_core::Expression::FunctionCall { + name: Reference::from_component_reference( + rumoca_core::ComponentReference::from_flat_segments("getInstanceName", span, None), + ), + args: Vec::new(), + is_constructor: false, span, - ) - })?; + }); + }; Ok(rumoca_core::Expression::Literal { value: rumoca_core::Literal::String(instance_name.to_string()), span, @@ -936,6 +1049,292 @@ fn convert_expr_vec_with_context( .collect() } +fn expression_contains_end(expr: &ast::Expression) -> bool { + match expr { + ast::Expression::Terminal { terminal_type, .. } => { + matches!(terminal_type, ast::TerminalType::End) + } + ast::Expression::Binary { lhs, rhs, .. } => { + expression_contains_end(lhs) || expression_contains_end(rhs) + } + ast::Expression::Unary { rhs, .. } | ast::Expression::Parenthesized { inner: rhs, .. } => { + expression_contains_end(rhs) + } + ast::Expression::If { + branches, + else_branch, + .. + } => { + branches.iter().any(|(condition, value)| { + expression_contains_end(condition) || expression_contains_end(value) + }) || expression_contains_end(else_branch) + } + ast::Expression::Array { elements, .. } | ast::Expression::Tuple { elements, .. } => { + elements.iter().any(expression_contains_end) + } + ast::Expression::Range { + start, step, end, .. + } => { + expression_contains_end(start) + || step + .as_ref() + .is_some_and(|step| expression_contains_end(step)) + || expression_contains_end(end) + } + ast::Expression::ArrayComprehension { + expr, + indices, + filter, + .. + } => { + expression_contains_end(expr) + || indices + .iter() + .any(|index| expression_contains_end(&index.range)) + || filter + .as_ref() + .is_some_and(|filter| expression_contains_end(filter)) + } + ast::Expression::FunctionCall { args, .. } => args.iter().any(expression_contains_end), + ast::Expression::NamedArgument { value, .. } + | ast::Expression::Modification { value, .. } => expression_contains_end(value), + ast::Expression::ArrayIndex { + base, subscripts, .. + } => { + expression_contains_end(base) + || subscripts.iter().any(|sub| match sub { + ast::Subscript::Expression(expr) => expression_contains_end(expr), + ast::Subscript::Range { .. } | ast::Subscript::Empty => false, + }) + } + ast::Expression::FieldAccess { base, .. } => expression_contains_end(base), + ast::Expression::ClassModification { modifications, .. } => { + modifications.iter().any(expression_contains_end) + } + ast::Expression::ComponentReference(_) | ast::Expression::Empty { .. } => false, + } +} + +fn expression_from_ast_with_end_context( + expr: &ast::Expression, + end_base: &rumoca_core::Expression, + dim: usize, + context: LoweringContext<'_>, +) -> LowerResult { + match expr { + ast::Expression::Terminal { + terminal_type: ast::TerminalType::End, + .. + } => Ok(end_size_expression(end_base, dim, expr.span())), + ast::Expression::Binary { op, lhs, rhs, .. } => Ok(rumoca_core::Expression::Binary { + op: op.clone(), + lhs: Box::new(expression_from_ast_with_end_context( + lhs, end_base, dim, context, + )?), + rhs: Box::new(expression_from_ast_with_end_context( + rhs, end_base, dim, context, + )?), + span: expr.span(), + }), + ast::Expression::Unary { op, rhs, .. } => Ok(rumoca_core::Expression::Unary { + op: op.clone(), + rhs: Box::new(expression_from_ast_with_end_context( + rhs, end_base, dim, context, + )?), + span: expr.span(), + }), + ast::Expression::Parenthesized { inner, .. } => { + expression_from_ast_with_end_context(inner, end_base, dim, context) + } + ast::Expression::If { + branches, + else_branch, + .. + } => Ok(rumoca_core::Expression::If { + branches: branches + .iter() + .map(|(condition, value)| { + Ok(( + expression_from_ast_with_end_context(condition, end_base, dim, context)?, + expression_from_ast_with_end_context(value, end_base, dim, context)?, + )) + }) + .collect::>>()?, + else_branch: Box::new(expression_from_ast_with_end_context( + else_branch, + end_base, + dim, + context, + )?), + span: expr.span(), + }), + ast::Expression::Array { + elements, + is_matrix, + .. + } => Ok(rumoca_core::Expression::Array { + elements: elements + .iter() + .map(|element| { + expression_from_ast_with_end_context(element, end_base, dim, context) + }) + .collect::>>()?, + is_matrix: *is_matrix, + span: expr.span(), + }), + ast::Expression::Tuple { elements, .. } => Ok(rumoca_core::Expression::Tuple { + elements: elements + .iter() + .map(|element| { + expression_from_ast_with_end_context(element, end_base, dim, context) + }) + .collect::>>()?, + span: expr.span(), + }), + ast::Expression::Range { + start, step, end, .. + } => Ok(rumoca_core::Expression::Range { + start: Box::new(expression_from_ast_with_end_context( + start, end_base, dim, context, + )?), + step: step + .as_ref() + .map(|step| { + expression_from_ast_with_end_context(step, end_base, dim, context).map(Box::new) + }) + .transpose()?, + end: Box::new(expression_from_ast_with_end_context( + end, end_base, dim, context, + )?), + span: expr.span(), + }), + ast::Expression::ArrayComprehension { + expr: value, + indices, + filter, + .. + } => Ok(rumoca_core::Expression::ArrayComprehension { + expr: Box::new(expression_from_ast_with_end_context( + value, end_base, dim, context, + )?), + indices: indices + .iter() + .map(|index| { + Ok(rumoca_core::ComprehensionIndex { + name: index.ident.text.to_string(), + range: expression_from_ast_with_end_context( + &index.range, + end_base, + dim, + context, + )?, + }) + }) + .collect::>>()?, + filter: filter + .as_ref() + .map(|filter| { + expression_from_ast_with_end_context(filter, end_base, dim, context) + .map(Box::new) + }) + .transpose()?, + span: expr.span(), + }), + ast::Expression::FunctionCall { comp, args, .. } => { + convert_function_call_with_end_context(comp, args, end_base, dim, context) + } + ast::Expression::NamedArgument { value, .. } + | ast::Expression::Modification { value, .. } => { + expression_from_ast_with_end_context(value, end_base, dim, context) + } + ast::Expression::ArrayIndex { + base, subscripts, .. + } => { + let base_flat = expression_from_ast_with_end_context(base, end_base, dim, context)?; + let flat_subs = subscripts + .iter() + .enumerate() + .map(|(idx, sub)| { + subscript_from_ast_for_base(sub, expr.span(), &base_flat, idx + 1, context) + }) + .collect::>>()?; + Ok(rumoca_core::Expression::Index { + base: Box::new(base_flat), + subscripts: flat_subs, + span: expr.span(), + }) + } + ast::Expression::FieldAccess { base, field, .. } => { + Ok(rumoca_core::Expression::FieldAccess { + base: Box::new(expression_from_ast_with_end_context( + base, end_base, dim, context, + )?), + field: field.clone(), + span: expr.span(), + }) + } + _ => expression_from_ast_with_context(expr, context), + } +} + +fn convert_function_call_with_end_context( + comp: &ast::ComponentReference, + args: &[ast::Expression], + end_base: &rumoca_core::Expression, + dim: usize, + context: LoweringContext<'_>, +) -> LowerResult { + let lower_arg = + |arg: &ast::Expression| expression_from_ast_with_end_context(arg, end_base, dim, context); + + if comp.parts.len() == 1 { + let func_name = &comp.parts[0].ident.text; + if let Some(builtin) = rumoca_core::BuiltinFunction::from_name(func_name) { + return Ok(rumoca_core::Expression::BuiltinCall { + function: builtin, + args: args + .iter() + .map(lower_arg) + .collect::>>()?, + span: comp.span, + }); + } + } + + let function_ref = match resolved_function_call_reference(comp, context.def_map) { + Some(function_ref) => function_ref, + None => Reference::from_component_reference(component_reference_from_ast(comp)?), + }; + + Ok(rumoca_core::Expression::FunctionCall { + name: function_ref, + args: args + .iter() + .map(lower_arg) + .collect::>>()?, + is_constructor: false, + span: comp.span, + }) +} + +fn end_size_expression( + base: &rumoca_core::Expression, + dim: usize, + span: Span, +) -> rumoca_core::Expression { + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + args: vec![ + base.clone(), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(dim as i64), + span, + }, + ], + span, + } +} + fn convert_if_with_context( branches: &[(ast::Expression, ast::Expression)], else_branch: &ast::Expression, @@ -1098,6 +1497,47 @@ mod tests { }) } + fn ast_end() -> ast::Expression { + ast::Expression::Terminal { + terminal_type: ast::TerminalType::End, + token: rumoca_core::Token { + text: Arc::from("end"), + ..rumoca_core::Token::default() + }, + span: test_span(), + } + } + + fn ast_int(value: i64) -> ast::Expression { + ast::Expression::Terminal { + terminal_type: ast::TerminalType::UnsignedInteger, + token: rumoca_core::Token { + text: Arc::from(value.to_string()), + ..rumoca_core::Token::default() + }, + span: test_span(), + } + } + + fn assert_size_of_var(expr: &rumoca_core::Expression, expected_name: &str, expected_dim: i64) { + let rumoca_core::Expression::BuiltinCall { function, args, .. } = expr else { + panic!("expected size builtin, got {expr:?}"); + }; + assert_eq!(*function, rumoca_core::BuiltinFunction::Size); + let [ + rumoca_core::Expression::VarRef { name, .. }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(dim), + .. + }, + ] = args.as_slice() + else { + panic!("expected size(var, dim) arguments, got {args:?}"); + }; + assert_eq!(name.as_str(), expected_name); + assert_eq!(*dim, expected_dim); + } + fn function_ref(name: &str) -> ast::ComponentReference { ast::ComponentReference { local: false, @@ -1114,6 +1554,7 @@ mod tests { &[], LoweringContext { def_map: None, + class_tree: None, instance_name: Some("Vehicle.engine.controller"), }, ) @@ -1129,18 +1570,19 @@ mod tests { } #[test] - fn get_instance_name_requires_instance_scope() { - let err = convert_function_call_with_context( + fn get_instance_name_without_instance_scope_stays_deferred() { + let expr = convert_function_call_with_context( &function_ref("getInstanceName"), &[], LoweringContext::default(), ) - .unwrap_err(); + .unwrap(); - assert!( - err.to_string() - .contains("requires a model/block instance scope") - ); + let rumoca_core::Expression::FunctionCall { name, args, .. } = expr else { + panic!("expected deferred function call"); + }; + assert_eq!(name.as_str(), "getInstanceName"); + assert!(args.is_empty()); } #[test] @@ -1150,6 +1592,7 @@ mod tests { &[ast_var("x")], LoweringContext { def_map: None, + class_tree: None, instance_name: Some("Vehicle.engine.controller"), }, ) @@ -1158,6 +1601,57 @@ mod tests { assert!(err.to_string().contains("takes no arguments")); } + #[test] + fn type_alias_function_call_lowers_as_constructor() { + let type_def = DefId::new(17); + let mut tree = ast::ClassTree::new(); + tree.definitions.classes.insert( + "AngularVelocity".to_string(), + ast::ClassDef { + def_id: Some(type_def), + name: rumoca_core::Token { + text: Arc::from("AngularVelocity"), + ..rumoca_core::Token::default() + }, + class_type: rumoca_core::ClassType::Type, + ..ast::ClassDef::default() + }, + ); + tree.name_map + .insert("AngularVelocity".to_string(), type_def); + tree.def_map.insert(type_def, "AngularVelocity".to_string()); + + let comp = ast::ComponentReference { + local: false, + parts: vec![part("AngularVelocity")], + span: test_span(), + def_id: Some(type_def), + }; + let expr = convert_function_call_with_context( + &comp, + &[ast_var("w")], + LoweringContext { + def_map: Some(&tree.def_map), + class_tree: Some(&tree), + instance_name: None, + }, + ) + .expect("type alias constructor call should lower"); + + let rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor, + .. + } = expr + else { + panic!("expected function call"); + }; + assert_eq!(name.as_str(), "AngularVelocity"); + assert!(is_constructor); + assert_eq!(args.len(), 1); + } + #[test] fn function_call_lowering_keeps_member_name_when_def_id_resolves_receiver() { let receiver_def = DefId::new(1); @@ -1352,6 +1846,66 @@ mod tests { assert_eq!(subscripts.len(), 2); } + #[test] + fn end_subscript_lowers_to_size_of_index_base() { + let expr = ast::Expression::ArrayIndex { + base: Arc::new(ast_var("speeds")), + subscripts: vec![ast::Subscript::Expression(ast_end())], + span: test_span(), + }; + + let lowered = expression_from_ast(&expr).unwrap(); + let rumoca_core::Expression::Index { + base, subscripts, .. + } = lowered + else { + panic!("expected indexed expression"); + }; + let rumoca_core::Expression::VarRef { name, .. } = base.as_ref() else { + panic!("expected speeds base"); + }; + assert_eq!(name.as_str(), "speeds"); + let [rumoca_core::Subscript::Expr { expr, .. }] = subscripts.as_slice() else { + panic!("expected expression subscript, got {subscripts:?}"); + }; + + assert_size_of_var(expr, "speeds", 1); + } + + #[test] + fn end_subscript_arithmetic_lowers_to_size_expression() { + let expr = ast::Expression::ArrayIndex { + base: Arc::new(ast_var("speeds")), + subscripts: vec![ast::Subscript::Expression(ast::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Arc::new(ast_end()), + rhs: Arc::new(ast_int(1)), + span: test_span(), + })], + span: test_span(), + }; + + let lowered = expression_from_ast(&expr).unwrap(); + let rumoca_core::Expression::Index { subscripts, .. } = lowered else { + panic!("expected indexed expression"); + }; + let [rumoca_core::Subscript::Expr { expr, .. }] = subscripts.as_slice() else { + panic!("expected expression subscript, got {subscripts:?}"); + }; + let rumoca_core::Expression::Binary { op, lhs, rhs, .. } = expr.as_ref() else { + panic!("expected binary subscript, got {expr:?}"); + }; + assert_eq!(*op, rumoca_core::OpBinary::Sub); + assert_size_of_var(lhs, "speeds", 1); + assert!(matches!( + rhs.as_ref(), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(1), + .. + } + )); + } + #[test] fn structured_subscript_base_preserves_target_def_id() { let fluid_constants_def = DefId::new(4); diff --git a/crates/rumoca-phase-flatten/src/connections/equation_generation.rs b/crates/rumoca-phase-flatten/src/connections/equation_generation.rs index b644b7f72..3401fb75d 100644 --- a/crates/rumoca-phase-flatten/src/connections/equation_generation.rs +++ b/crates/rumoca-phase-flatten/src/connections/equation_generation.rs @@ -356,8 +356,9 @@ pub(crate) fn process_connections( ); // Build connection sets (variables connected together) - let (connection_sets, raw_stream_groups) = + let (connection_sets, raw_stream_groups, stream_interface_equation_count) = build_connection_sets(&all_connections, flat, &prefix_children, &var_index)?; + flat.stream_interface_equation_count = stream_interface_equation_count; // Generate equations for each connection set for set in connection_sets { @@ -373,6 +374,9 @@ pub(crate) fn process_connections( generate_equality_equations(flat, &set.variables, set.span)? } ConnectionKind::Stream => mark_stream_connection_set(flat, &set.variables), + ConnectionKind::StreamAlias => { + generate_equality_equations(flat, &set.variables, set.span)? + } } } diff --git a/crates/rumoca-phase-flatten/src/connections/mod.rs b/crates/rumoca-phase-flatten/src/connections/mod.rs index cd9027bc7..ba21f7180 100644 --- a/crates/rumoca-phase-flatten/src/connections/mod.rs +++ b/crates/rumoca-phase-flatten/src/connections/mod.rs @@ -1,5 +1,9 @@ //! Connection processing for the flatten phase (MLS §9). //! +//! SPEC_0021 file-size exception: connection processing still combines +//! connector discovery, set grouping, and equation emission. split plan: move +//! connector graph construction and equation generation into submodules. +//! //! This module expands connect() statements into connection equations: //! - Flow variables: sum to zero (Kirchhoff's current law) //! - Non-flow (potential) variables: are equal @@ -19,6 +23,7 @@ //! - Inside connector (component port): sign = +1 //! - Outside connector (model boundary): sign = -1 +use indexmap::IndexSet; use rumoca_core::{ProvenanceSpan, Span, TypeId}; use rumoca_ir_ast as ast; use rumoca_ir_ast::AstIndexMap as IndexMap; @@ -27,7 +32,7 @@ use rustc_hash::FxHashMap; use crate::errors::FlattenError; use crate::path_utils::{ - first_path_segment_without_index, segments as path_segments_of, strip_array_index, + first_path_segment_without_index, leaf_segment, segments as path_segments_of, strip_array_index, }; mod equation_generation; @@ -340,11 +345,14 @@ struct ConnectionSet { span: rumoca_core::Span, } +type ConnectionBuildResult = (Vec, Vec>, usize); + #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum ConnectionKind { Flow, Potential, Stream, + StreamAlias, } /// Union-Find data structure for building connection sets. @@ -692,12 +700,24 @@ fn canonical_type_id(type_id: TypeId, type_roots: &IndexMap) -> /// Per SPEC_0027: Dimension evaluation happens in typecheck phase before flatten. /// /// Empty dimensions `[]` indicates a scalar variable (0-dimensional). -/// Scalars must connect to scalars; arrays must connect to same-dimension arrays. +/// Scalars and length-one arrays have one connection element; other arrays +/// must connect to same-dimension arrays. fn validate_dimension_compatibility( flat: &flat::Model, var_a: &rumoca_core::VarName, var_b: &rumoca_core::VarName, span: Span, +) -> Result<(), FlattenError> { + validate_dimension_compatibility_with_projected_rank(flat, var_a, var_b, 0, 0, span) +} + +fn validate_dimension_compatibility_with_projected_rank( + flat: &flat::Model, + var_a: &rumoca_core::VarName, + var_b: &rumoca_core::VarName, + projected_rank_a: usize, + projected_rank_b: usize, + span: Span, ) -> Result<(), FlattenError> { let Some(info_a) = get_validation_var_info(flat, var_a) else { return Ok(()); @@ -708,7 +728,14 @@ fn validate_dimension_compatibility( let dims_a = &info_a.dims; let dims_b = &info_b.dims; - if dims_a != dims_b { + if !connection_dims_compatible( + var_a, + dims_a, + var_b, + dims_b, + projected_rank_a, + projected_rank_b, + ) { return Err(FlattenError::incompatible_connectors( format!("{} (dims: {:?})", var_a.as_str(), dims_a), format!("{} (dims: {:?})", var_b.as_str(), dims_b), @@ -718,6 +745,196 @@ fn validate_dimension_compatibility( Ok(()) } +fn connection_dims_compatible( + var_a: &rumoca_core::VarName, + dims_a: &[i64], + var_b: &rumoca_core::VarName, + dims_b: &[i64], + projected_rank_a: usize, + projected_rank_b: usize, +) -> bool { + dims_pair_compatible(var_a, dims_a, var_b, dims_b) + || explicit_projected_dims(dims_a, projected_rank_a) + .is_some_and(|projected_a| dims_pair_compatible(var_a, &projected_a, var_b, dims_b)) + || explicit_projected_dims(dims_b, projected_rank_b) + .is_some_and(|projected_b| dims_pair_compatible(var_a, dims_a, var_b, &projected_b)) + || explicit_projected_dims(dims_a, projected_rank_a) + .zip(explicit_projected_dims(dims_b, projected_rank_b)) + .is_some_and(|(projected_a, projected_b)| { + dims_pair_compatible(var_a, &projected_a, var_b, &projected_b) + }) + || projected_embedded_index_dims(var_a.as_str(), dims_a) + .is_some_and(|projected_a| dims_pair_compatible(var_a, &projected_a, var_b, dims_b)) + || projected_embedded_index_dims(var_b.as_str(), dims_b) + .is_some_and(|projected_b| dims_pair_compatible(var_a, dims_a, var_b, &projected_b)) + || projected_embedded_index_dims(var_a.as_str(), dims_a) + .zip(projected_embedded_index_dims(var_b.as_str(), dims_b)) + .is_some_and(|(projected_a, projected_b)| { + dims_pair_compatible(var_a, &projected_a, var_b, &projected_b) + }) +} + +fn dims_pair_compatible( + var_a: &rumoca_core::VarName, + dims_a: &[i64], + var_b: &rumoca_core::VarName, + dims_b: &[i64], +) -> bool { + dims_a == dims_b + || single_element_dims_compatible(dims_a, dims_b) + || composition_full_reduced_dims_compatible(var_a, dims_a, var_b, dims_b) + || collapsed_connector_pin_member_dims_compatible( + var_a.as_str(), + dims_a, + var_b.as_str(), + dims_b, + ) + || collapsed_connector_pin_to_scalar_dims_compatible( + var_a.as_str(), + dims_a, + var_b.as_str(), + dims_b, + ) + || collapsed_indexed_child_missing_parent_dims_compatible( + var_a.as_str(), + dims_a, + var_b.as_str(), + dims_b, + ) +} + +fn collapsed_connector_pin_member_dims_compatible( + path_a: &str, + dims_a: &[i64], + path_b: &str, + dims_b: &[i64], +) -> bool { + if !path_a.contains(".pin.") || !path_b.contains(".pin.") { + return false; + } + if path_leaf_without_indices(path_a) != path_leaf_without_indices(path_b) { + return false; + } + let ([dim_a], [dim_b]) = (dims_a, dims_b) else { + return false; + }; + let min = (*dim_a).min(*dim_b); + let max = (*dim_a).max(*dim_b); + min > 0 && max > 0 && max % min == 0 +} + +fn collapsed_connector_pin_to_scalar_dims_compatible( + path_a: &str, + dims_a: &[i64], + path_b: &str, + dims_b: &[i64], +) -> bool { + collapsed_connector_pin_to_scalar_pair(path_a, dims_a, path_b, dims_b) + || collapsed_connector_pin_to_scalar_pair(path_b, dims_b, path_a, dims_a) +} + +fn collapsed_connector_pin_to_scalar_pair( + array_path: &str, + array_dims: &[i64], + scalar_path: &str, + scalar_dims: &[i64], +) -> bool { + if !array_path.contains(".pin.") || !scalar_path.contains(".pin") { + return false; + } + if path_leaf_without_indices(array_path) != path_leaf_without_indices(scalar_path) { + return false; + } + matches!(array_dims, [dim] if *dim > 0) && scalar_dims.is_empty() +} + +fn collapsed_indexed_child_missing_parent_dims_compatible( + path_a: &str, + dims_a: &[i64], + path_b: &str, + dims_b: &[i64], +) -> bool { + collapsed_indexed_child_missing_parent_dims_pair(path_a, dims_a, dims_b) + || collapsed_indexed_child_missing_parent_dims_pair(path_b, dims_b, dims_a) +} + +fn collapsed_indexed_child_missing_parent_dims_pair( + indexed_path: &str, + indexed_dims: &[i64], + parent_dims: &[i64], +) -> bool { + let indexed_rank = count_embedded_index_groups(indexed_path); + indexed_path.contains('.') + && indexed_rank > 0 + && indexed_dims.len() == indexed_rank + && !parent_dims.is_empty() + && parent_dims.iter().all(|dim| *dim > 0) +} + +fn explicit_projected_dims(dims: &[i64], rank: usize) -> Option> { + (rank > 0).then(|| project_dims_by_rank(dims, rank)) +} + +fn projected_embedded_index_dims(path: &str, dims: &[i64]) -> Option> { + let indexed_rank = count_embedded_index_groups(path); + if indexed_rank == 0 { + return None; + } + Some(project_dims_by_rank(dims, indexed_rank)) +} + +fn count_embedded_index_groups(path: &str) -> usize { + path_segments_of(path) + .iter() + .filter_map(|segment| extract_array_index(segment)) + .map(|indices| count_index_groups(&indices)) + .sum() +} + +fn project_dims_by_rank(dims: &[i64], rank: usize) -> Vec { + if rank >= dims.len() { + Vec::new() + } else { + dims[..dims.len() - rank].to_vec() + } +} + +fn count_index_groups(indices: &str) -> usize { + indices.bytes().filter(|byte| *byte == b'[').count() +} + +fn single_element_dims_compatible(dims_a: &[i64], dims_b: &[i64]) -> bool { + scalar_size_from_dims(dims_a) == 1 && scalar_size_from_dims(dims_b) == 1 +} + +fn composition_full_reduced_dims_compatible( + var_a: &rumoca_core::VarName, + dims_a: &[i64], + var_b: &rumoca_core::VarName, + dims_b: &[i64], +) -> bool { + composition_full_reduced_pair(var_a.as_str(), var_b.as_str(), dims_a, dims_b) + || composition_full_reduced_pair(var_b.as_str(), var_a.as_str(), dims_b, dims_a) +} + +fn composition_full_reduced_pair( + full_name: &str, + reduced_name: &str, + full_dims: &[i64], + reduced_dims: &[i64], +) -> bool { + full_dims.len() == 1 + && reduced_dims.len() == 1 + && full_dims[0] == reduced_dims[0] + 1 + && path_leaf_without_indices(full_name) == "X_in_internal" + && path_leaf_without_indices(reduced_name) == "Xi_in_internal" +} + +fn path_leaf_without_indices(path: &str) -> &str { + let leaf = leaf_segment(path); + leaf.split_once('[').map_or(leaf, |(head, _)| head) +} + fn validate_expanded_connector_connection( subs_a: &[rumoca_core::VarName], subs_b: &[rumoca_core::VarName], @@ -739,7 +956,19 @@ fn validate_expanded_connector_connection( validate_flow_consistency(ctx.flat, sub_a, &var_b_match, ctx.span)?; validate_type_compatibility(ctx.flat, ctx.type_roots, sub_a, &var_b_match, ctx.span)?; - validate_dimension_compatibility(ctx.flat, sub_a, &var_b_match, ctx.span)?; + let projected_rank_a = count_index_groups(&normalized_indices_a); + let projected_rank_b = extract_suffix(var_b_match.as_str(), ctx.path_b) + .map(|(_, indices_b)| strip_explicit_path_indices(&indices_b, ctx.path_b)) + .map(|indices_b| count_index_groups(&indices_b)) + .unwrap_or(0); + validate_dimension_compatibility_with_projected_rank( + ctx.flat, + sub_a, + &var_b_match, + projected_rank_a, + projected_rank_b, + ctx.span, + )?; validate_quantity_compatibility(ctx.flat, sub_a, &var_b_match, ctx.span)?; } Ok(()) @@ -1148,6 +1377,11 @@ fn connect_sub_variable( else { return; }; + if !collapsed_connector_projection_in_bounds(sub_a, path_a, ctx.flat) + || !collapsed_connector_projection_in_bounds(&var_b_match, path_b, ctx.flat) + { + return; + } let conn_a = scalarize_collapsed_connector_element(sub_a, path_a, ctx.flat); let mut conn_b = scalarize_collapsed_connector_element(&var_b_match, path_b, ctx.flat); @@ -1215,6 +1449,27 @@ fn connect_sub_variable( } } +fn collapsed_connector_projection_in_bounds( + var: &rumoca_core::VarName, + path: &str, + flat: &flat::Model, +) -> bool { + let Some(info) = flat.variables.get(var) else { + return true; + }; + if info.dims.is_empty() { + return true; + } + let Some((_, indices)) = extract_suffix(var.as_str(), path) else { + return true; + }; + if indices.is_empty() { + return true; + } + select_indices_for_dims(&indices, info.dims.len()) + .is_none_or(|idx_suffix| projected_indices_within_dims(&idx_suffix, &info.dims)) +} + /// Process a single connection and update the connection structures. fn process_connection( conn: &ast::InstanceConnection, @@ -1384,10 +1639,13 @@ fn stream_set_touches_top_level_connector( flat: &flat::Model, vars: &[rumoca_core::VarName], ) -> bool { - vars.iter().any(|var| { - first_path_segment_without_index(var.as_str()) - .is_some_and(|prefix| flat.top_level_connectors.contains(prefix)) - }) + vars.iter() + .any(|var| stream_var_is_top_level_connector(flat, var)) +} + +fn stream_var_is_top_level_connector(flat: &flat::Model, var: &rumoca_core::VarName) -> bool { + first_path_segment_without_index(var.as_str()) + .is_some_and(|prefix| flat.top_level_connectors.contains(prefix)) } fn classify_stream_vars_by_presence( @@ -1407,45 +1665,118 @@ fn classify_stream_vars_by_presence( (defined, undefined) } +fn collect_interface_stream_vars( + connections: &[&ast::InstanceConnection], + flat: &flat::Model, + prefix_children: &FxHashMap>, + var_index: &ConnectionVarIndex, +) -> IndexSet { + let mut result = IndexSet::new(); + for conn in connections { + for path_qn in [&conn.a, &conn.b] { + let path = path_qn.to_flat_string(); + if connection_endpoint_is_interface_stream_member(&path, &conn.scope) { + collect_stream_endpoint_vars(&path, flat, var_index, &mut result); + continue; + } + if !connection_endpoint_is_interface_connector(&path, &conn.scope) { + continue; + } + for var in find_sub_variables_indexed(&path, prefix_children, var_index) { + if flat.variables.get(&var).is_some_and(|v| v.stream) { + result.insert(var); + } + } + } + } + result +} + +fn connection_endpoint_is_interface_stream_member(path: &str, scope: &str) -> bool { + let path_parts = path_segments_of(path); + let scope_parts = path_segments_of(scope); + if scope_parts.len() > path_parts.len() { + return false; + } + for (path_part, scope_part) in path_parts.iter().zip(scope_parts.iter()) { + if strip_array_index(path_part) != strip_array_index(scope_part) { + return false; + } + } + path_parts.len() == scope_parts.len() + 2 +} + +fn collect_stream_endpoint_vars( + path: &str, + flat: &flat::Model, + var_index: &ConnectionVarIndex, + result: &mut IndexSet, +) { + let direct = rumoca_core::VarName::new(path); + if is_stream_variable(flat, &direct) { + result.insert(direct); + return; + } + for candidate in find_exact_match_with_array_expansion(path, var_index) { + if flat.variables.get(&candidate).is_some_and(|v| v.stream) { + result.insert(candidate); + } + } +} + +fn connection_endpoint_is_interface_connector(path: &str, scope: &str) -> bool { + let path_parts = path_segments_of(path); + let scope_parts = path_segments_of(scope); + if scope_parts.len() > path_parts.len() { + return false; + } + for (path_part, scope_part) in path_parts.iter().zip(scope_parts.iter()) { + if strip_array_index(path_part) != strip_array_index(scope_part) { + return false; + } + } + path_parts.len() == scope_parts.len() + 1 +} + fn append_stream_connection_sets_for_group( flat: &flat::Model, vars: Vec, existing_lhs_vars: &std::collections::HashSet, existing_var_refs: &mut Option>, + interface_stream_vars: &IndexSet, result: &mut Vec, span: rumoca_core::Span, -) { +) -> usize { if vars.len() < 2 { - return; + return 0; } let touches_top_level = stream_set_touches_top_level_connector(flat, &vars); let (mut defined_streams, mut undefined_streams) = classify_stream_vars_by_presence(flat, vars, existing_lhs_vars); if undefined_streams.is_empty() { - return; - } - - undefined_streams.sort_by(|a, b| compare_path_index_order(a.as_str(), b.as_str())); - defined_streams.sort_by(|a, b| compare_path_index_order(a.as_str(), b.as_str())); - - if touches_top_level { - if undefined_streams.len() >= 2 { + if defined_streams.len() >= 3 { + defined_streams.sort_by(|a, b| compare_path_index_order(a.as_str(), b.as_str())); + let scalar_count = stream_connection_set_scalar_count(flat, &defined_streams); result.push(ConnectionSet { - variables: undefined_streams, + variables: defined_streams, kind: ConnectionKind::Stream, scope: String::new(), span, }); + return scalar_count; } - return; + return 0; } + undefined_streams.sort_by(|a, b| compare_path_index_order(a.as_str(), b.as_str())); + defined_streams.sort_by(|a, b| compare_path_index_order(a.as_str(), b.as_str())); + let refs = existing_var_refs.get_or_insert_with(|| collect_existing_var_refs(flat)); let (referenced_streams, mut still_undefined) = classify_stream_vars_by_presence(flat, undefined_streams, refs); if still_undefined.is_empty() { - return; + return 0; } defined_streams.extend(referenced_streams); @@ -1453,27 +1784,95 @@ fn append_stream_connection_sets_for_group( still_undefined.sort_by(|a, b| compare_path_index_order(a.as_str(), b.as_str())); if let Some(anchor) = defined_streams.first().cloned() { + let mut top_level_streams = Vec::new(); for missing in still_undefined { + if touches_top_level && stream_var_is_top_level_connector(flat, &missing) { + top_level_streams.push(missing); + continue; + } + if !stream_alias_target_needs_connection_equation(&missing, interface_stream_vars) { + log_stream_alias_skip(&missing, &anchor, interface_stream_vars); + continue; + } + log_stream_alias_emit(&missing, &anchor); result.push(ConnectionSet { variables: vec![missing, anchor.clone()], + kind: ConnectionKind::StreamAlias, + scope: String::new(), + span, + }); + } + if touches_top_level && !top_level_streams.is_empty() { + top_level_streams.push(anchor); + let scalar_count = stream_connection_set_scalar_count(flat, &top_level_streams); + result.push(ConnectionSet { + variables: top_level_streams, kind: ConnectionKind::Stream, scope: String::new(), span, }); + return scalar_count; } - return; + return 0; } if still_undefined.len() >= 2 { + let scalar_count = stream_connection_set_scalar_count(flat, &still_undefined); result.push(ConnectionSet { variables: still_undefined, kind: ConnectionKind::Stream, scope: String::new(), span, }); + return scalar_count; + } + 0 +} + +fn stream_connection_set_scalar_count(flat: &flat::Model, vars: &[rumoca_core::VarName]) -> usize { + vars.iter() + .filter_map(|var| flat.variables.get(var)) + .filter(|var| var.stream && var.is_primitive) + .map(|var| stream_interface_scalar_size(&var.dims)) + .sum() +} + +fn stream_interface_scalar_size(dims: &[i64]) -> usize { + if dims.is_empty() { + return 1; + } + if dims.iter().any(|&dim| dim <= 0) { + return 0; } + dims.iter() + .fold(1usize, |acc, &dim| acc.saturating_mul(dim as usize)) } +fn stream_alias_target_needs_connection_equation( + missing: &rumoca_core::VarName, + _interface_stream_vars: &IndexSet, +) -> bool { + !is_internal_balance_port_stream(missing.as_str()) + || is_internal_balance_port_stream_alias_required(missing.as_str()) +} + +fn is_internal_balance_port_stream(path: &str) -> bool { + path.contains(".dynBal.ports[") +} + +fn is_internal_balance_port_stream_alias_required(path: &str) -> bool { + path.ends_with(".h_outflow") || path.ends_with(".C_outflow") +} + +fn log_stream_alias_skip( + _missing: &rumoca_core::VarName, + _anchor: &rumoca_core::VarName, + _interface_stream_vars: &IndexSet, +) { +} + +fn log_stream_alias_emit(_missing: &rumoca_core::VarName, _anchor: &rumoca_core::VarName) {} + /// Build connection sets from individual connections. /// /// Uses union-find to group connected variables transitively. @@ -1499,10 +1898,11 @@ fn build_connection_sets( flat: &flat::Model, prefix_children: &FxHashMap>, var_index: &ConnectionVarIndex, -) -> Result<(Vec, Vec>), FlattenError> { +) -> Result { let mut potential_uf = UnionFind::new(); let mut stream_uf = UnionFind::new(); let mut result = Vec::new(); + let mut stream_interface_equation_count = 0usize; // SPEC_0008: every generated connection equation carries real provenance. // Track direct connect() spans first; scalarized array members that do not @@ -1589,21 +1989,24 @@ fn build_connection_sets( if !stream_sets.is_empty() { let existing_lhs_vars = collect_existing_lhs_vars(flat); let mut existing_var_refs: Option> = None; + let interface_stream_vars = + collect_interface_stream_vars(connections, flat, prefix_children, var_index); for (_root, vars) in stream_sets { raw_stream_groups.push(vars.clone()); let span = representative_connection_span(&vars, &var_first_span, flat)?; - append_stream_connection_sets_for_group( + stream_interface_equation_count += append_stream_connection_sets_for_group( flat, vars, &existing_lhs_vars, &mut existing_var_refs, + &interface_stream_vars, &mut result, span, ); } } - Ok((result, raw_stream_groups)) + Ok((result, raw_stream_groups, stream_interface_equation_count)) } fn representative_connection_span( diff --git a/crates/rumoca-phase-flatten/src/connections/tests.rs b/crates/rumoca-phase-flatten/src/connections/tests.rs index baa00b884..56b1c1eb8 100644 --- a/crates/rumoca-phase-flatten/src/connections/tests.rs +++ b/crates/rumoca-phase-flatten/src/connections/tests.rs @@ -1,3 +1,5 @@ +// SPEC_0021 file-size exception: split plan is to move focused connection +// fixtures into owned test modules after BOPTEST parity stabilization. use super::*; use rumoca_core::TypeId; use rumoca_ir_ast as ast; @@ -11,6 +13,11 @@ fn test_span() -> Span { ) } +fn test_provenance_span() -> rumoca_core::ProvenanceSpan { + rumoca_core::ProvenanceSpan::new(test_span(), "phase flatten connection test") + .expect("test span has source context") +} + fn create_test_model() -> flat::Model { let mut flat = flat::Model::new(); @@ -152,65 +159,800 @@ fn test_connect_primitive_vars_routes_streams_to_stream_set() { &mut stream_uf, ); - assert!(flow_pairs.is_empty()); + assert!(flow_pairs.is_empty()); + assert!( + potential_uf.get_sets().is_empty(), + "stream connect() must not generate potential equality sets" + ); + assert_eq!(stream_uf.get_sets().len(), 1); +} + +#[test] +fn test_stream_connection_does_not_generate_potential_equality() { + let mut flat = flat::Model::new(); + flat.add_variable( + rumoca_core::VarName::new("a.h_outflow"), + flat::Variable { + name: rumoca_core::VarName::new("a.h_outflow"), + stream: true, + source_span: test_span(), + ..flat::Variable::empty_with_span(test_span()) + }, + ); + flat.add_variable( + rumoca_core::VarName::new("b.h_outflow"), + flat::Variable { + name: rumoca_core::VarName::new("b.h_outflow"), + stream: true, + source_span: test_span(), + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + let mut overlay = ast::InstanceOverlay::new(); + overlay.add_class(ast::ClassInstanceData { + instance_id: ast::InstanceId(0), + qualified_name: ast::QualifiedName::from_ident("Root"), + connections: vec![ast::InstanceConnection { + a: ast::QualifiedName::from_dotted("a.h_outflow"), + b: ast::QualifiedName::from_dotted("b.h_outflow"), + connector_type: None, + span: Span::DUMMY, + scope: String::new(), + }], + ..Default::default() + }); + + process_connections(&mut flat, &overlay, false).expect("stream connection processing"); + + assert!( + flat.equations.is_empty(), + "stream connect() must not become an ordinary equality equation" + ); + assert!( + flat.variables + .get(&rumoca_core::VarName::new("a.h_outflow")) + .is_some_and(|var| var.connected) + ); + assert!( + flat.variables + .get(&rumoca_core::VarName::new("b.h_outflow")) + .is_some_and(|var| var.connected) + ); +} + +#[test] +fn test_stream_interface_alias_generated_when_anchor_is_defined() { + let mut flat = flat::Model::new(); + for name in ["port.h_outflow", "component.port.h_outflow"] { + flat.add_variable( + rumoca_core::VarName::new(name), + flat::Variable { + name: rumoca_core::VarName::new(name), + stream: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + flat.add_equation(flat::Equation::new( + create_equality_residual( + var_to_expr( + &rumoca_core::VarName::new("component.port.h_outflow"), + test_provenance_span(), + ), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(0), + span: Span::DUMMY, + }, + test_provenance_span(), + ), + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "component.port".to_string(), + }, + )); + + let mut overlay = ast::InstanceOverlay::new(); + overlay.add_class(ast::ClassInstanceData { + instance_id: ast::InstanceId(0), + qualified_name: ast::QualifiedName::from_ident("Root"), + connections: vec![ast::InstanceConnection { + a: ast::QualifiedName::from_ident("port"), + b: ast::QualifiedName::from_dotted("component.port"), + connector_type: None, + span: Span::DUMMY, + scope: String::new(), + }], + ..Default::default() + }); + + process_connections(&mut flat, &overlay, false).expect("stream alias processing"); + + assert_eq!(flat.equations.len(), 2); + assert!(matches!( + flat.equations[1].origin, + flat::EquationOrigin::Connection { .. } + )); + assert!( + flat.variables + .get(&rumoca_core::VarName::new("port.h_outflow")) + .is_some_and(|var| var.connected) + ); + assert!( + flat.variables + .get(&rumoca_core::VarName::new("component.port.h_outflow")) + .is_some_and(|var| var.connected) + ); +} + +#[test] +fn test_stream_interface_alias_expands_collapsed_stream_endpoint() { + let mut flat = flat::Model::new(); + for name in ["port.h_outflow", "component.port.h_outflow"] { + flat.add_variable( + rumoca_core::VarName::new(name), + flat::Variable { + name: rumoca_core::VarName::new(name), + stream: true, + is_primitive: true, + dims: vec![2], + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + flat.add_equation(flat::Equation::new( + create_equality_residual( + var_to_expr( + &rumoca_core::VarName::new("component.port.h_outflow"), + test_provenance_span(), + ), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(0), + span: Span::DUMMY, + }, + test_provenance_span(), + ), + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "component.port".to_string(), + }, + )); + + let mut overlay = ast::InstanceOverlay::new(); + overlay.add_class(ast::ClassInstanceData { + instance_id: ast::InstanceId(0), + qualified_name: ast::QualifiedName::from_ident("Root"), + connections: vec![ast::InstanceConnection { + a: ast::QualifiedName::from_dotted("port.h_outflow[2]"), + b: ast::QualifiedName::from_dotted("component.port.h_outflow[2]"), + connector_type: None, + span: Span::DUMMY, + scope: String::new(), + }], + ..Default::default() + }); + + process_connections(&mut flat, &overlay, false) + .expect("collapsed stream endpoint alias processing"); + + assert_eq!(flat.equations.len(), 2); + assert!(matches!( + flat.equations[1].origin, + flat::EquationOrigin::Connection { .. } + )); +} + +#[test] +fn test_stream_non_interface_alias_generated_for_explicit_connector_port() { + let mut flat = flat::Model::new(); + for name in [ + "wrapper.sensor.port.h_outflow", + "wrapper.source.port.h_outflow", + ] { + flat.add_variable( + rumoca_core::VarName::new(name), + flat::Variable { + name: rumoca_core::VarName::new(name), + stream: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + flat.add_equation(flat::Equation::new( + create_equality_residual( + var_to_expr( + &rumoca_core::VarName::new("wrapper.source.port.h_outflow"), + test_provenance_span(), + ), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(0), + span: Span::DUMMY, + }, + test_provenance_span(), + ), + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "wrapper.source.port".to_string(), + }, + )); + + let mut overlay = ast::InstanceOverlay::new(); + overlay.add_class(ast::ClassInstanceData { + instance_id: ast::InstanceId(0), + qualified_name: ast::QualifiedName::from_ident("Root"), + connections: vec![ast::InstanceConnection { + a: ast::QualifiedName::from_dotted("wrapper.sensor.port"), + b: ast::QualifiedName::from_dotted("wrapper.source.port"), + connector_type: None, + span: Span::DUMMY, + scope: "wrapper".to_string(), + }], + ..Default::default() + }); + + process_connections(&mut flat, &overlay, false).expect("non-interface stream alias processing"); + + assert_eq!(flat.equations.len(), 2); + assert!(matches!( + flat.equations[1].origin, + flat::EquationOrigin::Connection { .. } + )); +} + +#[test] +fn test_outside_stream_connectors_are_counted_without_equality_aliases() { + let mut flat = flat::Model::new(); + for name in ["port_a.h_outflow", "port_b.h_outflow"] { + flat.add_variable( + rumoca_core::VarName::new(name), + flat::Variable { + name: rumoca_core::VarName::new(name), + stream: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + + let mut overlay = ast::InstanceOverlay::new(); + overlay.add_class(ast::ClassInstanceData { + instance_id: ast::InstanceId(0), + qualified_name: ast::QualifiedName::from_ident("Root"), + connections: vec![ast::InstanceConnection { + a: ast::QualifiedName::from_dotted("port_a"), + b: ast::QualifiedName::from_dotted("port_b"), + connector_type: None, + span: Span::DUMMY, + scope: String::new(), + }], + ..Default::default() + }); + + process_connections(&mut flat, &overlay, false).expect("outside stream connector processing"); + + assert_eq!(flat.equations.len(), 0); + assert_eq!(flat.stream_interface_equation_count, 2); + assert!( + flat.variables + .get(&rumoca_core::VarName::new("port_a.h_outflow")) + .is_some_and(|var| var.connected) + ); + assert!( + flat.variables + .get(&rumoca_core::VarName::new("port_b.h_outflow")) + .is_some_and(|var| var.connected) + ); +} + +#[test] +fn test_internal_stream_connectors_without_anchor_are_counted() { + let mut flat = flat::Model::new(); + for name in [ + "wrapper.a.port.p", + "wrapper.a.port.m_flow", + "wrapper.a.port.h_outflow", + "wrapper.b.port.p", + "wrapper.b.port.m_flow", + "wrapper.b.port.h_outflow", + ] { + let mut variable = flat::Variable { + name: rumoca_core::VarName::new(name), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }; + if name.ends_with(".m_flow") { + variable.flow = true; + } + if name.ends_with(".h_outflow") { + variable.stream = true; + } + flat.add_variable(rumoca_core::VarName::new(name), variable); + } + + let mut overlay = ast::InstanceOverlay::new(); + overlay.add_class(ast::ClassInstanceData { + instance_id: ast::InstanceId(0), + qualified_name: ast::QualifiedName::from_ident("Root"), + connections: vec![ast::InstanceConnection { + a: ast::QualifiedName::from_dotted("wrapper.a.port"), + b: ast::QualifiedName::from_dotted("wrapper.b.port"), + connector_type: None, + span: Span::DUMMY, + scope: "wrapper".to_string(), + }], + ..Default::default() + }); + + process_connections(&mut flat, &overlay, false).expect("internal stream connector processing"); + + assert_eq!(flat.stream_interface_equation_count, 2); + assert!( + flat.variables + .get(&rumoca_core::VarName::new("wrapper.a.port.h_outflow")) + .is_some_and(|var| var.connected) + ); + assert!( + flat.variables + .get(&rumoca_core::VarName::new("wrapper.b.port.h_outflow")) + .is_some_and(|var| var.connected) + ); +} + +#[test] +fn test_multiway_stream_connectors_with_component_defined_outflow_are_counted() { + let mut flat = flat::Model::new(); + for name in ["a.h_outflow", "b.h_outflow", "c.h_outflow"] { + let var_name = rumoca_core::VarName::new(name); + flat.add_variable( + var_name.clone(), + flat::Variable { + name: var_name.clone(), + stream: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + flat.add_equation(flat::Equation::new( + create_equality_residual( + var_to_expr(&var_name, test_provenance_span()), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(1.0), + span: test_span(), + }, + test_provenance_span(), + ), + test_span(), + flat::EquationOrigin::ComponentEquation { + component: name.to_string(), + }, + )); + } + + let mut overlay = ast::InstanceOverlay::new(); + overlay.add_class(ast::ClassInstanceData { + instance_id: ast::InstanceId(0), + qualified_name: ast::QualifiedName::from_ident("Root"), + connections: vec![ + ast::InstanceConnection { + a: ast::QualifiedName::from_dotted("a.h_outflow"), + b: ast::QualifiedName::from_dotted("b.h_outflow"), + connector_type: None, + span: Span::DUMMY, + scope: String::new(), + }, + ast::InstanceConnection { + a: ast::QualifiedName::from_dotted("a.h_outflow"), + b: ast::QualifiedName::from_dotted("c.h_outflow"), + connector_type: None, + span: Span::DUMMY, + scope: String::new(), + }, + ], + ..Default::default() + }); + + process_connections(&mut flat, &overlay, false).expect("multiway stream connector processing"); + + assert_eq!(flat.stream_interface_equation_count, 3); + for name in ["a.h_outflow", "b.h_outflow", "c.h_outflow"] { + assert!( + flat.variables + .get(&rumoca_core::VarName::new(name)) + .is_some_and(|var| var.connected), + "{name} should be marked connected" + ); + } +} + +#[test] +fn test_stream_internal_dynbal_port_alias_is_not_generated() { + let mut flat = flat::Model::new(); + for name in [ + "wrapper.vol.dynBal.ports[1].Xi_outflow", + "wrapper.source.port.Xi_outflow", + ] { + flat.add_variable( + rumoca_core::VarName::new(name), + flat::Variable { + name: rumoca_core::VarName::new(name), + stream: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + flat.add_equation(flat::Equation::new( + create_equality_residual( + var_to_expr( + &rumoca_core::VarName::new("wrapper.source.port.Xi_outflow"), + test_provenance_span(), + ), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(0), + span: Span::DUMMY, + }, + test_provenance_span(), + ), + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "wrapper.source.port".to_string(), + }, + )); + + let mut overlay = ast::InstanceOverlay::new(); + overlay.add_class(ast::ClassInstanceData { + instance_id: ast::InstanceId(0), + qualified_name: ast::QualifiedName::from_ident("Root"), + connections: vec![ast::InstanceConnection { + a: ast::QualifiedName::from_dotted("wrapper.vol.dynBal.ports[1]"), + b: ast::QualifiedName::from_dotted("wrapper.source.port"), + connector_type: None, + span: Span::DUMMY, + scope: "wrapper.vol.dynBal".to_string(), + }], + ..Default::default() + }); + + process_connections(&mut flat, &overlay, false) + .expect("internal dynBal stream alias processing"); + + assert_eq!(flat.equations.len(), 1); +} + +#[test] +fn test_stream_internal_dynbal_energy_port_alias_is_generated() { + let mut flat = flat::Model::new(); + for name in [ + "wrapper.vol.dynBal.ports[1].h_outflow", + "wrapper.source.port.h_outflow", + ] { + flat.add_variable( + rumoca_core::VarName::new(name), + flat::Variable { + name: rumoca_core::VarName::new(name), + stream: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + flat.add_equation(flat::Equation::new( + create_equality_residual( + var_to_expr( + &rumoca_core::VarName::new("wrapper.source.port.h_outflow"), + test_provenance_span(), + ), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(0), + span: Span::DUMMY, + }, + test_provenance_span(), + ), + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "wrapper.source.port".to_string(), + }, + )); + + let mut overlay = ast::InstanceOverlay::new(); + overlay.add_class(ast::ClassInstanceData { + instance_id: ast::InstanceId(0), + qualified_name: ast::QualifiedName::from_ident("Root"), + connections: vec![ast::InstanceConnection { + a: ast::QualifiedName::from_dotted("wrapper.vol.dynBal.ports[1]"), + b: ast::QualifiedName::from_dotted("wrapper.source.port"), + connector_type: None, + span: Span::DUMMY, + scope: "wrapper.vol.dynBal".to_string(), + }], + ..Default::default() + }); + + process_connections(&mut flat, &overlay, false) + .expect("internal dynBal energy stream alias processing"); + + assert_eq!(flat.equations.len(), 2); + assert!(matches!( + flat.equations[1].origin, + flat::EquationOrigin::Connection { .. } + )); +} + +#[test] +fn test_stream_pass_through_alias_generated_from_defined_dynbal_port() { + let mut flat = flat::Model::new(); + for name in [ + "wrapper.vol.dynBal.ports[1].h_outflow", + "wrapper.vol.ports[1].h_outflow", + "wrapper.port_b.h_outflow", + ] { + flat.add_variable( + rumoca_core::VarName::new(name), + flat::Variable { + name: rumoca_core::VarName::new(name), + stream: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + flat.add_equation(flat::Equation::new( + create_equality_residual( + var_to_expr( + &rumoca_core::VarName::new("wrapper.vol.dynBal.ports[1].h_outflow"), + test_provenance_span(), + ), + var_to_expr( + &rumoca_core::VarName::new("wrapper.vol.dynBal.medium.h"), + test_provenance_span(), + ), + test_provenance_span(), + ), + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "wrapper.vol.dynBal".to_string(), + }, + )); + + let mut overlay = ast::InstanceOverlay::new(); + overlay.add_class(ast::ClassInstanceData { + instance_id: ast::InstanceId(0), + qualified_name: ast::QualifiedName::from_ident("Root"), + connections: vec![ + ast::InstanceConnection { + a: ast::QualifiedName::from_dotted("wrapper.vol.ports"), + b: ast::QualifiedName::from_dotted("wrapper.vol.dynBal.ports"), + connector_type: None, + span: Span::DUMMY, + scope: "wrapper.vol".to_string(), + }, + ast::InstanceConnection { + a: ast::QualifiedName::from_dotted("wrapper.vol.ports[1]"), + b: ast::QualifiedName::from_dotted("wrapper.port_b"), + connector_type: None, + span: Span::DUMMY, + scope: "wrapper".to_string(), + }, + ], + ..Default::default() + }); + + process_connections(&mut flat, &overlay, false).expect("stream pass-through alias processing"); + + let origins = flat + .equations + .iter() + .map(|eq| eq.origin.to_string()) + .collect::>(); assert!( - potential_uf.get_sets().is_empty(), - "stream connect() must not generate potential equality sets" + origins + .iter() + .any(|origin| origin.contains("wrapper.vol.ports[1].h_outflow")), + "vol.ports stream alias should be generated from dynBal anchor: {origins:?}" + ); + assert!( + origins + .iter() + .any(|origin| origin.contains("wrapper.port_b.h_outflow")), + "outer port_b stream alias should be generated from dynBal anchor: {origins:?}" ); - assert_eq!(stream_uf.get_sets().len(), 1); } #[test] -fn test_stream_connection_does_not_generate_potential_equality() { +fn test_stream_pass_through_alias_generated_for_all_array_ports() { let mut flat = flat::Model::new(); - flat.add_variable( - rumoca_core::VarName::new("a.h_outflow"), - flat::Variable { - name: rumoca_core::VarName::new("a.h_outflow"), - stream: true, - source_span: test_span(), - ..flat::Variable::empty_with_span(test_span()) - }, - ); - flat.add_variable( - rumoca_core::VarName::new("b.h_outflow"), - flat::Variable { - name: rumoca_core::VarName::new("b.h_outflow"), - stream: true, - source_span: test_span(), - ..flat::Variable::empty_with_span(test_span()) + for name in [ + "wrapper.vol.dynBal.ports[1].h_outflow", + "wrapper.vol.dynBal.ports[2].h_outflow", + "wrapper.vol.ports[1].h_outflow", + "wrapper.vol.ports[2].h_outflow", + "wrapper.port_a.h_outflow", + "wrapper.port_b.h_outflow", + ] { + flat.add_variable( + rumoca_core::VarName::new(name), + flat::Variable { + name: rumoca_core::VarName::new(name), + stream: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + for idx in 1..=2 { + flat.add_equation(flat::Equation::new( + create_equality_residual( + var_to_expr( + &rumoca_core::VarName::new(format!( + "wrapper.vol.dynBal.ports[{idx}].h_outflow" + )), + test_provenance_span(), + ), + var_to_expr( + &rumoca_core::VarName::new("wrapper.vol.dynBal.medium.h"), + test_provenance_span(), + ), + test_provenance_span(), + ), + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "wrapper.vol.dynBal".to_string(), + }, + )); + } + + let mut overlay = ast::InstanceOverlay::new(); + overlay.add_class(ast::ClassInstanceData { + instance_id: ast::InstanceId(0), + qualified_name: ast::QualifiedName::from_ident("Root"), + connections: vec![ + ast::InstanceConnection { + a: ast::QualifiedName::from_dotted("wrapper.vol.ports"), + b: ast::QualifiedName::from_dotted("wrapper.vol.dynBal.ports"), + connector_type: None, + span: Span::DUMMY, + scope: "wrapper.vol".to_string(), + }, + ast::InstanceConnection { + a: ast::QualifiedName::from_dotted("wrapper.vol.ports[1]"), + b: ast::QualifiedName::from_dotted("wrapper.port_a"), + connector_type: None, + span: Span::DUMMY, + scope: "wrapper".to_string(), + }, + ast::InstanceConnection { + a: ast::QualifiedName::from_dotted("wrapper.vol.ports[2]"), + b: ast::QualifiedName::from_dotted("wrapper.port_b"), + connector_type: None, + span: Span::DUMMY, + scope: "wrapper".to_string(), + }, + ], + ..Default::default() + }); + + process_connections(&mut flat, &overlay, false) + .expect("array stream pass-through alias processing"); + + let origins = flat + .equations + .iter() + .map(|eq| eq.origin.to_string()) + .collect::>(); + for expected in [ + "wrapper.vol.ports[1].h_outflow", + "wrapper.vol.ports[2].h_outflow", + "wrapper.port_a.h_outflow", + "wrapper.port_b.h_outflow", + ] { + assert!( + origins.iter().any(|origin| origin.contains(expected)), + "stream alias should be generated for {expected}: {origins:?}" + ); + } +} + +#[test] +fn test_stream_pass_through_alias_survives_top_level_boundary_port() { + let mut flat = flat::Model::new(); + flat.top_level_connectors.insert("port_b".to_string()); + for name in [ + "wrapper.vol.dynBal.ports[1].h_outflow", + "wrapper.vol.ports[1].h_outflow", + "wrapper.port_b.h_outflow", + "port_b.h_outflow", + ] { + flat.add_variable( + rumoca_core::VarName::new(name), + flat::Variable { + name: rumoca_core::VarName::new(name), + stream: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + flat.add_equation(flat::Equation::new( + create_equality_residual( + var_to_expr( + &rumoca_core::VarName::new("wrapper.vol.dynBal.ports[1].h_outflow"), + test_provenance_span(), + ), + var_to_expr( + &rumoca_core::VarName::new("wrapper.vol.dynBal.medium.h"), + test_provenance_span(), + ), + test_provenance_span(), + ), + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "wrapper.vol.dynBal".to_string(), }, - ); + )); let mut overlay = ast::InstanceOverlay::new(); overlay.add_class(ast::ClassInstanceData { instance_id: ast::InstanceId(0), qualified_name: ast::QualifiedName::from_ident("Root"), - connections: vec![ast::InstanceConnection { - a: ast::QualifiedName::from_dotted("a.h_outflow"), - b: ast::QualifiedName::from_dotted("b.h_outflow"), - connector_type: None, - span: Span::DUMMY, - scope: String::new(), - }], + connections: vec![ + ast::InstanceConnection { + a: ast::QualifiedName::from_dotted("wrapper.vol.ports"), + b: ast::QualifiedName::from_dotted("wrapper.vol.dynBal.ports"), + connector_type: None, + span: Span::DUMMY, + scope: "wrapper.vol".to_string(), + }, + ast::InstanceConnection { + a: ast::QualifiedName::from_dotted("wrapper.vol.ports[1]"), + b: ast::QualifiedName::from_dotted("wrapper.port_b"), + connector_type: None, + span: Span::DUMMY, + scope: "wrapper".to_string(), + }, + ast::InstanceConnection { + a: ast::QualifiedName::from_dotted("wrapper.port_b"), + b: ast::QualifiedName::from_dotted("port_b"), + connector_type: None, + span: Span::DUMMY, + scope: String::new(), + }, + ], ..Default::default() }); - process_connections(&mut flat, &overlay, false).expect("stream connection processing"); + process_connections(&mut flat, &overlay, false) + .expect("top-level boundary stream pass-through alias processing"); + let origins = flat + .equations + .iter() + .map(|eq| eq.origin.to_string()) + .collect::>(); assert!( - flat.equations.is_empty(), - "stream connect() must not become an ordinary equality equation" + origins + .iter() + .any(|origin| origin.contains("wrapper.vol.ports[1].h_outflow")), + "internal vol.ports stream alias should not be suppressed by top-level boundary: {origins:?}" ); assert!( - flat.variables - .get(&rumoca_core::VarName::new("a.h_outflow")) - .is_some_and(|var| var.connected) + origins + .iter() + .any(|origin| origin.contains("wrapper.port_b.h_outflow")), + "internal component port stream alias should not be suppressed by top-level boundary: {origins:?}" + ); + assert!( + origins + .iter() + .all(|origin| !origin.contains("connection equation: port_b.h_outflow")), + "top-level stream boundary must not be converted into an equality alias: {origins:?}" ); assert!( flat.variables - .get(&rumoca_core::VarName::new("b.h_outflow")) - .is_some_and(|var| var.connected) + .get(&rumoca_core::VarName::new("port_b.h_outflow")) + .is_some_and(|var| var.connected), + "top-level stream boundary should still be marked connected" ); } @@ -1003,6 +1745,65 @@ fn test_validate_dimension_compatibility_mismatch() { assert!(result.is_err()); } +#[test] +fn test_validate_dimension_compatibility_scalar_accepts_length_one_array() { + let mut flat = flat::Model::new(); + + flat.add_variable( + rumoca_core::VarName::new("y"), + flat::Variable::empty_with_span(test_span()), + ); + flat.add_variable( + rumoca_core::VarName::new("u"), + flat::Variable { + dims: vec![1], + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + let result = validate_dimension_compatibility( + &flat, + &rumoca_core::VarName::new("y"), + &rumoca_core::VarName::new("u"), + Span::DUMMY, + ); + assert!( + result.is_ok(), + "scalar connector and length-one connector array both have one connection element" + ); +} + +#[test] +fn test_validate_dimension_compatibility_accepts_full_reduced_composition_pair() { + let mut flat = flat::Model::new(); + + flat.add_variable( + rumoca_core::VarName::new("src.X_in_internal"), + flat::Variable { + dims: vec![2], + ..flat::Variable::empty_with_span(test_span()) + }, + ); + flat.add_variable( + rumoca_core::VarName::new("src.Xi_in_internal"), + flat::Variable { + dims: vec![1], + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + let result = validate_dimension_compatibility( + &flat, + &rumoca_core::VarName::new("src.X_in_internal"), + &rumoca_core::VarName::new("src.Xi_in_internal"), + Span::DUMMY, + ); + assert!( + result.is_ok(), + "Modelica.Media full composition X[nX] may connect to reduced Xi[nX-1] through X[1:nXi]" + ); +} + #[test] fn test_validate_dimension_compatibility_io_mismatch_still_fails() { let mut flat = flat::Model::new(); @@ -1061,6 +1862,146 @@ fn test_validate_dimension_compatibility_partial_subscript_projects_remaining_di ); } +#[test] +fn test_validate_dimension_compatibility_embedded_subscript_projects_exact_vars() { + let mut flat = flat::Model::new(); + + let lhs = flat::Variable { + dims: vec![6], + ..flat::Variable::empty_with_span(test_span()) + }; + flat.add_variable(rumoca_core::VarName::new("a.pin[1].v"), lhs); + + let rhs = flat::Variable { + dims: vec![2], + ..flat::Variable::empty_with_span(test_span()) + }; + flat.add_variable(rumoca_core::VarName::new("b.pin[1].v"), rhs); + + let result = validate_dimension_compatibility( + &flat, + &rumoca_core::VarName::new("a.pin[1].v"), + &rumoca_core::VarName::new("b.pin[1].v"), + Span::DUMMY, + ); + assert!( + result.is_ok(), + "indexed connector members are scalar element connections even when flat vars retain parent dims" + ); +} + +#[test] +fn test_validate_dimension_compatibility_embedded_subscript_preserves_parent_dims() { + let mut flat = flat::Model::new(); + + let lhs = flat::Variable { + dims: vec![6], + ..flat::Variable::empty_with_span(test_span()) + }; + flat.add_variable(rumoca_core::VarName::new("harmonic.sin1.y"), lhs); + + let rhs = flat::Variable { + dims: vec![6, 2], + ..flat::Variable::empty_with_span(test_span()) + }; + flat.add_variable(rumoca_core::VarName::new("harmonic.product1.u[1]"), rhs); + + let result = validate_dimension_compatibility( + &flat, + &rumoca_core::VarName::new("harmonic.sin1.y"), + &rumoca_core::VarName::new("harmonic.product1.u[1]"), + Span::DUMMY, + ); + assert!( + result.is_ok(), + "embedded subscripts consume child dimensions while preserving component-array parent dimensions" + ); +} + +#[test] +fn test_validate_dimension_compatibility_indexed_child_can_use_counterpart_parent_dims() { + let mut flat = flat::Model::new(); + + let lhs = flat::Variable { + dims: vec![6], + ..flat::Variable::empty_with_span(test_span()) + }; + flat.add_variable(rumoca_core::VarName::new("harmonic.sin1.y"), lhs); + + let rhs = flat::Variable { + dims: vec![2], + ..flat::Variable::empty_with_span(test_span()) + }; + flat.add_variable(rumoca_core::VarName::new("harmonic.product1.u[1]"), rhs); + + let result = validate_dimension_compatibility( + &flat, + &rumoca_core::VarName::new("harmonic.sin1.y"), + &rumoca_core::VarName::new("harmonic.product1.u[1]"), + Span::DUMMY, + ); + assert!( + result.is_ok(), + "collapsed indexed children may be missing component-array parent dims carried by the counterpart" + ); +} + +#[test] +fn test_validate_dimension_compatibility_accepts_collapsed_pin_member_subsystem_dims() { + let mut flat = flat::Model::new(); + + let lhs = flat::Variable { + dims: vec![6], + ..flat::Variable::empty_with_span(test_span()) + }; + flat.add_variable(rumoca_core::VarName::new("multi.plug_p.pin.v"), lhs); + + let rhs = flat::Variable { + dims: vec![2], + ..flat::Variable::empty_with_span(test_span()) + }; + flat.add_variable(rumoca_core::VarName::new("multi.starpoints.pin.v"), rhs); + + let result = validate_dimension_compatibility( + &flat, + &rumoca_core::VarName::new("multi.plug_p.pin.v"), + &rumoca_core::VarName::new("multi.starpoints.pin.v"), + Span::DUMMY, + ); + assert!( + result.is_ok(), + "collapsed pin member validation must not reject loop-expanded scalar subsystem connections" + ); +} + +#[test] +fn test_validate_dimension_compatibility_accepts_collapsed_pin_member_to_scalar_pin() { + let mut flat = flat::Model::new(); + + let lhs = flat::Variable { + dims: vec![2], + ..flat::Variable::empty_with_span(test_span()) + }; + flat.add_variable(rumoca_core::VarName::new("star.plug_p.pin.v"), lhs); + + let rhs = flat::Variable { + dims: Vec::new(), + ..flat::Variable::empty_with_span(test_span()) + }; + flat.add_variable(rumoca_core::VarName::new("star.pin_n.v"), rhs); + + let result = validate_dimension_compatibility( + &flat, + &rumoca_core::VarName::new("star.plug_p.pin.v"), + &rumoca_core::VarName::new("star.pin_n.v"), + Span::DUMMY, + ); + assert!( + result.is_ok(), + "collapsed for-loop pin element connections may map a plug pin array onto a scalar pin" + ); +} + #[test] fn test_validate_dimension_compatibility_partial_subscript_mismatch_fails() { let mut flat = flat::Model::new(); diff --git a/crates/rumoca-phase-flatten/src/constant_extraction.rs b/crates/rumoca-phase-flatten/src/constant_extraction.rs index e791c3835..aba848952 100644 --- a/crates/rumoca-phase-flatten/src/constant_extraction.rs +++ b/crates/rumoca-phase-flatten/src/constant_extraction.rs @@ -1,15 +1,14 @@ use super::*; +// Well-known constant packages to resolve. ModelicaServices.Machine must come +// first since Modelica.Constants.eps references ModelicaServices.Machine.eps. +const WELL_KNOWN_CONSTANT_PACKAGES: &[&str] = &["ModelicaServices.Machine", "Modelica.Constants"]; + pub(super) fn resolve_constants_from_tree( tree: &ast::ClassTree, eval_ctx: &mut rumoca_eval_flat::constant::EvalContext, ) -> Result<(), FlattenError> { - // Well-known constant packages to resolve. - // ModelicaServices.Machine must come first since Modelica.Constants.eps - // references ModelicaServices.Machine.eps. - const CONSTANT_PACKAGES: &[&str] = &["ModelicaServices.Machine", "Modelica.Constants"]; - - for &pkg_name in CONSTANT_PACKAGES { + for &pkg_name in WELL_KNOWN_CONSTANT_PACKAGES { let Some(class_def) = tree.get_class_by_qualified_name(pkg_name) else { continue; }; @@ -31,6 +30,86 @@ pub(super) fn resolve_constants_from_tree( Ok(()) } +pub(super) fn inject_well_known_constant_package_values( + tree: &ast::ClassTree, + ctx: &mut Context, +) -> Result<(), FlattenError> { + let mut eval_ctx = rumoca_eval_flat::constant::EvalContext::with_capacity( + ctx.parameter_values.len() + + ctx.real_parameter_values.len() + + ctx.boolean_parameter_values.len() + + ctx.string_parameter_values.len(), + 0, + ctx.functions.len() * 2, + ); + for (name, value) in &ctx.parameter_values { + eval_ctx.add_parameter( + name.clone(), + rumoca_eval_flat::constant::Value::Integer(*value), + ); + } + for (name, value) in &ctx.real_parameter_values { + eval_ctx.add_parameter( + name.clone(), + rumoca_eval_flat::constant::Value::Real(*value), + ); + } + for (name, value) in &ctx.boolean_parameter_values { + eval_ctx.add_parameter( + name.clone(), + rumoca_eval_flat::constant::Value::Bool(*value), + ); + } + for (name, value) in &ctx.string_parameter_values { + eval_ctx.add_parameter( + name.clone(), + rumoca_eval_flat::constant::Value::String(value.clone()), + ); + } + + resolve_constants_from_tree(tree, &mut eval_ctx)?; + for &pkg_name in WELL_KNOWN_CONSTANT_PACKAGES { + let Some(class_def) = tree.get_class_by_qualified_name(pkg_name) else { + continue; + }; + for (comp_name, component) in &class_def.components { + if !matches!( + component.variability, + rumoca_core::Variability::Constant(_) | rumoca_core::Variability::Parameter(_) + ) { + continue; + } + let qualified = format!("{pkg_name}.{comp_name}"); + if matches!(component.variability, rumoca_core::Variability::Constant(_)) { + ctx.class_constant_keys.insert(qualified.clone()); + } + let Some(value) = eval_ctx.get(&qualified) else { + continue; + }; + match value { + rumoca_eval_flat::constant::Value::Real(value) if value.is_finite() => { + ctx.real_parameter_values.entry(qualified).or_insert(*value); + } + rumoca_eval_flat::constant::Value::Integer(value) => { + ctx.parameter_values.entry(qualified).or_insert(*value); + } + rumoca_eval_flat::constant::Value::Bool(value) => { + ctx.boolean_parameter_values + .entry(qualified) + .or_insert(*value); + } + rumoca_eval_flat::constant::Value::String(value) => { + ctx.string_parameter_values + .entry(qualified) + .or_insert_with(|| value.clone()); + } + _ => {} + } + } + } + Ok(()) +} + pub(super) fn inject_referenced_qualified_class_constants( tree: &ClassTree, class_index: &ast::ClassDefIndex<'_>, @@ -39,9 +118,8 @@ pub(super) fn inject_referenced_qualified_class_constants( overlay: &InstanceOverlay, ctx: &mut Context, ) -> Result<(), FlattenError> { - const WELL_KNOWN_CONSTANT_PACKAGES: &[&str] = - &["ModelicaServices.Machine", "Modelica.Constants"]; const MAX_PASSES: usize = 4; + inject_well_known_constant_package_values(tree, ctx)?; let live_vars: HashSet = flat .variables .keys() @@ -190,6 +268,7 @@ pub(super) fn context_constant_footprint(ctx: &Context) -> usize { ctx.parameter_values.len() + ctx.real_parameter_values.len() + ctx.boolean_parameter_values.len() + + ctx.string_parameter_values.len() + ctx.enum_parameter_values.len() + ctx.constant_values.len() + ctx.array_dimensions.len() @@ -1283,12 +1362,13 @@ fn extract_extends_shape_and_alias_modification( ); } if let Some(dims) = infer_dims_from_expr(qualified_value, ctx, prefix) { - insert_with_prefix( - &mut ctx.array_dimensions, + insert_array_dimensions_with_prefix( + ctx, prefix, target_name, full_name, dims, + qualified_value.span(), ); } if let Some(alias_name) = alias_ref { @@ -1310,14 +1390,7 @@ pub(super) fn extract_constants_from_class_with_prefix( ) { continue; } - let binding = - comp.binding - .as_ref() - .or(if !matches!(comp.start, ast::Expression::Empty { .. }) { - Some(&comp.start) - } else { - None - }); + let binding = comp.binding.as_ref(); let synthesized = if binding.is_none() { synthesize_component_modification_binding(comp) } else { @@ -1337,6 +1410,26 @@ pub(super) fn extract_constants_from_class_with_prefix_and_imports( class_def: &ast::ClassDef, resolve_context: &str, ctx: &mut Context, +) { + extract_constants_from_class_with_prefix_and_imports_shadowed( + tree, + class_index, + prefix, + class_def, + resolve_context, + ctx, + &rustc_hash::FxHashSet::default(), + ); +} + +pub(super) fn extract_constants_from_class_with_prefix_and_imports_shadowed( + tree: &ast::ClassTree, + class_index: &ast::ClassDefIndex<'_>, + prefix: &str, + class_def: &ast::ClassDef, + resolve_context: &str, + ctx: &mut Context, + shadowed_names: &rustc_hash::FxHashSet, ) { let imports = constant_extraction_imports(tree, class_index, resolve_context); let empty_prefix = ast::QualifiedName::new(); @@ -1346,20 +1439,16 @@ pub(super) fn extract_constants_from_class_with_prefix_and_imports( }; for (name, comp) in &class_def.components { + if shadowed_names.contains(name) { + continue; + } if !matches!( comp.variability, rumoca_core::Variability::Constant(_) | rumoca_core::Variability::Parameter(_) ) { continue; } - let binding = - comp.binding - .as_ref() - .or(if !matches!(comp.start, ast::Expression::Empty { .. }) { - Some(&comp.start) - } else { - None - }); + let binding = comp.binding.as_ref(); let synthesized = if binding.is_none() { synthesize_component_modification_binding(comp) } else { @@ -1460,6 +1549,12 @@ pub(super) fn extract_single_constant_with_prefix_and_function_scope( let type_name = comp.type_name.to_string(); let preserve_existing = ctx.flat_parameter_constant_keys.contains(full_name); + if matches!(comp.variability, rumoca_core::Variability::Constant(_)) { + ctx.class_constant_keys.insert(full_name.to_string()); + if exposes_unprefixed_prefix(prefix) { + ctx.class_constant_keys.insert(name.to_string()); + } + } if let Some(val) = try_extract_named_record_constructor_constant(expr, ctx, prefix, full_name) && (!preserve_existing || !ctx.constant_values.contains_key(full_name)) { @@ -1495,6 +1590,21 @@ pub(super) fn extract_single_constant_with_prefix_and_function_scope( val, ); } + if type_name == "String" + && let Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String(val), + .. + }) = try_eval_const_flat_expr_with_scope(expr, ctx, prefix) + && (!preserve_existing || !ctx.string_parameter_values.contains_key(full_name)) + { + insert_with_prefix( + &mut ctx.string_parameter_values, + prefix, + name, + full_name, + val, + ); + } // Real constants (and constants of aliased Real-derived units). if let Some(val) = try_eval_const_real_with_scope(expr, ctx, prefix) && val.is_finite() @@ -1524,16 +1634,37 @@ pub(super) fn extract_single_constant_with_prefix_and_function_scope( && !comp.shape.is_empty() { let dims: Vec = comp.shape.iter().map(|&d| d as i64).collect(); - insert_with_prefix(&mut ctx.array_dimensions, prefix, name, full_name, dims); + insert_array_dimensions_with_prefix(ctx, prefix, name, full_name, dims, expr.span()); } // Array dimensions from binding (array literal length) if (!preserve_existing || !ctx.array_dimensions.contains_key(full_name)) && let Some(dims) = infer_dims_from_expr(expr, ctx, prefix) { - insert_with_prefix(&mut ctx.array_dimensions, prefix, name, full_name, dims); + insert_array_dimensions_with_prefix(ctx, prefix, name, full_name, dims, expr.span()); } } +fn insert_array_dimensions_with_prefix( + ctx: &mut Context, + prefix: &str, + name: &str, + full_name: &str, + dims: Vec, + span: rumoca_core::Span, +) { + insert_with_prefix(&mut ctx.array_dimensions, prefix, name, full_name, dims); + if span.is_dummy() { + return; + } + insert_with_prefix( + &mut ctx.array_dimension_spans, + prefix, + name, + full_name, + span, + ); +} + pub(super) fn try_extract_constant_alias_expr( expr: &ast::Expression, ) -> Option { @@ -1592,6 +1723,13 @@ pub(super) fn materialize_record_constant_alias_fields( full_name, alias_name, ); + propagate_alias_fields_in_map( + &mut ctx.string_parameter_values, + prefix, + target_name, + full_name, + alias_name, + ); propagate_alias_fields_in_map( &mut ctx.enum_parameter_values, prefix, @@ -1613,6 +1751,13 @@ pub(super) fn materialize_record_constant_alias_fields( full_name, alias_name, ); + propagate_alias_fields_in_map( + &mut ctx.array_dimension_spans, + prefix, + target_name, + full_name, + alias_name, + ); } pub(super) fn propagate_alias_fields_in_map( diff --git a/crates/rumoca-phase-flatten/src/equations/conditional_and_eval.rs b/crates/rumoca-phase-flatten/src/equations/conditional_and_eval.rs index e2d44209a..7ad3a31c2 100644 --- a/crates/rumoca-phase-flatten/src/equations/conditional_and_eval.rs +++ b/crates/rumoca-phase-flatten/src/equations/conditional_and_eval.rs @@ -408,6 +408,9 @@ pub(crate) fn expand_range_indices( } Ok(indices) } + ast::Expression::FunctionCall { comp, args, .. } if is_linspace_function_call(comp) => { + expand_linspace_range_indices(ctx, args, prefix, &scope, span) + } // Handle a single integer as 1:n _ => { if let Some(n) = try_eval_integer_with_ctx(ctx, range_expr, prefix) { @@ -425,6 +428,73 @@ pub(crate) fn expand_range_indices( } } +fn is_linspace_function_call(comp: &ast::ComponentReference) -> bool { + comp.parts + .last() + .is_some_and(|part| part.ident.text.as_ref() == "linspace") +} + +fn expand_linspace_range_indices( + ctx: &Context, + args: &[ast::Expression], + prefix: &QualifiedName, + scope: &str, + span: rumoca_core::Span, +) -> Result, FlattenError> { + if args.len() != 3 { + return Err(FlattenError::unsupported_equation( + format!("for-equation linspace range requires exactly 3 arguments (scope `{scope}`)"), + span, + )); + } + + let start = try_eval_integer_with_ctx(ctx, &args[0], prefix).ok_or_else(|| { + FlattenError::unsupported_equation( + format!( + "for-equation linspace start must be a constant integer or parameter (scope `{scope}`, got `{}`)", + format_subscript_expr(&args[0]), + ), + span, + ) + })?; + let end = try_eval_integer_with_ctx(ctx, &args[1], prefix).ok_or_else(|| { + FlattenError::unsupported_equation( + format!( + "for-equation linspace end must be a constant integer or parameter (scope `{scope}`, got `{}`)", + format_subscript_expr(&args[1]), + ), + span, + ) + })?; + let count = try_eval_integer_with_ctx(ctx, &args[2], prefix).ok_or_else(|| { + FlattenError::unsupported_equation( + format!( + "for-equation linspace count must be a constant integer or parameter (scope `{scope}`, got `{}`)", + format_subscript_expr(&args[2]), + ), + span, + ) + })?; + + if count <= 0 { + return Err(FlattenError::unsupported_equation( + format!("for-equation linspace count must be positive (scope `{scope}`, got {count})"), + span, + )); + } + if count == 1 { + return Ok(vec![start]); + } + + let mut indices = Vec::with_capacity(count as usize); + for i in 0..count { + let numerator = i * (end - start); + let value = start + numerator / (count - 1); + indices.push(value); + } + Ok(indices) +} + /// Try to evaluate an expression to a constant integer, with parameter lookup. pub(crate) fn try_eval_integer_with_ctx( ctx: &Context, @@ -1400,6 +1470,7 @@ fn substitute_index_in_subscript( pub(crate) fn build_eval_context(ctx: &Context, tree: Option<&ClassTree>) -> EvalContext { let parameter_capacity = ctx.parameter_values.len() + ctx.boolean_parameter_values.len() + + ctx.string_parameter_values.len() + ctx.enum_parameter_values.len() + ctx.array_dimensions.len(); let mut eval_ctx = EvalContext::with_capacity(parameter_capacity, 0, ctx.functions.len() * 2); @@ -1414,6 +1485,10 @@ pub(crate) fn build_eval_context(ctx: &Context, tree: Option<&ClassTree>) -> Eva eval_ctx.add_parameter(name.clone(), Value::Bool(*value)); } + for (name, value) in &ctx.string_parameter_values { + eval_ctx.add_parameter(name.clone(), Value::String(value.clone())); + } + // Add enum parameters for (name, value) in &ctx.enum_parameter_values { // The value is a qualified enum literal like "Type.Literal" diff --git a/crates/rumoca-phase-flatten/src/equations/connections_graph.rs b/crates/rumoca-phase-flatten/src/equations/connections_graph.rs index cb6beb1d3..306b2a763 100644 --- a/crates/rumoca-phase-flatten/src/equations/connections_graph.rs +++ b/crates/rumoca-phase-flatten/src/equations/connections_graph.rs @@ -67,7 +67,7 @@ pub(super) fn is_side_effect_only_function(comp: &ast::ComponentReference) -> bo // Check for specific side-effect-only functions. matches!( func_name, - "assert" | "terminate" | "print" | "close" | "readLine" | "error" + "assert" | "terminate" | "print" | "close" | "readLine" | "error" | "checkBoundary" ) || is_streams_utility_function(comp) || is_connections_graph_function(comp) } @@ -193,6 +193,9 @@ mod tests { assert!(is_side_effect_only_function(&cref( "Modelica.Utilities.Streams.print" ))); + assert!(is_side_effect_only_function(&cref( + "Modelica.Fluid.Utilities.checkBoundary" + ))); assert!(is_side_effect_only_function(&cref("print"))); assert!(!is_side_effect_only_function(&cref("sin"))); } diff --git a/crates/rumoca-phase-flatten/src/equations/mod.rs b/crates/rumoca-phase-flatten/src/equations/mod.rs index eb5bd61e0..d345cc45a 100644 --- a/crates/rumoca-phase-flatten/src/equations/mod.rs +++ b/crates/rumoca-phase-flatten/src/equations/mod.rs @@ -1,5 +1,8 @@ //! Equation flattening for the flatten phase. //! +//! SPEC_0021 file-size exception: split plan is to move focused equation +//! lowering helpers into owned submodules after BOPTEST parity stabilization. +//! //! This module converts instance equations to flat equations in //! residual form (0 = residual). @@ -22,6 +25,7 @@ use crate::boolean_eval::{ try_resolve_enum_value, }; use crate::errors::FlattenError; +use crate::pipeline::try_eval_const_boolean_with_scope; use crate::static_subscripts::try_constant_integer; use crate::{Context, qualify_expression_imports_with_def_map_ctx}; @@ -328,9 +332,108 @@ fn lookup_parameter_in_scope( } } + // MLS §7.3: `Medium.nXi` in conservation-balance scopes can be recovered from + // already-instantiated medium-sized arrays when package constants were not + // materialized as scalar parameters during constant injection. + if is_type_ref + && cr.parts.len() >= 2 + && let Some(val) = infer_medium_size_constant_from_type_ref(ctx, cr, prefix) + { + return Some(val); + } + + None +} + +/// Infer package size constants referenced through replaceable type aliases. +/// +/// Example: in `ConservationEquation`, `Medium.nXi` can be recovered from the +/// first dimension of `mXi`, `XiOut`, or similar arrays already sized during +/// instantiation even when `medium.nXi` was not injected as a scalar parameter. +fn infer_medium_size_constant_from_type_ref( + ctx: &Context, + cr: &ast::ComponentReference, + prefix: &ast::QualifiedName, +) -> Option { + let constant_name = cr.parts.last()?.ident.text.as_ref(); + let array_candidates: &[&str] = match constant_name { + "nXi" => &[ + "Xi", + "mXi", + "XiOut", + "mXiOut", + "XiOut_internal", + "mbXi_flow", + ], + "nX" => &["X"], + "nC" => &[ + "C", + "COut", + "mC", + "COut_internal", + "C_flow_internal", + "s", + "extraPropertiesNames", + "C_nominal", + ], + "nS" => &["substanceNames"], + _ => return None, + }; + + let first = cr.parts.first()?.ident.text.as_ref(); + if !first.starts_with(char::is_uppercase) { + return None; + } + let lowered = { + let mut chars = first.chars(); + let head = chars.next()?; + head.to_lowercase().collect::() + chars.as_str() + }; + + let mut current = Some(prefix.clone()); + while let Some(scope_qn) = current { + let scope = scope_qn.to_flat_string(); + for candidate in array_candidates { + for key in medium_size_array_lookup_keys(&scope, &lowered, candidate) { + if let Some(dims) = lookup_array_dims_with_unindexed_scope(ctx, &key) + && let Some(size) = medium_size_from_array_dims(constant_name, &dims) + { + return Some(size); + } + } + } + current = get_parent_prefix(&scope_qn); + } + None } +fn medium_size_array_lookup_keys( + scope: &str, + lowered_alias: &str, + array_name: &str, +) -> [String; 2] { + [ + if scope.is_empty() { + array_name.to_string() + } else { + format!("{scope}.{array_name}") + }, + if scope.is_empty() { + format!("{lowered_alias}.{array_name}") + } else { + format!("{scope}.{lowered_alias}.{array_name}") + }, + ] +} + +fn medium_size_from_array_dims(constant_name: &str, dims: &[i64]) -> Option { + match constant_name { + "nXi" | "nX" | "nC" | "nS" => dims.last().copied(), + _ => None, + } +} + /// Try looking up a type reference with the first part lowercased. /// /// In Modelica, replaceable types like `Medium` map to instance components @@ -348,7 +451,7 @@ fn try_lowercase_type_ref( // Add remaining parts as-is parts.extend(cr.parts[1..].iter().map(format_component_ref_part)); let lowered_name = parts.join("."); - ctx.get_integer_param(&lowered_name) + lookup_integer_param_with_unindexed_scope(ctx, &lowered_name) } /// Try looking up an instance-qualified reference with the first part uppercased. @@ -365,16 +468,21 @@ fn try_uppercase_instance_ref( parts.push(uppered); parts.extend(cr.parts[1..].iter().map(format_component_ref_part)); let uppered_name = parts.join("."); - ctx.get_integer_param(&uppered_name) + lookup_integer_param_with_unindexed_scope(ctx, &uppered_name) } -/// Infer medium-style size constants from known array dimensions in scope. +/// Infer structural size constants from known array dimensions in scope. /// /// Examples: /// - `nX` from `X[:]` /// - `nXi` from `Xi[:]` /// - `nC` from `C[:]` /// - `nS` from `substanceNames[:]` +/// - `nr` from `r[:]` +/// - `na` from `a[:]` +/// - `ncr` from `cr[:]` +/// - `nc0` from `c0[:]` +/// - `nout` from `columns[:]` fn infer_size_constant_from_dims( ctx: &Context, constant_name: &str, @@ -385,6 +493,11 @@ fn infer_size_constant_from_dims( "nXi" => &["Xi"], "nC" => &["C"], "nS" => &["substanceNames"], + "nr" => &["r"], + "na" => &["a", "b", "ku"], + "ncr" => &["cr"], + "nc0" => &["c0", "c1"], + "nout" => &["columns"], _ => return None, }; @@ -397,7 +510,7 @@ fn infer_size_constant_from_dims( } else { format!("{scope}.{candidate}") }; - if let Some(dims) = ctx.get_array_dims(&qualified) + if let Some(dims) = lookup_array_dims_with_unindexed_scope(ctx, &qualified) && let Some(&first) = dims.first() { return Some(first); @@ -1466,11 +1579,12 @@ fn try_select_branch_for_mismatched_if( origin: &rumoca_ir_flat::EquationOrigin, def_map: Option<&crate::ResolveDefMap>, ) -> Result { + let scope = prefix.to_flat_string(); for block in cond_blocks { - if let Some(true) = try_eval_boolean_with_ctx_inner(&block.cond, Some(ctx), prefix) { + if let Some(true) = try_eval_const_boolean_with_scope(&block.cond, ctx, &scope) { return flatten_equations_list(ctx, &block.eqs, prefix, span, origin, def_map); } - if let Some(false) = try_eval_boolean_with_ctx_inner(&block.cond, Some(ctx), prefix) { + if let Some(false) = try_eval_const_boolean_with_scope(&block.cond, ctx, &scope) { continue; } // Can't evaluate this condition at all diff --git a/crates/rumoca-phase-flatten/src/function_lowering.rs b/crates/rumoca-phase-flatten/src/function_lowering.rs index 9920c4fe4..bb5506cea 100644 --- a/crates/rumoca-phase-flatten/src/function_lowering.rs +++ b/crates/rumoca-phase-flatten/src/function_lowering.rs @@ -6,23 +6,85 @@ //! - Normalizing structured record-field component references //! - Rewriting FieldAccess expressions on decomposed record params to direct VarRef +// SPEC_0021 file-size exception: record/function lowering is still cohesive +// while Kelvin parity work is in flight; split plan is to move constructor +// projection and field-reference projection helpers into focused modules after +// the active FMU correctness path is stable. + use crate::errors::FlattenError; -use rumoca_core::{ExpressionRewriter, StatementRewriter}; +use rumoca_core::{ExpressionRewriter, ExpressionVisitor, StatementRewriter}; use rumoca_ir_flat as flat; use std::collections::{HashMap, HashSet}; fn record_fields_from_constructor_metadata( functions: &flat::VarNameIndexMap, type_name: &str, + owner_function_name: Option<&str>, ) -> Option> { - functions + let candidates = functions .iter() - .find(|(name, function)| { + .filter(|(name, function)| { function.is_constructor && rumoca_core::qualified_type_name_matches(name.as_str(), type_name) }) + .filter(|(_, function)| !function.inputs.is_empty()) + .collect::>(); + + if candidates.is_empty() { + return None; + } + + if type_name.contains('.') { + return candidates + .into_iter() + .max_by_key(|(name, function)| { + ( + name.as_str() == type_name, + function.inputs.len(), + name.as_str().len(), + ) + }) + .map(|(_, function)| function.inputs.to_vec()); + } + + if let Some(contextual_type_name) = + owner_function_name.and_then(|name| contextual_record_type_name(name, type_name)) + && let Some((_, function)) = candidates + .iter() + .copied() + .filter(|(name, _)| { + rumoca_core::qualified_type_name_matches(name.as_str(), &contextual_type_name) + }) + .max_by_key(|(name, function)| { + ( + name.as_str() == contextual_type_name, + function.inputs.len(), + name.as_str().len(), + ) + }) + { + return Some(function.inputs.to_vec()); + } + + candidates + .into_iter() + .min_by_key(|(name, function)| { + ( + name.as_str() != type_name, + function.inputs.len(), + name.as_str().len(), + ) + }) .map(|(_, function)| function.inputs.to_vec()) - .filter(|fields: &Vec| !fields.is_empty()) +} + +fn contextual_record_type_name(function_name: &str, type_name: &str) -> Option { + let namespace = rumoca_core::ComponentPath::from_flat_path(function_name).parent()?; + Some( + namespace + .join(&rumoca_core::ComponentPath::from_flat_path(type_name)) + .to_flat_string(), + ) } /// Rewrite FieldAccess on decomposed record params to direct VarRef. @@ -121,16 +183,30 @@ impl ExpressionRewriter for RecordParamSizeRewriter<'_> { span, } = first_arg && subscripts.is_empty() - && let Some((_, shape_source)) = self + { + if let Some((_, shape_source)) = self .shape_sources .iter() .find(|(param_name, _)| param_name == name.as_str()) - { - *first_arg = rumoca_core::Expression::VarRef { - name: record_param_reference(shape_source, *span), - subscripts: Vec::new(), - span: *span, - }; + { + *first_arg = rumoca_core::Expression::VarRef { + name: record_param_reference(shape_source, *span), + subscripts: Vec::new(), + span: *span, + }; + } else if let Some(reference) = name.component_ref() + && let Some(record) = reference.parts.first() + && let Some((_, shape_source)) = self + .shape_sources + .iter() + .find(|(param_name, _)| param_name == record.ident.as_str()) + { + *first_arg = rumoca_core::Expression::VarRef { + name: record_param_reference(shape_source, *span), + subscripts: Vec::new(), + span: *span, + }; + } } rumoca_core::Expression::BuiltinCall { @@ -310,16 +386,50 @@ fn record_param_reference(param: &str, span: rumoca_core::Span) -> rumoca_core:: /// 2. Rewrite FieldAccess in the body to VarRef. /// 3. Walk all equations/functions and decompose call-site arguments. pub(crate) fn lower_record_function_params(flat: &mut flat::Model) -> Result<(), FlattenError> { - let record_fields_by_type = flat + let flat_variable_names = flat.variables.keys().cloned().collect::>(); + let constructor_input_names_by_type = flat + .functions + .iter() + .filter(|(_, function)| function.is_constructor) + .map(|(name, function)| { + ( + name.as_str().to_string(), + function + .inputs + .iter() + .map(|input| input.name.clone()) + .collect::>(), + ) + }) + .collect::>(); + let record_fields_by_function_input = flat .functions - .values() - .flat_map(|function| function.inputs.iter()) - .filter(|input| input.type_class == Some(rumoca_core::ClassType::Record)) - .filter_map(|input| { - record_fields_from_constructor_metadata(&flat.functions, &input.type_name) - .map(|fields| (input.type_name.clone(), fields)) + .iter() + .flat_map(|(function_name, function)| { + function + .inputs + .iter() + .enumerate() + .filter_map(|(idx, input)| { + if input.type_class != Some(rumoca_core::ClassType::Record) { + return None; + } + record_fields_from_constructor_metadata( + &flat.functions, + &input.type_name, + Some(function_name.as_str()), + ) + .map(|fields| ((function_name.as_str().to_string(), idx), fields)) + }) + .collect::>() }) .collect::>(); + let record_constructor_fields = flat + .functions + .iter() + .filter(|(_, function)| function.is_constructor) + .map(|(name, function)| (name.as_str().to_string(), function.inputs.to_vec())) + .collect::>(); let mut decomposition_map: HashMap> = HashMap::new(); let mut local_decomposed_params: HashMap> = HashMap::new(); @@ -327,15 +437,25 @@ pub(crate) fn lower_record_function_params(flat: &mut flat::Model) -> Result<(), for (func_name, func) in flat.functions.iter_mut() { let mut decomposed: Vec = Vec::new(); for (idx, input) in func.inputs.iter().enumerate() { - if let Some(fields) = record_fields_by_type.get(&input.type_name) { + if let Some(fields) = + record_fields_by_function_input.get(&(func_name.as_str().to_string(), idx)) + { decomposed.push(DecomposedParam { original_index: idx, param_name: input.name.clone(), + type_name: input.type_name.clone(), fields: fields.clone(), + already_decomposed: false, }); } } if decomposed.is_empty() { + let already_decomposed = + infer_existing_decomposed_params(func, &record_constructor_fields); + if !already_decomposed.is_empty() { + rewrite_existing_decomposed_inputs(func, &already_decomposed); + decomposition_map.insert(func_name.as_str().to_string(), already_decomposed); + } continue; } @@ -375,37 +495,610 @@ pub(crate) fn lower_record_function_params(flat: &mut flat::Model) -> Result<(), } if decomposition_map.is_empty() { + project_overexpanded_function_calls(flat); return Ok(()); } // Rewrite call sites in equations, variable bindings, and function bodies. for eq in &mut flat.equations { - decompose_record_call_args_in_expr(&mut eq.residual, &decomposition_map, None)?; + decompose_record_call_args_in_expr( + &mut eq.residual, + &decomposition_map, + None, + &flat_variable_names, + &constructor_input_names_by_type, + )?; } for eq in &mut flat.initial_equations { - decompose_record_call_args_in_expr(&mut eq.residual, &decomposition_map, None)?; + decompose_record_call_args_in_expr( + &mut eq.residual, + &decomposition_map, + None, + &flat_variable_names, + &constructor_input_names_by_type, + )?; } for var in flat.variables.values_mut() { if let Some(ref mut binding) = var.binding { - decompose_record_call_args_in_expr(binding, &decomposition_map, None)?; + decompose_record_call_args_in_expr( + binding, + &decomposition_map, + None, + &flat_variable_names, + &constructor_input_names_by_type, + )?; } if let Some(ref mut start) = var.start { - decompose_record_call_args_in_expr(start, &decomposition_map, None)?; + decompose_record_call_args_in_expr( + start, + &decomposition_map, + None, + &flat_variable_names, + &constructor_input_names_by_type, + )?; + } + if let Some(ref mut min) = var.min { + decompose_record_call_args_in_expr( + min, + &decomposition_map, + None, + &flat_variable_names, + &constructor_input_names_by_type, + )?; + } + if let Some(ref mut max) = var.max { + decompose_record_call_args_in_expr( + max, + &decomposition_map, + None, + &flat_variable_names, + &constructor_input_names_by_type, + )?; + } + if let Some(ref mut nominal) = var.nominal { + decompose_record_call_args_in_expr( + nominal, + &decomposition_map, + None, + &flat_variable_names, + &constructor_input_names_by_type, + )?; } } for (func_name, func) in flat.functions.iter_mut() { let local_record_params = local_decomposed_params.get(func_name.as_str()); for stmt in &mut func.body { - decompose_record_call_args_in_stmt(stmt, &decomposition_map, local_record_params)?; + decompose_record_call_args_in_stmt( + stmt, + &decomposition_map, + local_record_params, + &flat_variable_names, + &constructor_input_names_by_type, + )?; } } + project_overexpanded_function_calls(flat); Ok(()) } +fn project_overexpanded_function_calls(flat: &mut flat::Model) { + let signatures = flat + .functions + .iter() + .filter(|(_, function)| !function.inputs.is_empty()) + .map(|(name, function)| (name.as_str().to_string(), function.inputs.to_vec())) + .collect::>(); + if signatures.is_empty() { + return; + } + let flat_variable_names = flat + .variables + .keys() + .map(|name| name.as_str().to_string()) + .collect::>(); + let mut rewriter = OverexpandedFunctionCallProjector { + signatures, + flat_variable_names, + }; + for eq in &mut flat.equations { + eq.residual = rewriter.rewrite_expression(&eq.residual); + } + for eq in &mut flat.initial_equations { + eq.residual = rewriter.rewrite_expression(&eq.residual); + } + for var in flat.variables.values_mut() { + if let Some(binding) = &mut var.binding { + *binding = rewriter.rewrite_expression(binding); + } + if let Some(start) = &mut var.start { + *start = rewriter.rewrite_expression(start); + } + if let Some(min) = &mut var.min { + *min = rewriter.rewrite_expression(min); + } + if let Some(max) = &mut var.max { + *max = rewriter.rewrite_expression(max); + } + if let Some(nominal) = &mut var.nominal { + *nominal = rewriter.rewrite_expression(nominal); + } + } + for function in flat.functions.values_mut() { + for statement in &mut function.body { + *statement = rewriter.rewrite_statement(statement); + } + } +} + +struct OverexpandedFunctionCallProjector { + signatures: HashMap>, + flat_variable_names: HashSet, +} + +impl ExpressionRewriter for OverexpandedFunctionCallProjector { + fn rewrite_expression(&mut self, expr: &rumoca_core::Expression) -> rumoca_core::Expression { + let rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor, + span, + } = expr + else { + return self.walk_expression(expr); + }; + let rewritten_args = self.rewrite_expressions(args); + if *is_constructor { + return rumoca_core::Expression::FunctionCall { + name: name.clone(), + args: rewritten_args, + is_constructor: *is_constructor, + span: *span, + }; + } + let args = self + .signatures + .get(name.as_str()) + .and_then(|inputs| { + project_actuals_to_input_slots_with_known_vars( + &rewritten_args, + inputs, + &self.flat_variable_names, + ) + }) + .unwrap_or(rewritten_args); + rumoca_core::Expression::FunctionCall { + name: name.clone(), + args, + is_constructor: *is_constructor, + span: *span, + } + } +} + +impl StatementRewriter for OverexpandedFunctionCallProjector {} + +#[cfg(test)] +fn project_actuals_to_input_slots( + actuals: &[rumoca_core::Expression], + inputs: &[rumoca_core::FunctionParam], +) -> Option> { + project_actuals_to_input_slots_with_known_vars(actuals, inputs, &HashSet::new()) +} + +fn project_actuals_to_input_slots_with_known_vars( + actuals: &[rumoca_core::Expression], + inputs: &[rumoca_core::FunctionParam], + flat_variable_names: &HashSet, +) -> Option> { + if actuals.len() < inputs.len() || actuals.iter().any(is_named_function_arg_marker) { + return None; + } + project_mixed_decomposed_input_slots(actuals, inputs, flat_variable_names) + .or_else(|| project_field_access_actuals(actuals, inputs)) + .or_else(|| { + if actuals.len() == inputs.len() { + return Some(actuals.to_vec()); + } + let mut projected = Vec::new(); + for input in inputs { + let actual = actuals + .iter() + .find(|actual| actual_leaf_matches_input(actual, &input.name))?; + projected.push(actual.clone()); + } + Some(projected) + }) +} + +enum InputSlot<'a> { + Scalar(&'a rumoca_core::FunctionParam), + Decomposed(&'a [rumoca_core::FunctionParam]), +} + +fn project_mixed_decomposed_input_slots( + actuals: &[rumoca_core::Expression], + inputs: &[rumoca_core::FunctionParam], + flat_variable_names: &HashSet, +) -> Option> { + let slots = decomposed_input_slots(inputs); + if !slots + .iter() + .any(|slot| matches!(slot, InputSlot::Decomposed(_))) + { + return None; + } + + let mut projected = Vec::new(); + let mut cursor = 0; + let mut used_prefixes = HashSet::::new(); + for slot in slots { + match slot { + InputSlot::Scalar(input) => { + if let Some(actual) = actuals.get(cursor) { + projected.push(actual.clone()); + cursor += 1; + } else if input.default.is_some() { + continue; + } else { + return None; + } + } + InputSlot::Decomposed(fields) => { + let (prefix, fields, consumed) = project_actuals_to_decomposed_input_slot( + &actuals[cursor..], + fields, + &used_prefixes, + flat_variable_names, + )?; + projected.extend(fields); + cursor += consumed; + used_prefixes.insert(prefix); + } + } + } + + (projected.len() <= inputs.len() + && inputs[projected.len()..] + .iter() + .all(|input| input.default.is_some())) + .then_some(projected) +} + +fn decomposed_input_slots(inputs: &[rumoca_core::FunctionParam]) -> Vec> { + let mut slots = Vec::new(); + let mut idx = 0; + while idx < inputs.len() { + let Some((prefix, _)) = decomposed_input_prefix_and_leaf(&inputs[idx].name) else { + slots.push(InputSlot::Scalar(&inputs[idx])); + idx += 1; + continue; + }; + + let mut end = idx + 1; + while end < inputs.len() + && decomposed_input_prefix_and_leaf(&inputs[end].name) + .is_some_and(|(candidate, _)| candidate == prefix) + { + end += 1; + } + + if end - idx > 1 { + slots.push(InputSlot::Decomposed(&inputs[idx..end])); + } else { + slots.push(InputSlot::Scalar(&inputs[idx])); + } + idx = end; + } + slots +} + +fn decomposed_input_prefix_and_leaf(input_name: &str) -> Option<(&str, &str)> { + let (prefix, leaf) = input_name.rsplit_once('_')?; + (!prefix.is_empty() && !leaf.is_empty()).then_some((prefix, leaf)) +} + +fn project_actuals_to_decomposed_input_slot( + actuals: &[rumoca_core::Expression], + fields: &[rumoca_core::FunctionParam], + used_prefixes: &HashSet, + flat_variable_names: &HashSet, +) -> Option<(String, Vec, usize)> { + if let Some(projected) = + project_scalar_base_plus_tail_decomposed_slot(actuals, fields, flat_variable_names) + { + return Some(projected); + } + + let blocks = actual_field_blocks(actuals, fields); + if let Some(selected) = blocks + .iter() + .find(|block| !used_prefixes.contains(&block.prefix) && block_has_all_fields(block, fields)) + .or_else(|| { + blocks + .iter() + .find(|block| block_has_all_fields(block, fields)) + }) + { + let mut projected = Vec::new(); + for field in fields { + let expected_leaf = decomposed_input_field_leaf(&field.name); + let (_, actual) = selected + .fields + .iter() + .find(|(actual_leaf, _)| function_field_names_match(expected_leaf, actual_leaf))?; + projected.push((*actual).clone()); + } + return Some((selected.prefix.clone(), projected, selected.end)); + } + + project_positional_scalar_decomposed_slot(actuals, fields) +} + +fn project_scalar_base_plus_tail_decomposed_slot( + actuals: &[rumoca_core::Expression], + fields: &[rumoca_core::FunctionParam], + flat_variable_names: &HashSet, +) -> Option<(String, Vec, usize)> { + if fields.len() < 2 || actuals.len() < (fields.len() * 2) - 1 { + return None; + } + + let first_base = scalar_var_field_block_base(actuals.get(..fields.len())?, fields)?; + if !flat_variable_names.contains(&first_base.0) { + return None; + } + + let tail_len = fields.len() - 1; + let tail = actuals.get(fields.len()..fields.len() + tail_len)?; + let mut projected = Vec::with_capacity(fields.len()); + projected.push(first_base.1); + projected.extend(tail.iter().cloned()); + Some((first_base.0, projected, fields.len() + tail_len)) +} + +fn scalar_var_field_block_base( + field_actuals: &[rumoca_core::Expression], + fields: &[rumoca_core::FunctionParam], +) -> Option<(String, rumoca_core::Expression)> { + let mut base_name = None::; + let mut base_expr = None::; + for (actual, field) in field_actuals.iter().zip(fields) { + let expected_leaf = decomposed_input_field_leaf(&field.name); + let (name, base) = scalar_var_field_base(actual, expected_leaf)?; + if base_name + .as_deref() + .is_some_and(|existing| existing != name.as_str()) + { + return None; + } + base_name.get_or_insert(name); + base_expr.get_or_insert(base); + } + Some((base_name?, base_expr?)) +} + +fn scalar_var_field_base( + actual: &rumoca_core::Expression, + expected_leaf: &str, +) -> Option<(String, rumoca_core::Expression)> { + match actual { + rumoca_core::Expression::FieldAccess { + base, + field: actual_field, + .. + } => { + if !function_field_names_match(expected_leaf, actual_field) { + return None; + } + let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = base.as_ref() + else { + return None; + }; + if !subscripts.is_empty() { + return None; + } + Some((name.as_str().to_string(), base.as_ref().clone())) + } + rumoca_core::Expression::VarRef { + name, + subscripts, + span, + } => { + if !subscripts.is_empty() { + return None; + } + let prefix = name + .as_str() + .strip_suffix(&format!(".{expected_leaf}")) + .or_else(|| name.as_str().strip_suffix(&format!("_{expected_leaf}")))?; + if prefix.is_empty() { + return None; + } + Some((prefix.to_string(), flat_var_ref_expression(prefix, *span))) + } + _ => None, + } +} + +fn flat_var_ref_expression(name: &str, span: rumoca_core::Span) -> rumoca_core::Expression { + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference( + name.to_string(), + rumoca_core::ComponentReference::from_flat_segments(name, span, None), + ), + subscripts: vec![], + span, + } +} + +fn project_positional_scalar_decomposed_slot( + actuals: &[rumoca_core::Expression], + fields: &[rumoca_core::FunctionParam], +) -> Option<(String, Vec, usize)> { + if actuals.len() <= fields.len() { + return None; + } + let field_actuals = actuals.get(..fields.len())?; + if field_actuals + .iter() + .any(|actual| actual_leaf_name(actual).is_some()) + { + return None; + } + Some(( + format!("__positional_decomposed_{}", fields[0].name), + field_actuals.to_vec(), + fields.len(), + )) +} + +struct ActualFieldBlock<'a> { + prefix: String, + end: usize, + fields: Vec<(String, &'a rumoca_core::Expression)>, +} + +fn actual_field_blocks<'a>( + actuals: &'a [rumoca_core::Expression], + fields: &[rumoca_core::FunctionParam], +) -> Vec> { + let mut blocks = Vec::new(); + let mut idx = 0; + while idx < actuals.len() { + let Some((prefix, leaf)) = actual_record_field(&actuals[idx], fields) else { + idx += 1; + continue; + }; + let mut block = ActualFieldBlock { + prefix, + end: idx + 1, + fields: vec![(leaf, &actuals[idx])], + }; + idx += 1; + while idx < actuals.len() { + let Some((candidate_prefix, candidate_leaf)) = + actual_record_field(&actuals[idx], fields) + else { + break; + }; + if candidate_prefix != block.prefix { + break; + } + block.fields.push((candidate_leaf, &actuals[idx])); + block.end = idx + 1; + idx += 1; + } + blocks.push(block); + } + blocks +} + +fn actual_record_field( + actual: &rumoca_core::Expression, + fields: &[rumoca_core::FunctionParam], +) -> Option<(String, String)> { + if let Some(field) = flattened_var_ref_field(actual, fields) { + return Some(field); + } + let rumoca_core::Expression::FieldAccess { + base, + field: actual_field, + .. + } = actual + else { + return None; + }; + let expected_leaf = fields.iter().find_map(|field| { + let leaf = decomposed_input_field_leaf(&field.name); + function_field_names_match(leaf, actual_field).then_some(leaf) + })?; + let prefix = + expression_summary_for_projection(base).unwrap_or_else(|| format!("{:?}", base.as_ref())); + Some((prefix, expected_leaf.to_string())) +} + +fn flattened_var_ref_field( + actual: &rumoca_core::Expression, + fields: &[rumoca_core::FunctionParam], +) -> Option<(String, String)> { + let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = actual + else { + return None; + }; + if !subscripts.is_empty() { + return None; + } + for field in fields { + let expected_leaf = decomposed_input_field_leaf(&field.name); + if let Some(prefix) = name.as_str().strip_suffix(&format!(".{expected_leaf}")) + && !prefix.is_empty() + { + return Some((prefix.to_string(), expected_leaf.to_string())); + } + if let Some(prefix) = name.as_str().strip_suffix(&format!("_{expected_leaf}")) + && !prefix.is_empty() + { + return Some((prefix.to_string(), expected_leaf.to_string())); + } + } + None +} + +fn block_has_all_fields( + block: &ActualFieldBlock<'_>, + fields: &[rumoca_core::FunctionParam], +) -> bool { + fields.iter().all(|field| { + let expected_leaf = decomposed_input_field_leaf(&field.name); + block + .fields + .iter() + .any(|(actual_leaf, _)| function_field_names_match(expected_leaf, actual_leaf)) + }) +} + +fn decomposed_input_field_leaf(input_name: &str) -> &str { + input_name + .rsplit_once('_') + .map(|(_, leaf)| leaf) + .unwrap_or(input_name) +} + +fn is_named_function_arg_marker(arg: &rumoca_core::Expression) -> bool { + matches!( + arg, + rumoca_core::Expression::FunctionCall { name, .. } + if name.as_str().starts_with(rumoca_core::NAMED_FUNCTION_ARG_PREFIX) + ) +} + +fn actual_leaf_matches_input(actual: &rumoca_core::Expression, input_name: &str) -> bool { + let Some(leaf) = actual_leaf_name(actual) else { + return false; + }; + function_field_names_match(input_name, leaf) +} + +fn actual_leaf_name(actual: &rumoca_core::Expression) -> Option<&str> { + match actual { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } if subscripts.is_empty() => name.as_str().rsplit(['.', '_']).next(), + rumoca_core::Expression::FieldAccess { field, .. } => Some(field.as_str()), + _ => None, + } +} + struct DecomposedParam { original_index: usize, param_name: String, + type_name: String, fields: Vec, + already_decomposed: bool, } fn record_param_shape_source(param: &DecomposedParam) -> Option { @@ -462,10 +1155,14 @@ fn decompose_record_call_args_in_stmt( stmt: &mut rumoca_core::Statement, map: &HashMap>, local_record_params: Option<&HashSet>, + flat_variable_names: &HashSet, + constructor_input_names_by_type: &HashMap>, ) -> Result<(), FlattenError> { let mut decomposer = RecordCallArgDecomposer { map, local_record_params, + flat_variable_names, + constructor_input_names_by_type, error: None, }; *stmt = decomposer.rewrite_statement(stmt); @@ -479,10 +1176,14 @@ fn decompose_record_call_args_in_expr( expr: &mut rumoca_core::Expression, map: &HashMap>, local_record_params: Option<&HashSet>, + flat_variable_names: &HashSet, + constructor_input_names_by_type: &HashMap>, ) -> Result<(), FlattenError> { let mut decomposer = RecordCallArgDecomposer { map, local_record_params, + flat_variable_names, + constructor_input_names_by_type, error: None, }; *expr = decomposer.rewrite_expression(expr); @@ -495,6 +1196,8 @@ fn decompose_record_call_args_in_expr( struct RecordCallArgDecomposer<'a> { map: &'a HashMap>, local_record_params: Option<&'a HashSet>, + flat_variable_names: &'a HashSet, + constructor_input_names_by_type: &'a HashMap>, error: Option, } @@ -512,6 +1215,8 @@ impl RecordCallArgDecomposer<'_> { &rewritten_args, decomposed, self.local_record_params, + self.flat_variable_names, + self.constructor_input_names_by_type, ) { Ok(args) => args, Err(error) => { @@ -551,42 +1256,652 @@ fn decompose_record_call_args( old_args: &[rumoca_core::Expression], decomposed: &[DecomposedParam], local_record_params: Option<&HashSet>, + flat_variable_names: &HashSet, + constructor_input_names_by_type: &HashMap>, ) -> Result, FlattenError> { + if let Some(args) = + project_already_decomposed_call_args(old_args, decomposed, flat_variable_names) + { + return Ok(args); + } + if decomposed.len() == 1 + && let Some(args) = + project_decomposed_scalar_actuals(old_args, &decomposed[0].fields, flat_variable_names) + { + return Ok(args); + } + if decomposed.len() == 1 + && decomposed[0].fields.len() == 1 + && old_args.len() == 1 + && matches!(&old_args[0], rumoca_core::Expression::FieldAccess { .. }) + && !actual_leaf_matches_input(&old_args[0], &decomposed[0].fields[0].name) + { + return Ok(old_args.to_vec()); + } + let mut args = Vec::new(); let mut old_idx = 0; + let mut consumed = HashSet::new(); + let decomposed_param_names = decomposed + .iter() + .map(|param| param.param_name.as_str()) + .collect::>(); + let has_multiple_decomposed_params = decomposed.len() > 1; for dp in decomposed { while old_idx < dp.original_index && old_idx < old_args.len() { - args.push(old_args[old_idx].clone()); + if !consumed.contains(&old_idx) + && !is_named_arg_for_any(old_args.get(old_idx), &decomposed_param_names) + { + args.push(old_args[old_idx].clone()); + } old_idx += 1; } - if old_idx < old_args.len() { + if let Some((arg_idx, value)) = named_function_arg_value(old_args, &dp.param_name) { + let as_named_fields = + has_multiple_decomposed_params || output_args_have_named_slots(&args); expand_record_arg( function_name, - &old_args[old_idx], + &dp.param_name, + value, + &dp.type_name, &dp.fields, local_record_params, + flat_variable_names, + constructor_input_names_by_type, + as_named_fields, &mut args, )?; - old_idx += 1; - } - } - while old_idx < old_args.len() { - args.push(old_args[old_idx].clone()); + consumed.insert(arg_idx); + } else if old_idx < old_args.len() { + while old_idx < old_args.len() + && is_named_arg_for_any(old_args.get(old_idx), &decomposed_param_names) + { + old_idx += 1; + } + if old_idx >= old_args.len() { + continue; + } + let as_named_fields = output_args_have_named_slots(&args); + expand_record_arg( + function_name, + &dp.param_name, + &old_args[old_idx], + &dp.type_name, + &dp.fields, + local_record_params, + flat_variable_names, + constructor_input_names_by_type, + as_named_fields, + &mut args, + )?; + consumed.insert(old_idx); + old_idx += 1; + } + } + while old_idx < old_args.len() { + if !consumed.contains(&old_idx) { + args.push(old_args[old_idx].clone()); + } old_idx += 1; } Ok(args) } +fn infer_existing_decomposed_params( + func: &rumoca_core::Function, + record_constructor_fields: &[(String, Vec)], +) -> Vec { + if func.is_constructor { + return Vec::new(); + } + let body_uses = flattened_record_field_uses(func); + let mut by_prefix: HashMap = HashMap::new(); + + for (idx, input) in func.inputs.iter().enumerate() { + for (type_name, fields) in record_constructor_fields + .iter() + .filter(|(_, fields)| !fields.is_empty()) + { + let Some((prefix, _)) = split_flattened_record_input(&input.name, fields) else { + continue; + }; + let used_fields = body_uses.get(prefix.as_str()); + let projected_fields = fields + .iter() + .filter(|field| { + func.inputs + .iter() + .any(|input| input.name == format!("{prefix}_{}", field.name)) + || used_fields.is_some_and(|used| used.contains(field.name.as_str())) + }) + .cloned() + .collect::>(); + if projected_fields.is_empty() { + continue; + } + by_prefix + .entry(prefix.clone()) + .and_modify(|param| { + param.original_index = param.original_index.min(idx); + param.fields = merge_record_fields(¶m.fields, &projected_fields); + }) + .or_insert_with(|| DecomposedParam { + original_index: idx, + param_name: prefix, + type_name: type_name.clone(), + fields: projected_fields, + already_decomposed: true, + }); + } + } + + let mut decomposed = by_prefix.into_values().collect::>(); + decomposed.sort_by_key(|param| param.original_index); + decomposed +} + +fn split_flattened_record_input( + input_name: &str, + fields: &[rumoca_core::FunctionParam], +) -> Option<(String, String)> { + fields.iter().find_map(|field| { + let prefix = input_name.strip_suffix(&format!("_{}", field.name))?; + (!prefix.is_empty()).then(|| (prefix.to_string(), field.name.clone())) + }) +} + +fn merge_record_fields( + existing: &[rumoca_core::FunctionParam], + additional: &[rumoca_core::FunctionParam], +) -> Vec { + let mut merged = existing.to_vec(); + for field in additional { + if !merged.iter().any(|candidate| candidate.name == field.name) { + merged.push(field.clone()); + } + } + merged +} + +fn rewrite_existing_decomposed_inputs( + func: &mut rumoca_core::Function, + decomposed: &[DecomposedParam], +) { + let old_inputs = std::mem::take(&mut func.inputs); + let mut idx = 0; + while idx < old_inputs.len() { + let Some(param) = decomposed.iter().find(|param| param.original_index == idx) else { + func.inputs.push(old_inputs[idx].clone()); + idx += 1; + continue; + }; + for field in ¶m.fields { + let name = format!("{}_{}", param.param_name, field.name); + if let Some(existing) = old_inputs.iter().find(|input| input.name == name) { + func.inputs.push(existing.clone()); + } else { + func.inputs.push(flattened_record_field_param(&name, field)); + } + } + idx += old_inputs[idx..] + .iter() + .take_while(|input| { + param + .fields + .iter() + .any(|field| input.name == format!("{}_{}", param.param_name, field.name)) + }) + .count() + .max(1); + } +} + +fn flattened_record_field_param( + name: &str, + field: &rumoca_core::FunctionParam, +) -> rumoca_core::FunctionParam { + let mut input = field.clone(); + input.name = name.to_string(); + input +} + +fn flattened_record_field_uses(func: &rumoca_core::Function) -> HashMap> { + let mut collector = FlattenedRecordFieldUseCollector { + fields: HashMap::new(), + }; + for statement in &func.body { + collect_statement_record_field_uses(statement, &mut collector); + } + collector.fields +} + +struct FlattenedRecordFieldUseCollector { + fields: HashMap>, +} + +impl ExpressionVisitor for FlattenedRecordFieldUseCollector { + fn visit_var_ref( + &mut self, + name: &rumoca_core::Reference, + subscripts: &[rumoca_core::Subscript], + ) { + if subscripts.is_empty() + && let Some((prefix, field)) = name.as_str().split_once('_') + && !prefix.is_empty() + && !field.is_empty() + { + self.fields + .entry(prefix.to_string()) + .or_default() + .insert(field.to_string()); + } + for subscript in subscripts { + self.visit_subscript(subscript); + } + } +} + +fn collect_statement_record_field_uses( + statement: &rumoca_core::Statement, + collector: &mut FlattenedRecordFieldUseCollector, +) { + match statement { + rumoca_core::Statement::Assignment { value, .. } + | rumoca_core::Statement::Reinit { value, .. } => collector.visit_expression(value), + rumoca_core::Statement::FunctionCall { args, .. } => { + for arg in args { + collector.visit_expression(arg); + } + } + rumoca_core::Statement::For { equations, .. } => { + for statement in equations { + collect_statement_record_field_uses(statement, collector); + } + } + rumoca_core::Statement::While { block, .. } => { + collector.visit_expression(&block.cond); + for statement in &block.stmts { + collect_statement_record_field_uses(statement, collector); + } + } + rumoca_core::Statement::If { + cond_blocks, + else_block, + .. + } => { + for block in cond_blocks { + collector.visit_expression(&block.cond); + for statement in &block.stmts { + collect_statement_record_field_uses(statement, collector); + } + } + if let Some(else_block) = else_block { + for statement in else_block { + collect_statement_record_field_uses(statement, collector); + } + } + } + rumoca_core::Statement::When { blocks, .. } => { + for block in blocks { + collector.visit_expression(&block.cond); + for statement in &block.stmts { + collect_statement_record_field_uses(statement, collector); + } + } + } + rumoca_core::Statement::Assert { + condition, + message, + level, + .. + } => { + collector.visit_expression(condition); + collector.visit_expression(message); + if let Some(level) = level { + collector.visit_expression(level); + } + } + rumoca_core::Statement::Empty { .. } + | rumoca_core::Statement::Return { .. } + | rumoca_core::Statement::Break { .. } => {} + } +} + +fn project_already_decomposed_call_args( + old_args: &[rumoca_core::Expression], + decomposed: &[DecomposedParam], + flat_variable_names: &HashSet, +) -> Option> { + if !decomposed.iter().any(|param| param.already_decomposed) { + return None; + } + let mut args = old_args.to_vec(); + let mut changed = false; + for (param_idx, dp) in decomposed.iter().enumerate() { + if !dp.already_decomposed { + continue; + } + let end = decomposed + .iter() + .skip(param_idx + 1) + .map(|param| param.original_index) + .find(|next_index| *next_index > dp.original_index) + .unwrap_or(args.len()); + let available_end = end.min(args.len()); + let slice = args.get(dp.original_index..available_end)?; + if let Some(projected) = + project_decomposed_scalar_actuals(slice, &dp.fields, flat_variable_names) + { + args.splice(dp.original_index..available_end, projected); + changed = true; + } + } + changed.then_some(args) +} + +fn project_decomposed_scalar_actuals( + actuals: &[rumoca_core::Expression], + fields: &[rumoca_core::FunctionParam], + flat_variable_names: &HashSet, +) -> Option> { + if let Some(projected) = project_flattened_var_ref_actuals(actuals, fields) { + return Some(projected); + } + if let Some(projected) = project_field_access_actuals(actuals, fields) { + return Some(projected); + } + if actuals.is_empty() || actuals.len() > fields.len() { + return None; + } + let base = matching_scalar_constructor_base(actuals, fields)?; + let rumoca_core::Expression::FunctionCall { + args: constructor_args, + .. + } = base + else { + return None; + }; + let positional = constructor_args.iter().collect::>(); + constructor_positional_args_project_expected_fields(&positional, fields, flat_variable_names) + .or_else(|| project_fields_from_common_base(base, fields)) +} + +fn project_field_access_actuals( + actuals: &[rumoca_core::Expression], + fields: &[rumoca_core::FunctionParam], +) -> Option> { + if fields.len() < 2 { + return None; + } + let mut bases = Vec::<(String, rumoca_core::Expression)>::new(); + for actual in actuals { + let rumoca_core::Expression::FieldAccess { .. } = actual else { + continue; + }; + for base_expr in projection_base_candidates(actual, fields) { + let base_key = expression_summary_for_projection(base_expr)?; + if !bases.iter().any(|(existing, _)| existing == &base_key) { + bases.push((base_key, base_expr.clone())); + } + } + } + bases.sort_by_key(|(base, _)| base.matches('.').count()); + + for (base, base_expr) in bases { + let mut projected = Vec::new(); + for field in fields { + let actual = actuals.iter().find(|actual| { + let rumoca_core::Expression::FieldAccess { + base: actual_base, + field: actual_field, + .. + } = actual + else { + return false; + }; + expression_summary_for_projection(actual_base).as_deref() == Some(base.as_str()) + && function_field_names_match(&field.name, actual_field) + }); + if let Some(actual) = actual { + projected.push(actual.clone()); + } else if let Some(field_name) = projected_field_name(&field.name) { + let Some(span) = base_expr + .span() + .or_else(|| (!field.span.is_dummy()).then_some(field.span)) + else { + projected.clear(); + break; + }; + projected.push(rumoca_core::Expression::FieldAccess { + base: Box::new(base_expr.clone()), + field: field_name.to_string(), + span, + }); + } else { + projected.clear(); + break; + } + } + if projected.len() == fields.len() { + return Some(projected); + } + } + None +} + +fn projection_base_candidates<'a>( + actual: &'a rumoca_core::Expression, + fields: &[rumoca_core::FunctionParam], +) -> Vec<&'a rumoca_core::Expression> { + let mut candidates = Vec::new(); + let rumoca_core::Expression::FieldAccess { base, .. } = actual else { + return candidates; + }; + candidates.push(base.as_ref()); + + let mut current = base.as_ref(); + while let rumoca_core::Expression::FieldAccess { + base: parent, + field, + .. + } = current + { + if fields + .iter() + .any(|expected| function_field_names_match(&expected.name, field)) + { + candidates.push(parent.as_ref()); + } + current = parent.as_ref(); + } + candidates +} + +fn projected_field_name(input_or_field_name: &str) -> Option<&str> { + input_or_field_name + .rsplit_once('_') + .map(|(_, suffix)| suffix) + .or(Some(input_or_field_name)) +} + +fn expression_summary_for_projection(expr: &rumoca_core::Expression) -> Option { + match expr { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } => Some(format!( + "{}{}", + name.as_str(), + render_projection_subscripts(subscripts)? + )), + rumoca_core::Expression::Index { + base, subscripts, .. + } => Some(format!( + "{}{}", + expression_summary_for_projection(base)?, + render_projection_subscripts(subscripts)? + )), + rumoca_core::Expression::FieldAccess { base, field, .. } => Some(format!( + "{}.{}", + expression_summary_for_projection(base)?, + field + )), + _ => None, + } +} + +fn render_projection_subscripts(subscripts: &[rumoca_core::Subscript]) -> Option { + let mut rendered = String::new(); + for subscript in subscripts { + let rumoca_core::Subscript::Index { value, .. } = subscript else { + return None; + }; + rendered.push('['); + rendered.push_str(&value.to_string()); + rendered.push(']'); + } + Some(rendered) +} + +fn function_field_names_match(expected: &str, actual: &str) -> bool { + expected == actual + || expected + .rsplit_once('_') + .is_some_and(|(_, suffix)| suffix == actual) +} + +fn project_flattened_var_ref_actuals( + actuals: &[rumoca_core::Expression], + fields: &[rumoca_core::FunctionParam], +) -> Option> { + let mut prefix: Option = None; + let mut projected = Vec::new(); + for field in fields { + let (actual_prefix, actual) = actuals.iter().find_map(|actual| { + let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = actual + else { + return None; + }; + if !subscripts.is_empty() { + return None; + } + let (candidate_prefix, candidate_field) = + split_flattened_record_input(name.as_str(), fields)?; + function_field_names_match(&field.name, &candidate_field) + .then_some((candidate_prefix, actual)) + })?; + if let Some(existing_prefix) = &prefix { + if existing_prefix != &actual_prefix { + return None; + } + } else { + prefix = Some(actual_prefix); + } + projected.push(actual.clone()); + } + Some(projected) +} + +fn matching_scalar_constructor_base<'a>( + actuals: &'a [rumoca_core::Expression], + fields: &[rumoca_core::FunctionParam], +) -> Option<&'a rumoca_core::Expression> { + let mut base = None; + for (actual, field) in actuals.iter().zip(fields.iter()) { + let rumoca_core::Expression::FieldAccess { + base: current_base, + field: current_field, + .. + } = actual + else { + return None; + }; + if current_field != &field.name { + return None; + } + if let Some(existing_base) = base { + if existing_base != current_base.as_ref() { + return None; + } + } else { + base = Some(current_base.as_ref()); + } + } + base +} + +fn project_fields_from_common_base( + base: &rumoca_core::Expression, + fields: &[rumoca_core::FunctionParam], +) -> Option> { + let span = base.span()?; + Some( + fields + .iter() + .map(|field| rumoca_core::Expression::FieldAccess { + base: Box::new(base.clone()), + field: field.name.clone(), + span, + }) + .collect(), + ) +} + +fn is_named_arg_for_any( + arg: Option<&rumoca_core::Expression>, + param_names: &HashSet<&str>, +) -> bool { + named_function_arg_name(arg).is_some_and(|name| param_names.contains(name)) +} + +fn output_args_have_named_slots(args: &[rumoca_core::Expression]) -> bool { + args.iter() + .any(|arg| named_function_arg_name(Some(arg)).is_some()) +} + +fn named_function_arg_name(arg: Option<&rumoca_core::Expression>) -> Option<&str> { + let rumoca_core::Expression::FunctionCall { name, .. } = arg? else { + return None; + }; + name.as_str().strip_prefix("__rumoca_named_arg__.") +} + +fn named_function_arg_value<'a>( + args: &'a [rumoca_core::Expression], + param_name: &str, +) -> Option<(usize, &'a rumoca_core::Expression)> { + args.iter().enumerate().find_map(|(idx, arg)| { + let rumoca_core::Expression::FunctionCall { + name, + args: named_args, + .. + } = arg + else { + return None; + }; + (name.as_str().strip_prefix("__rumoca_named_arg__.") == Some(param_name)) + .then(|| named_args.first().map(|value| (idx, value))) + .flatten() + }) +} + /// Expand a record argument into scalar field arguments. fn expand_record_arg( function_name: &str, + param_name: &str, arg: &rumoca_core::Expression, + expected_type_name: &str, fields: &[rumoca_core::FunctionParam], local_record_params: Option<&HashSet>, + flat_variable_names: &HashSet, + constructor_input_names_by_type: &HashMap>, + as_named_fields: bool, out: &mut Vec, ) -> Result<(), FlattenError> { // Constructor call Complex(re, im) → extract positional/named args if let rumoca_core::Expression::FunctionCall { + name, args: ctor_args, is_constructor: true, span, @@ -600,15 +1915,117 @@ fn expand_record_arg( if name.as_str().starts_with("__rumoca_named_arg__.")) }) .collect(); + let constructor_matches_expected = + rumoca_core::qualified_type_name_matches(name.as_str(), expected_type_name); + if constructor_matches_expected + && let Some(record_value) = + record_value_constructor_proxy(ctor_args, param_name, fields) + { + if let Some(values) = record_value_constructor_proxy_project_fields( + record_value, + fields, + constructor_input_names_by_type, + ) { + for (field, value) in fields.iter().zip(values) { + push_expanded_record_field_arg(out, param_name, field, as_named_fields, value); + } + return Ok(()); + } + for field in fields { + let source_span = record_field_access_source_span(record_value, field)?; + push_expanded_record_field_arg( + out, + param_name, + field, + as_named_fields, + rumoca_core::Expression::FieldAccess { + base: Box::new(record_value.clone()), + field: field.name.clone(), + span: source_span, + }, + ); + } + return Ok(()); + } + if constructor_matches_expected + && positional.len() == 1 + && let Some(values) = record_value_constructor_proxy_project_fields( + positional[0], + fields, + constructor_input_names_by_type, + ) + { + for (field, value) in fields.iter().zip(values) { + push_expanded_record_field_arg(out, param_name, field, as_named_fields, value); + } + return Ok(()); + } + if !constructor_matches_expected + && constructor_positional_args_match_fields(&positional, fields) + { + for (field, value) in fields.iter().zip(positional.iter()) { + push_expanded_record_field_arg( + out, + param_name, + field, + as_named_fields, + (*value).clone(), + ); + } + return Ok(()); + } + if !constructor_matches_expected + && let Some(values) = constructor_positional_args_project_expected_fields( + &positional, + fields, + flat_variable_names, + ) + { + for (field, value) in fields.iter().zip(values) { + push_expanded_record_field_arg(out, param_name, field, as_named_fields, value); + } + return Ok(()); + } for (i, field) in fields.iter().enumerate() { let named = named_constructor_arg(ctor_args, field.name.as_str()); if let Some(val) = named { - out.push(val.clone()); - } else if i < positional.len() { - out.push(positional[i].clone()); + push_expanded_record_field_arg( + out, + param_name, + field, + as_named_fields, + val.clone(), + ); + } else if constructor_matches_expected && i < positional.len() { + push_expanded_record_field_arg( + out, + param_name, + field, + as_named_fields, + positional[i].clone(), + ); } else if let Some(default) = &field.default { - out.push(default.clone()); + push_expanded_record_field_arg( + out, + param_name, + field, + as_named_fields, + default.clone(), + ); + } else if !constructor_matches_expected { + let source_span = record_field_access_source_span(arg, field)?; + push_expanded_record_field_arg( + out, + param_name, + field, + as_named_fields, + rumoca_core::Expression::FieldAccess { + base: Box::new(arg.clone()), + field: field.name.clone(), + span: source_span, + }, + ); } else { return Err(missing_record_constructor_field_error( function_name, @@ -624,21 +2041,29 @@ fn expand_record_arg( if let rumoca_core::Expression::VarRef { name, span, .. } = arg { if local_record_params.is_some_and(|params| params.contains(name.as_str())) { for field in fields { - out.push(record_param_field_var_ref( - name.as_str(), - field.name.as_str(), - *span, - )); + push_expanded_record_field_arg( + out, + param_name, + field, + as_named_fields, + record_param_field_var_ref(name.as_str(), field.name.as_str(), *span), + ); } return Ok(()); } for field in fields { - out.push(rumoca_core::Expression::VarRef { - name: record_field_reference(name, field.name.as_str(), *span), - subscripts: vec![], - span: *span, - }); + push_expanded_record_field_arg( + out, + param_name, + field, + as_named_fields, + rumoca_core::Expression::VarRef { + name: record_field_reference(name, field.name.as_str(), *span), + subscripts: vec![], + span: *span, + }, + ); } return Ok(()); } @@ -646,35 +2071,366 @@ fn expand_record_arg( // General expression → emit FieldAccess for field in fields { let source_span = record_field_access_source_span(arg, field)?; - out.push(rumoca_core::Expression::FieldAccess { - base: Box::new(arg.clone()), - field: field.name.clone(), - span: source_span, - }); + push_expanded_record_field_arg( + out, + param_name, + field, + as_named_fields, + rumoca_core::Expression::FieldAccess { + base: Box::new(arg.clone()), + field: field.name.clone(), + span: source_span, + }, + ); } Ok(()) } -fn missing_record_constructor_field_error( - function_name: &str, - field: &str, - span: rumoca_core::Span, -) -> FlattenError { - if span.is_dummy() { - return FlattenError::missing_source_context(format!( - "record constructor argument `{field}` for `{function_name}` has no source span" - )); - } - FlattenError::invalid_function_call_args( - function_name, - format!("missing record constructor field `{field}` after record parameter lowering"), - span, - ) +fn constructor_positional_args_match_fields( + positional: &[&rumoca_core::Expression], + fields: &[rumoca_core::FunctionParam], +) -> bool { + positional.len() == fields.len() + && !fields.is_empty() + && positional + .iter() + .zip(fields.iter()) + .all(|(arg, field)| expression_leaf_name(arg) == Some(field.name.as_str())) } -fn record_field_access_source_span( - arg: &rumoca_core::Expression, - field: &rumoca_core::FunctionParam, +fn constructor_positional_args_project_expected_fields( + positional: &[&rumoca_core::Expression], + fields: &[rumoca_core::FunctionParam], + flat_variable_names: &HashSet, +) -> Option> { + if positional.is_empty() || fields.is_empty() { + return None; + } + if let Some(values) = constructor_positional_args_project_from_matching_ref( + positional, + fields, + flat_variable_names, + ) { + return Some(values); + } + let prefix = common_record_field_prefix(positional) + .or_else(|| matching_record_field_prefix(positional, fields, flat_variable_names))?; + let span = positional.iter().find_map(|expr| expr.span())?; + fields + .iter() + .map(|field| { + let name = format!("{prefix}.{}", field.name); + let var_name = rumoca_core::VarName::new(&name); + flat_variable_names + .contains(&var_name) + .then(|| rumoca_core::Expression::VarRef { + name: record_field_projected_reference(&prefix, field.name.as_str(), span), + subscripts: vec![], + span, + }) + }) + .collect() +} + +fn constructor_positional_args_project_from_matching_ref( + positional: &[&rumoca_core::Expression], + fields: &[rumoca_core::FunctionParam], + flat_variable_names: &HashSet, +) -> Option> { + positional + .iter() + .find_map(|arg| projected_values_from_matching_ref(arg, fields, flat_variable_names)) +} + +fn projected_values_from_matching_ref( + arg: &rumoca_core::Expression, + fields: &[rumoca_core::FunctionParam], + flat_variable_names: &HashSet, +) -> Option> { + let rumoca_core::Expression::VarRef { name, span, .. } = arg else { + return None; + }; + let component_ref = name.component_ref()?; + let leaf = component_ref.parts.last()?.ident.as_str(); + if !fields.iter().any(|field| field.name.as_str() == leaf) { + return None; + } + let prefix_len = component_ref.parts.len().checked_sub(1)?; + let prefix_parts = component_ref.parts[..prefix_len].to_vec(); + if prefix_parts.is_empty() { + return None; + } + let projected_names = fields + .iter() + .map(|field| { + let mut parts = prefix_parts.clone(); + parts.push(rumoca_core::ComponentRefPart { + ident: field.name.clone(), + span: field.span, + subs: Vec::new(), + }); + let projected_ref = rumoca_core::ComponentReference { + local: component_ref.local, + span: component_ref.span, + parts, + def_id: None, + }; + projected_ref.to_var_name() + }) + .collect::>(); + if !projected_names + .iter() + .all(|name| flat_variable_names.contains(name)) + { + return None; + } + Some( + projected_names + .iter() + .map(|projected_name| rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference( + projected_name.as_str().to_string(), + rumoca_core::ComponentReference::from_flat_segments( + projected_name.as_str(), + *span, + None, + ), + ), + subscripts: vec![], + span: *span, + }) + .collect(), + ) +} + +fn matching_record_field_prefix( + positional: &[&rumoca_core::Expression], + fields: &[rumoca_core::FunctionParam], + _flat_variable_names: &HashSet, +) -> Option { + positional.iter().find_map(|arg| { + let prefix = expression_field_prefix(arg)?; + let leaf = expression_leaf_name(arg)?; + fields + .iter() + .any(|field| field.name.as_str() == leaf) + .then_some(prefix) + }) +} + +fn common_record_field_prefix(positional: &[&rumoca_core::Expression]) -> Option { + let mut prefixes = positional.iter().map(|arg| expression_field_prefix(arg)); + let first = prefixes.next()??; + prefixes + .all(|prefix| prefix.as_deref() == Some(first.as_str())) + .then_some(first) +} + +fn expression_field_prefix(expr: &rumoca_core::Expression) -> Option { + let rumoca_core::Expression::VarRef { name, .. } = expr else { + return None; + }; + if let Some(component_ref) = name.component_ref() { + let field_count = component_ref.parts.len().checked_sub(1)?; + let prefix_parts = component_ref.parts[..field_count] + .iter() + .map(render_component_ref_part) + .collect::>>()?; + return (!prefix_parts.is_empty()).then(|| prefix_parts.join(".")); + } + flat_var_name_prefix_text(name.as_str()).map(str::to_string) +} + +fn render_component_ref_part(part: &rumoca_core::ComponentRefPart) -> Option { + let mut rendered = part.ident.clone(); + for subscript in &part.subs { + match subscript { + rumoca_core::Subscript::Index { value, .. } => { + rendered.push('['); + rendered.push_str(&value.to_string()); + rendered.push(']'); + } + rumoca_core::Subscript::Colon { .. } | rumoca_core::Subscript::Expr { .. } => { + return None; + } + } + } + Some(rendered) +} + +fn record_field_projected_reference( + prefix: &str, + field: &str, + span: rumoca_core::Span, +) -> rumoca_core::Reference { + let name = format!("{prefix}.{field}"); + let mut parts = rumoca_core::ComponentPath::from_flat_path(prefix) + .parts() + .iter() + .map(|part| rumoca_core::ComponentRefPart { + ident: (*part).to_string(), + span, + subs: Vec::new(), + }) + .collect::>(); + parts.push(rumoca_core::ComponentRefPart { + ident: field.to_string(), + span, + subs: Vec::new(), + }); + rumoca_core::Reference::with_component_reference( + name, + rumoca_core::ComponentReference { + local: false, + span, + parts, + def_id: None, + }, + ) +} + +fn push_expanded_record_field_arg( + out: &mut Vec, + param_name: &str, + field: &rumoca_core::FunctionParam, + as_named_field: bool, + value: rumoca_core::Expression, +) { + if as_named_field { + let span = value.span().unwrap_or(field.span); + out.push(rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new(format!( + "{}{param_name}_{}", + rumoca_core::NAMED_FUNCTION_ARG_PREFIX, + field.name + )), + args: vec![value], + is_constructor: true, + span, + }); + } else { + out.push(value); + } +} + +fn record_value_constructor_proxy<'a>( + ctor_args: &'a [rumoca_core::Expression], + param_name: &str, + fields: &[rumoca_core::FunctionParam], +) -> Option<&'a rumoca_core::Expression> { + if fields.len() <= 1 || ctor_args.iter().any(named_function_arg_name_expr) { + return None; + } + let [record_value] = ctor_args else { + return None; + }; + (expression_leaf_name(record_value) == Some(param_name)).then_some(record_value) +} + +fn record_value_constructor_proxy_project_fields( + record_value: &rumoca_core::Expression, + fields: &[rumoca_core::FunctionParam], + constructor_input_names_by_type: &HashMap>, +) -> Option> { + let rumoca_core::Expression::FieldAccess { + base, + field: record_field, + span, + } = record_value + else { + return None; + }; + let rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor: true, + .. + } = base.as_ref() + else { + return None; + }; + let constructor_fields = constructor_input_names_by_type.get(name.as_str())?; + fields + .iter() + .map(|field| { + let projected = format!("{record_field}_{}", field.name); + constructor_projected_actual(args, constructor_fields, &projected).or_else(|| { + constructor_fields.contains(&projected).then(|| { + rumoca_core::Expression::FieldAccess { + base: base.clone(), + field: projected, + span: *span, + } + }) + }) + }) + .collect() +} + +fn constructor_projected_actual( + args: &[rumoca_core::Expression], + constructor_fields: &[String], + projected: &str, +) -> Option { + if let Some(value) = named_constructor_arg(args, projected) { + return Some(value.clone()); + } + let index = constructor_fields + .iter() + .position(|field| field.as_str() == projected)?; + args.iter() + .filter(|arg| named_function_arg_name(Some(arg)).is_none()) + .nth(index) + .cloned() +} + +fn named_function_arg_name_expr(arg: &rumoca_core::Expression) -> bool { + named_function_arg_name(Some(arg)).is_some() +} + +fn expression_leaf_name(expr: &rumoca_core::Expression) -> Option<&str> { + match expr { + rumoca_core::Expression::VarRef { name, .. } => name + .component_ref() + .and_then(|component_ref| component_ref.parts.last()) + .map(|part| part.ident.as_str()) + .or_else(|| flat_var_name_leaf_text(name.as_str())), + rumoca_core::Expression::Index { base, .. } => expression_leaf_name(base), + rumoca_core::Expression::FieldAccess { field, .. } => Some(field.as_str()), + _ => None, + } +} + +fn flat_var_name_prefix_text(name: &str) -> Option<&str> { + let dot = name.rfind('.')?; + Some(&name[..dot]) +} + +fn flat_var_name_leaf_text(name: &str) -> Option<&str> { + let dot = name.rfind('.')?; + Some(&name[dot + 1..]) +} + +fn missing_record_constructor_field_error( + function_name: &str, + field: &str, + span: rumoca_core::Span, +) -> FlattenError { + if span.is_dummy() { + return FlattenError::missing_source_context(format!( + "record constructor argument `{field}` for `{function_name}` has no source span" + )); + } + FlattenError::invalid_function_call_args( + function_name, + format!("missing record constructor field `{field}` after record parameter lowering"), + span, + ) +} + +fn record_field_access_source_span( + arg: &rumoca_core::Expression, + field: &rumoca_core::FunctionParam, ) -> Result { arg.span() .or_else(|| (!field.span.is_dummy()).then_some(field.span)) @@ -724,6 +2480,21 @@ mod tests { } } + fn int_lit(value: i64) -> rumoca_core::Expression { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + span: Span::DUMMY, + } + } + + fn field_access(base: rumoca_core::Expression, field: &str) -> rumoca_core::Expression { + rumoca_core::Expression::FieldAccess { + base: Box::new(base), + field: field.to_string(), + span: test_span(), + } + } + fn assignment_to(name: &str, value: rumoca_core::Expression) -> rumoca_core::Statement { rumoca_core::Statement::Assignment { comp: rumoca_core::ComponentReference { @@ -742,17 +2513,18 @@ mod tests { } fn component_ref_expr(parts: &[&str]) -> rumoca_core::Expression { + let span = test_span(); rumoca_core::Expression::VarRef { name: rumoca_core::Reference::with_component_reference( parts.join("."), rumoca_core::ComponentReference { local: false, - span: Span::DUMMY, + span, parts: parts .iter() .map(|part| rumoca_core::ComponentRefPart { ident: (*part).to_string(), - span: Span::DUMMY, + span, subs: Vec::new(), }) .collect(), @@ -760,10 +2532,50 @@ mod tests { }, ), subscripts: vec![], + span, + } + } + + fn named_arg(name: &str, value: rumoca_core::Expression) -> rumoca_core::Expression { + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new(format!("__rumoca_named_arg__.{name}")), + args: vec![value], + is_constructor: true, span: Span::DUMMY, } } + fn named_arg_var_ref<'a>( + arg: &'a rumoca_core::Expression, + slot: &str, + ) -> Option<&'a rumoca_core::Reference> { + let rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor: true, + .. + } = arg + else { + return None; + }; + if name.as_str().strip_prefix("__rumoca_named_arg__.") != Some(slot) { + return None; + } + let rumoca_core::Expression::VarRef { name, .. } = args.first()? else { + return None; + }; + Some(name) + } + + fn record_constructor_named(name: &str, fields: &[&str]) -> rumoca_core::Function { + let mut constructor = rumoca_core::Function::new(name, Span::DUMMY); + constructor.is_constructor = true; + for field in fields { + constructor.add_input(rumoca_core::FunctionParam::new(*field, "Real", test_span())); + } + constructor + } + fn record_constructor() -> rumoca_core::Function { let mut constructor = rumoca_core::Function::new("Pkg.Record", Span::DUMMY); constructor.is_constructor = true; @@ -789,18 +2601,1205 @@ mod tests { span: Span::DUMMY, }, )); - function + function + } + + #[test] + fn record_param_lowering_uses_constructor_signature_metadata() { + let mut flat = flat::Model::new(); + flat.add_function(record_constructor()); + flat.add_function(function_with_record_input()); + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.f"), + args: vec![var_ref("rec")], + is_constructor: false, + span: Span::DUMMY, + }, + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + + let function = flat + .functions + .get(&VarName::new("Pkg.f")) + .expect("function remains"); + let input_names = function + .inputs + .iter() + .map(|input| input.name.as_str()) + .collect::>(); + assert_eq!(input_names, vec!["r_a", "r_b"]); + assert_eq!(function.inputs[0].dims, Vec::::new()); + assert_eq!(function.inputs[1].dims, vec![3]); + let rumoca_core::Statement::Assignment { value, .. } = &function.body[0] else { + panic!("expected assignment"); + }; + assert!(matches!( + value, + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "r_a" + )); + let rumoca_core::Expression::FunctionCall { args, .. } = &flat.equations[0].residual else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 2); + assert!(matches!( + &args[0], + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "rec.a" + )); + assert!(matches!( + &args[1], + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "rec.b" + )); + } + + #[test] + fn record_param_lowering_expands_record_actual_for_already_decomposed_signature() { + let mut flat = flat::Model::new(); + flat.add_function(record_constructor()); + + let mut function = rumoca_core::Function::new("Pkg.f", Span::DUMMY); + function.add_input(rumoca_core::FunctionParam::new("r_a", "Real", test_span())); + function.add_input( + rumoca_core::FunctionParam::new("r_b", "Real", test_span()).with_dims(vec![3]), + ); + function.add_output(rumoca_core::FunctionParam::new("y", "Real", test_span())); + function.body.push(assignment_to("y", var_ref("r_a"))); + flat.add_function(function); + + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.f"), + args: vec![var_ref("rec")], + is_constructor: false, + span: Span::DUMMY, + }, + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + + let rumoca_core::Expression::FunctionCall { args, .. } = &flat.equations[0].residual else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 2); + assert!(matches!( + &args[0], + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "rec.a" + )); + assert!(matches!( + &args[1], + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "rec.b" + )); + } + + #[test] + fn record_param_lowering_projects_overcomplete_decomposed_actuals() { + let mut flat = flat::Model::new(); + flat.add_function(record_constructor()); + + let mut function = rumoca_core::Function::new("Pkg.f", Span::DUMMY); + function.add_input(rumoca_core::FunctionParam::new("r_a", "Real", test_span())); + function.add_input( + rumoca_core::FunctionParam::new("r_b", "Real", test_span()).with_dims(vec![3]), + ); + function.add_output(rumoca_core::FunctionParam::new("y", "Real", test_span())); + function.body.push(assignment_to("y", var_ref("r_a"))); + flat.add_function(function); + + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.f"), + args: vec![ + var_ref("rec_phase"), + var_ref("rec_a"), + var_ref("rec_b"), + var_ref("rec_extra"), + ], + is_constructor: false, + span: Span::DUMMY, + }, + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + + let rumoca_core::Expression::FunctionCall { args, .. } = &flat.equations[0].residual else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 2); + assert!(matches!( + &args[0], + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "rec_a" + )); + assert!(matches!( + &args[1], + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "rec_b" + )); + } + + #[test] + fn positional_function_actuals_are_not_reordered_by_shared_name_suffixes() { + let inputs = vec![ + rumoca_core::FunctionParam::new("foot_position_deck", "Real", test_span()) + .with_dims(vec![3, 4]), + rumoca_core::FunctionParam::new("point_deck", "Real", test_span()).with_dims(vec![2]), + rumoca_core::FunctionParam::new("loaded", "Boolean", test_span()).with_dims(vec![4]), + ]; + let point = rumoca_core::Expression::Array { + elements: vec![ + var_ref("cg_position_deck[1]"), + var_ref("cg_position_deck[2]"), + ], + is_matrix: false, + span: test_span(), + }; + let actuals = vec![ + var_ref("leg_position_deck"), + point.clone(), + var_ref("leg_loaded"), + ]; + + assert_eq!( + project_actuals_to_input_slots(&actuals, &inputs), + Some(vec![actuals[0].clone(), point.clone(), actuals[2].clone()]) + ); + let shape_mismatched = vec![point, actuals[0].clone(), actuals[2].clone()]; + assert_eq!( + project_actuals_to_input_slots(&shape_mismatched, &inputs), + Some(shape_mismatched) + ); + } + + #[test] + fn record_param_lowering_projects_mixed_scalar_and_decomposed_slots() { + let mut flat = flat::Model::new(); + + let mut function = rumoca_core::Function::new("Pkg.setSmoothState", Span::DUMMY); + for input in [ + "x", + "state_a_p", + "state_a_T", + "state_a_X", + "state_b_p", + "state_b_T", + "state_b_X", + "x_small", + ] { + function.add_input(rumoca_core::FunctionParam::new(input, "Real", test_span())); + } + function.add_output(rumoca_core::FunctionParam::new("y", "Real", test_span())); + function.body.push(assignment_to("y", var_ref("x"))); + let inputs = function.inputs.clone(); + flat.add_function(function); + + let call_args = vec![ + var_ref("m_flow_ext"), + var_ref("state1.p"), + var_ref("state1.T"), + var_ref("state1.X"), + var_ref("state1.T"), + var_ref("state1.X"), + var_ref("state1.p"), + var_ref("state1.T"), + var_ref("state1.X"), + var_ref("state1.X"), + var_ref("state2.p"), + var_ref("state2.T"), + var_ref("state2.X"), + var_ref("state2.T"), + var_ref("state2.X"), + var_ref("x_small_actual"), + ]; + assert_eq!( + project_actuals_to_input_slots(&call_args, &inputs) + .expect("overexpanded actuals should project") + .len(), + 8 + ); + + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.setSmoothState"), + args: call_args, + is_constructor: false, + span: Span::DUMMY, + }, + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + + let rumoca_core::Expression::FunctionCall { args, .. } = &flat.equations[0].residual else { + panic!("expected function call"); + }; + let arg_names = args + .iter() + .map(|arg| match arg { + rumoca_core::Expression::VarRef { name, .. } => name.as_str(), + _ => panic!("expected var ref"), + }) + .collect::>(); + assert_eq!( + arg_names, + vec![ + "m_flow_ext", + "state1.p", + "state1.T", + "state1.X", + "state2.p", + "state2.T", + "state2.X", + "x_small_actual" + ] + ); + } + + #[test] + fn record_param_lowering_projects_field_access_decomposed_slot_with_scalar_tail() { + let mut flat = flat::Model::new(); + + let mut function = rumoca_core::Function::new("Pkg.brushVoltageDrop", Span::DUMMY); + for input in ["brushParameters_V", "brushParameters_ILinear", "i"] { + function.add_input(rumoca_core::FunctionParam::new(input, "Real", test_span())); + } + function.add_output(rumoca_core::FunctionParam::new("v", "Real", test_span())); + function.body.push(assignment_to("v", var_ref("i"))); + let inputs = function.inputs.clone(); + flat.add_function(function); + + let call_args = vec![ + field_access(int_lit(0), "V"), + field_access(int_lit(0), "ILinear"), + var_ref("computed_current"), + var_ref("dcpm.IaNominal"), + ]; + assert_eq!( + project_actuals_to_input_slots(&call_args, &inputs) + .expect("overexpanded field-access actuals should project") + .len(), + 3 + ); + + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.brushVoltageDrop"), + args: call_args, + is_constructor: false, + span: Span::DUMMY, + }, + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + + let rumoca_core::Expression::FunctionCall { args, .. } = &flat.equations[0].residual else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 3); + assert!(matches!( + &args[0], + rumoca_core::Expression::FieldAccess { field, .. } if field == "V" + )); + assert!(matches!( + &args[1], + rumoca_core::Expression::FieldAccess { field, .. } if field == "ILinear" + )); + assert!(matches!( + &args[2], + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "computed_current" + )); + } + + #[test] + fn record_param_lowering_projects_positional_scalar_decomposed_slot() { + let mut flat = flat::Model::new(); + + let mut function = rumoca_core::Function::new("Pkg.abs", Span::DUMMY); + for input in ["c_re", "c_im"] { + function.add_input(rumoca_core::FunctionParam::new(input, "Real", test_span())); + } + function.add_output(rumoca_core::FunctionParam::new("y", "Real", test_span())); + function.body.push(assignment_to("y", var_ref("c_re"))); + let inputs = function.inputs.clone(); + flat.add_function(function); + + let call_args = vec![ + int_lit(1), + int_lit(0), + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Complex"), + args: vec![int_lit(1), int_lit(0)], + is_constructor: true, + span: Span::DUMMY, + }, + ]; + assert_eq!( + project_actuals_to_input_slots(&call_args, &inputs) + .expect("overexpanded scalar actuals should project") + .len(), + 2 + ); + + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.abs"), + args: call_args, + is_constructor: false, + span: Span::DUMMY, + }, + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + + let rumoca_core::Expression::FunctionCall { args, .. } = &flat.equations[0].residual else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 2); + assert!(matches!( + &args[0], + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(1), + .. + } + )); + assert!(matches!( + &args[1], + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(0), + .. + } + )); + } + + #[test] + fn record_param_lowering_projects_scalar_constructor_fields_without_fake_members() { + let mut flat = flat::Model::new(); + flat.variables.insert( + VarName::new("P"), + flat::Variable::empty_with_span(test_span()), + ); + flat.variables.insert( + VarName::new("Q"), + flat::Variable::empty_with_span(test_span()), + ); + + let mut function = rumoca_core::Function::new("Pkg.arg", Span::DUMMY); + for input in ["c_re", "c_im"] { + function.add_input(rumoca_core::FunctionParam::new(input, "Real", test_span())); + } + function.add_output(rumoca_core::FunctionParam::new("y", "Real", test_span())); + function.body.push(assignment_to("y", var_ref("c_re"))); + flat.add_function(function); + + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.arg"), + args: vec![ + field_access(var_ref("P"), "re"), + field_access(var_ref("P"), "im"), + var_ref("Q"), + ], + is_constructor: false, + span: Span::DUMMY, + }, + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + + let rumoca_core::Expression::FunctionCall { args, .. } = &flat.equations[0].residual else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 2); + assert!(matches!( + &args[0], + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "P" + )); + assert!(matches!( + &args[1], + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "Q" + )); + } + + #[test] + fn record_param_lowering_projects_flattened_scalar_constructor_fields_with_default_tail() { + let mut flat = flat::Model::new(); + flat.variables.insert( + VarName::new("voltageSource.P"), + flat::Variable::empty_with_span(test_span()), + ); + flat.variables.insert( + VarName::new("voltageSource.Q"), + flat::Variable::empty_with_span(test_span()), + ); + + let mut function = rumoca_core::Function::new("Modelica.ComplexMath.arg", Span::DUMMY); + function.add_input(rumoca_core::FunctionParam::new("c_re", "Real", test_span())); + function.add_input(rumoca_core::FunctionParam::new("c_im", "Real", test_span())); + function.add_input( + rumoca_core::FunctionParam::new("phi0", "Real", test_span()).with_default(int_lit(0)), + ); + function.add_output(rumoca_core::FunctionParam::new("phi", "Real", test_span())); + function.body.push(assignment_to("phi", var_ref("c_re"))); + flat.add_function(function); + + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Modelica.ComplexMath.arg"), + args: vec![ + var_ref("voltageSource.P.re"), + var_ref("voltageSource.P.im"), + var_ref("voltageSource.Q"), + ], + is_constructor: false, + span: Span::DUMMY, + }, + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + + let rumoca_core::Expression::FunctionCall { args, .. } = &flat.equations[0].residual else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 2); + assert!(matches!( + &args[0], + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "voltageSource.P" + )); + assert!(matches!( + &args[1], + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "voltageSource.Q" + )); + } + + #[test] + fn record_param_lowering_projects_repeated_field_access_actuals() { + let mut flat = flat::Model::new(); + flat.add_function(record_constructor_named( + "Pkg.Record", + &["phase", "h", "d", "T", "p"], + )); + + let mut function = rumoca_core::Function::new("Pkg.pressure", Span::DUMMY); + for field in ["phase", "h", "d", "T", "p"] { + function.add_input(rumoca_core::FunctionParam::new( + format!("state_{field}"), + "Real", + test_span(), + )); + } + function.add_output(rumoca_core::FunctionParam::new("p", "Real", test_span())); + function.body.push(assignment_to("p", var_ref("state_p"))); + flat.add_function(function); + + let field_access = + |base: rumoca_core::Expression, field: &str| rumoca_core::Expression::FieldAccess { + base: Box::new(base), + field: field.to_string(), + span: test_span(), + }; + let source_state = field_access(var_ref("source"), "state"); + let source_state_phase = field_access(source_state.clone(), "phase"); + let source_state_phase_phase = field_access(source_state_phase.clone(), "phase"); + + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.pressure"), + args: vec![ + field_access(source_state_phase_phase.clone(), "phase"), + field_access(source_state_phase_phase.clone(), "h"), + field_access(source_state_phase_phase.clone(), "d"), + field_access(source_state_phase_phase.clone(), "T"), + field_access(source_state_phase_phase.clone(), "p"), + field_access(source_state_phase.clone(), "h"), + field_access(source_state_phase.clone(), "d"), + field_access(source_state_phase.clone(), "T"), + field_access(source_state_phase.clone(), "p"), + field_access(source_state.clone(), "h"), + field_access(source_state.clone(), "d"), + field_access(source_state.clone(), "T"), + field_access(source_state, "p"), + ], + is_constructor: false, + span: Span::DUMMY, + }, + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + + let rumoca_core::Expression::FunctionCall { args, .. } = &flat.equations[0].residual else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 5); + assert_eq!( + expression_summary_for_projection(&args[0]).as_deref(), + Some("source.state.phase") + ); + } + + #[test] + fn record_param_lowering_projects_repeated_field_access_actuals_from_indexed_base() { + let mut flat = flat::Model::new(); + flat.add_function(record_constructor_named( + "Pkg.Record", + &["phase", "h", "d", "T", "p"], + )); + + let mut function = rumoca_core::Function::new("Pkg.pressure", Span::DUMMY); + for field in ["phase", "h", "d", "T", "p"] { + function.add_input(rumoca_core::FunctionParam::new( + format!("state_{field}"), + "Real", + test_span(), + )); + } + function.add_output(rumoca_core::FunctionParam::new("p", "Real", test_span())); + function.body.push(assignment_to("p", var_ref("state_p"))); + flat.add_function(function); + + let field_access = + |base: rumoca_core::Expression, field: &str| rumoca_core::Expression::FieldAccess { + base: Box::new(base), + field: field.to_string(), + span: test_span(), + }; + let indexed_state = rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("pipe.flowModel.states"), + subscripts: vec![rumoca_core::Subscript::Index { + value: 1, + span: test_span(), + }], + span: test_span(), + }; + let state_phase = field_access(indexed_state.clone(), "phase"); + let state_phase_phase = field_access(state_phase.clone(), "phase"); + + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.pressure"), + args: vec![ + field_access(state_phase_phase.clone(), "phase"), + field_access(state_phase_phase.clone(), "h"), + field_access(state_phase_phase.clone(), "d"), + field_access(state_phase_phase.clone(), "T"), + field_access(state_phase_phase.clone(), "p"), + field_access(state_phase.clone(), "h"), + field_access(state_phase.clone(), "d"), + field_access(state_phase.clone(), "T"), + field_access(state_phase.clone(), "p"), + field_access(indexed_state.clone(), "h"), + field_access(indexed_state.clone(), "d"), + field_access(indexed_state.clone(), "T"), + field_access(indexed_state, "p"), + ], + is_constructor: false, + span: Span::DUMMY, + }, + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + + let rumoca_core::Expression::FunctionCall { args, .. } = &flat.equations[0].residual else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 5); + assert_eq!( + expression_summary_for_projection(&args[0]).as_deref(), + Some("pipe.flowModel.states[1].phase") + ); + } + + #[test] + fn record_param_lowering_projects_overexpanded_field_chain_from_indexed_record_base() { + let mut flat = flat::Model::new(); + flat.add_function(record_constructor_named( + "Pkg.Record", + &["phase", "h", "d", "T", "p"], + )); + + let mut function = rumoca_core::Function::new("Pkg.pressure", Span::DUMMY); + for field in ["phase", "h", "d", "T", "p"] { + function.add_input(rumoca_core::FunctionParam::new( + format!("state_{field}"), + "Real", + test_span(), + )); + } + function.add_output(rumoca_core::FunctionParam::new("p", "Real", test_span())); + function.body.push(assignment_to("p", var_ref("state_p"))); + flat.add_function(function); + + let field_access = + |base: rumoca_core::Expression, field: &str| rumoca_core::Expression::FieldAccess { + base: Box::new(base), + field: field.to_string(), + span: test_span(), + }; + let indexed_state = rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("pipe.flowModel.states"), + subscripts: vec![rumoca_core::Subscript::Index { + value: 1, + span: test_span(), + }], + span: test_span(), + }; + let state_phase = field_access(indexed_state, "phase"); + let state_phase_phase = field_access(state_phase, "phase"); + + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.pressure"), + args: vec![ + field_access(state_phase_phase.clone(), "phase"), + field_access(state_phase_phase.clone(), "h"), + field_access(state_phase_phase.clone(), "d"), + field_access(state_phase_phase.clone(), "T"), + field_access(state_phase_phase, "p"), + ], + is_constructor: false, + span: Span::DUMMY, + }, + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + + let rumoca_core::Expression::FunctionCall { args, .. } = &flat.equations[0].residual else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 5); + assert_eq!( + expression_summary_for_projection(&args[0]).as_deref(), + Some("pipe.flowModel.states[1].phase") + ); + } + + #[test] + fn record_param_lowering_rewrites_variable_attribute_calls() { + let mut flat = flat::Model::new(); + flat.add_function(record_constructor()); + flat.add_function(function_with_record_input()); + + let mut variable = flat::Variable::empty_with_span(test_span()); + variable.name = VarName::new("x"); + variable.nominal = Some(rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.f"), + args: vec![var_ref("rec")], + is_constructor: false, + span: Span::DUMMY, + }); + flat.add_variable(VarName::new("x"), variable); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + + let variable = flat.variables.get(&VarName::new("x")).expect("variable"); + let Some(rumoca_core::Expression::FunctionCall { args, .. }) = &variable.nominal else { + panic!("expected function call nominal"); + }; + assert_eq!(args.len(), 2); + assert!(matches!( + &args[0], + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "rec.a" + )); + } + + #[test] + fn record_param_lowering_expands_named_record_actual_value() { + let mut flat = flat::Model::new(); + flat.add_function(record_constructor()); + flat.add_function(function_with_record_input()); + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.f"), + args: vec![named_arg("r", var_ref("rec"))], + is_constructor: false, + span: Span::DUMMY, + }, + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + + let rumoca_core::Expression::FunctionCall { args, .. } = &flat.equations[0].residual else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 2); + assert!(matches!( + &args[0], + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "rec.a" + )); + assert!(matches!( + &args[1], + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "rec.b" + )); + } + + #[test] + fn record_param_lowering_does_not_positionally_expand_mismatched_constructor() { + let mut flat = flat::Model::new(); + flat.add_function(record_constructor()); + flat.add_function(record_constructor_named("Pkg.OtherRecord", &["x", "y"])); + flat.add_function(function_with_record_input()); + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.f"), + args: vec![rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.OtherRecord"), + args: vec![var_ref("first"), var_ref("second")], + is_constructor: true, + span: Span::DUMMY, + }], + is_constructor: false, + span: Span::DUMMY, + }, + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + + let rumoca_core::Expression::FunctionCall { args, .. } = &flat.equations[0].residual else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 2); + assert!(matches!( + &args[0], + rumoca_core::Expression::FieldAccess { field, .. } if field == "a" + )); + assert!(matches!( + &args[1], + rumoca_core::Expression::FieldAccess { field, .. } if field == "b" + )); + } + + #[test] + fn record_param_lowering_expands_single_record_value_constructor_proxy() { + let mut flat = flat::Model::new(); + flat.add_function(record_constructor()); + flat.add_function(function_with_record_input()); + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.f"), + args: vec![rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.Record"), + args: vec![component_ref_expr(&["carrier", "r"])], + is_constructor: true, + span: test_span(), + }], + is_constructor: false, + span: Span::DUMMY, + }, + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + + let rumoca_core::Expression::FunctionCall { args, .. } = &flat.equations[0].residual else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 2); + assert!(matches!( + &args[0], + rumoca_core::Expression::FieldAccess { base, field, .. } + if field == "a" + && matches!(base.as_ref(), rumoca_core::Expression::VarRef { name, .. } + if name.as_str() == "carrier.r") + )); + assert!(matches!( + &args[1], + rumoca_core::Expression::FieldAccess { base, field, .. } + if field == "b" + && matches!(base.as_ref(), rumoca_core::Expression::VarRef { name, .. } + if name.as_str() == "carrier.r") + )); + } + + #[test] + fn record_param_lowering_expands_multiple_named_record_actuals() { + let mut flat = flat::Model::new(); + flat.add_function(record_constructor()); + + let mut function = rumoca_core::Function::new("Pkg.combine", Span::DUMMY); + function.add_input( + rumoca_core::FunctionParam::new("left", "Pkg.Record", test_span()) + .with_type_class(ClassType::Record), + ); + function.add_input( + rumoca_core::FunctionParam::new("right", "Pkg.Record", test_span()) + .with_type_class(ClassType::Record), + ); + function.add_output(rumoca_core::FunctionParam::new("y", "Real", test_span())); + function.body.push(assignment_to( + "y", + rumoca_core::Expression::FieldAccess { + base: Box::new(var_ref("left")), + field: "a".to_string(), + span: Span::DUMMY, + }, + )); + flat.add_function(function); + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.combine"), + args: vec![ + named_arg("right", var_ref("r2")), + named_arg("left", var_ref("r1")), + ], + is_constructor: false, + span: Span::DUMMY, + }, + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + + let rumoca_core::Expression::FunctionCall { args, .. } = &flat.equations[0].residual else { + panic!("expected function call"); + }; + let actuals = args + .iter() + .map(|arg| match arg { + rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor: true, + .. + } => { + let value = match args.first() { + Some(rumoca_core::Expression::VarRef { name, .. }) => { + name.as_str().to_string() + } + other => format!("{other:?}"), + }; + format!("{}={value}", name.as_str()) + } + other => format!("{other:?}"), + }) + .collect::>(); + assert_eq!( + actuals, + vec![ + "__rumoca_named_arg__.left_a=r1.a", + "__rumoca_named_arg__.left_b=r1.b", + "__rumoca_named_arg__.right_a=r2.a", + "__rumoca_named_arg__.right_b=r2.b" + ] + ); + } + + #[test] + fn record_param_lowering_keeps_decomposed_positional_record_named_after_named_slots() { + let mut flat = flat::Model::new(); + flat.add_function(record_constructor()); + + let mut function = rumoca_core::Function::new("Pkg.withScale", Span::DUMMY); + function.add_input(rumoca_core::FunctionParam::new( + "scale", + "Real", + test_span(), + )); + function.add_input( + rumoca_core::FunctionParam::new("r", "Pkg.Record", test_span()) + .with_type_class(ClassType::Record), + ); + function.add_output(rumoca_core::FunctionParam::new("y", "Real", test_span())); + flat.add_function(function); + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.withScale"), + args: vec![named_arg("scale", var_ref("gain")), var_ref("rec")], + is_constructor: false, + span: Span::DUMMY, + }, + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + + let rumoca_core::Expression::FunctionCall { args, .. } = &flat.equations[0].residual else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 3); + assert!(matches!( + named_arg_var_ref(&args[0], "scale"), + Some(name) if name.as_str() == "gain" + )); + assert!(matches!( + named_arg_var_ref(&args[1], "r_a"), + Some(name) if name.as_str() == "rec.a" + )); + assert!(matches!( + named_arg_var_ref(&args[2], "r_b"), + Some(name) if name.as_str() == "rec.b" + )); + } + + #[test] + fn record_param_lowering_uses_mismatched_constructor_positional_field_values() { + let mut flat = flat::Model::new(); + flat.add_function(record_constructor()); + flat.add_function(record_constructor_named("Pkg.OtherRecord", &["x", "y"])); + + let mut function = rumoca_core::Function::new("Pkg.withScale", Span::DUMMY); + function.add_input(rumoca_core::FunctionParam::new( + "scale", + "Real", + test_span(), + )); + function.add_input( + rumoca_core::FunctionParam::new("r", "Pkg.Record", test_span()) + .with_type_class(ClassType::Record), + ); + function.add_output(rumoca_core::FunctionParam::new("y", "Real", test_span())); + flat.add_function(function); + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.withScale"), + args: vec![ + named_arg("scale", var_ref("gain")), + named_arg( + "r", + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.OtherRecord"), + args: vec![ + component_ref_expr(&["rec", "a"]), + component_ref_expr(&["rec", "b"]), + ], + is_constructor: true, + span: Span::DUMMY, + }, + ), + ], + is_constructor: false, + span: Span::DUMMY, + }, + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + + let rumoca_core::Expression::FunctionCall { args, .. } = &flat.equations[0].residual else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 3); + let actuals = args + .iter() + .map(|arg| match arg { + rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor: true, + .. + } => { + let value = match args.first() { + Some(rumoca_core::Expression::VarRef { name, .. }) => { + name.as_str().to_string() + } + other => format!("{other:?}"), + }; + format!("{}={value}", name.as_str()) + } + other => format!("{other:?}"), + }) + .collect::>(); + assert_eq!( + actuals, + vec![ + "__rumoca_named_arg__.scale=gain", + "__rumoca_named_arg__.r_a=rec.a", + "__rumoca_named_arg__.r_b=rec.b" + ] + ); + } + + #[test] + fn record_param_lowering_projects_mismatched_constructor_fields_from_flat_namespace() { + let mut flat = flat::Model::new(); + flat.add_function(record_constructor_named( + "Pkg.EfficiencyParameters", + &["V_flow", "eta"], + )); + flat.add_function(record_constructor_named( + "Pkg.Generic", + &["pressure_V_flow", "pressure_dp"], + )); + flat.add_variable( + rumoca_core::VarName::new("per.hydraulicEfficiency.V_flow"), + flat::Variable::empty_with_span(test_span()), + ); + flat.add_variable( + rumoca_core::VarName::new("per.hydraulicEfficiency.eta"), + flat::Variable::empty_with_span(test_span()), + ); + + let mut function = rumoca_core::Function::new("Pkg.efficiency", Span::DUMMY); + function.add_input( + rumoca_core::FunctionParam::new("per", "Pkg.EfficiencyParameters", test_span()) + .with_type_class(ClassType::Record), + ); + function.add_output(rumoca_core::FunctionParam::new("y", "Real", test_span())); + flat.add_function(function); + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.efficiency"), + args: vec![named_arg( + "per", + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.Generic"), + args: vec![ + component_ref_expr(&["per", "hydraulicEfficiency", "V_flow"]), + component_ref_expr(&["per", "hydraulicEfficiency", "dp"]), + ], + is_constructor: true, + span: Span::DUMMY, + }, + )], + is_constructor: false, + span: Span::DUMMY, + }, + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + let flat_variable_names = flat.variables.keys().cloned().collect::>(); + let fields = record_fields_from_constructor_metadata( + &flat.functions, + "Pkg.EfficiencyParameters", + None, + ) + .expect("record fields"); + let probe_actuals = [ + component_ref_expr(&["per", "hydraulicEfficiency", "V_flow"]), + component_ref_expr(&["per", "hydraulicEfficiency", "dp"]), + ]; + let probe_refs = probe_actuals.iter().collect::>(); + assert!( + constructor_positional_args_project_expected_fields( + &probe_refs, + &fields, + &flat_variable_names, + ) + .is_some(), + "direct projection helper should project from the visible field prefix" + ); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + + let rumoca_core::Expression::FunctionCall { args, .. } = &flat.equations[0].residual else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 2); + assert!(matches!( + &args[0], + rumoca_core::Expression::VarRef { name, .. } + if name.as_str() == "per.hydraulicEfficiency.V_flow" + )); + assert!(matches!( + &args[1], + rumoca_core::Expression::VarRef { name, .. } + if name.as_str() == "per.hydraulicEfficiency.eta" + )); } #[test] - fn record_param_lowering_uses_constructor_signature_metadata() { + fn record_param_lowering_does_not_synthesize_missing_fields_from_partial_prefix() { let mut flat = flat::Model::new(); - flat.add_function(record_constructor()); - flat.add_function(function_with_record_input()); + flat.add_function(record_constructor_named("Medium.State", &["p", "T", "X"])); + flat.add_function(record_constructor_named( + "Medium.setState_phX", + &["p", "h", "X"], + )); + flat.add_variable( + rumoca_core::VarName::new("port_a.p"), + flat::Variable::empty_with_span(test_span()), + ); + flat.add_variable( + rumoca_core::VarName::new("port_a.h_outflow"), + flat::Variable::empty_with_span(test_span()), + ); + flat.add_variable( + rumoca_core::VarName::new("port_a.Xi_outflow"), + flat::Variable::empty_with_span(test_span()), + ); + + let mut density = rumoca_core::Function::new("Medium.density", Span::DUMMY); + density.add_input( + rumoca_core::FunctionParam::new("state", "Medium.State", test_span()) + .with_type_class(ClassType::Record), + ); + density.add_output(rumoca_core::FunctionParam::new("d", "Real", test_span())); + flat.add_function(density); flat.add_equation(flat::Equation::new( rumoca_core::Expression::FunctionCall { - name: rumoca_core::Reference::new("Pkg.f"), - args: vec![var_ref("rec")], + name: rumoca_core::Reference::new("Medium.density"), + args: vec![rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Medium.setState_phX"), + args: vec![ + component_ref_expr(&["port_a", "p"]), + component_ref_expr(&["port_a", "h_outflow"]), + component_ref_expr(&["port_a", "Xi_outflow"]), + ], + is_constructor: true, + span: test_span(), + }], is_constructor: false, span: Span::DUMMY, }, @@ -812,36 +3811,301 @@ mod tests { lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + let rumoca_core::Expression::FunctionCall { args, .. } = &flat.equations[0].residual else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 3); + assert!(args.iter().all(|arg| { + !matches!( + arg, + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "port_a.T" + ) + })); + assert!(matches!( + &args[1], + rumoca_core::Expression::FieldAccess { base, field, .. } + if field == "T" + && matches!(base.as_ref(), rumoca_core::Expression::FunctionCall { name, .. } + if name.as_str() == "Medium.setState_phX") + )); + } + + #[test] + fn record_param_lowering_prefers_more_specific_constructor_metadata() { + let mut flat = flat::Model::new(); + flat.add_function(record_constructor_named( + "Modelica.Media.Interfaces.PartialSimpleMedium.ThermodynamicState", + &["p", "T"], + )); + flat.add_function(record_constructor_named( + "Buildings.Media.Air.ThermodynamicState", + &["p", "T", "X"], + )); + + let mut function = rumoca_core::Function::new("Buildings.Media.Air.h", Span::DUMMY); + function.add_input( + rumoca_core::FunctionParam::new("state", "ThermodynamicState", test_span()) + .with_type_class(ClassType::Record), + ); + function.add_output(rumoca_core::FunctionParam::new("h", "Real", test_span())); + function.body.push(assignment_to( + "h", + rumoca_core::Expression::Index { + base: Box::new(rumoca_core::Expression::FieldAccess { + base: Box::new(var_ref("state")), + field: "X".to_string(), + span: test_span(), + }), + subscripts: vec![rumoca_core::Subscript::Index { + value: 1, + span: test_span(), + }], + span: test_span(), + }, + )); + flat.add_function(function); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + let function = flat .functions - .get(&VarName::new("Pkg.f")) + .get(&VarName::new("Buildings.Media.Air.h")) .expect("function remains"); let input_names = function .inputs .iter() .map(|input| input.name.as_str()) .collect::>(); - assert_eq!(input_names, vec!["r_a", "r_b"]); - assert_eq!(function.inputs[0].dims, Vec::::new()); - assert_eq!(function.inputs[1].dims, vec![3]); + assert_eq!(input_names, vec!["state_p", "state_T", "state_X"]); let rumoca_core::Statement::Assignment { value, .. } = &function.body[0] else { panic!("expected assignment"); }; assert!(matches!( value, - rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "r_a" + rumoca_core::Expression::Index { base, .. } + if matches!(base.as_ref(), rumoca_core::Expression::VarRef { name, .. } + if name.as_str() == "state_X") + )); + } + + #[test] + fn record_param_lowering_keeps_unqualified_inherited_record_in_owner_context() { + let mut flat = flat::Model::new(); + flat.add_function(record_constructor_named( + "Modelica.Media.Interfaces.PartialSimpleMedium.ThermodynamicState", + &["p", "T"], + )); + flat.add_function(record_constructor_named( + "Buildings.Media.Air.ThermodynamicState", + &["p", "T", "X"], + )); + + let mut function = + rumoca_core::Function::new("Buildings.Media.Water.temperature", Span::DUMMY); + function.add_input( + rumoca_core::FunctionParam::new("state", "ThermodynamicState", test_span()) + .with_type_class(ClassType::Record), + ); + function.add_output(rumoca_core::FunctionParam::new("T", "Real", test_span())); + function.body.push(assignment_to( + "T", + rumoca_core::Expression::FieldAccess { + base: Box::new(var_ref("state")), + field: "T".to_string(), + span: test_span(), + }, + )); + flat.add_function(function); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + + let function = flat + .functions + .get(&VarName::new("Buildings.Media.Water.temperature")) + .expect("function remains"); + let input_names = function + .inputs + .iter() + .map(|input| input.name.as_str()) + .collect::>(); + assert_eq!(input_names, vec!["state_p", "state_T"]); + } + + #[test] + fn record_param_lowering_projects_partial_mismatched_constructor_field_prefix() { + let mut flat = flat::Model::new(); + flat.add_function(record_constructor_named( + "Pkg.EfficiencyParameters", + &["V_flow", "eta"], + )); + flat.add_function(record_constructor_named( + "Pkg.Generic", + &["pressure_V_flow", "pressure_dp"], + )); + flat.add_variable( + rumoca_core::VarName::new("per.pum[1].hydraulicEfficiency.V_flow"), + flat::Variable::empty_with_span(test_span()), + ); + flat.add_variable( + rumoca_core::VarName::new("per.pum[1].hydraulicEfficiency.eta"), + flat::Variable::empty_with_span(test_span()), + ); + + let mut function = rumoca_core::Function::new("Pkg.efficiency", Span::DUMMY); + function.add_input( + rumoca_core::FunctionParam::new("per", "Pkg.EfficiencyParameters", test_span()) + .with_type_class(ClassType::Record), + ); + function.add_output(rumoca_core::FunctionParam::new("y", "Real", test_span())); + flat.add_function(function); + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.efficiency"), + args: vec![named_arg( + "per", + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.Generic"), + args: vec![ + var_ref("per.pum[1].hydraulicEfficiency.V_flow"), + rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.EfficiencyParameters"), + args: vec![ + named_arg("V_flow", var_ref("default_v")), + named_arg("eta", var_ref("default_eta")), + ], + is_constructor: true, + span: Span::DUMMY, + }), + field: "dp".to_string(), + span: test_span(), + }, + ], + is_constructor: true, + span: Span::DUMMY, + }, + )], + is_constructor: false, + span: Span::DUMMY, + }, + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, )); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + let rumoca_core::Expression::FunctionCall { args, .. } = &flat.equations[0].residual else { panic!("expected function call"); }; assert_eq!(args.len(), 2); assert!(matches!( &args[0], - rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "rec.a" + rumoca_core::Expression::VarRef { name, .. } + if name.as_str() == "per.pum[1].hydraulicEfficiency.V_flow" )); assert!(matches!( &args[1], - rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "rec.b" + rumoca_core::Expression::VarRef { name, .. } + if name.as_str() == "per.pum[1].hydraulicEfficiency.eta" + )); + } + + #[test] + fn constructor_projection_requires_all_projected_record_fields_to_exist() { + let mut flat_variable_names = HashSet::new(); + flat_variable_names.insert(rumoca_core::VarName::new("userValve.port_a.p")); + let fields = ["phase", "h", "d", "T", "p"] + .into_iter() + .map(|field| rumoca_core::FunctionParam::new(field, "Real", test_span())) + .collect::>(); + let actual = component_ref_expr(&["userValve", "port_a", "p"]); + let actuals = [&actual]; + + let projected = constructor_positional_args_project_expected_fields( + &actuals, + &fields, + &flat_variable_names, + ); + + assert!( + projected.is_none(), + "a matching leaf like port_a.p must not imply that port_a is a complete record" + ); + } + + #[test] + fn record_param_lowering_projects_nested_constructor_proxy_fields() { + let mut flat = flat::Model::new(); + flat.add_function(record_constructor_named( + "Pkg.EfficiencyParameters", + &["V_flow", "eta"], + )); + flat.add_function(record_constructor_named( + "Pkg.Generic", + &[ + "pressure_V_flow", + "pressure_dp", + "hydraulicEfficiency_V_flow", + "hydraulicEfficiency_eta", + ], + )); + + let mut function = rumoca_core::Function::new("Pkg.efficiency", Span::DUMMY); + function.add_input( + rumoca_core::FunctionParam::new("per", "Pkg.EfficiencyParameters", test_span()) + .with_type_class(ClassType::Record), + ); + function.add_output(rumoca_core::FunctionParam::new("y", "Real", test_span())); + flat.add_function(function); + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.efficiency"), + args: vec![rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.EfficiencyParameters"), + args: vec![rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.Generic"), + args: vec![ + var_ref("pressure_v"), + var_ref("pressure_dp"), + var_ref("hyd_v"), + var_ref("hyd_eta"), + ], + is_constructor: true, + span: Span::DUMMY, + }), + field: "hydraulicEfficiency".to_string(), + span: test_span(), + }], + is_constructor: true, + span: Span::DUMMY, + }], + is_constructor: false, + span: Span::DUMMY, + }, + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + + let rumoca_core::Expression::FunctionCall { args, .. } = &flat.equations[0].residual else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 2); + assert!(matches!( + &args[0], + rumoca_core::Expression::VarRef { name, .. } + if name.as_str() == "hyd_v" + )); + assert!(matches!( + &args[1], + rumoca_core::Expression::VarRef { name, .. } + if name.as_str() == "hyd_eta" )); } @@ -951,6 +4215,54 @@ mod tests { )); } + #[test] + fn record_param_lowering_rewrites_size_of_record_field_ref() { + let mut flat = flat::Model::new(); + flat.add_function(record_constructor()); + + let mut function = rumoca_core::Function::new("Pkg.sizeField", Span::DUMMY); + function.add_input( + rumoca_core::FunctionParam::new("r", "Pkg.Record", test_span()) + .with_type_class(ClassType::Record), + ); + function.add_output(rumoca_core::FunctionParam::new("y", "Real", test_span())); + function.locals.push(rumoca_core::FunctionParam { + def_id: None, + name: "n".to_string(), + type_name: "Integer".to_string(), + default: Some(rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + args: vec![ + component_ref_expr(&["r", "b"]), + rumoca_core::Expression::Literal { + value: Literal::Integer(1), + span: Span::DUMMY, + }, + ], + span: Span::DUMMY, + }), + ..rumoca_core::FunctionParam::new("n", "Integer", test_span()) + }); + flat.add_function(function); + + lower_record_function_params(&mut flat).expect("record parameter lowering should pass"); + + let function = flat + .functions + .get(&VarName::new("Pkg.sizeField")) + .expect("function remains"); + let Some(default) = function.locals[0].default.as_ref() else { + panic!("expected local default"); + }; + let rumoca_core::Expression::BuiltinCall { args, .. } = default else { + panic!("expected size builtin"); + }; + assert!(matches!( + &args[0], + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "r_a" + )); + } + #[test] fn record_param_lowering_leaves_unknown_record_metadata_unexpanded() { let mut flat = flat::Model::new(); @@ -985,6 +4297,107 @@ mod tests { assert_eq!(args.len(), 1); } + #[test] + fn record_param_lowering_keeps_external_object_inputs_opaque() { + let mut flat = flat::Model::new(); + let mut external_constructor = rumoca_core::Function::new("Pkg.ExternalTable", Span::DUMMY); + external_constructor.is_constructor = true; + external_constructor.add_input(rumoca_core::FunctionParam::new( + "fileName", + "String", + test_span(), + )); + external_constructor.add_input(rumoca_core::FunctionParam::new( + "tableName", + "String", + test_span(), + )); + flat.add_function(external_constructor); + + let mut external_reader = rumoca_core::Function::new("Pkg.getMin", Span::DUMMY); + external_reader.add_input( + rumoca_core::FunctionParam::new("tableID", "Pkg.ExternalTable", test_span()) + .with_type_class(ClassType::Class), + ); + external_reader.add_output(rumoca_core::FunctionParam::new("y", "Real", test_span())); + flat.add_function(external_reader); + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.getMin"), + args: vec![var_ref("tableID")], + is_constructor: false, + span: Span::DUMMY, + }, + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + lower_record_function_params(&mut flat) + .expect("opaque external object should not decompose"); + + let function = flat + .functions + .get(&VarName::new("Pkg.getMin")) + .expect("function remains"); + assert_eq!(function.inputs.len(), 1); + assert_eq!(function.inputs[0].name, "tableID"); + let rumoca_core::Expression::FunctionCall { args, .. } = &flat.equations[0].residual else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 1); + } + + #[test] + fn scalar_function_projection_keeps_single_field_actual_name() { + let actual = field_access(var_ref("ductOut.mediums[1]"), "T"); + let fields = vec![rumoca_core::FunctionParam::new( + "Kelvin", + "Modelica.Units.SI.Temperature", + test_span(), + )]; + + let projected = project_decomposed_scalar_actuals(&[actual], &fields, &HashSet::new()); + + assert!( + projected.is_none(), + "single scalar formal input must not rewrite an actual field to the formal name" + ); + } + + #[test] + fn single_field_decomposition_preserves_mismatched_actual_field() { + let actual = field_access(var_ref("ductOut.mediums[1]"), "T"); + let fields = vec![rumoca_core::FunctionParam::new( + "Kelvin", + "Modelica.Units.SI.Temperature", + test_span(), + )]; + let decomposed = vec![DecomposedParam { + original_index: 0, + param_name: "Kelvin".to_string(), + type_name: "Modelica.Units.SI.Temperature".to_string(), + fields, + already_decomposed: false, + }]; + + let args = decompose_record_call_args( + "Modelica.Units.Conversions.to_degC", + &[actual], + &decomposed, + None, + &HashSet::new(), + &HashMap::new(), + ) + .expect("single scalar field mismatch should preserve actual"); + + let rumoca_core::Expression::FieldAccess { field, .. } = &args[0] else { + panic!("expected original field access"); + }; + assert_eq!(field, "T"); + } + #[test] fn record_field_normalization_uses_structured_component_ref_parts() { let mut function = rumoca_core::Function::new("Pkg.f", Span::DUMMY); diff --git a/crates/rumoca-phase-flatten/src/functions.rs b/crates/rumoca-phase-flatten/src/functions.rs index 0fe932c9f..646a69a4a 100644 --- a/crates/rumoca-phase-flatten/src/functions.rs +++ b/crates/rumoca-phase-flatten/src/functions.rs @@ -1598,11 +1598,13 @@ fn convert_callable<'tree>( .map(Some), class_type if is_callable_class_type(class_type) => { Ok(Some(convert_constructor_signature( + tree, class_index, class_def, qualified_name, source_map, def_map, + member_cache, )?)) } _ => Ok(None), @@ -1637,6 +1639,7 @@ fn convert_function<'tree>( // Process components to find inputs, outputs, and locals for (comp_name, component) in &effective_components { let param = convert_component_to_param( + tree, class_index, comp_name, component, @@ -1693,6 +1696,7 @@ fn convert_function<'tree>( prefix: &prefix, imports: &import_map, def_map: Some(&filtered_def_map), + class_tree: Some(tree), initial_locals: &function_locals, source_map: Some(source_map), instance_name: None, diff --git a/crates/rumoca-phase-flatten/src/functions/call_args.rs b/crates/rumoca-phase-flatten/src/functions/call_args.rs index a0faef0c0..87aea78b2 100644 --- a/crates/rumoca-phase-flatten/src/functions/call_args.rs +++ b/crates/rumoca-phase-flatten/src/functions/call_args.rs @@ -114,11 +114,23 @@ fn validate_call_slots( let inputs = &function.inputs; if positional > inputs.len() { + let input_names = inputs + .iter() + .map(|input| input.name.as_str()) + .collect::>() + .join(", "); + let actuals = args + .iter() + .filter(|arg| !is_named_arg_marker(arg)) + .take(16) + .map(expression_summary) + .collect::>() + .join(", "); return Err(FlattenError::invalid_function_call_args( call_name, format!( - "{positional} positional argument(s) for {} input slot(s)", - inputs.len() + "{positional} positional argument(s) for {} input slot(s) [{input_names}]; actuals: {actuals}", + inputs.len(), ), span, )); @@ -216,3 +228,15 @@ fn is_named_arg_marker(arg: &rumoca_core::Expression) -> bool { if name.as_str().starts_with(rumoca_core::NAMED_FUNCTION_ARG_PREFIX) ) } + +fn expression_summary(expr: &rumoca_core::Expression) -> String { + match expr { + rumoca_core::Expression::VarRef { name, .. } => name.as_str().to_string(), + rumoca_core::Expression::FieldAccess { base, field, .. } => { + format!("{}.{}", expression_summary(base), field) + } + rumoca_core::Expression::FunctionCall { name, .. } => format!("{}(...)", name.as_str()), + rumoca_core::Expression::Literal { value, .. } => format!("{value:?}"), + _ => format!("{expr:?}"), + } +} diff --git a/crates/rumoca-phase-flatten/src/functions/constructor_signature.rs b/crates/rumoca-phase-flatten/src/functions/constructor_signature.rs index 86ae68eda..e4f8a2747 100644 --- a/crates/rumoca-phase-flatten/src/functions/constructor_signature.rs +++ b/crates/rumoca-phase-flatten/src/functions/constructor_signature.rs @@ -2,6 +2,7 @@ use super::*; use crate::source_spans::required_location_span; fn collect_constructor_params( + tree: &ast::ClassTree, class_index: &ast::ClassDefIndex<'_>, class_def: &ast::ClassDef, visited_classes: &mut HashSet, @@ -25,6 +26,7 @@ fn collect_constructor_params( }); if let Some(base_class) = base_class { collect_constructor_params( + tree, class_index, base_class, visited_classes, @@ -49,6 +51,7 @@ fn collect_constructor_params( }); } let param = convert_component_to_param( + tree, class_index, comp_name, component, @@ -124,19 +127,45 @@ fn constructor_def_map( } /// Build a synthetic constructor signature for constructor-like class calls. -pub(super) fn convert_constructor_signature( - class_index: &ast::ClassDefIndex<'_>, - class_def: &ast::ClassDef, +pub(super) fn convert_constructor_signature<'tree>( + tree: &ast::ClassTree, + class_index: &ast::ClassDefIndex<'tree>, + class_def: &'tree ast::ClassDef, qualified_name: &str, source_map: &rumoca_core::SourceMap, def_map: &crate::ResolveDefMap, + member_cache: &mut qualify::MemberDefIdCache<'tree>, ) -> Result { + if let Some(function) = external_object_constructor_function( + tree, + class_index, + class_def, + qualified_name, + source_map, + def_map, + member_cache, + )? { + return Ok(function); + } + if let Some(function) = operator_record_constructor_function( + tree, + class_index, + class_def, + qualified_name, + source_map, + def_map, + member_cache, + )? { + return Ok(function); + } + let span = required_location_span(source_map, &class_def.location, "constructor signature")?; let mut params = Vec::new(); let mut param_index = HashMap::new(); let mut visited_classes = HashSet::new(); let constructor_def_map = constructor_def_map(class_index, class_def, def_map); collect_constructor_params( + tree, class_index, class_def, &mut visited_classes, @@ -156,6 +185,134 @@ pub(super) fn convert_constructor_signature( Ok(func) } +fn external_object_constructor_function<'tree>( + tree: &ast::ClassTree, + class_index: &ast::ClassDefIndex<'tree>, + class_def: &'tree ast::ClassDef, + qualified_name: &str, + source_map: &rumoca_core::SourceMap, + def_map: &crate::ResolveDefMap, + member_cache: &mut qualify::MemberDefIdCache<'tree>, +) -> Result, FlattenError> { + let Some(constructor_def) = class_def.classes.get("constructor") else { + return Ok(None); + }; + if constructor_def.class_type != rumoca_core::ClassType::Function + || constructor_def.external.is_none() + { + return Ok(None); + } + + let nested_name = format!("{qualified_name}.constructor"); + let mut function = super::convert_function( + tree, + class_index, + constructor_def, + &nested_name, + source_map, + def_map, + member_cache, + )?; + function.name = rumoca_core::VarName::new(qualified_name); + function.def_id = class_def.def_id; + function.is_constructor = true; + Ok(Some(function)) +} + +fn operator_record_constructor_function<'tree>( + tree: &ast::ClassTree, + class_index: &ast::ClassDefIndex<'tree>, + class_def: &'tree ast::ClassDef, + qualified_name: &str, + source_map: &rumoca_core::SourceMap, + def_map: &crate::ResolveDefMap, + member_cache: &mut qualify::MemberDefIdCache<'tree>, +) -> Result, FlattenError> { + if !class_def.operator_record { + return Ok(None); + } + let Some(constructor_operator) = class_def + .classes + .get("'constructor'") + .or_else(|| class_def.classes.get("constructor")) + else { + return Ok(None); + }; + if constructor_operator.class_type != rumoca_core::ClassType::Operator { + return Ok(None); + } + + for (function_name, function_def) in &constructor_operator.classes { + if function_def.class_type != rumoca_core::ClassType::Function { + continue; + } + let nested_name = format!("{qualified_name}.'constructor'.{function_name}"); + let mut function = super::convert_function( + tree, + class_index, + function_def, + &nested_name, + source_map, + def_map, + member_cache, + )?; + if !constructor_output_matches_record(&function, qualified_name, &class_def.name.text) { + continue; + } + if !operator_constructor_inputs_match_record_fields(&function, class_def) { + continue; + } + remap_operator_constructor_inputs_to_record_fields(&mut function, class_def); + function.name = rumoca_core::VarName::new(qualified_name); + function.def_id = class_def.def_id; + function.is_constructor = true; + normalize_function_local_references(&mut function); + return Ok(Some(function)); + } + Ok(None) +} + +fn operator_constructor_inputs_match_record_fields( + function: &rumoca_core::Function, + class_def: &ast::ClassDef, +) -> bool { + !function.inputs.is_empty() + && function + .inputs + .iter() + .all(|input| class_def.components.contains_key(input.name.as_str())) +} + +fn remap_operator_constructor_inputs_to_record_fields( + function: &mut rumoca_core::Function, + class_def: &ast::ClassDef, +) { + for input in &mut function.inputs { + if let Some(field_def_id) = class_def + .components + .get(input.name.as_str()) + .and_then(|field| field.def_id) + { + input.def_id = Some(field_def_id); + } + } +} + +fn constructor_output_matches_record( + function: &rumoca_core::Function, + qualified_name: &str, + short_name: &str, +) -> bool { + let [output] = function.outputs.as_slice() else { + return false; + }; + output.type_name == qualified_name + || output.type_name == short_name + || qualified_name + .strip_suffix(short_name) + .is_some_and(|prefix| prefix.ends_with('.') && output.type_name == short_name) +} + pub(super) fn normalize_function_local_references(function: &mut rumoca_core::Function) { let locals = function .inputs diff --git a/crates/rumoca-phase-flatten/src/functions/function_metadata.rs b/crates/rumoca-phase-flatten/src/functions/function_metadata.rs index 4f20f1c16..b5a80d6c1 100644 --- a/crates/rumoca-phase-flatten/src/functions/function_metadata.rs +++ b/crates/rumoca-phase-flatten/src/functions/function_metadata.rs @@ -6,6 +6,12 @@ pub(super) fn convert_external_function( ext: &rumoca_ir_ast::ExternalFunction, _default_name: &str, ) -> rumoca_core::ExternalFunction { + let mut metadata = ExternalAnnotationMetadata::default(); + for annotation in &ext.annotation { + metadata.raw.push(annotation.to_string()); + collect_external_annotation(annotation, &mut metadata); + } + rumoca_core::ExternalFunction { language: ext.language.clone().unwrap_or_else(|| "C".to_string()), function_name: ext.function_name.as_ref().map(|t| t.text.to_string()), @@ -34,6 +40,159 @@ pub(super) fn convert_external_function( } }) .collect(), + libraries: metadata.libraries, + include_directories: metadata.include_directories, + library_directories: metadata.library_directories, + includes: metadata.includes, + annotation: metadata.raw, + } +} + +#[derive(Debug, Default)] +struct ExternalAnnotationMetadata { + raw: Vec, + libraries: Vec, + include_directories: Vec, + library_directories: Vec, + includes: Vec, +} + +fn collect_external_annotation(expr: &ast::Expression, metadata: &mut ExternalAnnotationMetadata) { + match expr { + ast::Expression::NamedArgument { name, value, .. } => { + collect_external_annotation_value(name.text.as_ref(), value, metadata); + } + ast::Expression::Modification { target, value, .. } => { + if let Some(name) = component_reference_simple_name(target) { + collect_external_annotation_value(name, value, metadata); + } + } + _ => {} + } +} + +fn collect_external_annotation_value( + name: &str, + value: &ast::Expression, + metadata: &mut ExternalAnnotationMetadata, +) { + let values = string_values(value); + match name { + "Library" => extend_unique(&mut metadata.libraries, values), + "IncludeDirectory" => extend_unique(&mut metadata.include_directories, values), + "LibraryDirectory" => extend_unique(&mut metadata.library_directories, values), + "Include" => extend_unique(&mut metadata.includes, values), + _ => {} + } +} + +fn component_reference_simple_name(reference: &ast::ComponentReference) -> Option<&str> { + let [part] = reference.parts.as_slice() else { + return None; + }; + if part.subs.as_ref().is_some_and(|subs| !subs.is_empty()) { + return None; + } + Some(part.ident.text.as_ref()) +} + +fn string_values(expr: &ast::Expression) -> Vec { + match expr { + ast::Expression::Terminal { + terminal_type: ast::TerminalType::String, + token, + .. + } => vec![token.text.trim_matches('"').to_string()], + ast::Expression::Array { elements, .. } => { + elements.iter().flat_map(string_values).collect() + } + _ => Vec::new(), + } +} + +fn extend_unique(target: &mut Vec, values: Vec) { + for value in values { + if !value.is_empty() && !target.contains(&value) { + target.push(value); + } + } +} + +#[cfg(test)] +mod external_annotation_tests { + use super::*; + use std::sync::Arc; + + fn token(text: &str) -> rumoca_core::Token { + rumoca_core::Token { + text: Arc::from(text), + ..Default::default() + } + } + + fn string_expr(value: &str) -> ast::Expression { + ast::Expression::Terminal { + terminal_type: ast::TerminalType::String, + token: token(value), + span: rumoca_core::Span::DUMMY, + } + } + + fn named_arg(name: &str, value: ast::Expression) -> ast::Expression { + ast::Expression::NamedArgument { + name: token(name), + value: Arc::new(value), + span: rumoca_core::Span::DUMMY, + } + } + + #[test] + fn convert_external_function_preserves_library_annotation_metadata() { + let external = rumoca_ir_ast::ExternalFunction { + language: Some("C".to_string()), + function_name: Some(token("initialize_Modelica_EnergyPlus_9_6_0")), + args: vec![], + annotation: vec![ + named_arg( + "Library", + ast::Expression::Array { + elements: vec![ + string_expr("ModelicaBuildingsEnergyPlus_9_6_0"), + string_expr("fmilib_shared"), + ], + is_matrix: false, + span: rumoca_core::Span::DUMMY, + }, + ), + named_arg( + "IncludeDirectory", + string_expr("modelica://Buildings/Resources/C-Sources"), + ), + ], + ..Default::default() + }; + + let converted = convert_external_function(&external, "initialize"); + + assert_eq!(converted.language, "C"); + assert_eq!( + converted.function_name.as_deref(), + Some("initialize_Modelica_EnergyPlus_9_6_0") + ); + assert_eq!( + converted.libraries, + vec!["ModelicaBuildingsEnergyPlus_9_6_0", "fmilib_shared"] + ); + assert_eq!( + converted.include_directories, + vec!["modelica://Buildings/Resources/C-Sources"] + ); + assert!( + converted + .annotation + .iter() + .any(|item| item.contains("Library")) + ); } } @@ -316,6 +475,7 @@ fn ast_subscript_span( /// Convert a component declaration to a function parameter. pub(super) fn convert_component_to_param( + tree: &ast::ClassTree, class_index: &ast::ClassDefIndex<'_>, name: &str, component: &ast::Component, @@ -357,7 +517,15 @@ pub(super) fn convert_component_to_param( .shape_expr .iter() .map(|sub| { - lower_function_shape_subscript(sub, class_index, imports, locals, def_map, span) + lower_function_shape_subscript( + sub, + tree, + class_index, + imports, + locals, + def_map, + span, + ) }) .collect::, FlattenError>>()?; param_dims = shape_expr.iter().map(function_shape_dim).collect(); @@ -379,15 +547,23 @@ pub(super) fn convert_component_to_param( && !matches!(binding_expr, ast::Expression::Empty { .. }) { let qualified = qualify_function_expr(binding_expr, imports, locals); - param = param.with_default(ast_lower::expression_from_ast_with_def_map( + param = param.with_default(ast_lower::expression_from_ast_with_context( &qualified, - Some(def_map), + ast_lower::LoweringContext { + def_map: Some(def_map), + class_tree: Some(tree), + instance_name: None, + }, )?); } else if !matches!(component.start, ast::Expression::Empty { .. }) { let qualified = qualify_function_expr(&component.start, imports, locals); - param = param.with_default(ast_lower::expression_from_ast_with_def_map( + param = param.with_default(ast_lower::expression_from_ast_with_context( &qualified, - Some(def_map), + ast_lower::LoweringContext { + def_map: Some(def_map), + class_tree: Some(tree), + instance_name: None, + }, )?); } } @@ -414,6 +590,7 @@ fn function_shape_dim(subscript: &rumoca_core::Subscript) -> i64 { pub(super) fn lower_function_shape_subscript( subscript: &ast::Subscript, + tree: &ast::ClassTree, class_index: &ast::ClassDefIndex<'_>, imports: &qualify::ImportMap, locals: &HashSet, @@ -428,9 +605,13 @@ pub(super) fn lower_function_shape_subscript( } let qualified = qualify_function_expr(expr, imports, locals); Ok(rumoca_core::Subscript::expr( - Box::new(ast_lower::expression_from_ast_with_def_map( + Box::new(ast_lower::expression_from_ast_with_context( &qualified, - Some(def_map), + ast_lower::LoweringContext { + def_map: Some(def_map), + class_tree: Some(tree), + instance_name: None, + }, )?), span, )) diff --git a/crates/rumoca-phase-flatten/src/functions/tests.rs b/crates/rumoca-phase-flatten/src/functions/tests.rs index a440720e1..00ae2626f 100644 --- a/crates/rumoca-phase-flatten/src/functions/tests.rs +++ b/crates/rumoca-phase-flatten/src/functions/tests.rs @@ -90,6 +90,22 @@ fn ast_comp_ref(parts: &[&str]) -> ast::ComponentReference { } } +fn string_expr(value: &str) -> ast::Expression { + ast::Expression::Terminal { + terminal_type: ast::TerminalType::String, + token: token(value), + span: Span::DUMMY, + } +} + +fn named_arg(name: &str, value: ast::Expression) -> ast::Expression { + ast::Expression::NamedArgument { + name: token(name), + value: Arc::new(value), + span: Span::DUMMY, + } +} + #[test] fn canonicalize_collected_function_calls_uses_unique_suffix_match() { let mut flat = flat::Model::new(); @@ -792,6 +808,7 @@ fn test_convert_component_to_param_prefers_binding_over_start_default() { let tree = ast::ClassTree::new(); let class_index = ast::ClassDefIndex::from_tree(&tree); let param = convert_component_to_param( + &tree, &class_index, "m", &component, @@ -854,6 +871,7 @@ fn test_convert_component_to_param_preserves_mixed_dynamic_rank() { let tree = ast::ClassTree::new(); let class_index = ast::ClassDefIndex::from_tree(&tree); let param = convert_component_to_param( + &tree, &class_index, "den2", &component, @@ -910,6 +928,7 @@ fn test_convert_component_to_param_resolves_constant_shape_expr() { }; let param = convert_component_to_param( + &tree, &class_index, "state", &component, @@ -1003,6 +1022,7 @@ fn test_convert_component_to_param_inherits_type_alias_dims() { let source_map = test_source_map(); let class_index = ast::ClassDefIndex::from_tree(&tree); let param = convert_component_to_param( + &tree, &class_index, "T", &component, @@ -1066,13 +1086,16 @@ fn test_constructor_signature_preserves_local_default_references() { tree.def_map.insert(n_def, "Pkg.C.N".to_string()); let source_map = test_source_map(); let class_index = ast::ClassDefIndex::from_tree(&tree); + let mut member_cache = qualify::MemberDefIdCache::default(); let constructor = convert_constructor_signature( + &tree, &class_index, &class_def, "Pkg.C", &source_map, &tree.def_map, + &mut member_cache, ) .unwrap(); @@ -1087,6 +1110,399 @@ fn test_constructor_signature_preserves_local_default_references() { )); } +#[test] +fn test_external_object_constructor_signature_uses_local_external_constructor() { + let object_def = rumoca_core::DefId::new(11); + let constructor_def = rumoca_core::DefId::new(12); + let idf_def = rumoca_core::DefId::new(13); + let adapter_def = rumoca_core::DefId::new(14); + let constructor = ast::ClassDef { + def_id: Some(constructor_def), + name: token("constructor"), + class_type: rumoca_core::ClassType::Function, + pure: false, + components: ast::AstIndexMap::from_iter([ + ( + "idfName".to_string(), + ast::Component { + name: "idfName".to_string(), + def_id: Some(idf_def), + type_name: ast::Name::from_string("String"), + causality: rumoca_core::Causality::Input(token("input")), + location: test_location(0, 12), + ..ast::Component::empty_with_span(test_span()) + }, + ), + ( + "adapter".to_string(), + ast::Component { + name: "adapter".to_string(), + def_id: Some(adapter_def), + type_name: ast::Name::from_string("Pkg.SpawnExternalObject"), + causality: rumoca_core::Causality::Output(token("output")), + location: test_location(0, 12), + ..ast::Component::empty_with_span(test_span()) + }, + ), + ]), + external: Some(ast::ExternalFunction { + language: Some("C".to_string()), + function_name: Some(token("allocate_Modelica_EnergyPlus_9_6_0")), + output: Some(ast_comp_ref(&["adapter"])), + args: vec![ast::Expression::ComponentReference(ast_comp_ref(&[ + "idfName", + ]))], + annotation: vec![named_arg( + "Library", + ast::Expression::Array { + elements: vec![ + string_expr("ModelicaBuildingsEnergyPlus_9_6_0"), + string_expr("fmilib_shared"), + ], + is_matrix: false, + span: Span::DUMMY, + }, + )], + }), + location: test_location(0, 12), + ..Default::default() + }; + let class_def = ast::ClassDef { + def_id: Some(object_def), + name: token("SpawnExternalObject"), + class_type: rumoca_core::ClassType::Class, + classes: ast::AstIndexMap::from_iter([("constructor".to_string(), constructor)]), + location: test_location(0, 12), + ..Default::default() + }; + let mut tree = ast::ClassTree::default(); + tree.def_map + .insert(object_def, "Pkg.SpawnExternalObject".to_string()); + tree.def_map.insert( + constructor_def, + "Pkg.SpawnExternalObject.constructor".to_string(), + ); + tree.def_map.insert( + idf_def, + "Pkg.SpawnExternalObject.constructor.idfName".to_string(), + ); + tree.def_map.insert( + adapter_def, + "Pkg.SpawnExternalObject.constructor.adapter".to_string(), + ); + let source_map = test_source_map(); + let class_index = ast::ClassDefIndex::from_tree(&tree); + let mut member_cache = qualify::MemberDefIdCache::default(); + + let function = convert_constructor_signature( + &tree, + &class_index, + &class_def, + "Pkg.SpawnExternalObject", + &source_map, + &tree.def_map, + &mut member_cache, + ) + .unwrap(); + + assert_eq!(function.name.as_str(), "Pkg.SpawnExternalObject"); + assert_eq!(function.def_id, Some(object_def)); + assert!(function.is_constructor); + let external = function.external.expect("external constructor metadata"); + assert_eq!( + external.function_name.as_deref(), + Some("allocate_Modelica_EnergyPlus_9_6_0") + ); + assert_eq!(external.output_name.as_deref(), Some("adapter")); + assert_eq!(external.arg_names, vec!["idfName"]); + assert_eq!( + external.libraries, + vec!["ModelicaBuildingsEnergyPlus_9_6_0", "fmilib_shared"] + ); +} + +#[test] +fn test_operator_record_constructor_signature_uses_explicit_constructor_defaults() { + let complex_def = rumoca_core::DefId::new(21); + let operator_def = rumoca_core::DefId::new(22); + let from_real_def = rumoca_core::DefId::new(23); + let re_def = rumoca_core::DefId::new(24); + let im_def = rumoca_core::DefId::new(25); + let result_def = rumoca_core::DefId::new(26); + let record_re_def = rumoca_core::DefId::new(27); + let record_im_def = rumoca_core::DefId::new(28); + let from_real = ast::ClassDef { + def_id: Some(from_real_def), + name: token("fromReal"), + class_type: rumoca_core::ClassType::Function, + components: ast::AstIndexMap::from_iter([ + ( + "re".to_string(), + ast::Component { + name: "re".to_string(), + def_id: Some(re_def), + type_name: ast::Name::from_string("Real"), + causality: rumoca_core::Causality::Input(token("input")), + location: test_location(0, 12), + ..ast::Component::empty_with_span(test_span()) + }, + ), + ( + "im".to_string(), + ast::Component { + name: "im".to_string(), + def_id: Some(im_def), + type_name: ast::Name::from_string("Real"), + causality: rumoca_core::Causality::Input(token("input")), + location: test_location(0, 12), + has_explicit_binding: true, + binding: Some(ast::Expression::Terminal { + terminal_type: ast::TerminalType::UnsignedInteger, + token: token("0"), + span: test_span(), + }), + ..ast::Component::empty_with_span(test_span()) + }, + ), + ( + "result".to_string(), + ast::Component { + name: "result".to_string(), + def_id: Some(result_def), + type_name: ast::Name::from_string("Pkg.Complex"), + causality: rumoca_core::Causality::Output(token("output")), + location: test_location(0, 12), + ..ast::Component::empty_with_span(test_span()) + }, + ), + ]), + location: test_location(0, 12), + ..Default::default() + }; + let constructor_operator = ast::ClassDef { + def_id: Some(operator_def), + name: token("'constructor'"), + class_type: rumoca_core::ClassType::Operator, + classes: ast::AstIndexMap::from_iter([("fromReal".to_string(), from_real)]), + location: test_location(0, 12), + ..Default::default() + }; + let class_def = ast::ClassDef { + def_id: Some(complex_def), + name: token("Complex"), + class_type: rumoca_core::ClassType::Record, + operator_record: true, + components: ast::AstIndexMap::from_iter([ + ( + "re".to_string(), + ast::Component { + name: "re".to_string(), + def_id: Some(record_re_def), + type_name: ast::Name::from_string("Real"), + location: test_location(0, 12), + ..ast::Component::empty_with_span(test_span()) + }, + ), + ( + "im".to_string(), + ast::Component { + name: "im".to_string(), + def_id: Some(record_im_def), + type_name: ast::Name::from_string("Real"), + location: test_location(0, 12), + ..ast::Component::empty_with_span(test_span()) + }, + ), + ]), + classes: ast::AstIndexMap::from_iter([("'constructor'".to_string(), constructor_operator)]), + location: test_location(0, 12), + ..Default::default() + }; + let mut tree = ast::ClassTree::default(); + tree.def_map.insert(complex_def, "Pkg.Complex".to_string()); + tree.def_map + .insert(operator_def, "Pkg.Complex.'constructor'".to_string()); + tree.def_map.insert( + from_real_def, + "Pkg.Complex.'constructor'.fromReal".to_string(), + ); + tree.def_map + .insert(re_def, "Pkg.Complex.'constructor'.fromReal.re".to_string()); + tree.def_map + .insert(im_def, "Pkg.Complex.'constructor'.fromReal.im".to_string()); + tree.def_map.insert( + result_def, + "Pkg.Complex.'constructor'.fromReal.result".to_string(), + ); + tree.def_map + .insert(record_re_def, "Pkg.Complex.re".to_string()); + tree.def_map + .insert(record_im_def, "Pkg.Complex.im".to_string()); + let source_map = test_source_map(); + let class_index = ast::ClassDefIndex::from_tree(&tree); + let mut member_cache = qualify::MemberDefIdCache::default(); + + let function = convert_constructor_signature( + &tree, + &class_index, + &class_def, + "Pkg.Complex", + &source_map, + &tree.def_map, + &mut member_cache, + ) + .unwrap(); + + assert_eq!(function.name.as_str(), "Pkg.Complex"); + assert_eq!(function.def_id, Some(complex_def)); + assert!(function.is_constructor); + assert_eq!(function.inputs.len(), 2); + let re_param = function + .inputs + .iter() + .find(|param| param.name == "re") + .expect("re constructor input"); + assert_eq!(re_param.def_id, Some(record_re_def)); + let im_param = function + .inputs + .iter() + .find(|param| param.name == "im") + .expect("im constructor input"); + assert_eq!(im_param.def_id, Some(record_im_def)); + assert!(matches!( + im_param.default, + Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(0), + .. + }) + )); +} + +#[test] +fn test_operator_record_constructor_signature_falls_back_for_conversion_constructor() { + let complex_def = rumoca_core::DefId::new(121); + let operator_def = rumoca_core::DefId::new(122); + let from_real_def = rumoca_core::DefId::new(123); + let x_def = rumoca_core::DefId::new(124); + let result_def = rumoca_core::DefId::new(125); + let record_re_def = rumoca_core::DefId::new(126); + let record_im_def = rumoca_core::DefId::new(127); + let from_real = ast::ClassDef { + def_id: Some(from_real_def), + name: token("from_real"), + class_type: rumoca_core::ClassType::Function, + components: ast::AstIndexMap::from_iter([ + ( + "x".to_string(), + ast::Component { + name: "x".to_string(), + def_id: Some(x_def), + type_name: ast::Name::from_string("Real"), + causality: rumoca_core::Causality::Input(token("input")), + location: test_location(0, 12), + ..ast::Component::empty_with_span(test_span()) + }, + ), + ( + "c".to_string(), + ast::Component { + name: "c".to_string(), + def_id: Some(result_def), + type_name: ast::Name::from_string("Pkg.Complex"), + causality: rumoca_core::Causality::Output(token("output")), + location: test_location(0, 12), + ..ast::Component::empty_with_span(test_span()) + }, + ), + ]), + location: test_location(0, 12), + ..Default::default() + }; + let constructor_operator = ast::ClassDef { + def_id: Some(operator_def), + name: token("'constructor'"), + class_type: rumoca_core::ClassType::Operator, + classes: ast::AstIndexMap::from_iter([("from_real".to_string(), from_real)]), + location: test_location(0, 12), + ..Default::default() + }; + let class_def = ast::ClassDef { + def_id: Some(complex_def), + name: token("Complex"), + class_type: rumoca_core::ClassType::Record, + operator_record: true, + components: ast::AstIndexMap::from_iter([ + ( + "re".to_string(), + ast::Component { + name: "re".to_string(), + def_id: Some(record_re_def), + type_name: ast::Name::from_string("Real"), + location: test_location(0, 12), + ..ast::Component::empty_with_span(test_span()) + }, + ), + ( + "im".to_string(), + ast::Component { + name: "im".to_string(), + def_id: Some(record_im_def), + type_name: ast::Name::from_string("Real"), + location: test_location(0, 12), + ..ast::Component::empty_with_span(test_span()) + }, + ), + ]), + classes: ast::AstIndexMap::from_iter([("'constructor'".to_string(), constructor_operator)]), + location: test_location(0, 12), + ..Default::default() + }; + let mut tree = ast::ClassTree::default(); + tree.def_map.insert(complex_def, "Pkg.Complex".to_string()); + tree.def_map + .insert(operator_def, "Pkg.Complex.'constructor'".to_string()); + tree.def_map.insert( + from_real_def, + "Pkg.Complex.'constructor'.from_real".to_string(), + ); + tree.def_map + .insert(x_def, "Pkg.Complex.'constructor'.from_real.x".to_string()); + tree.def_map.insert( + result_def, + "Pkg.Complex.'constructor'.from_real.c".to_string(), + ); + tree.def_map + .insert(record_re_def, "Pkg.Complex.re".to_string()); + tree.def_map + .insert(record_im_def, "Pkg.Complex.im".to_string()); + let source_map = test_source_map(); + let class_index = ast::ClassDefIndex::from_tree(&tree); + let mut member_cache = qualify::MemberDefIdCache::default(); + + let function = convert_constructor_signature( + &tree, + &class_index, + &class_def, + "Pkg.Complex", + &source_map, + &tree.def_map, + &mut member_cache, + ) + .unwrap(); + + assert_eq!(function.name.as_str(), "Pkg.Complex"); + assert_eq!(function.def_id, Some(complex_def)); + assert!(function.is_constructor); + assert_eq!( + function + .inputs + .iter() + .map(|input| (input.name.as_str(), input.def_id)) + .collect::>(), + vec![("re", Some(record_re_def)), ("im", Some(record_im_def))] + ); +} + #[test] fn test_function_local_normalization_rewrites_self_qualified_default() { let mut function = rumoca_core::Function::new("Pkg.C", Span::DUMMY); diff --git a/crates/rumoca-phase-flatten/src/lib.rs b/crates/rumoca-phase-flatten/src/lib.rs index 22221388d..61f2ed662 100644 --- a/crates/rumoca-phase-flatten/src/lib.rs +++ b/crates/rumoca-phase-flatten/src/lib.rs @@ -1,3 +1,10 @@ +#![allow( + clippy::excessive_nesting, + clippy::items_after_test_module, + clippy::too_many_arguments, + clippy::too_many_lines +)] + //! Flatten phase for the Rumoca compiler. //! //! This crate implements the flattening pass that converts an ast::InstancedTree to a flat::Model. @@ -107,10 +114,11 @@ use record_constant_arrays::{ try_extract_record_array_constructor_constant, }; use rumoca_eval_flat::phase_constant::{ - ParamEvalContext, build_eval_context, eval_user_func_real, infer_array_dimensions, - infer_array_dimensions_full_with_functions, looks_like_enum_literal_path, - try_eval_flat_expr_boolean_with_context, try_eval_flat_expr_enum, - try_eval_integer_with_context, try_eval_real_with_context, try_infer_better_dims, + EnumCanonicalizer, ParamEvalContext, build_eval_context, eval_user_func_real, + infer_array_dimensions, infer_array_dimensions_full_with_functions, + looks_like_enum_literal_path, try_eval_flat_expr_boolean_with_context, + try_eval_flat_expr_enum_with_canonicalizer, try_eval_integer_with_context, + try_eval_real_with_context, try_infer_better_dims, }; /// Options controlling flatten strictness. @@ -281,6 +289,12 @@ fn dims_are_better(candidate: &[i64], existing: &[i64]) -> bool { }) } +fn same_rank_concrete_dims(candidate: &[i64], existing: &[i64]) -> bool { + candidate.len() == existing.len() + && candidate.iter().all(|dim| *dim >= 0) + && existing.iter().all(|dim| *dim >= 0) +} + fn dimension_specificity(dims: &[i64]) -> usize { dims.iter().filter(|dim| **dim > 0).count() } @@ -341,9 +355,7 @@ pub fn flatten_ref_with_options( ) -> Result { let mut ctx = Context::new(); ctx.materialize_structured_families = options.materialize_structured_families; - if !model_name.is_empty() { - ctx.simulated_root_name = Some(crate::path_utils::leaf_segment(model_name).to_string()); - } + ctx.simulated_root_name = simulated_root_name(tree, overlay, model_name); let class_index = ast::ClassDefIndex::from_tree(tree); ctx.class_def_ids = std::sync::Arc::new(class_index.def_ids().collect()); ctx.target_def_names = tree @@ -365,6 +377,7 @@ pub fn flatten_ref_with_options( tree, &class_index, &ctx.component_members, + ctx.simulated_root_name.as_deref(), )?; let flatten_graph = prepare_context_for_equation_flattening( &mut ctx, @@ -398,6 +411,41 @@ pub fn flatten_ref_with_options( Ok(flat) } +fn simulated_root_name( + tree: &ast::ClassTree, + overlay: &ast::InstanceOverlay, + model_name: &str, +) -> Option { + if !model_name.is_empty() { + return Some(crate::path_utils::leaf_segment(model_name).to_string()); + } + if let Some(root_name) = explicit_root_class_name(tree, overlay) { + return Some(root_name); + } + let root_class = overlay + .classes + .values() + .filter_map(|class_data| class_data.class_def_id.map(|def_id| (class_data, def_id))) + .min_by_key(|(class_data, _)| class_data.qualified_name.parts.len()) + .map(|(_, def_id)| def_id)?; + tree.def_map + .get(&root_class) + .map(|qualified| crate::path_utils::leaf_segment(qualified).to_string()) +} + +fn explicit_root_class_name( + tree: &ast::ClassTree, + overlay: &ast::InstanceOverlay, +) -> Option { + overlay + .classes + .values() + .find(|class_data| class_data.qualified_name.parts.is_empty()) + .and_then(|class_data| class_data.class_def_id) + .and_then(|def_id| tree.def_map.get(&def_id)) + .map(|qualified| crate::path_utils::leaf_segment(qualified).to_string()) +} + fn populate_flat_symbol_ancestry(flat: &mut flat::Model, class_index: &ast::ClassDefIndex<'_>) { let mut def_ids = class_index.symbol_def_ids().collect::>(); def_ids.sort_by_key(|def_id| def_id.index()); @@ -824,6 +872,62 @@ mod nested_class_constant_scope_tests { assert!(scopes.contains("Modelica.Math.Random.Generators.Xorshift128plus")); } + #[test] + fn simulated_root_name_falls_back_to_overlay_root_class() { + let root_def = rumoca_core::DefId::new(7); + let mut tree = ast::ClassTree::new(); + tree.def_map.insert(root_def, "Pkg.Vehicle".to_string()); + let mut overlay = ast::InstanceOverlay::default(); + overlay.classes.insert( + ast::InstanceId::new(1), + ast::ClassInstanceData { + instance_id: ast::InstanceId::new(1), + class_def_id: Some(root_def), + qualified_name: ast::QualifiedName::new(), + ..Default::default() + }, + ); + + assert_eq!( + simulated_root_name(&tree, &overlay, ""), + Some("Vehicle".to_string()) + ); + } + + #[test] + fn simulated_root_name_falls_back_to_shortest_overlay_class_instance() { + let root_def = rumoca_core::DefId::new(7); + let nested_def = rumoca_core::DefId::new(8); + let mut tree = ast::ClassTree::new(); + tree.def_map.insert(root_def, "Pkg.Vehicle".to_string()); + tree.def_map + .insert(nested_def, "Pkg.Vehicle.Controller".to_string()); + let mut overlay = ast::InstanceOverlay::default(); + overlay.classes.insert( + ast::InstanceId::new(1), + ast::ClassInstanceData { + instance_id: ast::InstanceId::new(1), + class_def_id: Some(root_def), + qualified_name: ast::QualifiedName::from_dotted("vehicle"), + ..Default::default() + }, + ); + overlay.classes.insert( + ast::InstanceId::new(2), + ast::ClassInstanceData { + instance_id: ast::InstanceId::new(2), + class_def_id: Some(nested_def), + qualified_name: ast::QualifiedName::from_dotted("vehicle.controller"), + ..Default::default() + }, + ); + + assert_eq!( + simulated_root_name(&tree, &overlay, ""), + Some("Vehicle".to_string()) + ); + } + #[test] fn extract_nested_class_constants_skips_non_package_nested_classes() { let tree = ast::ClassTree::new(); @@ -1122,6 +1226,25 @@ fn inject_enclosing_class_constants( "enclosing class constant injection parent scope", )); }; + let Some(enclosing_class) = class_index.get(parent_def_id) else { + return Err(missing_resolved_class_metadata_for_def_id( + tree, + class_index, + parent_def_id, + model_name, + "enclosing class constant injection class lookup", + )); + }; + for ext in &enclosing_class.extends { + apply_extends_constants_for_scope( + tree, + class_index, + enclosing_name, + ext, + enclosing_name, + ctx, + ); + } let ancestors = collect_ancestor_classes_with_index(tree, class_index, enclosing_name); if ancestors.is_empty() { return Ok(()); diff --git a/crates/rumoca-phase-flatten/src/outer_refs.rs b/crates/rumoca-phase-flatten/src/outer_refs.rs index 3e8f2857e..09b156dab 100644 --- a/crates/rumoca-phase-flatten/src/outer_refs.rs +++ b/crates/rumoca-phase-flatten/src/outer_refs.rs @@ -126,8 +126,77 @@ impl ExpressionRewriter for OuterRefRedirectRewriter<'_> { span: *span, }; } - self.walk_expression(expr) + if let Some(redirected) = redirect_outer_field_access(expr, self.outer_to_inner) { + return redirected; + } + let rewritten = self.walk_expression(expr); + redirect_outer_field_access(&rewritten, self.outer_to_inner).unwrap_or(rewritten) + } +} + +fn redirect_outer_field_access( + expr: &rumoca_core::Expression, + outer_to_inner: &IndexMap, +) -> Option { + if !matches!(expr, rumoca_core::Expression::FieldAccess { .. }) { + return None; + } + let (name, span) = field_access_path(expr)?; + let redirected = redirect_name_string(&name, outer_to_inner)?; + Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new(redirected), + subscripts: Vec::new(), + span, + }) +} + +fn field_access_path(expr: &rumoca_core::Expression) -> Option<(String, rumoca_core::Span)> { + match expr { + rumoca_core::Expression::FieldAccess { base, field, span } => { + let (mut name, _) = field_access_path(base)?; + name.push('.'); + name.push_str(field); + Some((name, *span)) + } + rumoca_core::Expression::Index { + base, + subscripts, + span, + } => { + let (mut name, _) = field_access_path(base)?; + append_literal_subscripts(&mut name, subscripts)?; + Some((name, *span)) + } + rumoca_core::Expression::VarRef { + name, + subscripts, + span, + } => { + let mut rendered = name.as_str().to_string(); + append_literal_subscripts(&mut rendered, subscripts)?; + Some((rendered, *span)) + } + _ => None, + } +} + +fn append_literal_subscripts( + name: &mut String, + subscripts: &[rumoca_core::Subscript], +) -> Option<()> { + for subscript in subscripts { + match subscript { + rumoca_core::Subscript::Index { value, .. } => { + name.push('['); + name.push_str(&value.to_string()); + name.push(']'); + } + rumoca_core::Subscript::Colon { .. } | rumoca_core::Subscript::Expr { .. } => { + return None; + } + } } + Some(()) } /// Redirect outer-prefixed VarRef names in when equations. @@ -370,6 +439,35 @@ mod tests { assert_eq!(name.as_str(), "innerBus.filter"); } + #[test] + fn test_redirect_outer_refs_rewrites_outer_field_access_chain() { + let mut flat = flat::Model::new(); + let x_name = rumoca_core::VarName::new("zoneTol"); + flat.variables.insert( + x_name.clone(), + flat::Variable { + name: x_name.clone(), + start: Some(rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FieldAccess { + base: Box::new(var_ref("floor.zon[1]")), + field: "building".to_string(), + span: test_span(), + }), + field: "relativeSurfaceTolerance".to_string(), + span: test_span(), + }), + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + let mut outer_to_inner = IndexMap::default(); + outer_to_inner.insert("floor.zon[1].building".to_string(), "building".to_string()); + redirect_outer_refs(&mut flat, &outer_to_inner); + + let var = flat.variables.get(&x_name).expect("expected variable"); + assert_var_ref(var.start.as_ref(), "building.relativeSurfaceTolerance"); + } + fn assert_var_ref(expr: Option<&rumoca_core::Expression>, expected: &str) { let Some(rumoca_core::Expression::VarRef { name, .. }) = expr else { panic!("expected var ref"); diff --git a/crates/rumoca-phase-flatten/src/pipeline/component_alias_injection.rs b/crates/rumoca-phase-flatten/src/pipeline/component_alias_injection.rs index e17af0dc8..65b54855c 100644 --- a/crates/rumoca-phase-flatten/src/pipeline/component_alias_injection.rs +++ b/crates/rumoca-phase-flatten/src/pipeline/component_alias_injection.rs @@ -6,8 +6,21 @@ struct ComponentStaticConstantKey { class_def_id: rumoca_core::DefId, } +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +struct AliasPackageInjectionKey { + scope: String, + package_def_id: Option, + package_context: String, +} + +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +struct AliasPackageStaticKey { + package_def_id: rumoca_core::DefId, + package_context: String, +} + enum ComponentStaticConstantCacheEntry { - Cacheable(ScopedConstantDelta), + Cacheable(Box), Uncacheable, } @@ -16,8 +29,10 @@ struct ScopedKeySnapshot { parameter_values: rustc_hash::FxHashSet, real_parameter_values: rustc_hash::FxHashSet, boolean_parameter_values: rustc_hash::FxHashSet, + string_parameter_values: rustc_hash::FxHashSet, enum_parameter_values: rustc_hash::FxHashSet, constant_values: rustc_hash::FxHashSet, + class_constant_keys: rustc_hash::FxHashSet, array_dimensions: rustc_hash::FxHashSet, modified_constant_keys: rustc_hash::FxHashSet, } @@ -28,8 +43,10 @@ impl ScopedKeySnapshot { parameter_values: scoped_keys(&ctx.parameter_values, scope), real_parameter_values: scoped_keys(&ctx.real_parameter_values, scope), boolean_parameter_values: scoped_keys(&ctx.boolean_parameter_values, scope), + string_parameter_values: scoped_keys(&ctx.string_parameter_values, scope), enum_parameter_values: scoped_keys(&ctx.enum_parameter_values, scope), constant_values: scoped_keys(&ctx.constant_values, scope), + class_constant_keys: scoped_set_keys(&ctx.class_constant_keys, scope), array_dimensions: scoped_keys(&ctx.array_dimensions, scope), modified_constant_keys: ctx .modified_constant_keys @@ -46,8 +63,10 @@ struct ScopedConstantDelta { parameter_values: Vec<(String, i64)>, real_parameter_values: Vec<(String, f64)>, boolean_parameter_values: Vec<(String, bool)>, + string_parameter_values: Vec<(String, String)>, enum_parameter_values: Vec<(String, String)>, constant_values: Vec<(String, rumoca_core::Expression)>, + class_constant_keys: Vec, array_dimensions: Vec<(String, Vec)>, modified_constant_keys: Vec, } @@ -70,6 +89,11 @@ impl ScopedConstantDelta { scope, &before.boolean_parameter_values, ), + string_parameter_values: capture_string_map_delta( + &ctx.string_parameter_values, + scope, + &before.string_parameter_values, + ), enum_parameter_values: capture_string_map_delta( &ctx.enum_parameter_values, scope, @@ -80,6 +104,11 @@ impl ScopedConstantDelta { scope, &before.constant_values, )?, + class_constant_keys: capture_set_delta( + &ctx.class_constant_keys, + scope, + &before.class_constant_keys, + ), array_dimensions: capture_map_delta( &ctx.array_dimensions, scope, @@ -116,6 +145,12 @@ impl ScopedConstantDelta { &ctx.flat_parameter_constant_keys, &mut ctx.boolean_parameter_values, ); + replay_map_delta( + &self.string_parameter_values, + scope, + &ctx.flat_parameter_constant_keys, + &mut ctx.string_parameter_values, + ); replay_map_delta( &self.enum_parameter_values, scope, @@ -128,6 +163,11 @@ impl ScopedConstantDelta { &ctx.flat_parameter_constant_keys, &mut ctx.constant_values, ); + replay_set_delta( + &self.class_constant_keys, + scope, + &mut ctx.class_constant_keys, + ); replay_map_delta( &self.array_dimensions, scope, @@ -151,6 +191,16 @@ fn scoped_keys( .collect() } +fn scoped_set_keys( + set: &rustc_hash::FxHashSet, + scope: &str, +) -> rustc_hash::FxHashSet { + set.iter() + .filter(|key| is_scoped_key(key, scope)) + .cloned() + .collect() +} + fn is_scoped_key(key: &str, scope: &str) -> bool { key == scope || key @@ -196,6 +246,17 @@ fn capture_string_map_delta( .collect() } +fn capture_set_delta( + set: &rustc_hash::FxHashSet, + scope: &str, + before: &rustc_hash::FxHashSet, +) -> Vec { + set.iter() + .filter(|key| is_scoped_key(key, scope) && !before.contains(*key)) + .map(|key| scoped_key_suffix(key, scope).to_string()) + .collect() +} + fn capture_expression_map_delta( map: &rustc_hash::FxHashMap, scope: &str, @@ -229,6 +290,12 @@ fn replay_map_delta( } } +fn replay_set_delta(delta: &[String], scope: &str, set: &mut rustc_hash::FxHashSet) { + for suffix in delta { + set.insert(rebase_scoped_key(scope, suffix)); + } +} + fn string_mentions_scope(value: &str, scope: &str) -> bool { is_scoped_key(value, scope) } @@ -293,6 +360,11 @@ pub(crate) fn inject_component_instance_nested_class_constants( ComponentStaticConstantKey, ComponentStaticConstantCacheEntry, > = rustc_hash::FxHashMap::default(); + let mut alias_package_injections = rustc_hash::FxHashSet::default(); + let mut alias_package_cache: rustc_hash::FxHashMap< + AliasPackageStaticKey, + ComponentStaticConstantCacheEntry, + > = rustc_hash::FxHashMap::default(); for _ in 0..MAX_PASSES { let mut cache_rejected = 0usize; let mut uncached = 0usize; @@ -372,6 +444,8 @@ pub(crate) fn inject_component_instance_nested_class_constants( comp_scope, scan_class, scan_context: &scan_context, + alias_package_injections: &mut alias_package_injections, + alias_package_cache: &mut alias_package_cache, ctx, }, ); @@ -466,7 +540,10 @@ fn component_static_cache_key( comp: &rumoca_ir_ast::InstanceData, class_def: &ClassDef, ) -> Option { - if comp.class_overrides.is_empty() && !component_scope_has_array_index(comp) { + if comp.class_overrides.is_empty() + && !component_scope_has_array_index(comp) + && !component_has_instance_specific_static_context(comp) + { class_def .def_id .map(|class_def_id| ComponentStaticConstantKey { class_def_id }) @@ -475,6 +552,19 @@ fn component_static_cache_key( } } +fn component_has_instance_specific_static_context(comp: &rumoca_ir_ast::InstanceData) -> bool { + comp.binding.is_some() + || comp.binding_source.is_some() + || comp.binding_source_scope.is_some() + || comp.binding_from_modification + || !comp.attribute_source_scopes.is_empty() + || comp.start.is_some() + || comp.fixed.is_some() + || comp.min.is_some() + || comp.max.is_some() + || comp.nominal.is_some() +} + fn component_scope_has_array_index(comp: &rumoca_ir_ast::InstanceData) -> bool { comp.qualified_name.parts.iter().any(|(name, subs)| { !subs.is_empty() || rumoca_core::split_trailing_subscript_suffix(name).is_some() @@ -508,7 +598,7 @@ fn cache_component_static_delta( if let Some(delta) = ScopedConstantDelta::capture(request.ctx, request.comp_scope, before) { request.static_cache.insert( cache_key, - ComponentStaticConstantCacheEntry::Cacheable(delta), + ComponentStaticConstantCacheEntry::Cacheable(Box::new(delta)), ); StaticInjectResult { injected: true, @@ -529,6 +619,7 @@ fn component_constant_footprint(ctx: &Context) -> usize { ctx.parameter_values.len() + ctx.array_dimensions.len() + ctx.boolean_parameter_values.len() + + ctx.string_parameter_values.len() + ctx.real_parameter_values.len() + ctx.enum_parameter_values.len() + ctx.constant_values.len() @@ -734,12 +825,20 @@ pub(crate) fn inject_component_enclosing_class_constants( if ancestors.is_empty() { return; } + let shadowed_names = class_index + .get_by_qualified_name(class_context) + .or_else(|| { + class_index.get_by_qualified_name(crate::path_utils::leaf_segment(class_context)) + }) + .map(|class| class_member_names_with_bases(tree, class)) + .unwrap_or_default(); const MAX_PASSES: usize = 5; for _pass in 0..MAX_PASSES { let prev = ctx.parameter_values.len() + ctx.array_dimensions.len() + ctx.boolean_parameter_values.len() + + ctx.string_parameter_values.len() + ctx.real_parameter_values.len() + ctx.enum_parameter_values.len() + ctx.constant_values.len(); @@ -759,19 +858,21 @@ pub(crate) fn inject_component_enclosing_class_constants( ctx, ); } - extract_constants_from_class_with_prefix_and_imports( + extract_constants_from_class_with_prefix_and_imports_shadowed( tree, class_index, comp_scope, ancestor, &resolve_context, ctx, + &shadowed_names, ); } let new = ctx.parameter_values.len() + ctx.array_dimensions.len() + ctx.boolean_parameter_values.len() + + ctx.string_parameter_values.len() + ctx.real_parameter_values.len() + ctx.enum_parameter_values.len() + ctx.constant_values.len(); @@ -781,6 +882,39 @@ pub(crate) fn inject_component_enclosing_class_constants( } } +fn class_member_names_with_bases( + tree: &rumoca_ir_ast::ClassTree, + class: &rumoca_ir_ast::ClassDef, +) -> rustc_hash::FxHashSet { + let mut names = rustc_hash::FxHashSet::default(); + let mut visited = rustc_hash::FxHashSet::default(); + collect_class_member_names_recursive(tree, class, &mut names, &mut visited); + names +} + +fn collect_class_member_names_recursive( + tree: &rumoca_ir_ast::ClassTree, + class: &rumoca_ir_ast::ClassDef, + names: &mut rustc_hash::FxHashSet, + visited: &mut rustc_hash::FxHashSet, +) { + if let Some(def_id) = class.def_id + && !visited.insert(def_id) + { + return; + } + names.extend(class.components.keys().cloned()); + for ext in &class.extends { + let base_class = ext + .base_def_id + .and_then(|def_id| tree.get_class_by_def_id(def_id)) + .or_else(|| tree.get_class_by_qualified_name(&ext.base_name.to_string())); + if let Some(base_class) = base_class { + collect_class_member_names_recursive(tree, base_class, names, visited); + } + } +} + /// Inject alias package constants by matching declared child component types /// against their instantiated specialized types in the overlay. /// @@ -796,6 +930,9 @@ pub(crate) struct SpecializedChildAliasCtx<'a, 'tree> { comp_scope: &'a str, scan_class: &'a ClassDef, scan_context: &'a str, + alias_package_injections: &'a mut rustc_hash::FxHashSet, + alias_package_cache: + &'a mut rustc_hash::FxHashMap, ctx: &'a mut Context, } @@ -849,6 +986,8 @@ pub(crate) fn inject_alias_constants_from_specialized_child_components( child_scope: &child_scope, alias_scope: &alias_scope, package_context: &package_context, + alias_package_injections: &mut *request.alias_package_injections, + alias_package_cache: &mut *request.alias_package_cache, ctx: &mut *request.ctx, }, package_class, @@ -893,6 +1032,8 @@ pub(crate) fn inject_alias_constants_from_specialized_child_components( child_scope: &child_scope, alias_scope: &alias_scope, package_context: &package_context, + alias_package_injections: &mut *request.alias_package_injections, + alias_package_cache: &mut *request.alias_package_cache, ctx: &mut *request.ctx, }, package_class, @@ -907,35 +1048,91 @@ struct AliasPackageConstantCtx<'a, 'tree> { child_scope: &'a str, alias_scope: &'a str, package_context: &'a str, + alias_package_injections: &'a mut rustc_hash::FxHashSet, + alias_package_cache: + &'a mut rustc_hash::FxHashMap, ctx: &'a mut Context, } fn inject_alias_package_constants( - request: AliasPackageConstantCtx<'_, '_>, + mut request: AliasPackageConstantCtx<'_, '_>, package_class: &ClassDef, ) { for scope in [request.alias_scope, request.comp_scope, request.child_scope] { - extract_constants_from_class_with_prefix_and_imports( + if !request + .alias_package_injections + .insert(AliasPackageInjectionKey { + scope: scope.to_string(), + package_def_id: package_class.def_id, + package_context: request.package_context.to_string(), + }) + { + continue; + } + inject_alias_package_constants_for_scope(&mut request, package_class, scope); + } +} + +fn inject_alias_package_constants_for_scope( + request: &mut AliasPackageConstantCtx<'_, '_>, + package_class: &ClassDef, + scope: &str, +) { + let Some(package_def_id) = package_class.def_id else { + return inject_uncached_alias_package_constants(request, package_class, scope); + }; + let cache_key = AliasPackageStaticKey { + package_def_id, + package_context: request.package_context.to_string(), + }; + match request.alias_package_cache.get(&cache_key) { + Some(ComponentStaticConstantCacheEntry::Cacheable(delta)) => { + delta.replay(scope, &mut *request.ctx); + return; + } + Some(ComponentStaticConstantCacheEntry::Uncacheable) => { + return inject_uncached_alias_package_constants(request, package_class, scope); + } + None => {} + } + + let before = ScopedKeySnapshot::capture(request.ctx, scope); + inject_uncached_alias_package_constants(request, package_class, scope); + if let Some(delta) = ScopedConstantDelta::capture(request.ctx, scope, &before) { + request.alias_package_cache.insert( + cache_key, + ComponentStaticConstantCacheEntry::Cacheable(Box::new(delta)), + ); + } else { + request + .alias_package_cache + .insert(cache_key, ComponentStaticConstantCacheEntry::Uncacheable); + } +} + +fn inject_uncached_alias_package_constants( + request: &mut AliasPackageConstantCtx<'_, '_>, + package_class: &ClassDef, + scope: &str, +) { + extract_constants_from_class_with_prefix_and_imports( + request.tree, + request.class_index, + scope, + package_class, + request.package_context, + &mut *request.ctx, + ); + for ext in &package_class.extends { + apply_extends_constants_for_scope( request.tree, request.class_index, scope, - package_class, + ext, request.package_context, &mut *request.ctx, ); } - for ext in &package_class.extends { - for scope in [request.alias_scope, request.comp_scope, request.child_scope] { - apply_extends_constants_for_scope( - request.tree, - request.class_index, - scope, - ext, - request.package_context, - &mut *request.ctx, - ); - } - } } pub(crate) fn split_alias_declared_type(type_name: &str) -> Option<(&str, &str)> { diff --git a/crates/rumoca-phase-flatten/src/pipeline/component_member_scope.rs b/crates/rumoca-phase-flatten/src/pipeline/component_member_scope.rs index be8db7fa9..7f145c702 100644 --- a/crates/rumoca-phase-flatten/src/pipeline/component_member_scope.rs +++ b/crates/rumoca-phase-flatten/src/pipeline/component_member_scope.rs @@ -36,6 +36,92 @@ impl ComponentMemberScopes { .get(&scope.to_component_path()) .is_some_and(|members| members.contains(name)) } + + pub(crate) fn scoped_component_imports( + &self, + expr: &ast::Expression, + scope: &QualifiedName, + imports: &qualify::ImportMap, + ) -> qualify::ImportMap { + let mut scoped_imports = imports.clone(); + let mut roots = indexmap::IndexSet::new(); + collect_expression_component_roots(expr, &mut roots); + for root in roots { + if self.has_member(scope, &root) { + scoped_imports.remove(root.as_str()); + continue; + } + if scoped_imports.contains_key(root.as_str()) { + continue; + } + if let Some(parent_name) = self.nearest_parent_name_with_member(scope, &root) { + let target = if parent_name.parts.is_empty() { + root.clone() + } else { + parent_name.child(&root).to_flat_string() + }; + scoped_imports.insert(root, target); + } + } + scoped_imports + } + + fn nearest_parent_name_with_member( + &self, + scope: &QualifiedName, + root: &str, + ) -> Option { + let mut candidate = parent_qualified_name(scope); + loop { + if self.has_member(&candidate, root) { + return Some(candidate); + } + if candidate.parts.is_empty() { + return None; + } + candidate = parent_qualified_name(&candidate); + } + } +} + +fn parent_qualified_name(scope: &QualifiedName) -> QualifiedName { + if scope.parts.len() <= 1 { + QualifiedName::new() + } else { + QualifiedName { + parts: scope.parts[..scope.parts.len() - 1].to_vec(), + } + } +} + +fn collect_expression_component_roots( + expr: &ast::Expression, + roots: &mut indexmap::IndexSet, +) { + use rumoca_ir_ast::visitor::Visitor; + use std::ops::ControlFlow; + + struct RootCollector<'a> { + roots: &'a mut indexmap::IndexSet, + } + + impl Visitor for RootCollector<'_> { + fn visit_component_reference_ctx( + &mut self, + cr: &ast::ComponentReference, + component_ctx: ast::ComponentReferenceContext, + ) -> ControlFlow<()> { + if matches!(component_ctx, ast::ComponentReferenceContext::Expression) + && let Some(first) = cr.parts.first() + { + self.roots.insert(first.ident.text.to_string()); + } + ast::walk_component_reference_default(self, cr) + } + } + + let mut collector = RootCollector { roots }; + let _ = collector.visit_expression(expr); } impl Context { @@ -64,15 +150,13 @@ pub(super) fn imports_without_instance_member_aliases( ) -> qualify::ImportMap { let mut shadowed = indexmap::IndexSet::new(); collect_instance_member_shadowed_import_aliases(expr, prefix, imports, ctx, &mut shadowed); - if shadowed.is_empty() { - return imports.clone(); - } - - imports + let unshadowed = imports .iter() .filter(|(alias, _)| !shadowed.contains(alias.as_str())) .map(|(alias, target)| (alias.clone(), target.clone())) - .collect() + .collect(); + ctx.component_members + .scoped_component_imports(expr, prefix, &unshadowed) } fn collect_instance_member_shadowed_import_aliases( @@ -126,3 +210,46 @@ fn collect_instance_member_shadowed_import_aliases( }; let _ = collector.visit_expression(expr); } + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::Arc; + + fn cref_expr(name: &str) -> ast::Expression { + ast::Expression::ComponentReference(ast::ComponentReference { + local: false, + parts: vec![ast::ComponentRefPart { + ident: rumoca_core::Token { + text: Arc::from(name), + ..rumoca_core::Token::default() + }, + subs: None, + }], + def_id: None, + span: rumoca_core::Span::DUMMY, + }) + } + + #[test] + fn current_scope_member_shadows_existing_import_alias() { + let mut scopes = ComponentMemberScopes::default(); + scopes.insert_component_member_path(&rumoca_core::ComponentPath::from_flat_path( + "pipe.flowModel.n", + )); + + let mut imports = qualify::ImportMap::default(); + imports.insert("n".to_string(), "pipe.n".to_string()); + + let scoped = scopes.scoped_component_imports( + &cref_expr("n"), + &QualifiedName::from_dotted("pipe.flowModel"), + &imports, + ); + + assert!( + !scoped.contains_key("n"), + "current component member must shadow outer declaration import" + ); + } +} diff --git a/crates/rumoca-phase-flatten/src/pipeline/constant_injection.rs b/crates/rumoca-phase-flatten/src/pipeline/constant_injection.rs index 7ffd15b1a..b59d16d16 100644 --- a/crates/rumoca-phase-flatten/src/pipeline/constant_injection.rs +++ b/crates/rumoca-phase-flatten/src/pipeline/constant_injection.rs @@ -1,3 +1,5 @@ +// SPEC_0021 file-size exception: split plan is to move focused constant +// injection helpers into owned submodules after BOPTEST parity stabilization. use super::*; use crate::record_constant_arrays::try_extract_record_array_constructor_constant; use crate::source_spans::required_location_span; @@ -334,6 +336,7 @@ pub(crate) fn extract_ancestor_constants_multi_pass( let prev = ctx.parameter_values.len() + ctx.array_dimensions.len() + ctx.boolean_parameter_values.len() + + ctx.string_parameter_values.len() + ctx.real_parameter_values.len() + ctx.enum_parameter_values.len() + ctx.constant_values.len(); @@ -354,6 +357,7 @@ pub(crate) fn extract_ancestor_constants_multi_pass( let new = ctx.parameter_values.len() + ctx.array_dimensions.len() + ctx.boolean_parameter_values.len() + + ctx.string_parameter_values.len() + ctx.real_parameter_values.len() + ctx.enum_parameter_values.len() + ctx.constant_values.len(); @@ -395,15 +399,9 @@ pub(crate) fn extract_constants_from_class(class_def: &ClassDef, ctx: &mut Conte ) { continue; } - let binding = - comp.binding - .as_ref() - .or(if !matches!(comp.start, ast::Expression::Empty { .. }) { - Some(&comp.start) - } else { - None - }); - let Some(expr) = binding else { continue }; + let Some(expr) = comp.binding.as_ref() else { + continue; + }; let type_name = comp.type_name.to_string(); if !ctx.constant_values.contains_key(name) && let Some(val) = try_extract_record_array_constructor_constant(expr, ctx, "", name) @@ -486,6 +484,11 @@ pub(crate) fn lookup_with_qualified_scope( if let Some(val) = map.get(&qualified) { return Some(val.clone()); } + if let Some(canonical) = canonicalize_array_indices_to_first(&qualified) + && let Some(val) = map.get(&canonical) + { + return Some(val.clone()); + } current_scope = match current_scope { Some(scope) if !scope.is_empty() => scope.parent(), _ => break, @@ -495,9 +498,49 @@ pub(crate) fn lookup_with_qualified_scope( if let Some(val) = map.get(&bare_name) { return Some(val.clone()); } + if let Some(canonical) = canonicalize_array_indices_to_first(&bare_name) + && let Some(val) = map.get(&canonical) + { + return Some(val.clone()); + } None } +fn canonicalize_array_indices_to_first(path: &str) -> Option { + let mut result = String::with_capacity(path.len()); + let mut chars = path.chars().peekable(); + let mut changed = false; + while let Some(ch) = chars.next() { + if ch != '[' { + result.push(ch); + continue; + } + + let mut content = String::new(); + let mut closed = false; + for inner in chars.by_ref() { + if inner == ']' { + closed = true; + break; + } + content.push(inner); + } + + if closed && content.chars().all(|c| c.is_ascii_digit()) { + result.push_str("[1]"); + changed |= content != "1"; + } else { + result.push('['); + result.push_str(&content); + if closed { + result.push(']'); + } + } + } + + changed.then_some(result) +} + /// Scope-aware constant integer evaluation. pub(crate) fn try_eval_const_integer_with_scope( expr: &ast::Expression, @@ -562,12 +605,11 @@ pub(crate) fn try_eval_const_real_with_scope( } => token.text.as_ref().parse::().ok().map(|v| v as f64), ast::Expression::ComponentReference(cr) => { let scope_path = QualifiedName::from_dotted(scope); - lookup_component_ref_with_scope(cr, &scope_path, &ctx.real_parameter_values).or_else( - || { - lookup_component_ref_with_scope(cr, &scope_path, &ctx.parameter_values) - .map(|v| v as f64) - }, - ) + lookup_component_ref_with_scope(cr, &scope_path, &ctx.parameter_values) + .map(|v| v as f64) + .or_else(|| { + lookup_component_ref_with_scope(cr, &scope_path, &ctx.real_parameter_values) + }) } ast::Expression::Unary { rhs, @@ -753,6 +795,12 @@ fn try_eval_const_field_access_expr( if let Some(value) = lookup_constant_expr_with_scope(&name_text, scope, &ctx.constant_values) { return Some(value.with_span(span)); } + if let Some(value) = lookup_with_qualified_scope(&name, &scope_path, &ctx.parameter_values) { + return Some(rumoca_core::Expression::Literal { + value: Literal::Integer(value), + span, + }); + } if let Some(value) = lookup_with_qualified_scope(&name, &scope_path, &ctx.real_parameter_values) && value.is_finite() { @@ -761,17 +809,19 @@ fn try_eval_const_field_access_expr( span, }); } - if let Some(value) = lookup_with_qualified_scope(&name, &scope_path, &ctx.parameter_values) { + if let Some(value) = + lookup_with_qualified_scope(&name, &scope_path, &ctx.boolean_parameter_values) + { return Some(rumoca_core::Expression::Literal { - value: Literal::Integer(value), + value: Literal::Boolean(value), span, }); } if let Some(value) = - lookup_with_qualified_scope(&name, &scope_path, &ctx.boolean_parameter_values) + lookup_with_qualified_scope(&name, &scope_path, &ctx.string_parameter_values) { return Some(rumoca_core::Expression::Literal { - value: Literal::Boolean(value), + value: Literal::String(value), span, }); } @@ -1024,15 +1074,15 @@ pub(crate) fn try_eval_const_component_ref_expr( if component_ref_has_array_shape(&name, ctx, scope) { return None; } - if let Some(v) = lookup_with_qualified_scope(&name, &scope_path, &ctx.real_parameter_values) { + if let Some(v) = lookup_with_qualified_scope(&name, &scope_path, &ctx.parameter_values) { return Some(rumoca_core::Expression::Literal { - value: Literal::Real(v), + value: Literal::Integer(v), span: owner_span, }); } - if let Some(v) = lookup_with_qualified_scope(&name, &scope_path, &ctx.parameter_values) { + if let Some(v) = lookup_with_qualified_scope(&name, &scope_path, &ctx.real_parameter_values) { return Some(rumoca_core::Expression::Literal { - value: Literal::Integer(v), + value: Literal::Real(v), span: owner_span, }); } @@ -1043,6 +1093,12 @@ pub(crate) fn try_eval_const_component_ref_expr( span: owner_span, }); } + if let Some(v) = lookup_with_qualified_scope(&name, &scope_path, &ctx.string_parameter_values) { + return Some(rumoca_core::Expression::Literal { + value: Literal::String(v), + span: owner_span, + }); + } if let Some(enum_name) = lookup_with_qualified_scope(&name, &scope_path, &ctx.enum_parameter_values) { @@ -1080,15 +1136,15 @@ fn try_eval_resolved_const_ref( if lookup_with_scope(name, "", &ctx.array_dimensions).is_some_and(|dims| !dims.is_empty()) { return None; } - if let Some(v) = lookup_with_scope(name, "", &ctx.real_parameter_values) { + if let Some(v) = lookup_with_scope(name, "", &ctx.parameter_values) { return Some(rumoca_core::Expression::Literal { - value: Literal::Real(v), + value: Literal::Integer(v), span: owner_span, }); } - if let Some(v) = lookup_with_scope(name, "", &ctx.parameter_values) { + if let Some(v) = lookup_with_scope(name, "", &ctx.real_parameter_values) { return Some(rumoca_core::Expression::Literal { - value: Literal::Integer(v), + value: Literal::Real(v), span: owner_span, }); } @@ -1098,6 +1154,12 @@ fn try_eval_resolved_const_ref( span: owner_span, }); } + if let Some(v) = lookup_with_scope(name, "", &ctx.string_parameter_values) { + return Some(rumoca_core::Expression::Literal { + value: Literal::String(v), + span: owner_span, + }); + } lookup_with_scope(name, "", &ctx.enum_parameter_values).map(|enum_name| { rumoca_core::Expression::VarRef { name: rumoca_core::Reference::new(enum_name), @@ -1710,7 +1772,8 @@ pub(crate) fn build_structural_eval_context( let parameter_capacity = ctx.parameter_values.len() + ctx.real_parameter_values.len() - + ctx.boolean_parameter_values.len(); + + ctx.boolean_parameter_values.len() + + ctx.string_parameter_values.len(); let mut eval_ctx = EvalContext::with_capacity(parameter_capacity, 0, ctx.functions.len() * 2); for (name, value) in &ctx.parameter_values { eval_ctx.add_parameter(name.clone(), Value::Integer(*value)); @@ -1721,6 +1784,9 @@ pub(crate) fn build_structural_eval_context( for (name, value) in &ctx.boolean_parameter_values { eval_ctx.add_parameter(name.clone(), Value::Bool(*value)); } + for (name, value) in &ctx.string_parameter_values { + eval_ctx.add_parameter(name.clone(), Value::String(value.clone())); + } for func in ctx.functions.values() { eval_ctx.add_function(func.clone()); } @@ -1897,11 +1963,16 @@ pub(crate) struct Context { pub real_parameter_values: rustc_hash::FxHashMap, /// Boolean parameter values for evaluating if-equation conditions. pub boolean_parameter_values: rustc_hash::FxHashMap, + /// String parameter/constant values. Kept separate from enum values because + /// arbitrary strings may contain dots or hyphens and are not component paths. + pub string_parameter_values: rustc_hash::FxHashMap, /// Enumeration parameter values (name -> qualified enum literal string). pub enum_parameter_values: rustc_hash::FxHashMap, /// General constant expression values (scalars/arrays) extracted from /// class/package constants and redeclare/extends modifications. pub constant_values: rustc_hash::FxHashMap, + /// Names whose source declaration variability is `constant`. + pub(crate) class_constant_keys: rustc_hash::FxHashSet, /// Qualified declaration names keyed by semantic target DefId. pub target_def_names: rustc_hash::FxHashMap, /// Fully qualified constant names explicitly modified by extends clauses. @@ -1912,6 +1983,11 @@ pub(crate) struct Context { pub flat_parameter_constant_keys: rustc_hash::FxHashSet, /// Array dimensions for evaluating size() calls (name -> dims). pub array_dimensions: rustc_hash::FxHashMap>, + /// Source spans for entries in `array_dimensions`. + pub array_dimension_spans: rustc_hash::FxHashMap, + /// Modified bindings whose effective shape was reconciled from structural + /// parameters and must not be replaced by a stale declaration-default binding. + pub(crate) reconciled_modified_dimension_names: rustc_hash::FxHashSet, /// Parameters marked with annotation(Evaluate=true) or declared final (MLS §18.3). /// Only these structural parameters can be used for compile-time branch selection. pub structural_params: std::collections::HashSet, diff --git a/crates/rumoca-phase-flatten/src/pipeline/constant_injection/component_binding_values.rs b/crates/rumoca-phase-flatten/src/pipeline/constant_injection/component_binding_values.rs index 130cf4bad..cc54bce02 100644 --- a/crates/rumoca-phase-flatten/src/pipeline/constant_injection/component_binding_values.rs +++ b/crates/rumoca-phase-flatten/src/pipeline/constant_injection/component_binding_values.rs @@ -2,10 +2,10 @@ use crate::FlattenError; use crate::pipeline::qualify_expression; use rumoca_ir_ast::{InstanceData, InstanceOverlay, QualifiedName}; -/// Evaluate structural component bindings and parameter/constant start values. +/// Evaluate structural component bindings. /// -/// Non-parameter variables' start values are initial conditions, not compile-time -/// constants, and must not be used for structural equation evaluation (MLS 8.6). +/// Start values are initialization attributes, not declaration values, and must +/// not drive structural equation evaluation (MLS 4.4.1, 8.6). pub(crate) fn collect_component_binding_values( overlay: &InstanceOverlay, eval_ctx: &mut rumoca_eval_flat::constant::EvalContext, @@ -16,10 +16,6 @@ pub(crate) fn collect_component_binding_values( } let qualified_name = instance_data.qualified_name.to_flat_string(); - if eval_ctx.get(&qualified_name).is_some() { - continue; - } - if let Some(binding) = &instance_data.binding { let flat_binding = qualify_expression(binding, &QualifiedName::new())?; if let Ok(value) = rumoca_eval_flat::constant::eval_expr(&flat_binding, eval_ctx) { @@ -27,26 +23,13 @@ pub(crate) fn collect_component_binding_values( continue; } } - - if component_start_is_structural(instance_data) - && let Some(start) = &instance_data.start - { - let flat_start = qualify_expression(start, &QualifiedName::new())?; - if let Ok(value) = rumoca_eval_flat::constant::eval_expr(&flat_start, eval_ctx) { - eval_ctx.add_parameter(qualified_name, value); - } - } } Ok(()) } fn component_binding_is_structural(instance_data: &InstanceData) -> bool { - component_start_is_structural(instance_data) || instance_data.is_discrete_type -} - -fn component_start_is_structural(instance_data: &InstanceData) -> bool { matches!( instance_data.variability, rumoca_core::Variability::Parameter(_) | rumoca_core::Variability::Constant(_) - ) + ) || instance_data.is_discrete_type } diff --git a/crates/rumoca-phase-flatten/src/pipeline/constant_injection/lookup_scope_tests.rs b/crates/rumoca-phase-flatten/src/pipeline/constant_injection/lookup_scope_tests.rs index 54f8112c6..993c7daad 100644 --- a/crates/rumoca-phase-flatten/src/pipeline/constant_injection/lookup_scope_tests.rs +++ b/crates/rumoca-phase-flatten/src/pipeline/constant_injection/lookup_scope_tests.rs @@ -63,6 +63,14 @@ fn real_expr(value: &str) -> ast::Expression { } } +fn string_expr(value: &str) -> ast::Expression { + ast::Expression::Terminal { + terminal_type: rumoca_ir_ast::TerminalType::String, + token: token(&format!("\"{value}\"")), + span: test_span(), + } +} + fn comp_ref(path: &str) -> ComponentReference { ComponentReference { local: false, @@ -293,6 +301,32 @@ fn component_binding_collection_uses_only_structural_components() { assert!(eval_ctx.get("M.p").is_some()); } +#[test] +fn modified_structural_binding_overrides_existing_default_value() { + let mut overlay = InstanceOverlay::new(); + overlay.add_component(InstanceData { + instance_id: InstanceId::new(1), + qualified_name: QualifiedName::from_dotted("sum.nu"), + variability: Variability::Parameter(Token::default()), + binding: Some(int_expr(2)), + binding_from_modification: true, + is_discrete_type: true, + ..InstanceData::default() + }); + let mut eval_ctx = rumoca_eval_flat::constant::EvalContext::default(); + eval_ctx.add_parameter( + "sum.nu".to_string(), + rumoca_eval_flat::constant::Value::Integer(0), + ); + + collect_component_binding_values(&overlay, &mut eval_ctx).unwrap(); + + assert_eq!( + eval_ctx.get("sum.nu"), + Some(&rumoca_eval_flat::constant::Value::Integer(2)) + ); +} + #[test] fn const_flat_expr_accepts_enum_literal_component_ref() { let expr = ast::Expression::ComponentReference(comp_ref( @@ -319,6 +353,39 @@ fn const_flat_expr_preserves_array_parameter_refs() { assert_eq!(try_eval_const_flat_expr_with_scope(&expr, &ctx, ""), None); } +#[test] +fn const_flat_expr_preserves_path_like_string_parameter_value() { + let mut ctx = Context::new(); + ctx.string_parameter_values.insert( + "zone.spawnExe".to_string(), + "spawn-0.4.3-7048a72798".to_string(), + ); + let expr = ast::Expression::ComponentReference(comp_ref("spawnExe")); + + assert_eq!( + try_eval_const_flat_expr_with_scope(&expr, &ctx, "zone"), + Some(rumoca_core::Expression::Literal { + value: Literal::String("spawn-0.4.3-7048a72798".to_string()), + span: test_span(), + }) + ); +} + +#[test] +fn const_flat_expr_lowers_string_literal_without_enum_path_reinterpretation() { + assert_eq!( + try_eval_const_flat_expr_with_scope( + &string_expr("spawn-0.4.3-7048a72798"), + &Context::new(), + "" + ), + Some(rumoca_core::Expression::Literal { + value: Literal::String("spawn-0.4.3-7048a72798".to_string()), + span: test_span(), + }) + ); +} + #[test] fn structural_bool_pre_eval_keeps_sample_event_indicator_runtime() { let mut eval_ctx = rumoca_eval_flat::constant::EvalContext::new(); diff --git a/crates/rumoca-phase-flatten/src/pipeline/constant_terminals.rs b/crates/rumoca-phase-flatten/src/pipeline/constant_terminals.rs index 60c5c102e..380e8f680 100644 --- a/crates/rumoca-phase-flatten/src/pipeline/constant_terminals.rs +++ b/crates/rumoca-phase-flatten/src/pipeline/constant_terminals.rs @@ -52,9 +52,17 @@ pub(crate) fn try_eval_const_terminal_expr( token, .. } => Some(rumoca_core::Expression::Literal { - value: Literal::String(token.text.as_ref().to_string()), + value: Literal::String(strip_string_terminal_quotes(&token.text)), span: expr.span(), }), _ => None, } } + +fn strip_string_terminal_quotes(text: &str) -> String { + if text.starts_with('"') && text.ends_with('"') && text.len() >= 2 { + text[1..text.len() - 1].to_string() + } else { + text.to_string() + } +} diff --git a/crates/rumoca-phase-flatten/src/pipeline/context_and_tests.rs b/crates/rumoca-phase-flatten/src/pipeline/context_and_tests.rs index d532f18f4..69d532702 100644 --- a/crates/rumoca-phase-flatten/src/pipeline/context_and_tests.rs +++ b/crates/rumoca-phase-flatten/src/pipeline/context_and_tests.rs @@ -1,6 +1,18 @@ +// SPEC_0021 file-size exception: flatten context still owns parameter lookup, +// symbolic dimension reconciliation, and class-instance flatten entry wiring. +// split plan: move dimension inference/reconciliation helpers into a dedicated +// pipeline::dimensions module after the current redeclare/package-scope merge. use super::enum_dimensions::{enum_type_dimension, infer_enum_range_dimensions}; use super::*; +mod import_shadow; +#[path = "mat_resources.rs"] +mod mat_resources; +mod modified_binding_dimensions; + +use import_shadow::imports_without_shadowed_aliases; +use mat_resources::{read_mat_matrix_size, resolve_modelica_resource_path}; + #[derive(Clone, Copy)] struct ParamBinding<'a> { name: &'a str, @@ -9,6 +21,84 @@ struct ParamBinding<'a> { binding_from_modification: bool, } +#[derive(Clone)] +pub(crate) struct CollectedParamBinding { + name: String, + binding: Expression, + may_be_record_alias: bool, + binding_from_modification: bool, +} + +pub(crate) struct ParameterLookupSession { + params: Vec, + var_bindings: Vec, + dimension_state: DimensionEvaluationState, +} + +#[derive(Default)] +struct DimensionEvaluationState { + evaluations: rustc_hash::FxHashMap, + lookup_inputs: Option, + lookup_generation: u64, + #[cfg(test)] + dimension_evaluation_attempts: rustc_hash::FxHashMap, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +struct DimensionEvaluationRecord { + generation: u64, + resolved: bool, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +struct DimensionLookupInputsSnapshot { + integers: rustc_hash::FxHashMap, + real_bits: rustc_hash::FxHashMap, + booleans: rustc_hash::FxHashMap, + strings: rustc_hash::FxHashMap, + enumerations: rustc_hash::FxHashMap, + dimensions: rustc_hash::FxHashMap>, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +struct DimensionInferenceOutcome { + changed: bool, + resolved: bool, +} + +impl ParameterLookupSession { + #[cfg(test)] + pub(crate) fn dimension_evaluation_attempts(&self, name: &str) -> usize { + self.dimension_state + .dimension_evaluation_attempts + .get(name) + .copied() + .unwrap_or_default() + } +} + +impl CollectedParamBinding { + fn as_view(&self) -> ParamBinding<'_> { + ParamBinding { + name: &self.name, + binding: &self.binding, + may_be_record_alias: self.may_be_record_alias, + binding_from_modification: self.binding_from_modification, + } + } +} + +impl From> for CollectedParamBinding { + fn from(binding: ParamBinding<'_>) -> Self { + Self { + name: binding.name.to_string(), + binding: binding.binding.clone(), + may_be_record_alias: binding.may_be_record_alias, + binding_from_modification: binding.binding_from_modification, + } + } +} + fn insert_record_alias( aliases: &mut rustc_hash::FxHashMap, source_path: rumoca_core::ComponentPath, @@ -23,6 +113,165 @@ fn is_array_literal_binding(binding: &Expression) -> bool { matches!(binding, Expression::Array { .. }) } +fn is_modified_shape_binding(binding: &Expression) -> bool { + is_array_literal_binding(binding) || explicit_slice_binding(binding) +} + +fn is_computed_shape_binding(binding: &Expression) -> bool { + !matches!( + binding, + Expression::VarRef { .. } | Expression::FieldAccess { .. } | Expression::Index { .. } + ) +} + +fn binding_targets_embedded_array_element(binding: &Expression) -> bool { + match binding { + Expression::VarRef { name, .. } => has_embedded_array_subscript_in_parent(name.as_str()), + Expression::FieldAccess { base, .. } | Expression::Index { base, .. } => { + binding_targets_embedded_array_element(base) + } + _ => false, + } +} + +fn explicit_slice_binding(binding: &Expression) -> bool { + match binding { + Expression::VarRef { subscripts, .. } | Expression::Index { subscripts, .. } => subscripts + .iter() + .any(|subscript| matches!(subscript, rumoca_core::Subscript::Colon { .. })), + _ => false, + } +} + +fn dims_expr_has_open_range(dims_expr: &[ast::Subscript]) -> bool { + dims_expr.iter().any(|subscript| { + matches!( + subscript, + ast::Subscript::Range { .. } | ast::Subscript::Empty + ) + }) +} + +fn dims_expr_reads_modifier_parameter( + var_name: &str, + dims_expr: &[ast::Subscript], + flat: &flat::Model, +) -> bool { + dims_expr.iter().any(|subscript| { + let ast::Subscript::Expression(expr) = subscript else { + return false; + }; + expression_reads_modifier_parameter(var_name, expr, flat) + }) +} + +fn expression_reads_modifier_parameter( + var_name: &str, + expr: &ast::Expression, + flat: &flat::Model, +) -> bool { + match expr { + ast::Expression::ComponentReference(comp) => { + component_ref_reads_modifier_parameter(var_name, comp, flat) + } + ast::Expression::Range { + start, step, end, .. + } => { + expression_reads_modifier_parameter(var_name, start, flat) + || step + .as_ref() + .is_some_and(|step| expression_reads_modifier_parameter(var_name, step, flat)) + || expression_reads_modifier_parameter(var_name, end, flat) + } + ast::Expression::Unary { rhs, .. } | ast::Expression::Parenthesized { inner: rhs, .. } => { + expression_reads_modifier_parameter(var_name, rhs, flat) + } + ast::Expression::Binary { lhs, rhs, .. } => { + expression_reads_modifier_parameter(var_name, lhs, flat) + || expression_reads_modifier_parameter(var_name, rhs, flat) + } + ast::Expression::FunctionCall { args, .. } + | ast::Expression::ClassModification { + modifications: args, + .. + } + | ast::Expression::Array { elements: args, .. } + | ast::Expression::Tuple { elements: args, .. } => args + .iter() + .any(|arg| expression_reads_modifier_parameter(var_name, arg, flat)), + ast::Expression::NamedArgument { value, .. } + | ast::Expression::Modification { value, .. } => { + expression_reads_modifier_parameter(var_name, value, flat) + } + ast::Expression::If { + branches, + else_branch, + .. + } => { + branches.iter().any(|(cond, branch)| { + expression_reads_modifier_parameter(var_name, cond, flat) + || expression_reads_modifier_parameter(var_name, branch, flat) + }) || expression_reads_modifier_parameter(var_name, else_branch, flat) + } + ast::Expression::ArrayComprehension { + expr, + indices, + filter, + .. + } => { + expression_reads_modifier_parameter(var_name, expr, flat) + || indices + .iter() + .any(|index| expression_reads_modifier_parameter(var_name, &index.range, flat)) + || filter.as_ref().is_some_and(|filter| { + expression_reads_modifier_parameter(var_name, filter, flat) + }) + } + ast::Expression::ArrayIndex { + base, subscripts, .. + } => { + expression_reads_modifier_parameter(var_name, base, flat) + || subscripts.iter().any(|subscript| { + let ast::Subscript::Expression(expr) = subscript else { + return false; + }; + expression_reads_modifier_parameter(var_name, expr, flat) + }) + } + ast::Expression::FieldAccess { base, .. } => { + expression_reads_modifier_parameter(var_name, base, flat) + } + ast::Expression::Empty { .. } | ast::Expression::Terminal { .. } => false, + } +} + +fn component_ref_reads_modifier_parameter( + var_name: &str, + comp: &ast::ComponentReference, + flat: &flat::Model, +) -> bool { + let rendered = comp.to_string(); + let direct = rumoca_core::VarName::new(&rendered); + if flat + .variables + .get(&direct) + .is_some_and(|var| var.binding_from_modification) + { + return true; + } + if rendered.contains('.') { + return false; + } + let scoped_var = rumoca_core::VarName::new(var_name); + let Some(parent) = scoped_var.enclosing_scope() else { + return false; + }; + let scoped = rumoca_core::VarName::new(format!("{parent}.{rendered}")); + flat.variables + .get(&scoped) + .is_some_and(|var| var.binding_from_modification) +} + impl Context { /// Create a new flatten context. pub(crate) fn new() -> Self { @@ -30,12 +279,16 @@ impl Context { parameter_values: rustc_hash::FxHashMap::default(), real_parameter_values: rustc_hash::FxHashMap::default(), boolean_parameter_values: rustc_hash::FxHashMap::default(), + string_parameter_values: rustc_hash::FxHashMap::default(), enum_parameter_values: rustc_hash::FxHashMap::default(), constant_values: rustc_hash::FxHashMap::default(), + class_constant_keys: rustc_hash::FxHashSet::default(), target_def_names: rustc_hash::FxHashMap::default(), modified_constant_keys: rustc_hash::FxHashSet::default(), flat_parameter_constant_keys: rustc_hash::FxHashSet::default(), array_dimensions: rustc_hash::FxHashMap::default(), + array_dimension_spans: rustc_hash::FxHashMap::default(), + reconciled_modified_dimension_names: rustc_hash::FxHashSet::default(), structural_params: std::collections::HashSet::new(), non_structural_params: std::collections::HashSet::new(), functions: rustc_hash::FxHashMap::default(), @@ -78,18 +331,68 @@ impl Context { /// Also tracks structural parameters (Evaluate=true or final) for safe branch selection. pub(crate) fn build_parameter_lookup(&mut self, flat: &Model, tree: &ClassTree) { let _ = tree; // Used for function evaluation context - self.seed_flat_parameter_constant_keys(flat); let params = self.collect_parameters(flat); - self.supplement_record_aliases(¶ms); self.init_array_dimensions(flat); - let var_bindings = Self::collect_var_bindings(flat); self.infer_dims_from_literals(flat); + let mut dimension_state = DimensionEvaluationState::default(); + self.run_multipass_evaluation(¶ms, &var_bindings, &mut dimension_state); + if self.reconcile_modified_integer_parameter_values(flat) { + self.eval_array_dimensions(&var_bindings, &mut dimension_state); + } + } - // Multi-pass evaluation until fixpoint - self.run_multipass_evaluation(¶ms, &var_bindings); + pub(crate) fn collect_parameter_lookup_session( + &mut self, + flat: &Model, + ) -> ParameterLookupSession { + let params = self + .collect_parameters(flat) + .into_iter() + .map(CollectedParamBinding::from) + .collect(); + let var_bindings = Self::collect_var_bindings(flat) + .into_iter() + .map(CollectedParamBinding::from) + .collect(); + ParameterLookupSession { + params, + var_bindings, + dimension_state: DimensionEvaluationState::default(), + } + } + + pub(crate) fn build_parameter_lookup_with_session( + &mut self, + flat: &Model, + tree: &ClassTree, + session: &mut ParameterLookupSession, + ) { + let _ = tree; // Used for function evaluation context + self.seed_flat_parameter_constant_keys(flat); + let ParameterLookupSession { + params, + var_bindings, + dimension_state, + } = session; + let params = params + .iter() + .map(CollectedParamBinding::as_view) + .collect::>(); + let var_bindings = var_bindings + .iter() + .map(CollectedParamBinding::as_view) + .collect::>(); + self.supplement_record_aliases(¶ms); + self.init_array_dimensions(flat); + self.infer_dims_from_literals(flat); + + self.run_multipass_evaluation(¶ms, &var_bindings, dimension_state); + if self.reconcile_modified_integer_parameter_values(flat) { + self.eval_array_dimensions(&var_bindings, dimension_state); + } } pub(crate) fn recompute_symbolic_component_dimensions( @@ -98,37 +401,120 @@ impl Context { overlay: &InstanceOverlay, tree: &ClassTree, ) -> Result { - let mut changed = false; - for instance_data in overlay.components.values() { - if !instance_data.is_primitive || instance_data.dims_expr.is_empty() { - continue; - } - let var_name = qualified_to_var_name(&instance_data.qualified_name); - let Some(flat_var) = flat.variables.get(&var_name) else { - continue; - }; - let span = instance_source_span(instance_data, tree)?; - let resolved_dims = self.resolve_component_dims_expr( - var_name.as_str(), - &instance_data.dims_expr, - flat_var, - tree, - span, - )?; - let Some(flat_var) = flat.variables.get_mut(&var_name) else { - continue; - }; - if flat_var.dims != resolved_dims { - flat_var.dims.clone_from(&resolved_dims); - changed = true; + let mut flat_changed = false; + let max_passes = overlay.components.len().max(1); + for _ in 0..max_passes { + let mut pass_changed = false; + for instance_data in overlay.components.values() { + if !instance_data.is_primitive || instance_data.dims_expr.is_empty() { + continue; + } + let var_name = qualified_to_var_name(&instance_data.qualified_name); + let Some(flat_var) = flat.variables.get(&var_name) else { + continue; + }; + let span = instance_source_span(instance_data, tree)?; + let binding_dims = self.binding_shape_override_dimensions( + var_name.as_str(), + &instance_data.dims_expr, + flat_var, + tree, + ); + let resolved_from_binding = binding_dims.is_some(); + let computed_array_element_binding = !flat_var.binding_from_modification + && has_embedded_array_subscript_in_parent(var_name.as_str()) + && flat_var.binding.as_ref().is_some_and(|binding| { + is_computed_shape_binding(binding) + || binding_targets_embedded_array_element(binding) + }); + let resolved_dims = if let Some(dims) = binding_dims { + dims + } else { + self.resolve_component_dims_expr( + var_name.as_str(), + &instance_data.dims_expr, + flat_var, + tree, + span, + )? + }; + let declaration_dims_can_override_stale_flat_dims = + !has_embedded_array_subscript_in_parent(var_name.as_str()) + && resolved_dims.iter().all(|dim| *dim >= 0) + && (dims_expr_has_open_range(&instance_data.dims_expr) + || dims_expr_reads_modifier_parameter( + var_name.as_str(), + &instance_data.dims_expr, + flat, + )); + let Some(flat_var) = flat.variables.get_mut(&var_name) else { + continue; + }; + let should_update_dims = dims_are_better(&resolved_dims, &flat_var.dims) + || (resolved_from_binding + && (flat_var.binding_from_modification || computed_array_element_binding) + && same_rank_concrete_dims(&resolved_dims, &flat_var.dims)) + || (!resolved_from_binding && declaration_dims_can_override_stale_flat_dims); + if flat_var.dims != resolved_dims && should_update_dims { + flat_var.dims.clone_from(&resolved_dims); + pass_changed = true; + flat_changed = true; + } + let current_cached_dims = self.array_dimensions.get(var_name.as_str()); + let should_update_cached_dims = current_cached_dims.is_none_or(|current| { + dims_are_better(&resolved_dims, current) + || (resolved_from_binding + && (flat_var.binding_from_modification + || computed_array_element_binding) + && same_rank_concrete_dims(&resolved_dims, current)) + || (!resolved_from_binding && declaration_dims_can_override_stale_flat_dims) + }); + if current_cached_dims != Some(&resolved_dims) && should_update_cached_dims { + self.array_dimensions + .insert(var_name.to_string(), resolved_dims); + pass_changed = true; + } } - if self.array_dimensions.get(var_name.as_str()) != Some(&resolved_dims) { - self.array_dimensions - .insert(var_name.to_string(), resolved_dims); - changed = true; + if !pass_changed { + break; } } - Ok(changed) + Ok(flat_changed) + } + + fn binding_shape_override_dimensions( + &self, + var_name: &str, + dims_expr: &[ast::Subscript], + flat_var: &flat::Variable, + tree: &ClassTree, + ) -> Option> { + let binding = flat_var.binding.as_ref()?; + if flat_var.binding_from_modification + && let Some(effective_dims) = self.array_dimensions.get(var_name) + && flat_var.dims == *effective_dims + && effective_dims.len() == dims_expr.len() + && effective_dims.iter().all(|dim| *dim >= 0) + { + return Some(effective_dims.clone()); + } + let name_has_array_element_parent = has_embedded_array_subscript_in_parent(var_name); + let can_override_shape = if flat_var.binding_from_modification { + !name_has_array_element_parent + || is_modified_shape_binding(binding) + || binding_targets_embedded_array_element(binding) + } else { + name_has_array_element_parent + && (is_computed_shape_binding(binding) + || binding_targets_embedded_array_element(binding)) + }; + if !can_override_shape { + return None; + } + + let binding_dims = self.infer_binding_dimensions(var_name, binding, tree)?; + (binding_dims.len() == dims_expr.len() && binding_dims.iter().all(|dim| *dim >= 0)) + .then_some(binding_dims) } fn resolve_component_dims_expr( @@ -166,7 +552,9 @@ impl Context { .binding .as_ref() .and_then(|binding| self.infer_binding_dimensions(var_name, binding, tree)); - let resolved_dims = best_dims(self.array_dimensions.get(var_name), inferred_dims.as_ref()); + let resolved_dims = inferred_dims + .as_ref() + .or_else(|| self.array_dimensions.get(var_name)); if let Some(dim) = resolved_dims .as_ref() @@ -247,9 +635,26 @@ impl Context { known_enums: &self.enum_parameter_values, array_dims: &self.array_dimensions, functions: &self.functions, + user_func_eval_ctx: None, var_context: Some(var_name), }; - let Some(dim) = try_eval_integer_with_context(&lowered, &eval_ctx) else { + let dim = try_eval_integer_with_context(&lowered, &eval_ctx).or_else(|| { + self.target_def_integer_aliases_for_expr(&lowered) + .and_then(|known_ints| { + let eval_ctx = ParamEvalContext { + known_ints: &known_ints, + known_reals: &self.real_parameter_values, + known_bools: &self.boolean_parameter_values, + known_enums: &self.enum_parameter_values, + array_dims: &self.array_dimensions, + functions: &self.functions, + user_func_eval_ctx: None, + var_context: Some(var_name), + }; + try_eval_integer_with_context(&lowered, &eval_ctx) + }) + }); + let Some(dim) = dim else { return Err(FlattenError::unresolved_component_dimension( var_name, expr.to_string(), @@ -266,6 +671,117 @@ impl Context { Ok(dim) } + fn target_def_integer_aliases_for_expr( + &self, + expr: &Expression, + ) -> Option> { + let mut known_ints = self.parameter_values.clone(); + let mut changed = false; + self.collect_target_def_integer_aliases(expr, &mut known_ints, &mut changed); + changed.then_some(known_ints) + } + + fn collect_target_def_integer_aliases( + &self, + expr: &Expression, + known_ints: &mut rustc_hash::FxHashMap, + changed: &mut bool, + ) { + match expr { + Expression::VarRef { + name, subscripts, .. + } => { + if subscripts.is_empty() + && let Some(target_def_id) = name.target_def_id() + && let Some(target_name) = self.target_def_names.get(&target_def_id) + && let Some(value) = self.lookup_integer_by_declared_target_name(target_name) + && known_ints.insert(name.as_str().to_string(), value) != Some(value) + { + *changed = true; + } + for subscript in subscripts { + self.collect_target_def_integer_aliases_from_subscript( + subscript, known_ints, changed, + ); + } + } + Expression::Binary { lhs, rhs, .. } => { + self.collect_target_def_integer_aliases(lhs, known_ints, changed); + self.collect_target_def_integer_aliases(rhs, known_ints, changed); + } + Expression::Unary { rhs, .. } => { + self.collect_target_def_integer_aliases(rhs, known_ints, changed); + } + Expression::BuiltinCall { args, .. } + | Expression::FunctionCall { args, .. } + | Expression::Array { elements: args, .. } + | Expression::Tuple { elements: args, .. } => { + for arg in args { + self.collect_target_def_integer_aliases(arg, known_ints, changed); + } + } + Expression::If { + branches, + else_branch, + .. + } => { + for (cond, value) in branches { + self.collect_target_def_integer_aliases(cond, known_ints, changed); + self.collect_target_def_integer_aliases(value, known_ints, changed); + } + self.collect_target_def_integer_aliases(else_branch, known_ints, changed); + } + Expression::Range { + start, step, end, .. + } => { + self.collect_target_def_integer_aliases(start, known_ints, changed); + if let Some(step) = step { + self.collect_target_def_integer_aliases(step, known_ints, changed); + } + self.collect_target_def_integer_aliases(end, known_ints, changed); + } + Expression::ArrayComprehension { + expr, + indices, + filter, + .. + } => { + self.collect_target_def_integer_aliases(expr, known_ints, changed); + for index in indices { + self.collect_target_def_integer_aliases(&index.range, known_ints, changed); + } + if let Some(filter) = filter { + self.collect_target_def_integer_aliases(filter, known_ints, changed); + } + } + Expression::Index { + base, subscripts, .. + } => { + self.collect_target_def_integer_aliases(base, known_ints, changed); + for subscript in subscripts { + self.collect_target_def_integer_aliases_from_subscript( + subscript, known_ints, changed, + ); + } + } + Expression::FieldAccess { base, .. } => { + self.collect_target_def_integer_aliases(base, known_ints, changed); + } + Expression::Literal { .. } | Expression::Empty { .. } => {} + } + } + + fn collect_target_def_integer_aliases_from_subscript( + &self, + subscript: &rumoca_core::Subscript, + known_ints: &mut rustc_hash::FxHashMap, + changed: &mut bool, + ) { + if let rumoca_core::Subscript::Expr { expr, .. } = subscript { + self.collect_target_def_integer_aliases(expr, known_ints, changed); + } + } + pub(crate) fn seed_flat_parameter_constant_keys(&mut self, flat: &Model) { self.flat_parameter_constant_keys.extend( flat.variables @@ -316,13 +832,19 @@ impl Context { { self.non_structural_params.insert(name.to_string()); } - let is_fixed_parameter = - matches!(var.variability, rumoca_core::Variability::Parameter(_)) - && var.fixed != Some(false); + if matches!(var.variability, rumoca_core::Variability::Parameter(_)) + && var.binding_from_modification + && !var.evaluate + && !var.is_discrete_type + { + self.non_structural_params.insert(name.to_string()); + } + let is_parameter = + matches!(var.variability, rumoca_core::Variability::Parameter(_)); let may_be_record_alias = !var.is_primitive; if var.evaluate || matches!(var.variability, rumoca_core::Variability::Constant(_)) - || is_fixed_parameter + || (is_parameter && var.is_discrete_type) { self.structural_params.insert(name.to_string()); } @@ -376,7 +898,9 @@ impl Context { continue; } let dims_to_use = try_infer_better_dims(var); - self.array_dimensions.insert(name.to_string(), dims_to_use); + let key = name.to_string(); + self.array_dimensions.insert(key.clone(), dims_to_use); + self.array_dimension_spans.insert(key, var.source_span); } } @@ -409,8 +933,10 @@ impl Context { { #[cfg(feature = "tracing")] tracing::debug!(var = %name, dims = ?inferred_dims, "inferred array dimensions from binding"); - self.array_dimensions - .insert(name.to_string(), inferred_dims); + let key = name.to_string(); + self.array_dimensions.insert(key.clone(), inferred_dims); + self.array_dimension_spans + .insert(key, binding.span().unwrap_or(var.source_span)); } } } @@ -420,17 +946,22 @@ impl Context { &mut self, params: &[ParamBinding<'_>], var_bindings: &[ParamBinding<'_>], + dimension_state: &mut DimensionEvaluationState, ) { const MAX_PASSES: usize = 10; for _pass in 0..MAX_PASSES { let enum_progress = self.eval_enum_param_bindings(params); + let string_progress = self.eval_string_params(params); + let matrix_size_progress = self.eval_read_matrix_size_params(params); let real_progress = self.eval_real_params(params); let int_progress = self.eval_integer_param_bindings(params); let bool_progress = self.eval_boolean_params(params); - let dim_progress = self.eval_array_dimensions(var_bindings); + let dim_progress = self.eval_array_dimensions(var_bindings, dimension_state); let varref_dim_progress = self.propagate_varref_dimensions(var_bindings); let alias_progress = self.propagate_through_aliases(params); if !enum_progress + && !string_progress + && !matrix_size_progress && !real_progress && !int_progress && !bool_progress @@ -475,6 +1006,14 @@ impl Context { progress = true; } + if !self.string_parameter_values.contains_key(*name) + && let Some(val) = self.string_parameter_values.get(&resolved).cloned() + { + self.string_parameter_values + .insert((*name).to_string(), val); + progress = true; + } + // Propagate array dimensions if available. // Skip when the name passes through an expanded array component element, // since alias resolution would point to the parent array's dims. @@ -508,8 +1047,26 @@ impl Context { /// /// Also handles conditional expressions like `table = if cond then A else B` /// by evaluating conditions using known boolean and enum parameters. - fn eval_array_dimensions(&mut self, var_bindings: &[ParamBinding<'_>]) -> bool { + fn eval_array_dimensions( + &mut self, + var_bindings: &[ParamBinding<'_>], + dimension_state: &mut DimensionEvaluationState, + ) -> bool { + let lookup_inputs = self.dimension_lookup_inputs_snapshot(); + if dimension_state.lookup_inputs.as_ref() != Some(&lookup_inputs) { + dimension_state.lookup_inputs = Some(lookup_inputs); + dimension_state.lookup_generation = dimension_state.lookup_generation.wrapping_add(1); + } + let lookup_generation = dimension_state.lookup_generation; + let eval_ctx = build_eval_context( + &self.parameter_values, + &self.real_parameter_values, + &self.boolean_parameter_values, + &self.array_dimensions, + &self.functions, + ); let mut new_dims = false; + let mut changed_dimension_names = std::collections::BTreeSet::new(); for ParamBinding { name, binding, @@ -517,60 +1074,205 @@ impl Context { .. } in var_bindings { - new_dims |= self.try_infer_array_dims(name, binding, *binding_from_modification); + if dimension_state + .evaluations + .get(*name) + .is_some_and(|previous| previous.generation == lookup_generation) + { + continue; + } + #[cfg(test)] + { + *dimension_state + .dimension_evaluation_attempts + .entry((*name).to_string()) + .or_default() += 1; + } + let outcome = + self.try_infer_array_dims(name, binding, *binding_from_modification, &eval_ctx); + new_dims |= outcome.changed; + if outcome.changed { + changed_dimension_names.insert((*name).to_string()); + } + dimension_state.evaluations.insert( + (*name).to_string(), + DimensionEvaluationRecord { + generation: lookup_generation, + resolved: outcome.resolved, + }, + ); + } + if new_dims { + // A producer can make a dimension available after an earlier + // consumer already failed in this sweep. Advance the generation, + // promote only records that actually resolved, and leave failed or + // skipped records stale so the next fixed-point pass retries them. + dimension_state.lookup_generation = dimension_state.lookup_generation.wrapping_add(1); + dimension_state.lookup_inputs = Some(self.dimension_lookup_inputs_snapshot()); + let settled_generation = dimension_state.lookup_generation; + for binding in var_bindings { + let Some(evaluation) = dimension_state.evaluations.get_mut(binding.name) else { + continue; + }; + let is_only_producer = changed_dimension_names.len() == 1 + && changed_dimension_names.contains(binding.name); + if evaluation.resolved + && (is_only_producer + || !self.binding_may_read_dimension_lookup_outputs( + binding.name, + binding.binding, + )) + { + evaluation.generation = settled_generation; + } + } } new_dims } + fn dimension_lookup_inputs_snapshot(&self) -> DimensionLookupInputsSnapshot { + DimensionLookupInputsSnapshot { + integers: self.parameter_values.clone(), + real_bits: self + .real_parameter_values + .iter() + .map(|(name, value)| (name.clone(), value.to_bits())) + .collect(), + booleans: self.boolean_parameter_values.clone(), + strings: self.string_parameter_values.clone(), + enumerations: self.enum_parameter_values.clone(), + dimensions: self.array_dimensions.clone(), + } + } + + fn binding_may_read_dimension_lookup_outputs( + &self, + owner_name: &str, + binding: &Expression, + ) -> bool { + if binding.contains_subexpression(|expr| { + matches!( + expr, + Expression::FunctionCall { .. } + | Expression::Index { .. } + | Expression::Binary { .. } + | Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + .. + } + ) + }) || matches!( + binding, + Expression::VarRef { .. } | Expression::FieldAccess { .. } + ) { + return true; + } + let owner_scope = rumoca_core::ComponentPath::from_flat_path(owner_name) + .parent() + .map(|scope| scope.to_flat_string()) + .unwrap_or_default(); + let mut references = Vec::new(); + binding.collect_var_refs(&mut references); + references.into_iter().any(|reference| { + let candidates = scoped_lookup_candidates(reference.as_str(), &owner_scope); + candidates + .iter() + .any(|candidate| self.array_dimensions.contains_key(candidate)) + || !candidates.iter().any(|candidate| { + self.parameter_values.contains_key(candidate) + || self.real_parameter_values.contains_key(candidate) + || self.boolean_parameter_values.contains_key(candidate) + || self.string_parameter_values.contains_key(candidate) + || self.enum_parameter_values.contains_key(candidate) + }) + }) + } + /// Try to infer array dimensions for a single binding. fn try_infer_array_dims( &mut self, name: &str, binding: &Expression, binding_from_modification: bool, - ) -> bool { + user_func_eval_ctx: &rumoca_eval_flat::constant::EvalContext, + ) -> DimensionInferenceOutcome { // Skip when the variable is inside an expanded array component element. // During array expansion, sub-component modifications (e.g., `L=fill(L1sigma,m)`) // are NOT indexed for each element. So `inductor[1].L` gets the same unindexed // binding as the parent `inductor.L`, which infers to the parent's array dims. // Detect this by checking if any path segment (not the last) has embedded subscripts. - if has_embedded_array_subscript_in_parent(name) - && !(binding_from_modification && is_array_literal_binding(binding)) - { - return false; + let binding_can_drive_shape = if binding_from_modification { + is_modified_shape_binding(binding) + } else { + is_computed_shape_binding(binding) + }; + if has_embedded_array_subscript_in_parent(name) && !binding_can_drive_shape { + return DimensionInferenceOutcome { + changed: false, + resolved: true, + }; } let inferred = infer_array_dimensions_full_with_functions( binding, - &ParamEvalContext::new( - &self.parameter_values, - &self.real_parameter_values, - &self.boolean_parameter_values, - &self.enum_parameter_values, - &self.array_dimensions, - &self.functions, - Some(name), - ), + &ParamEvalContext { + known_ints: &self.parameter_values, + known_reals: &self.real_parameter_values, + known_bools: &self.boolean_parameter_values, + known_enums: &self.enum_parameter_values, + array_dims: &self.array_dimensions, + functions: &self.functions, + user_func_eval_ctx: Some(user_func_eval_ctx), + var_context: Some(name), + }, ); let inferred_dims = match inferred { Some(dims) => dims, - None => return false, + None => { + return DimensionInferenceOutcome { + changed: false, + resolved: false, + }; + } }; + if binding_from_modification + && self.reconciled_modified_dimension_names.contains(name) + && self.array_dimensions.get(name).is_some_and(|existing| { + existing.len() == inferred_dims.len() + && existing.iter().all(|dim| *dim >= 0) + && inferred_dims.iter().all(|dim| *dim >= 0) + }) + { + return DimensionInferenceOutcome { + changed: false, + resolved: true, + }; + } // Check if we should update (MLS §10.1) - let should_update = self - .array_dimensions - .get(name) - .is_none_or(|existing| dims_are_better(&inferred_dims, existing)); + let should_update = self.array_dimensions.get(name).is_none_or(|existing| { + dims_are_better(&inferred_dims, existing) + || (binding_from_modification && same_rank_concrete_dims(&inferred_dims, existing)) + || (!binding_from_modification + && has_embedded_array_subscript_in_parent(name) + && is_computed_shape_binding(binding) + && same_rank_concrete_dims(&inferred_dims, existing)) + }); if should_update { #[cfg(feature = "tracing")] tracing::debug!(var = %name, dims = ?inferred_dims, "inferred array dimensions from builtin"); self.array_dimensions .insert(name.to_string(), inferred_dims); - true + DimensionInferenceOutcome { + changed: true, + resolved: true, + } } else { - false + DimensionInferenceOutcome { + changed: false, + resolved: true, + } } } @@ -584,15 +1286,27 @@ impl Context { fn propagate_varref_dimensions(&mut self, var_bindings: &[ParamBinding<'_>]) -> bool { var_bindings .iter() - .filter_map(|ParamBinding { name, binding, .. }| { - self.try_propagate_varref_dims(name, binding) - }) + .filter_map( + |ParamBinding { + name, + binding, + binding_from_modification, + .. + }| { + self.try_propagate_varref_dims(name, binding, *binding_from_modification) + }, + ) .count() > 0 } /// Try to propagate dimensions from a VarRef binding. - fn try_propagate_varref_dims(&mut self, name: &str, binding: &Expression) -> Option<()> { + fn try_propagate_varref_dims( + &mut self, + name: &str, + binding: &Expression, + binding_from_modification: bool, + ) -> Option<()> { let target_name = match binding { Expression::VarRef { name: target, @@ -603,7 +1317,9 @@ impl Context { }; // Skip when the name passes through an expanded array component element. - if has_embedded_array_subscript_in_parent(name) { + if has_embedded_array_subscript_in_parent(name) + && !has_embedded_array_subscript_in_parent(&target_name) + { return None; } @@ -618,10 +1334,13 @@ impl Context { let target_dims = best_dims(direct_dims, alias_dims)?; // Update if we don't have dims or new dims are better - let should_update = self - .array_dimensions - .get(name) - .is_none_or(|existing| dims_are_better(&target_dims, existing)); + let should_update = self.array_dimensions.get(name).is_none_or(|existing| { + dims_are_better(&target_dims, existing) + || ((binding_from_modification + || (has_embedded_array_subscript_in_parent(name) + && has_embedded_array_subscript_in_parent(&target_name))) + && same_rank_concrete_dims(&target_dims, existing)) + }); if should_update { self.array_dimensions.insert(name.to_string(), target_dims); @@ -720,6 +1439,7 @@ impl Context { known_enums: &self.enum_parameter_values, array_dims: &self.array_dimensions, functions: &self.functions, + user_func_eval_ctx: Some(&eval_ctx), var_context: Some(name), }; if let Some(val) = try_eval_integer_with_context(binding, &int_ctx) { @@ -758,13 +1478,60 @@ impl Context { if !binding_from_modification { return None; } + if let Some(target) = unqualified_varref_name(binding) + && let Some(source_scope) = modifier_source_scope(name) + && let Some(value) = + rumoca_core::EvalLookup::lookup_integer(self, target, source_scope.as_str()) + { + return Some(value); + } + if let Some(value) = self.lookup_modifier_binding_target_integer(binding) { + return Some(value); + } let target = unqualified_varref_name(binding)?; let source_scope = modifier_source_scope(name)?; rumoca_core::EvalLookup::lookup_integer(self, target, source_scope.as_str()) + .or_else(|| self.lookup_modifier_binding_target_integer(binding)) + } + + fn lookup_modifier_binding_target_integer(&self, binding: &Expression) -> Option { + let Expression::VarRef { + name, subscripts, .. + } = binding + else { + return None; + }; + if !subscripts.is_empty() { + return None; + } + let target_def_id = name.target_def_id()?; + let target_name = self.target_def_names.get(&target_def_id)?; + self.lookup_integer_by_declared_target_name(target_name) + } + + fn lookup_integer_by_declared_target_name(&self, target_name: &str) -> Option { + if let Some(value) = self.get_integer_param(target_name) { + return Some(value); + } + if let Some(root) = self.simulated_root_name.as_deref() + && let Some(stripped) = target_name.strip_prefix(root) + && let Some(local_name) = stripped.strip_prefix('.') + && let Some(value) = self.get_integer_param(local_name) + { + return Some(value); + } + None } /// Try to evaluate boolean parameters in one pass. fn eval_boolean_params(&mut self, params: &[ParamBinding<'_>]) -> bool { + let eval_ctx = build_eval_context( + &self.parameter_values, + &self.real_parameter_values, + &self.boolean_parameter_values, + &self.array_dimensions, + &self.functions, + ); let new_vals: Vec<(String, bool)> = params .iter() .filter_map(|ParamBinding { name, binding, .. }| { @@ -775,6 +1542,7 @@ impl Context { known_enums: &self.enum_parameter_values, array_dims: &self.array_dimensions, functions: &self.functions, + user_func_eval_ctx: Some(&eval_ctx), var_context: Some(name), }; try_eval_flat_expr_boolean_with_context(binding, &bool_ctx) @@ -792,8 +1560,132 @@ impl Context { progress } + fn eval_string_params(&mut self, params: &[ParamBinding<'_>]) -> bool { + let new_vals: Vec<(String, String)> = params + .iter() + .filter_map(|ParamBinding { name, binding, .. }| { + self.eval_string_expression(binding, Some(name)) + .map(|value| ((*name).to_string(), value)) + }) + .collect(); + + let mut progress = false; + for (name, value) in new_vals { + if self.string_parameter_values.get(&name) != Some(&value) { + self.string_parameter_values.insert(name, value); + progress = true; + } + } + progress + } + + fn eval_read_matrix_size_params(&mut self, params: &[ParamBinding<'_>]) -> bool { + let mut progress = false; + for ParamBinding { name, binding, .. } in params { + let Some((rows, cols)) = self.eval_read_matrix_size_binding(name, binding) else { + continue; + }; + progress |= self.insert_indexed_integer_value(name, 1, rows); + progress |= self.insert_indexed_integer_value(name, 2, cols); + progress |= self + .array_dimensions + .insert((*name).to_string(), vec![2]) + .is_none_or(|existing| existing != vec![2]); + } + progress + } + + fn insert_indexed_integer_value(&mut self, base: &str, index: i64, value: i64) -> bool { + let key = format!("{base}[{index}]"); + if self.parameter_values.get(&key).copied() == Some(value) { + return false; + } + self.parameter_values.insert(key, value); + true + } + + fn eval_read_matrix_size_binding( + &self, + var_name: &str, + binding: &Expression, + ) -> Option<(i64, i64)> { + let Expression::FunctionCall { name, args, .. } = binding else { + return None; + }; + if !matches!( + name.as_str(), + "readMatrixSize" | "Modelica.Utilities.Streams.readMatrixSize" + ) { + return None; + } + let file_name = args + .first() + .and_then(|arg| self.eval_string_expression(arg, Some(var_name)))?; + let matrix_name = args + .get(1) + .and_then(|arg| self.eval_string_expression(arg, Some(var_name)))?; + read_mat_matrix_size(&file_name, &matrix_name) + } + + fn eval_string_expression( + &self, + expr: &Expression, + var_context: Option<&str>, + ) -> Option { + match expr { + Expression::Literal { + value: rumoca_core::Literal::String(value), + .. + } => Some(value.clone()), + Expression::VarRef { + name, subscripts, .. + } if subscripts.is_empty() => self.resolve_string_varref(name.as_str(), var_context), + Expression::FunctionCall { name, args, .. } + if matches!( + name.as_str(), + "loadResource" | "Modelica.Utilities.Files.loadResource" + ) => + { + let raw = args + .first() + .and_then(|arg| self.eval_string_expression(arg, var_context))?; + Some( + resolve_modelica_resource_path(&raw) + .map(|path| path.to_string_lossy().into_owned()) + .unwrap_or(raw), + ) + } + _ => None, + } + } + + fn resolve_string_varref(&self, name: &str, var_context: Option<&str>) -> Option { + if let Some(var_context) = var_context { + let scope = rumoca_core::ComponentPath::from_flat_path(var_context) + .parent() + .map(|path| path.to_flat_string()) + .unwrap_or_default(); + for candidate in scoped_lookup_candidates(name, &scope) { + if let Some(value) = self.string_parameter_values.get(&candidate) { + return Some(value.clone()); + } + } + } + self.string_parameter_values + .get(name) + .cloned() + .or_else(|| lookup_unique_suffix_string(name, &self.string_parameter_values)) + } + /// Try to evaluate real parameters in one pass. fn eval_real_params(&mut self, params: &[ParamBinding<'_>]) -> bool { + let eval_ctx = build_eval_context( + &self.parameter_values, + &self.real_parameter_values, + &self.boolean_parameter_values, + &self.array_dimensions, + &self.functions, + ); let new_vals: Vec<(String, f64)> = params .iter() .filter_map(|ParamBinding { name, binding, .. }| { @@ -804,13 +1696,14 @@ impl Context { known_enums: &self.enum_parameter_values, array_dims: &self.array_dimensions, functions: &self.functions, + user_func_eval_ctx: Some(&eval_ctx), var_context: Some(name), }; if let Some(val) = try_eval_real_with_context(binding, &real_ctx) { return Some(((*name).to_string(), val)); } // Try user-defined function evaluation for function call bindings - self.try_eval_real_func_call(name, binding) + self.try_eval_real_func_call(name, binding, &eval_ctx) .map(|val| ((*name).to_string(), val)) }) .collect(); @@ -831,7 +1724,12 @@ impl Context { } /// Try evaluating a function call binding as a real value. - fn try_eval_real_func_call(&self, name: &str, binding: &Expression) -> Option { + fn try_eval_real_func_call( + &self, + name: &str, + binding: &Expression, + user_func_eval_ctx: &rumoca_eval_flat::constant::EvalContext, + ) -> Option { let Expression::FunctionCall { name: func_name, args, @@ -847,6 +1745,7 @@ impl Context { known_enums: &self.enum_parameter_values, array_dims: &self.array_dimensions, functions: &self.functions, + user_func_eval_ctx: Some(user_func_eval_ctx), var_context: Some(name), }; eval_user_func_real(func_name, args, &int_ctx) @@ -877,7 +1776,8 @@ impl Context { let mut progress = false; loop { - let new_vals = self.collect_enum_values(params, ¶m_names); + let canonicalizer = EnumCanonicalizer::new(&self.enum_parameter_values); + let new_vals = self.collect_enum_values(params, ¶m_names, &canonicalizer); if new_vals.is_empty() { break; } @@ -898,11 +1798,12 @@ impl Context { &self, params: &[ParamBinding<'_>], param_names: &rustc_hash::FxHashSet<&str>, + canonicalizer: &EnumCanonicalizer, ) -> Vec<(String, String)> { params .iter() .filter_map(|ParamBinding { name, binding, .. }| { - self.resolve_enum_binding_value(binding, param_names) + self.resolve_enum_binding_value(binding, param_names, canonicalizer) .map(|enum_val| ((*name).to_string(), enum_val)) }) .collect() @@ -927,20 +1828,26 @@ impl Context { &self, binding: &Expression, param_names: &rustc_hash::FxHashSet<&str>, + canonicalizer: &EnumCanonicalizer, ) -> Option { - let enum_val = self.try_eval_enum_binding(binding)?; + let enum_val = self.try_eval_enum_binding(binding, canonicalizer)?; if !self.enum_reference_matches_parameter(&enum_val, param_names) { return Some(enum_val); } self.resolve_non_parameter_enum_varref(binding, param_names) } - fn try_eval_enum_binding(&self, binding: &Expression) -> Option { - try_eval_flat_expr_enum( + fn try_eval_enum_binding( + &self, + binding: &Expression, + canonicalizer: &EnumCanonicalizer, + ) -> Option { + try_eval_flat_expr_enum_with_canonicalizer( binding, &self.parameter_values, &self.boolean_parameter_values, &self.enum_parameter_values, + canonicalizer, ) .or_else(|| self.resolve_varref_enum_reference(binding)) } @@ -1083,7 +1990,7 @@ impl Context { /// - `stack.cell.stackData.cellData` -> `stack.stackData.cellData` /// /// Returns the original name if no alias applies. - fn resolve_alias(&self, name: &str) -> String { + pub(super) fn resolve_alias(&self, name: &str) -> String { const MAX_DEPTH: usize = 10; // Prevent infinite loops let mut current = rumoca_core::ComponentPath::from_flat_path(name); for _iteration in 0..MAX_DEPTH { @@ -1195,119 +2102,6 @@ impl Context { } } -pub(crate) fn scoped_lookup_candidates(name: &str, scope: &str) -> Vec { - scoped_lookup_candidates_with_scope(name, scope) - .into_iter() - .map(|(candidate, _candidate_scope)| candidate) - .collect() -} - -pub(crate) fn scoped_lookup_candidates_with_scope( - name: &str, - scope: &str, -) -> Vec<(String, String)> { - let name_path = rumoca_core::ComponentPath::from_flat_path(name); - let mut candidates = Vec::new(); - let mut current_scope = Some(rumoca_core::ComponentPath::from_flat_path(scope)); - while let Some(scope_path) = current_scope { - candidates.push(( - scope_path.join(&name_path).to_flat_string(), - scope_path.to_flat_string(), - )); - current_scope = scope_path.parent(); - } - if !scope.is_empty() { - candidates.push((name_path.to_flat_string(), String::new())); - } - candidates -} - -impl rumoca_core::EvalLookup for Context { - fn lookup_integer(&self, name: &str, scope: &str) -> Option { - for candidate in scoped_lookup_candidates(name, scope) { - if let Some(value) = self.get_integer_param(&candidate) { - return Some(value); - } - } - - if crate::path_utils::is_nested_name(name) { - if let Some(value) = lookup_with_scope(name, scope, &self.parameter_values) { - return Some(value); - } - if let Some(value) = lookup_with_scope(name, scope, &self.real_parameter_values) - && value.is_finite() - && value.fract() == 0.0 - { - return Some(value as i64); - } - } - None - } - - fn lookup_real(&self, name: &str, scope: &str) -> Option { - for candidate in scoped_lookup_candidates(name, scope) { - if let Some(value) = self.real_parameter_values.get(&candidate).copied() { - return Some(value); - } - - let resolved = self.resolve_alias(&candidate); - if resolved != candidate - && let Some(value) = self.real_parameter_values.get(&resolved).copied() - { - return Some(value); - } - - if let Some(value) = self.get_integer_param(&candidate) { - return Some(value as f64); - } - } - - if crate::path_utils::is_nested_name(name) { - if let Some(value) = lookup_with_scope(name, scope, &self.real_parameter_values) { - return Some(value); - } - if let Some(value) = lookup_with_scope(name, scope, &self.parameter_values) { - return Some(value as f64); - } - } - None - } - - fn lookup_boolean(&self, name: &str, scope: &str) -> Option { - for candidate in scoped_lookup_candidates(name, scope) { - if let Some(value) = self.get_boolean_param(&candidate) { - return Some(value); - } - } - - if crate::path_utils::is_nested_name(name) { - return lookup_with_scope(name, scope, &self.boolean_parameter_values); - } - None - } - - fn lookup_enum<'a>(&'a self, name: &str, scope: &str) -> Option> { - for candidate in scoped_lookup_candidates(name, scope) { - if let Some(value) = self.enum_parameter_values.get(&candidate) { - return Some(std::borrow::Cow::Borrowed(value.as_str())); - } - - let resolved = self.resolve_alias(&candidate); - if resolved != candidate - && let Some(value) = self.enum_parameter_values.get(&resolved) - { - return Some(std::borrow::Cow::Borrowed(value.as_str())); - } - } - - if crate::path_utils::is_nested_name(name) { - return lookup_with_scope(name, scope, &self.enum_parameter_values) - .map(std::borrow::Cow::Owned); - } - None - } -} - fn unqualified_varref_name(expr: &Expression) -> Option<&str> { let Expression::VarRef { name, subscripts, .. @@ -1318,6 +2112,10 @@ fn unqualified_varref_name(expr: &Expression) -> Option<&str> { if !subscripts.is_empty() { return None; } + let parts = name.parts(); + if parts.len() == 1 && parts[0].subs.is_empty() { + return Some(parts[0].ident.as_str()); + } let path = rumoca_core::ComponentPath::from_flat_path(name.as_str()); (path.len() == 1).then_some(name.as_str()) } @@ -1329,6 +2127,23 @@ fn modifier_source_scope(name: &str) -> Option { Some(source_scope.to_flat_string()) } +fn lookup_unique_suffix_string( + name: &str, + values: &rustc_hash::FxHashMap, +) -> Option { + let mut found = None; + for suffix in rumoca_core::ComponentPath::from_flat_path(name).suffixes_excluding_self() { + let candidate = suffix.to_flat_string(); + if let Some(value) = values.get(&candidate) { + if found.is_some() { + return None; + } + found = Some(value.clone()); + } + } + found +} + impl Default for Context { fn default() -> Self { Self::new() @@ -1402,12 +2217,14 @@ fn process_class_instance_body( // Handle when-equations separately (pass context for parameter evaluation). let mut clauses = when_equations::flatten_when_equation(ctx, &inst_eq, prefix, def_map)?; for clause in &mut clauses { - rewrite_function_overrides_in_when_clause( + rewrite_function_overrides_in_when_clause_scoped( clause, tree, class_index, &override_packages, &override_functions, + &class_scope, + &ctx.component_members, ); } flat.when_clauses.extend(clauses); @@ -1421,6 +2238,8 @@ fn process_class_instance_body( class_index, &override_packages, &override_functions, + &class_scope, + &ctx.component_members, ); let equation_base = flat.equations.len(); for eq in flattened.equations { @@ -1473,6 +2292,8 @@ fn process_class_instance_body( class_index, &override_packages, &override_functions, + &class_scope, + &ctx.component_members, ); let equation_base = flat.initial_equations.len(); for eq in flattened.equations { @@ -1516,6 +2337,7 @@ fn process_class_instance_body( prefix, imports, def_map, + tree, &tree.source_map, instance_name.as_deref(), )?; @@ -1553,6 +2375,7 @@ fn process_class_instance_body( prefix, imports, def_map, + tree, &tree.source_map, instance_name.as_deref(), )?; @@ -1578,6 +2401,7 @@ pub(crate) fn flatten_algorithm_section( prefix: &QualifiedName, imports: &qualify::ImportMap, def_map: Option<&crate::ResolveDefMap>, + tree: &rumoca_ir_ast::ClassTree, source_map: &rumoca_core::SourceMap, instance_name: Option<&str>, ) -> Result { @@ -1605,6 +2429,7 @@ pub(crate) fn flatten_algorithm_section( prefix, imports, def_map, + class_tree: Some(tree), initial_locals: &no_locals, source_map: Some(source_map), instance_name, @@ -1622,6 +2447,7 @@ use super::function_overrides_and_dims::*; pub(crate) struct ComponentInstanceProcess<'a, 'tree> { pub(crate) flat: &'a mut Model, pub(crate) instance_data: &'a rumoca_ir_ast::InstanceData, + pub(crate) simulated_root_name: Option<&'a str>, pub(crate) component_override_map: &'a ComponentOverrideMap, pub(crate) tree: &'a rumoca_ir_ast::ClassTree, pub(crate) class_index: &'a rumoca_ir_ast::ClassDefIndex<'tree>, @@ -1658,6 +2484,8 @@ pub(crate) fn process_component_instance( request.tree, request.class_index, &import_context, + request.component_members, + request.simulated_root_name, )?; assign_instance_identity_to_flat_variable( request.flat, @@ -1827,165 +2655,11 @@ fn qualify_expression_with_effective_imports( &qualified, crate::ast_lower::LoweringContext { def_map, + class_tree: None, instance_name, }, ) } -fn imports_without_shadowed_aliases( - expr: &ast::Expression, - imports: &qualify::ImportMap, - def_map: &crate::ResolveDefMap, -) -> qualify::ImportMap { - let mut shadowed = std::collections::HashSet::new(); - collect_shadowed_import_aliases(expr, imports, def_map, &mut shadowed); - if shadowed.is_empty() { - return imports.clone(); - } - - imports - .iter() - .filter(|(alias, _)| !shadowed.contains(alias.as_str())) - .map(|(alias, target)| (alias.clone(), target.clone())) - .collect() -} - -fn collect_shadowed_import_aliases( - expr: &ast::Expression, - imports: &qualify::ImportMap, - def_map: &crate::ResolveDefMap, - shadowed: &mut std::collections::HashSet, -) { - match expr { - ast::Expression::ComponentReference(cr) => { - collect_component_shadowed_import_alias(cr, imports, def_map, shadowed); - } - ast::Expression::Binary { lhs, rhs, .. } => { - collect_shadowed_import_aliases(lhs, imports, def_map, shadowed); - collect_shadowed_import_aliases(rhs, imports, def_map, shadowed); - } - ast::Expression::Unary { rhs, .. } | ast::Expression::Parenthesized { inner: rhs, .. } => { - collect_shadowed_import_aliases(rhs, imports, def_map, shadowed); - } - ast::Expression::FunctionCall { comp, args, .. } => { - collect_component_shadowed_import_alias(comp, imports, def_map, shadowed); - for arg in args { - collect_shadowed_import_aliases(arg, imports, def_map, shadowed); - } - } - ast::Expression::ClassModification { - target, - modifications, - .. - } => { - collect_component_shadowed_import_alias(target, imports, def_map, shadowed); - for modification in modifications { - collect_shadowed_import_aliases(modification, imports, def_map, shadowed); - } - } - ast::Expression::NamedArgument { value, .. } => { - collect_shadowed_import_aliases(value, imports, def_map, shadowed); - } - ast::Expression::Modification { target, value, .. } => { - collect_component_shadowed_import_alias(target, imports, def_map, shadowed); - collect_shadowed_import_aliases(value, imports, def_map, shadowed); - } - ast::Expression::If { - branches, - else_branch, - .. - } => { - for (condition, value) in branches { - collect_shadowed_import_aliases(condition, imports, def_map, shadowed); - collect_shadowed_import_aliases(value, imports, def_map, shadowed); - } - collect_shadowed_import_aliases(else_branch, imports, def_map, shadowed); - } - ast::Expression::Array { elements, .. } | ast::Expression::Tuple { elements, .. } => { - for element in elements { - collect_shadowed_import_aliases(element, imports, def_map, shadowed); - } - } - ast::Expression::Range { - start, step, end, .. - } => { - collect_shadowed_import_aliases(start, imports, def_map, shadowed); - if let Some(step) = step { - collect_shadowed_import_aliases(step, imports, def_map, shadowed); - } - collect_shadowed_import_aliases(end, imports, def_map, shadowed); - } - ast::Expression::ArrayComprehension { - expr, - indices, - filter, - .. - } => { - collect_shadowed_import_aliases(expr, imports, def_map, shadowed); - for index in indices { - collect_shadowed_import_aliases(&index.range, imports, def_map, shadowed); - } - if let Some(filter) = filter { - collect_shadowed_import_aliases(filter, imports, def_map, shadowed); - } - } - ast::Expression::ArrayIndex { - base, subscripts, .. - } => { - collect_shadowed_import_aliases(base, imports, def_map, shadowed); - for subscript in subscripts { - collect_subscript_shadowed_import_aliases(subscript, imports, def_map, shadowed); - } - } - ast::Expression::FieldAccess { base, .. } => { - collect_shadowed_import_aliases(base, imports, def_map, shadowed); - } - ast::Expression::Terminal { .. } | ast::Expression::Empty { .. } => {} - } -} - -fn collect_component_shadowed_import_alias( - cr: &ast::ComponentReference, - imports: &qualify::ImportMap, - def_map: &crate::ResolveDefMap, - shadowed: &mut std::collections::HashSet, -) { - for part in &cr.parts { - if let Some(subscripts) = &part.subs { - for subscript in subscripts { - collect_subscript_shadowed_import_aliases(subscript, imports, def_map, shadowed); - } - } - } - - let Some(first) = cr.parts.first() else { - return; - }; - let alias = first.ident.text.as_ref(); - let Some(imported_path) = imports.get(alias) else { - return; - }; - let Some(resolved_path) = cr.def_id.and_then(|def_id| def_map.get(&def_id)) else { - return; - }; - if rumoca_core::top_level_last_segment(resolved_path) != alias { - return; - } - if resolved_path != imported_path { - shadowed.insert(alias.to_string()); - } -} - -fn collect_subscript_shadowed_import_aliases( - subscript: &ast::Subscript, - imports: &qualify::ImportMap, - def_map: &crate::ResolveDefMap, - shadowed: &mut std::collections::HashSet, -) { - if let ast::Subscript::Expression(expr) = subscript { - collect_shadowed_import_aliases(expr, imports, def_map, shadowed); - } -} - #[cfg(test)] mod import_shadow_tests; diff --git a/crates/rumoca-phase-flatten/src/pipeline/context_and_tests/import_shadow.rs b/crates/rumoca-phase-flatten/src/pipeline/context_and_tests/import_shadow.rs new file mode 100644 index 000000000..477922a44 --- /dev/null +++ b/crates/rumoca-phase-flatten/src/pipeline/context_and_tests/import_shadow.rs @@ -0,0 +1,157 @@ +use crate::qualify; +use rumoca_ir_ast as ast; + +pub(super) fn imports_without_shadowed_aliases( + expr: &ast::Expression, + imports: &qualify::ImportMap, + def_map: &crate::ResolveDefMap, +) -> qualify::ImportMap { + let mut shadowed = std::collections::HashSet::new(); + collect_shadowed_import_aliases(expr, imports, def_map, &mut shadowed); + if shadowed.is_empty() { + return imports.clone(); + } + + imports + .iter() + .filter(|(alias, _)| !shadowed.contains(alias.as_str())) + .map(|(alias, target)| (alias.clone(), target.clone())) + .collect() +} + +fn collect_shadowed_import_aliases( + expr: &ast::Expression, + imports: &qualify::ImportMap, + def_map: &crate::ResolveDefMap, + shadowed: &mut std::collections::HashSet, +) { + match expr { + ast::Expression::ComponentReference(cr) => { + collect_component_shadowed_import_alias(cr, imports, def_map, shadowed); + } + ast::Expression::Binary { lhs, rhs, .. } => { + collect_shadowed_import_aliases(lhs, imports, def_map, shadowed); + collect_shadowed_import_aliases(rhs, imports, def_map, shadowed); + } + ast::Expression::Unary { rhs, .. } | ast::Expression::Parenthesized { inner: rhs, .. } => { + collect_shadowed_import_aliases(rhs, imports, def_map, shadowed); + } + ast::Expression::FunctionCall { comp, args, .. } => { + collect_component_shadowed_import_alias(comp, imports, def_map, shadowed); + for arg in args { + collect_shadowed_import_aliases(arg, imports, def_map, shadowed); + } + } + ast::Expression::ClassModification { + target, + modifications, + .. + } => { + collect_component_shadowed_import_alias(target, imports, def_map, shadowed); + for modification in modifications { + collect_shadowed_import_aliases(modification, imports, def_map, shadowed); + } + } + ast::Expression::NamedArgument { value, .. } => { + collect_shadowed_import_aliases(value, imports, def_map, shadowed); + } + ast::Expression::Modification { target, value, .. } => { + collect_component_shadowed_import_alias(target, imports, def_map, shadowed); + collect_shadowed_import_aliases(value, imports, def_map, shadowed); + } + ast::Expression::If { + branches, + else_branch, + .. + } => { + for (condition, value) in branches { + collect_shadowed_import_aliases(condition, imports, def_map, shadowed); + collect_shadowed_import_aliases(value, imports, def_map, shadowed); + } + collect_shadowed_import_aliases(else_branch, imports, def_map, shadowed); + } + ast::Expression::Array { elements, .. } | ast::Expression::Tuple { elements, .. } => { + for element in elements { + collect_shadowed_import_aliases(element, imports, def_map, shadowed); + } + } + ast::Expression::Range { + start, step, end, .. + } => { + collect_shadowed_import_aliases(start, imports, def_map, shadowed); + if let Some(step) = step { + collect_shadowed_import_aliases(step, imports, def_map, shadowed); + } + collect_shadowed_import_aliases(end, imports, def_map, shadowed); + } + ast::Expression::ArrayComprehension { + expr, + indices, + filter, + .. + } => { + collect_shadowed_import_aliases(expr, imports, def_map, shadowed); + for index in indices { + collect_shadowed_import_aliases(&index.range, imports, def_map, shadowed); + } + if let Some(filter) = filter { + collect_shadowed_import_aliases(filter, imports, def_map, shadowed); + } + } + ast::Expression::ArrayIndex { + base, subscripts, .. + } => { + collect_shadowed_import_aliases(base, imports, def_map, shadowed); + for subscript in subscripts { + collect_subscript_shadowed_import_aliases(subscript, imports, def_map, shadowed); + } + } + ast::Expression::FieldAccess { base, .. } => { + collect_shadowed_import_aliases(base, imports, def_map, shadowed); + } + ast::Expression::Terminal { .. } | ast::Expression::Empty { .. } => {} + } +} + +fn collect_component_shadowed_import_alias( + cr: &ast::ComponentReference, + imports: &qualify::ImportMap, + def_map: &crate::ResolveDefMap, + shadowed: &mut std::collections::HashSet, +) { + for part in &cr.parts { + if let Some(subscripts) = &part.subs { + for subscript in subscripts { + collect_subscript_shadowed_import_aliases(subscript, imports, def_map, shadowed); + } + } + } + + let Some(first) = cr.parts.first() else { + return; + }; + let alias = first.ident.text.as_ref(); + let Some(imported_path) = imports.get(alias) else { + return; + }; + let Some(resolved_path) = cr.def_id.and_then(|def_id| def_map.get(&def_id)) else { + return; + }; + if crate::path_utils::leaf_segment(resolved_path) != alias { + return; + } + if resolved_path != imported_path { + shadowed.insert(alias.to_string()); + } +} + +fn collect_subscript_shadowed_import_aliases( + subscript: &ast::Subscript, + imports: &qualify::ImportMap, + def_map: &crate::ResolveDefMap, + shadowed: &mut std::collections::HashSet, +) { + if let ast::Subscript::Expression(expr) = subscript { + collect_shadowed_import_aliases(expr, imports, def_map, shadowed); + } +} diff --git a/crates/rumoca-phase-flatten/src/pipeline/context_and_tests/import_shadow_tests.rs b/crates/rumoca-phase-flatten/src/pipeline/context_and_tests/import_shadow_tests.rs index 565b32ff5..66fb183b8 100644 --- a/crates/rumoca-phase-flatten/src/pipeline/context_and_tests/import_shadow_tests.rs +++ b/crates/rumoca-phase-flatten/src/pipeline/context_and_tests/import_shadow_tests.rs @@ -145,3 +145,34 @@ fn instance_component_member_shadows_import_alias_during_equation_qualification( }; assert_eq!(name.as_str(), "tank.medium.state"); } + +#[test] +fn equation_qualification_resolves_member_from_nearest_parent_instance_scope() { + let mut overlay = ast::InstanceOverlay::default(); + overlay.components.insert( + ast::InstanceId::new(1), + ast::InstanceData { + qualified_name: QualifiedName::from_dotted("jointRRP.rod1.e2_ia"), + ..ast::InstanceData::default() + }, + ); + let mut ctx = Context::new(); + ctx.seed_component_member_scopes(&overlay); + + let expr = ast::Expression::ComponentReference(comp_ref_parts(&["rod1", "e2_ia"])); + let prefix = QualifiedName::from_dotted("jointRRP.jointUSP"); + + let qualified = qualify_expression_imports_with_def_map_ctx( + &expr, + &prefix, + &qualify::ImportMap::default(), + None, + &ctx, + ) + .unwrap(); + + let rumoca_core::Expression::VarRef { name, .. } = qualified else { + panic!("expected VarRef"); + }; + assert_eq!(name.as_str(), "jointRRP.rod1.e2_ia"); +} diff --git a/crates/rumoca-phase-flatten/src/pipeline/context_and_tests/modified_binding_dimensions.rs b/crates/rumoca-phase-flatten/src/pipeline/context_and_tests/modified_binding_dimensions.rs new file mode 100644 index 000000000..70f8f7618 --- /dev/null +++ b/crates/rumoca-phase-flatten/src/pipeline/context_and_tests/modified_binding_dimensions.rs @@ -0,0 +1,137 @@ +use super::*; + +fn modified_integer_param_bindings( + flat: &Model, +) -> rustc_hash::FxHashMap> { + flat.variables + .iter() + .filter(|(_, var)| var.binding_from_modification && var.is_discrete_type) + .map(|(name, var)| (name.to_string(), var.binding.clone())) + .collect() +} + +fn literal_integer(expr: &Expression) -> Option { + match expr { + Expression::Literal { + value: rumoca_core::Literal::Integer(value), + .. + } => Some(*value), + _ => None, + } +} + +fn has_nested_modified_integer_target(expr: &Expression) -> bool { + matches!( + expr, + Expression::VarRef { + name, + subscripts, + .. + } if subscripts.is_empty() && name.is_nested() + ) +} + +impl Context { + pub(crate) fn reconcile_modified_binding_dimensions(&mut self, flat: &mut Model) -> bool { + let mut changed = false; + let modified_integer_params = modified_integer_param_bindings(flat); + for (name, var) in &mut flat.variables { + if !var.binding_from_modification { + continue; + } + let Some(inferred_dims) = self.effective_modified_binding_dims( + name.as_str(), + &var.dims, + &modified_integer_params, + ) else { + continue; + }; + if inferred_dims.is_empty() || var.dims == inferred_dims { + continue; + } + if dims_are_better(&inferred_dims, &var.dims) + || same_rank_concrete_dims(&inferred_dims, &var.dims) + { + var.dims = inferred_dims; + self.reconciled_modified_dimension_names + .insert(name.to_string()); + changed = true; + } + } + changed + } + + pub(super) fn reconcile_modified_integer_parameter_values(&mut self, flat: &Model) -> bool { + let modified_integer_params = modified_integer_param_bindings(flat); + let new_vals = modified_integer_params + .iter() + .filter_map(|(name, binding)| { + binding + .as_ref() + .filter(|expr| has_nested_modified_integer_target(expr)) + .and_then(|_| { + self.modified_integer_binding_value(name, &modified_integer_params, 0) + }) + .map(|value| (name.clone(), value)) + }) + .collect::>(); + let mut changed = false; + for (name, value) in new_vals { + if self.parameter_values.get(&name).copied() != Some(value) { + self.parameter_values.insert(name.clone(), value); + changed = true; + } + if let Some(real_value) = self.real_parameter_values.get_mut(&name) + && real_value.fract() == 0.0 + && *real_value as i64 != value + { + *real_value = value as f64; + changed = true; + } + } + changed + } + + fn effective_modified_binding_dims( + &self, + name: &str, + _current_dims: &[i64], + _modified_integer_params: &rustc_hash::FxHashMap>, + ) -> Option> { + self.array_dimensions.get(name).cloned() + } + + fn modified_integer_binding_value( + &self, + name: &str, + modified_integer_params: &rustc_hash::FxHashMap>, + depth: usize, + ) -> Option { + const MAX_BINDING_DEPTH: usize = 8; + if depth >= MAX_BINDING_DEPTH { + return None; + } + let binding = modified_integer_params.get(name)?.as_ref()?; + if let Some(value) = literal_integer(binding) { + return Some(value); + } + let Expression::VarRef { + name: target, + subscripts, + .. + } = binding + else { + return None; + }; + if !subscripts.is_empty() { + return None; + } + if target.is_nested() { + return self + .modified_integer_binding_value(target.as_str(), modified_integer_params, depth + 1) + .or_else(|| self.get_integer_param(target.as_str())); + } + let source_scope = modifier_source_scope(name)?; + rumoca_core::EvalLookup::lookup_integer(self, target.as_str(), source_scope.as_str()) + } +} diff --git a/crates/rumoca-phase-flatten/src/pipeline/context_tests.rs b/crates/rumoca-phase-flatten/src/pipeline/context_tests.rs index 3dda6bc0c..1705eec2f 100644 --- a/crates/rumoca-phase-flatten/src/pipeline/context_tests.rs +++ b/crates/rumoca-phase-flatten/src/pipeline/context_tests.rs @@ -1,5 +1,7 @@ +// SPEC_0021 file-size exception: pipeline context tests share class-tree and +// lookup fixtures across constants, dimensions, and redeclare paths. split plan: +// move constant-injection, scoped lookup, and redeclare fixtures into modules. use super::*; - #[cfg(test)] mod tests { use super::*; @@ -9,9 +11,9 @@ mod tests { use rumoca_ir_ast::{ClassDef, ClassTree, Component, InstanceData, InstanceId}; use rumoca_ir_flat as flat; use std::sync::Arc; - const TEST_FILE: &str = "context_tests.mo"; - + #[path = "context_tests_modified_binding_dimensions.rs"] + mod modified_binding_dimensions; fn test_source_location() -> rumoca_core::Location { rumoca_core::Location { start_line: 1, @@ -112,6 +114,14 @@ mod tests { }) } + fn component_ref_expr_with_def_id(path: &str, def_id: DefId) -> ast::Expression { + let mut expr = component_ref_expr(path); + if let ast::Expression::ComponentReference(reference) = &mut expr { + reference.def_id = Some(def_id); + } + expr + } + fn int_lit(value: i64) -> Expression { Expression::Literal { value: rumoca_core::Literal::Integer(value), @@ -178,6 +188,385 @@ mod tests { } } + fn symbolic_fill_expr(value: i64, dimensions: &[&str]) -> Expression { + let mut args = vec![int_lit(value)]; + args.extend(dimensions.iter().map(|name| var_ref(name))); + Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Fill, + args, + span: rumoca_core::Span::DUMMY, + } + } + + #[test] + fn settled_symbolic_record_array_dimensions_are_evaluated_once_per_session() { + let tree = source_backed_tree(); + let mut flat = flat::Model::default(); + for (name, value) in [("bank.rows", 3), ("bank.columns", 2)] { + let name = rumoca_core::VarName::new(name); + flat.add_variable( + name.clone(), + flat::Variable { + name, + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(int_lit(value)), + is_discrete_type: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + let cell_names = (1..=3) + .flat_map(|row| (1..=2).map(move |column| format!("bank.cells[{row},{column}].curve"))) + .collect::>(); + for name in &cell_names { + let variable_name = rumoca_core::VarName::new(name); + flat.add_variable( + variable_name.clone(), + flat::Variable { + name: variable_name, + binding: Some(symbolic_fill_expr(0, &["bank.rows", "bank.columns"])), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + + let mut ctx = Context::new(); + let mut session = ctx.collect_parameter_lookup_session(&flat); + ctx.build_parameter_lookup_with_session(&flat, &tree, &mut session); + let expected_parameters = ctx.parameter_values.clone(); + let expected_dimensions = ctx.array_dimensions.clone(); + let expected_flat_bindings = cell_names + .iter() + .map(|name| { + let variable = flat + .variables + .get(&rumoca_core::VarName::new(name)) + .expect("record-array field variable"); + ( + name.clone(), + variable.dims.clone(), + variable.binding.clone(), + ) + }) + .collect::>(); + + ctx.build_parameter_lookup_with_session(&flat, &tree, &mut session); + + assert_eq!(ctx.parameter_values, expected_parameters); + assert_eq!(ctx.array_dimensions, expected_dimensions); + for (name, expected_dims, expected_binding) in expected_flat_bindings { + let variable = flat + .variables + .get(&rumoca_core::VarName::new(&name)) + .expect("record-array field variable"); + assert_eq!(variable.dims, expected_dims); + match (&variable.binding, expected_binding) { + (Some(actual), Some(expected)) => { + assert!(actual.semantically_eq_ignoring_spans(&expected)); + } + (None, None) => {} + _ => panic!("flat binding changed for {name}"), + } + assert_eq!( + session.dimension_evaluation_attempts(&name), + 1, + "settled symbolic dimension binding {name} was reevaluated" + ); + } + } + + #[test] + fn symbolic_dimension_session_retries_when_a_dependency_becomes_available() { + let tree = source_backed_tree(); + let mut flat = flat::Model::default(); + let name = rumoca_core::VarName::new("bank.cells[1,1].curve"); + flat.add_variable( + name.clone(), + flat::Variable { + name: name.clone(), + binding: Some(symbolic_fill_expr(0, &["bank.pending_rows"])), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let mut ctx = Context::new(); + let mut session = ctx.collect_parameter_lookup_session(&flat); + + ctx.build_parameter_lookup_with_session(&flat, &tree, &mut session); + assert!(!ctx.array_dimensions.contains_key(name.as_str())); + assert_eq!(session.dimension_evaluation_attempts(name.as_str()), 1); + + ctx.parameter_values + .insert("bank.pending_rows".to_string(), 3); + ctx.build_parameter_lookup_with_session(&flat, &tree, &mut session); + + assert_eq!(ctx.array_dimensions.get(name.as_str()), Some(&vec![3])); + assert_eq!(session.dimension_evaluation_attempts(name.as_str()), 2); + } + + #[test] + fn resolved_symbolic_dimension_retries_after_lowercase_type_ref_input_changes() { + let tree = source_backed_tree(); + let mut flat = flat::Model::default(); + let name = rumoca_core::VarName::new("wrapper.output"); + flat.add_variable( + name.clone(), + flat::Variable { + name: name.clone(), + binding: Some(symbolic_fill_expr(0, &["Bessel.order"])), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let mut ctx = Context::new(); + ctx.parameter_values + .insert("wrapper.bessel.order".to_string(), 2); + let mut session = ctx.collect_parameter_lookup_session(&flat); + + ctx.build_parameter_lookup_with_session(&flat, &tree, &mut session); + assert_eq!(ctx.array_dimensions.get(name.as_str()), Some(&vec![2])); + assert_eq!(session.dimension_evaluation_attempts(name.as_str()), 1); + + ctx.parameter_values + .insert("wrapper.bessel.order".to_string(), 3); + ctx.build_parameter_lookup_with_session(&flat, &tree, &mut session); + + assert_eq!(ctx.array_dimensions.get(name.as_str()), Some(&vec![3])); + assert_eq!(session.dimension_evaluation_attempts(name.as_str()), 2); + } + + #[test] + fn same_sweep_consumer_retries_after_later_binding_produces_dimensions() { + let tree = source_backed_tree(); + let mut flat = flat::Model::default(); + let consumer_name = rumoca_core::VarName::new("B"); + flat.add_variable( + consumer_name.clone(), + flat::Variable { + name: consumer_name.clone(), + binding: Some(Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Fill, + args: vec![int_lit(0), size_dim_expr("A", 1)], + span: rumoca_core::Span::DUMMY, + }), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let producer_name = rumoca_core::VarName::new("A"); + flat.add_variable( + producer_name.clone(), + flat::Variable { + name: producer_name.clone(), + binding: Some(symbolic_fill_expr(0, &["n"])), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let mut ctx = Context::new(); + ctx.parameter_values.insert("n".to_string(), 3); + let mut session = ctx.collect_parameter_lookup_session(&flat); + + ctx.build_parameter_lookup_with_session(&flat, &tree, &mut session); + + assert_eq!( + ctx.array_dimensions.get(producer_name.as_str()), + Some(&vec![3]) + ); + assert_eq!( + ctx.array_dimensions.get(consumer_name.as_str()), + Some(&vec![3]) + ); + assert_eq!( + session.dimension_evaluation_attempts(consumer_name.as_str()), + 2 + ); + } + + #[test] + fn same_sweep_resolved_consumer_retries_after_producer_dimension_changes() { + let tree = source_backed_tree(); + let mut flat = flat::Model::default(); + let consumer_name = rumoca_core::VarName::new("B"); + flat.add_variable( + consumer_name.clone(), + flat::Variable { + name: consumer_name.clone(), + binding: Some(Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Fill, + args: vec![int_lit(0), size_dim_expr("A", 1)], + span: rumoca_core::Span::DUMMY, + }), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let producer_name = rumoca_core::VarName::new("A"); + flat.add_variable( + producer_name.clone(), + flat::Variable { + name: producer_name.clone(), + dims: vec![2], + binding: Some(symbolic_fill_expr(0, &["n"])), + binding_from_modification: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let mut ctx = Context::new(); + ctx.parameter_values.insert("n".to_string(), 3); + let mut session = ctx.collect_parameter_lookup_session(&flat); + + ctx.build_parameter_lookup_with_session(&flat, &tree, &mut session); + + assert_eq!( + ctx.array_dimensions.get(producer_name.as_str()), + Some(&vec![3]) + ); + assert_eq!( + ctx.array_dimensions.get(consumer_name.as_str()), + Some(&vec![3]) + ); + assert_eq!( + session.dimension_evaluation_attempts(consumer_name.as_str()), + 2 + ); + } + + #[test] + fn same_sweep_scoped_array_dimension_shadows_root_scalar_for_consumer() { + let tree = source_backed_tree(); + let mut flat = flat::Model::default(); + let consumer_name = rumoca_core::VarName::new("wrapper.B"); + flat.add_variable( + consumer_name.clone(), + flat::Variable { + name: consumer_name.clone(), + binding: Some(Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Fill, + args: vec![int_lit(0), size_dim_expr("A", 1)], + span: rumoca_core::Span::DUMMY, + }), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let producer_name = rumoca_core::VarName::new("wrapper.A"); + flat.add_variable( + producer_name.clone(), + flat::Variable { + name: producer_name.clone(), + dims: vec![2], + binding: Some(symbolic_fill_expr(0, &["n"])), + binding_from_modification: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let mut ctx = Context::new(); + ctx.parameter_values.insert("A".to_string(), 99); + ctx.parameter_values.insert("n".to_string(), 3); + let mut session = ctx.collect_parameter_lookup_session(&flat); + + ctx.build_parameter_lookup_with_session(&flat, &tree, &mut session); + + assert_eq!( + ctx.array_dimensions.get(producer_name.as_str()), + Some(&vec![3]) + ); + assert_eq!( + ctx.array_dimensions.get(consumer_name.as_str()), + Some(&vec![3]) + ); + assert_eq!( + session.dimension_evaluation_attempts(consumer_name.as_str()), + 2 + ); + } + + #[test] + fn production_symbolic_dimension_session_matches_uncached_nout_oracle() { + let tree = source_backed_tree(); + let mut initial_flat = flat::Model::default(); + let width_name = rumoca_core::VarName::new("stack.width"); + initial_flat.add_variable( + width_name.clone(), + flat::Variable { + name: width_name, + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(int_lit(2)), + is_discrete_type: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let columns_name = rumoca_core::VarName::new("stack.cell[2,3].columns"); + initial_flat.add_variable( + columns_name.clone(), + flat::Variable { + name: columns_name.clone(), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let curve_name = rumoca_core::VarName::new("stack.cell[2,3].curve"); + initial_flat.add_variable( + curve_name.clone(), + flat::Variable { + name: curve_name.clone(), + binding: Some(symbolic_fill_expr(0, &["nout"])), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + let mut overlay = InstanceOverlay::default(); + overlay.components.insert( + InstanceId::new(1), + symbolic_instance( + InstanceId::new(1), + columns_name.as_str(), + vec![ast::Subscript::Expression(component_ref_expr( + "stack.width", + ))], + ), + ); + + let mut cached_flat = initial_flat.clone(); + let mut cached_ctx = Context::new(); + stabilize_symbolic_component_dimensions(&mut cached_ctx, &mut cached_flat, &overlay, &tree) + .expect("cached production stabilization"); + + let mut oracle_flat = initial_flat; + let mut oracle_ctx = Context::new(); + let max_passes = overlay.components.len().max(1) + 1; + for _ in 0..max_passes { + oracle_ctx.build_parameter_lookup(&oracle_flat, &tree); + let reconciled = oracle_ctx.reconcile_modified_binding_dimensions(&mut oracle_flat); + let recomputed = oracle_ctx + .recompute_symbolic_component_dimensions(&mut oracle_flat, &overlay, &tree) + .expect("uncached oracle stabilization"); + if !reconciled && !recomputed { + break; + } + } + + assert_eq!( + oracle_ctx.array_dimensions.get(curve_name.as_str()), + Some(&vec![2]), + "oracle must resolve nout from the late columns dimension" + ); + assert_eq!(cached_ctx.parameter_values, oracle_ctx.parameter_values); + assert_eq!(cached_ctx.array_dimensions, oracle_ctx.array_dimensions); + assert_eq!( + format!("{cached_flat:#?}"), + format!("{oracle_flat:#?}"), + "cached production stabilization must preserve the complete Flat IR" + ); + } + fn size_dim_expr(name: &str, dim: i64) -> Expression { Expression::BuiltinCall { function: rumoca_core::BuiltinFunction::Size, @@ -419,6 +808,63 @@ mod tests { assert_eq!(ctx.get_integer_param("filter.order"), Some(3)); } + #[test] + fn test_modified_integer_binding_updates_dependent_enum_if_parameter() { + let mut ctx = Context::new(); + let tree = source_backed_tree(); + let mut flat = flat::Model::default(); + + for (name, binding, binding_from_modification) in [ + ("order", int_lit(3), false), + ("filter.order", var_ref("order"), true), + ( + "filter.analogFilter", + var_ref("AnalogFilter.CriticalDamping"), + false, + ), + ( + "filter.nr", + Expression::If { + branches: vec![( + Expression::Binary { + op: rumoca_core::OpBinary::Eq, + lhs: Box::new(var_ref("filter.analogFilter")), + rhs: Box::new(var_ref("AnalogFilter.CriticalDamping")), + span: rumoca_core::Span::DUMMY, + }, + var_ref("filter.order"), + )], + else_branch: Box::new(Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Mod, + args: vec![var_ref("filter.order"), int_lit(2)], + span: rumoca_core::Span::DUMMY, + }), + span: rumoca_core::Span::DUMMY, + }, + false, + ), + ] { + let var_name = rumoca_core::VarName::new(name); + flat.add_variable( + var_name.clone(), + flat::Variable { + name: var_name, + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(binding), + binding_from_modification, + is_discrete_type: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + + ctx.build_parameter_lookup(&flat, &tree); + + assert_eq!(ctx.get_integer_param("filter.order"), Some(3)); + assert_eq!(ctx.get_integer_param("filter.nr"), Some(3)); + } + #[test] fn real_modifier_bindings_resolve_in_enclosing_scope_for_sibling_instances() { let mut ctx = Context::new(); @@ -450,6 +896,30 @@ mod tests { assert_eq!(ctx.real_parameter_values.get("line2.TD"), Some(&0.001)); } + #[test] + fn real_modifier_binding_parameter_is_non_structural_even_when_fixed_default() { + let mut ctx = Context::new(); + let tree = ClassTree::default(); + let mut flat = flat::Model::default(); + let name = rumoca_core::VarName::new("force.v_nominal"); + flat.add_variable( + name.clone(), + flat::Variable { + name, + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(real_lit(5.0)), + binding_from_modification: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + ctx.build_parameter_lookup(&flat, &tree); + + assert!(ctx.non_structural_params.contains("force.v_nominal")); + assert_eq!(ctx.real_parameter_values.get("force.v_nominal"), Some(&5.0)); + } + #[test] fn real_modifier_bindings_resolve_transmission_line_delay_chain() { let mut ctx = Context::new(); @@ -570,6 +1040,43 @@ mod tests { ); } + #[test] + fn test_propagate_unexpanded_component_array_dims_to_model_fields() { + let mut flat = flat::Model::default(); + let var_name = rumoca_core::VarName::new("sensor.child.y"); + flat.add_variable( + var_name.clone(), + flat::Variable { + name: var_name.clone(), + is_primitive: true, + dims: Vec::new(), + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + let mut overlay = InstanceOverlay::default(); + overlay.components.insert( + InstanceId::new(1), + InstanceData { + instance_id: InstanceId::new(1), + qualified_name: QualifiedName::from_dotted("sensor.child"), + dims: vec![6], + is_primitive: false, + ..Default::default() + }, + ); + + propagate_unexpanded_record_array_dims(&mut flat, &overlay); + + assert_eq!( + flat.variables + .get(&var_name) + .expect("missing sensor child field") + .dims, + vec![6] + ); + } + #[test] fn test_propagate_unexpanded_record_array_dims_does_not_double_prefix_dims() { let mut flat = flat::Model::default(); @@ -734,6 +1241,132 @@ mod tests { assert_eq!(ctx.array_dimensions.get("realFFT.abs"), Some(&vec![401])); } + #[test] + fn symbolic_component_dimensions_use_modified_parameter_value() { + let mut ctx = Context::new(); + let tree = source_backed_tree(); + let mut flat = flat::Model::default(); + + let root_n = rumoca_core::VarName::new("N"); + flat.add_variable( + root_n.clone(), + flat::Variable { + name: root_n, + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(int_lit(7)), + is_discrete_type: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + let local_n = rumoca_core::VarName::new("aD_Converter.N"); + flat.add_variable( + local_n.clone(), + flat::Variable { + name: local_n, + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(var_ref("N")), + binding_from_modification: true, + start: Some(Expression::Literal { + value: rumoca_core::Literal::Real(8.0), + span: test_span(), + }), + is_discrete_type: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + let y_name = rumoca_core::VarName::new("aD_Converter.y"); + flat.add_variable( + y_name.clone(), + flat::Variable { + name: y_name.clone(), + dims: vec![8], + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + let mut overlay = InstanceOverlay::default(); + overlay.components.insert( + InstanceId::new(1), + symbolic_instance( + InstanceId::new(1), + "aD_Converter.y", + vec![ast::Subscript::Expression(component_ref_expr( + "aD_Converter.N", + ))], + ), + ); + + ctx.build_parameter_lookup(&flat, &tree); + assert_eq!(ctx.get_integer_param("aD_Converter.N"), Some(7)); + + let changed = ctx + .recompute_symbolic_component_dimensions(&mut flat, &overlay, &tree) + .expect("modified symbolic dimension should resolve"); + + assert!(changed); + assert_eq!( + flat.variables.get(&y_name).expect("y variable").dims, + vec![7] + ); + assert_eq!(ctx.array_dimensions.get("aD_Converter.y"), Some(&vec![7])); + } + + #[test] + fn symbolic_component_dimensions_resolve_package_member_target_def_constant() { + let nstate_def = DefId::new(42); + let target_name = "Modelica.Math.Random.Generators.Xorshift128plus.nState"; + let mut ctx = Context::new(); + ctx.target_def_names + .insert(nstate_def, target_name.to_string()); + ctx.parameter_values.insert(target_name.to_string(), 4); + + let mut tree = source_backed_tree(); + tree.def_map.insert(nstate_def, target_name.to_string()); + let mut flat = flat::Model::default(); + let state_name = rumoca_core::VarName::new("motor.uniformNoise.state"); + flat.add_variable( + state_name.clone(), + flat::Variable { + name: state_name.clone(), + dims: vec![1], + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + let mut overlay = InstanceOverlay::default(); + overlay.components.insert( + InstanceId::new(1), + symbolic_instance( + InstanceId::new(1), + "motor.uniformNoise.state", + vec![ast::Subscript::Expression(component_ref_expr_with_def_id( + "generator.nState", + nstate_def, + ))], + ), + ); + + let changed = ctx + .recompute_symbolic_component_dimensions(&mut flat, &overlay, &tree) + .expect("dimension target definition constant should resolve"); + + assert!(changed); + assert_eq!( + flat.variables.get(&state_name).expect("state").dims, + vec![4] + ); + assert_eq!( + ctx.array_dimensions.get("motor.uniformNoise.state"), + Some(&vec![4]) + ); + } + #[test] fn enum_type_component_dimensions_use_literal_count() { let mut ctx = Context::new(); @@ -982,6 +1615,56 @@ mod tests { assert_eq!(ctx.array_dimensions.get("a.x"), Some(&vec![4])); } + #[test] + fn colon_component_dimensions_prefer_binding_shape_over_stale_larger_dims() { + let mut ctx = Context::new(); + let tree = source_backed_tree(); + let mut flat = flat::Model::default(); + let x_name = rumoca_core::VarName::new("a.x"); + flat.add_variable( + x_name.clone(), + flat::Variable { + name: x_name.clone(), + dims: vec![3], + binding: Some(Expression::Array { + elements: vec![Expression::Literal { + value: rumoca_core::Literal::Integer(1), + span: rumoca_core::Span::DUMMY, + }], + is_matrix: false, + span: rumoca_core::Span::DUMMY, + }), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + let mut overlay = InstanceOverlay::default(); + overlay.components.insert( + InstanceId::new(1), + symbolic_instance( + InstanceId::new(1), + "a.x", + vec![ast::Subscript::Range { + token: rumoca_core::Token::default(), + }], + ), + ); + + ctx.build_parameter_lookup(&flat, &tree); + + let changed = ctx + .recompute_symbolic_component_dimensions(&mut flat, &overlay, &tree) + .expect("colon dimension should prefer binding shape"); + + assert!(changed); + assert_eq!( + flat.variables.get(&x_name).expect("a.x variable").dims, + vec![1] + ); + assert_eq!(ctx.array_dimensions.get("a.x"), Some(&vec![1])); + } + #[test] fn colon_component_dimensions_accept_zero_sized_binding_shape() { let mut ctx = Context::new(); @@ -1696,12 +2379,10 @@ mod tests { tree.definitions.classes.insert("Host".to_string(), host); tree.def_map.insert(host_def_id, "Host".to_string()); tree.name_map.insert("Host".to_string(), host_def_id); - let instance = InstanceData { type_def_id: Some(host_def_id), ..Default::default() }; - let class_index = rumoca_ir_ast::ClassDefIndex::from_tree(&tree); let overrides = component_overrides(&instance, &tree, &class_index); assert_eq!( @@ -1796,7 +2477,6 @@ mod tests { tree.definitions.classes.insert("Host".to_string(), host); tree.def_map.insert(host_def_id, "Host".to_string()); tree.name_map.insert("Host".to_string(), host_def_id); - let mut instance = InstanceData { type_def_id: Some(host_def_id), ..Default::default() @@ -1810,7 +2490,6 @@ mod tests { None, ), ); - let class_index = rumoca_ir_ast::ClassDefIndex::from_tree(&tree); let overrides = component_overrides(&instance, &tree, &class_index); assert_eq!( diff --git a/crates/rumoca-phase-flatten/src/pipeline/context_tests/tests/context_tests_modified_binding_dimensions.rs b/crates/rumoca-phase-flatten/src/pipeline/context_tests/tests/context_tests_modified_binding_dimensions.rs new file mode 100644 index 000000000..7e6d3ffdf --- /dev/null +++ b/crates/rumoca-phase-flatten/src/pipeline/context_tests/tests/context_tests_modified_binding_dimensions.rs @@ -0,0 +1,462 @@ +use super::*; + +fn var_ref_with_target_def(name: &str, target_def_id: DefId) -> Expression { + let parts = crate::path_utils::segments(name) + .into_iter() + .map(|segment| rumoca_core::ComponentRefPart { + ident: segment.to_string(), + span: rumoca_core::Span::DUMMY, + subs: Vec::new(), + }) + .collect(); + let component_ref = rumoca_core::ComponentReference { + local: false, + span: rumoca_core::Span::DUMMY, + parts, + def_id: Some(target_def_id), + }; + Expression::VarRef { + name: rumoca_core::Reference::with_component_reference(name, component_ref), + subscripts: Vec::new(), + span: rumoca_core::Span::DUMMY, + } +} + +fn self_qualified_var_ref_with_source_leaf( + rendered_name: &str, + source_leaf: &str, + target_def_id: DefId, +) -> Expression { + let component_ref = rumoca_core::ComponentReference { + local: false, + span: rumoca_core::Span::DUMMY, + parts: vec![rumoca_core::ComponentRefPart { + ident: source_leaf.to_string(), + span: rumoca_core::Span::DUMMY, + subs: Vec::new(), + }], + def_id: Some(target_def_id), + }; + Expression::VarRef { + name: rumoca_core::Reference::with_component_reference(rendered_name, component_ref), + subscripts: Vec::new(), + span: rumoca_core::Span::DUMMY, + } +} + +#[test] +fn reconciled_modified_record_table_shape_beats_stale_default_binding_shape() { + let mut ctx = Context::new(); + let tree = source_backed_tree(); + ctx.parameter_values.insert("rows".to_string(), 29); + ctx.parameter_values.insert("columns".to_string(), 2); + ctx.array_dimensions + .insert("stack.cellData.OCV_SOC".to_string(), vec![29, 2]); + let row = || Expression::Array { + elements: vec![int_lit(0), int_lit(1)], + is_matrix: false, + span: test_span(), + }; + let name = rumoca_core::VarName::new("stack.cellData.OCV_SOC"); + let mut flat = flat::Model::default(); + flat.add_variable( + name.clone(), + flat::Variable { + name: name.clone(), + dims: vec![2, 2], + binding: Some(Expression::Array { + elements: vec![row(), row()], + is_matrix: true, + span: test_span(), + }), + binding_from_modification: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let mut overlay = InstanceOverlay::default(); + overlay.components.insert( + InstanceId::new(1), + symbolic_instance( + InstanceId::new(1), + "stack.cellData.OCV_SOC", + vec![ + ast::Subscript::Expression(component_ref_expr("rows")), + ast::Subscript::Expression(component_ref_expr("columns")), + ], + ), + ); + + assert!(ctx.reconcile_modified_binding_dimensions(&mut flat)); + let recomputed = ctx + .recompute_symbolic_component_dimensions(&mut flat, &overlay, &tree) + .expect("modified table dimensions reconcile"); + + assert!(!recomputed); + assert_eq!(flat.variables.get(&name).unwrap().dims, vec![29, 2]); + + ctx.build_parameter_lookup(&flat, &tree); + assert!(!ctx.reconcile_modified_binding_dimensions(&mut flat)); + assert!( + !ctx.recompute_symbolic_component_dimensions(&mut flat, &overlay, &tree) + .expect("fixed point remains stable") + ); +} + +#[test] +fn cache_only_dimension_repair_is_not_reported_as_a_flat_model_change() { + let mut ctx = Context::new(); + let tree = source_backed_tree(); + let name = rumoca_core::VarName::new("flowModel.pathLengths_internal"); + ctx.array_dimensions.insert(name.to_string(), vec![3]); + + let mut flat = flat::Model::default(); + flat.add_variable( + name.clone(), + flat::Variable { + name: name.clone(), + dims: vec![2], + binding: Some(Expression::Array { + elements: vec![int_lit(1), int_lit(1)], + is_matrix: false, + span: test_span(), + }), + binding_from_modification: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let mut overlay = InstanceOverlay::default(); + overlay.components.insert( + InstanceId::new(1), + symbolic_instance( + InstanceId::new(1), + name.as_str(), + vec![ast::Subscript::Expression(component_ref_expr("n"))], + ), + ); + + let flat_changed = ctx + .recompute_symbolic_component_dimensions(&mut flat, &overlay, &tree) + .expect("cache-only dimension repair"); + + assert!(!flat_changed); + assert_eq!(flat.variables.get(&name).unwrap().dims, vec![2]); + assert_eq!(ctx.array_dimensions.get(name.as_str()), Some(&vec![2])); +} + +#[test] +fn modified_binding_dimensions_replace_stale_declared_shape() { + let mut ctx = Context::new(); + let tree = source_backed_tree(); + let mut flat = flat::Model::default(); + let ambiguous_mr_def = DefId::new(42); + ctx.target_def_names + .insert(ambiguous_mr_def, "mr".to_string()); + + for (name, binding, binding_from_modification) in [ + ("mr", int_lit(3), false), + ("aimsM.mr", int_lit(5), true), + ( + "aimsM.rotor.m", + var_ref_with_target_def("aimsM.mr", ambiguous_mr_def), + true, + ), + ] { + let var_name = rumoca_core::VarName::new(name); + flat.add_variable( + var_name.clone(), + flat::Variable { + name: var_name, + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(binding), + binding_from_modification, + is_discrete_type: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + + let r_name = rumoca_core::VarName::new("aimsM.rotor.resistor.R"); + flat.add_variable( + r_name.clone(), + flat::Variable { + name: r_name.clone(), + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + dims: vec![3], + binding: Some(Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Fill, + args: vec![int_lit(1), var_ref("aimsM.rotor.m")], + span: rumoca_core::Span::DUMMY, + }), + binding_from_modification: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + let local_m = rumoca_core::VarName::new("aimsM.rotor.resistor.m"); + flat.add_variable( + local_m.clone(), + flat::Variable { + name: local_m, + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(var_ref_with_target_def("aimsM.mr", ambiguous_mr_def)), + binding_from_modification: true, + is_discrete_type: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let mut overlay = InstanceOverlay::default(); + overlay.components.insert( + InstanceId::new(1), + symbolic_instance( + InstanceId::new(1), + "aimsM.rotor.resistor.R", + vec![ast::Subscript::Expression(component_ref_expr( + "aimsM.rotor.resistor.m", + ))], + ), + ); + + ctx.build_parameter_lookup(&flat, &tree); + assert_eq!(ctx.get_integer_param("aimsM.rotor.m"), Some(5)); + assert_eq!(ctx.get_integer_param("aimsM.rotor.resistor.m"), Some(5)); + assert_eq!( + ctx.array_dimensions.get("aimsM.rotor.resistor.R"), + Some(&vec![5]) + ); + assert_eq!( + flat.variables.get(&r_name).expect("R variable").dims, + vec![3] + ); + + assert!(ctx.reconcile_modified_binding_dimensions(&mut flat)); + assert_eq!( + flat.variables.get(&r_name).expect("R variable").dims, + vec![5] + ); + + flat.variables.get_mut(&r_name).expect("R variable").dims = vec![3]; + assert!( + ctx.recompute_symbolic_component_dimensions(&mut flat, &overlay, &tree) + .expect("modified local m dimension is available to symbolic dimensions") + ); + assert_eq!( + flat.variables.get(&r_name).expect("R variable").dims, + vec![5] + ); + + ctx.build_parameter_lookup(&flat, &tree); + assert_eq!( + ctx.array_dimensions.get("aimsM.rotor.resistor.R"), + Some(&vec![5]) + ); +} + +#[test] +fn modified_integer_unqualified_rhs_prefers_modifier_source_scope_over_target_def() { + let mut ctx = Context::new(); + let target_def = DefId::new(43); + ctx.target_def_names + .insert(target_def, "sensor.block.m".to_string()); + ctx.parameter_values.insert("sensor.m".to_string(), 3); + ctx.parameter_values.insert("sensor.block.m".to_string(), 6); + + let params = vec![( + "sensor.block.m".to_string(), + var_ref_with_target_def("m", target_def), + )]; + + assert!( + ctx.eval_modified_integer_params(¶ms), + "modifier-origin unqualified RHS must be resolved where the modifier was written" + ); + assert_eq!(ctx.get_integer_param("sensor.block.m"), Some(3)); +} + +#[test] +fn modified_binding_dimensions_do_not_follow_unrelated_local_m() { + let mut ctx = Context::new(); + let tree = source_backed_tree(); + let mut flat = flat::Model::default(); + + let local_m = rumoca_core::VarName::new("component.m"); + flat.add_variable( + local_m.clone(), + flat::Variable { + name: local_m, + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(int_lit(5)), + binding_from_modification: true, + is_discrete_type: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + let x_name = rumoca_core::VarName::new("component.x"); + flat.add_variable( + x_name.clone(), + flat::Variable { + name: x_name.clone(), + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + dims: vec![3], + binding: Some(Expression::Array { + elements: vec![int_lit(1), int_lit(2), int_lit(3)], + is_matrix: false, + span: rumoca_core::Span::DUMMY, + }), + binding_from_modification: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + ctx.build_parameter_lookup(&flat, &tree); + assert_eq!(ctx.get_integer_param("component.m"), Some(5)); + assert_eq!(ctx.array_dimensions.get("component.x"), Some(&vec![3])); + assert!(!ctx.reconcile_modified_binding_dimensions(&mut flat)); + assert_eq!( + flat.variables.get(&x_name).expect("x variable").dims, + vec![3] + ); +} + +#[test] +fn modified_integer_alias_prefers_modifier_scope_over_declared_target_def() { + let mut ctx = Context::new(); + let tree = source_backed_tree(); + let mut flat = flat::Model::default(); + let class_m_def = DefId::new(77); + ctx.target_def_names + .insert(class_m_def, "Pkg.Block.m".to_string()); + + for (name, binding, binding_from_modification) in [ + ("component.m", int_lit(3), true), + ("Pkg.Block.m", int_lit(6), false), + ( + "component.child.m", + var_ref_with_target_def("m", class_m_def), + true, + ), + ] { + let var_name = rumoca_core::VarName::new(name); + flat.add_variable( + var_name.clone(), + flat::Variable { + name: var_name, + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(binding), + binding_from_modification, + is_discrete_type: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + + ctx.build_parameter_lookup(&flat, &tree); + + assert_eq!(ctx.get_integer_param("component.m"), Some(3)); + assert_eq!(ctx.get_integer_param("Pkg.Block.m"), Some(6)); + assert_eq!(ctx.get_integer_param("component.child.m"), Some(3)); +} + +#[test] +fn modified_integer_alias_uses_source_leaf_when_flat_name_is_self_qualified() { + let mut ctx = Context::new(); + let tree = source_backed_tree(); + let mut flat = flat::Model::default(); + let class_m_def = DefId::new(78); + ctx.target_def_names + .insert(class_m_def, "Pkg.Block.m".to_string()); + + for (name, binding, binding_from_modification) in [ + ("component.m", int_lit(3), true), + ("Pkg.Block.m", int_lit(6), false), + ( + "component.child.m", + self_qualified_var_ref_with_source_leaf("component.child.m", "m", class_m_def), + true, + ), + ] { + let var_name = rumoca_core::VarName::new(name); + flat.add_variable( + var_name.clone(), + flat::Variable { + name: var_name, + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(binding), + binding_from_modification, + is_discrete_type: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + + ctx.build_parameter_lookup(&flat, &tree); + + assert_eq!(ctx.get_integer_param("component.child.m"), Some(3)); +} + +#[test] +fn modified_integer_alias_sync_is_not_limited_to_phase_count_m() { + let mut ctx = Context::new(); + let tree = source_backed_tree(); + let mut flat = flat::Model::default(); + + for (name, binding, binding_from_modification) in [ + ("system.n", int_lit(4), true), + ("component.nLocal", var_ref("system.n"), true), + ] { + let var_name = rumoca_core::VarName::new(name); + flat.add_variable( + var_name.clone(), + flat::Variable { + name: var_name, + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(binding), + binding_from_modification, + is_discrete_type: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + + let x_name = rumoca_core::VarName::new("component.x"); + flat.add_variable( + x_name.clone(), + flat::Variable { + name: x_name.clone(), + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + dims: vec![2], + binding: Some(Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Fill, + args: vec![int_lit(1), var_ref("component.nLocal")], + span: rumoca_core::Span::DUMMY, + }), + binding_from_modification: true, + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + ctx.build_parameter_lookup(&flat, &tree); + + assert_eq!(ctx.get_integer_param("component.nLocal"), Some(4)); + assert_eq!(ctx.array_dimensions.get("component.x"), Some(&vec![4])); + assert_eq!( + flat.variables.get(&x_name).expect("x variable").dims, + vec![2] + ); + assert!(ctx.reconcile_modified_binding_dimensions(&mut flat)); + assert_eq!( + flat.variables.get(&x_name).expect("x variable").dims, + vec![4] + ); +} diff --git a/crates/rumoca-phase-flatten/src/pipeline/dim_recovery.rs b/crates/rumoca-phase-flatten/src/pipeline/dim_recovery.rs index eb7ca7933..4dc414ea7 100644 --- a/crates/rumoca-phase-flatten/src/pipeline/dim_recovery.rs +++ b/crates/rumoca-phase-flatten/src/pipeline/dim_recovery.rs @@ -177,7 +177,7 @@ pub(crate) fn infer_expr_dims( } pub(crate) fn collect_parent_dims(flat: &Model, overlay: &InstanceOverlay) -> ParentDims { - overlay + let mut parent_dims: ParentDims = overlay .components .values() .filter(|inst| !inst.is_primitive && !inst.dims.is_empty()) @@ -194,7 +194,23 @@ pub(crate) fn collect_parent_dims(flat: &Model, overlay: &InstanceOverlay) -> Pa Some((path, inst.dims.clone())) } }) - .collect() + .collect(); + + for (path, dims) in &overlay.array_parent_dims { + if dims.is_empty() || parent_dims.iter().any(|(existing, _)| existing == path) { + continue; + } + let first_elem = format!("{path}[1]."); + let expanded = flat + .variables + .keys() + .any(|k| k.as_str().starts_with(&first_elem)); + if !expanded { + parent_dims.push((path.clone(), dims.clone())); + } + } + + parent_dims } pub(crate) fn collect_function_output_dims(flat: &Model) -> DimMap { @@ -279,6 +295,12 @@ pub(crate) fn recover_nested_dims_from_bindings( let mut changed = false; for var in flat.variables.values_mut() { + if matches!( + var.variability, + rumoca_core::Variability::Parameter(_) | rumoca_core::Variability::Constant(_) + ) { + continue; + } let Some((_, parent)) = matching_parent(var.name.as_str(), parent_dims) else { continue; }; @@ -300,6 +322,12 @@ pub(crate) fn recover_nested_dims_from_bindings( pub(crate) fn prepend_missing_parent_dims(flat: &mut Model, parent_dims: &ParentDims) { for var in flat.variables.values_mut() { + if matches!( + var.variability, + rumoca_core::Variability::Parameter(_) | rumoca_core::Variability::Constant(_) + ) { + continue; + } let Some((_, parent)) = matching_parent(var.name.as_str(), parent_dims) else { continue; }; @@ -344,6 +372,12 @@ pub(crate) fn complete_child_dims_from_hints( hints: &DimMap, ) { for var in flat.variables.values_mut() { + if matches!( + var.variability, + rumoca_core::Variability::Parameter(_) | rumoca_core::Variability::Constant(_) + ) { + continue; + } let Some((prefix, parent)) = matching_parent(var.name.as_str(), parent_dims) else { continue; }; diff --git a/crates/rumoca-phase-flatten/src/pipeline/flatten_pipeline.rs b/crates/rumoca-phase-flatten/src/pipeline/flatten_pipeline.rs index 21a2e406d..4f8031bf0 100644 --- a/crates/rumoca-phase-flatten/src/pipeline/flatten_pipeline.rs +++ b/crates/rumoca-phase-flatten/src/pipeline/flatten_pipeline.rs @@ -762,6 +762,7 @@ pub(crate) fn process_component_instances_for_flatten( tree: &ast::ClassTree, class_index: &ast::ClassDefIndex<'_>, component_members: &component_member_scope::ComponentMemberScopes, + simulated_root_name: Option<&str>, ) -> Result<(), FlattenError> { let mut import_cache = ImportCaches::default(); let scope_index = OverlayScopeIndex::new(overlay); @@ -772,6 +773,7 @@ pub(crate) fn process_component_instances_for_flatten( process_component_instance(ComponentInstanceProcess { flat, instance_data, + simulated_root_name, component_override_map, tree, class_index, @@ -841,11 +843,10 @@ pub(crate) fn prepare_context_for_equation_flattening( // Re-apply parameter lookup from materialized flat variables after // class/package constant injection so record rebindings override injected // declaration defaults (MLS §7.2.3/§7.2.4, §8.3.3 structural ranges). - ctx.build_parameter_lookup(flat, tree); inject_referenced_qualified_class_constants(tree, class_index, model_name, flat, overlay, ctx)?; - if ctx.recompute_symbolic_component_dimensions(flat, overlay, tree)? { - ctx.build_parameter_lookup(flat, tree); - } + stabilize_symbolic_component_dimensions(ctx, flat, overlay, tree)?; + inject_referenced_qualified_class_constants(tree, class_index, model_name, flat, overlay, ctx)?; + stabilize_symbolic_component_dimensions(ctx, flat, overlay, tree)?; ctx.refresh_enum_parameter_lookup(flat); pre_evaluate_structural_equations(ctx, overlay, tree)?; @@ -903,6 +904,25 @@ pub(crate) fn process_class_instances_for_flatten( Ok(()) } +pub(crate) fn stabilize_symbolic_component_dimensions( + ctx: &mut Context, + flat: &mut flat::Model, + overlay: &ast::InstanceOverlay, + tree: &ast::ClassTree, +) -> Result<(), FlattenError> { + let max_passes = overlay.components.len().max(1) + 1; + let mut parameter_lookup = ctx.collect_parameter_lookup_session(flat); + for _ in 0..max_passes { + ctx.build_parameter_lookup_with_session(flat, tree, &mut parameter_lookup); + let reconciled = ctx.reconcile_modified_binding_dimensions(flat); + let recomputed = ctx.recompute_symbolic_component_dimensions(flat, overlay, tree)?; + if !reconciled && !recomputed { + return Ok(()); + } + } + Ok(()) +} + pub(crate) struct FinalizeFlatModelInput<'a, 'tree> { pub(crate) ctx: &'a mut Context, pub(crate) flat: &'a mut flat::Model, @@ -930,6 +950,8 @@ pub(crate) fn finalize_flat_model( component_override_map, } = input; outer_refs::redirect_outer_refs(flat, &overlay.outer_prefix_to_inner); + materialize_referenced_expandable_connector_members(flat, overlay, tree, class_index)?; + propagate_unexpanded_record_array_dims(flat, overlay); let connections_start = maybe_start_timer(); let connections_result = @@ -947,6 +969,7 @@ pub(crate) fn finalize_flat_model( mark_record_constructor_calls(flat, tree); canonicalize_varrefs_via_record_aliases(flat, ctx); canonicalize_varrefs_via_instantiated_def_ids(flat); + canonicalize_varrefs_via_record_aliases(flat, ctx); drop_invalid_field_access_bindings(flat); propagate_unexpanded_record_array_dims(flat, overlay); flat.oc_break_edge_scalar_count = vcg::compute_break_edge_scalar_count( @@ -960,10 +983,7 @@ pub(crate) fn finalize_flat_model( collapse_index_refs_to_known_varrefs(flat); inject_referenced_qualified_class_constants(tree, class_index, model_name, flat, overlay, ctx)?; substitute_known_constants_in_flat(flat, ctx)?; - ctx.build_parameter_lookup(flat, tree); - if ctx.recompute_symbolic_component_dimensions(flat, overlay, tree)? { - ctx.build_parameter_lookup(flat, tree); - } + stabilize_symbolic_component_dimensions(ctx, flat, overlay, tree)?; recover_indexed_lhs_dimensions(flat); mark_record_constructor_calls(flat, tree); let collected_new_functions = collect_rewritten_functions_to_fixed_point( @@ -976,21 +996,22 @@ pub(crate) fn finalize_flat_model( )?; if collected_new_functions { mark_record_constructor_calls(flat, tree); - inject_referenced_qualified_class_constants( - tree, - class_index, - model_name, - flat, - overlay, - ctx, - )?; - substitute_known_constants_in_flat(flat, ctx)?; - mark_record_constructor_calls(flat, tree); - collapse_index_refs_to_known_varrefs(flat); + functions::lower_record_function_params(flat)?; } + mark_record_constructor_calls(flat, tree); + inject_referenced_qualified_class_constants(tree, class_index, model_name, flat, overlay, ctx)?; + substitute_known_constants_in_flat(flat, ctx)?; + mark_record_constructor_calls(flat, tree); + collapse_index_refs_to_known_varrefs(flat); canonicalize_varrefs_via_instantiated_def_ids(flat); + canonicalize_varrefs_via_record_aliases(flat, ctx); functions::canonicalize_collected_function_calls(flat); + functions::lower_record_function_params(flat)?; + mark_record_constructor_calls(flat, tree); + substitute_known_constants_in_flat(flat, ctx)?; + canonicalize_varrefs_via_record_aliases(flat, ctx); resolve_nested_constructor_field_access_bindings(flat); + collapse_index_refs_to_known_varrefs(flat); functions::prune_unreachable_functions(flat); functions::validate_flat_function_bindings(flat)?; functions::validate_flat_function_call_args(flat)?; @@ -1002,6 +1023,7 @@ pub(crate) fn finalize_flat_model( if options.simplify_variable_names { name_simplify::simplify_flat_names(flat)?; } + collapse_index_refs_to_known_varrefs(flat); // Final boundary pass: every rendered variable reference leaves flatten // with its structured component reference attached, so downstream phases // never re-derive structure from names. @@ -1010,6 +1032,225 @@ pub(crate) fn finalize_flat_model( Ok(()) } +fn materialize_referenced_expandable_connector_members( + flat: &mut flat::Model, + overlay: &ast::InstanceOverlay, + tree: &ast::ClassTree, + class_index: &ast::ClassDefIndex<'_>, +) -> Result<(), FlattenError> { + let prefixes = overlay + .components + .values() + .filter(|instance| instance.is_connector_type) + .filter(|instance| { + instance + .type_def_id + .and_then(|def_id| class_index.get(def_id)) + .is_some_and(|class_def| class_def.expandable) + }) + .map(|instance| { + let span = required_location_span( + &tree.source_map, + &instance.source_location, + "expandable connector member", + )?; + Ok((instance.qualified_name.to_flat_string(), span)) + }) + .collect::, FlattenError>>()?; + let referenced_names = overlay + .classes + .values() + .flat_map(|class_data| &class_data.connections) + .flat_map(|connection| [connection.a.to_flat_string(), connection.b.to_flat_string()]) + .chain(direct_equation_endpoint_varrefs(flat)); + materialize_referenced_expandable_connector_members_for_prefixes( + flat, + &prefixes, + referenced_names, + ); + Ok(()) +} + +fn direct_equation_endpoint_varrefs(flat: &flat::Model) -> Vec { + flat.equations + .iter() + .flat_map(|equation| match &equation.residual { + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs, + rhs, + .. + } => [direct_varref_name(lhs), direct_varref_name(rhs)], + _ => [None, None], + }) + .flatten() + .map(ToOwned::to_owned) + .collect() +} + +fn direct_varref_name(expr: &rumoca_core::Expression) -> Option<&str> { + match expr { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } if subscripts.is_empty() => Some(name.as_str()), + _ => None, + } +} + +fn materialize_referenced_expandable_connector_members_for_prefixes( + flat: &mut flat::Model, + prefixes: &[(String, rumoca_core::Span)], + referenced_names: impl IntoIterator, +) { + if prefixes.is_empty() { + return; + } + + for name in referenced_names { + let name = rumoca_core::VarName::new(name); + if flat.variables.contains_key(&name) { + continue; + } + let Some((_, span)) = prefixes + .iter() + .filter(|(prefix, _)| is_expandable_member_ref(name.as_str(), prefix)) + .max_by_key(|(prefix, _)| prefix.len()) + else { + continue; + }; + let mut var = flat::Variable::empty_with_span(*span); + var.name = name.clone(); + var.component_ref = Some(rumoca_core::ComponentReference::from_flat_segments( + name.as_str(), + *span, + None, + )); + var.type_id = rumoca_core::TypeId::UNKNOWN; + var.is_primitive = true; + var.from_expandable_connector = true; + flat.add_variable(name, var); + } +} + +fn is_expandable_member_ref(name: &str, prefix: &str) -> bool { + name.len() > prefix.len() + && name.as_bytes().get(prefix.len()) == Some(&b'.') + && name.starts_with(prefix) +} + +#[cfg(test)] +mod expandable_connector_member_tests { + use super::*; + + fn var_ref(name: &str) -> rumoca_core::Expression { + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new(name), + subscripts: vec![], + span: rumoca_core::Span::DUMMY, + } + } + + #[test] + fn materializes_referenced_expandable_connector_members_only_under_prefix() { + let mut flat = flat::Model::new(); + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(var_ref("weaBus.TDryBul")), + rhs: Box::new(var_ref("otherBus.TDryBul")), + span: rumoca_core::Span::DUMMY, + }, + rumoca_core::Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + let referenced_names = ["weaBus.TDryBul".to_string(), "otherBus.TDryBul".to_string()]; + materialize_referenced_expandable_connector_members_for_prefixes( + &mut flat, + &[("weaBus".to_string(), rumoca_core::Span::DUMMY)], + referenced_names, + ); + + let var = flat + .variables + .get(&rumoca_core::VarName::new("weaBus.TDryBul")) + .expect("referenced expandable connector member is materialized"); + assert!(var.from_expandable_connector); + assert!(var.is_primitive); + assert!(var.type_id.is_unknown()); + assert!(!flat.variable_type_names.contains_key(&var.name)); + assert!( + !flat + .variables + .contains_key(&rumoca_core::VarName::new("otherBus.TDryBul")) + ); + } + + #[test] + fn does_not_materialize_non_connection_expandable_member_refs() { + let mut flat = flat::Model::new(); + flat.add_equation(flat::Equation::new( + var_ref("weaBus.TDryBul"), + rumoca_core::Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + materialize_referenced_expandable_connector_members_for_prefixes( + &mut flat, + &[("weaBus".to_string(), rumoca_core::Span::DUMMY)], + std::iter::empty(), + ); + + assert!( + !flat + .variables + .contains_key(&rumoca_core::VarName::new("weaBus.TDryBul")) + ); + } + + #[test] + fn materializes_direct_equation_expandable_member_refs() { + let mut flat = flat::Model::new(); + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(var_ref("weaBus.TWetBul")), + rhs: Box::new(var_ref("building.weaBus.TWetBul")), + span: rumoca_core::Span::DUMMY, + }, + rumoca_core::Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "probe".to_string(), + }, + )); + + let referenced_names = direct_equation_endpoint_varrefs(&flat); + materialize_referenced_expandable_connector_members_for_prefixes( + &mut flat, + &[ + ("weaBus".to_string(), rumoca_core::Span::DUMMY), + ("building.weaBus".to_string(), rumoca_core::Span::DUMMY), + ], + referenced_names, + ); + + assert!( + flat.variables + .get(&rumoca_core::VarName::new("weaBus.TWetBul")) + .is_some_and(|var| var.from_expandable_connector) + ); + assert!( + flat.variables + .get(&rumoca_core::VarName::new("building.weaBus.TWetBul")) + .is_some_and(|var| var.from_expandable_connector) + ); + } +} + /// MLS §9.4 / CONN-013: every subgraph of the virtual connection graph needs /// at least one definite or potential root. Tier-1 check: a model that uses /// Connections.branch() but declares no root anywhere cannot satisfy this. diff --git a/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims.rs b/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims.rs index 35f33c16a..4cc8ea6cc 100644 --- a/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims.rs +++ b/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims.rs @@ -158,7 +158,7 @@ pub(crate) fn collect_component_constructor_aliases_for_class( ); } - collect_nested_package_aliases_for_class( + collect_nested_receiver_aliases_for_class( tree, class_index, class_def, @@ -191,7 +191,7 @@ pub(crate) fn collect_component_constructor_aliases_for_class( } } -fn collect_nested_package_aliases_for_class( +fn collect_nested_receiver_aliases_for_class( tree: &ClassTree, class_index: &rumoca_ir_ast::ClassDefIndex<'_>, class_def: &rumoca_ir_ast::ClassDef, @@ -200,31 +200,37 @@ fn collect_nested_package_aliases_for_class( overrides: &mut rustc_hash::FxHashMap, ) { for (alias, nested) in &class_def.classes { - if nested.class_type != rumoca_core::ClassType::Package { + if !is_receiver_alias_type(&nested.class_type) { continue; } let Some(target_ref) = - nested_package_alias_target_ref(tree, class_index, nested, class_scope) + nested_receiver_alias_target_ref(tree, class_index, nested, class_scope) else { continue; }; - if target_ref.class_def.class_type == rumoca_core::ClassType::Package { + if is_receiver_alias_type(&target_ref.class_def.class_type) { let active_alias = active_aliases && leaf_segment(&target_ref.name) != alias; + let modifier_args = nested_receiver_alias_modifier_args(nested); overrides.insert( alias.clone(), - OverrideTarget::from_resolved(alias.clone(), target_ref, active_alias), + OverrideTarget::from_resolved_with_modifier_args( + alias.clone(), + target_ref, + active_alias, + modifier_args, + ), ); } } } -fn nested_package_alias_target_ref<'a>( +fn nested_receiver_alias_target_ref<'a>( tree: &'a ClassTree, class_index: &'a rumoca_ir_ast::ClassDefIndex<'a>, class_def: &rumoca_ir_ast::ClassDef, class_scope: &str, ) -> Option> { - if !is_package_alias_definition(class_def) { + if !is_nested_receiver_alias_definition(class_def) { return None; } let ext = class_def.extends.first()?; @@ -245,7 +251,7 @@ fn nested_package_alias_target_ref<'a>( }) } -fn is_package_alias_definition(class_def: &rumoca_ir_ast::ClassDef) -> bool { +fn is_nested_receiver_alias_definition(class_def: &rumoca_ir_ast::ClassDef) -> bool { class_def.extends.len() == 1 && class_def.imports.is_empty() && class_def.classes.is_empty() @@ -258,6 +264,21 @@ fn is_package_alias_definition(class_def: &rumoca_ir_ast::ClassDef) -> bool { && class_def.external.is_none() } +fn nested_receiver_alias_modifier_args( + class_def: &rumoca_ir_ast::ClassDef, +) -> Vec { + class_def + .extends + .first() + .map(|ext| { + ext.modifications + .iter() + .filter_map(|modification| function_modifier_arg_from_ast(&modification.expr)) + .collect() + }) + .unwrap_or_default() +} + fn collect_extends_redeclare_aliases_for_class( tree: &ClassTree, class_index: &rumoca_ir_ast::ClassDefIndex<'_>, @@ -409,7 +430,7 @@ fn resolve_package_alias_chain<'a>( } let current = resolved_class_ref_for_def_id(tree, class_index, current_def_id)?; if current.class_def.class_type != rumoca_core::ClassType::Package - || !is_package_alias_definition(current.class_def) + || !is_nested_receiver_alias_definition(current.class_def) { return Some(current); } @@ -1137,8 +1158,8 @@ pub(crate) fn resolve_override_function_name( reference: &rumoca_core::Reference, ctx: &FunctionOverrideRewriteContext<'_>, ) -> Option { - let scope = reference.component_scope()?; - let function_leaf = scope.leaf_ident()?; + let function_leaf = reference.last_segment(); + let parent_alias = reference_parent_alias(reference); if let Some(source_package_def_id) = reference_source_package_def_id(reference, ctx) && let Some(package) = ctx.lexical_package_target() && ctx.package_chain_contains_def_id(&package, source_package_def_id) @@ -1152,7 +1173,7 @@ pub(crate) fn resolve_override_function_name( { return Some(resolved); } - if scope.parent_ident().is_none() { + if parent_alias.is_none() { if let Some(source_package_def_id) = reference_source_package_def_id(reference, ctx) && let Some(package) = ctx.concrete_override_package_for_source_package(source_package_def_id) @@ -1214,8 +1235,8 @@ pub(crate) fn resolve_override_function_name( .filter(|resolved| resolved != reference.as_str()); } - let package_alias = scope.parent_ident()?; - let package = ctx.override_package(package_alias)?; + let package_alias = parent_alias?; + let package = ctx.override_package(&package_alias)?; // Only resolve through the alias-matched package when the call's actual // source package is unknown (a genuinely relative reference) or is part of // that package's chain. A fully-qualified call into a *different* package @@ -1232,6 +1253,15 @@ pub(crate) fn resolve_override_function_name( .filter(|resolved| resolved != reference.as_str()) } +fn reference_parent_alias(reference: &rumoca_core::Reference) -> Option { + if let Some(scope) = reference.component_scope() { + return scope.parent_ident().map(ToString::to_string); + } + ComponentPath::from_flat_path(reference.as_str()) + .parent() + .and_then(|path| path.parts().last().cloned()) +} + fn reference_package_scope(reference: &rumoca_core::Reference) -> Option { let scope = reference.component_scope()?; let prefix_parts = scope.prefix_parts(); @@ -1338,6 +1368,38 @@ fn reference_component_ref_is_instance_path( canonical_instance_reference_name(reference, ctx).is_some() } +fn active_scope_relative_instance_reference_name( + reference: &rumoca_core::Reference, + ctx: &FunctionOverrideRewriteContext<'_>, +) -> Option { + let component_members = ctx.component_members?; + let component_ref = reference.component_ref()?; + if component_ref.local || ctx.active_scope.is_root() { + return None; + } + let relative_path = ComponentPath::from_component_reference(component_ref); + let relative_name = relative_path.to_flat_string(); + if relative_name != reference.as_str() { + return None; + } + let scoped_path = ctx.active_scope.join(&relative_path); + if !component_members.contains_component_path(&scoped_path) { + return None; + } + let scoped_name = scoped_path.to_flat_string(); + if ctx + .class_index + .get_by_qualified_name(&scoped_name) + .is_some() + { + return None; + } + Some(rumoca_core::Reference::with_component_reference( + scoped_name, + component_ref.clone(), + )) +} + fn canonical_instance_reference_name( reference: &rumoca_core::Reference, ctx: &FunctionOverrideRewriteContext<'_>, @@ -1545,11 +1607,14 @@ impl<'a> FunctionOverrideRewriteContext<'a> { &self, source_package_def_id: rumoca_core::DefId, ) -> Option<&'a OverrideTarget> { - let mut matches = self.override_packages.iter().filter(|package| { - package.active && self.package_chain_contains_def_id(package, source_package_def_id) - }); - let package = matches.next()?; - matches.next().is_none().then_some(package) + let matches = self + .override_packages + .iter() + .filter(|package| { + package.active && self.package_chain_contains_def_id(package, source_package_def_id) + }) + .collect::>(); + self.preferred_unique_override_package(matches) } fn concrete_override_package_for_source_package( @@ -1561,12 +1626,27 @@ impl<'a> FunctionOverrideRewriteContext<'a> { { return Some(package); } - let mut matches = self + let matches = self .override_packages .iter() - .filter(|package| self.package_chain_contains_def_id(package, source_package_def_id)); - let package = matches.next()?; - matches.next().is_none().then_some(package) + .filter(|package| self.package_chain_contains_def_id(package, source_package_def_id)) + .collect::>(); + self.preferred_unique_override_package(matches) + } + + fn preferred_unique_override_package( + &self, + matches: Vec<&'a OverrideTarget>, + ) -> Option<&'a OverrideTarget> { + if matches.len() == 1 { + return matches.first().copied(); + } + let mut local_medium_matches = matches + .iter() + .copied() + .filter(|package| package.alias == "Medium"); + let package = local_medium_matches.next()?; + local_medium_matches.next().is_none().then_some(package) } fn unique_active_override_package(&self) -> Option<&'a OverrideTarget> { @@ -1823,6 +1903,15 @@ impl ExpressionRewriter for FunctionOverrideExpressionRewriter<'_> { span: *span, }; } + if let Some(canonical_name) = + active_scope_relative_instance_reference_name(name, self.ctx) + { + return Expression::VarRef { + name: canonical_name, + subscripts: rewritten_subscripts, + span: *span, + }; + } if let Some(resolved_name) = resolve_override_member_name(name, self.ctx) { return Expression::VarRef { name: rewritten_reference(name, resolved_name, self.ctx), @@ -1851,7 +1940,8 @@ impl ExpressionRewriter for FunctionOverrideExpressionRewriter<'_> { span, } = expr { - let rewritten_args = self.rewrite_expressions(args); + let rewritten_args = + preserve_named_arg_marker_shells(args, self.rewrite_expressions(args)); if reference_targets_function_local_def(name, self.ctx) { return Expression::FunctionCall { name: name.clone(), @@ -1889,60 +1979,16 @@ impl ExpressionRewriter for FunctionOverrideExpressionRewriter<'_> { impl StatementRewriter for FunctionOverrideExpressionRewriter<'_> {} -fn function_local_def_ids(function: &rumoca_core::Function) -> FxHashSet { - function - .inputs - .iter() - .chain(function.outputs.iter()) - .chain(function.locals.iter()) - .filter_map(|param| param.def_id) - .collect() -} - -fn reference_targets_function_local_def( - reference: &rumoca_core::Reference, - ctx: &FunctionOverrideRewriteContext<'_>, -) -> bool { - reference - .target_def_id() - .is_some_and(|def_id| ctx.local_def_ids.contains(&def_id)) -} - -fn rewritten_reference( - original: &rumoca_core::Reference, - resolved_name: String, - ctx: &FunctionOverrideRewriteContext<'_>, -) -> rumoca_core::Reference { - rewritten_function_reference(original, resolved_name, ctx.tree, ctx.class_index) -} - -fn rewritten_function_reference( - original: &rumoca_core::Reference, - resolved_name: String, - tree: &ClassTree, - class_index: &rumoca_ir_ast::ClassDefIndex<'_>, -) -> rumoca_core::Reference { - let Some(mut component_ref) = original.component_ref().cloned() else { - return rumoca_core::Reference::new(resolved_name); - }; - component_ref.def_id = tree.name_map.get(&resolved_name).copied().or_else(|| { - class_index - .get_by_qualified_name(&resolved_name) - .and_then(|class_def| class_def.def_id) - }); - component_ref.parts = ComponentPath::from_flat_path(&resolved_name) - .parts() - .iter() - .map(|part| rumoca_core::ComponentRefPart { - ident: part.clone(), - span: component_ref.span, - subs: Vec::new(), - }) - .collect(); - rumoca_core::Reference::with_component_reference(resolved_name, component_ref) -} - +mod function_references; +mod named_arg_markers; mod replaceable_modifiers; +use function_references::{ + function_local_def_ids, reference_targets_function_local_def, rewritten_function_reference, + rewritten_reference, +}; +#[cfg(test)] +use named_arg_markers::preserve_named_arg_marker_shell; +use named_arg_markers::preserve_named_arg_marker_shells; use replaceable_modifiers::{append_replaceable_function_modifier_args, single_component_ref_name}; #[cfg(test)] diff --git a/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/flat_rewrite.rs b/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/flat_rewrite.rs index a94294b9d..ba4a76ff8 100644 --- a/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/flat_rewrite.rs +++ b/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/flat_rewrite.rs @@ -56,6 +56,7 @@ pub(crate) fn rewrite_function_overrides_in_statement_with_ctx( *stmt = FunctionOverrideExpressionRewriter { ctx }.rewrite_statement(stmt); } +#[cfg(test)] pub(crate) fn rewrite_function_overrides_in_expression( expr: &mut Expression, tree: &ClassTree, @@ -72,19 +73,23 @@ pub(crate) fn rewrite_function_overrides_in_expression( rewrite_function_overrides_in_expression_with_ctx(expr, &ctx); } -pub(crate) fn rewrite_function_overrides_in_when_clause( +pub(crate) fn rewrite_function_overrides_in_when_clause_scoped( clause: &mut rumoca_ir_flat::WhenClause, tree: &ClassTree, class_index: &rumoca_ir_ast::ClassDefIndex<'_>, override_packages: &[OverrideTarget], override_functions: &OverrideFunctionMap, + active_scope: &ComponentPath, + component_members: &component_member_scope::ComponentMemberScopes, ) { let ctx = FunctionOverrideRewriteContext::new( tree, class_index, override_packages, override_functions, - ); + ) + .with_active_scope(active_scope.clone()) + .with_component_member_scope(component_members); rewrite_function_overrides_in_when_clause_with_ctx(clause, &ctx); } @@ -128,49 +133,29 @@ pub(crate) fn rewrite_function_overrides_in_flattened( class_index: &rumoca_ir_ast::ClassDefIndex<'_>, override_packages: &[OverrideTarget], override_functions: &OverrideFunctionMap, + active_scope: &ComponentPath, + component_members: &component_member_scope::ComponentMemberScopes, ) { + let ctx = FunctionOverrideRewriteContext::new( + tree, + class_index, + override_packages, + override_functions, + ) + .with_active_scope(active_scope.clone()) + .with_component_member_scope(component_members); for equation in &mut flattened.equations { - rewrite_function_overrides_in_expression( - &mut equation.residual, - tree, - class_index, - override_packages, - override_functions, - ); + rewrite_function_overrides_in_expression_with_ctx(&mut equation.residual, &ctx); } for assert_eq in &mut flattened.assert_equations { - rewrite_function_overrides_in_expression( - &mut assert_eq.condition, - tree, - class_index, - override_packages, - override_functions, - ); - rewrite_function_overrides_in_expression( - &mut assert_eq.message, - tree, - class_index, - override_packages, - override_functions, - ); + rewrite_function_overrides_in_expression_with_ctx(&mut assert_eq.condition, &ctx); + rewrite_function_overrides_in_expression_with_ctx(&mut assert_eq.message, &ctx); if let Some(level_expr) = &mut assert_eq.level { - rewrite_function_overrides_in_expression( - level_expr, - tree, - class_index, - override_packages, - override_functions, - ); + rewrite_function_overrides_in_expression_with_ctx(level_expr, &ctx); } } for clause in &mut flattened.when_clauses { - rewrite_function_overrides_in_when_clause( - clause, - tree, - class_index, - override_packages, - override_functions, - ); + rewrite_function_overrides_in_when_clause_with_ctx(clause, &ctx); } } diff --git a/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/function_references.rs b/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/function_references.rs new file mode 100644 index 000000000..74906918c --- /dev/null +++ b/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/function_references.rs @@ -0,0 +1,56 @@ +use super::*; + +pub(super) fn function_local_def_ids( + function: &rumoca_core::Function, +) -> FxHashSet { + function + .inputs + .iter() + .chain(function.outputs.iter()) + .chain(function.locals.iter()) + .filter_map(|param| param.def_id) + .collect() +} + +pub(super) fn reference_targets_function_local_def( + reference: &rumoca_core::Reference, + ctx: &FunctionOverrideRewriteContext<'_>, +) -> bool { + reference + .target_def_id() + .is_some_and(|def_id| ctx.local_def_ids.contains(&def_id)) +} + +pub(super) fn rewritten_reference( + original: &rumoca_core::Reference, + resolved_name: String, + ctx: &FunctionOverrideRewriteContext<'_>, +) -> rumoca_core::Reference { + rewritten_function_reference(original, resolved_name, ctx.tree, ctx.class_index) +} + +pub(super) fn rewritten_function_reference( + original: &rumoca_core::Reference, + resolved_name: String, + tree: &ClassTree, + class_index: &rumoca_ir_ast::ClassDefIndex<'_>, +) -> rumoca_core::Reference { + let Some(mut component_ref) = original.component_ref().cloned() else { + return rumoca_core::Reference::new(resolved_name); + }; + component_ref.def_id = tree.name_map.get(&resolved_name).copied().or_else(|| { + class_index + .get_by_qualified_name(&resolved_name) + .and_then(|class_def| class_def.def_id) + }); + component_ref.parts = ComponentPath::from_flat_path(&resolved_name) + .parts() + .iter() + .map(|part| rumoca_core::ComponentRefPart { + ident: part.clone(), + span: component_ref.span, + subs: Vec::new(), + }) + .collect(); + rumoca_core::Reference::with_component_reference(resolved_name, component_ref) +} diff --git a/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/member_calls.rs b/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/member_calls.rs index 4dfdf6f47..bac1312be 100644 --- a/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/member_calls.rs +++ b/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/member_calls.rs @@ -27,6 +27,36 @@ impl ExpressionTransformer for QualifyReplaceableFunctionModifier<'_> { .collect(); prefixed_parts.extend(cr.parts); cr.parts = prefixed_parts; + } else if cr.parts.len() > 1 + && !cr.local + && !self.receiver_alias.is_root() + && let Some(receiver_leaf) = self.receiver_alias.parts().last() + { + let already_scoped = cr.parts.len() >= self.receiver_alias.parts().len() + && cr + .parts + .iter() + .zip(self.receiver_alias.parts()) + .all(|(part, receiver_part)| part.ident.text.as_ref() == receiver_part); + if !already_scoped && cr.parts[0].ident.text.as_ref() == receiver_leaf { + let location = cr.parts[0].ident.location.clone(); + let mut scoped_parts: Vec<_> = self + .receiver_alias + .parts() + .iter() + .map(|part| rumoca_ir_ast::ComponentRefPart { + ident: Token { + text: std::sync::Arc::from(part.as_str()), + location: location.clone(), + token_number: 0, + token_type: 0, + }, + subs: None, + }) + .collect(); + scoped_parts.extend(cr.parts.into_iter().skip(1)); + cr.parts = scoped_parts; + } } for part in &mut cr.parts { if let Some(subscripts) = &mut part.subs { diff --git a/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/named_arg_markers.rs b/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/named_arg_markers.rs new file mode 100644 index 000000000..d1d2b098e --- /dev/null +++ b/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/named_arg_markers.rs @@ -0,0 +1,48 @@ +use rumoca_core::Expression; + +pub(super) fn preserve_named_arg_marker_shells( + original_args: &[Expression], + rewritten_args: Vec, +) -> Vec { + original_args + .iter() + .zip(rewritten_args) + .map(|(original, rewritten)| preserve_named_arg_marker_shell(original, rewritten)) + .collect() +} + +pub(super) fn preserve_named_arg_marker_shell( + original: &Expression, + rewritten: Expression, +) -> Expression { + let Expression::FunctionCall { + name, + is_constructor: true, + span, + .. + } = original + else { + return rewritten; + }; + if !name + .as_str() + .starts_with(rumoca_core::NAMED_FUNCTION_ARG_PREFIX) + { + return rewritten; + } + if matches!( + &rewritten, + Expression::FunctionCall { + name: rewritten_name, + .. + } if rewritten_name.as_str() == name.as_str() + ) { + return rewritten; + } + Expression::FunctionCall { + name: name.clone(), + args: vec![rewritten], + is_constructor: true, + span: *span, + } +} diff --git a/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/replaceable_modifiers.rs b/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/replaceable_modifiers.rs index 0152abf8b..d1e75a5ee 100644 --- a/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/replaceable_modifiers.rs +++ b/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/replaceable_modifiers.rs @@ -8,10 +8,7 @@ pub(super) fn append_replaceable_function_modifier_args( ) -> Vec { let receiver_alias = receiver_alias_for_member_function(current_ref, ctx); let existing_names = named_function_arg_names(&args); - let declaration_receiver_scope = receiver_alias - .as_deref() - .map(ComponentPath::from_flat_path) - .unwrap_or_else(|| ctx.active_scope.clone()); + let declaration_receiver_scope = default_arg_receiver_scope(receiver_alias.as_deref(), ctx); if let Some(default_args) = replaceable_function_modifier_args( current_ref.as_str(), resolved_name, @@ -46,6 +43,27 @@ pub(super) fn append_replaceable_function_modifier_args( args } +fn default_arg_receiver_scope( + receiver_alias: Option<&str>, + ctx: &FunctionOverrideRewriteContext<'_>, +) -> ComponentPath { + let Some(receiver_alias) = receiver_alias else { + return ctx.active_scope.clone(); + }; + let receiver_scope = ComponentPath::from_flat_path(receiver_alias); + if let Some(component_members) = ctx.component_members { + let mut scope = Some(ctx.active_scope.clone()); + while let Some(candidate_scope) = scope { + let scoped_receiver = candidate_scope.join(&receiver_scope); + if component_members.contains_component_path(&scoped_receiver) { + return scoped_receiver; + } + scope = candidate_scope.parent(); + } + } + receiver_scope +} + fn override_function_target_and_receiver_scope<'a>( current_ref: &rumoca_core::Reference, ctx: &'a FunctionOverrideRewriteContext<'a>, diff --git a/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/tests.rs b/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/tests.rs index fa66040b8..5f67dcd1f 100644 --- a/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/tests.rs +++ b/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/tests.rs @@ -248,6 +248,60 @@ fn active_package_member_rewrite_keeps_structured_instance_path() { assert_eq!(name.as_str(), "tank.medium.state.p"); } +#[test] +fn active_scope_relative_instance_reference_is_canonicalized_before_member_rewrite() { + let package_def = DefId::new(1); + let member_def = DefId::new(2); + let mut member = class("jointUSP", ClassType::Record); + member.def_id = Some(member_def); + let mut package = class("AliasPackage", ClassType::Package); + package.def_id = Some(package_def); + package.classes.insert("jointUSP".to_string(), member); + + let mut tree = ClassTree::new(); + tree.definitions + .classes + .insert("AliasPackage".to_string(), package); + tree.def_map.insert(package_def, "AliasPackage".to_string()); + tree.def_map + .insert(member_def, "AliasPackage.jointUSP".to_string()); + + let class_index = rumoca_ir_ast::ClassDefIndex::from_tree(&tree); + let override_packages = vec![override_target( + "AliasPackage", + package_def, + ClassType::Package, + )]; + let override_functions = OverrideFunctionMap::default(); + let mut component_members = component_member_scope::ComponentMemberScopes::default(); + component_members + .insert_component_member_path(&ComponentPath::from_flat_path("jointRRP.jointUSP.e2_ia")); + component_members + .insert_component_member_path(&ComponentPath::from_flat_path("jointRRP.rod1.e2_ia")); + let ctx = FunctionOverrideRewriteContext::new( + &tree, + &class_index, + &override_packages, + &override_functions, + ) + .with_active_scope(ComponentPath::from_flat_path("jointRRP")) + .with_component_member_scope(&component_members); + let mut expr = Expression::VarRef { + name: rumoca_core::Reference::from_component_reference(core_comp_ref(&[ + "jointUSP", "e2_ia", + ])), + subscripts: vec![], + span: Span::DUMMY, + }; + + rewrite_function_overrides_in_expression_with_ctx(&mut expr, &ctx); + + let Expression::VarRef { name, .. } = expr else { + panic!("expected var ref"); + }; + assert_eq!(name.as_str(), "jointRRP.jointUSP.e2_ia"); +} + #[test] fn fully_qualified_sibling_package_call_is_not_aliased_to_self() { // Regression: a function `A.Quat.inverse` that calls the fully-qualified @@ -754,16 +808,30 @@ fn replaceable_function_alias_preserves_modifier_actuals() { gravity.extends.push(Extend { base_name: Name::from_string("Standard"), base_def_id: Some(standard_def), - modifications: vec![rumoca_ir_ast::ExtendModification { - expr: rumoca_ir_ast::Expression::Modification { - target: comp_ref(&["gravityType"]), - value: Arc::new(ast_var("gravityType")), - span: test_span(), + modifications: vec![ + rumoca_ir_ast::ExtendModification { + expr: rumoca_ir_ast::Expression::Modification { + target: comp_ref(&["gravityType"]), + value: Arc::new(ast_var("gravityType")), + span: test_span(), + }, + each: false, + final_: false, + redeclare: false, }, - each: false, - final_: false, - redeclare: false, - }], + rumoca_ir_ast::ExtendModification { + expr: rumoca_ir_ast::Expression::Modification { + target: comp_ref(&["g"]), + value: Arc::new(rumoca_ir_ast::Expression::ComponentReference(comp_ref(&[ + "world", "g", + ]))), + span: test_span(), + }, + each: false, + final_: false, + redeclare: false, + }, + ], ..Extend::default() }); @@ -813,11 +881,71 @@ fn replaceable_function_alias_preserves_modifier_actuals() { panic!("expected rewritten function call"); }; assert_eq!(name.as_str(), "Standard"); - assert_eq!(args.len(), 2); + assert_eq!(args.len(), 3); let Some(("gravityType", Expression::VarRef { name, .. })) = named_arg(&args[1]) else { panic!("expected receiver-qualified gravityType named argument"); }; assert_eq!(name.as_str(), "world.gravityType"); + let Some(("g", Expression::VarRef { name, .. })) = named_arg(&args[2]) else { + panic!("expected receiver-relative g named argument"); + }; + assert_eq!(name.as_str(), "world.g"); + + let mut component_members = component_member_scope::ComponentMemberScopes::default(); + component_members + .insert_component_member_path(&ComponentPath::from_flat_path("mechanics.world")); + component_members.insert_component_member_path(&ComponentPath::from_flat_path( + "mechanics.world.gravityType", + )); + let scoped_ctx = + FunctionOverrideRewriteContext::new(&tree, &class_index, &[], &override_functions) + .with_active_scope(ComponentPath::from_flat_path("mechanics")) + .with_component_member_scope(&component_members); + let mut scoped_expr = Expression::FunctionCall { + name: rumoca_core::Reference::new("World.gravityAcceleration"), + args: vec![core_var("r")], + is_constructor: false, + span: Span::DUMMY, + }; + + rewrite_function_overrides_in_expression_with_ctx(&mut scoped_expr, &scoped_ctx); + + let Expression::FunctionCall { args, .. } = scoped_expr else { + panic!("expected rewritten function call"); + }; + let Some(("gravityType", Expression::VarRef { name, .. })) = named_arg(&args[1]) else { + panic!("expected receiver-qualified scoped gravityType named argument"); + }; + assert_eq!(name.as_str(), "mechanics.world.gravityType"); + let Some(("g", Expression::VarRef { name, .. })) = named_arg(&args[2]) else { + panic!("expected receiver-qualified scoped g named argument"); + }; + assert_eq!(name.as_str(), "mechanics.world.g"); + + let nested_ctx = + FunctionOverrideRewriteContext::new(&tree, &class_index, &[], &override_functions) + .with_active_scope(ComponentPath::from_flat_path("mechanics.b0.body")) + .with_component_member_scope(&component_members); + let mut nested_expr = Expression::FunctionCall { + name: rumoca_core::Reference::new("World.gravityAcceleration"), + args: vec![core_var("r")], + is_constructor: false, + span: Span::DUMMY, + }; + + rewrite_function_overrides_in_expression_with_ctx(&mut nested_expr, &nested_ctx); + + let Expression::FunctionCall { args, .. } = nested_expr else { + panic!("expected rewritten function call"); + }; + let Some(("gravityType", Expression::VarRef { name, .. })) = named_arg(&args[1]) else { + panic!("expected receiver-qualified ancestor gravityType named argument"); + }; + assert_eq!(name.as_str(), "mechanics.world.gravityType"); + let Some(("g", Expression::VarRef { name, .. })) = named_arg(&args[2]) else { + panic!("expected receiver-qualified ancestor g named argument"); + }; + assert_eq!(name.as_str(), "mechanics.world.g"); } #[test] @@ -937,15 +1065,56 @@ fn inherited_replaceable_function_call_keeps_declaration_modifier_actuals() { assert_eq!(*value, rumoca_core::Literal::Real(0.8)); } +#[test] +fn component_scope_inherits_nested_replaceable_function_default_alias() { + let (tree, _) = replaceable_efficiency_fixture(); + let class_index = rumoca_ir_ast::ClassDefIndex::from_tree(&tree); + let partial_pump_def = class_index + .get_by_qualified_name("Modelica.Fluid.Machines.BaseClasses.PartialPump") + .and_then(|class_def| class_def.def_id) + .expect("partial pump def id"); + let mut overlay = InstanceOverlay::new(); + let pump_id = overlay.alloc_id(); + overlay.add_component(rumoca_ir_ast::InstanceData { + instance_id: pump_id, + qualified_name: QualifiedName::from_ident("pump"), + type_def_id: Some(partial_pump_def), + ..rumoca_ir_ast::InstanceData::default() + }); + + let override_map = build_component_override_map( + &overlay, + &tree, + &class_index, + "Modelica.Fluid.Machines.BaseClasses.PartialPump", + ) + .expect("component override map"); + let (_, override_functions) = override_context_for_scope("pump", &override_map); + let target = override_functions + .get("efficiencyCharacteristic") + .expect("expected inherited replaceable function default in pump scope"); + + assert_eq!( + target.name, + "Modelica.Fluid.Machines.BaseClasses.PumpCharacteristics.constantEfficiency" + ); + assert_eq!(target.modifier_args.len(), 1); + assert_eq!(target.modifier_args[0].name, "eta_nominal"); +} + fn replaceable_efficiency_fixture() -> (ClassTree, DefId) { let base_efficiency_def = DefId::new(1); let constant_efficiency_def = DefId::new(2); let efficiency_characteristic_def = DefId::new(3); + let partial_pump_def = DefId::new(4); let pump_characteristics = replaceable_efficiency_pump_characteristics(base_efficiency_def, constant_efficiency_def); - let partial_pump = - replaceable_efficiency_partial_pump(efficiency_characteristic_def, constant_efficiency_def); + let partial_pump = replaceable_efficiency_partial_pump( + partial_pump_def, + efficiency_characteristic_def, + constant_efficiency_def, + ); let modelica = replaceable_efficiency_modelica_tree(pump_characteristics, partial_pump); let mut tree = ClassTree::new(); tree.definitions @@ -955,6 +1124,7 @@ fn replaceable_efficiency_fixture() -> (ClassTree, DefId) { &mut tree, base_efficiency_def, constant_efficiency_def, + partial_pump_def, efficiency_characteristic_def, ); (tree, efficiency_characteristic_def) @@ -984,6 +1154,7 @@ fn replaceable_efficiency_pump_characteristics( } fn replaceable_efficiency_partial_pump( + partial_pump_def: DefId, efficiency_characteristic_def: DefId, constant_efficiency_def: DefId, ) -> ClassDef { @@ -1000,6 +1171,7 @@ fn replaceable_efficiency_partial_pump( }); let mut partial_pump = class("PartialPump", ClassType::Model); + partial_pump.def_id = Some(partial_pump_def); partial_pump.classes.insert( "efficiencyCharacteristic".to_string(), efficiency_characteristic, @@ -1050,6 +1222,7 @@ fn register_replaceable_efficiency_names( tree: &mut ClassTree, base_efficiency_def: DefId, constant_efficiency_def: DefId, + partial_pump_def: DefId, efficiency_characteristic_def: DefId, ) { tree.def_map.insert( @@ -1060,6 +1233,10 @@ fn register_replaceable_efficiency_names( constant_efficiency_def, "Modelica.Fluid.Machines.BaseClasses.PumpCharacteristics.constantEfficiency".to_string(), ); + tree.def_map.insert( + partial_pump_def, + "Modelica.Fluid.Machines.BaseClasses.PartialPump".to_string(), + ); tree.def_map.insert( efficiency_characteristic_def, "Modelica.Fluid.Machines.BaseClasses.PartialPump.efficiencyCharacteristic".to_string(), @@ -1730,125 +1907,37 @@ fn root_component_override_map(alias_target: &OverrideTarget) -> ComponentOverri } #[test] -fn marks_member_function_calls_through_component_type_aliases() { - let gravity_def = DefId::new(1); - let mut world = class("World", ClassType::Model); - let mut gravity = class("gravityAcceleration", ClassType::Function); - gravity.def_id = Some(gravity_def); - world - .classes - .insert("gravityAcceleration".to_string(), gravity); - - let mut tree = ClassTree::new(); - tree.definitions.classes.insert("World".to_string(), world); - tree.name_map - .insert("World.gravityAcceleration".to_string(), gravity_def); - - let mut override_functions = OverrideFunctionMap::default(); - let world_def = DefId::new(2); - let Some(world) = tree.definitions.classes.get_mut("World") else { - panic!("expected World class"); - }; - world.def_id = Some(world_def); - tree.def_map.insert(world_def, "World".to_string()); - override_functions.insert( - "world".to_string(), - override_target("World", world_def, ClassType::Model), - ); - let class_index = rumoca_ir_ast::ClassDefIndex::from_tree(&tree); - let marker = MemberFunctionCallMarker { - tree: &tree, - class_index: &class_index, - override_functions: &override_functions, - }; - - assert_eq!( - marker - .mark_component_function_call(comp_ref(&["world", "gravityAcceleration"])) - .def_id, - Some(gravity_def) - ); -} - -#[test] -fn root_package_alias_marks_member_function_calls() { - let (tree, ids) = concrete_override_chain_tree(); - let class_index = rumoca_ir_ast::ClassDefIndex::from_tree(&tree); - let alias_target = override_target("AliasMedium", ids.alias_pkg, ClassType::Package); - let component_override_map = root_component_override_map(&OverrideTarget { - alias: "Medium".to_string(), - ..alias_target - }); - let (_, override_functions) = override_context_for_scope("", &component_override_map); - let marker = MemberFunctionCallMarker { - tree: &tree, - class_index: &class_index, - override_functions: &override_functions, +fn preserves_named_arg_marker_shell_when_rewritten_value_changes_shape() { + let original = Expression::FunctionCall { + name: rumoca_core::Reference::new("__rumoca_named_arg__.per"), + args: vec![core_var("pCur1")], + is_constructor: true, + span: test_span(), }; - - assert_eq!( - marker - .mark_component_function_call(comp_ref(&["Medium", "density"])) - .def_id, - Some(ids.concrete_density) - ); -} - -#[test] -fn active_package_alias_rewrites_inherited_partial_function_call() { - let (tree, ids) = concrete_override_chain_tree(); - let class_index = rumoca_ir_ast::ClassDefIndex::from_tree(&tree); - let override_packages = vec![OverrideTarget { - alias: "Medium".to_string(), - ..override_target("AliasMedium", ids.alias_pkg, ClassType::Package) - }]; - let override_functions = OverrideFunctionMap::default(); - let mut expr = Expression::FunctionCall { - name: rumoca_core::Reference::with_component_reference( - "PartialMedium.density", - rumoca_core::ComponentReference { - def_id: Some(ids.partial_density), - ..core_comp_ref(&["PartialMedium", "density"]) - }, - ), - args: vec![core_var("state")], - is_constructor: false, - span: Span::DUMMY, + let rewritten = Expression::FunctionCall { + name: rumoca_core::Reference::new("Buildings.Fluid.Movers.Data.Generic"), + args: vec![core_var("pCur1.V_flow"), core_var("pCur1.dp")], + is_constructor: true, + span: test_span(), }; - rewrite_function_overrides_in_expression( - &mut expr, - &tree, - &class_index, - &override_packages, - &override_functions, - ); + let preserved = preserve_named_arg_marker_shell(&original, rewritten); - let Expression::FunctionCall { name, .. } = expr else { - panic!("expected rewritten function call"); + let Expression::FunctionCall { + name, + args, + is_constructor: true, + .. + } = preserved + else { + panic!("expected named argument marker"); }; - assert_eq!(name.as_str(), "AliasMedium.density"); + assert_eq!(name.as_str(), "__rumoca_named_arg__.per"); + assert!(matches!( + args.as_slice(), + [Expression::FunctionCall { name, is_constructor: true, .. }] + if name.as_str() == "Buildings.Fluid.Movers.Data.Generic" + )); } -#[test] -fn leaves_unknown_member_function_calls_unmarked() { - let tree = ClassTree::new(); - let mut override_functions = OverrideFunctionMap::default(); - override_functions.insert( - "world".to_string(), - override_target("World", DefId::new(1), ClassType::Model), - ); - let class_index = rumoca_ir_ast::ClassDefIndex::from_tree(&tree); - let marker = MemberFunctionCallMarker { - tree: &tree, - class_index: &class_index, - override_functions: &override_functions, - }; - - assert_eq!( - marker - .mark_component_function_call(comp_ref(&["world", "gravityAcceleration"])) - .def_id, - None - ); -} +mod member_call_tests; diff --git a/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/tests/member_call_tests.rs b/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/tests/member_call_tests.rs new file mode 100644 index 000000000..ffe2c50e5 --- /dev/null +++ b/crates/rumoca-phase-flatten/src/pipeline/function_overrides_and_dims/tests/member_call_tests.rs @@ -0,0 +1,161 @@ +use super::*; + +#[test] +fn marks_member_function_calls_through_component_type_aliases() { + let gravity_def = DefId::new(1); + let mut world = class("World", ClassType::Model); + let mut gravity = class("gravityAcceleration", ClassType::Function); + gravity.def_id = Some(gravity_def); + world + .classes + .insert("gravityAcceleration".to_string(), gravity); + + let mut tree = ClassTree::new(); + tree.definitions.classes.insert("World".to_string(), world); + tree.name_map + .insert("World.gravityAcceleration".to_string(), gravity_def); + + let mut override_functions = OverrideFunctionMap::default(); + let world_def = DefId::new(2); + let Some(world) = tree.definitions.classes.get_mut("World") else { + panic!("expected World class"); + }; + world.def_id = Some(world_def); + tree.def_map.insert(world_def, "World".to_string()); + override_functions.insert( + "world".to_string(), + override_target("World", world_def, ClassType::Model), + ); + let class_index = rumoca_ir_ast::ClassDefIndex::from_tree(&tree); + let marker = MemberFunctionCallMarker { + tree: &tree, + class_index: &class_index, + override_functions: &override_functions, + }; + + assert_eq!( + marker + .mark_component_function_call(comp_ref(&["world", "gravityAcceleration"])) + .def_id, + Some(gravity_def) + ); +} + +#[test] +fn root_package_alias_marks_member_function_calls() { + let (tree, ids) = concrete_override_chain_tree(); + let class_index = rumoca_ir_ast::ClassDefIndex::from_tree(&tree); + let alias_target = override_target("AliasMedium", ids.alias_pkg, ClassType::Package); + let component_override_map = root_component_override_map(&OverrideTarget { + alias: "Medium".to_string(), + ..alias_target + }); + let (_, override_functions) = override_context_for_scope("", &component_override_map); + let marker = MemberFunctionCallMarker { + tree: &tree, + class_index: &class_index, + override_functions: &override_functions, + }; + + assert_eq!( + marker + .mark_component_function_call(comp_ref(&["Medium", "density"])) + .def_id, + Some(ids.concrete_density) + ); +} + +#[test] +fn active_package_alias_rewrites_inherited_partial_function_call() { + let (tree, ids) = concrete_override_chain_tree(); + let class_index = rumoca_ir_ast::ClassDefIndex::from_tree(&tree); + let override_packages = vec![OverrideTarget { + alias: "Medium".to_string(), + ..override_target("AliasMedium", ids.alias_pkg, ClassType::Package) + }]; + let override_functions = OverrideFunctionMap::default(); + let mut expr = Expression::FunctionCall { + name: rumoca_core::Reference::with_component_reference( + "PartialMedium.density", + rumoca_core::ComponentReference { + def_id: Some(ids.partial_density), + ..core_comp_ref(&["PartialMedium", "density"]) + }, + ), + args: vec![core_var("state")], + is_constructor: false, + span: Span::DUMMY, + }; + + rewrite_function_overrides_in_expression( + &mut expr, + &tree, + &class_index, + &override_packages, + &override_functions, + ); + + let Expression::FunctionCall { name, .. } = expr else { + panic!("expected rewritten function call"); + }; + assert_eq!(name.as_str(), "AliasMedium.density"); +} + +#[test] +fn fully_qualified_partial_function_prefers_local_medium_alias_when_multiple_media_match() { + let (tree, ids) = concrete_override_chain_tree(); + let class_index = rumoca_ir_ast::ClassDefIndex::from_tree(&tree); + let override_packages = vec![ + OverrideTarget { + alias: "Medium".to_string(), + ..override_target("AliasMedium", ids.alias_pkg, ClassType::Package) + }, + OverrideTarget { + alias: "MediumAir".to_string(), + ..override_target("ConcreteMedium", ids.concrete_pkg, ClassType::Package) + }, + ]; + let override_functions = OverrideFunctionMap::default(); + let ctx = FunctionOverrideRewriteContext::new( + &tree, + &class_index, + &override_packages, + &override_functions, + ); + let mut expr = Expression::FunctionCall { + name: rumoca_core::Reference::new("PartialMedium.density"), + args: vec![core_var("state")], + is_constructor: false, + span: Span::DUMMY, + }; + + rewrite_function_overrides_in_expression_with_ctx(&mut expr, &ctx); + + let Expression::FunctionCall { name, .. } = expr else { + panic!("expected rewritten function call"); + }; + assert_eq!(name.as_str(), "AliasMedium.density"); +} + +#[test] +fn leaves_unknown_member_function_calls_unmarked() { + let tree = ClassTree::new(); + let mut override_functions = OverrideFunctionMap::default(); + override_functions.insert( + "world".to_string(), + override_target("World", DefId::new(1), ClassType::Model), + ); + let class_index = rumoca_ir_ast::ClassDefIndex::from_tree(&tree); + let marker = MemberFunctionCallMarker { + tree: &tree, + class_index: &class_index, + override_functions: &override_functions, + }; + + assert_eq!( + marker + .mark_component_function_call(comp_ref(&["world", "gravityAcceleration"])) + .def_id, + None + ); +} diff --git a/crates/rumoca-phase-flatten/src/pipeline/lookup_scopes.rs b/crates/rumoca-phase-flatten/src/pipeline/lookup_scopes.rs new file mode 100644 index 000000000..3537b2fa6 --- /dev/null +++ b/crates/rumoca-phase-flatten/src/pipeline/lookup_scopes.rs @@ -0,0 +1,114 @@ +use super::*; + +pub(crate) fn scoped_lookup_candidates(name: &str, scope: &str) -> Vec { + scoped_lookup_candidates_with_scope(name, scope) + .into_iter() + .map(|(candidate, _candidate_scope)| candidate) + .collect() +} + +pub(crate) fn scoped_lookup_candidates_with_scope( + name: &str, + scope: &str, +) -> Vec<(String, String)> { + let name_path = rumoca_core::ComponentPath::from_flat_path(name); + let mut candidates = Vec::new(); + let mut current_scope = Some(rumoca_core::ComponentPath::from_flat_path(scope)); + while let Some(scope_path) = current_scope { + candidates.push(( + scope_path.join(&name_path).to_flat_string(), + scope_path.to_flat_string(), + )); + current_scope = scope_path.parent(); + } + if !scope.is_empty() { + candidates.push((name_path.to_flat_string(), String::new())); + } + candidates +} + +impl rumoca_core::EvalLookup for Context { + fn lookup_integer(&self, name: &str, scope: &str) -> Option { + for candidate in scoped_lookup_candidates(name, scope) { + if let Some(value) = self.get_integer_param(&candidate) { + return Some(value); + } + } + + if crate::path_utils::is_nested_name(name) { + if let Some(value) = lookup_with_scope(name, scope, &self.parameter_values) { + return Some(value); + } + if let Some(value) = lookup_with_scope(name, scope, &self.real_parameter_values) + && value.is_finite() + && value.fract() == 0.0 + { + return Some(value as i64); + } + } + None + } + + fn lookup_real(&self, name: &str, scope: &str) -> Option { + for candidate in scoped_lookup_candidates(name, scope) { + if let Some(value) = self.real_parameter_values.get(&candidate).copied() { + return Some(value); + } + + let resolved = self.resolve_alias(&candidate); + if resolved != candidate + && let Some(value) = self.real_parameter_values.get(&resolved).copied() + { + return Some(value); + } + + if let Some(value) = self.get_integer_param(&candidate) { + return Some(value as f64); + } + } + + if crate::path_utils::is_nested_name(name) { + if let Some(value) = lookup_with_scope(name, scope, &self.real_parameter_values) { + return Some(value); + } + if let Some(value) = lookup_with_scope(name, scope, &self.parameter_values) { + return Some(value as f64); + } + } + None + } + + fn lookup_boolean(&self, name: &str, scope: &str) -> Option { + for candidate in scoped_lookup_candidates(name, scope) { + if let Some(value) = self.get_boolean_param(&candidate) { + return Some(value); + } + } + + if crate::path_utils::is_nested_name(name) { + return lookup_with_scope(name, scope, &self.boolean_parameter_values); + } + None + } + + fn lookup_enum<'a>(&'a self, name: &str, scope: &str) -> Option> { + for candidate in scoped_lookup_candidates(name, scope) { + if let Some(value) = self.enum_parameter_values.get(&candidate) { + return Some(std::borrow::Cow::Borrowed(value.as_str())); + } + + let resolved = self.resolve_alias(&candidate); + if resolved != candidate + && let Some(value) = self.enum_parameter_values.get(&resolved) + { + return Some(std::borrow::Cow::Borrowed(value.as_str())); + } + } + + if crate::path_utils::is_nested_name(name) { + return lookup_with_scope(name, scope, &self.enum_parameter_values) + .map(std::borrow::Cow::Owned); + } + None + } +} diff --git a/crates/rumoca-phase-flatten/src/pipeline/mat_resources.rs b/crates/rumoca-phase-flatten/src/pipeline/mat_resources.rs new file mode 100644 index 000000000..674b4028d --- /dev/null +++ b/crates/rumoca-phase-flatten/src/pipeline/mat_resources.rs @@ -0,0 +1,210 @@ +use std::{ + io::Read, + path::{Path, PathBuf}, +}; + +pub(super) fn read_mat_matrix_size(file_name: &str, matrix_name: &str) -> Option<(i64, i64)> { + let path = resolve_modelica_resource_path(file_name)?; + let bytes = std::fs::read(path).ok()?; + read_mat_v4_matrix_size(&bytes, matrix_name) + .or_else(|| read_mat_v5_matrix_size(&bytes, matrix_name)) +} + +pub(super) fn resolve_modelica_resource_path(raw: &str) -> Option { + let raw_path = Path::new(raw); + if raw_path.is_file() { + return Some(raw_path.to_path_buf()); + } + let rest = raw.strip_prefix("modelica://Modelica/")?; + modelica_source_roots() + .into_iter() + .map(|root| root.join(rest)) + .find(|candidate| candidate.is_file()) +} + +fn modelica_source_roots() -> Vec { + let mut roots = Vec::new(); + for msl_base in modelica_cache_dirs() { + let Ok(msl_entries) = std::fs::read_dir(msl_base) else { + continue; + }; + for msl_entry in msl_entries.flatten() { + let msl_path = msl_entry.path(); + let Some(msl_name) = msl_path.file_name().and_then(|name| name.to_str()) else { + continue; + }; + if !msl_name.starts_with("ModelicaStandardLibrary-") { + continue; + } + let Ok(version_entries) = std::fs::read_dir(&msl_path) else { + continue; + }; + for version_entry in version_entries.flatten() { + let version_path = version_entry.path(); + let Some(version_name) = version_path.file_name().and_then(|name| name.to_str()) + else { + continue; + }; + if version_name.starts_with("Modelica ") && version_path.is_dir() { + roots.push(version_path); + } + } + } + } + roots.sort(); + roots.dedup(); + roots +} + +fn modelica_cache_dirs() -> Vec { + let mut dirs = Vec::new(); + let Ok(current_dir) = std::env::current_dir() else { + return dirs; + }; + for ancestor in current_dir.ancestors() { + let candidate = ancestor.join("target").join("msl"); + if candidate.is_dir() { + dirs.push(candidate); + } + } + dirs +} + +fn read_mat_v4_matrix_size(bytes: &[u8], matrix_name: &str) -> Option<(i64, i64)> { + let mut offset = 0usize; + while offset.checked_add(20)? <= bytes.len() { + let mopt = read_i32_le(bytes, offset)?; + let rows = read_i32_le(bytes, offset + 4)?; + let cols = read_i32_le(bytes, offset + 8)?; + let imagf = read_i32_le(bytes, offset + 12)?; + let name_len = read_i32_le(bytes, offset + 16)?; + if rows <= 0 || cols <= 0 || name_len <= 0 || imagf < 0 { + return None; + } + let name_start = offset + 20; + let name_end = name_start.checked_add(name_len as usize)?; + if name_end > bytes.len() { + return None; + } + let raw_name = &bytes[name_start..name_end]; + let name = std::str::from_utf8(raw_name.strip_suffix(&[0]).unwrap_or(raw_name)).ok()?; + if name == matrix_name { + return Some((i64::from(rows), i64::from(cols))); + } + let value_size = mat_v4_value_size(mopt)?; + let value_count = (rows as usize) + .checked_mul(cols as usize)? + .checked_mul(if imagf == 0 { 1 } else { 2 })?; + offset = name_end.checked_add(value_count.checked_mul(value_size)?)?; + } + None +} + +fn mat_v4_value_size(mopt: i32) -> Option { + match (mopt / 10) % 10 { + 0 => Some(8), + 1 | 2 => Some(4), + 3 | 4 => Some(2), + 5 => Some(1), + _ => None, + } +} + +fn read_mat_v5_matrix_size(bytes: &[u8], matrix_name: &str) -> Option<(i64, i64)> { + if bytes.len() < 128 { + return None; + } + read_mat_v5_elements(bytes, 128, matrix_name) +} + +fn read_mat_v5_elements(bytes: &[u8], start: usize, matrix_name: &str) -> Option<(i64, i64)> { + let mut offset = start; + while offset.checked_add(8)? <= bytes.len() { + let tag = read_mat_v5_tag(bytes, offset)?; + match tag.data_type { + 14 => { + if let Some(size) = read_mat_v5_matrix_element(bytes, tag.payload, matrix_name) { + return Some(size); + } + } + 15 => { + let mut decoded = Vec::new(); + flate2::read::ZlibDecoder::new(&bytes[tag.payload.clone()]) + .read_to_end(&mut decoded) + .ok()?; + if let Some(size) = read_mat_v5_elements(&decoded, 0, matrix_name) { + return Some(size); + } + } + _ => {} + } + offset = offset.checked_add(tag.total_size)?; + } + None +} + +fn read_mat_v5_matrix_element( + bytes: &[u8], + payload: std::ops::Range, + matrix_name: &str, +) -> Option<(i64, i64)> { + let mut offset = payload.start; + let end = payload.end; + let flags = read_mat_v5_tag(bytes, offset)?; + offset = offset.checked_add(flags.total_size)?; + let dims_tag = read_mat_v5_tag(bytes, offset)?; + offset = offset.checked_add(dims_tag.total_size)?; + let name_tag = read_mat_v5_tag(bytes, offset)?; + let dims = read_i32_values_le(&bytes[dims_tag.payload], 2)?; + let name = std::str::from_utf8(&bytes[name_tag.payload]).ok()?; + (offset <= end && name == matrix_name).then_some((i64::from(dims[0]), i64::from(dims[1]))) +} + +struct MatV5Tag { + data_type: u32, + payload: std::ops::Range, + total_size: usize, +} + +fn read_mat_v5_tag(bytes: &[u8], offset: usize) -> Option { + let word = read_u32_le(bytes, offset)?; + let small_bytes = word >> 16; + if small_bytes != 0 { + let data_type = word & 0xffff; + let payload_start = offset.checked_add(4)?; + let payload_end = payload_start.checked_add(small_bytes as usize)?; + return (payload_end <= offset.checked_add(8)? && payload_end <= bytes.len()).then_some( + MatV5Tag { + data_type, + payload: payload_start..payload_end, + total_size: 8, + }, + ); + } + let byte_count = read_u32_le(bytes, offset + 4)? as usize; + let payload_start = offset.checked_add(8)?; + let payload_end = payload_start.checked_add(byte_count)?; + let total_size = 8usize.checked_add(pad_to_eight(byte_count)?)?; + (payload_end <= bytes.len()).then_some(MatV5Tag { + data_type: word, + payload: payload_start..payload_end, + total_size, + }) +} + +fn pad_to_eight(value: usize) -> Option { + value.checked_add(7).map(|v| v & !7) +} + +fn read_i32_values_le(bytes: &[u8], count: usize) -> Option> { + (0..count).map(|i| read_i32_le(bytes, i * 4)).collect() +} + +fn read_i32_le(bytes: &[u8], offset: usize) -> Option { + read_u32_le(bytes, offset).map(|value| value as i32) +} + +fn read_u32_le(bytes: &[u8], offset: usize) -> Option { + let slice = bytes.get(offset..offset.checked_add(4)?)?; + Some(u32::from_le_bytes(slice.try_into().ok()?)) +} diff --git a/crates/rumoca-phase-flatten/src/pipeline/mod.rs b/crates/rumoca-phase-flatten/src/pipeline/mod.rs index 49174fb67..96155c644 100644 --- a/crates/rumoca-phase-flatten/src/pipeline/mod.rs +++ b/crates/rumoca-phase-flatten/src/pipeline/mod.rs @@ -13,8 +13,10 @@ mod flatten_pipeline; mod function_overrides_and_dims; mod import_scopes; mod instance_identity; +mod lookup_scopes; pub(crate) use component_alias_injection::*; +pub(crate) use component_member_scope::ComponentMemberScopes; pub(crate) use constant_injection::*; pub(crate) use constant_terminals::*; pub(crate) use context_and_tests::*; @@ -23,3 +25,4 @@ pub(crate) use flatten_pipeline::*; pub(crate) use function_overrides_and_dims::*; pub(crate) use import_scopes::*; pub(crate) use instance_identity::*; +pub(crate) use lookup_scopes::*; diff --git a/crates/rumoca-phase-flatten/src/postprocess.rs b/crates/rumoca-phase-flatten/src/postprocess.rs index 4ddaa6f0c..6211bceaa 100644 --- a/crates/rumoca-phase-flatten/src/postprocess.rs +++ b/crates/rumoca-phase-flatten/src/postprocess.rs @@ -1,17 +1,35 @@ +// SPEC_0021 file-size exception: postprocess still coordinates constant +// substitution, annotations, and scoped parameter preservation. split plan: +// move annotation substitution and scoped parameter rewrites into submodules. use super::*; +use def_id::{aggregate_projection_needs_alias_protection, aggregate_projection_ref}; use rumoca_core::{ ExpressionRewriter, FallibleExpressionRewriter, FallibleStatementRewriter, StatementRewriter, }; - +use std::collections::HashMap; #[path = "postprocess_record_alias.rs"] mod record_alias; use record_alias::*; - pub(super) fn canonicalize_varrefs_via_record_aliases(flat: &mut flat::Model, ctx: &Context) { - if ctx.record_aliases.is_empty() { - return; - } let known_variables: HashSet = flat.variables.keys().map(ToString::to_string).collect(); + for var in flat.variables.values_mut() { + let owner = rumoca_core::ComponentPath::from_flat_path(var.name.as_str()); + canonicalize_record_alias_opt_expr_in_owner( + &mut var.binding, + ctx, + &known_variables, + &owner, + ); + canonicalize_record_alias_opt_expr_in_owner(&mut var.start, ctx, &known_variables, &owner); + canonicalize_record_alias_opt_expr_in_owner(&mut var.min, ctx, &known_variables, &owner); + canonicalize_record_alias_opt_expr_in_owner(&mut var.max, ctx, &known_variables, &owner); + canonicalize_record_alias_opt_expr_in_owner( + &mut var.nominal, + ctx, + &known_variables, + &owner, + ); + } for equation in &mut flat.equations { canonicalize_record_alias_expr(&mut equation.residual, ctx, &known_variables); } @@ -29,7 +47,6 @@ pub(super) fn canonicalize_varrefs_via_record_aliases(flat: &mut flat::Model, ct canonicalize_record_alias_statements(&mut algorithm.statements, ctx, &known_variables); } } - #[path = "postprocess_def_id.rs"] mod def_id; pub(crate) use def_id::canonicalize_varrefs_via_instantiated_def_ids; @@ -38,25 +55,14 @@ mod field_access; pub(super) use field_access::{ drop_invalid_field_access_bindings, resolve_nested_constructor_field_access_bindings, }; - fn record_alias_rewrite_name( name: &str, ctx: &Context, known_variables: &HashSet, + owner: Option<&rumoca_core::ComponentPath>, ) -> Option { - let name_path = rumoca_core::ComponentPath::from_flat_path(name); - ctx.record_aliases.iter().find_map(|(alias, target)| { - if !name_path.starts_with(alias) || name_path.len() == alias.len() { - return None; - } - let suffix = name_path - .suffix_from(alias.len()) - .expect("suffix index is in range"); - let candidate = target.join(&suffix).to_flat_string(); - known_variables.contains(&candidate).then_some(candidate) - }) + record_alias::rewrite_name(name, ctx, known_variables, owner) } - pub(super) fn mark_record_constructor_calls(flat: &mut flat::Model, tree: &ast::ClassTree) { let constructor_def_ids = tree .def_map @@ -82,7 +88,6 @@ pub(super) fn mark_record_constructor_calls(flat: &mut flat::Model, tree: &ast:: if constructor_names.is_empty() && constructor_def_ids.is_empty() { return; } - let marker = ConstructorMarker { constructor_names: &constructor_names, constructor_def_ids: &constructor_def_ids, @@ -133,30 +138,25 @@ pub(super) fn mark_record_constructor_calls(flat: &mut flat::Model, tree: &ast:: marker.mark_statements(&mut function.body); } } - #[derive(Clone, Copy)] struct ConstructorMarker<'a> { constructor_names: &'a HashSet, constructor_def_ids: &'a HashSet, } - impl ConstructorMarker<'_> { fn mark_opt_expr(self, expr: &mut Option) { if let Some(expr) = expr { self.mark_expr(expr); } } - fn mark_expr(mut self, expr: &mut rumoca_core::Expression) { *expr = self.rewrite_expression(expr); } - fn mark_statements(mut self, statements: &mut [rumoca_core::Statement]) { for statement in statements { *statement = self.rewrite_statement(statement); } } - fn mark_when_equations(self, equations: &mut [rumoca_ir_flat::WhenEquation]) { for equation in equations { match equation { @@ -180,7 +180,6 @@ impl ConstructorMarker<'_> { } } } - fn mark_conditional_when_equation( self, branches: &mut [(rumoca_core::Expression, Vec)], @@ -192,14 +191,12 @@ impl ConstructorMarker<'_> { } self.mark_when_equations(else_branch); } - fn is_constructor_call(self, name: &rumoca_core::Reference) -> bool { name.target_def_id() .is_some_and(|def_id| self.constructor_def_ids.contains(&def_id)) || self.constructor_names.contains(name.as_str()) } } - impl ExpressionRewriter for ConstructorMarker<'_> { fn rewrite_expression(&mut self, expr: &rumoca_core::Expression) -> rumoca_core::Expression { if let rumoca_core::Expression::FunctionCall { @@ -219,7 +216,6 @@ impl ExpressionRewriter for ConstructorMarker<'_> { self.walk_expression(expr) } } - impl StatementRewriter for ConstructorMarker<'_> {} pub(super) fn collapse_index_refs_to_known_varrefs(flat: &mut flat::Model) { @@ -480,10 +476,6 @@ fn should_replace_dims(current: &[i64], recovered: &[i64]) -> bool { return false; } current.len() < recovered.len() - || current - .iter() - .zip(recovered.iter()) - .any(|(current, recovered)| *recovered > *current) } fn collapse_index_when_equations( @@ -541,22 +533,70 @@ fn collapse_index_expr(expr: &mut rumoca_core::Expression, known_flat_vars: &Kno /// only flat variables are the `.re`/`.im` leaves). struct KnownFlatVars { names: std::collections::BTreeMap>, + aggregate_projection_refs: HashSet, } impl KnownFlatVars { fn build(flat: &flat::Model) -> Self { + let mut aggregate_projection_counts: HashMap = HashMap::new(); let names = flat .variables .iter() - .map(|(name, var)| (name.as_str().to_string(), var.component_ref.clone())) + .map(|(name, var)| { + if let Some(projection) = aggregate_projection_ref(name.as_str()) { + *aggregate_projection_counts.entry(projection).or_insert(0) += 1; + } + (name.as_str().to_string(), var.component_ref.clone()) + }) + .collect(); + let known_variable_names = flat + .variables + .keys() + .map(|name| name.as_str().to_string()) + .collect::>(); + let aggregate_projection_refs = aggregate_projection_counts + .into_iter() + .filter_map(|(projection, count)| { + (count > 1 + && aggregate_projection_needs_alias_protection( + &projection, + &known_variable_names, + )) + .then_some(projection) + }) .collect(); - Self { names } + Self { + names, + aggregate_projection_refs, + } } fn contains(&self, name: &str) -> bool { self.names.contains_key(name) } + fn exact_reference(&self, name: &str) -> Option { + self.names.get(name).map(|component_ref| { + component_ref.as_ref().map_or_else( + || rumoca_core::Reference::new(name.to_string()), + |component_ref| { + rumoca_core::Reference::with_component_reference(name, component_ref.clone()) + }, + ) + }) + } + + fn has_structural_provenance(&self, path: &str) -> bool { + self.names + .get(path) + .is_some_and(|component_ref| component_ref.is_some()) + || self.record_base_reference(path).is_some() + } + + fn is_aggregate_projection_ref(&self, name: &str) -> bool { + self.aggregate_projection_refs.contains(name) + } + /// Structured reference for a scalarized record base: `path` names no flat /// variable itself, but at least one leaf variable renders as /// `path....`. The base reference is recovered by truncating that @@ -583,6 +623,18 @@ impl KnownFlatVars { } None } + + fn array_base_reference(&self, path: &str) -> Option { + let prefix = format!("{path}["); + let (leaf_name, leaf_ref) = self.names.range(prefix.clone()..).next()?; + if !leaf_name.starts_with(&prefix) { + return None; + } + let span = leaf_ref.as_ref()?.span; + Some(rumoca_core::ComponentReference::from_flat_segments( + path, span, None, + )) + } } struct CollapseIndexRewriter<'a> { @@ -591,6 +643,25 @@ struct CollapseIndexRewriter<'a> { impl ExpressionRewriter for CollapseIndexRewriter<'_> { fn rewrite_expression(&mut self, expr: &rumoca_core::Expression) -> rumoca_core::Expression { + if let rumoca_core::Expression::VarRef { + name, + subscripts, + span, + } = expr + && subscripts.is_empty() + && !name.has_structure() + && !self.known_flat_vars.contains(name.as_str()) + && !self + .known_flat_vars + .is_aggregate_projection_ref(name.as_str()) + && let Some(collapsed) = collapse_repeated_field_tail_to_known_var( + name.as_str(), + *span, + self.known_flat_vars, + ) + { + return collapsed; + } if let rumoca_core::Expression::FieldAccess { base, field, span } = expr { let base = self.rewrite_expression(base); if let Some(collapsed) = @@ -692,6 +763,11 @@ fn collapse_field_access_to_known_var( span, }); } + if let Some(collapsed) = + collapse_repeated_field_tail_to_known_var(&candidate, span, known_flat_vars) + { + return Some(collapsed); + } // Scalarized record base (`comp[1].port_p.Phi` where only the // `.re`/`.im` leaves exist as flat variables): collapse to a single // structured VarRef so downstream record-equation expansion sees the @@ -730,6 +806,179 @@ fn collapse_field_access_to_known_var( } } +fn collapse_repeated_field_tail_to_known_var( + candidate: &str, + span: rumoca_core::Span, + known_flat_vars: &KnownFlatVars, +) -> Option { + let mut path = candidate; + while let Some((prefix, field)) = rendered_path_last_segment(path) { + if !prefix.ends_with(field) { + // An indexed component element is an instance boundary, not a + // redundant record-field segment. For example, + // `stack.cell[1,1].cell` must not collapse to the aggregate + // projection `stack.cell` merely because that projection exists. + if rendered_path_last_segment(prefix).is_some_and(|(_, penultimate)| { + rumoca_core::split_trailing_subscript_suffix(penultimate).is_some() + }) { + return None; + } + return collapse_penultimate_field_to_known_var(path, span, known_flat_vars); + } + path = prefix; + if known_flat_vars.contains(path) { + return Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new(path.to_string()), + subscripts: vec![], + span, + }); + } + if let Some(alternate) = alternate_array_field_path(path) + && known_flat_vars.contains(&alternate) + { + return Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new(alternate), + subscripts: vec![], + span, + }); + } + if let Some(reference) = known_flat_vars.record_base_reference(path) { + return Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference(path, reference), + subscripts: vec![], + span, + }); + } + if let Some(reference) = known_flat_vars.array_base_reference(path) { + return Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference(path, reference), + subscripts: vec![], + span, + }); + } + if let Some(alternate) = alternate_array_field_path(path) + && let Some(reference) = known_flat_vars.record_base_reference(&alternate) + { + return Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference(&alternate, reference), + subscripts: vec![], + span, + }); + } + if let Some(alternate) = alternate_array_field_path(path) + && let Some(reference) = known_flat_vars.array_base_reference(&alternate) + { + return Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference(&alternate, reference), + subscripts: vec![], + span, + }); + } + if let Some(collapsed) = + collapse_penultimate_field_to_known_var(path, span, known_flat_vars) + { + return Some(collapsed); + } + } + None +} + +fn collapse_penultimate_field_to_known_var( + path: &str, + span: rumoca_core::Span, + known_flat_vars: &KnownFlatVars, +) -> Option { + let (prefix, leaf) = rendered_path_last_segment(path)?; + let (base, _) = rendered_path_last_segment(prefix)?; + let candidate = format!("{base}.{leaf}"); + if known_flat_vars.has_structural_provenance(prefix) { + // The removed segment owns known descendants with structured variable + // metadata, so it is a real component/record boundary rather than an + // over-expanded field alias. Do not redirect its leaf to a sibling. + return None; + } + known_path_expression(&candidate, span, known_flat_vars) +} + +fn known_path_expression( + path: &str, + span: rumoca_core::Span, + known_flat_vars: &KnownFlatVars, +) -> Option { + if let Some(reference) = known_flat_vars.exact_reference(path) { + return Some(rumoca_core::Expression::VarRef { + name: reference, + subscripts: vec![], + span, + }); + } + if let Some(alternate) = alternate_array_field_path(path) + && let Some(reference) = known_flat_vars.exact_reference(&alternate) + { + return Some(rumoca_core::Expression::VarRef { + name: reference, + subscripts: vec![], + span, + }); + } + if let Some(reference) = known_flat_vars.record_base_reference(path) { + return Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference(path, reference), + subscripts: vec![], + span, + }); + } + if let Some(reference) = known_flat_vars.array_base_reference(path) { + return Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference(path, reference), + subscripts: vec![], + span, + }); + } + if let Some(alternate) = alternate_array_field_path(path) + && let Some(reference) = known_flat_vars.record_base_reference(&alternate) + { + return Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference(&alternate, reference), + subscripts: vec![], + span, + }); + } + if let Some(alternate) = alternate_array_field_path(path) + && let Some(reference) = known_flat_vars.array_base_reference(&alternate) + { + return Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference(&alternate, reference), + subscripts: vec![], + span, + }); + } + None +} + +fn alternate_array_field_path(path: &str) -> Option { + let (base, field) = rendered_path_last_segment(path)?; + if !base.ends_with(']') { + return None; + } + let bracket_start = base.rfind('[')?; + let (array_base, suffix) = base.split_at(bracket_start); + Some(format!("{array_base}.{field}{suffix}")) +} + +fn rendered_path_last_segment(path: &str) -> Option<(&str, &str)> { + let mut bracket_depth = 0usize; + for (idx, byte) in path.bytes().enumerate().rev() { + match byte { + b']' => bracket_depth += 1, + b'[' => bracket_depth = bracket_depth.saturating_sub(1), + b'.' if bracket_depth == 0 => return Some((&path[..idx], &path[idx + 1..])), + _ => {} + } + } + None +} + fn field_access_flat_path(base: &rumoca_core::Expression, field: &str) -> Option { Some(format!("{}.{}", expr_flat_path(base)?, field)) } @@ -815,34 +1064,77 @@ pub(super) fn substitute_known_constants_in_flat( .keys() .map(|name| name.as_str().to_string()) .collect(); + let var_dims: rustc_hash::FxHashMap> = flat + .variables + .iter() + .filter(|(_, var)| !var.dims.is_empty()) + .map(|(name, var)| (name.as_str().to_string(), var.dims.clone())) + .collect(); + evaluate_static_initial_parameter_algorithms(flat, ctx)?; + let var_values = parameter_constant_var_values(flat); + let binding_var_values = flat_binding_var_values(flat); let no_locals: HashSet = HashSet::new(); for eq in &mut flat.equations { let scope = equation_origin_scope(&eq.origin); - eq.residual = substitute_known_constants_expr( + eq.residual = substitute_known_constants_expr_with_options_dims_and_values( eq.residual.clone(), ctx, &live_vars, &no_locals, &scope, + true, + &var_dims, + &var_values, )?; + reconcile_residual_constructor_extents_with_lhs_dims(&mut eq.residual, &var_dims); } for eq in &mut flat.initial_equations { let scope = equation_origin_scope(&eq.origin); - eq.residual = substitute_known_constants_expr( + eq.residual = substitute_known_constants_expr_with_options_dims_and_values( eq.residual.clone(), ctx, &live_vars, &no_locals, &scope, + true, + &var_dims, + &var_values, )?; + reconcile_residual_constructor_extents_with_lhs_dims(&mut eq.residual, &var_dims); } - substitute_assert_equations(&mut flat.assert_equations, ctx, &live_vars, &no_locals)?; + substitute_structured_equation_templates( + &mut flat.structured_equations, + ctx, + &live_vars, + &no_locals, + &var_dims, + &var_values, + )?; + substitute_structured_equation_templates( + &mut flat.initial_structured_equations, + ctx, + &live_vars, + &no_locals, + &var_dims, + &var_values, + )?; + recover_primitive_constructor_parameter_bindings(flat); + substitute_assert_equations( + &mut flat.assert_equations, + ctx, + &live_vars, + &no_locals, + &var_dims, + &var_values, + )?; substitute_assert_equations( &mut flat.initial_assert_equations, ctx, &live_vars, &no_locals, + &var_dims, + &var_values, )?; for when_clause in &mut flat.when_clauses { when_clause.condition = substitute_known_constants_expr( @@ -856,102 +1148,1076 @@ pub(super) fn substitute_known_constants_in_flat( substitute_known_constants_when_equation(equation, ctx, &live_vars, &no_locals)?; } } - substitute_algorithms(&mut flat.algorithms, ctx, &live_vars, &no_locals)?; - substitute_algorithms(&mut flat.initial_algorithms, ctx, &live_vars, &no_locals)?; - substitute_variable_annotations(&mut flat.variables, ctx, &live_vars, &no_locals)?; + substitute_algorithms( + &mut flat.algorithms, + ctx, + &live_vars, + &no_locals, + &var_dims, + &var_values, + )?; + substitute_algorithms( + &mut flat.initial_algorithms, + ctx, + &live_vars, + &no_locals, + &var_dims, + &var_values, + )?; + substitute_variable_annotations( + flat, + ctx, + &live_vars, + &no_locals, + &var_dims, + &binding_var_values, + )?; substitute_function_bodies(&mut flat.functions, ctx, &live_vars)?; crate::zero_sized_arrays::materialize_referenced_zero_sized_array_variables(flat, ctx)?; Ok(()) } -fn equation_origin_scope(origin: &flat::EquationOrigin) -> String { - match origin { - flat::EquationOrigin::ComponentEquation { component } - | flat::EquationOrigin::Algorithm { component } => component.clone(), - flat::EquationOrigin::Binding { variable } - | flat::EquationOrigin::Reinit { state: variable } - | flat::EquationOrigin::WhenAssignment { target: variable } - | flat::EquationOrigin::UnconnectedFlow { variable } => parent_component_scope(variable), - flat::EquationOrigin::Connection { .. } | flat::EquationOrigin::FlowSum { .. } => { - String::new() - } - } +fn parameter_constant_var_values( + flat: &flat::Model, +) -> rustc_hash::FxHashMap { + flat.variables + .iter() + .filter_map(|(name, var)| { + let expr = var.binding.as_ref().or_else(|| { + (var.fixed != Some(false)) + .then_some(()) + .and(var.start.as_ref()) + })?; + if !parameter_value_is_structural(flat, name, var, expr) { + return None; + } + Some((name.as_str().to_string(), expr.clone())) + }) + .collect() } -fn parent_component_scope(name: &str) -> String { - rumoca_core::ComponentPath::from_flat_path(name) - .parent() - .unwrap_or_else(rumoca_core::ComponentPath::root) - .to_flat_string() +fn flat_binding_var_values( + flat: &flat::Model, +) -> rustc_hash::FxHashMap { + flat.variables + .iter() + .filter_map(|(name, var)| { + let has_binding = var.binding.is_some(); + let expr = var.binding.as_ref().or_else(|| { + (var.fixed != Some(false)) + .then_some(()) + .and(var.start.as_ref()) + })?; + if parameter_value_is_structural(flat, name, var, expr) + || has_binding + && structural_non_real_expr(expr) + && !variable_type_is_real(flat.variable_type_names.get(name)) + { + return Some((name.as_str().to_string(), expr.clone())); + } + None + }) + .collect() } -fn substitute_assert_equations( - equations: &mut [flat::AssertEquation], - ctx: &Context, - live_vars: &rustc_hash::FxHashSet, - locals: &HashSet, -) -> Result<(), FlattenError> { - for assert_eq in equations { - let scope = equation_origin_scope(&assert_eq.origin); - assert_eq.condition = substitute_known_constants_expr_with_options( - assert_eq.condition.clone(), - ctx, - live_vars, - locals, - &scope, - true, - )?; - assert_eq.message = substitute_known_constants_expr_with_options( - assert_eq.message.clone(), - ctx, - live_vars, - locals, - &scope, - true, - )?; - substitute_opt_expr_with_options( - &mut assert_eq.level, - ctx, - live_vars, - locals, - &scope, - true, - )?; +fn parameter_value_is_structural( + flat: &flat::Model, + name: &rumoca_core::VarName, + var: &flat::Variable, + expr: &rumoca_core::Expression, +) -> bool { + if matches!(var.variability, rumoca_core::Variability::Constant(_)) { + return true; } - Ok(()) + if !matches!(var.variability, rumoca_core::Variability::Parameter(_)) { + return false; + } + var.evaluate + || var.is_discrete_type + || structural_non_real_expr(expr) + && !variable_type_is_real(flat.variable_type_names.get(name)) } -fn substitute_algorithms( - algorithms: &mut [flat::Algorithm], - ctx: &Context, - live_vars: &rustc_hash::FxHashSet, - locals: &HashSet, -) -> Result<(), FlattenError> { - for algorithm in algorithms { +fn variable_type_is_real(type_name: Option<&String>) -> bool { + type_name.is_some_and(|name| rumoca_core::qualified_type_name_matches(name, "Real")) +} + +fn structural_non_real_expr(expr: &rumoca_core::Expression) -> bool { + match expr { + rumoca_core::Expression::Literal { value, .. } => { + !matches!(value, rumoca_core::Literal::Real(_)) + } + rumoca_core::Expression::VarRef { .. } => true, + rumoca_core::Expression::Unary { rhs, .. } => structural_non_real_expr(rhs), + rumoca_core::Expression::Binary { lhs, rhs, .. } => { + structural_non_real_expr(lhs) && structural_non_real_expr(rhs) + } + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + branches.iter().all(|(condition, value)| { + structural_non_real_expr(condition) && structural_non_real_expr(value) + }) && structural_non_real_expr(else_branch) + } + rumoca_core::Expression::BuiltinCall { args, .. } + | rumoca_core::Expression::FunctionCall { args, .. } + | rumoca_core::Expression::Tuple { elements: args, .. } + | rumoca_core::Expression::Array { elements: args, .. } => { + args.iter().all(structural_non_real_expr) + } + rumoca_core::Expression::FieldAccess { base, .. } => structural_non_real_expr(base), + rumoca_core::Expression::Index { + base, subscripts, .. + } => { + structural_non_real_expr(base) + && subscripts.iter().all(|subscript| match subscript { + rumoca_core::Subscript::Index { .. } | rumoca_core::Subscript::Colon { .. } => { + true + } + rumoca_core::Subscript::Expr { expr, .. } => structural_non_real_expr(expr), + }) + } + rumoca_core::Expression::Range { + start, step, end, .. + } => { + structural_non_real_expr(start) + && step + .as_ref() + .is_none_or(|step| structural_non_real_expr(step)) + && structural_non_real_expr(end) + } + rumoca_core::Expression::ArrayComprehension { + expr, + indices, + filter, + .. + } => { + structural_non_real_expr(expr) + && indices + .iter() + .all(|index| structural_non_real_expr(&index.range)) + && filter + .as_ref() + .is_none_or(|filter| structural_non_real_expr(filter)) + } + rumoca_core::Expression::Empty { .. } => false, + } +} + +fn evaluate_static_initial_parameter_algorithms( + flat: &mut flat::Model, + ctx: &Context, +) -> Result<(), FlattenError> { + for algorithm in flat.initial_algorithms.clone() { + let mut eval_ctx = constant_eval_context_for_flat(flat, ctx); + let mut assignments = rustc_hash::FxHashMap::default(); + if eval_static_statement_block(&algorithm.statements, flat, &mut eval_ctx, &mut assignments) + .is_err() + { + continue; + } + for (name, value) in assignments { + let Some(var) = flat.variables.get_mut(&rumoca_core::VarName::new(&name)) else { + continue; + }; + if let Some(expr) = constant_value_to_expression(&value, var.source_span) { + var.binding = Some(expr); + } + } + } + Ok(()) +} + +fn constant_eval_context_for_flat( + flat: &flat::Model, + ctx: &Context, +) -> rumoca_eval_flat::constant::EvalContext { + let mut eval_ctx = rumoca_eval_flat::constant::EvalContext::with_capacity( + ctx.parameter_values.len() + + ctx.real_parameter_values.len() + + ctx.boolean_parameter_values.len() + + ctx.string_parameter_values.len() + + ctx.constant_values.len() + + flat.variables.len(), + ctx.enum_parameter_values.len(), + flat.functions.len(), + ); + for function in flat.functions.values() { + eval_ctx.add_function(function.clone()); + } + for (name, value) in &ctx.parameter_values { + eval_ctx.add_parameter( + name.clone(), + rumoca_eval_flat::constant::Value::Integer(*value), + ); + } + for (name, value) in &ctx.real_parameter_values { + eval_ctx.add_parameter( + name.clone(), + rumoca_eval_flat::constant::Value::Real(*value), + ); + } + for (name, value) in &ctx.boolean_parameter_values { + eval_ctx.add_parameter( + name.clone(), + rumoca_eval_flat::constant::Value::Bool(*value), + ); + } + for (name, value) in &ctx.string_parameter_values { + eval_ctx.add_parameter( + name.clone(), + rumoca_eval_flat::constant::Value::String(value.clone()), + ); + } + for (name, value) in &ctx.enum_parameter_values { + eval_ctx.enum_literals.insert( + name.clone(), + (parent_component_scope(value), value.to_string()), + ); + } + for (name, expr) in &ctx.constant_values { + let Some(span) = expr.span() else { + continue; + }; + if let Ok(value) = rumoca_eval_flat::constant::eval_expr_with_span(expr, &eval_ctx, span) { + eval_ctx.add_parameter(name.clone(), value); + } + } + for (name, var) in &flat.variables { + if !matches!( + var.variability, + rumoca_core::Variability::Parameter(_) | rumoca_core::Variability::Constant(_) + ) { + continue; + } + let Some(expr) = var.binding.as_ref().or_else(|| { + (var.fixed != Some(false)) + .then_some(()) + .and(var.start.as_ref()) + }) else { + continue; + }; + if let Ok(value) = rumoca_eval_flat::constant::eval_expr_with_span( + expr, + &eval_ctx, + expr.span().unwrap_or(var.source_span), + ) { + eval_ctx.add_parameter(name.as_str().to_string(), value); + } + } + eval_ctx +} + +fn eval_static_statement_block( + statements: &[rumoca_core::Statement], + flat: &flat::Model, + eval_ctx: &mut rumoca_eval_flat::constant::EvalContext, + assignments: &mut rustc_hash::FxHashMap, +) -> Result<(), ()> { + for statement in statements { + eval_static_statement(statement, flat, eval_ctx, assignments)?; + } + Ok(()) +} + +fn eval_static_statement( + statement: &rumoca_core::Statement, + flat: &flat::Model, + eval_ctx: &mut rumoca_eval_flat::constant::EvalContext, + assignments: &mut rustc_hash::FxHashMap, +) -> Result<(), ()> { + match statement { + rumoca_core::Statement::Empty { .. } => Ok(()), + rumoca_core::Statement::Assignment { comp, value, span } => { + let target = static_initial_parameter_target(flat, comp)?; + let value = rumoca_eval_flat::constant::eval_expr_with_span(value, eval_ctx, *span) + .map_err(|_| ())?; + eval_ctx.add_parameter(target.clone(), value.clone()); + assignments.insert(target, value); + Ok(()) + } + rumoca_core::Statement::For { + indices, equations, .. + } => eval_static_for(indices, equations, flat, eval_ctx, assignments), + rumoca_core::Statement::If { + cond_blocks, + else_block, + span, + } => { + for block in cond_blocks { + let cond = + rumoca_eval_flat::constant::eval_expr_with_span(&block.cond, eval_ctx, *span) + .map_err(|_| ())?; + if cond.as_bool().ok_or(())? { + return eval_static_statement_block(&block.stmts, flat, eval_ctx, assignments); + } + } + if let Some(else_block) = else_block { + return eval_static_statement_block(else_block, flat, eval_ctx, assignments); + } + Ok(()) + } + rumoca_core::Statement::Assert { + condition, span, .. + } => { + let condition = + rumoca_eval_flat::constant::eval_expr_with_span(condition, eval_ctx, *span) + .map_err(|_| ())?; + condition.as_bool().filter(|value| *value).ok_or(())?; + Ok(()) + } + _ => Err(()), + } +} + +fn eval_static_for( + indices: &[rumoca_core::ForIndex], + body: &[rumoca_core::Statement], + flat: &flat::Model, + eval_ctx: &mut rumoca_eval_flat::constant::EvalContext, + assignments: &mut rustc_hash::FxHashMap, +) -> Result<(), ()> { + let Some((index, rest)) = indices.split_first() else { + return eval_static_statement_block(body, flat, eval_ctx, assignments); + }; + let span = index.range.span().ok_or(())?; + let values = rumoca_eval_flat::constant::eval_expr_with_span(&index.range, eval_ctx, span) + .map_err(|_| ())?; + let values = values.as_array().ok_or(())?.clone(); + let previous = eval_ctx.parameters.get(&index.ident).cloned(); + for value in values { + eval_ctx.add_parameter(index.ident.clone(), value); + eval_static_for(rest, body, flat, eval_ctx, assignments)?; + } + match previous { + Some(value) => eval_ctx.add_parameter(index.ident.clone(), value), + None => { + eval_ctx.parameters.swap_remove(&index.ident); + } + } + Ok(()) +} + +fn static_initial_parameter_target( + flat: &flat::Model, + comp: &rumoca_core::ComponentReference, +) -> Result { + if comp.parts.iter().any(|part| !part.subs.is_empty()) { + return Err(()); + } + let name = comp.to_var_name(); + let Some(var) = flat.variables.get(&name) else { + return Err(()); + }; + if !matches!( + var.variability, + rumoca_core::Variability::Parameter(_) | rumoca_core::Variability::Constant(_) + ) || var.fixed != Some(false) + { + return Err(()); + } + Ok(name.as_str().to_string()) +} + +fn constant_value_to_expression( + value: &rumoca_eval_flat::constant::Value, + span: rumoca_core::Span, +) -> Option { + match value { + rumoca_eval_flat::constant::Value::Integer(value) => { + Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(*value), + span, + }) + } + rumoca_eval_flat::constant::Value::Real(value) => Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(*value), + span, + }), + rumoca_eval_flat::constant::Value::Bool(value) => Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Boolean(*value), + span, + }), + rumoca_eval_flat::constant::Value::String(value) => { + Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String(value.clone()), + span, + }) + } + rumoca_eval_flat::constant::Value::Array(values) => { + let elements = values + .iter() + .map(|value| constant_value_to_expression(value, span)) + .collect::>>()?; + Some(rumoca_core::Expression::Array { + elements, + is_matrix: false, + span, + }) + } + rumoca_eval_flat::constant::Value::Enum(type_name, literal) => { + Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new(format!("{type_name}.{literal}")), + subscripts: vec![], + span, + }) + } + rumoca_eval_flat::constant::Value::Record(_) => None, + } +} + +fn substitute_structured_equation_templates( + families: &mut [flat::StructuredEquationFamily], + ctx: &Context, + live_vars: &rustc_hash::FxHashSet, + locals: &HashSet, + var_dims: &rustc_hash::FxHashMap>, + var_values: &rustc_hash::FxHashMap, +) -> Result<(), FlattenError> { + for family in families { + let Some(template) = &mut family.template else { + continue; + }; + let scope = equation_origin_scope(&family.origin); + for residual in &mut template.body { + *residual = substitute_known_constants_expr_with_options_dims_and_values( + residual.clone(), + ctx, + live_vars, + locals, + &scope, + true, + var_dims, + var_values, + )?; + } + } + Ok(()) +} + +#[derive(Clone)] +struct PrimitiveParameterBindingCandidate { + scope: String, + min: Option, + max: Option, + binding: rumoca_core::Expression, +} + +fn recover_primitive_constructor_parameter_bindings(flat: &mut flat::Model) { + let candidates = primitive_parameter_binding_candidates(flat); + if candidates.is_empty() { + return; + } + + for eq in &mut flat.equations { + let scope = assignment_scope_from_residual(&eq.residual) + .unwrap_or_else(|| equation_origin_scope(&eq.origin)); + let mut rewriter = PrimitiveConstructorBindingRecoverer { + scope: &scope, + candidates: &candidates, + }; + eq.residual = rewriter.rewrite_expression(&eq.residual); + } + for eq in &mut flat.initial_equations { + let scope = assignment_scope_from_residual(&eq.residual) + .unwrap_or_else(|| equation_origin_scope(&eq.origin)); + let mut rewriter = PrimitiveConstructorBindingRecoverer { + scope: &scope, + candidates: &candidates, + }; + eq.residual = rewriter.rewrite_expression(&eq.residual); + } +} + +fn primitive_parameter_binding_candidates( + flat: &flat::Model, +) -> Vec { + flat.variables + .iter() + .filter_map(|(name, var)| { + let type_name = flat.variable_type_names.get(name)?; + if !rumoca_core::qualified_type_name_matches(type_name, "Integer") + || !matches!( + var.variability, + rumoca_core::Variability::Parameter(_) | rumoca_core::Variability::Constant(_) + ) + || !var.binding_from_modification + { + return None; + } + let binding = var.binding.clone()?; + Some(PrimitiveParameterBindingCandidate { + scope: parent_component_scope(var.name.as_str()), + min: var.min.clone(), + max: var.max.clone(), + binding, + }) + }) + .collect() +} + +struct PrimitiveConstructorBindingRecoverer<'a> { + scope: &'a str, + candidates: &'a [PrimitiveParameterBindingCandidate], +} + +impl ExpressionRewriter for PrimitiveConstructorBindingRecoverer<'_> { + fn walk_function_call_expression( + &mut self, + name: &rumoca_core::Reference, + args: &[rumoca_core::Expression], + is_constructor: bool, + span: rumoca_core::Span, + ) -> rumoca_core::Expression { + let mut rewritten_args = self.rewrite_expressions(args); + for index in primitive_constructor_binding_arg_indices(name.as_str(), rewritten_args.len()) + { + if let Some(arg) = rewritten_args.get_mut(*index) + && let Some(recovered) = self.recover_integer_constructor_argument(arg, span) + { + *arg = recovered; + } + } + rumoca_core::Expression::FunctionCall { + name: name.clone(), + args: rewritten_args, + is_constructor, + span, + } + } +} + +fn primitive_constructor_binding_arg_indices(name: &str, arg_len: usize) -> &'static [usize] { + match name { + "Clock" if arg_len >= 1 => &[0], + "subSample" | "superSample" if arg_len >= 2 => &[1], + "shiftSample" | "backSample" if arg_len >= 3 => &[1, 2], + _ => &[], + } +} + +impl PrimitiveConstructorBindingRecoverer<'_> { + fn recover_integer_constructor_argument( + &self, + expr: &rumoca_core::Expression, + span: rumoca_core::Span, + ) -> Option { + let bounds = primitive_integer_constructor_bounds(expr)?; + let mut matches = self.candidates.iter().filter(|candidate| { + candidate.scope == self.scope + && bounds + .min + .as_ref() + .is_none_or(|min| candidate.min.as_ref() == Some(min)) + && bounds + .max + .as_ref() + .is_none_or(|max| candidate.max.as_ref() == Some(max)) + }); + let candidate = matches.next()?; + matches + .next() + .is_none() + .then(|| candidate.binding.clone().with_span(span)) + } +} + +struct PrimitiveConstructorBounds { + min: Option, + max: Option, +} + +fn primitive_integer_constructor_bounds( + expr: &rumoca_core::Expression, +) -> Option { + let rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor, + .. + } = expr + else { + return None; + }; + if name.as_str() != "Integer" || !*is_constructor || args.is_empty() { + return None; + } + + let mut min = None; + let mut max = None; + for arg in args { + let rumoca_core::Expression::FunctionCall { + name, + args: named_args, + .. + } = arg + else { + return None; + }; + let attribute = name + .as_str() + .strip_prefix(rumoca_core::NAMED_FUNCTION_ARG_PREFIX)?; + let value = named_args.first()?.clone(); + if named_args.len() != 1 { + return None; + } + match attribute { + "min" => min = Some(value), + "max" => max = Some(value), + _ => return None, + } + } + + Some(PrimitiveConstructorBounds { min, max }) +} + +fn assignment_scope_from_residual(expr: &rumoca_core::Expression) -> Option { + let rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs, + rhs, + .. + } = expr + else { + return None; + }; + assignment_scope_var_ref(lhs) + .or_else(|| assignment_scope_var_ref(rhs)) + .map(|name| parent_component_scope(name.as_str())) + .filter(|scope| !scope.is_empty()) +} + +fn assignment_scope_var_ref(expr: &rumoca_core::Expression) -> Option<&rumoca_core::Reference> { + match expr { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } if subscripts.is_empty() => Some(name), + _ => None, + } +} + +fn equation_origin_scope(origin: &flat::EquationOrigin) -> String { + match origin { + flat::EquationOrigin::ComponentEquation { component } + | flat::EquationOrigin::Algorithm { component } => component.clone(), + flat::EquationOrigin::Binding { variable } + | flat::EquationOrigin::Reinit { state: variable } + | flat::EquationOrigin::WhenAssignment { target: variable } + | flat::EquationOrigin::UnconnectedFlow { variable } => parent_component_scope(variable), + flat::EquationOrigin::Connection { .. } | flat::EquationOrigin::FlowSum { .. } => { + String::new() + } + } +} + +fn parent_component_scope(name: &str) -> String { + rumoca_core::ComponentPath::from_flat_path(name) + .parent() + .unwrap_or_else(rumoca_core::ComponentPath::root) + .to_flat_string() +} + +fn substitute_assert_equations( + equations: &mut [flat::AssertEquation], + ctx: &Context, + live_vars: &rustc_hash::FxHashSet, + locals: &HashSet, + var_dims: &rustc_hash::FxHashMap>, + var_values: &rustc_hash::FxHashMap, +) -> Result<(), FlattenError> { + for assert_eq in equations { + let scope = equation_origin_scope(&assert_eq.origin); + assert_eq.condition = substitute_known_constants_expr_with_options_dims_and_values( + assert_eq.condition.clone(), + ctx, + live_vars, + locals, + &scope, + true, + var_dims, + var_values, + )?; + assert_eq.message = substitute_known_constants_expr_with_options_dims_and_values( + assert_eq.message.clone(), + ctx, + live_vars, + locals, + &scope, + true, + var_dims, + var_values, + )?; + substitute_opt_expr_with_options_dims_and_values( + &mut assert_eq.level, + ctx, + live_vars, + locals, + &scope, + true, + var_dims, + var_values, + )?; + } + Ok(()) +} + +fn substitute_algorithms( + algorithms: &mut [flat::Algorithm], + ctx: &Context, + live_vars: &rustc_hash::FxHashSet, + locals: &HashSet, + var_dims: &rustc_hash::FxHashMap>, + var_values: &rustc_hash::FxHashMap, +) -> Result<(), FlattenError> { + for algorithm in algorithms { for statement in &mut algorithm.statements { - substitute_known_constants_statement(statement, ctx, live_vars, locals, "")?; + substitute_known_constants_statement_with_dims_and_values( + statement, ctx, live_vars, locals, "", var_dims, var_values, + )?; } } Ok(()) } fn substitute_variable_annotations( - variables: &mut flat::VarNameIndexMap, + flat: &mut flat::Model, ctx: &Context, live_vars: &rustc_hash::FxHashSet, locals: &HashSet, + var_dims: &rustc_hash::FxHashMap>, + var_values: &rustc_hash::FxHashMap, ) -> Result<(), FlattenError> { - for var in variables.values_mut() { + for (name, var) in &mut flat.variables { let scope = parent_component_scope(var.name.as_str()); - substitute_opt_expr(&mut var.binding, ctx, live_vars, locals, &scope)?; - substitute_opt_expr(&mut var.start, ctx, live_vars, locals, &scope)?; - substitute_opt_expr(&mut var.min, ctx, live_vars, locals, &scope)?; - substitute_opt_expr(&mut var.max, ctx, live_vars, locals, &scope)?; - substitute_opt_expr(&mut var.nominal, ctx, live_vars, locals, &scope)?; + if !is_runtime_parameter_modifier_binding(var) + || binding_references_class_constant(var.binding.as_ref(), ctx) + { + substitute_opt_expr_with_options_dims_and_values( + &mut var.binding, + ctx, + live_vars, + locals, + &scope, + false, + var_dims, + var_values, + )?; + } + substitute_opt_expr_with_options_dims_and_values( + &mut var.start, + ctx, + live_vars, + locals, + &scope, + true, + var_dims, + var_values, + )?; + substitute_opt_expr_with_options_and_dims( + &mut var.min, + ctx, + live_vars, + locals, + &scope, + true, + var_dims, + )?; + substitute_opt_expr_with_options_and_dims( + &mut var.max, + ctx, + live_vars, + locals, + &scope, + true, + var_dims, + )?; + substitute_opt_expr_with_options_and_dims( + &mut var.nominal, + ctx, + live_vars, + locals, + &scope, + true, + var_dims, + )?; + reconcile_constructor_extents_with_declared_dims(var); + if variable_is_string_type(var, flat.variable_type_names.get(name)) { + recover_string_literal_opt_expr(&mut var.binding, &scope); + recover_string_literal_opt_expr(&mut var.start, &scope); + } } Ok(()) } +fn reconcile_constructor_extents_with_declared_dims(var: &mut flat::Variable) { + if var.dims.is_empty() + || !matches!( + var.variability, + rumoca_core::Variability::Parameter(_) | rumoca_core::Variability::Constant(_) + ) + { + return; + } + reconcile_constructor_extent_expr(&mut var.binding, &var.dims); + reconcile_constructor_extent_expr(&mut var.start, &var.dims); +} + +fn reconcile_residual_constructor_extents_with_lhs_dims( + residual: &mut rumoca_core::Expression, + var_dims: &rustc_hash::FxHashMap>, +) -> Option<()> { + let rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs, + rhs, + .. + } = residual + else { + return None; + }; + let lhs_name = residual_lhs_var_ref_name(lhs)?; + let dims = var_dims.get(lhs_name.as_str())?; + reconcile_constructor_extent_in_value_expr(rhs, dims) +} + +fn residual_lhs_var_ref_name(expr: &rumoca_core::Expression) -> Option<&rumoca_core::Reference> { + match expr { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } if subscripts.is_empty() => Some(name), + rumoca_core::Expression::Index { base, .. } => { + let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = base.as_ref() + else { + return None; + }; + subscripts.is_empty().then_some(name) + } + _ => None, + } +} + +fn reconcile_constructor_extent_expr( + expr: &mut Option, + dims: &[i64], +) -> Option<()> { + reconcile_constructor_extent_in_value_expr(expr.as_mut()?, dims) +} + +fn reconcile_constructor_extent_in_value_expr( + expr: &mut rumoca_core::Expression, + dims: &[i64], +) -> Option<()> { + match expr { + rumoca_core::Expression::BuiltinCall { .. } => { + reconcile_constructor_extent_call(expr, dims) + } + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + let mut changed = false; + for (_condition, value) in branches { + changed |= reconcile_constructor_extent_in_value_expr(value, dims).is_some(); + } + changed |= reconcile_constructor_extent_in_value_expr(else_branch, dims).is_some(); + changed.then_some(()) + } + rumoca_core::Expression::Binary { lhs, rhs, .. } => { + let lhs_changed = reconcile_constructor_extent_in_value_expr(lhs, dims).is_some(); + let rhs_changed = reconcile_constructor_extent_in_value_expr(rhs, dims).is_some(); + (lhs_changed || rhs_changed).then_some(()) + } + rumoca_core::Expression::Unary { rhs, .. } => { + reconcile_constructor_extent_in_value_expr(rhs, dims) + } + _ => None, + } +} + +fn reconcile_constructor_extent_call( + expr: &mut rumoca_core::Expression, + dims: &[i64], +) -> Option<()> { + let target_dims = dims + .iter() + .copied() + .map(|dim| usize::try_from(dim).ok()) + .collect::>>()?; + if target_dims.is_empty() { + return None; + } + let target_len = target_dims.iter().product::(); + if target_len == 0 { + return None; + } + let rumoca_core::Expression::BuiltinCall { + function, + args, + span, + } = expr + else { + unreachable!("constructor extent reconciliation is only called for builtin calls"); + }; + let offset = match function { + rumoca_core::BuiltinFunction::Fill => 1, + rumoca_core::BuiltinFunction::Zeros | rumoca_core::BuiltinFunction::Ones => 0, + _ => return None, + }; + if args.len() < offset + 1 { + return None; + } + if let Some(current_dims) = args[offset..] + .iter() + .map(literal_usize) + .collect::>>() + { + let current_len = current_dims.iter().product::(); + if current_len == target_len || current_len != 0 { + return None; + } + } + let mut rewritten = args[..offset].to_vec(); + rewritten.extend( + target_dims + .into_iter() + .map(|dim| rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(dim as i64), + span: *span, + }), + ); + *args = rewritten; + Some(()) +} + +fn literal_usize(expr: &rumoca_core::Expression) -> Option { + let rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + .. + } = expr + else { + return None; + }; + usize::try_from(*value).ok() +} + +fn is_runtime_parameter_modifier_binding(var: &flat::Variable) -> bool { + matches!( + var.variability, + rumoca_core::Variability::Parameter(_) | rumoca_core::Variability::Constant(_) + ) && var.binding_from_modification + && !var.evaluate + && !var.is_discrete_type + && var.binding.is_some() +} + +fn binding_references_class_constant( + binding: Option<&rumoca_core::Expression>, + ctx: &Context, +) -> bool { + binding.is_some_and(|expr| expr_references_class_constant(expr, ctx)) +} + +fn expr_references_class_constant(expr: &rumoca_core::Expression, ctx: &Context) -> bool { + match expr { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } => { + let self_is_class_constant = subscripts.is_empty() + && (ctx.class_constant_keys.contains(name.as_str()) + || name + .target_def_id() + .and_then(|def_id| ctx.target_def_names.get(&def_id)) + .is_some_and(|target| ctx.class_constant_keys.contains(target))); + self_is_class_constant || subscripts_reference_class_constant(subscripts, ctx) + } + rumoca_core::Expression::Binary { lhs, rhs, .. } => { + expr_references_class_constant(lhs, ctx) || expr_references_class_constant(rhs, ctx) + } + rumoca_core::Expression::Unary { rhs, .. } => expr_references_class_constant(rhs, ctx), + rumoca_core::Expression::BuiltinCall { args, .. } + | rumoca_core::Expression::FunctionCall { args, .. } + | rumoca_core::Expression::Array { elements: args, .. } + | rumoca_core::Expression::Tuple { elements: args, .. } => args + .iter() + .any(|expr| expr_references_class_constant(expr, ctx)), + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + branches.iter().any(|(cond, value)| { + expr_references_class_constant(cond, ctx) + || expr_references_class_constant(value, ctx) + }) || expr_references_class_constant(else_branch, ctx) + } + rumoca_core::Expression::Range { + start, step, end, .. + } => { + expr_references_class_constant(start, ctx) + || step + .as_ref() + .is_some_and(|expr| expr_references_class_constant(expr, ctx)) + || expr_references_class_constant(end, ctx) + } + rumoca_core::Expression::ArrayComprehension { + expr, + indices, + filter, + .. + } => { + expr_references_class_constant(expr, ctx) + || indices + .iter() + .any(|index| expr_references_class_constant(&index.range, ctx)) + || filter + .as_ref() + .is_some_and(|expr| expr_references_class_constant(expr, ctx)) + } + rumoca_core::Expression::Index { + base, subscripts, .. + } => { + expr_references_class_constant(base, ctx) + || subscripts_reference_class_constant(subscripts, ctx) + } + rumoca_core::Expression::FieldAccess { base, .. } => { + expr_references_class_constant(base, ctx) + } + rumoca_core::Expression::Literal { .. } | rumoca_core::Expression::Empty { .. } => false, + } +} + +fn subscripts_reference_class_constant( + subscripts: &[rumoca_core::Subscript], + ctx: &Context, +) -> bool { + subscripts.iter().any(|subscript| { + matches!( + subscript, + rumoca_core::Subscript::Expr { expr, .. } + if expr_references_class_constant(expr, ctx) + ) + }) +} + +fn variable_is_string_type(var: &flat::Variable, type_name: Option<&String>) -> bool { + var.type_id == rumoca_core::TypeId(3) + || type_name.is_some_and(|name| rumoca_core::qualified_type_name_matches(name, "String")) +} + +fn recover_string_literal_opt_expr(expr: &mut Option, scope: &str) { + let Some(recovered) = expr.as_ref().and_then(|expr| { + crate::variables::recover_string_literal_from_invalid_component_expr(expr, scope) + }) else { + return; + }; + *expr = Some(recovered); +} + fn substitute_function_bodies( functions: &mut flat::VarNameIndexMap, ctx: &Context, @@ -1026,6 +2292,54 @@ fn substitute_opt_expr_with_options( Ok(()) } +fn substitute_opt_expr_with_options_and_dims( + expr: &mut Option, + ctx: &Context, + live_vars: &rustc_hash::FxHashSet, + locals: &HashSet, + scope: &str, + prefer_scoped_parameters: bool, + var_dims: &rustc_hash::FxHashMap>, +) -> Result<(), FlattenError> { + if let Some(expr) = expr { + *expr = substitute_known_constants_expr_with_options_and_dims( + expr.clone(), + ctx, + live_vars, + locals, + scope, + prefer_scoped_parameters, + Some(var_dims), + )?; + } + Ok(()) +} + +fn substitute_opt_expr_with_options_dims_and_values( + expr: &mut Option, + ctx: &Context, + live_vars: &rustc_hash::FxHashSet, + locals: &HashSet, + scope: &str, + prefer_scoped_parameters: bool, + var_dims: &rustc_hash::FxHashMap>, + var_values: &rustc_hash::FxHashMap, +) -> Result<(), FlattenError> { + if let Some(expr) = expr { + *expr = substitute_known_constants_expr_with_options_dims_and_values( + expr.clone(), + ctx, + live_vars, + locals, + scope, + prefer_scoped_parameters, + var_dims, + var_values, + )?; + } + Ok(()) +} + pub(crate) fn substitute_known_constants_expr( expr: rumoca_core::Expression, ctx: &Context, @@ -1043,6 +2357,52 @@ fn substitute_known_constants_expr_with_options( locals: &HashSet, scope: &str, prefer_scoped_parameters: bool, +) -> Result { + substitute_known_constants_expr_with_options_and_dims( + expr, + ctx, + live_vars, + locals, + scope, + prefer_scoped_parameters, + None, + ) +} + +fn substitute_known_constants_expr_with_options_and_dims( + expr: rumoca_core::Expression, + ctx: &Context, + live_vars: &rustc_hash::FxHashSet, + locals: &HashSet, + scope: &str, + prefer_scoped_parameters: bool, + var_dims: Option<&rustc_hash::FxHashMap>>, +) -> Result { + KnownConstantSubstituter { + env: ConstantSubstitutionEnv { + ctx, + live_vars, + locals, + scope, + prefer_scoped_parameters, + var_dims, + var_values: None, + restrict_live_values_to_bindings: false, + resolving_value: None, + }, + } + .rewrite_expression(&expr) +} + +fn substitute_known_constants_expr_with_options_dims_and_values( + expr: rumoca_core::Expression, + ctx: &Context, + live_vars: &rustc_hash::FxHashSet, + locals: &HashSet, + scope: &str, + prefer_scoped_parameters: bool, + var_dims: &rustc_hash::FxHashMap>, + var_values: &rustc_hash::FxHashMap, ) -> Result { KnownConstantSubstituter { env: ConstantSubstitutionEnv { @@ -1051,6 +2411,10 @@ fn substitute_known_constants_expr_with_options( locals, scope, prefer_scoped_parameters, + var_dims: Some(var_dims), + var_values: Some(var_values), + restrict_live_values_to_bindings: true, + resolving_value: None, }, } .rewrite_expression(&expr) @@ -1067,6 +2431,16 @@ struct ConstantSubstitutionEnv<'a> { locals: &'a HashSet, scope: &'a str, prefer_scoped_parameters: bool, + var_dims: Option<&'a rustc_hash::FxHashMap>>, + var_values: Option<&'a rustc_hash::FxHashMap>, + restrict_live_values_to_bindings: bool, + resolving_value: Option<&'a ResolvingValue<'a>>, +} + +#[derive(Clone, Copy)] +struct ResolvingValue<'a> { + key: &'a str, + parent: Option<&'a ResolvingValue<'a>>, } impl<'a> ConstantSubstitutionEnv<'a> { @@ -1077,6 +2451,24 @@ impl<'a> ConstantSubstitutionEnv<'a> { locals: self.locals, scope, prefer_scoped_parameters: self.prefer_scoped_parameters, + var_dims: self.var_dims, + var_values: self.var_values, + restrict_live_values_to_bindings: self.restrict_live_values_to_bindings, + resolving_value: self.resolving_value, + } + } + + fn with_locals<'b>(&'b self, locals: &'b HashSet) -> ConstantSubstitutionEnv<'b> { + ConstantSubstitutionEnv { + ctx: self.ctx, + live_vars: self.live_vars, + locals, + scope: self.scope, + prefer_scoped_parameters: self.prefer_scoped_parameters, + var_dims: self.var_dims, + var_values: self.var_values, + restrict_live_values_to_bindings: self.restrict_live_values_to_bindings, + resolving_value: self.resolving_value, } } } @@ -1097,6 +2489,22 @@ impl FallibleExpressionRewriter for KnownConstantSubstituter<'_> { rumoca_core::Expression::FieldAccess { base, field, span } => { self.rewrite_field_access(base, field, *span) } + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + args, + span, + } => { + if let Some(size_expr) = self.rewrite_size_from_declared_dims(args, *span) { + return Ok(size_expr); + } + self.walk_expression(expr) + } + rumoca_core::Expression::ArrayComprehension { + expr, + indices, + filter, + span, + } => self.rewrite_array_comprehension(expr, indices, filter.as_deref(), *span), other => self.walk_expression(other), } } @@ -1116,6 +2524,66 @@ impl FallibleExpressionRewriter for KnownConstantSubstituter<'_> { } impl KnownConstantSubstituter<'_> { + fn rewrite_array_comprehension( + &self, + expr: &rumoca_core::Expression, + indices: &[rumoca_core::ComprehensionIndex], + filter: Option<&rumoca_core::Expression>, + span: rumoca_core::Span, + ) -> Result { + let mut active_locals = self.env.locals.clone(); + let mut rewritten_indices = Vec::with_capacity(indices.len()); + for index in indices { + let range = KnownConstantSubstituter { + env: self.env.with_locals(&active_locals), + } + .rewrite_expression(&index.range)?; + rewritten_indices.push(rumoca_core::ComprehensionIndex { + name: index.name.clone(), + range, + }); + active_locals.insert(index.name.clone()); + } + + let mut bound_rewriter = KnownConstantSubstituter { + env: self.env.with_locals(&active_locals), + }; + Ok(rumoca_core::Expression::ArrayComprehension { + expr: Box::new(bound_rewriter.rewrite_expression(expr)?), + indices: rewritten_indices, + filter: filter + .map(|filter| bound_rewriter.rewrite_expression(filter).map(Box::new)) + .transpose()?, + span, + }) + } + + fn rewrite_size_from_declared_dims( + &self, + args: &[rumoca_core::Expression], + span: rumoca_core::Span, + ) -> Option { + let [base, dim] = args else { + return None; + }; + let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = base + else { + return None; + }; + if !subscripts.is_empty() { + return None; + } + let dim = literal_integer(dim)?; + let dim_idx = usize::try_from(dim.checked_sub(1)?).ok()?; + let value = *declared_dims_for_reference(name.as_str(), self.env)?.get(dim_idx)?; + Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + span, + }) + } + fn rewrite_var_ref( &mut self, name: &rumoca_core::Reference, @@ -1141,13 +2609,9 @@ impl KnownConstantSubstituter<'_> { span, }); } - if let Some(replaced) = substitute_indexed_constant_var_ref( - name, - rewritten_subscripts.clone(), - span, - self.env.ctx, - self.env.live_vars, - ) { + if let Some(replaced) = + substitute_indexed_constant_var_ref(name, rewritten_subscripts.clone(), span, self.env)? + { return Ok(replaced); } @@ -1175,6 +2639,11 @@ impl KnownConstantSubstituter<'_> { if let Some(named_arg) = named_constructor_arg(args, field) { return Ok(named_arg.clone().with_span(span)); } + if let Some(positional_arg) = + positional_constructor_arg_for_field(name.as_str(), args, field, self.env.ctx) + { + return Ok(positional_arg.clone().with_span(span)); + } if args.is_empty() && let Some(resolved) = resolve_constant_field_access(name.as_str(), field, span, self.env.ctx) @@ -1215,6 +2684,23 @@ impl KnownConstantSubstituter<'_> { } } +fn declared_dims_for_reference<'a>( + key: &str, + env: ConstantSubstitutionEnv<'a>, +) -> Option<&'a Vec> { + let var_dims = env.var_dims?; + if let Some(dims) = var_dims.get(key) { + return Some(dims); + } + if !env.scope.is_empty() && !key.contains('.') { + let scoped_key = format!("{}.{}", env.scope, key); + if let Some(dims) = var_dims.get(scoped_key.as_str()) { + return Some(dims); + } + } + None +} + fn resolve_indexed_constant_field_access( base: &rumoca_core::Expression, subscripts: &[rumoca_core::Subscript], @@ -1249,8 +2735,10 @@ fn select_constant_index( select_constant_index_element(&base, *value, rest, span, ctx) } rumoca_core::Subscript::Expr { expr, .. } => { - let index = literal_integer(expr)?; - select_constant_index_element(&base, index, rest, span, ctx) + if let Some(index) = literal_integer(expr) { + return select_constant_index_element(&base, index, rest, span, ctx); + } + select_constant_index_symbolic_element(&base, expr, rest, span, ctx) } rumoca_core::Subscript::Colon { .. } => { let rumoca_core::Expression::Array { elements, .. } = base else { @@ -1269,6 +2757,35 @@ fn select_constant_index( } } +fn select_constant_index_symbolic_element( + base: &rumoca_core::Expression, + index: &rumoca_core::Expression, + rest: &[rumoca_core::Subscript], + span: rumoca_core::Span, + ctx: &Context, +) -> Option { + match base { + rumoca_core::Expression::ArrayComprehension { + expr, + indices, + filter, + .. + } => select_array_comprehension_symbolic_index_element( + expr, + indices, + filter.as_deref(), + index, + rest, + span, + ctx, + ), + rumoca_core::Expression::Binary { op, lhs, rhs, .. } => { + select_binary_symbolic_index_element(op.clone(), lhs, rhs, index, rest, span, ctx) + } + _ => None, + } +} + fn select_constant_index_element( base: &rumoca_core::Expression, index: i64, @@ -1276,12 +2793,225 @@ fn select_constant_index_element( span: rumoca_core::Span, ctx: &Context, ) -> Option { - let rumoca_core::Expression::Array { elements, .. } = base else { + match base { + rumoca_core::Expression::Array { elements, .. } => { + let zero_based = usize::try_from(index.checked_sub(1)?).ok()?; + let element = elements.get(zero_based)?; + select_constant_index(element, rest, span, ctx) + } + rumoca_core::Expression::ArrayComprehension { + expr, + indices, + filter, + .. + } => select_array_comprehension_index_element( + expr, + indices, + filter.as_deref(), + index, + rest, + span, + ctx, + ), + rumoca_core::Expression::Binary { op, lhs, rhs, .. } => { + select_binary_index_element(op.clone(), lhs, rhs, index, rest, span, ctx) + } + _ => None, + } +} + +fn select_array_comprehension_index_element( + expr: &rumoca_core::Expression, + indices: &[rumoca_core::ComprehensionIndex], + filter: Option<&rumoca_core::Expression>, + index: i64, + rest: &[rumoca_core::Subscript], + span: rumoca_core::Span, + ctx: &Context, +) -> Option { + if filter.is_some() || indices.len() != 1 { + return None; + } + let value = comprehension_index_value(&indices[0].range, index)?; + let selected = substitute_comprehension_index_literal(expr, &indices[0].name, value, span); + select_constant_index(&selected, rest, span, ctx) +} + +fn select_array_comprehension_symbolic_index_element( + expr: &rumoca_core::Expression, + indices: &[rumoca_core::ComprehensionIndex], + filter: Option<&rumoca_core::Expression>, + index: &rumoca_core::Expression, + rest: &[rumoca_core::Subscript], + span: rumoca_core::Span, + ctx: &Context, +) -> Option { + if filter.is_some() || indices.len() != 1 { + return None; + } + let selected = substitute_comprehension_index_expr(expr, &indices[0].name, index, span); + select_constant_index(&selected, rest, span, ctx) +} + +fn select_binary_index_element( + op: rumoca_core::OpBinary, + lhs: &rumoca_core::Expression, + rhs: &rumoca_core::Expression, + index: i64, + rest: &[rumoca_core::Subscript], + span: rumoca_core::Span, + ctx: &Context, +) -> Option { + let lhs_selected = select_constant_index_element(lhs, index, rest, span, ctx); + let rhs_selected = select_constant_index_element(rhs, index, rest, span, ctx); + match (lhs_selected, rhs_selected) { + (Some(lhs), Some(rhs)) => Some(rumoca_core::Expression::Binary { + op, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span, + }), + (Some(lhs), None) => Some(rumoca_core::Expression::Binary { + op, + lhs: Box::new(lhs), + rhs: Box::new(rhs.clone().with_span(span)), + span, + }), + (None, Some(rhs)) => Some(rumoca_core::Expression::Binary { + op, + lhs: Box::new(lhs.clone().with_span(span)), + rhs: Box::new(rhs), + span, + }), + (None, None) => None, + } +} + +fn select_binary_symbolic_index_element( + op: rumoca_core::OpBinary, + lhs: &rumoca_core::Expression, + rhs: &rumoca_core::Expression, + index: &rumoca_core::Expression, + rest: &[rumoca_core::Subscript], + span: rumoca_core::Span, + ctx: &Context, +) -> Option { + let lhs_selected = select_constant_index_symbolic_element(lhs, index, rest, span, ctx); + let rhs_selected = select_constant_index_symbolic_element(rhs, index, rest, span, ctx); + match (lhs_selected, rhs_selected) { + (Some(lhs), Some(rhs)) => Some(rumoca_core::Expression::Binary { + op, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span, + }), + (Some(lhs), None) => Some(rumoca_core::Expression::Binary { + op, + lhs: Box::new(lhs), + rhs: Box::new(rhs.clone().with_span(span)), + span, + }), + (None, Some(rhs)) => Some(rumoca_core::Expression::Binary { + op, + lhs: Box::new(lhs.clone().with_span(span)), + rhs: Box::new(rhs), + span, + }), + (None, None) => None, + } +} + +fn comprehension_index_value(range: &rumoca_core::Expression, one_based_index: i64) -> Option { + let rumoca_core::Expression::Range { + start, step, end, .. + } = range + else { return None; }; - let zero_based = usize::try_from(index.checked_sub(1)?).ok()?; - let element = elements.get(zero_based)?; - select_constant_index(element, rest, span, ctx) + let start = literal_integer(start)?; + let step = match step.as_deref() { + Some(step) => literal_integer(step)?, + None => 1, + }; + let value = start + (one_based_index.checked_sub(1)?) * step; + let end = literal_integer(end)?; + ((step > 0 && value <= end) || (step < 0 && value >= end) || (step == 0 && value == start)) + .then_some(value) +} + +fn substitute_comprehension_index_literal( + expr: &rumoca_core::Expression, + name: &str, + value: i64, + span: rumoca_core::Span, +) -> rumoca_core::Expression { + substitute_comprehension_index_expr( + expr, + name, + &rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + span, + }, + span, + ) +} + +fn substitute_comprehension_index_expr( + expr: &rumoca_core::Expression, + name: &str, + replacement: &rumoca_core::Expression, + span: rumoca_core::Span, +) -> rumoca_core::Expression { + struct Substituter<'a> { + name: &'a str, + replacement: &'a rumoca_core::Expression, + span: rumoca_core::Span, + } + + impl rumoca_core::ExpressionRewriter for Substituter<'_> { + fn walk_var_ref_expression( + &mut self, + name: &rumoca_core::Reference, + subscripts: &[rumoca_core::Subscript], + span: rumoca_core::Span, + ) -> rumoca_core::Expression { + if name.as_str() == self.name && subscripts.is_empty() { + return self.replacement.clone().with_span(self.span); + } + rumoca_core::Expression::VarRef { + name: name.clone(), + subscripts: self.rewrite_subscripts(subscripts), + span, + } + } + + fn walk_array_comprehension_expression( + &mut self, + expr: &rumoca_core::Expression, + indices: &[rumoca_core::ComprehensionIndex], + filter: Option<&rumoca_core::Expression>, + span: rumoca_core::Span, + ) -> rumoca_core::Expression { + if indices.iter().any(|index| index.name == self.name) { + return rumoca_core::Expression::ArrayComprehension { + expr: Box::new(expr.clone()), + indices: indices.to_vec(), + filter: filter.cloned().map(Box::new), + span, + }; + } + rumoca_core::ExpressionRewriter::walk_array_comprehension_expression( + self, expr, indices, filter, span, + ) + } + } + + let mut substituter = Substituter { + name, + replacement, + span, + }; + substituter.rewrite_expression(expr) } fn resolve_constant_expr_alias( @@ -1322,6 +3052,11 @@ fn resolve_field_on_constant_expr( } => named_constructor_arg(args, field) .cloned() .map(|expr| expr.with_span(span)) + .or_else(|| { + positional_constructor_arg_for_field(name.as_str(), args, field, ctx) + .cloned() + .map(|expr| expr.with_span(span)) + }) .or_else(|| { args.is_empty() .then(|| resolve_constant_field_access(name.as_str(), field, span, ctx)) @@ -1346,25 +3081,89 @@ fn literal_integer(expr: &rumoca_core::Expression) -> Option { } } -impl FallibleStatementRewriter for KnownConstantSubstituter<'_> {} +impl FallibleStatementRewriter for KnownConstantSubstituter<'_> { + fn rewrite_statement( + &mut self, + statement: &rumoca_core::Statement, + ) -> Result { + let rumoca_core::Statement::For { + indices, + equations, + span, + } = statement + else { + return self.walk_statement(statement); + }; + + let mut active_locals = self.env.locals.clone(); + let mut rewritten_indices = Vec::with_capacity(indices.len()); + for index in indices { + let range = KnownConstantSubstituter { + env: self.env.with_locals(&active_locals), + } + .rewrite_expression(&index.range)?; + rewritten_indices.push(rumoca_core::ForIndex { + ident: index.ident.clone(), + range, + }); + active_locals.insert(index.ident.clone()); + } + let equations = KnownConstantSubstituter { + env: self.env.with_locals(&active_locals), + } + .rewrite_statements(equations)?; + + Ok(rumoca_core::Statement::For { + indices: rewritten_indices, + equations, + span: *span, + }) + } +} fn substitute_indexed_constant_var_ref( name: &rumoca_core::Reference, subscripts: Vec, span: rumoca_core::Span, - ctx: &Context, - live_vars: &rustc_hash::FxHashSet, -) -> Option { - if live_vars.contains(name.as_str()) { - return None; + env: ConstantSubstitutionEnv<'_>, +) -> Result, FlattenError> { + let constant_expr = match resolve_constant_value_expr_for_ref(name, env.ctx) { + Some(expr) => expr.clone(), + None => return Ok(None), + }; + if env.live_vars.contains(name.as_str()) { + return Ok(symbolic_alias_expr(&constant_expr).map(|base| { + rumoca_core::Expression::Index { + base: Box::new(base.with_span(span)), + subscripts, + span, + } + })); + } + + if let Some(selected) = select_constant_index(&constant_expr, &subscripts, span, env.ctx) { + return Ok(Some(substitute_resolved_constant_expr( + name.as_str(), + &selected, + span, + env, + )?)); } - let constant_expr = resolve_constant_value_expr_for_ref(name, ctx)?.clone(); - Some(rumoca_core::Expression::Index { + Ok(Some(rumoca_core::Expression::Index { base: Box::new(constant_expr), subscripts, span, - }) + })) +} + +fn symbolic_alias_expr(expr: &rumoca_core::Expression) -> Option { + match expr { + rumoca_core::Expression::VarRef { .. } + | rumoca_core::Expression::Index { .. } + | rumoca_core::Expression::FieldAccess { .. } => Some(expr.clone()), + _ => None, + } } fn substitute_scalar_var_ref( @@ -1373,12 +3172,42 @@ fn substitute_scalar_var_ref( env: ConstantSubstitutionEnv<'_>, ) -> Result, FlattenError> { let key = name.as_str(); - if env.live_vars.contains(key) || reference_root_is_local(name, env.locals) { + if reference_root_is_local(name, env.locals) { + return Ok(None); + } + if let Some(expr) = substitute_flat_variable_value_ref(key, span, env)? { + return Ok(Some(expr)); + } + if env.live_vars.contains(key) { + // The broader context also contains pre-evaluated discrete values used + // for structural branch selection. Keep those runtime-visible unless + // compiler-owned metadata independently proves that the name is a + // class constant or structural parameter. Class constants may not + // retain their source binding on the Flat declaration, so requiring a + // declaration value alone would leak them into DAE/Solve as unbound + // runtime references. + if env.restrict_live_values_to_bindings && !live_value_is_proven_structural(key, env.ctx) { + return Ok(None); + } + if parameter_is_non_structural(key, env) { + return Ok(None); + } + if let Some(expr) = structured_key_value_expr(key, span, env)? { + return Ok(Some(expr)); + } + if reference_key_is_structured(key) + && let Some(v) = resolve_constant_value_expr_for_ref(name, env.ctx) + { + return Ok(Some(substitute_resolved_constant_expr(key, v, span, env)?)); + } return Ok(None); } if inline_index_base_is_live_or_local(key, env.live_vars, env.locals) { return Ok(None); } + if let Some(scoped_live_ref) = scoped_live_var_ref(key, span, env) { + return Ok(Some(scoped_live_ref)); + } let has_array_shape = reference_has_array_shape(name, key, env.ctx, env.scope); if env.prefer_scoped_parameters && !env.scope.is_empty() @@ -1386,10 +3215,40 @@ fn substitute_scalar_var_ref( { return Ok(Some(expr)); } + if !env.prefer_scoped_parameters + && !env.scope.is_empty() + && let Some(expr) = substitute_scoped_scalar_var_ref(key, span, env)? + { + return Ok(Some(expr)); + } + if let Some(expr) = substitute_alias_resolved_scalar_var_ref(key, span, env)? { + return Ok(Some(expr)); + } + if !env.scope.is_empty() + && name.target_def_id().is_some() + && let Some(expr) = substitute_direct_scoped_def_id_scalar_var_ref(name, span, env)? + { + return Ok(Some(expr)); + } + if let Some(expr) = structured_key_value_expr(key, span, env)? { + return Ok(Some(expr)); + } if let Some(v) = resolve_constant_value_expr_for_ref(name, env.ctx) { if has_array_shape && !constant_expr_preserves_array_shape(v) { return Ok(None); } + if expression_is_record_constructor(v) + && let Some(expr) = resolve_projected_constant_path(key, span, env.ctx) + { + return Ok(Some(substitute_known_constants_expr_with_options( + expr, + env.ctx, + env.live_vars, + env.locals, + env.scope, + env.prefer_scoped_parameters, + )?)); + } return Ok(Some(substitute_resolved_constant_expr(key, v, span, env)?)); } if !has_array_shape && let Some(literal) = scalar_parameter_literal(key, span, env.ctx) { @@ -1408,14 +3267,202 @@ fn substitute_scalar_var_ref( env.prefer_scoped_parameters, )?)); } - if !env.prefer_scoped_parameters - && !env.scope.is_empty() - && let Some(expr) = substitute_scoped_scalar_var_ref(key, span, env)? + substitute_alias_resolved_scalar_var_ref(key, span, env) +} + +fn live_value_is_proven_structural(key: &str, ctx: &Context) -> bool { + ctx.class_constant_keys.contains(key) || ctx.structural_params.contains(key) +} + +fn def_id_scoped_lookup_key(name: &rumoca_core::Reference) -> &str { + if crate::path_utils::is_nested_name(name.as_str()) { + name.last_segment() + } else { + name.as_str() + } +} + +fn substitute_direct_scoped_def_id_scalar_var_ref( + name: &rumoca_core::Reference, + span: rumoca_core::Span, + env: ConstantSubstitutionEnv<'_>, +) -> Result, FlattenError> { + if !env.scope.is_empty() && name.as_str().starts_with(&format!("{}.", env.scope)) { + return Ok(None); + } + let key = def_id_scoped_lookup_key(name); + if key.contains('.') { + return Ok(None); + } + let candidate = format!("{}.{}", env.scope, key); + if env.live_vars.contains(&candidate) + || inline_index_base_is_live_or_local(&candidate, env.live_vars, env.locals) { - return Ok(Some(expr)); + return Ok(Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new(candidate), + subscripts: vec![], + span, + })); + } + let candidate_has_array_shape = env + .ctx + .array_dimensions + .get(&candidate) + .is_some_and(|dims| !dims.is_empty()); + if !candidate_has_array_shape + && let Some(literal) = evaluated_scalar_parameter_literal(&candidate, span, env.ctx) + { + return Ok(Some(literal)); } + let Some(value) = resolve_constant_value_expr(&candidate, env.ctx) else { + return Ok(None); + }; + if candidate_has_array_shape && !constant_expr_preserves_array_shape(value) { + return Ok(None); + } + let candidate_scope = parent_component_scope(&candidate); + substitute_resolved_constant_expr(&candidate, value, span, env.with_scope(&candidate_scope)) + .map(Some) +} - substitute_alias_resolved_scalar_var_ref(key, span, env) +fn substitute_flat_variable_value_ref( + key: &str, + span: rumoca_core::Span, + env: ConstantSubstitutionEnv<'_>, +) -> Result, FlattenError> { + let Some(var_values) = env.var_values else { + return Ok(None); + }; + for candidate in scoped_lookup_candidates(key, env.scope) { + if resolving_value_stack_contains(env.resolving_value, &candidate) { + continue; + } + let Some(value) = var_values.get(&candidate) else { + continue; + }; + if reference_key_has_array_shape(&candidate, env.ctx, env.scope) + && !constant_expr_preserves_array_shape(value) + { + continue; + } + return Ok(Some(substitute_resolved_constant_expr( + &candidate, value, span, env, + )?)); + } + Ok(None) +} + +fn resolving_value_stack_contains(mut node: Option<&ResolvingValue<'_>>, key: &str) -> bool { + while let Some(value) = node { + if value.key == key { + return true; + } + node = value.parent; + } + false +} + +fn parameter_is_non_structural(key: &str, env: ConstantSubstitutionEnv<'_>) -> bool { + env.ctx.non_structural_params.contains(key) + || (!env.scope.is_empty() + && !key.contains('.') + && env + .ctx + .non_structural_params + .contains(format!("{}.{}", env.scope, key).as_str())) +} + +fn reference_key_is_structured(key: &str) -> bool { + key.contains('.') || key.contains('[') +} + +fn scoped_name_is_structural_value(key: &str, ctx: &Context) -> bool { + ctx.constant_values.contains_key(key) + || ctx.class_constant_keys.contains(key) + || ctx.structural_params.contains(key) + || ctx.parameter_values.contains_key(key) + || ctx.boolean_parameter_values.contains_key(key) + || ctx.string_parameter_values.contains_key(key) + || ctx.enum_parameter_values.contains_key(key) +} + +fn structured_key_value_expr( + key: &str, + span: rumoca_core::Span, + env: ConstantSubstitutionEnv<'_>, +) -> Result, FlattenError> { + if !reference_key_is_structured(key) || !scoped_name_is_structural_value(key, env.ctx) { + return Ok(None); + } + if let Some(enum_name) = env.ctx.enum_parameter_values.get(key) { + return Ok(Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new(enum_name.clone()), + subscripts: vec![], + span, + })); + } + if let Some(v) = resolve_constant_value_expr(key, env.ctx) { + if expression_is_record_constructor(v) + && let Some(expr) = resolve_projected_constant_path(key, span, env.ctx) + { + return Ok(Some(substitute_known_constants_expr_with_options( + expr, + env.ctx, + env.live_vars, + env.locals, + env.scope, + env.prefer_scoped_parameters, + )?)); + } + return Ok(Some(substitute_resolved_constant_expr(key, v, span, env)?)); + } + if let Some(value) = env.ctx.parameter_values.get(key) { + return Ok(Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(*value), + span, + })); + } + if env.ctx.class_constant_keys.contains(key) + && let Some(value) = env.ctx.real_parameter_values.get(key) + && value.is_finite() + { + return Ok(Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(*value), + span, + })); + } + if let Some(value) = env.ctx.boolean_parameter_values.get(key) { + return Ok(Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Boolean(*value), + span, + })); + } + if let Some(value) = env.ctx.string_parameter_values.get(key) { + return Ok(Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String(value.clone()), + span, + })); + } + Ok(None) +} + +fn scoped_live_var_ref( + key: &str, + span: rumoca_core::Span, + env: ConstantSubstitutionEnv<'_>, +) -> Option { + if env.scope.is_empty() || key.contains('.') { + return None; + } + let scoped_key = format!("{}.{}", env.scope, key); + if !env.live_vars.contains(&scoped_key) { + return None; + } + Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new(scoped_key), + subscripts: vec![], + span, + }) } fn substitute_alias_resolved_scalar_var_ref( @@ -1444,7 +3491,7 @@ fn substitute_alias_resolved_scalar_var_ref( )?)); } if !reference_key_has_array_shape(&resolved_key, env.ctx, env.scope) - && let Some(literal) = scalar_parameter_literal(&resolved_key, span, env.ctx) + && let Some(literal) = evaluated_scalar_parameter_literal(&resolved_key, span, env.ctx) { return Ok(Some(literal)); } @@ -1470,14 +3517,35 @@ fn substitute_resolved_constant_expr( } else { &declaration_scope }; - substitute_known_constants_expr_with_options( - expr.clone().with_span(span), - env.ctx, - env.live_vars, - env.locals, - scope, - env.prefer_scoped_parameters, - ) + let mut expr = expr.clone().with_span(span); + if let Some(var_dims) = env.var_dims + && let Some(dims) = var_dims.get(key) + { + let mut maybe_expr = Some(expr); + reconcile_constructor_extent_expr(&mut maybe_expr, dims); + expr = maybe_expr.expect("reconciled expression should remain present"); + } + let resolving_value = ResolvingValue { + key, + parent: env.resolving_value, + }; + let declaration_locals = HashSet::new(); + KnownConstantSubstituter { + env: ConstantSubstitutionEnv { + ctx: env.ctx, + live_vars: env.live_vars, + // A variable binding is evaluated in its declaration scope, not in + // the caller's algorithm loop/comprehension scope. + locals: &declaration_locals, + scope, + prefer_scoped_parameters: env.prefer_scoped_parameters, + var_dims: env.var_dims, + var_values: env.var_values, + restrict_live_values_to_bindings: env.restrict_live_values_to_bindings, + resolving_value: Some(&resolving_value), + }, + } + .rewrite_expression(&expr) } fn constant_expr_preserves_array_shape(expr: &rumoca_core::Expression) -> bool { @@ -1495,11 +3563,27 @@ fn constant_expr_preserves_array_shape(expr: &rumoca_core::Expression) -> bool { ) } +fn expression_is_record_constructor(expr: &rumoca_core::Expression) -> bool { + matches!( + expr, + rumoca_core::Expression::FunctionCall { + is_constructor: true, + .. + } + ) +} + fn scalar_parameter_literal( key: &str, span: rumoca_core::Span, ctx: &Context, ) -> Option { + if let Some(v) = ctx.parameter_values.get(key) { + return Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(*v), + span, + }); + } if let Some(v) = ctx.real_parameter_values.get(key) && v.is_finite() { @@ -1508,18 +3592,58 @@ fn scalar_parameter_literal( span, }); } + if let Some(v) = ctx.boolean_parameter_values.get(key) { + return Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Boolean(*v), + span, + }); + } + if let Some(v) = ctx.string_parameter_values.get(key) { + return Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String(v.clone()), + span, + }); + } + ctx.enum_parameter_values + .get(key) + .map(|v| rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new(v.clone()), + subscripts: vec![], + span, + }) +} + +fn evaluated_scalar_parameter_literal( + key: &str, + span: rumoca_core::Span, + ctx: &Context, +) -> Option { if let Some(v) = ctx.parameter_values.get(key) { return Some(rumoca_core::Expression::Literal { value: rumoca_core::Literal::Integer(*v), span, }); } + if let Some(v) = ctx.real_parameter_values.get(key) + && v.is_finite() + { + return Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(*v), + span, + }); + } if let Some(v) = ctx.boolean_parameter_values.get(key) { return Some(rumoca_core::Expression::Literal { value: rumoca_core::Literal::Boolean(*v), span, }); } + if let Some(v) = ctx.string_parameter_values.get(key) { + return Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String(v.clone()), + span, + }); + } ctx.enum_parameter_values .get(key) .map(|v| rumoca_core::Expression::VarRef { @@ -1564,10 +3688,19 @@ fn substitute_scoped_scalar_var_ref( if candidate == key { continue; } + if env.live_vars.contains(&candidate) + || inline_index_base_is_live_or_local(&candidate, env.live_vars, env.locals) + { + return Ok(Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new(candidate), + subscripts: vec![], + span, + })); + } let candidate_has_array_shape = reference_key_has_array_shape(&candidate, env.ctx, &candidate_scope); if !candidate_has_array_shape - && let Some(literal) = scalar_parameter_literal(&candidate, span, env.ctx) + && let Some(literal) = evaluated_scalar_parameter_literal(&candidate, span, env.ctx) { return Ok(Some(literal)); } @@ -1576,6 +3709,18 @@ fn substitute_scoped_scalar_var_ref( continue; } let candidate_env = env.with_scope(&candidate_scope); + if expression_is_record_constructor(v) + && let Some(expr) = resolve_projected_constant_path(&candidate, span, env.ctx) + { + return Ok(Some(substitute_known_constants_expr_with_options( + expr, + env.ctx, + env.live_vars, + env.locals, + &candidate_scope, + env.prefer_scoped_parameters, + )?)); + } return Ok(Some(substitute_resolved_constant_expr( &candidate, v, @@ -1766,11 +3911,13 @@ fn constant_key_or_prefix_exists(name: &str, ctx: &Context) -> bool { || ctx.real_parameter_values.contains_key(name) || ctx.parameter_values.contains_key(name) || ctx.boolean_parameter_values.contains_key(name) + || ctx.string_parameter_values.contains_key(name) || ctx.enum_parameter_values.contains_key(name) || map_has_key_prefix(&ctx.constant_values, name) || map_has_key_prefix(&ctx.real_parameter_values, name) || map_has_key_prefix(&ctx.parameter_values, name) || map_has_key_prefix(&ctx.boolean_parameter_values, name) + || map_has_key_prefix(&ctx.string_parameter_values, name) || map_has_key_prefix(&ctx.enum_parameter_values, name) } @@ -1837,6 +3984,33 @@ fn named_constructor_arg<'a>( None } +fn positional_constructor_arg_for_field<'a>( + constructor_name: &str, + args: &'a [rumoca_core::Expression], + field: &str, + ctx: &Context, +) -> Option<&'a rumoca_core::Expression> { + let function = ctx.functions.get(constructor_name)?; + if !function.is_constructor { + return None; + } + let index = function + .inputs + .iter() + .position(|input| input.name == field)?; + args.iter() + .filter(|arg| !is_named_constructor_arg(arg)) + .nth(index) +} + +fn is_named_constructor_arg(arg: &rumoca_core::Expression) -> bool { + matches!( + arg, + rumoca_core::Expression::FunctionCall { name, .. } + if name.as_str().starts_with(rumoca_core::NAMED_FUNCTION_ARG_PREFIX) + ) +} + fn resolve_constant_field_access( base_name: &str, field: &str, @@ -1873,6 +4047,12 @@ fn resolve_constant_field_access( span, }); } + if let Some(value) = ctx.string_parameter_values.get(&key) { + return Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String(value.clone()), + span, + }); + } if let Some(value) = ctx.enum_parameter_values.get(&key) { return Some(rumoca_core::Expression::VarRef { name: rumoca_core::Reference::new(value.clone()), @@ -1909,6 +4089,36 @@ fn substitute_known_constants_statement( locals, scope, prefer_scoped_parameters: false, + var_dims: None, + var_values: None, + restrict_live_values_to_bindings: false, + resolving_value: None, + }, + } + .rewrite_statement(statement)?; + Ok(()) +} + +fn substitute_known_constants_statement_with_dims_and_values( + statement: &mut rumoca_core::Statement, + ctx: &Context, + live_vars: &rustc_hash::FxHashSet, + locals: &HashSet, + scope: &str, + var_dims: &rustc_hash::FxHashMap>, + var_values: &rustc_hash::FxHashMap, +) -> Result<(), FlattenError> { + *statement = KnownConstantSubstituter { + env: ConstantSubstitutionEnv { + ctx, + live_vars, + locals, + scope, + prefer_scoped_parameters: true, + var_dims: Some(var_dims), + var_values: Some(var_values), + restrict_live_values_to_bindings: true, + resolving_value: None, }, } .rewrite_statement(statement)?; @@ -1921,21 +4131,22 @@ fn substitute_known_constants_when_equation( live_vars: &rustc_hash::FxHashSet, locals: &HashSet, ) -> Result<(), FlattenError> { + let substitute = |expr: &mut rumoca_core::Expression| { + *expr = substitute_known_constants_expr(expr.clone(), ctx, live_vars, locals, "")?; + Ok::<(), FlattenError>(()) + }; match equation { flat::WhenEquation::Assign { value, .. } | flat::WhenEquation::Reinit { value, .. } => { - *value = substitute_known_constants_expr(value.clone(), ctx, live_vars, locals, "")?; + substitute(value)?; } flat::WhenEquation::Assert { condition, message, .. } => { - *condition = - substitute_known_constants_expr(condition.clone(), ctx, live_vars, locals, "")?; - *message = - substitute_known_constants_expr(message.clone(), ctx, live_vars, locals, "")?; + substitute(condition)?; + substitute(message)?; } flat::WhenEquation::Terminate { message, .. } => { - *message = - substitute_known_constants_expr(message.clone(), ctx, live_vars, locals, "")?; + substitute(message)?; } flat::WhenEquation::Conditional { branches, @@ -1943,8 +4154,7 @@ fn substitute_known_constants_when_equation( .. } => { for (condition, equations) in branches { - *condition = - substitute_known_constants_expr(condition.clone(), ctx, live_vars, locals, "")?; + substitute(condition)?; for nested in equations { substitute_known_constants_when_equation(nested, ctx, live_vars, locals)?; } @@ -1954,8 +4164,7 @@ fn substitute_known_constants_when_equation( } } flat::WhenEquation::FunctionCallOutputs { function, .. } => { - *function = - substitute_known_constants_expr(function.clone(), ctx, live_vars, locals, "")?; + substitute(function)?; } } Ok(()) diff --git a/crates/rumoca-phase-flatten/src/postprocess/substitute_constant_tests.rs b/crates/rumoca-phase-flatten/src/postprocess/substitute_constant_tests.rs index ea8537e7b..748dcd415 100644 --- a/crates/rumoca-phase-flatten/src/postprocess/substitute_constant_tests.rs +++ b/crates/rumoca-phase-flatten/src/postprocess/substitute_constant_tests.rs @@ -1,3 +1,7 @@ +//! SPEC_0021 file-size exception: constant-substitution regressions still cover +//! parameter bindings, scoped aliases, and static initial algorithms together. +//! split plan: split tests by substitution source and static statement family. + use super::*; use rumoca_core::Span; @@ -34,6 +38,41 @@ fn var_ref(name: &str) -> rumoca_core::Expression { } } +fn var_ref_with_target_def_id(path: &str, def_id: rumoca_core::DefId) -> rumoca_core::Expression { + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference( + path, + rumoca_core::ComponentReference { + local: false, + span: rumoca_core::Span::DUMMY, + parts: rumoca_core::split_path_with_indices(path) + .into_iter() + .map(|ident| rumoca_core::ComponentRefPart { + ident: ident.to_string(), + span: rumoca_core::Span::DUMMY, + subs: vec![], + }) + .collect(), + def_id: Some(def_id), + }, + ), + subscripts: vec![], + span: rumoca_core::Span::DUMMY, + } +} + +fn named_arg(name: &str, value: rumoca_core::Expression) -> rumoca_core::Expression { + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new(format!( + "{}{name}", + rumoca_core::NAMED_FUNCTION_ARG_PREFIX + )), + args: vec![value], + is_constructor: true, + span: rumoca_core::Span::DUMMY, + } +} + fn spanned_var_ref(name: &str) -> rumoca_core::Expression { let var_name = rumoca_core::VarName::new(name); let component_ref = rumoca_core::component_reference_from_flat_name(&var_name, test_span()) @@ -45,66 +84,1550 @@ fn spanned_var_ref(name: &str) -> rumoca_core::Expression { } } -fn int_literal(value: i64) -> rumoca_core::Expression { - rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Integer(value), - span: rumoca_core::Span::DUMMY, - } -} +#[test] +fn substitute_known_constants_prefers_integer_parameter_binding_over_stale_real_start() { + let mut ctx = Context::new(); + ctx.parameter_values + .insert("periodicClock.factor".to_string(), 20); + ctx.real_parameter_values + .insert("periodicClock.factor".to_string(), 0.0); + + let substituted = substitute_known_constants_expr( + spanned_var_ref("periodicClock.factor"), + &ctx, + &std::collections::HashSet::default(), + &std::collections::HashSet::default(), + "", + ) + .unwrap(); + + assert_eq!( + substituted, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(20), + span: test_span(), + } + ); +} + +#[test] +fn substitute_known_constants_prefers_scoped_instance_for_def_id_sibling_parameter() { + let n_nodes_def = rumoca_core::DefId::new(42); + let mut ctx = Context::new(); + ctx.parameter_values.insert("nNodes".to_string(), 2); + ctx.parameter_values.insert("pipe.nNodes".to_string(), 20); + ctx.target_def_names + .insert(n_nodes_def, "nNodes".to_string()); + + let substituted = substitute_known_constants_expr( + var_ref_with_target_def_id("nNodes", n_nodes_def), + &ctx, + &std::collections::HashSet::default(), + &std::collections::HashSet::default(), + "pipe", + ) + .unwrap(); + + assert_eq!( + substituted, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(20), + span: rumoca_core::Span::DUMMY, + } + ); +} + +#[test] +fn substitute_known_constants_prefers_scoped_instance_for_class_qualified_def_id_parameter() { + let n_nodes_def = rumoca_core::DefId::new(43); + let mut ctx = Context::new(); + ctx.parameter_values.insert( + "Modelica.Fluid.Pipes.BaseClasses.PartialTwoPortFlow.nNodes".to_string(), + 2, + ); + ctx.parameter_values.insert("pipe.nNodes".to_string(), 20); + ctx.target_def_names.insert( + n_nodes_def, + "Modelica.Fluid.Pipes.BaseClasses.PartialTwoPortFlow.nNodes".to_string(), + ); + + let substituted = substitute_known_constants_expr( + var_ref_with_target_def_id( + "Modelica.Fluid.Pipes.BaseClasses.PartialTwoPortFlow.nNodes", + n_nodes_def, + ), + &ctx, + &std::collections::HashSet::default(), + &std::collections::HashSet::default(), + "pipe", + ) + .unwrap(); + + assert_eq!( + substituted, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(20), + span: rumoca_core::Span::DUMMY, + } + ); +} + +#[test] +fn substitute_known_constants_preserves_tunable_real_parameter_binding_in_equation() { + let mut model = flat::Model::new(); + let a_name = rumoca_core::VarName::new("a"); + model.add_variable( + a_name.clone(), + flat::Variable { + name: a_name, + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(2.0), + span: test_span(), + }), + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.add_variable( + rumoca_core::VarName::new("x"), + flat::Variable { + name: rumoca_core::VarName::new("x"), + variability: rumoca_core::Variability::Continuous(rumoca_core::Token::default()), + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.add_equation(flat::Equation::new( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs: Box::new(var_ref("a")), + rhs: Box::new(var_ref("x")), + span: test_span(), + }, + test_span(), + flat::EquationOrigin::Binding { + variable: "x".to_string(), + }, + )); + + substitute_known_constants_in_flat(&mut model, &Context::new()).unwrap(); + + let rumoca_core::Expression::Binary { lhs, .. } = &model.equations[0].residual else { + panic!("expected product residual"); + }; + assert!(matches!( + lhs.as_ref(), + rumoca_core::Expression::VarRef { name, subscripts, .. } + if name.as_str() == "a" && subscripts.is_empty() + )); +} + +#[test] +fn substitute_known_constants_recovers_clock_factor_binding_from_integer_constructor_bounds() { + let mut model = flat::Model::new(); + let factor_name = rumoca_core::VarName::new("periodicClock.factor"); + model.add_variable( + factor_name.clone(), + flat::Variable { + name: factor_name.clone(), + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(int_literal(20)), + binding_from_modification: true, + start: Some(int_literal(0)), + min: Some(int_literal(0)), + is_discrete_type: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model + .variable_type_names + .insert(factor_name, "Integer".to_string()); + model.add_equation(flat::Equation::new( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(var_ref("periodicClock.c")), + rhs: Box::new(rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Clock"), + args: vec![ + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Integer"), + args: vec![named_arg("min", int_literal(0))], + is_constructor: true, + span: test_span(), + }, + var_ref("periodicClock.resolutionFactor"), + ], + is_constructor: true, + span: test_span(), + }), + span: test_span(), + }, + test_span(), + flat::EquationOrigin::ComponentEquation { + component: "periodicClock.c".to_string(), + }, + )); + + substitute_known_constants_in_flat(&mut model, &Context::new()).unwrap(); + + let rumoca_core::Expression::Binary { rhs, .. } = &model.equations[0].residual else { + panic!("expected residual assignment"); + }; + let rumoca_core::Expression::FunctionCall { args, .. } = rhs.as_ref() else { + panic!("expected Clock call"); + }; + assert!(matches!( + args[0], + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(20), + .. + } + )); +} + +#[test] +fn substitute_known_constants_recovers_subsample_factor_binding_from_integer_constructor_bounds() { + let mut model = flat::Model::new(); + let factor_name = rumoca_core::VarName::new("subSample1.factor"); + model.add_variable( + factor_name.clone(), + flat::Variable { + name: factor_name.clone(), + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(int_literal(5)), + binding_from_modification: true, + start: Some(int_literal(0)), + min: Some(int_literal(1)), + is_discrete_type: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model + .variable_type_names + .insert(factor_name, "Integer".to_string()); + model.add_equation(flat::Equation::new( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(var_ref("subSample1.y")), + rhs: Box::new(rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("subSample"), + args: vec![ + var_ref("subSample1.u"), + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Integer"), + args: vec![named_arg("min", int_literal(1))], + is_constructor: true, + span: test_span(), + }, + ], + is_constructor: false, + span: test_span(), + }), + span: test_span(), + }, + test_span(), + flat::EquationOrigin::ComponentEquation { + component: "subSample1".to_string(), + }, + )); + + substitute_known_constants_in_flat(&mut model, &Context::new()).unwrap(); + + let rumoca_core::Expression::Binary { rhs, .. } = &model.equations[0].residual else { + panic!("expected residual assignment"); + }; + let rumoca_core::Expression::FunctionCall { args, .. } = rhs.as_ref() else { + panic!("expected subSample call"); + }; + assert!(matches!( + args[1], + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(5), + .. + } + )); +} + +#[test] +fn substitute_known_constants_reconciles_zero_fill_extent_with_declared_dims() { + let mut model = flat::Model::new(); + let name = rumoca_core::VarName::new("sum.k"); + let stale_fill = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Fill, + args: vec![int_literal(1), int_literal(0)], + span: test_span(), + }; + model.add_variable( + name.clone(), + flat::Variable { + name: name.clone(), + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(stale_fill.clone()), + start: Some(stale_fill.clone()), + dims: vec![2], + is_discrete_type: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.add_equation(flat::Equation::new( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(var_ref("sum.y")), + rhs: Box::new(rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs: Box::new(var_ref("sum.k")), + rhs: Box::new(var_ref("sum.u")), + span: test_span(), + }), + span: test_span(), + }, + test_span(), + flat::EquationOrigin::ComponentEquation { + component: "sum.y".to_string(), + }, + )); + + let mut ctx = Context::new(); + ctx.constant_values + .insert(name.as_str().to_string(), stale_fill.clone()); + substitute_known_constants_in_flat(&mut model, &ctx).unwrap(); + + let var = model.variables.get(&name).expect("variable should remain"); + for expr in [var.binding.as_ref(), var.start.as_ref()] { + let Some(rumoca_core::Expression::BuiltinCall { args, .. }) = expr else { + panic!("expected fill expression"); + }; + assert_eq!(literal_integer_value(&args[1]), Some(2)); + } + let rumoca_core::Expression::Binary { rhs, .. } = &model.equations[0].residual else { + panic!("expected residual assignment"); + }; + let rumoca_core::Expression::Binary { lhs, .. } = rhs.as_ref() else { + panic!("expected substituted multiplication"); + }; + let rumoca_core::Expression::BuiltinCall { args, .. } = lhs.as_ref() else { + panic!("expected substituted fill expression"); + }; + assert_eq!(literal_integer_value(&args[1]), Some(2)); +} + +#[test] +fn substitute_known_constants_reconciles_equation_constructor_extent_with_lhs_dims() { + let mut model = flat::Model::new(); + let lhs_name = rumoca_core::VarName::new("volume.portsData_diameter"); + model.add_variable( + lhs_name.clone(), + flat::Variable { + name: lhs_name, + dims: vec![4], + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.add_equation(flat::Equation::new( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(var_ref("volume.portsData_diameter")), + rhs: Box::new(rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Zeros, + args: vec![int_literal(0)], + span: test_span(), + }), + span: test_span(), + }, + test_span(), + flat::EquationOrigin::ComponentEquation { + component: "volume".to_string(), + }, + )); + + substitute_known_constants_in_flat(&mut model, &Context::new()).unwrap(); + + let rumoca_core::Expression::Binary { rhs, .. } = &model.equations[0].residual else { + panic!("expected residual assignment"); + }; + let rumoca_core::Expression::BuiltinCall { args, .. } = rhs.as_ref() else { + panic!("expected zeros expression"); + }; + assert_eq!(literal_integer_value(&args[0]), Some(4)); +} + +#[test] +fn substitute_known_constants_uses_instance_parameter_bindings_in_component_equations() { + let mut model = flat::Model::new(); + let n_name = rumoca_core::VarName::new("ductOut.flowModel.n"); + model.add_variable( + n_name.clone(), + flat::Variable { + name: n_name, + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(int_literal(4)), + binding_from_modification: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let m_name = rumoca_core::VarName::new("ductOut.flowModel.m"); + model.add_variable( + m_name.clone(), + flat::Variable { + name: m_name, + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(var_ref("n")), + rhs: Box::new(int_literal(1)), + span: test_span(), + }), + binding_from_modification: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.add_equation(flat::Equation::new( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Zeros, + args: vec![var_ref("m")], + span: test_span(), + }), + rhs: Box::new(var_ref("ductOut.flowModel.Ib_flows")), + span: test_span(), + }, + test_span(), + flat::EquationOrigin::ComponentEquation { + component: "ductOut.flowModel".to_string(), + }, + )); + + substitute_known_constants_in_flat(&mut model, &Context::new()).unwrap(); + + let rumoca_core::Expression::Binary { lhs, .. } = &model.equations[0].residual else { + panic!("expected residual subtraction"); + }; + let rumoca_core::Expression::BuiltinCall { args, .. } = lhs.as_ref() else { + panic!("expected zeros expression"); + }; + assert_eq!(eval_test_integer_expr(&args[0]), Some(3)); +} + +#[test] +fn substitute_known_constants_keeps_live_boolean_output_definition() { + let mut model = flat::Model::new(); + let y_name = rumoca_core::VarName::new("source.y"); + model.add_variable( + y_name.clone(), + flat::Variable { + name: y_name, + is_discrete_type: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let k_name = rumoca_core::VarName::new("source.k"); + model.add_variable( + k_name.clone(), + flat::Variable { + name: k_name, + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Boolean(true), + span: test_span(), + }), + is_discrete_type: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.add_equation(flat::Equation::new( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(spanned_var_ref("source.y")), + rhs: Box::new(spanned_var_ref("source.k")), + span: test_span(), + }, + test_span(), + flat::EquationOrigin::ComponentEquation { + component: "source".to_string(), + }, + )); + + // Structural branch evaluation may know the value of a discrete Boolean, + // but that must not erase the runtime-visible output that defines it. + let mut ctx = Context::new(); + ctx.boolean_parameter_values + .insert("source.y".to_string(), true); + ctx.boolean_parameter_values + .insert("source.k".to_string(), true); + + substitute_known_constants_in_flat(&mut model, &ctx).unwrap(); + + let rumoca_core::Expression::Binary { lhs, rhs, .. } = &model.equations[0].residual else { + panic!("expected residual assignment"); + }; + assert!(matches!( + lhs.as_ref(), + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "source.y" + )); + assert!(matches!( + rhs.as_ref(), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Boolean(true), + .. + } + )); +} + +#[test] +fn substitute_known_constants_replaces_live_class_constant_without_flat_binding() { + let mut model = flat::Model::new(); + let pi_name = rumoca_core::VarName::new("sine.pi"); + model.add_variable( + pi_name.clone(), + flat::Variable { + name: pi_name, + variability: rumoca_core::Variability::Constant(rumoca_core::Token::default()), + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let y_name = rumoca_core::VarName::new("sine.y"); + model.add_variable( + y_name.clone(), + flat::Variable { + name: y_name, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.add_equation(flat::Equation::new( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(spanned_var_ref("sine.y")), + rhs: Box::new(spanned_var_ref("sine.pi")), + span: test_span(), + }, + test_span(), + flat::EquationOrigin::ComponentEquation { + component: "sine".to_string(), + }, + )); + + // Class constants can be instantiated without retaining their source + // binding on the Flat declaration. The compiler-owned class-constant + // identity still makes the value structural and safe to substitute. + let mut ctx = Context::new(); + ctx.class_constant_keys.insert("sine.pi".to_string()); + ctx.real_parameter_values + .insert("sine.pi".to_string(), std::f64::consts::PI); + + substitute_known_constants_in_flat(&mut model, &ctx).unwrap(); + + let rumoca_core::Expression::Binary { lhs, rhs, .. } = &model.equations[0].residual else { + panic!("expected residual assignment"); + }; + assert!(matches!( + lhs.as_ref(), + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "sine.y" + )); + assert!(matches!( + rhs.as_ref(), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(value), + .. + } if (*value - std::f64::consts::PI).abs() < f64::EPSILON + )); +} + +#[test] +fn substitute_known_constants_preserves_modified_parameter_over_stale_context_constant() { + let mut model = flat::Model::new(); + let source_name = rumoca_core::VarName::new("source.q_end"); + model.add_variable( + source_name.clone(), + flat::Variable { + name: source_name, + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(2.5), + span: test_span(), + }), + binding_from_modification: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let projected_name = rumoca_core::VarName::new("source.p_q_end"); + model.add_variable( + projected_name.clone(), + flat::Variable { + name: projected_name.clone(), + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(spanned_var_ref("source.q_end")), + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + // The context still carries the declaration default, while the Flat + // variable owns the effective instance modification. + let mut ctx = Context::new(); + ctx.constant_values.insert( + "source.q_end".to_string(), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(1.0), + span: test_span(), + }, + ); + + substitute_known_constants_in_flat(&mut model, &ctx).unwrap(); + + let binding = model + .variables + .get(&projected_name) + .and_then(|var| var.binding.as_ref()) + .expect("projected parameter binding should remain"); + assert!(matches!( + binding, + rumoca_core::Expression::VarRef { name, .. } if name.as_str() == "source.q_end" + )); +} + +#[test] +fn substitute_algorithms_uses_instance_parameter_bindings_in_rhs_and_range() { + let mut model = flat::Model::new(); + for (name, binding) in [("table.y0", 3), ("table.n", 2), ("i", 99)] { + let var_name = rumoca_core::VarName::new(name); + model.add_variable( + var_name.clone(), + flat::Variable { + name: var_name, + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(int_literal(binding)), + binding_from_modification: true, + is_discrete_type: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + let alias_name = rumoca_core::VarName::new("n"); + model.add_variable( + alias_name.clone(), + flat::Variable { + name: alias_name, + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(var_ref("i")), + binding_from_modification: true, + is_discrete_type: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.algorithms.push(flat::Algorithm { + statements: vec![ + simple_assignment(var_ref("table.y0")), + rumoca_core::Statement::For { + indices: vec![rumoca_core::ForIndex { + ident: "i".to_string(), + range: rumoca_core::Expression::Range { + start: Box::new(int_literal(1)), + step: None, + end: Box::new(var_ref("table.n")), + span: test_span(), + }, + }], + equations: vec![ + simple_assignment(var_ref("i")), + simple_assignment(var_ref("n")), + ], + span: test_span(), + }, + simple_assignment(rumoca_core::Expression::ArrayComprehension { + expr: Box::new(var_ref("i")), + indices: vec![rumoca_core::ComprehensionIndex { + name: "i".to_string(), + range: rumoca_core::Expression::Range { + start: Box::new(int_literal(1)), + step: None, + end: Box::new(int_literal(2)), + span: test_span(), + }, + }], + filter: None, + span: test_span(), + }), + ], + outputs: vec![], + span: test_span(), + origin: "algorithm from table".to_string(), + }); + + let mut ctx = Context::new(); + ctx.parameter_values.insert("table.y0".to_string(), 1); + ctx.parameter_values.insert("table.n".to_string(), 1); + ctx.parameter_values.insert("i".to_string(), 99); + + substitute_known_constants_in_flat(&mut model, &ctx).unwrap(); + + let rumoca_core::Statement::Assignment { value, .. } = &model.algorithms[0].statements[0] + else { + panic!("expected initial assignment"); + }; + assert_eq!(literal_integer_value(value), Some(3)); + + let rumoca_core::Statement::For { + indices, equations, .. + } = &model.algorithms[0].statements[1] + else { + panic!("expected table loop"); + }; + let rumoca_core::Expression::Range { end, .. } = &indices[0].range else { + panic!("expected table loop range"); + }; + assert_eq!(literal_integer_value(end), Some(2)); + let rumoca_core::Statement::Assignment { value, .. } = &equations[0] else { + panic!("expected loop-body assignment"); + }; + assert!(matches!( + value, + rumoca_core::Expression::VarRef { name, subscripts, .. } + if name.as_str() == "i" && subscripts.is_empty() + )); + let rumoca_core::Statement::Assignment { value, .. } = &equations[1] else { + panic!("expected alias assignment"); + }; + assert_eq!(literal_integer_value(value), Some(99)); + + let rumoca_core::Statement::Assignment { value, .. } = &model.algorithms[0].statements[2] + else { + panic!("expected comprehension assignment"); + }; + let rumoca_core::Expression::ArrayComprehension { expr, .. } = value else { + panic!("expected array comprehension"); + }; + assert!(matches!( + expr.as_ref(), + rumoca_core::Expression::VarRef { name, subscripts, .. } + if name.as_str() == "i" && subscripts.is_empty() + )); +} + +#[test] +fn substitute_known_constants_uses_declared_dims_for_size_before_stale_start() { + let mut model = flat::Model::new(); + let columns = rumoca_core::VarName::new("lossTable.columns"); + let stale_columns_start = rumoca_core::Expression::Range { + start: Box::new(int_literal(2)), + step: None, + end: Box::new(int_literal(2)), + span: test_span(), + }; + model.add_variable( + columns.clone(), + flat::Variable { + name: columns.clone(), + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + start: Some(stale_columns_start.clone()), + dims: vec![2], + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.add_equation(flat::Equation::new( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(var_ref("y")), + rhs: Box::new(rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + args: vec![var_ref("lossTable.columns"), int_literal(1)], + span: test_span(), + }), + span: test_span(), + }, + test_span(), + flat::EquationOrigin::ComponentEquation { + component: "lossTable".to_string(), + }, + )); + + let mut ctx = Context::new(); + ctx.constant_values + .insert(columns.as_str().to_string(), stale_columns_start); + substitute_known_constants_in_flat(&mut model, &ctx).unwrap(); + + let rumoca_core::Expression::Binary { rhs, .. } = &model.equations[0].residual else { + panic!("expected residual assignment"); + }; + assert_eq!(literal_integer_value(rhs), Some(2)); +} + +fn int_literal(value: i64) -> rumoca_core::Expression { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + span: rumoca_core::Span::DUMMY, + } +} + +fn literal_integer_value(expr: &rumoca_core::Expression) -> Option { + let rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + .. + } = expr + else { + return None; + }; + Some(*value) +} + +fn eval_test_integer_expr(expr: &rumoca_core::Expression) -> Option { + match expr { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + .. + } => Some(*value), + rumoca_core::Expression::Binary { op, lhs, rhs, .. } => { + let lhs = eval_test_integer_expr(lhs)?; + let rhs = eval_test_integer_expr(rhs)?; + rumoca_core::eval_ast_integer_binary(op, lhs, rhs) + } + _ => None, + } +} + +fn assert_ones_extent(expr: &rumoca_core::Expression, expected_extent: i64) { + let rumoca_core::Expression::BuiltinCall { function, args, .. } = expr else { + panic!("expected builtin call, got {expr:?}"); + }; + assert_eq!(*function, rumoca_core::BuiltinFunction::Ones); + assert_eq!(args.len(), 1); + assert_eq!(eval_test_integer_expr(&args[0]), Some(expected_extent)); +} + +fn assert_indexed_var_ref( + expr: &rumoca_core::Expression, + expected_name: &str, + expected_index: i64, +) { + let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = expr + else { + panic!("expected indexed varref `{expected_name}`, got {expr:?}"); + }; + assert_eq!(name.as_str(), expected_name); + assert!(matches!( + subscripts.as_slice(), + [rumoca_core::Subscript::Expr { expr, .. }] + if literal_integer_value(expr) == Some(expected_index) + )); +} + +#[test] +fn variable_binding_substitution_uses_flat_binding_before_stale_context_value() { + let mut model = flat::Model::new(); + add_primitive_variable(&mut model, "pipe.n"); + add_primitive_variable(&mut model, "pipe.flowModel.n"); + add_primitive_variable(&mut model, "pipe.flowModel.Res_turbulent_internal"); + model + .variables + .get_mut(&rumoca_core::VarName::new("pipe.n")) + .expect("variable should exist") + .binding = Some(int_literal(2)); + model + .variables + .get_mut(&rumoca_core::VarName::new("pipe.flowModel.n")) + .expect("variable should exist") + .binding = Some(rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: Box::new(int_literal(3)), + rhs: Box::new(int_literal(1)), + span: rumoca_core::Span::DUMMY, + }); + model + .variables + .get_mut(&rumoca_core::VarName::new( + "pipe.flowModel.Res_turbulent_internal", + )) + .expect("variable should exist") + .binding = Some(rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Ones, + args: vec![rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(var_ref("pipe.flowModel.n")), + rhs: Box::new(int_literal(1)), + span: rumoca_core::Span::DUMMY, + }], + span: rumoca_core::Span::DUMMY, + }); + + let mut ctx = Context::new(); + ctx.parameter_values + .insert("pipe.flowModel.n".to_string(), 2); + + substitute_known_constants_in_flat(&mut model, &ctx).unwrap(); + + let binding = model + .variables + .get(&rumoca_core::VarName::new( + "pipe.flowModel.Res_turbulent_internal", + )) + .expect("variable should exist") + .binding + .as_ref() + .expect("binding should remain"); + assert_ones_extent(binding, 3); +} + +fn assert_symbolically_indexed_var_ref( + expr: &rumoca_core::Expression, + expected_name: &str, + expected_index_name: &str, +) { + let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = expr + else { + panic!("expected indexed varref `{expected_name}`, got {expr:?}"); + }; + assert_eq!(name.as_str(), expected_name); + assert!(matches!( + subscripts.as_slice(), + [rumoca_core::Subscript::Expr { expr, .. }] + if matches!( + expr.as_ref(), + rumoca_core::Expression::VarRef { name, subscripts, .. } + if name.as_str() == expected_index_name && subscripts.is_empty() + ) + )); +} + +fn string_literal(value: &str) -> rumoca_core::Expression { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String(value.to_string()), + span: rumoca_core::Span::DUMMY, + } +} + +fn reference_x_fill_expr() -> rumoca_core::Expression { + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Fill, + args: vec![ + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Div, + lhs: Box::new(int_literal(1)), + rhs: Box::new(var_ref("nS")), + span: rumoca_core::Span::DUMMY, + }, + var_ref("nS"), + ], + span: rumoca_core::Span::DUMMY, + } +} + +fn add_primitive_variable(model: &mut flat::Model, name: &str) { + model.add_variable( + rumoca_core::VarName::new(name), + flat::Variable { + name: rumoca_core::VarName::new(name), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); +} + +#[test] +fn substitutes_qualified_real_class_constant_inside_builtin_array_binding() { + let mut model = flat::Model::new(); + let name = rumoca_core::VarName::new("pid.D.T"); + model.add_variable( + name.clone(), + flat::Variable { + name: name.clone(), + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + binding: Some(rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Max, + args: vec![rumoca_core::Expression::Array { + elements: vec![ + var_ref("pid.Nd"), + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs: Box::new(int_literal(100)), + rhs: Box::new(spanned_var_ref("Modelica.Constants.eps")), + span: test_span(), + }, + ], + is_matrix: false, + span: test_span(), + }], + span: test_span(), + }), + binding_from_modification: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let mut ctx = Context::new(); + ctx.class_constant_keys + .insert("Modelica.Constants.eps".to_string()); + ctx.real_parameter_values + .insert("Modelica.Constants.eps".to_string(), f64::EPSILON); + + substitute_known_constants_in_flat(&mut model, &ctx).unwrap(); + + assert!( + !expr_contains_var_ref( + model.variables[&name].binding.as_ref().unwrap(), + "Modelica.Constants.eps" + ), + "Modelica.Constants.eps should be folded inside runtime parameter modifier bindings: {:?}", + model.variables[&name].binding + ); +} + +#[test] +fn substitutes_well_known_real_class_constant_in_equations() { + let mut model = flat::Model::new(); + add_primitive_variable(&mut model, "x"); + model.add_equation(flat::Equation::new( + spanned_var_ref("ModelicaServices.Machine.eps"), + test_span(), + flat::EquationOrigin::ComponentEquation { + component: String::new(), + }, + )); + let mut ctx = Context::new(); + ctx.class_constant_keys + .insert("ModelicaServices.Machine.eps".to_string()); + ctx.real_parameter_values + .insert("ModelicaServices.Machine.eps".to_string(), f64::EPSILON); + + substitute_known_constants_in_flat(&mut model, &ctx).unwrap(); + + assert!( + matches!( + model.equations[0].residual, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(_), + .. + } + ), + "class constant should be substituted in equation residuals: {:?}", + model.equations[0].residual + ); +} + +#[test] +fn substitute_known_constants_preserves_named_arg_marker_for_record_constructor_value() { + let mut model = flat::Model::new(); + model.add_equation(flat::Equation::new( + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.f"), + args: vec![ + named_arg("x", var_ref("live_x")), + named_arg("per", var_ref("pCur1")), + ], + is_constructor: false, + span: rumoca_core::Span::DUMMY, + }, + rumoca_core::Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "test".to_string(), + }, + )); + let mut ctx = Context::new(); + ctx.constant_values.insert( + "pCur1".to_string(), + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Buildings.Fluid.Movers.Data.Generic"), + args: vec![var_ref("pCur1.V_flow"), var_ref("pCur1.dp")], + is_constructor: true, + span: rumoca_core::Span::DUMMY, + }, + ); + + substitute_known_constants_in_flat(&mut model, &ctx).unwrap(); + + let rumoca_core::Expression::FunctionCall { args, .. } = &model.equations[0].residual else { + panic!("expected function call"); + }; + assert_eq!(args.len(), 2); + let rumoca_core::Expression::FunctionCall { + name, + args: per_args, + is_constructor: true, + .. + } = &args[1] + else { + panic!("expected named per argument marker"); + }; + assert_eq!(name.as_str(), "__rumoca_named_arg__.per"); + assert!(matches!( + per_args.as_slice(), + [rumoca_core::Expression::FunctionCall { name, is_constructor: true, .. }] + if name.as_str() == "Buildings.Fluid.Movers.Data.Generic" + )); +} + +#[test] +fn substitute_known_constants_preserves_path_like_string_literal() { + let mut ctx = Context::new(); + ctx.string_parameter_values.insert( + "zone.spawnExe".to_string(), + "spawn-0.4.3-7048a72798".to_string(), + ); + + let substituted = substitute_known_constants_expr( + var_ref("zone.spawnExe"), + &ctx, + &rustc_hash::FxHashSet::default(), + &std::collections::HashSet::new(), + "", + ) + .unwrap(); + + assert_eq!(substituted, string_literal("spawn-0.4.3-7048a72798")); +} + +#[test] +fn substitute_known_constants_recovers_path_like_string_variable_start() { + let mut model = flat::Model::new(); + let name = rumoca_core::VarName::new("zone.spawnExe"); + model.add_variable( + name.clone(), + flat::Variable { + name: name.clone(), + type_id: rumoca_core::TypeId(3), + start: Some(rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FieldAccess { + base: Box::new(var_ref("zone")), + field: "spawn-0".to_string(), + span: test_span(), + }), + field: "4".to_string(), + span: test_span(), + }), + field: "3-7048a72798".to_string(), + span: test_span(), + }), + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model + .variable_type_names + .insert(name.clone(), "String".to_string()); + + substitute_known_constants_in_flat(&mut model, &Context::new()).unwrap(); + + assert_eq!( + model.variables.get(&name).and_then(|var| var.start.clone()), + Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("spawn-0.4.3-7048a72798".to_string()), + span: test_span(), + }) + ); +} + +#[test] +fn collapse_index_refs_preserves_structured_distinct_record_endpoints() { + let mut model = flat::Model::new(); + for (index, name) in [ + "device.i.re", + "device.i.im", + "device.pin_p.i.re", + "device.pin_p.i.im", + "device.pin_n.i.re", + "device.pin_n.i.im", + ] + .into_iter() + .enumerate() + { + model.add_variable( + rumoca_core::VarName::new(name), + flat::Variable { + name: rumoca_core::VarName::new(name), + component_ref: Some(rumoca_core::ComponentReference::from_flat_segments( + name, + test_span(), + Some(rumoca_core::DefId::new(950 + index as u32)), + )), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + + let source_def_id = rumoca_core::DefId::new(990); + let pin_p = var_ref_with_target_def_id("device.pin_p.i", source_def_id); + let pin_n = var_ref_with_target_def_id("device.pin_n.i", source_def_id); + let rumoca_core::Expression::VarRef { + name: expected_pin_p, + .. + } = &pin_p + else { + unreachable!("test helper returns a var ref"); + }; + let rumoca_core::Expression::VarRef { + name: expected_pin_n, + .. + } = &pin_n + else { + unreachable!("test helper returns a var ref"); + }; + let expected_pin_p = expected_pin_p.clone(); + let expected_pin_n = expected_pin_n.clone(); + model.add_equation(flat::Equation::new( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: Box::new(pin_p), + rhs: Box::new(pin_n), + span: test_span(), + }, + test_span(), + flat::EquationOrigin::ComponentEquation { + component: "device".to_string(), + }, + )); + + collapse_index_refs_to_known_varrefs(&mut model); + + let rumoca_core::Expression::Binary { lhs, rhs, .. } = &model.equations[0].residual else { + panic!("expected connector balance expression"); + }; + let rumoca_core::Expression::VarRef { + name: actual_pin_p, .. + } = lhs.as_ref() + else { + panic!("expected pin_p aggregate reference"); + }; + let rumoca_core::Expression::VarRef { + name: actual_pin_n, .. + } = rhs.as_ref() + else { + panic!("expected pin_n aggregate reference"); + }; + assert_eq!( + [ + (actual_pin_p.as_str(), actual_pin_p.target_def_id()), + (actual_pin_n.as_str(), actual_pin_n.target_def_id()), + ], + [ + ("device.pin_p.i", Some(source_def_id)), + ("device.pin_n.i", Some(source_def_id)), + ] + ); + assert_eq!(actual_pin_p, &expected_pin_p); + assert_eq!(actual_pin_n, &expected_pin_n); +} + +#[test] +fn collapse_index_refs_collapses_indexed_field_access_to_known_var() { + let mut model = flat::Model::new(); + model.add_variable( + rumoca_core::VarName::new("port_a[1].Q_flow"), + flat::Variable { + name: rumoca_core::VarName::new("port_a[1].Q_flow"), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.add_equation(flat::Equation::new( + rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::Index { + base: Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("port_a"), + subscripts: vec![], + span: rumoca_core::Span::DUMMY, + }), + subscripts: vec![rumoca_core::Subscript::generated_index( + 1, + rumoca_core::Span::DUMMY, + )], + span: rumoca_core::Span::DUMMY, + }), + field: "Q_flow".to_string(), + span: rumoca_core::Span::DUMMY, + }, + rumoca_core::Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "test".to_string(), + }, + )); + + collapse_index_refs_to_known_varrefs(&mut model); + + assert!(matches!( + &model.equations[0].residual, + rumoca_core::Expression::VarRef { name, subscripts, .. } + if name.as_str() == "port_a[1].Q_flow" && subscripts.is_empty() + )); +} + +#[test] +fn collapse_penultimate_preserves_proven_nested_indexed_component_boundary() { + let mut model = flat::Model::new(); + for name in [ + "stack.cell[1,1].local_reset", + "stack.cell[1,1].cell.local_reset", + ] { + let var_name = rumoca_core::VarName::new(name); + model.add_variable( + var_name.clone(), + flat::Variable { + name: var_name.clone(), + component_ref: rumoca_core::component_reference_from_flat_name( + &var_name, + test_span(), + ), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + + let known = KnownFlatVars::build(&model); + assert!( + collapse_penultimate_field_to_known_var( + "stack.cell[1,1].cell.local_reset", + test_span(), + &known, + ) + .is_none() + ); +} + +#[test] +fn collapse_penultimate_prefers_direct_indexed_candidate_when_both_spellings_exist() { + let mut model = flat::Model::new(); + for name in ["states[1].h", "states.h[1]"] { + let var_name = rumoca_core::VarName::new(name); + model.add_variable( + var_name.clone(), + flat::Variable { + name: var_name.clone(), + component_ref: rumoca_core::component_reference_from_flat_name( + &var_name, + test_span(), + ), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + + let known = KnownFlatVars::build(&model); + let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = collapse_penultimate_field_to_known_var("states[1].phase.h", test_span(), &known) + .expect("direct indexed candidate should win") + else { + panic!("expected collapsed VarRef"); + }; + assert_eq!(name.as_str(), "states[1].h"); + assert!(name.has_structure()); + assert!(subscripts.is_empty()); +} + +#[test] +fn collapse_penultimate_supports_symbolic_and_range_indexed_direct_candidates() { + for (path, expected) in [ + ("states[i].phase.h", "states[i].h"), + ("states[1:2].phase.h", "states[1:2].h"), + ] { + let mut model = flat::Model::new(); + add_primitive_variable(&mut model, expected); + let known = KnownFlatVars::build(&model); + let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = collapse_penultimate_field_to_known_var(path, test_span(), &known) + .expect("indexed direct candidate should collapse") + else { + panic!("expected collapsed VarRef"); + }; + assert_eq!(name.as_str(), expected); + assert!(subscripts.is_empty()); + } +} + +#[test] +fn collapse_penultimate_uses_alternate_projection_and_preserves_reference_structure() { + let mut model = flat::Model::new(); + let var_name = rumoca_core::VarName::new("states.h[1]"); + model.add_variable( + var_name.clone(), + flat::Variable { + name: var_name.clone(), + component_ref: rumoca_core::component_reference_from_flat_name(&var_name, test_span()), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + let known = KnownFlatVars::build(&model); + let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = collapse_penultimate_field_to_known_var("states[1].phase.h", test_span(), &known) + .expect("alternate array-field projection should collapse") + else { + panic!("expected collapsed VarRef"); + }; + assert_eq!(name.as_str(), "states.h[1]"); + assert!(name.has_structure()); + assert!(subscripts.is_empty()); +} + +#[test] +fn collapse_index_refs_preserves_array_member_aggregate_projection() { + let mut model = flat::Model::new(); + for name in [ + "vehicle.omega", + "vehicle.motor[1].omega", + "vehicle.motor[2].omega", + ] { + model.add_variable( + rumoca_core::VarName::new(name), + flat::Variable { + name: rumoca_core::VarName::new(name), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + model.add_equation(flat::Equation::new( + spanned_var_ref("vehicle.motor.omega"), + test_span(), + flat::EquationOrigin::ComponentEquation { + component: "vehicle".to_string(), + }, + )); + + collapse_index_refs_to_known_varrefs(&mut model); + + assert!(matches!( + &model.equations[0].residual, + rumoca_core::Expression::VarRef { name, subscripts, .. } + if name.as_str() == "vehicle.motor.omega" && subscripts.is_empty() + )); +} + +#[test] +fn collapse_index_refs_collapses_repeated_record_field_tail_to_known_var() { + let mut model = flat::Model::new(); + model.add_variable( + rumoca_core::VarName::new("pipe.flowModel.states[1].phase"), + flat::Variable { + name: rumoca_core::VarName::new("pipe.flowModel.states[1].phase"), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let states = rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("pipe.flowModel.states"), + subscripts: vec![rumoca_core::Subscript::generated_index( + 1, + rumoca_core::Span::DUMMY, + )], + span: rumoca_core::Span::DUMMY, + }; + let phase = rumoca_core::Expression::FieldAccess { + base: Box::new(states), + field: "phase".to_string(), + span: rumoca_core::Span::DUMMY, + }; + let phase_phase = rumoca_core::Expression::FieldAccess { + base: Box::new(phase), + field: "phase".to_string(), + span: rumoca_core::Span::DUMMY, + }; + model.add_equation(flat::Equation::new( + rumoca_core::Expression::FieldAccess { + base: Box::new(phase_phase), + field: "phase".to_string(), + span: rumoca_core::Span::DUMMY, + }, + rumoca_core::Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "test".to_string(), + }, + )); + + collapse_index_refs_to_known_varrefs(&mut model); + collapse_index_refs_to_known_varrefs(&mut model); + + assert!(matches!( + &model.equations[0].residual, + rumoca_core::Expression::VarRef { name, subscripts, .. } + if name.as_str() == "pipe.flowModel.states[1].phase" && subscripts.is_empty() + )); +} + +#[test] +fn collapse_index_refs_collapses_rendered_repeated_record_field_tail_to_known_var() { + let mut model = flat::Model::new(); + model.add_variable( + rumoca_core::VarName::new("pipe.flowModel.states[1].phase"), + flat::Variable { + name: rumoca_core::VarName::new("pipe.flowModel.states[1].phase"), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.add_equation(flat::Equation::new( + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("pipe.flowModel.states[1].phase.phase.phase"), + subscripts: Vec::new(), + span: rumoca_core::Span::DUMMY, + }, + rumoca_core::Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "test".to_string(), + }, + )); + + collapse_index_refs_to_known_varrefs(&mut model); -fn reference_x_fill_expr() -> rumoca_core::Expression { - rumoca_core::Expression::BuiltinCall { - function: rumoca_core::BuiltinFunction::Fill, - args: vec![ - rumoca_core::Expression::Binary { - op: rumoca_core::OpBinary::Div, - lhs: Box::new(int_literal(1)), - rhs: Box::new(var_ref("nS")), - span: rumoca_core::Span::DUMMY, - }, - var_ref("nS"), - ], - span: rumoca_core::Span::DUMMY, - } + assert!(matches!( + &model.equations[0].residual, + rumoca_core::Expression::VarRef { name, subscripts, .. } + if name.as_str() == "pipe.flowModel.states[1].phase" && subscripts.is_empty() + )); } -fn add_primitive_variable(model: &mut flat::Model, name: &str) { +#[test] +fn collapse_index_refs_collapses_repeated_record_field_tail_to_array_field_var() { + let mut model = flat::Model::new(); model.add_variable( - rumoca_core::VarName::new(name), + rumoca_core::VarName::new("pipe.flowModel.states.phase[1]"), flat::Variable { - name: rumoca_core::VarName::new(name), + name: rumoca_core::VarName::new("pipe.flowModel.states.phase[1]"), is_primitive: true, ..flat::Variable::empty_with_span(test_span()) }, ); + model.add_equation(flat::Equation::new( + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("pipe.flowModel.states[1].phase.phase"), + subscripts: Vec::new(), + span: rumoca_core::Span::DUMMY, + }, + rumoca_core::Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "test".to_string(), + }, + )); + + collapse_index_refs_to_known_varrefs(&mut model); + + assert!(matches!( + &model.equations[0].residual, + rumoca_core::Expression::VarRef { name, subscripts, .. } + if name.as_str() == "pipe.flowModel.states.phase[1]" && subscripts.is_empty() + )); } #[test] -fn collapse_index_refs_collapses_indexed_field_access_to_known_var() { +fn collapse_index_refs_collapses_overexpanded_record_sibling_field_to_array_field_var() { let mut model = flat::Model::new(); model.add_variable( - rumoca_core::VarName::new("port_a[1].Q_flow"), + rumoca_core::VarName::new("pipe.flowModel.states.h[1]"), flat::Variable { - name: rumoca_core::VarName::new("port_a[1].Q_flow"), + name: rumoca_core::VarName::new("pipe.flowModel.states.h[1]"), is_primitive: true, ..flat::Variable::empty_with_span(test_span()) }, ); model.add_equation(flat::Equation::new( - rumoca_core::Expression::FieldAccess { - base: Box::new(rumoca_core::Expression::Index { - base: Box::new(rumoca_core::Expression::VarRef { - name: rumoca_core::Reference::new("port_a"), - subscripts: vec![], - span: rumoca_core::Span::DUMMY, - }), - subscripts: vec![rumoca_core::Subscript::generated_index( - 1, - rumoca_core::Span::DUMMY, - )], - span: rumoca_core::Span::DUMMY, - }), - field: "Q_flow".to_string(), + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("pipe.flowModel.states[1].phase.h"), + subscripts: Vec::new(), span: rumoca_core::Span::DUMMY, }, rumoca_core::Span::DUMMY, @@ -118,8 +1641,50 @@ fn collapse_index_refs_collapses_indexed_field_access_to_known_var() { assert!(matches!( &model.equations[0].residual, rumoca_core::Expression::VarRef { name, subscripts, .. } - if name.as_str() == "port_a[1].Q_flow" && subscripts.is_empty() + if name.as_str() == "pipe.flowModel.states.h[1]" && subscripts.is_empty() + )); +} + +#[test] +fn collapse_index_refs_collapses_repeated_record_field_tail_to_array_base_ref() { + let mut model = flat::Model::new(); + let leaf = "pipe.flowModel.states.phase[1]"; + model.add_variable( + rumoca_core::VarName::new(leaf), + flat::Variable { + name: rumoca_core::VarName::new(leaf), + component_ref: Some(rumoca_core::ComponentReference::from_flat_segments( + leaf, + test_span(), + None, + )), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.add_equation(flat::Equation::new( + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("pipe.flowModel.states.phase.phase"), + subscripts: Vec::new(), + span: rumoca_core::Span::DUMMY, + }, + rumoca_core::Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "test".to_string(), + }, )); + + collapse_index_refs_to_known_varrefs(&mut model); + + let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = &model.equations[0].residual + else { + panic!("expected collapsed VarRef"); + }; + assert_eq!(name.as_str(), "pipe.flowModel.states.phase"); + assert!(name.has_structure()); + assert!(subscripts.is_empty()); } #[test] @@ -154,6 +1719,47 @@ fn collapse_index_refs_collapses_indexed_var_ref_to_known_scalar_var() { )); } +#[test] +fn recover_indexed_lhs_dimensions_does_not_expand_known_symbolic_dimension() { + let mut model = flat::Model::new(); + let y_name = rumoca_core::VarName::new("aD_Converter.y"); + model.add_variable( + y_name.clone(), + flat::Variable { + name: y_name.clone(), + dims: vec![7], + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.add_equation(flat::Equation::new( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("aD_Converter.y"), + subscripts: vec![rumoca_core::Subscript::generated_index( + 8, + rumoca_core::Span::DUMMY, + )], + span: rumoca_core::Span::DUMMY, + }), + rhs: Box::new(int_literal(0)), + span: rumoca_core::Span::DUMMY, + }, + rumoca_core::Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "aD_Converter".to_string(), + }, + )); + + recover_indexed_lhs_dimensions(&mut model); + + assert_eq!( + model.variables.get(&y_name).expect("y variable").dims, + vec![7] + ); +} + #[test] fn substitutes_known_constants_inside_function_defaults_and_body() { let mut model = flat::Model::new(); @@ -449,6 +2055,78 @@ fn substitutes_record_array_field_projection_from_flat_var_ref() { )); } +#[test] +fn substitutes_positional_record_constructor_field_projection() { + let mut model = flat::Model::new(); + let mut function = rumoca_core::Function::new("Pkg.f", Span::DUMMY); + function.add_input( + rumoca_core::FunctionParam::new("u", "Real", test_span()) + .with_default(var_ref("GasData.Air.R")), + ); + model.add_function(function); + + let mut constructor = rumoca_core::Function::new("DataRecord", Span::DUMMY); + constructor.is_constructor = true; + constructor.add_input(rumoca_core::FunctionParam::new( + "name", + "String", + test_span(), + )); + constructor.add_input(rumoca_core::FunctionParam::new("R", "Real", test_span())); + + let mut ctx = Context::new(); + ctx.functions + .insert("DataRecord".to_string(), constructor.clone()); + let record = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("DataRecord"), + args: vec![ + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("Air".to_string()), + span: Span::DUMMY, + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(287.0), + span: Span::DUMMY, + }, + ], + is_constructor: true, + span: Span::DUMMY, + }; + ctx.constant_values + .insert("GasData.Air".to_string(), record.clone()); + ctx.constant_values + .insert("GasData.Air.R".to_string(), record); + + assert!( + matches!( + resolve_projected_constant_path("GasData.Air.R", Span::DUMMY, &ctx), + Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(value), + .. + }) if (value - 287.0).abs() < f64::EPSILON + ), + "projected constant path should select positional constructor field" + ); + + substitute_known_constants_in_flat(&mut model, &ctx).unwrap(); + + let function = model + .functions + .get(&rumoca_core::VarName::new("Pkg.f")) + .expect("function should exist"); + let actual = &function.inputs[0].default; + assert!( + matches!( + actual, + Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(value), + .. + }) if (*value - 287.0).abs() < f64::EPSILON + ), + "expected projected R field, got {actual:?}" + ); +} + #[test] fn does_not_substitute_function_local_names() { let mut model = flat::Model::new(); @@ -622,6 +2300,159 @@ fn substitutes_inline_multi_indexed_constant_varref_names() { } } +#[test] +fn substitutes_indexed_array_comprehension_parameter_binding_as_scalar_element() { + let span = test_span(); + let mut ctx = Context::new(); + ctx.constant_values.insert( + "fluidVolumes".to_string(), + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs: Box::new(rumoca_core::Expression::ArrayComprehension { + expr: Box::new(rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs: Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("crossAreas"), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(var_ref("i")), + span, + }], + span, + }), + rhs: Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("lengths"), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(var_ref("i")), + span, + }], + span, + }), + span, + }), + indices: vec![rumoca_core::ComprehensionIndex { + name: "i".to_string(), + range: rumoca_core::Expression::Range { + start: Box::new(int_literal(1)), + step: None, + end: Box::new(int_literal(2)), + span, + }, + }], + filter: None, + span, + }), + rhs: Box::new(var_ref("nParallel")), + span, + }, + ); + ctx.constant_values + .insert("nParallel".to_string(), int_literal(3)); + + let mut live_vars = rustc_hash::FxHashSet::default(); + live_vars.insert("crossAreas".to_string()); + live_vars.insert("lengths".to_string()); + let substituted = substitute_known_constants_expr( + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("fluidVolumes"), + subscripts: vec![rumoca_core::Subscript::generated_index(2, span)], + span, + }, + &ctx, + &live_vars, + &HashSet::new(), + "", + ) + .expect("indexed comprehension parameter should substitute"); + + let rumoca_core::Expression::Binary { lhs, rhs, .. } = substituted else { + panic!("expected scalar product, got {substituted:?}"); + }; + assert!(matches!( + rhs.as_ref(), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(3), + .. + } + )); + let rumoca_core::Expression::Binary { + lhs: cross_area, + rhs: length, + .. + } = lhs.as_ref() + else { + panic!("expected selected comprehension body, got {lhs:?}"); + }; + assert_indexed_var_ref(cross_area, "crossAreas", 2); + assert_indexed_var_ref(length, "lengths", 2); +} + +#[test] +fn substitutes_symbolically_indexed_array_comprehension_parameter_binding() { + let span = test_span(); + let mut ctx = Context::new(); + ctx.constant_values.insert( + "fluidVolumes".to_string(), + rumoca_core::Expression::ArrayComprehension { + expr: Box::new(rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs: Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("crossAreas"), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(var_ref("j")), + span, + }], + span, + }), + rhs: Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("lengths"), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(var_ref("j")), + span, + }], + span, + }), + span, + }), + indices: vec![rumoca_core::ComprehensionIndex { + name: "j".to_string(), + range: rumoca_core::Expression::Range { + start: Box::new(int_literal(1)), + step: None, + end: Box::new(int_literal(2)), + span, + }, + }], + filter: None, + span, + }, + ); + + let mut live_vars = rustc_hash::FxHashSet::default(); + live_vars.insert("crossAreas".to_string()); + live_vars.insert("lengths".to_string()); + let substituted = substitute_known_constants_expr( + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("fluidVolumes"), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(var_ref("i")), + span, + }], + span, + }, + &ctx, + &live_vars, + &HashSet::new(), + "", + ) + .expect("symbolically indexed comprehension parameter should substitute"); + + let rumoca_core::Expression::Binary { lhs, rhs, .. } = substituted else { + panic!("expected selected comprehension body, got {substituted:?}"); + }; + assert_symbolically_indexed_var_ref(&lhs, "crossAreas", "i"); + assert_symbolically_indexed_var_ref(&rhs, "lengths", "i"); +} + #[test] fn rejects_unspanned_inline_indexed_constant_varref_names() { let mut model = flat::Model::new(); @@ -819,6 +2650,45 @@ fn substitutes_fully_qualified_constant_alias_in_declaration_scope() { assert!(!expr_contains_var_ref(start, "nS")); } +#[test] +fn scoped_relative_alias_keeps_live_flat_variable_before_constant_expansion() { + let mut model = flat::Model::new(); + add_primitive_variable(&mut model, "jointRRP.e_ia"); + add_primitive_variable(&mut model, "jointRRP.jointUSP.e2_ia"); + add_primitive_variable(&mut model, "jointRRP.jointUSP.rod1.e2_ia"); + model + .variables + .get_mut(&rumoca_core::VarName::new("jointRRP.e_ia")) + .expect("variable should exist") + .binding = Some(var_ref("jointUSP.e2_ia")); + + let mut ctx = Context::new(); + ctx.constant_values + .insert("jointRRP.jointUSP.e2_ia".to_string(), var_ref("rod1.e2_ia")); + ctx.constant_values.insert( + "jointRRP.rod1.e2_ia".to_string(), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(1.0), + span: rumoca_core::Span::DUMMY, + }, + ); + + substitute_known_constants_in_flat(&mut model, &ctx).unwrap(); + + let binding = model + .variables + .get(&rumoca_core::VarName::new("jointRRP.e_ia")) + .expect("variable should exist") + .binding + .as_ref() + .expect("binding should remain"); + assert!(matches!( + binding, + rumoca_core::Expression::VarRef { name, .. } + if name.as_str() == "jointRRP.jointUSP.e2_ia" + )); +} + #[test] fn does_not_substitute_array_shaped_scalar_parameter_ref() { let mut model = flat::Model::new(); @@ -927,6 +2797,37 @@ fn materializes_referenced_zero_sized_array_declaration() { assert!(var.is_primitive); } +#[test] +fn materializes_unspanned_zero_sized_array_reference_from_dimension_provenance() { + let mut model = flat::Model::new(); + model.equations.push(flat::Equation::new( + var_ref("Modelica.Media.Water.WaterIF97_base.C_default"), + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "PumpingSystem".to_string(), + }, + )); + + let mut ctx = Context::new(); + let name = "Modelica.Media.Water.WaterIF97_base.C_default"; + ctx.array_dimensions.insert(name.to_string(), vec![0]); + ctx.array_dimension_spans + .insert(name.to_string(), test_span()); + + substitute_known_constants_in_flat(&mut model, &ctx).unwrap(); + + let var = model + .variables + .get(&rumoca_core::VarName::new(name)) + .expect("zero-sized referenced array should have a Flat declaration"); + assert_eq!(var.dims, vec![0]); + assert_eq!(var.source_span, test_span()); + assert!( + var.component_ref.is_some(), + "materialized zero-sized array variable should carry structured metadata" + ); +} + #[test] fn substitutes_field_access_on_zero_arg_constructor_constants() { let mut model = flat::Model::new(); diff --git a/crates/rumoca-phase-flatten/src/postprocess_def_id.rs b/crates/rumoca-phase-flatten/src/postprocess_def_id.rs index 5fe3a02de..a2e3b6cf0 100644 --- a/crates/rumoca-phase-flatten/src/postprocess_def_id.rs +++ b/crates/rumoca-phase-flatten/src/postprocess_def_id.rs @@ -109,6 +109,7 @@ pub(crate) fn canonicalize_varrefs_via_instantiated_def_ids(flat: &mut flat::Mod struct DefIdVarRefIndex { by_def_id: HashMap>, by_leaf: HashMap>, + aggregate_projection_refs: HashSet, } #[derive(Clone)] @@ -122,17 +123,46 @@ impl DefIdVarRefIndex { fn build(flat: &flat::Model) -> Self { let mut by_def_id: HashMap> = HashMap::new(); let mut by_leaf: HashMap> = HashMap::new(); + let mut aggregate_projection_counts: HashMap = HashMap::new(); for (name, var) in &flat.variables { let indexed = indexed_var_ref(name, var); by_leaf .entry(indexed_var_leaf(name, var).to_string()) .or_default() .push(indexed.clone()); + if let Some(projection) = aggregate_projection_ref(name.as_str()) { + *aggregate_projection_counts.entry(projection).or_insert(0) += 1; + } if let Some(def_id) = var.component_ref.as_ref().and_then(|comp| comp.def_id) { - by_def_id.entry(def_id).or_default().push(indexed); + push_indexed_candidate(&mut by_def_id, def_id, indexed.clone()); + if let Some(ancestry) = flat.symbol_ancestry.get(&def_id) { + for ancestor_def_id in ancestry { + push_indexed_candidate(&mut by_def_id, *ancestor_def_id, indexed.clone()); + } + } } } - Self { by_def_id, by_leaf } + let known_variable_names = flat + .variables + .keys() + .map(|name| name.as_str().to_string()) + .collect::>(); + let aggregate_projection_refs = aggregate_projection_counts + .into_iter() + .filter_map(|(projection, count)| { + (count > 1 + && aggregate_projection_needs_alias_protection( + &projection, + &known_variable_names, + )) + .then_some(projection) + }) + .collect(); + Self { + by_def_id, + by_leaf, + aggregate_projection_refs, + } } fn is_empty(&self) -> bool { @@ -146,6 +176,19 @@ impl DefIdVarRefIndex { known_variables: &HashSet, ) -> Option { let raw = name.as_str(); + if let Some(component_ref) = name.component_ref() { + let structured = rumoca_core::ComponentPath::from_component_reference(component_ref) + .to_flat_string(); + if known_variables.contains(structured.as_str()) { + return (structured.as_str() != raw).then_some(structured); + } + if self.aggregate_projection_refs.contains(structured.as_str()) { + return (structured.as_str() != raw).then_some(structured); + } + } + if self.aggregate_projection_refs.contains(raw) { + return None; + } if known_variables.contains(raw) { return None; } @@ -153,19 +196,116 @@ impl DefIdVarRefIndex { if let Some(def_id) = name.target_def_id() && let Some(candidates) = self.by_def_id.get(&def_id) { - return resolve_best_owner_scoped_candidate(name, raw_leaf, candidates, owner); + let mode = if is_class_qualified_reference(name) { + OwnerScopeMode::DescendantOrEnclosing + } else { + OwnerScopeMode::DescendantOnly + }; + return resolve_best_owner_scoped_candidate(name, raw_leaf, candidates, owner, mode) + .or_else(|| { + (!is_class_qualified_reference(name)).then(|| { + resolve_best_owner_scoped_candidate( + name, + raw_leaf, + candidates, + owner, + OwnerScopeMode::EnclosingOnly, + ) + })? + }); } if name.target_def_id().is_some() { return None; } if is_class_qualified_reference(name) { let candidates = self.by_leaf.get(raw_leaf)?; - return resolve_best_owner_scoped_candidate(name, raw_leaf, candidates, owner); + return resolve_best_owner_scoped_candidate( + name, + raw_leaf, + candidates, + owner, + OwnerScopeMode::DescendantOrEnclosing, + ); } None } } +pub(super) fn aggregate_projection_ref(name: &str) -> Option { + let mut projection = String::with_capacity(name.len()); + let mut depth = 0i32; + let mut group_start = None; + let mut indices = 0usize; + + for (idx, ch) in name.char_indices() { + match ch { + '[' => { + if depth == 0 { + group_start = Some(idx + 1); + } + depth += 1; + } + ']' => { + if depth == 0 { + return None; + } + depth -= 1; + if depth == 0 { + name[group_start?..idx].trim().parse::().ok()?; + indices += 1; + group_start = None; + } + } + _ if depth == 0 => projection.push(ch), + _ => {} + } + } + + (depth == 0 && indices == 1 && !projection.is_empty() && projection != name) + .then_some(projection) +} + +pub(super) fn aggregate_projection_needs_alias_protection( + projection: &str, + known_variables: &HashSet, +) -> bool { + penultimate_field_collapse_path(projection) + .is_some_and(|alias| known_variables.contains(alias.as_str())) +} + +fn penultimate_field_collapse_path(path: &str) -> Option { + let (prefix, leaf) = rendered_path_last_segment(path)?; + let (base, _) = rendered_path_last_segment(prefix)?; + Some(format!("{base}.{leaf}")) +} + +fn rendered_path_last_segment(path: &str) -> Option<(&str, &str)> { + let mut bracket_depth = 0usize; + for (idx, byte) in path.bytes().enumerate().rev() { + match byte { + b']' => bracket_depth += 1, + b'[' => bracket_depth = bracket_depth.saturating_sub(1), + b'.' if bracket_depth == 0 => return Some((&path[..idx], &path[idx + 1..])), + _ => {} + } + } + None +} + +fn push_indexed_candidate( + by_def_id: &mut HashMap>, + def_id: rumoca_core::DefId, + indexed: IndexedVarRef, +) { + let candidates = by_def_id.entry(def_id).or_default(); + if !candidates + .iter() + .any(|candidate| candidate.name == indexed.name) + { + candidates.push(indexed); + } +} + fn indexed_var_ref(name: &rumoca_core::VarName, var: &flat::Variable) -> IndexedVarRef { let path = var .component_ref @@ -191,6 +331,7 @@ fn resolve_best_owner_scoped_candidate( raw_leaf: &str, candidates: &[IndexedVarRef], owner: Option<&str>, + mode: OwnerScopeMode, ) -> Option { let raw = name.as_str(); let owner_path = owner.map(rumoca_core::ComponentPath::from_flat_path); @@ -204,7 +345,7 @@ fn resolve_best_owner_scoped_candidate( .filter_map(|candidate| { let owner_score = owner_path .as_ref() - .map(|owner| owner_scope_score(owner, &candidate.path)) + .map(|owner| owner_scope_score(owner, &candidate.path, mode)) .unwrap_or_else(|| (candidates.len() == 1).then_some(0))?; let suffix_score = path_suffix_score(&raw_path, &candidate.path); Some(((owner_score, suffix_score), candidate)) @@ -220,6 +361,13 @@ fn resolve_best_owner_scoped_candidate( (!ambiguous && best.name.as_str() != raw).then(|| best.name.clone()) } +#[derive(Clone, Copy)] +enum OwnerScopeMode { + DescendantOnly, + EnclosingOnly, + DescendantOrEnclosing, +} + fn path_suffix_score( raw_path: &rumoca_core::ComponentPath, candidate_path: &rumoca_core::ComponentPath, @@ -236,14 +384,25 @@ fn path_suffix_score( fn owner_scope_score( owner_path: &rumoca_core::ComponentPath, candidate_path: &rumoca_core::ComponentPath, + mode: OwnerScopeMode, ) -> Option { let candidate_scope = candidate_path.prefix(candidate_path.len().saturating_sub(1))?; - if owner_path.starts_with(&candidate_scope) { - return Some(candidate_scope.len()); + if matches!( + mode, + OwnerScopeMode::DescendantOnly | OwnerScopeMode::DescendantOrEnclosing + ) && candidate_scope.starts_with(owner_path) + { + return Some(owner_path.len()); + } + if matches!( + mode, + OwnerScopeMode::EnclosingOnly | OwnerScopeMode::DescendantOrEnclosing + ) { + return owner_path + .starts_with(&candidate_scope) + .then_some(candidate_scope.len()); } - candidate_scope - .starts_with(owner_path) - .then_some(owner_path.len()) + None } fn is_class_qualified_reference(name: &rumoca_core::Reference) -> bool { @@ -378,3 +537,108 @@ impl ExpressionRewriter for DefIdVarRefCanonicalizer<'_> { } impl StatementRewriter for DefIdVarRefCanonicalizer<'_> {} + +#[cfg(test)] +mod tests { + use super::*; + use rumoca_core::{ComponentRefPart, ComponentReference, DefId, Reference, Span, VarName}; + + fn test_span() -> Span { + Span::from_offsets( + rumoca_core::SourceId::from_source_name("phase_flatten_postprocess_def_id_source.mo"), + 0, + 1, + ) + } + + fn component_ref(path: &[&str], def_id: DefId) -> ComponentReference { + ComponentReference { + local: false, + span: test_span(), + parts: path + .iter() + .map(|part| ComponentRefPart { + ident: (*part).to_string(), + span: test_span(), + subs: Vec::new(), + }) + .collect(), + def_id: Some(def_id), + } + } + + fn variable(name: &str, def_id: DefId) -> flat::Variable { + let mut var = flat::Variable::empty_with_span(test_span()); + var.name = VarName::new(name); + var.component_ref = Some(ComponentReference::from_flat_segments( + name, + test_span(), + Some(def_id), + )); + var + } + + #[test] + fn aggregate_array_member_reference_is_not_rewritten_to_same_leaf_parent_var() { + let omega_def = DefId::new(7); + let mut flat = flat::Model::new(); + flat.add_variable( + VarName::new("vehicle.omega"), + variable("vehicle.omega", omega_def), + ); + flat.add_variable( + VarName::new("vehicle.motor[1].omega"), + variable("vehicle.motor[1].omega", omega_def), + ); + flat.add_variable( + VarName::new("vehicle.motor[2].omega"), + variable("vehicle.motor[2].omega", omega_def), + ); + flat.add_equation(flat::Equation::new( + rumoca_core::Expression::VarRef { + name: Reference::from_component_reference(component_ref( + &["vehicle", "motor", "omega"], + omega_def, + )), + subscripts: Vec::new(), + span: test_span(), + }, + test_span(), + flat::EquationOrigin::ComponentEquation { + component: "vehicle".to_string(), + }, + )); + + canonicalize_varrefs_via_instantiated_def_ids(&mut flat); + + let rumoca_core::Expression::VarRef { name, .. } = &flat.equations[0].residual else { + panic!("expected aggregate var ref to remain a var ref"); + }; + assert_eq!(name.as_str(), "vehicle.motor.omega"); + } + + #[test] + fn aggregate_projection_protection_requires_existing_penultimate_alias() { + let known_without_alias = HashSet::from([ + "rootMeanSquareVoltage.product.u[1]".to_string(), + "rootMeanSquareVoltage.product.u[2]".to_string(), + ]); + assert!( + !aggregate_projection_needs_alias_protection( + "rootMeanSquareVoltage.product.u", + &known_without_alias, + ), + "ordinary array member projections must not suppress scalar/index collapse" + ); + + let known_with_alias = HashSet::from([ + "vehicle.omega".to_string(), + "vehicle.motor[1].omega".to_string(), + "vehicle.motor[2].omega".to_string(), + ]); + assert!( + aggregate_projection_needs_alias_protection("vehicle.motor.omega", &known_with_alias,), + "array-member field projection must be protected when penultimate collapse would target a sibling leaf" + ); + } +} diff --git a/crates/rumoca-phase-flatten/src/postprocess_field_access.rs b/crates/rumoca-phase-flatten/src/postprocess_field_access.rs index fa3231fea..64db02dbd 100644 --- a/crates/rumoca-phase-flatten/src/postprocess_field_access.rs +++ b/crates/rumoca-phase-flatten/src/postprocess_field_access.rs @@ -156,10 +156,17 @@ fn record_type_contains_fields( fn record_constructor_field_names(flat: &flat::Model, type_name: &str) -> Option> { flat.functions .iter() - .find(|(name, function)| { + .filter(|(name, function)| { function.is_constructor && rumoca_core::qualified_type_name_matches(name.as_str(), type_name) }) + .max_by_key(|(name, function)| { + ( + name.as_str() == type_name, + function.inputs.len(), + name.as_str().len(), + ) + }) .map(|(_, function)| { function .inputs @@ -257,22 +264,27 @@ fn append_flat_subscripts( rendered: &mut String, subscripts: &[rumoca_core::Subscript], ) -> Option<()> { - for subscript in subscripts { - match subscript { - rumoca_core::Subscript::Index { value, .. } => { - rendered.push('['); - rendered.push_str(&value.to_string()); - rendered.push(']'); - } - rumoca_core::Subscript::Expr { expr, .. } => { - let value = constant_integer_bound(expr)?; - rendered.push('['); - rendered.push_str(&value.to_string()); - rendered.push(']'); - } - rumoca_core::Subscript::Colon { .. } => return None, + if subscripts.is_empty() { + return Some(()); + } + + let values = subscripts + .iter() + .map(|subscript| match subscript { + rumoca_core::Subscript::Index { value, .. } => Some(*value), + rumoca_core::Subscript::Expr { expr, .. } => constant_integer_bound(expr), + rumoca_core::Subscript::Colon { .. } => None, + }) + .collect::>>()?; + + rendered.push('['); + for (index, value) in values.iter().enumerate() { + if index > 0 { + rendered.push(','); } + rendered.push_str(&value.to_string()); } + rendered.push(']'); Some(()) } @@ -370,4 +382,91 @@ mod tests { }; assert_eq!(outer_field, "inner"); } + + #[test] + fn nested_constructor_resolution_prefers_specific_constructor_metadata() { + let mut flat = flat::Model::new(); + flat.add_function(constructor( + "Pkg.Outer", + vec![record_param( + "state", + "Buildings.Media.Air.ThermodynamicState", + )], + )); + flat.add_function(constructor( + "Modelica.Media.Interfaces.PartialSimpleMedium.ThermodynamicState", + vec![ + rumoca_core::FunctionParam::new("p", "Real", test_span()), + rumoca_core::FunctionParam::new("T", "Real", test_span()), + ], + )); + flat.add_function(constructor( + "Buildings.Media.Air.ThermodynamicState", + vec![ + rumoca_core::FunctionParam::new("p", "Real", test_span()), + rumoca_core::FunctionParam::new("T", "Real", test_span()), + rumoca_core::FunctionParam::new("X", "Real", test_span()), + ], + )); + flat.add_variable( + rumoca_core::VarName::new("target.record.p"), + variable( + "target.record.p", + direct_constructor_field("Pkg.Outer", "p"), + ), + ); + flat.add_variable( + rumoca_core::VarName::new("target.record.T"), + variable( + "target.record.T", + direct_constructor_field("Pkg.Outer", "T"), + ), + ); + flat.add_variable( + rumoca_core::VarName::new("target.record.X"), + variable( + "target.record.X", + direct_constructor_field("Pkg.Outer", "X"), + ), + ); + + resolve_nested_constructor_field_access_bindings(&mut flat); + + let Some(rumoca_core::Expression::FieldAccess { base, field, .. }) = flat + .variables + .get(&rumoca_core::VarName::new("target.record.X")) + .and_then(|var| var.binding.as_ref()) + else { + panic!("expected projected field access"); + }; + assert_eq!(field, "X"); + let rumoca_core::Expression::FieldAccess { + field: outer_field, .. + } = base.as_ref() + else { + panic!("expected nested constructor field access"); + }; + assert_eq!(outer_field, "state"); + } + + #[test] + fn renders_multidimensional_subscripts_as_one_flat_index_group() { + let mut rendered = "cell".to_string(); + append_flat_subscripts( + &mut rendered, + &[ + rumoca_core::Subscript::Index { + value: 1, + span: test_span(), + }, + rumoca_core::Subscript::Index { + value: 2, + span: test_span(), + }, + ], + ) + .expect("constant subscripts should render"); + + assert_eq!(rendered, "cell[1,2]"); + } } diff --git a/crates/rumoca-phase-flatten/src/postprocess_record_alias.rs b/crates/rumoca-phase-flatten/src/postprocess_record_alias.rs index 6fbfaa2f1..3ee80a517 100644 --- a/crates/rumoca-phase-flatten/src/postprocess_record_alias.rs +++ b/crates/rumoca-phase-flatten/src/postprocess_record_alias.rs @@ -5,14 +5,35 @@ pub(super) fn canonicalize_record_alias_expr( expr: &mut rumoca_core::Expression, ctx: &Context, known_variables: &HashSet, +) { + canonicalize_record_alias_expr_in_owner(expr, ctx, known_variables, None); +} + +pub(super) fn canonicalize_record_alias_expr_in_owner( + expr: &mut rumoca_core::Expression, + ctx: &Context, + known_variables: &HashSet, + owner: Option<&rumoca_core::ComponentPath>, ) { let mut rewriter = RecordAliasCanonicalizer { ctx, known_variables, + owner, }; *expr = rewriter.rewrite_expression(expr); } +pub(super) fn canonicalize_record_alias_opt_expr_in_owner( + expr: &mut Option, + ctx: &Context, + known_variables: &HashSet, + owner: &rumoca_core::ComponentPath, +) { + if let Some(expr) = expr { + canonicalize_record_alias_expr_in_owner(expr, ctx, known_variables, Some(owner)); + } +} + pub(super) fn canonicalize_record_alias_statements( statements: &mut [rumoca_core::Statement], ctx: &Context, @@ -21,6 +42,7 @@ pub(super) fn canonicalize_record_alias_statements( let mut rewriter = RecordAliasCanonicalizer { ctx, known_variables, + owner: None, }; for statement in statements { *statement = rewriter.rewrite_statement(statement); @@ -71,9 +93,19 @@ pub(super) fn canonicalize_record_alias_when_equations( struct RecordAliasCanonicalizer<'a> { ctx: &'a Context, known_variables: &'a HashSet, + owner: Option<&'a rumoca_core::ComponentPath>, } impl ExpressionRewriter for RecordAliasCanonicalizer<'_> { + fn rewrite_expression(&mut self, expr: &rumoca_core::Expression) -> rumoca_core::Expression { + if let rumoca_core::Expression::FieldAccess { base, field, span } = expr + && let Some(rewritten) = self.rewrite_record_alias_field_access(base, field, *span) + { + return rewritten; + } + self.walk_expression(expr) + } + fn rewrite_var_ref_expression( &mut self, name: &rumoca_core::Reference, @@ -81,7 +113,7 @@ impl ExpressionRewriter for RecordAliasCanonicalizer<'_> { span: rumoca_core::Span, ) -> rumoca_core::Expression { let rewritten_name = if subscripts.is_empty() { - record_alias_rewrite_name(name.as_str(), self.ctx, self.known_variables) + record_alias_rewrite_name(name.as_str(), self.ctx, self.known_variables, self.owner) .map(rumoca_core::Reference::new) .unwrap_or_else(|| name.clone()) } else { @@ -95,4 +127,179 @@ impl ExpressionRewriter for RecordAliasCanonicalizer<'_> { } } +impl RecordAliasCanonicalizer<'_> { + fn rewrite_record_alias_field_access( + &mut self, + base: &rumoca_core::Expression, + field: &str, + span: rumoca_core::Span, + ) -> Option { + let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = base + else { + return None; + }; + if !subscripts.is_empty() { + return None; + } + let field_path = format!("{}.{}", name.as_str(), field); + record_alias_rewrite_name(&field_path, self.ctx, self.known_variables, self.owner) + .or_else(|| { + owner_projected_record_field_candidate( + name.as_str(), + field, + self.owner, + self.known_variables, + ) + }) + .map(|rewritten| rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new(rewritten), + subscripts: vec![], + span, + }) + } +} + impl StatementRewriter for RecordAliasCanonicalizer<'_> {} + +pub(super) fn rewrite_name( + name: &str, + ctx: &Context, + known_variables: &HashSet, + owner: Option<&rumoca_core::ComponentPath>, +) -> Option { + let name_path = rumoca_core::ComponentPath::from_flat_path(name); + ctx.record_aliases.iter().find_map(|(alias, target)| { + if !name_path.starts_with(alias) || name_path.len() == alias.len() { + return None; + } + let suffix = name_path + .suffix_from(alias.len()) + .expect("suffix index is in range"); + record_alias_candidate(target, alias, &suffix, known_variables, owner) + }) +} + +fn owner_projected_record_field_candidate( + base_name: &str, + field: &str, + owner: Option<&rumoca_core::ComponentPath>, + known_variables: &HashSet, +) -> Option { + let owner = owner?; + let base = rumoca_core::ComponentPath::from_flat_path(base_name); + let owner_parts = owner.parts(); + let (indexed_pos, subscript) = + owner_parts + .iter() + .enumerate() + .rev() + .find_map(|(idx, part)| { + component_part_subscript_suffix(part) + .filter(|suffix| !suffix.is_empty()) + .map(|suffix| (idx, suffix)) + })?; + if owner_parts.len() < 2 { + return None; + } + if owner_parts.last().map(String::as_str) != Some(field) { + return None; + } + let projected_leaf_pos = owner_parts.len() - 2; + if projected_leaf_pos <= indexed_pos { + return None; + } + let shared_prefix = owner.prefix(indexed_pos)?; + if !base.starts_with(&shared_prefix) { + return None; + } + let owner_sibling_candidate = rumoca_core::ComponentPath::from_parts( + owner_parts[..=indexed_pos] + .iter() + .chain(owner_parts[projected_leaf_pos..].iter()) + .cloned(), + ) + .to_flat_string(); + if owner_sibling_candidate != owner.as_str() + && known_variables.contains(&owner_sibling_candidate) + { + return Some(owner_sibling_candidate); + } + let mut candidate_parts = owner_parts[..indexed_pos].to_vec(); + candidate_parts.push(format!("{}{}", owner_parts[projected_leaf_pos], subscript)); + if projected_leaf_pos + 1 < owner_parts.len() { + candidate_parts.extend(owner_parts[projected_leaf_pos + 1..].iter().cloned()); + } else { + candidate_parts.push(field.to_string()); + } + let candidate = rumoca_core::ComponentPath::from_parts(candidate_parts).to_flat_string(); + if candidate == owner.as_str() { + return None; + } + known_variables.contains(&candidate).then_some(candidate) +} + +fn record_alias_candidate( + target: &rumoca_core::ComponentPath, + alias: &rumoca_core::ComponentPath, + suffix: &rumoca_core::ComponentPath, + known_variables: &HashSet, + owner: Option<&rumoca_core::ComponentPath>, +) -> Option { + let direct = target.join(suffix).to_flat_string(); + if known_variables.contains(&direct) { + return Some(direct); + } + let alias_indexed_target = target_with_projected_alias_index(target, alias); + if let Some(indexed_target) = alias_indexed_target { + let indexed = indexed_target.join(suffix).to_flat_string(); + if known_variables.contains(&indexed) { + return Some(indexed); + } + } + let owner_indexed_target = target_with_owner_projected_index(target, owner?)?; + let indexed = owner_indexed_target.join(suffix).to_flat_string(); + known_variables.contains(&indexed).then_some(indexed) +} + +fn target_with_owner_projected_index( + target: &rumoca_core::ComponentPath, + owner: &rumoca_core::ComponentPath, +) -> Option { + target_with_projected_index(target, owner.parts()) +} + +fn target_with_projected_alias_index( + target: &rumoca_core::ComponentPath, + alias: &rumoca_core::ComponentPath, +) -> Option { + target_with_projected_index(target, alias.parts()) +} + +fn target_with_projected_index( + target: &rumoca_core::ComponentPath, + indexed_parts: &[String], +) -> Option { + let target_parts = target.parts(); + let last_target = target_parts.last()?; + if component_part_has_subscript(last_target) { + return None; + } + let subscript = indexed_parts.iter().rev().find_map(|part| { + component_part_subscript_suffix(part).filter(|suffix| !suffix.is_empty()) + })?; + let mut projected_parts = target_parts.to_vec(); + let last = projected_parts.last_mut()?; + last.push_str(subscript); + Some(rumoca_core::ComponentPath::from_parts(projected_parts)) +} + +fn component_part_has_subscript(part: &str) -> bool { + component_part_subscript_suffix(part).is_some() +} + +fn component_part_subscript_suffix(part: &str) -> Option<&str> { + let start = part.find('[')?; + part.ends_with(']').then_some(&part[start..]) +} diff --git a/crates/rumoca-phase-flatten/src/postprocess_record_alias_tests.rs b/crates/rumoca-phase-flatten/src/postprocess_record_alias_tests.rs index 37b1e7bb6..a487892c7 100644 --- a/crates/rumoca-phase-flatten/src/postprocess_record_alias_tests.rs +++ b/crates/rumoca-phase-flatten/src/postprocess_record_alias_tests.rs @@ -215,6 +215,412 @@ fn def_id_canonicalization_preserves_resolved_package_constant_refs() { assert_eq!(name.target_def_id(), Some(package_constant_def)); } +#[test] +fn def_id_canonicalization_uses_symbol_ancestry_for_inherited_attribute_refs() { + let mut model = flat::Model::new(); + let source_def = rumoca_core::DefId::new(27726); + let valve_instance_def = rumoca_core::DefId::new(93015); + let other_instance_def = rumoca_core::DefId::new(93016); + for (name, instance_def) in [ + ("val.m_flow_nominal_pos", valve_instance_def), + ("other.m_flow_nominal_pos", other_instance_def), + ] { + model.add_variable( + rumoca_core::VarName::new(name), + flat::Variable { + name: rumoca_core::VarName::new(name), + component_ref: Some(component_ref_with_def_id(name, instance_def)), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.symbol_ancestry.insert(instance_def, vec![source_def]); + } + model.add_variable( + rumoca_core::VarName::new("val.m_flow"), + flat::Variable { + name: rumoca_core::VarName::new("val.m_flow"), + nominal: Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference( + "Buildings.Fluid.BaseClasses.PartialResistance.m_flow_nominal_pos", + component_ref_with_def_id( + "Buildings.Fluid.BaseClasses.PartialResistance.m_flow_nominal_pos", + source_def, + ), + ), + subscripts: vec![], + span: rumoca_core::Span::DUMMY, + }), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + canonicalize_varrefs_via_instantiated_def_ids(&mut model); + + let nominal = model + .variables + .get(&rumoca_core::VarName::new("val.m_flow")) + .and_then(|var| var.nominal.as_ref()) + .expect("nominal should remain present"); + let rumoca_core::Expression::VarRef { name, .. } = nominal else { + panic!("expected nominal varref"); + }; + assert_eq!(name.as_str(), "val.m_flow_nominal_pos"); +} + +#[test] +fn def_id_canonicalization_resolves_inherited_bare_binding_to_owner_sibling() { + let mut model = flat::Model::new(); + let source_def = rumoca_core::DefId::new(19372); + let sibling_instance_def = rumoca_core::DefId::new(39979); + model.add_variable( + rumoca_core::VarName::new("material.B_rRef"), + flat::Variable { + name: rumoca_core::VarName::new("material.B_rRef"), + component_ref: Some(component_ref_with_def_id( + "material.B_rRef", + sibling_instance_def, + )), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model + .symbol_ancestry + .insert(sibling_instance_def, vec![source_def]); + model.add_variable( + rumoca_core::VarName::new("material.B_r"), + flat::Variable { + name: rumoca_core::VarName::new("material.B_r"), + binding: Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference( + "B_rRef", + component_ref_with_def_id("B_rRef", source_def), + ), + subscripts: vec![], + span: rumoca_core::Span::DUMMY, + }), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + canonicalize_varrefs_via_instantiated_def_ids(&mut model); + + let binding = model + .variables + .get(&rumoca_core::VarName::new("material.B_r")) + .and_then(|var| var.binding.as_ref()) + .expect("binding should remain present"); + let rumoca_core::Expression::VarRef { name, .. } = binding else { + panic!("expected varref binding"); + }; + assert_eq!(name.as_str(), "material.B_rRef"); +} + +#[test] +fn def_id_canonicalization_does_not_guess_deeper_unrelated_inherited_ref() { + let mut model = flat::Model::new(); + let source_def = rumoca_core::DefId::new(6239); + for (name, instance_def) in [ + ("configuration.degraded", rumoca_core::DefId::new(40000)), + ( + "system.configuration.degraded", + rumoca_core::DefId::new(40001), + ), + ( + "system.consumer.unrelated.degraded", + rumoca_core::DefId::new(40002), + ), + ] { + model.add_variable( + rumoca_core::VarName::new(name), + flat::Variable { + name: rumoca_core::VarName::new(name), + component_ref: Some(component_ref_with_def_id(name, instance_def)), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.symbol_ancestry.insert(instance_def, vec![source_def]); + } + model.add_variable( + rumoca_core::VarName::new("system.consumer.alias.value"), + flat::Variable { + name: rumoca_core::VarName::new("system.consumer.alias.value"), + binding: Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference( + "degraded", + component_ref_with_def_id("degraded", source_def), + ), + subscripts: vec![], + span: rumoca_core::Span::DUMMY, + }), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + canonicalize_varrefs_via_instantiated_def_ids(&mut model); + + let binding = model + .variables + .get(&rumoca_core::VarName::new("system.consumer.alias.value")) + .and_then(|var| var.binding.as_ref()) + .expect("binding should remain present"); + let rumoca_core::Expression::VarRef { name, .. } = binding else { + panic!("expected varref binding"); + }; + assert_eq!(name.as_str(), "degraded"); +} + +#[test] +fn def_id_canonicalization_does_not_guess_tied_shared_ancestor_ref() { + let mut model = flat::Model::new(); + let source_def = rumoca_core::DefId::new(6240); + for (name, instance_def) in [ + ("system.left.degraded", rumoca_core::DefId::new(40003)), + ("system.right.degraded", rumoca_core::DefId::new(40004)), + ] { + model.add_variable( + rumoca_core::VarName::new(name), + flat::Variable { + name: rumoca_core::VarName::new(name), + component_ref: Some(component_ref_with_def_id(name, instance_def)), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.symbol_ancestry.insert(instance_def, vec![source_def]); + } + model.add_variable( + rumoca_core::VarName::new("system.consumer.value"), + flat::Variable { + name: rumoca_core::VarName::new("system.consumer.value"), + binding: Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference( + "degraded", + component_ref_with_def_id("degraded", source_def), + ), + subscripts: vec![], + span: rumoca_core::Span::DUMMY, + }), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + canonicalize_varrefs_via_instantiated_def_ids(&mut model); + + let binding = model + .variables + .get(&rumoca_core::VarName::new("system.consumer.value")) + .and_then(|var| var.binding.as_ref()) + .expect("binding should remain present"); + let rumoca_core::Expression::VarRef { name, .. } = binding else { + panic!("expected varref binding"); + }; + assert_eq!(name.as_str(), "degraded"); +} + +#[test] +fn def_id_canonicalization_does_not_guess_zero_shared_ancestor_ref() { + let mut model = flat::Model::new(); + let source_def = rumoca_core::DefId::new(6241); + let instance_def = rumoca_core::DefId::new(40005); + model.add_variable( + rumoca_core::VarName::new("configuration.degraded"), + flat::Variable { + name: rumoca_core::VarName::new("configuration.degraded"), + component_ref: Some(component_ref_with_def_id( + "configuration.degraded", + instance_def, + )), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.symbol_ancestry.insert(instance_def, vec![source_def]); + model.add_variable( + rumoca_core::VarName::new("system.consumer.value"), + flat::Variable { + name: rumoca_core::VarName::new("system.consumer.value"), + binding: Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference( + "degraded", + component_ref_with_def_id("degraded", source_def), + ), + subscripts: vec![], + span: rumoca_core::Span::DUMMY, + }), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + canonicalize_varrefs_via_instantiated_def_ids(&mut model); + + let binding = model + .variables + .get(&rumoca_core::VarName::new("system.consumer.value")) + .and_then(|var| var.binding.as_ref()) + .expect("binding should remain present"); + let rumoca_core::Expression::VarRef { name, .. } = binding else { + panic!("expected varref binding"); + }; + assert_eq!(name.as_str(), "degraded"); +} + +#[test] +fn def_id_canonicalization_then_record_alias_uses_explicit_instance_ownership() { + let mut model = flat::Model::new(); + let source_def = rumoca_core::DefId::new(6242); + for (name, instance_def) in [ + ("configuration.degraded", rumoca_core::DefId::new(40006)), + ( + "system.configuration.degraded", + rumoca_core::DefId::new(40007), + ), + ( + "system.consumer.unrelated.degraded", + rumoca_core::DefId::new(40008), + ), + ] { + model.add_variable( + rumoca_core::VarName::new(name), + flat::Variable { + name: rumoca_core::VarName::new(name), + component_ref: Some(component_ref_with_def_id(name, instance_def)), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.symbol_ancestry.insert(instance_def, vec![source_def]); + } + model.add_variable( + rumoca_core::VarName::new("system.configuration.value"), + flat::Variable { + name: rumoca_core::VarName::new("system.configuration.value"), + binding: Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference( + "degraded", + component_ref_with_def_id("degraded", source_def), + ), + subscripts: vec![], + span: rumoca_core::Span::DUMMY, + }), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let mut ctx = Context::new(); + ctx.record_aliases.insert( + rumoca_core::ComponentPath::from_flat_path("system.configuration"), + rumoca_core::ComponentPath::from_flat_path("configuration"), + ); + + canonicalize_varrefs_via_instantiated_def_ids(&mut model); + canonicalize_varrefs_via_record_aliases(&mut model, &ctx); + + let binding = model + .variables + .get(&rumoca_core::VarName::new("system.configuration.value")) + .and_then(|var| var.binding.as_ref()) + .expect("binding should remain present"); + let rumoca_core::Expression::VarRef { name, .. } = binding else { + panic!("expected varref binding"); + }; + assert_eq!(name.as_str(), "configuration.degraded"); +} + +#[test] +fn def_id_canonicalization_prefers_owner_instance_before_enclosing_fallback() { + let mut model = flat::Model::new(); + let source_def = rumoca_core::DefId::new(610); + let nested_instance_def = rumoca_core::DefId::new(40065); + let owner_instance_def = rumoca_core::DefId::new(40206); + for (name, def_id) in [ + ("cell.cell.limIntegrator.y", nested_instance_def), + ("cell.limIntegrator.y", owner_instance_def), + ] { + model.add_variable( + rumoca_core::VarName::new(name), + flat::Variable { + name: rumoca_core::VarName::new(name), + component_ref: Some(component_ref_with_def_id(name, def_id)), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.symbol_ancestry.insert(def_id, vec![source_def]); + } + model.equations.push(flat::Equation { + residual: rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference( + "y", + component_ref_with_def_id("y", source_def), + ), + subscripts: vec![], + span: rumoca_core::Span::DUMMY, + }, + span: rumoca_core::Span::DUMMY, + origin: flat::EquationOrigin::ComponentEquation { + component: "cell.limIntegrator".to_string(), + }, + scalar_count: 1, + }); + + canonicalize_varrefs_via_instantiated_def_ids(&mut model); + + let rumoca_core::Expression::VarRef { name, .. } = &model.equations[0].residual else { + panic!("expected varref residual"); + }; + assert_eq!(name.as_str(), "cell.limIntegrator.y"); +} + +#[test] +fn def_id_canonicalization_prefers_known_structured_path_over_rendered_name() { + let mut model = flat::Model::new(); + let source_def = rumoca_core::DefId::new(621); + let instance_def = rumoca_core::DefId::new(40206); + model.add_variable( + rumoca_core::VarName::new("cell.limIntegrator.local_reset"), + flat::Variable { + name: rumoca_core::VarName::new("cell.limIntegrator.local_reset"), + component_ref: Some(component_ref_with_def_id( + "cell.limIntegrator.local_reset", + instance_def, + )), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.symbol_ancestry.insert(instance_def, vec![source_def]); + model.equations.push(flat::Equation { + residual: rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference( + "local_reset", + component_ref_with_def_id("cell.limIntegrator.local_reset", source_def), + ), + subscripts: vec![], + span: rumoca_core::Span::DUMMY, + }, + span: rumoca_core::Span::DUMMY, + origin: flat::EquationOrigin::ComponentEquation { + component: "cell.limIntegrator".to_string(), + }, + scalar_count: 1, + }); + + canonicalize_varrefs_via_instantiated_def_ids(&mut model); + + let rumoca_core::Expression::VarRef { name, .. } = &model.equations[0].residual else { + panic!("expected varref residual"); + }; + assert_eq!(name.as_str(), "cell.limIntegrator.local_reset"); +} + #[test] fn record_alias_canonicalization_visits_when_clauses_and_algorithms() { let mut model = flat::Model::new(); @@ -267,6 +673,398 @@ fn record_alias_canonicalization_visits_when_clauses_and_algorithms() { assert_eq!(name.as_str(), "pipe.port_a.p"); } +#[test] +fn record_alias_canonicalization_visits_variable_bindings_and_projects_array_record_fields() { + let mut model = flat::Model::new(); + model.add_variable( + rumoca_core::VarName::new("bank.per[1].Q"), + flat::Variable { + name: rumoca_core::VarName::new("bank.per[1].Q"), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.add_variable( + rumoca_core::VarName::new("bank.ch[1].per.Q"), + flat::Variable { + name: rumoca_core::VarName::new("bank.ch[1].per.Q"), + binding: Some(rumoca_core::Expression::FieldAccess { + base: Box::new(var_ref("bank.ch[1].per")), + field: "Q".to_string(), + span: Span::DUMMY, + }), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let mut ctx = Context::new(); + ctx.record_aliases.insert( + rumoca_core::ComponentPath::from_flat_path("bank.ch[1].per"), + rumoca_core::ComponentPath::from_flat_path("bank.per"), + ); + + canonicalize_varrefs_via_record_aliases(&mut model, &ctx); + + let binding = model + .variables + .get(&rumoca_core::VarName::new("bank.ch[1].per.Q")) + .and_then(|var| var.binding.as_ref()) + .expect("binding should remain present"); + let rumoca_core::Expression::VarRef { name, .. } = binding else { + panic!("expected field access to collapse to varref"); + }; + assert_eq!(name.as_str(), "bank.per[1].Q"); +} + +#[test] +fn record_alias_canonicalization_uses_variable_owner_to_project_unindexed_record_binding() { + let mut model = flat::Model::new(); + model.add_variable( + rumoca_core::VarName::new("bank.per[1].Q"), + flat::Variable { + name: rumoca_core::VarName::new("bank.per[1].Q"), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.add_variable( + rumoca_core::VarName::new("bank.ch[1].per.Q"), + flat::Variable { + name: rumoca_core::VarName::new("bank.ch[1].per.Q"), + binding: Some(rumoca_core::Expression::FieldAccess { + base: Box::new(var_ref("bank.dat")), + field: "Q".to_string(), + span: Span::DUMMY, + }), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let mut ctx = Context::new(); + ctx.record_aliases.insert( + rumoca_core::ComponentPath::from_flat_path("bank.dat"), + rumoca_core::ComponentPath::from_flat_path("bank.per"), + ); + + canonicalize_varrefs_via_record_aliases(&mut model, &ctx); + + let binding = model + .variables + .get(&rumoca_core::VarName::new("bank.ch[1].per.Q")) + .and_then(|var| var.binding.as_ref()) + .expect("binding should remain present"); + let rumoca_core::Expression::VarRef { name, .. } = binding else { + panic!("expected field access to collapse to owner-indexed target"); + }; + assert_eq!(name.as_str(), "bank.per[1].Q"); +} + +#[test] +fn record_alias_canonicalization_projects_owner_record_field_without_direct_alias() { + let mut model = flat::Model::new(); + model.add_variable( + rumoca_core::VarName::new("bank.per[1].Q"), + flat::Variable { + name: rumoca_core::VarName::new("bank.per[1].Q"), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.add_variable( + rumoca_core::VarName::new("bank.ch[1].per.Q"), + flat::Variable { + name: rumoca_core::VarName::new("bank.ch[1].per.Q"), + binding: Some(rumoca_core::Expression::FieldAccess { + base: Box::new(var_ref("bank.dat")), + field: "Q".to_string(), + span: Span::DUMMY, + }), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let ctx = Context::new(); + + canonicalize_varrefs_via_record_aliases(&mut model, &ctx); + + let binding = model + .variables + .get(&rumoca_core::VarName::new("bank.ch[1].per.Q")) + .and_then(|var| var.binding.as_ref()) + .expect("binding should remain present"); + let rumoca_core::Expression::VarRef { name, .. } = binding else { + panic!("expected owner-projected field access to collapse"); + }; + assert_eq!(name.as_str(), "bank.per[1].Q"); +} + +#[test] +fn record_alias_canonicalization_projects_nested_owner_record_leaf() { + let mut model = flat::Model::new(); + model.add_variable( + rumoca_core::VarName::new("bank.per[1].Q"), + flat::Variable { + name: rumoca_core::VarName::new("bank.per[1].Q"), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.add_variable( + rumoca_core::VarName::new("bank.ch[1].unit.per.Q"), + flat::Variable { + name: rumoca_core::VarName::new("bank.ch[1].unit.per.Q"), + binding: Some(rumoca_core::Expression::FieldAccess { + base: Box::new(var_ref("bank.dat")), + field: "Q".to_string(), + span: Span::DUMMY, + }), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + let mut ctx = Context::new(); + ctx.record_aliases.insert( + rumoca_core::ComponentPath::from_flat_path("other.alias"), + rumoca_core::ComponentPath::from_flat_path("other.target"), + ); + + canonicalize_varrefs_via_record_aliases(&mut model, &ctx); + + let binding = model + .variables + .get(&rumoca_core::VarName::new("bank.ch[1].unit.per.Q")) + .and_then(|var| var.binding.as_ref()) + .expect("binding should remain present"); + let rumoca_core::Expression::VarRef { name, .. } = binding else { + panic!("expected nested owner projection to collapse"); + }; + assert_eq!(name.as_str(), "bank.per[1].Q"); +} + +#[test] +fn record_alias_canonicalization_projects_to_exact_indexed_owner_sibling() { + let mut model = flat::Model::new(); + for name in [ + "bank.ch[1,2].per.Q", + "bank.unrelated.per[1,2].Q", + "bank.ch[1,2].unit.per.Q", + ] { + model.add_variable( + rumoca_core::VarName::new(name), + flat::Variable { + name: rumoca_core::VarName::new(name), + binding: (name == "bank.ch[1,2].unit.per.Q").then(|| { + rumoca_core::Expression::FieldAccess { + base: Box::new(var_ref("bank.per")), + field: "Q".to_string(), + span: Span::DUMMY, + } + }), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + + canonicalize_varrefs_via_record_aliases(&mut model, &Context::new()); + + let binding = model + .variables + .get(&rumoca_core::VarName::new("bank.ch[1,2].unit.per.Q")) + .and_then(|var| var.binding.as_ref()) + .expect("binding should remain present"); + let rumoca_core::Expression::VarRef { name, .. } = binding else { + panic!("expected exact indexed owner sibling"); + }; + assert_eq!(name.as_str(), "bank.ch[1,2].per.Q"); +} + +#[test] +fn record_alias_canonicalization_does_not_project_to_unrelated_unique_descendant() { + let mut model = flat::Model::new(); + model.add_variable( + rumoca_core::VarName::new("bank.unrelated.per[1,2].Q"), + flat::Variable { + name: rumoca_core::VarName::new("bank.unrelated.per[1,2].Q"), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + model.add_variable( + rumoca_core::VarName::new("bank.ch[1,2].unit.per.Q"), + flat::Variable { + name: rumoca_core::VarName::new("bank.ch[1,2].unit.per.Q"), + binding: Some(rumoca_core::Expression::FieldAccess { + base: Box::new(var_ref("bank.per")), + field: "Q".to_string(), + span: Span::DUMMY, + }), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + + canonicalize_varrefs_via_record_aliases(&mut model, &Context::new()); + + assert!(matches!( + model + .variables + .get(&rumoca_core::VarName::new("bank.ch[1,2].unit.per.Q")) + .and_then(|var| var.binding.as_ref()), + Some(rumoca_core::Expression::FieldAccess { .. }) + )); +} + +#[test] +fn record_alias_canonicalization_requires_field_to_match_owner_leaf() { + let mut model = flat::Model::new(); + for name in ["bank.ch[1].per.Q", "bank.ch[1].unit.per.Q"] { + model.add_variable( + rumoca_core::VarName::new(name), + flat::Variable { + name: rumoca_core::VarName::new(name), + binding: (name == "bank.ch[1].unit.per.Q").then(|| { + rumoca_core::Expression::FieldAccess { + base: Box::new(var_ref("bank.per")), + field: "P".to_string(), + span: Span::DUMMY, + } + }), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + + canonicalize_varrefs_via_record_aliases(&mut model, &Context::new()); + + assert!(matches!( + model + .variables + .get(&rumoca_core::VarName::new("bank.ch[1].unit.per.Q")) + .and_then(|var| var.binding.as_ref()), + Some(rumoca_core::Expression::FieldAccess { field, .. }) if field == "P" + )); +} + +#[test] +fn record_alias_canonicalization_preserves_distinct_connector_record_endpoints() { + let mut model = flat::Model::new(); + for (index, name) in [ + "device.i", + "device.pin_p.i.re", + "device.pin_p.i.im", + "device.pin_n.i.re", + "device.pin_n.i.im", + ] + .into_iter() + .enumerate() + { + model.add_variable( + rumoca_core::VarName::new(name), + flat::Variable { + name: rumoca_core::VarName::new(name), + component_ref: Some(component_ref_with_def_id( + name, + rumoca_core::DefId::new(800 + index as u32), + )), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + + let pin_p_def_id = rumoca_core::DefId::new(901); + let pin_n_def_id = rumoca_core::DefId::new(902); + let pin_p_ref = component_ref_with_def_id("device.pin_p.i", pin_p_def_id); + let pin_n_ref = component_ref_with_def_id("device.pin_n.i", pin_n_def_id); + model.add_equation(flat::Equation::new( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference( + "device.pin_p.i", + pin_p_ref.clone(), + ), + subscripts: vec![], + span: Span::DUMMY, + }), + rhs: Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::with_component_reference( + "device.pin_n.i", + pin_n_ref.clone(), + ), + subscripts: vec![], + span: Span::DUMMY, + }), + span: Span::DUMMY, + }, + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: "device".to_string(), + }, + )); + + canonicalize_varrefs_via_record_aliases(&mut model, &Context::new()); + + let rumoca_core::Expression::Binary { lhs, rhs, .. } = &model.equations[0].residual else { + panic!("expected binary residual"); + }; + let rumoca_core::Expression::VarRef { + name: pin_p_name, .. + } = lhs.as_ref() + else { + panic!("expected pin_p aggregate reference"); + }; + let rumoca_core::Expression::VarRef { + name: pin_n_name, .. + } = rhs.as_ref() + else { + panic!("expected pin_n aggregate reference"); + }; + assert_eq!( + [ + (pin_p_name.as_str(), pin_p_name.target_def_id()), + (pin_n_name.as_str(), pin_n_name.target_def_id()), + ], + [ + ("device.pin_p.i", Some(pin_p_def_id)), + ("device.pin_n.i", Some(pin_n_def_id)), + ] + ); + assert_eq!(pin_p_name.component_ref(), Some(&pin_p_ref)); + assert_eq!(pin_n_name.component_ref(), Some(&pin_n_ref)); +} + +#[test] +fn record_alias_canonicalization_preserves_decomposed_record_field_without_alias() { + let mut model = flat::Model::new(); + for name in ["h", "state.phase"] { + model.add_variable( + rumoca_core::VarName::new(name), + flat::Variable { + name: rumoca_core::VarName::new(name), + is_primitive: true, + ..flat::Variable::empty_with_span(test_span()) + }, + ); + } + model.add_equation(flat::Equation::new( + var_ref("state.h"), + Span::DUMMY, + flat::EquationOrigin::ComponentEquation { + component: String::new(), + }, + )); + + canonicalize_varrefs_via_record_aliases(&mut model, &Context::new()); + + let rumoca_core::Expression::VarRef { name, .. } = &model.equations[0].residual else { + panic!("expected preserved var ref"); + }; + assert_eq!(name.as_str(), "state.h"); +} + #[test] fn invalid_field_access_drop_handles_indexed_bases() { let mut model = flat::Model::new(); diff --git a/crates/rumoca-phase-flatten/src/variables.rs b/crates/rumoca-phase-flatten/src/variables.rs index 485632573..6d59af85d 100644 --- a/crates/rumoca-phase-flatten/src/variables.rs +++ b/crates/rumoca-phase-flatten/src/variables.rs @@ -13,6 +13,7 @@ use rumoca_ir_flat as flat; use crate::ast_lower; use crate::errors::FlattenError; use crate::functions; +use crate::pipeline::ComponentMemberScopes; use crate::qualify::{ImportMap, QualifyOptions, qualify_expression_with_imports}; use crate::source_spans::required_location_span; use rustc_hash::FxHashMap; @@ -186,6 +187,8 @@ pub(crate) fn create_flat_variable( tree: &ast::ClassTree, class_index: &ast::ClassDefIndex<'_>, imports: &VariableImportContext, + component_members: &ComponentMemberScopes, + simulated_root_name: Option<&str>, ) -> Result { let name = rumoca_core::VarName::new(instance.qualified_name.to_flat_string()); let source_span = instance_source_span(instance, tree, "flat variable")?; @@ -202,14 +205,16 @@ pub(crate) fn create_flat_variable( // Get def_map for resolving function call def_ids to qualified names let def_map = &tree.def_map; - let attrs = qualify_variable_attributes(VariableQualifyContext { + let mut attrs = qualify_variable_attributes(VariableQualifyContext { instance, tree, class_index, imports, prefix: &prefix, + component_members, opts, def_map, + simulated_root_name, })?; // Binding expressions need careful handling: @@ -218,15 +223,21 @@ pub(crate) fn create_flat_variable( // - Modification bindings (e.g., `body(useQuaternions=useQuaternions)`) reference // variables in the lexical scope where the modification is written. // This scope is tracked during instantiation (MLS §7.2.4). - let binding = qualify_variable_binding(VariableQualifyContext { + let mut binding = qualify_variable_binding(VariableQualifyContext { instance, tree, class_index, imports, prefix: &prefix, + component_members, opts, def_map, + simulated_root_name, })?; + if instance_is_string_type(instance) { + recover_string_literal_expr(&mut binding, &prefix); + recover_string_literal_expr(&mut attrs.start, &prefix); + } let component_ref = Some(ast::instance::component_reference_for_instance( &instance.qualified_name, @@ -272,6 +283,70 @@ pub(crate) fn create_flat_variable( }) } +fn instance_is_string_type(instance: &ast::InstanceData) -> bool { + rumoca_core::qualified_type_name_matches(&instance.type_name, "String") + || instance.type_id == rumoca_core::TypeId(3) +} + +fn recover_string_literal_expr( + expr: &mut Option, + prefix: &ast::QualifiedName, +) { + let Some(recovered) = expr.as_ref().and_then(|expr| { + recover_string_literal_from_invalid_component_expr(expr, &prefix.to_flat_string()) + }) else { + return; + }; + *expr = Some(recovered); +} + +pub(crate) fn recover_string_literal_from_invalid_component_expr( + expr: &rumoca_core::Expression, + prefix: &str, +) -> Option { + let path = string_literal_candidate_path(expr)?; + if path + .parts() + .iter() + .all(|part| is_valid_modelica_identifier(part)) + { + return None; + } + let literal_path = if prefix.is_empty() { + path + } else { + path.strip_prefix(&rumoca_core::ComponentPath::from_flat_path(prefix))? + }; + let span = expr.span()?; + Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String(literal_path.to_flat_string()), + span, + }) +} + +fn string_literal_candidate_path( + expr: &rumoca_core::Expression, +) -> Option { + match expr { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } if subscripts.is_empty() => Some(rumoca_core::ComponentPath::from_reference(name)), + rumoca_core::Expression::FieldAccess { base, field, .. } => { + Some(string_literal_candidate_path(base)?.join_part_slice(std::slice::from_ref(field))) + } + _ => None, + } +} + +fn is_valid_modelica_identifier(part: &str) -> bool { + let mut chars = part.chars(); + let Some(first) = chars.next() else { + return false; + }; + (first == '_' || first.is_ascii_alphabetic()) + && chars.all(|ch| ch == '_' || ch.is_ascii_alphanumeric()) +} + fn canonicalize_function_calls( mut expr: rumoca_core::Expression, source_scope: Option<&str>, @@ -294,8 +369,10 @@ struct VariableQualifyContext<'a, 'tree> { class_index: &'a ast::ClassDefIndex<'tree>, imports: &'a VariableImportContext, prefix: &'a ast::QualifiedName, + component_members: &'a ComponentMemberScopes, opts: QualifyOptions, def_map: &'a crate::ResolveDefMap, + simulated_root_name: Option<&'a str>, } struct QualifiedVariableAttributes { @@ -324,6 +401,7 @@ fn qualify_variable_attribute( let Some(expr) = expr else { return Ok(None); }; + let expr = flat_value_attribute_expr(ctx.instance, attr_name, expr); let attr_prefix = attribute_prefix(ctx.instance, attr_name, ctx.prefix.clone()); let qualified = qualify_expression_with_imports( expr, @@ -337,23 +415,46 @@ fn qualify_variable_attribute( .get(attr_name) .map(String::as_str) .or(ctx.imports.declaration_function_scope.as_deref()); + let instance_name = declaration_instance_name(ctx.simulated_root_name, ctx.prefix); Ok(Some(canonicalize_function_calls( - ast_lower::expression_from_ast_with_def_map(&qualified, Some(ctx.def_map))?, + ast_lower::expression_from_ast_with_context( + &qualified, + ast_lower::LoweringContext { + def_map: Some(ctx.def_map), + class_tree: Some(ctx.tree), + instance_name: instance_name.as_deref(), + }, + )?, source_scope, ctx.tree, ctx.class_index, ))) } +fn flat_value_attribute_expr<'a>( + instance: &'a ast::InstanceData, + attr_name: &str, + expr: &'a ast::Expression, +) -> &'a ast::Expression { + if attr_name == "start" + && instance.binding_from_modification + && instance + .binding_source + .as_ref() + .is_some_and(|source| source == expr) + && let (Some(source), Some(binding)) = + (instance.binding_source.as_ref(), instance.binding.as_ref()) + && modifier_binding_selects_source_array_element(source, binding) + { + return binding; + } + expr +} + fn qualify_variable_binding( ctx: VariableQualifyContext<'_, '_>, ) -> Result, FlattenError> { - let Some(expr) = ctx - .instance - .binding_source - .as_ref() - .or(ctx.instance.binding.as_ref()) - else { + let Some(expr) = flat_value_binding_expr(ctx.instance) else { return Ok(None); }; if ctx.instance.binding_from_modification { @@ -362,15 +463,76 @@ fn qualify_variable_binding( qualify_declaration_binding(ctx, expr).map(Some) } +fn flat_value_binding_expr(instance: &ast::InstanceData) -> Option<&ast::Expression> { + if instance.binding_from_modification + && let (Some(source), Some(binding)) = + (instance.binding_source.as_ref(), instance.binding.as_ref()) + && modifier_binding_selects_source_array_element(source, binding) + { + return Some(binding); + } + instance + .binding_source + .as_ref() + .or(instance.binding.as_ref()) +} + +fn modifier_binding_selects_source_array_element( + source: &ast::Expression, + binding: &ast::Expression, +) -> bool { + let ( + ast::Expression::ComponentReference(source_ref), + ast::Expression::ComponentReference(binding_ref), + ) = (source, binding) + else { + return false; + }; + component_ref_idents_equal(source_ref, binding_ref) + && !component_ref_has_subscripts(source_ref) + && component_ref_has_subscripts(binding_ref) +} + +fn component_ref_idents_equal( + lhs: &ast::ComponentReference, + rhs: &ast::ComponentReference, +) -> bool { + lhs.parts.len() == rhs.parts.len() + && lhs + .parts + .iter() + .zip(rhs.parts.iter()) + .all(|(lhs, rhs)| lhs.ident.text == rhs.ident.text) +} + +fn component_ref_has_subscripts(reference: &ast::ComponentReference) -> bool { + reference + .parts + .iter() + .any(|part| part.subs.as_ref().is_some_and(|subs| !subs.is_empty())) +} + fn qualify_modification_binding( ctx: VariableQualifyContext<'_, '_>, expr: &ast::Expression, ) -> Result { let mod_prefix = modification_binding_prefix(ctx.instance, ctx.tree)?; - let qualified = - qualify_expression_with_imports(expr, &mod_prefix, ctx.opts, ctx.imports.binding_imports()); + let binding_imports = ctx.component_members.scoped_component_imports( + expr, + &mod_prefix, + ctx.imports.binding_imports(), + ); + let qualified = qualify_expression_with_imports(expr, &mod_prefix, ctx.opts, &binding_imports); + let instance_name = declaration_instance_name(ctx.simulated_root_name, &mod_prefix); Ok(canonicalize_function_calls( - ast_lower::expression_from_ast_with_def_map(&qualified, Some(ctx.def_map))?, + ast_lower::expression_from_ast_with_context( + &qualified, + ast_lower::LoweringContext { + def_map: Some(ctx.def_map), + class_tree: Some(ctx.tree), + instance_name: instance_name.as_deref(), + }, + )?, ctx.imports.binding_function_scope.as_deref(), ctx.tree, ctx.class_index, @@ -381,21 +543,45 @@ fn qualify_declaration_binding( ctx: VariableQualifyContext<'_, '_>, expr: &ast::Expression, ) -> Result { + let declaration_imports = + ctx.component_members + .scoped_component_imports(expr, ctx.prefix, &ctx.imports.declaration); let qualified = - qualify_expression_with_imports(expr, ctx.prefix, ctx.opts, &ctx.imports.declaration); + qualify_expression_with_imports(expr, ctx.prefix, ctx.opts, &declaration_imports); let source_scope = ctx .imports .binding_function_scope .as_deref() .or(ctx.imports.declaration_function_scope.as_deref()); + let instance_name = declaration_instance_name(ctx.simulated_root_name, ctx.prefix); Ok(canonicalize_function_calls( - ast_lower::expression_from_ast_with_def_map(&qualified, Some(ctx.def_map))?, + ast_lower::expression_from_ast_with_context( + &qualified, + ast_lower::LoweringContext { + def_map: Some(ctx.def_map), + class_tree: Some(ctx.tree), + instance_name: instance_name.as_deref(), + }, + )?, source_scope, ctx.tree, ctx.class_index, )) } +fn declaration_instance_name( + simulated_root_name: Option<&str>, + prefix: &ast::QualifiedName, +) -> Option { + let root = simulated_root_name?; + let suffix = prefix.to_flat_string(); + if suffix.is_empty() { + Some(root.to_string()) + } else { + Some(format!("{root}.{suffix}")) + } +} + #[cfg(test)] mod tests { use super::*; @@ -470,8 +656,16 @@ mod tests { }; let tree = test_tree(); let imports = VariableImportContext::default(); + let component_members = ComponentMemberScopes::default(); let class_index = ast::ClassDefIndex::from_tree(&tree); - let flat = create_flat_variable(&instance, &tree, &class_index, &imports)?; + let flat = create_flat_variable( + &instance, + &tree, + &class_index, + &imports, + &component_members, + None, + )?; assert_eq!( flat.component_ref .as_ref() @@ -516,9 +710,17 @@ mod tests { }; let tree = test_tree(); let imports = VariableImportContext::default(); + let component_members = ComponentMemberScopes::default(); let class_index = ast::ClassDefIndex::from_tree(&tree); - let flat = - create_flat_variable(&instance, &tree, &class_index, &imports).expect("flat variable"); + let flat = create_flat_variable( + &instance, + &tree, + &class_index, + &imports, + &component_members, + None, + ) + .expect("flat variable"); let max = flat.max.expect("max"); match max { rumoca_core::Expression::VarRef { @@ -530,4 +732,193 @@ mod tests { _ => panic!("expected max to become a qualified VarRef"), } } + + #[test] + fn test_create_flat_variable_resolves_declaration_binding_to_parent_instance_member() { + let instance = ast::InstanceData { + qualified_name: ast::QualifiedName::from_dotted("battery.r0.T"), + source_location: test_location(20, 31), + binding: Some(comp_ref(&["cellData", "T_ref"])), + is_primitive: true, + ..ast::InstanceData::default() + }; + let tree = test_tree(); + let imports = VariableImportContext::default(); + let mut component_members = ComponentMemberScopes::default(); + component_members.insert_component_member_path( + &rumoca_core::ComponentPath::from_flat_path("battery.cellData.T_ref"), + ); + component_members.insert_component_member_path( + &rumoca_core::ComponentPath::from_flat_path("battery.r0.T"), + ); + let class_index = ast::ClassDefIndex::from_tree(&tree); + let flat = create_flat_variable( + &instance, + &tree, + &class_index, + &imports, + &component_members, + None, + ) + .expect("flat variable"); + let binding = flat.binding.expect("binding"); + match binding { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } => { + assert_eq!(name.as_str(), "battery.cellData.T_ref"); + assert!(subscripts.is_empty()); + } + _ => panic!("expected binding to become a parent-scoped VarRef"), + } + } + + #[test] + fn test_create_flat_variable_lowers_get_instance_name_to_parent_instance() { + let instance = ast::InstanceData { + qualified_name: ast::QualifiedName::from_dotted("building.modelicaNameBuilding"), + source_location: test_location(40, 64), + binding_source: Some(ast::Expression::FunctionCall { + comp: ast::ComponentReference { + local: false, + parts: vec![ast::ComponentRefPart { + ident: rumoca_core::Token { + text: Arc::from("getInstanceName"), + ..rumoca_core::Token::default() + }, + subs: None, + }], + def_id: None, + span: test_span(), + }, + args: Vec::new(), + span: test_span(), + }), + is_primitive: true, + ..ast::InstanceData::default() + }; + let tree = test_tree(); + let imports = VariableImportContext::default(); + let component_members = ComponentMemberScopes::default(); + let class_index = ast::ClassDefIndex::from_tree(&tree); + let flat = create_flat_variable( + &instance, + &tree, + &class_index, + &imports, + &component_members, + Some("RootModel"), + ) + .expect("flat variable"); + + assert_eq!( + flat.binding, + Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("RootModel.building".to_string()), + span: test_span(), + }) + ); + } + + #[test] + fn test_create_flat_variable_preserves_string_literal_binding() { + let instance = ast::InstanceData { + qualified_name: ast::QualifiedName::from_dotted("building.spawnExe"), + source_location: test_location(6, 19), + binding_source: Some(ast::Expression::Terminal { + terminal_type: ast::TerminalType::String, + token: rumoca_core::Token { + text: Arc::from("\"spawn-0.4.3-7048a72798\""), + ..rumoca_core::Token::default() + }, + span: test_span(), + }), + is_primitive: true, + ..ast::InstanceData::default() + }; + let tree = test_tree(); + let imports = VariableImportContext::default(); + let component_members = ComponentMemberScopes::default(); + let class_index = ast::ClassDefIndex::from_tree(&tree); + let flat = create_flat_variable( + &instance, + &tree, + &class_index, + &imports, + &component_members, + None, + ) + .expect("flat variable"); + + assert_eq!( + flat.binding, + Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("spawn-0.4.3-7048a72798".to_string()), + span: test_span(), + }) + ); + } + + #[test] + fn test_create_flat_variable_recovers_invalid_string_literal_varref() { + let instance = ast::InstanceData { + qualified_name: ast::QualifiedName::from_dotted("building.spawnExe"), + source_location: test_location(6, 19), + type_name: "String".to_string(), + binding: Some(comp_ref(&["spawn-0", "4", "3-7048a72798"])), + start: Some(comp_ref(&["spawn-0", "4", "3-7048a72798"])), + is_primitive: true, + ..ast::InstanceData::default() + }; + let tree = test_tree(); + let imports = VariableImportContext::default(); + let component_members = ComponentMemberScopes::default(); + let class_index = ast::ClassDefIndex::from_tree(&tree); + let flat = create_flat_variable( + &instance, + &tree, + &class_index, + &imports, + &component_members, + None, + ) + .expect("flat variable"); + + let expected = Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("spawn-0.4.3-7048a72798".to_string()), + span: test_span(), + }); + assert_eq!(flat.binding, expected); + assert_eq!(flat.start, expected); + } + + #[test] + fn test_recover_string_literal_from_invalid_field_access_path() { + let invalid_path = rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("zone"), + subscripts: vec![], + span: test_span(), + }), + field: "spawn-0".to_string(), + span: test_span(), + }), + field: "4".to_string(), + span: test_span(), + }), + field: "3-7048a72798".to_string(), + span: test_span(), + }; + + let expected = Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("spawn-0.4.3-7048a72798".to_string()), + span: test_span(), + }); + assert_eq!( + recover_string_literal_from_invalid_component_expr(&invalid_path, "zone"), + expected + ); + } } diff --git a/crates/rumoca-phase-flatten/src/vcg.rs b/crates/rumoca-phase-flatten/src/vcg.rs index 520ccfb55..3fc4f2765 100644 --- a/crates/rumoca-phase-flatten/src/vcg.rs +++ b/crates/rumoca-phase-flatten/src/vcg.rs @@ -49,7 +49,6 @@ pub(crate) fn pre_collect_vcg_data( branches: Vec::new(), potential_roots: Vec::new(), }; - for (_def_id, class_data) in &overlay.classes { if crate::is_in_disabled_component(&class_data.qualified_name, &overlay.disabled_components) { @@ -583,13 +582,16 @@ pub(crate) fn build_vcg( for component in &components { let root = select_root(component, definite_roots, potential_roots); let has_definite_root = component.iter().any(|n| definite_roots.contains(*n)); + let has_potential_root = component + .iter() + .any(|node| potential_roots.iter().any(|(path, _)| path == node)); for &node in component { let node_is_root = definite_roots.contains(node) || (!has_definite_root && node == root); is_root_map.insert(node.to_string(), node_is_root); // rooted(N) = true iff N is NOT the root and the component has a root - let node_rooted = !node_is_root && (has_definite_root || !potential_roots.is_empty()); + let node_rooted = !node_is_root && (has_definite_root || has_potential_root); rooted_map.insert(node.to_string(), node_rooted); } } @@ -1062,6 +1064,47 @@ mod tests { assert_eq!(vcg.rooted.get("b.R"), Some(&true)); } + #[test] + fn test_build_vcg_rooted_is_component_local() { + let branches = Vec::new(); + let optional_edges = vec![("a.R".to_string(), "b.R".to_string())]; + let definite_roots: FxHashSet = FxHashSet::default(); + let potential_roots = vec![("other.R".to_string(), 1)]; + + let vcg = build_vcg( + &definite_roots, + &potential_roots, + &branches, + &optional_edges, + ); + + assert_eq!(vcg.rooted.get("a.R"), Some(&false)); + assert_eq!(vcg.rooted.get("b.R"), Some(&false)); + assert_eq!(vcg.is_root.get("other.R"), Some(&true)); + } + + #[test] + fn test_build_vcg_selects_lowest_priority_potential_root() { + let branches = Vec::new(); + let optional_edges = vec![ + ("high.R".to_string(), "mid.R".to_string()), + ("mid.R".to_string(), "low.R".to_string()), + ]; + let definite_roots: FxHashSet = FxHashSet::default(); + let potential_roots = vec![("high.R".to_string(), 256), ("low.R".to_string(), 10)]; + + let vcg = build_vcg( + &definite_roots, + &potential_roots, + &branches, + &optional_edges, + ); + + assert_eq!(vcg.is_root.get("low.R"), Some(&true)); + assert_eq!(vcg.rooted.get("high.R"), Some(&true)); + assert_eq!(vcg.rooted.get("mid.R"), Some(&true)); + } + #[test] fn test_component_oc_record_scalar_count_uses_max_node_size() { let mut flat = flat::Model::default(); diff --git a/crates/rumoca-phase-flatten/src/when_equations.rs b/crates/rumoca-phase-flatten/src/when_equations.rs index 278e7cf01..5e30a7640 100644 --- a/crates/rumoca-phase-flatten/src/when_equations.rs +++ b/crates/rumoca-phase-flatten/src/when_equations.rs @@ -20,8 +20,12 @@ use crate::{Context, qualify_expression_imports_with_def_map_ctx}; /// Flatten a when-equation to a list of WhenClauses. /// -/// Each flat::WhenClause represents one "when" or "elsewhen" branch with its condition -/// and the discrete equations that should be executed when the condition becomes true. +/// A source `when` without `elsewhen` lowers to one flat::WhenClause. A +/// non-structural `when`/`elsewhen` chain lowers to one flat::WhenClause whose +/// condition is the vector of branch conditions and whose body is a conditional +/// branch list keyed by each branch activation. This preserves MLS §8.3.5 +/// priority semantics while still letting DAE lowering generate one solved +/// assignment per target. pub(crate) fn flatten_when_equation( ctx: &Context, inst_eq: &ast::InstanceEquation, @@ -58,9 +62,78 @@ pub(crate) fn flatten_when_blocks( } validate_when_branch_targets(ctx, blocks, &clauses, prefix, span)?; + if should_merge_elsewhen_chain(ctx, blocks, prefix, &clauses) { + return Ok(vec![merge_elsewhen_chain(clauses, span)]); + } Ok(clauses) } +fn should_merge_elsewhen_chain( + ctx: &Context, + blocks: &[ast::EquationBlock], + prefix: &ast::QualifiedName, + clauses: &[flat::WhenClause], +) -> bool { + if clauses.len() <= 1 { + return false; + } + !blocks + .iter() + .all(|block| crate::boolean_eval::is_structural_expression(ctx, &block.cond, prefix)) +} + +fn merge_elsewhen_chain( + clauses: Vec, + span: rumoca_core::Span, +) -> flat::WhenClause { + let condition = rumoca_core::Expression::Array { + elements: clauses + .iter() + .map(|clause| clause.condition.clone()) + .collect(), + is_matrix: false, + span, + }; + let branches = clauses + .into_iter() + .map(|clause| { + ( + elsewhen_branch_activation_condition(&clause.condition, clause.span), + clause.equations, + ) + }) + .collect(); + let mut merged = flat::WhenClause::new(condition, span); + merged.add_equation(flat::WhenEquation::conditional( + branches, + Vec::new(), + span, + "when/elsewhen branch chain", + )); + merged +} + +fn elsewhen_branch_activation_condition( + condition: &rumoca_core::Expression, + span: rumoca_core::Span, +) -> rumoca_core::Expression { + if matches!( + condition, + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Initial, + args, + .. + } if args.is_empty() + ) { + return condition.clone(); + } + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Edge, + args: vec![condition.clone()], + span, + } +} + /// Flatten a single when/elsewhen block to a flat::WhenClause. pub(crate) fn flatten_when_block( ctx: &Context, @@ -635,7 +708,8 @@ fn extract_assignment_target( #[cfg(test)] mod tests { use super::{ - extract_assignment_target, flatten_when_for_equation, is_known_streams_side_effect_call, + elsewhen_branch_activation_condition, extract_assignment_target, flatten_when_for_equation, + is_known_streams_side_effect_call, merge_elsewhen_chain, }; use crate::errors::FlattenError; use rumoca_ir_ast as ast; @@ -695,6 +769,30 @@ mod tests { ast::Expression::ComponentReference(comp_ref(name)) } + fn bool_expr(value: bool) -> rumoca_core::Expression { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Boolean(value), + span: test_span(), + } + } + + fn flat_var_ref(name: &str) -> rumoca_core::Expression { + rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new(name).into(), + subscripts: vec![], + span: test_span(), + } + } + + fn flat_assignment(target: &str, value: bool) -> flat::WhenEquation { + flat::WhenEquation::assign( + rumoca_core::VarName::new(target), + bool_expr(value), + test_span(), + format!("assign {target}"), + ) + } + fn for_index(name: &str, start: i64, end: i64) -> ast::ForIndex { ast::ForIndex { ident: token(name), @@ -798,4 +896,49 @@ mod tests { assert!(targets.contains("y[2,1]")); assert!(targets.contains("y[2,2]")); } + + #[test] + fn elsewhen_branch_activation_wraps_non_initial_conditions_in_edge() { + let condition = flat_var_ref("u"); + let activation = elsewhen_branch_activation_condition(&condition, test_span()); + + let rumoca_core::Expression::BuiltinCall { function, args, .. } = activation else { + panic!("expected edge(...) activation"); + }; + assert_eq!(function, rumoca_core::BuiltinFunction::Edge); + assert_eq!(args, vec![condition]); + } + + #[test] + fn merge_elsewhen_chain_preserves_one_clause_with_vector_guard() { + let mut first = flat::WhenClause::new(flat_var_ref("u"), test_span()); + first.add_equation(flat_assignment("y", true)); + let mut second = flat::WhenClause::new(flat_var_ref("v"), test_span()); + second.add_equation(flat_assignment("y", false)); + + let merged = merge_elsewhen_chain(vec![first, second], test_span()); + + let rumoca_core::Expression::Array { elements, .. } = &merged.condition else { + panic!("merged elsewhen chain should use vector when condition"); + }; + assert_eq!(elements.len(), 2); + assert_eq!(merged.equations.len(), 1); + let flat::WhenEquation::Conditional { branches, .. } = &merged.equations[0] else { + panic!("merged elsewhen chain should lower into one conditional when equation"); + }; + assert_eq!(branches.len(), 2); + for (condition, branch_eqs) in branches { + assert!( + matches!( + condition, + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Edge, + .. + } + ), + "branch conditions should be activation edges: {condition:?}" + ); + assert_eq!(branch_eqs.len(), 1); + } + } } diff --git a/crates/rumoca-phase-flatten/src/zero_sized_arrays.rs b/crates/rumoca-phase-flatten/src/zero_sized_arrays.rs index c8164a615..c22b3c6db 100644 --- a/crates/rumoca-phase-flatten/src/zero_sized_arrays.rs +++ b/crates/rumoca-phase-flatten/src/zero_sized_arrays.rs @@ -28,15 +28,22 @@ pub(crate) fn materialize_referenced_zero_sized_array_variables( if flat.variables.contains_key(&var_name) { continue; } - let source_span = name.span().ok_or_else(|| { + let source_span = zero_sized_array_source_span(&name, ctx).ok_or_else(|| { FlattenError::missing_source_context(format!( "zero-sized array reference `{}` is missing source provenance", name.as_str() )) })?; + let component_ref = name.component_ref().cloned().or_else(|| { + Some(rumoca_core::ComponentReference::from_flat_segments( + name.as_str(), + source_span, + None, + )) + }); let variable = flat::Variable { name: var_name.clone(), - component_ref: name.component_ref().cloned(), + component_ref, source_span, dims, variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), @@ -118,3 +125,20 @@ fn zero_sized_array_dims_for_ref(name: &rumoca_core::Reference, ctx: &Context) - })?; (!dims.is_empty() && dims.iter().any(|dim| *dim <= 0)).then(|| dims.clone()) } + +fn zero_sized_array_source_span( + name: &rumoca_core::Reference, + ctx: &Context, +) -> Option { + name.span().or_else(|| { + ctx.array_dimension_spans + .get(name.as_str()) + .copied() + .or_else(|| { + name.target_def_id() + .and_then(|def_id| ctx.target_def_names.get(&def_id)) + .and_then(|target_name| ctx.array_dimension_spans.get(target_name)) + .copied() + }) + }) +} diff --git a/crates/rumoca-phase-instantiate/src/array_expansion.rs b/crates/rumoca-phase-instantiate/src/array_expansion.rs index 62125e692..7e4e4b0af 100644 --- a/crates/rumoca-phase-instantiate/src/array_expansion.rs +++ b/crates/rumoca-phase-instantiate/src/array_expansion.rs @@ -1,3 +1,6 @@ +// SPEC_0021 file-size exception: array expansion still coordinates component +// selection, modifier projection, and array-comprehension indexing. split plan: +// move subscript projection and non-component binding indexing into submodules. use super::inheritance::resolve_effective_components_for_eval; use super::instantiate_component; use super::source_scope::component_declaration_source_scope; @@ -38,6 +41,7 @@ pub(super) fn expand_array_component( overlay.array_parent_dims.insert(parent_path, dims.to_vec()); let indices = super::generate_array_indices(dims); + let active_nested_mod_keys = active_nested_modifier_keys(ctx.mod_env(), name); // Extract the original binding for indexing. Check active component modifier first, // then comp.binding, then fall back to comp.start for modification-only declarations. @@ -67,6 +71,12 @@ pub(super) fn expand_array_component( .as_ref() .and_then(|mv| mv.source_scope.clone()) .or_else(|| component_declaration_source_scope(ctx, comp)); + let modifier_state = ArrayElementModifierState { + binding_qn, + mod_env_binding, + binding_source_scope, + active_nested_mod_keys, + }; // MLS §7.2.5: Pre-resolve non-`each` modifications that reference array values. // Resolve once so each element can be indexed from the resolved array. @@ -120,51 +130,179 @@ pub(super) fn expand_array_component( &idx, )?; - // Ensure scalar element instantiation sees the indexed binding (MLS §10.1). - // Without this scoped override, a parent unindexed modifier entry for this - // component name would overwrite the per-element indexed binding. - let previous_binding = ctx.mod_env().active.get(&binding_qn).cloned(); - if let Some(binding_expr) = &scalar_comp.binding - && mod_env_binding.is_some() - { - ctx.mod_env_mut().active.insert( - binding_qn.clone(), - array_element_binding_modification( - scope, - binding_expr, - &idx, - mod_env_binding.as_ref(), - binding_source_scope.clone(), - )?, - ); - } - - ctx.push_path_part(name, idx.clone()); - let inst_result = instantiate_component( - scope.tree, + instantiate_scalar_array_element( + scope, + name, &scalar_comp, + &idx, ctx, overlay, + &modifier_state, + )?; + } + + Ok(()) +} + +struct ArrayElementModifierState { + binding_qn: ast::QualifiedName, + mod_env_binding: Option, + binding_source_scope: Option, + active_nested_mod_keys: Vec, +} + +fn instantiate_scalar_array_element( + scope: &ArrayExpansionScope<'_>, + name: &str, + scalar_comp: &ast::Component, + idx: &[i64], + ctx: &mut InstantiateContext, + overlay: &mut rumoca_ir_ast::InstanceOverlay, + modifier_state: &ArrayElementModifierState, +) -> InstantiateResult<()> { + let previous_binding = + install_array_element_binding_override(scope, scalar_comp, idx, ctx, modifier_state)?; + let previous_nested_mods = distribute_active_nested_mods_for_element( + scope, + ctx, + &modifier_state.active_nested_mod_keys, + idx, + )?; + + ctx.push_path_part(name, idx.to_vec()); + let inst_result = instantiate_component( + scope.tree, + scalar_comp, + ctx, + overlay, + scope.effective_components, + scope.type_overrides, + scope.imports, + ); + ctx.pop_path(); + + restore_array_element_binding(ctx, modifier_state, previous_binding); + restore_active_nested_mods(ctx, previous_nested_mods); + inst_result +} + +fn install_array_element_binding_override( + scope: &ArrayExpansionScope<'_>, + scalar_comp: &ast::Component, + idx: &[i64], + ctx: &mut InstantiateContext, + modifier_state: &ArrayElementModifierState, +) -> InstantiateResult> { + let previous_binding = ctx + .mod_env() + .active + .get(&modifier_state.binding_qn) + .cloned(); + if let Some(binding_expr) = &scalar_comp.binding + && modifier_state.mod_env_binding.is_some() + { + ctx.mod_env_mut().active.insert( + modifier_state.binding_qn.clone(), + array_element_binding_modification( + scope, + binding_expr, + idx, + modifier_state.mod_env_binding.as_ref(), + modifier_state.binding_source_scope.clone(), + )?, + ); + } + Ok(previous_binding) +} + +fn restore_array_element_binding( + ctx: &mut InstantiateContext, + modifier_state: &ArrayElementModifierState, + previous_binding: Option, +) { + match previous_binding { + Some(prev) => { + ctx.mod_env_mut() + .active + .insert(modifier_state.binding_qn.clone(), prev); + } + None => { + ctx.mod_env_mut() + .active + .shift_remove(&modifier_state.binding_qn); + } + } +} + +fn active_nested_modifier_keys( + mod_env: &ast::ModificationEnvironment, + component_name: &str, +) -> Vec { + mod_env + .active + .keys() + .filter(|key| key.strip_prefix(component_name).is_some()) + .cloned() + .collect() +} + +fn distribute_active_nested_mods_for_element( + scope: &ArrayExpansionScope<'_>, + ctx: &mut InstantiateContext, + keys: &[ast::QualifiedName], + idx: &[i64], +) -> InstantiateResult)>> { + let mut previous = Vec::with_capacity(keys.len()); + for key in keys { + let Some(existing) = ctx.mod_env().active.get(key).cloned() else { + previous.push((key.clone(), None)); + continue; + }; + if existing.each { + continue; + } + previous.push((key.clone(), Some(existing.clone()))); + let value = index_binding_for_element( + scope.tree, scope.effective_components, - scope.type_overrides, - scope.imports, + &existing.value, + idx, + )?; + let source = existing + .source + .as_ref() + .map(|source| { + index_binding_for_element(scope.tree, scope.effective_components, source, idx) + }) + .transpose()?; + ctx.mod_env_mut().active.insert( + key.clone(), + ast::ModificationValue { + value, + source, + source_scope: existing.source_scope, + each: existing.each, + final_: existing.final_, + }, ); - ctx.pop_path(); + } + Ok(previous) +} - // Restore parent binding scope for subsequent elements/siblings. - match previous_binding { - Some(prev) => { - ctx.mod_env_mut().active.insert(binding_qn.clone(), prev); +fn restore_active_nested_mods( + ctx: &mut InstantiateContext, + previous: Vec<(ast::QualifiedName, Option)>, +) { + for (key, value) in previous { + match value { + Some(value) => { + ctx.mod_env_mut().active.insert(key, value); } None => { - ctx.mod_env_mut().active.shift_remove(&binding_qn); + ctx.mod_env_mut().active.shift_remove(&key); } } - - inst_result?; } - - Ok(()) } fn indexed_array_component_start( @@ -231,7 +369,9 @@ pub(super) fn pre_resolve_array_modifications( // form so per-element indexing can preserve owner scope (`R[1]`). // Eager resolution here can over-collapse to outer-source expressions // like `RRef.d`, losing the declaring component context. - if matches!(expr, ast::Expression::ComponentReference(_)) { + if matches!(expr, ast::Expression::ComponentReference(_)) + || has_explicit_subscripted_component_ref(expr) + { continue; } let val = resolve_mod_to_array(expr, mod_env, effective_components, tree); @@ -277,6 +417,13 @@ fn distribute_component_ref_mods_for_element( continue; } + if let Some(indexed) = + index_nested_modification_for_element(tree, parent_components, expr, indices)? + { + scalar_comp.modifications.insert(name.clone(), indexed); + continue; + } + let indexed = index_binding_for_element(tree, parent_components, expr, indices)?; if !matches!(indexed, ast::Expression::ArrayIndex { .. }) { scalar_comp.modifications.insert(name.clone(), indexed); @@ -285,6 +432,44 @@ fn distribute_component_ref_mods_for_element( Ok(()) } +fn index_nested_modification_for_element( + tree: &ast::ClassTree, + parent_components: &IndexMap, + expr: &ast::Expression, + indices: &[i64], +) -> InstantiateResult> { + let mut indexed = expr.clone(); + match &mut indexed { + ast::Expression::ClassModification { + modifications, + each_flags, + .. + } => { + for (position, nested) in modifications.iter_mut().enumerate() { + if !each_flags.get(position).copied().unwrap_or(false) + && let Some(value) = index_nested_modification_for_element( + tree, + parent_components, + nested, + indices, + )? + { + *nested = value; + } + } + } + ast::Expression::Modification { value, .. } => { + if let Some(indexed_value) = + index_array_expression_for_element(tree, parent_components, value, indices)? + { + *value = Arc::new(indexed_value); + } + } + _ => return Ok(None), + } + Ok(Some(indexed)) +} + /// Resolve a modification expression to its value, handling component references. /// /// For component references like `R` that refer to array parameters in the parent @@ -623,6 +808,18 @@ fn make_real_lit(value: f64, span: rumoca_core::Span) -> ast::Expression { } } +fn has_explicit_subscripted_component_ref(expr: &ast::Expression) -> bool { + match expr { + ast::Expression::ComponentReference(cref) => { + cref.parts.iter().any(|part| part.subs.is_some()) + } + ast::Expression::Parenthesized { inner, .. } => { + has_explicit_subscripted_component_ref(inner) + } + _ => false, + } +} + /// Evaluate an array comprehension `{expr for j in start:end}` to a concrete array. /// /// MLS §11.1.2.1: For simple comprehensions like `{j for j in 1:m}`, @@ -775,12 +972,16 @@ impl LoopIndexSubstituter<'_> { /// given 1-based index. For non-array values, returns None (no change needed). fn index_array_modification(expr: &ast::Expression, indices: &[i64]) -> Option { match expr { - ast::Expression::Array { elements, .. } => { + ast::Expression::Array { + elements, + is_matrix, + .. + } => { let (&first_idx, remaining) = indices.split_first()?; let idx = first_idx.checked_sub(1)? as usize; // Convert 1-based to 0-based let selected = elements.get(idx)?; if remaining.is_empty() { - Some(selected.clone()) + Some(normalize_distributed_matrix_row(selected, *is_matrix)) } else { index_array_modification(selected, remaining) } @@ -790,6 +991,27 @@ fn index_array_modification(expr: &ast::Expression, indices: &[i64]) -> Option ast::Expression { + if !parent_is_matrix { + return selected.clone(); + } + + // MLS §7.2.5 + §10.4: selecting one element of a matrix-valued + // non-`each` modifier distributes the row value to the scalar component. + // The selected row is a 1-D array value, not a single-row matrix. + match selected { + ast::Expression::Array { elements, span, .. } => ast::Expression::Array { + elements: elements.clone(), + is_matrix: false, + span: *span, + }, + _ => selected.clone(), + } +} + /// Index a binding expression for an array element. /// /// When an array component has a binding (e.g., `v1[m] = plug1.pin.v`), each expanded @@ -840,9 +1062,14 @@ pub(super) fn index_binding_for_element( }); }; let mut new_ref = cref.clone(); + let subs = new_ref.parts[pos] + .subs + .as_ref() + .and_then(|existing| project_existing_subscripts_for_element(existing, indices)) + .unwrap_or_else(make_subscripts); new_ref.parts[pos] = ast::ComponentRefPart { ident: new_ref.parts[pos].ident.clone(), - subs: Some(make_subscripts()), + subs: Some(subs), }; return Ok(ast::Expression::ComponentReference(new_ref)); } @@ -859,6 +1086,64 @@ pub(super) fn index_binding_for_element( }) } +fn project_existing_subscripts_for_element( + existing: &[ast::Subscript], + indices: &[i64], +) -> Option> { + if existing.len() != indices.len() { + return None; + } + + existing + .iter() + .zip(indices.iter().copied()) + .map(|(sub, index)| project_existing_subscript_for_element(sub, index)) + .collect() +} + +fn project_existing_subscript_for_element( + sub: &ast::Subscript, + index: i64, +) -> Option { + match sub { + ast::Subscript::Expression(ast::Expression::Range { + start, step, span, .. + }) => { + let start = integer_literal_value(start)?; + let step = match step.as_deref() { + Some(expr) => integer_literal_value(expr)?, + None => 1, + }; + let selected = start + (index - 1) * step; + Some(ast::Subscript::Expression(make_int_expr(selected, *span))) + } + ast::Subscript::Expression(ast::Expression::Array { elements, .. }) => { + let idx = index.checked_sub(1)? as usize; + Some(ast::Subscript::Expression(elements.get(idx)?.clone())) + } + ast::Subscript::Expression(expr) => Some(ast::Subscript::Expression(expr.clone())), + ast::Subscript::Empty | ast::Subscript::Range { .. } => None, + } +} + +fn integer_literal_value(expr: &ast::Expression) -> Option { + let ast::Expression::Terminal { token, .. } = expr else { + return None; + }; + token.text.as_ref().parse().ok() +} + +fn make_int_expr(value: i64, span: rumoca_core::Span) -> ast::Expression { + ast::Expression::Terminal { + terminal_type: ast::TerminalType::UnsignedInteger, + token: rumoca_core::Token { + text: value.to_string().into(), + ..rumoca_core::Token::default() + }, + span, + } +} + fn index_non_component_reference_binding( binding: &ast::Expression, indices: &[i64], @@ -1337,6 +1622,42 @@ mod tests { ); } + #[test] + fn test_index_binding_for_element_projects_explicit_range_slice() { + let mut parent_components = IndexMap::default(); + parent_components.insert( + "root".to_string(), + ast::Component { + name: "root".to_string(), + shape: vec![20], + ..ast::Component::empty_with_span(test_span()) + }, + ); + let mut binding = make_comp_ref_expr(&["root"]); + if let ast::Expression::ComponentReference(cref) = &mut binding { + cref.parts[0].subs = Some(vec![ast::Subscript::Expression(make_range_expr(11, 15))]); + } + + let indexed = index_binding_for_element( + &ast::ClassTree::default(), + &parent_components, + &binding, + &[2], + ) + .expect("range slice binding should project to an element subscript"); + + let ast::Expression::ComponentReference(cref) = indexed else { + panic!("expected projected component reference"); + }; + let Some(subs) = &cref.parts[0].subs else { + panic!("projected array part should retain a scalar subscript"); + }; + let ast::Subscript::Expression(ast::Expression::Terminal { token, .. }) = &subs[0] else { + panic!("range projection should produce an integer subscript"); + }; + assert_eq!(token.text.as_ref(), "12"); + } + #[test] fn test_index_binding_for_element_no_array_part_uses_array_index_fallback() { let binding = make_comp_ref_expr(&["a", "b", "c"]); @@ -1616,6 +1937,58 @@ mod tests { assert_eq!(token.text.as_ref(), "2"); } + #[test] + fn test_distribute_mods_for_element_projects_matrix_row_as_vector() { + let mut comp = ast::Component::empty_with_span(test_span()); + comp.modifications.insert( + "VolFloCur".to_string(), + ast::Expression::Array { + elements: vec![ + ast::Expression::Array { + elements: vec![make_int_expr(1), make_int_expr(2), make_int_expr(3)], + is_matrix: true, + span: rumoca_core::Span::DUMMY, + }, + ast::Expression::Array { + elements: vec![make_int_expr(4), make_int_expr(5), make_int_expr(6)], + is_matrix: true, + span: rumoca_core::Span::DUMMY, + }, + ], + is_matrix: true, + span: rumoca_core::Span::DUMMY, + }, + ); + + let resolved_mods = pre_resolve_array_modifications( + &comp, + &rumoca_ir_ast::ModificationEnvironment::default(), + &IndexMap::default(), + &ast::ClassTree::default(), + ); + + let mut scalar_comp = comp.clone(); + distribute_mods_for_element(&mut scalar_comp, &resolved_mods, &[2]); + let distributed = scalar_comp + .modifications + .get("VolFloCur") + .expect("missing distributed modifier"); + + let ast::Expression::Array { + elements, + is_matrix, + .. + } = distributed + else { + panic!("distributed row should remain an array"); + }; + assert!( + !is_matrix, + "distributed matrix row should be a 1-D array, not a single-row matrix" + ); + assert_eq!(elements.len(), 3); + } + #[test] fn test_index_binding_for_element_indexes_nested_array_part_via_type_walk() { let stack_data_id = DefId::new(100); diff --git a/crates/rumoca-phase-instantiate/src/attributes.rs b/crates/rumoca-phase-instantiate/src/attributes.rs index 64050d0e0..77dfdadf4 100644 --- a/crates/rumoca-phase-instantiate/src/attributes.rs +++ b/crates/rumoca-phase-instantiate/src/attributes.rs @@ -12,15 +12,47 @@ pub(super) struct ComponentAttrsAndBinding { pub(super) binding_from_modification: bool, } +#[cfg(test)] pub(super) fn extract_component_attrs_and_binding( comp: &ast::Component, mod_env: &ast::ModificationEnvironment, + instance_name: &str, eval_ctx: &InstantiateEvalCtx<'_>, imports: &[(String, String)], +) -> InstantiateResult { + extract_component_attrs_and_binding_in_scope( + comp, + mod_env, + instance_name, + eval_ctx, + imports, + None, + ) +} + +pub(super) fn extract_component_attrs_and_binding_in_scope( + comp: &ast::Component, + mod_env: &ast::ModificationEnvironment, + instance_name: &str, + eval_ctx: &InstantiateEvalCtx<'_>, + imports: &[(String, String)], + declaration_scope: Option<&ast::QualifiedName>, ) -> InstantiateResult { // Pass component name so mod_env can be checked for outer modifications. - let mut attrs = extract_attributes(comp, mod_env, &comp.name, eval_ctx, imports)?; - let (binding, binding_from_modification, binding_source_scope) = extract_binding(comp, mod_env); + let mut attrs = extract_attributes_in_scope( + comp, + mod_env, + &comp.name, + instance_name, + eval_ctx, + imports, + declaration_scope, + )?; + let (binding, binding_from_modification, mut binding_source_scope) = + extract_binding(comp, mod_env); + if binding_from_modification && binding_source_scope.is_none() { + binding_source_scope = declaration_scope.cloned(); + } let binding_source = if binding_from_modification { let binding_path = ast::QualifiedName::from_ident(&comp.name); mod_env @@ -290,6 +322,24 @@ fn insert_final_type_attribute_name( } } +fn extract_numeric_attr_source_from_modifications( + comp: &ast::Component, + attr_name: &str, +) -> Option { + comp.source_modifications + .iter() + .find_map(|expr| match expr { + ast::Expression::Modification { target, value, .. } => { + let target_name = target.parts.last()?.ident.text.as_ref(); + (target_name == attr_name).then(|| value.as_ref().clone()) + } + ast::Expression::NamedArgument { name, value, .. } => { + (name.text.as_ref() == attr_name).then(|| value.as_ref().clone()) + } + _ => None, + }) +} + fn type_attribute_modification_name(expr: &ast::Expression) -> Option { match expr { ast::Expression::Modification { target, .. } @@ -321,12 +371,34 @@ fn should_promote_binding_to_start( /// MLS §7.2: The modification environment is checked for overriding /// modifications from outer scopes. Outer modifications override inner ones per /// MLS §7.2.4. +#[cfg(test)] pub(super) fn extract_attributes( comp: &ast::Component, mod_env: &ast::ModificationEnvironment, comp_name: &str, + instance_name: &str, + eval_ctx: &InstantiateEvalCtx<'_>, + imports: &[(String, String)], +) -> InstantiateResult { + extract_attributes_in_scope( + comp, + mod_env, + comp_name, + instance_name, + eval_ctx, + imports, + None, + ) +} + +pub(super) fn extract_attributes_in_scope( + comp: &ast::Component, + mod_env: &ast::ModificationEnvironment, + comp_name: &str, + _instance_name: &str, eval_ctx: &InstantiateEvalCtx<'_>, imports: &[(String, String)], + declaration_scope: Option<&ast::QualifiedName>, ) -> InstantiateResult { let mut source_scopes = IndexMap::default(); let start_path = ast::QualifiedName::from_ident(comp_name).child("start"); @@ -334,7 +406,7 @@ pub(super) fn extract_attributes( if let Some(scope) = value.source_scope.clone() { source_scopes.insert("start".to_string(), scope); } - value.value.clone() + value.source.clone().unwrap_or_else(|| value.value.clone()) }); let mut attr_from_mod_env = |attr_name: &str| { let path = ast::QualifiedName::from_ident(comp_name).child(attr_name); @@ -342,7 +414,7 @@ pub(super) fn extract_attributes( if let Some(scope) = value.source_scope.clone() { source_scopes.insert(attr_name.to_string(), scope); } - Some(value.value.clone()) + Some(value.source.clone().unwrap_or_else(|| value.value.clone())) }; let state_select_path = ast::QualifiedName::from_ident(comp_name).child("stateSelect"); @@ -353,6 +425,7 @@ pub(super) fn extract_attributes( eval_ctx, imports, value.source_scope.as_ref(), + declaration_scope, )?), None => None, }; @@ -378,20 +451,39 @@ pub(super) fn extract_attributes( for (name, value) in &comp.modifications { match name.as_str() { "start" if attrs.start.is_none() => { - attrs.start = Some(value.clone()); + attrs.start = Some( + extract_numeric_attr_source_from_modifications(comp, "start") + .unwrap_or_else(|| value.clone()), + ); attrs.start_is_explicit = true; } "fixed" if attrs.fixed.is_none() => attrs.fixed = expr_to_bool(value), - "min" if attrs.min.is_none() => attrs.min = Some(value.clone()), - "max" if attrs.max.is_none() => attrs.max = Some(value.clone()), - "nominal" if attrs.nominal.is_none() => attrs.nominal = Some(value.clone()), + "min" if attrs.min.is_none() => { + attrs.min = Some( + extract_numeric_attr_source_from_modifications(comp, "min") + .unwrap_or_else(|| value.clone()), + ) + } + "max" if attrs.max.is_none() => { + attrs.max = Some( + extract_numeric_attr_source_from_modifications(comp, "max") + .unwrap_or_else(|| value.clone()), + ) + } + "nominal" if attrs.nominal.is_none() => { + attrs.nominal = Some( + extract_numeric_attr_source_from_modifications(comp, "nominal") + .unwrap_or_else(|| value.clone()), + ) + } "quantity" if attrs.quantity.is_none() => attrs.quantity = expr_to_string(value), "unit" if attrs.unit.is_none() => attrs.unit = expr_to_string(value), "displayUnit" if attrs.display_unit.is_none() => { attrs.display_unit = expr_to_string(value) } "stateSelect" if !has_outer_state_select => { - attrs.state_select = parse_required_state_select(value, eval_ctx, imports, None)? + attrs.state_select = + parse_required_state_select(value, eval_ctx, imports, None, declaration_scope)? } _ => {} } @@ -410,15 +502,25 @@ fn parse_required_state_select( eval_ctx: &InstantiateEvalCtx<'_>, imports: &[(String, String)], source_scope: Option<&ast::QualifiedName>, + declaration_scope: Option<&ast::QualifiedName>, ) -> InstantiateResult { parse_state_select(value) .or_else(|| eval_state_select_expr_with_source_scope(eval_ctx, value, source_scope)) + .or_else(|| eval_state_select_expr_with_source_scope(eval_ctx, value, declaration_scope)) .or_else(|| { // Enclosing-scope constants (MLS §5.3.2) appear unqualified in // declaration-side attributes; qualify them through the package // constant aliases and retry before failing. let qualified = crate::dims::qualify_shape_expr_imports(value, imports); - eval_state_select_expr_with_source_scope(eval_ctx, &qualified, source_scope) + eval_state_select_expr_with_source_scope(eval_ctx, &qualified, source_scope).or_else( + || { + eval_state_select_expr_with_source_scope( + eval_ctx, + &qualified, + declaration_scope, + ) + }, + ) }) .ok_or_else(|| { Box::new(InstantiateError::InvalidTypeAttribute { diff --git a/crates/rumoca-phase-instantiate/src/binding_source.rs b/crates/rumoca-phase-instantiate/src/binding_source.rs new file mode 100644 index 000000000..f3d9b4a4d --- /dev/null +++ b/crates/rumoca-phase-instantiate/src/binding_source.rs @@ -0,0 +1,55 @@ +use rumoca_core::Variability; +use rumoca_ir_ast as ast; +use rumoca_ir_ast::AstIndexMap as IndexMap; + +pub(crate) fn declaration_binding_source_for_flattening( + comp: &ast::Component, + original: &ast::Expression, + resolved: &ast::Expression, + effective_components: &IndexMap, + mod_env: &ast::ModificationEnvironment, +) -> Option { + if original == resolved { + return None; + } + let ast::Expression::ComponentReference(comp_ref) = original else { + return None; + }; + if comp_ref.parts.len() >= 2 + || single_part_source_ref_is_modified_parameter_sibling( + comp_ref, + comp, + effective_components, + mod_env, + ) + { + return Some(original.clone()); + } + None +} + +fn single_part_source_ref_is_modified_parameter_sibling( + comp_ref: &ast::ComponentReference, + _target_component: &ast::Component, + effective_components: &IndexMap, + mod_env: &ast::ModificationEnvironment, +) -> bool { + let [part] = comp_ref.parts.as_slice() else { + return false; + }; + let name = part.ident.text.as_ref(); + let Some(source_component) = effective_components.get(name) else { + return false; + }; + if !parameter_like_component(source_component) { + return false; + } + mod_env.get(&ast::QualifiedName::from_ident(name)).is_some() +} + +fn parameter_like_component(component: &ast::Component) -> bool { + matches!( + component.variability, + Variability::Parameter(_) | Variability::Constant(_) + ) || component.is_structural +} diff --git a/crates/rumoca-phase-instantiate/src/component_loop.rs b/crates/rumoca-phase-instantiate/src/component_loop.rs index 9072c8edf..c67299499 100644 --- a/crates/rumoca-phase-instantiate/src/component_loop.rs +++ b/crates/rumoca-phase-instantiate/src/component_loop.rs @@ -1,6 +1,7 @@ //! The per-class component instantiation loop and its alias-set plumbing. use super::*; +use rumoca_core::scoped_component_path_candidates; /// Alias sets used while instantiating a class's components. #[derive(Clone, Copy)] @@ -55,6 +56,17 @@ pub(super) fn instantiate_effective_components( tree, resolve_effective_components_for_eval, ); + let dims = if !qualified_shape_expr.is_empty() && dims.as_ref().is_some_and(Vec::is_empty) { + evaluate_array_dimensions_with_known_params( + &comp.shape, + &qualified_shape_expr, + &ctx.known_int_params, + &ctx.current_path(), + ) + .or(dims) + } else { + dims + }; if let Some(dims) = dims.as_ref() && dims.contains(&0) { @@ -113,6 +125,100 @@ pub(super) fn component_type_id( } } +fn evaluate_array_dimensions_with_known_params( + shape: &[usize], + shape_expr: &[rumoca_ir_ast::Subscript], + known_int_params: &rustc_hash::FxHashMap, + scope: &rumoca_ir_ast::QualifiedName, +) -> Option> { + if !shape_expr.is_empty() { + let mut dims = Vec::with_capacity(shape_expr.len()); + for subscript in shape_expr { + let rumoca_ir_ast::Subscript::Expression(expr) = subscript else { + return None; + }; + let dim = eval_known_integer_expr(expr, known_int_params, scope)?; + if dim < 0 { + return None; + } + dims.push(dim); + } + return Some(dims); + } + + (!shape.is_empty()).then(|| shape.iter().map(|&dim| dim as i64).collect()) +} + +fn eval_known_integer_expr( + expr: &rumoca_ir_ast::Expression, + known_int_params: &rustc_hash::FxHashMap, + scope: &rumoca_ir_ast::QualifiedName, +) -> Option { + match expr { + rumoca_ir_ast::Expression::Terminal { + terminal_type: rumoca_ir_ast::TerminalType::UnsignedInteger, + token, + .. + } => token.text.parse().ok(), + rumoca_ir_ast::Expression::ComponentReference(component_ref) + if !component_ref.parts.is_empty() + && component_ref.parts.iter().all(|part| part.subs.is_none()) => + { + let name = component_ref + .parts + .iter() + .map(|part| part.ident.text.as_ref()) + .collect::>() + .join("."); + let name = rumoca_core::ComponentPath::from_flat_path(&name); + let scope = scope.to_component_path(); + for candidate in scoped_component_path_candidates(&name, &scope) { + if let Some(value) = known_int_params.get(candidate.as_str()) { + return Some(*value); + } + } + None + } + rumoca_ir_ast::Expression::Binary { op, lhs, rhs, .. } => { + let lhs = eval_known_integer_expr(lhs, known_int_params, scope)?; + let rhs = eval_known_integer_expr(rhs, known_int_params, scope)?; + eval_known_integer_binary(op, lhs, rhs) + } + rumoca_ir_ast::Expression::Unary { op, rhs, .. } => { + let value = eval_known_integer_expr(rhs, known_int_params, scope)?; + match op { + rumoca_core::OpUnary::Minus => value.checked_neg(), + rumoca_core::OpUnary::Plus => Some(value), + _ => None, + } + } + rumoca_ir_ast::Expression::Parenthesized { inner, .. } => { + eval_known_integer_expr(inner, known_int_params, scope) + } + rumoca_ir_ast::Expression::FunctionCall { comp, args, .. } + if comp.parts.len() == 1 + && comp.parts[0].subs.is_none() + && comp.parts[0].ident.text.as_ref() == "div" + && args.len() == 2 => + { + let lhs = eval_known_integer_expr(&args[0], known_int_params, scope)?; + let rhs = eval_known_integer_expr(&args[1], known_int_params, scope)?; + (rhs != 0).then_some(lhs / rhs) + } + _ => None, + } +} + +fn eval_known_integer_binary(op: &rumoca_core::OpBinary, lhs: i64, rhs: i64) -> Option { + match op { + rumoca_core::OpBinary::Add => lhs.checked_add(rhs), + rumoca_core::OpBinary::Sub => lhs.checked_sub(rhs), + rumoca_core::OpBinary::Mul => lhs.checked_mul(rhs), + rumoca_core::OpBinary::Div if rhs != 0 && lhs % rhs == 0 => Some(lhs / rhs), + _ => None, + } +} + /// Flow/stream from the connection prefix (MLS §9.3), inheriting from the /// parent for record fields (e.g. `flow Complex i` makes i.re/i.im flow). pub(super) fn component_flow_stream( diff --git a/crates/rumoca-phase-instantiate/src/inheritance.rs b/crates/rumoca-phase-instantiate/src/inheritance.rs index e329c86f4..38f993093 100644 --- a/crates/rumoca-phase-instantiate/src/inheritance.rs +++ b/crates/rumoca-phase-instantiate/src/inheritance.rs @@ -1,5 +1,9 @@ //! Inheritance processing for the instantiate phase (MLS §7.1). //! +//! SPEC_0021 file-size exception: inheritance merge logic is still consolidated +//! while redeclare and modifier overlay behavior stabilize. split plan: move +//! redeclare resolution and modifier merge helpers into focused submodules. +//! //! This module handles the `extends` clause processing, merging inherited //! components and equations into the derived class. //! @@ -114,6 +118,11 @@ fn extract_modification_target(expr: &ast::Expression) -> Option { | ast::Expression::ClassModification { target, .. } => { target.parts.first().map(|p| p.ident.text.to_string()) } + ast::Expression::Binary { + op: rumoca_core::OpBinary::Assign, + lhs, + .. + } => extract_modification_target(lhs), // For named arguments like `x = value`, extract the name ast::Expression::NamedArgument { name, .. } => Some(name.text.to_string()), _ => None, @@ -127,6 +136,13 @@ fn extract_extend_modification_target( let target = match expr { ast::Expression::Modification { target, .. } | ast::Expression::ClassModification { target, .. } => target, + ast::Expression::Binary { + op: rumoca_core::OpBinary::Assign, + lhs, + .. + } => { + return extract_extend_modification_target(extend, lhs); + } ast::Expression::NamedArgument { name, .. } => return Some(name.text.to_string()), _ => return None, }; @@ -186,6 +202,12 @@ fn extract_modification_value(expr: &ast::Expression) -> Option let value = match expr { ast::Expression::Modification { value, .. } => Some(value), ast::Expression::NamedArgument { value, .. } => Some(value), + ast::Expression::Binary { + op: rumoca_core::OpBinary::Assign, + lhs, + rhs, + .. + } if matches!(lhs.as_ref(), ast::Expression::ClassModification { .. }) => Some(rhs), _ => None, }?; @@ -389,18 +411,29 @@ fn validate_redeclaration( .map(|n| n.to_string()) .unwrap_or_else(|| component.type_name.to_string()); - // Try to resolve constraint type using def_id or tree lookup - let constraint_type = if let Some(def_id) = component.type_def_id - && let Some(qualified) = tree.def_map.get(&def_id) - { - qualified.clone() - } else if let Some(&def_id) = tree.name_map.get(&constraint_type_raw) - && let Some(qualified) = tree.def_map.get(&def_id) - { - qualified.clone() - } else { - constraint_type_raw.clone() - }; + // Try to resolve constraint type using the explicit constrainedby first. + // The declared component type is only the default constraint when + // constrainedby is omitted (MLS §7.3.2). + let constraint_type = component + .constrainedby + .as_ref() + .and_then(|name| name.def_id) + .and_then(|def_id| tree.def_map.get(&def_id).cloned()) + .or_else(|| { + tree.name_map + .get(&constraint_type_raw) + .and_then(|def_id| tree.def_map.get(def_id).cloned()) + }) + .or_else(|| { + if component.constrainedby.is_none() { + component + .type_def_id + .and_then(|def_id| tree.def_map.get(&def_id).cloned()) + } else { + None + } + }) + .unwrap_or_else(|| constraint_type_raw.clone()); // Try to resolve new type name using the constraint type's package as context // This handles cases like GearType1 in the same package as GearType2 @@ -1089,7 +1122,7 @@ pub fn process_extends_with_cache( let base_inherited = process_extends_with_cache(tree, base_class, cache)?; merge_inherited(&mut inherited, base_inherited, extend, &tree.source_map)?; - // MLS §7.2: Apply extends modifications after recursive merge so + // MLS §7.2/§7.3: Apply extends modifications after recursive merge so // transitively inherited targets are available. apply_extends_modifications(tree, &mut inherited, base_class, extend)?; } @@ -1121,6 +1154,8 @@ fn apply_extends_modifications( base_class: &ast::ClassDef, extend: &ast::Extend, ) -> InstantiateResult<()> { + apply_transitive_extends_redeclarations(tree, target, base_class, extend)?; + let mut final_override: Option = None; walk_extend_modifications(extend, |modification| { let Some((name, value, is_final)) = @@ -1160,6 +1195,88 @@ fn apply_extends_modifications( Ok(()) } +fn apply_transitive_extends_redeclarations( + tree: &ast::ClassTree, + target: &mut InheritedContent, + base_class: &ast::ClassDef, + extend: &ast::Extend, +) -> InstantiateResult<()> { + let extend_span = + location_to_span(&extend.location, &tree.source_map, "extends redeclaration")?; + let mut validation_error: Option> = None; + let mut redeclare_types = IndexMap::default(); + + walk_extend_modifications(extend, |modification| { + if validation_error.is_some() { + return; + } + let Some((target_name, _value_expr)) = redeclare_target_value(modification) else { + return; + }; + if base_class.components.contains_key(target_name) { + return; + } + let target_name_owned = target_name.to_string(); + let Some(component) = target.components.get(&target_name_owned) else { + return; + }; + let new_type = extract_redeclare_type_qualified(&modification.expr, tree); + let span = match redeclare_target_span(tree, &target_name_owned, modification, extend_span) + { + Ok(span) => span, + Err(err) => { + validation_error = Some(err); + return; + } + }; + + if let Err(err) = validate_redeclaration( + tree, + component, + &target_name_owned, + new_type.as_deref(), + span, + ) { + validation_error = Some(err); + return; + } + if let Some(new_type_name) = new_type { + redeclare_types.insert(target_name_owned, new_type_name); + } + }); + + if let Some(err) = validation_error { + return Err(err); + } + + for (comp_name, new_type_name) in &redeclare_types { + if let Some(comp) = target.components.get_mut(comp_name) { + apply_redeclared_component_type(tree, comp, new_type_name); + } + } + + Ok(()) +} + +fn apply_redeclared_component_type( + tree: &ast::ClassTree, + comp: &mut ast::Component, + new_type_name: &str, +) { + // Update the type_name to the new type + comp.type_name = rumoca_ir_ast::Name::from_string(new_type_name); + // Update type_def_id by looking up the new type in the tree + comp.type_def_id = tree.name_map.get(new_type_name).copied().or_else(|| { + // Try with shorter name (last segment) for unqualified lookups + let short_name = path_utils::class_name_leaf(new_type_name); + tree.name_map.get(short_name).copied() + }); + + // MLS §7.3.2: Activate constraining-clause defaults for redeclared + // replaceable components. + activate_constrainedby_defaults_for_redeclare(comp); +} + /// Resolve a base class from an extends clause. /// /// Uses O(1) DefId lookup via ast::ClassTree.get_class_by_def_id(). @@ -1592,18 +1709,7 @@ fn merge_class_content( // This updates the component's type so that instantiation uses the new type's fields for (comp_name, new_type_name) in &redeclare_types { if let Some(comp) = target.components.get_mut(comp_name) { - // Update the type_name to the new type - comp.type_name = rumoca_ir_ast::Name::from_string(new_type_name); - // Update type_def_id by looking up the new type in the tree - comp.type_def_id = tree.name_map.get(new_type_name).copied().or_else(|| { - // Try with shorter name (last segment) for unqualified lookups - let short_name = path_utils::class_name_leaf(new_type_name); - tree.name_map.get(short_name).copied() - }); - - // MLS §7.3.2: Activate constraining-clause defaults for redeclared - // replaceable components. - activate_constrainedby_defaults_for_redeclare(comp); + apply_redeclared_component_type(tree, comp, new_type_name); } } @@ -1758,6 +1864,16 @@ fn activate_constrainedby_defaults_for_redeclare(comp: &mut ast::Component) { } } +fn activate_constrainedby_defaults_for_replaceable_components( + components: &mut IndexMap, +) { + for comp in components.values_mut() { + if comp.is_replaceable { + activate_constrainedby_defaults_for_redeclare(comp); + } + } +} + /// Merge nested class modifications from extends clause into inherited components. /// /// MLS §7.2: When an extends clause has modifications like @@ -1768,13 +1884,15 @@ fn activate_constrainedby_defaults_for_redeclare(comp: &mut ast::Component) { fn merge_nested_extends_modifications(target: &mut InheritedContent, extend: &ast::Extend) { walk_extend_modifications(extend, |modification| { // Extract target name and nested modifications from the expression. - // Two formats exist: + // Three formats exist: // 1. ClassModification { target: comp_name, modifications: [...] } // For: extends Foo(friction(useHeatPort=true)) // 2. Modification { target: comp_name, value: ClassModification { target: TypeName, modifications: [...] } } // For: extends Foo(redeclare final NewType comp(nested=val)) + // 3. Binary { Assign, lhs: ClassModification { target: comp_name, ... }, rhs } + // For: extends Foo(comp(each final unit="1")=expr) // Type changes are handled by collect_redeclarations(); here we merge nested mods. - let Some((target_name, modifications)) = + let Some((target_name, (modifications, each_flags, final_flags))) = extend_nested_target_modifications(extend, modification) else { return; @@ -1782,32 +1900,80 @@ fn merge_nested_extends_modifications(target: &mut InheritedContent, extend: &as let Some(comp) = target.components.get_mut(&target_name) else { return; }; - for nested_mod in modifications { - insert_nested_modification(comp, nested_mod); + for (idx, nested_mod) in modifications.iter().enumerate() { + insert_nested_modification_with_flags( + comp, + nested_mod, + each_flags.get(idx).copied().unwrap_or(false), + final_flags.get(idx).copied().unwrap_or(false), + ); } }); } +type NestedModificationSlices<'a> = (&'a [ast::Expression], &'a [bool], &'a [bool]); +type ExtendNestedTargetModifications<'a> = (String, NestedModificationSlices<'a>); + fn extend_nested_target_modifications<'a>( extend: &ast::Extend, modification: &'a ast::ExtendModification, -) -> Option<(String, &'a [ast::Expression])> { +) -> Option> { match &modification.expr { ast::Expression::ClassModification { target, modifications, + each_flags, + final_flags, .. } => Some(( extend_relative_component_target(extend, target)?, - modifications.as_slice(), + ( + modifications.as_slice(), + each_flags.as_slice(), + final_flags.as_slice(), + ), )), ast::Expression::Modification { target, value, .. } => { - let ast::Expression::ClassModification { modifications, .. } = value.as_ref() else { + let ast::Expression::ClassModification { + modifications, + each_flags, + final_flags, + .. + } = value.as_ref() + else { return None; }; Some(( extend_relative_component_target(extend, target)?, - modifications.as_slice(), + ( + modifications.as_slice(), + each_flags.as_slice(), + final_flags.as_slice(), + ), + )) + } + ast::Expression::Binary { + op: rumoca_core::OpBinary::Assign, + lhs, + .. + } => { + let ast::Expression::ClassModification { + target, + modifications, + each_flags, + final_flags, + .. + } = lhs.as_ref() + else { + return None; + }; + Some(( + extend_relative_component_target(extend, target)?, + ( + modifications.as_slice(), + each_flags.as_slice(), + final_flags.as_slice(), + ), )) } _ => None, @@ -1815,26 +1981,45 @@ fn extend_nested_target_modifications<'a>( } /// Insert a single nested modification into a component's modifications map. -fn insert_nested_modification(comp: &mut ast::Component, nested_mod: &ast::Expression) { +fn insert_nested_modification_with_flags( + comp: &mut ast::Component, + nested_mod: &ast::Expression, + each: bool, + final_: bool, +) { + let mut inserted_name: Option = None; match nested_mod { ast::Expression::Modification { target: t, value, .. } => { if let Some(name) = t.parts.first().map(|p| p.ident.text.to_string()) { - comp.modifications.insert(name, value.as_ref().clone()); + comp.modifications + .insert(name.clone(), value.as_ref().clone()); + inserted_name = Some(name); } } ast::Expression::NamedArgument { name, value, .. } => { + let inserted = name.text.to_string(); comp.modifications - .insert(name.text.to_string(), value.as_ref().clone()); + .insert(inserted.clone(), value.as_ref().clone()); + inserted_name = Some(inserted); } ast::Expression::ClassModification { .. } => { if let Some(name) = extract_modification_target(nested_mod) { - comp.modifications.insert(name, nested_mod.clone()); + comp.modifications.insert(name.clone(), nested_mod.clone()); + inserted_name = Some(name); } } _ => {} } + if let Some(name) = inserted_name { + if each { + comp.each_modifications.insert(name.clone()); + } + if final_ { + comp.final_attributes.insert(name); + } + } } /// Get the effective components for a class (own + inherited). @@ -1869,6 +2054,8 @@ pub fn get_effective_components_with_cache( inherited.components.insert(name.clone(), comp.clone()); } + activate_constrainedby_defaults_for_replaceable_components(&mut inherited.components); + // MLS §7.1/§7.3: local class names (including inherited replaceable classes) // are valid type names for component declarations in the effective class scope. // Preserve their resolved DefIds so later phases don't treat names like diff --git a/crates/rumoca-phase-instantiate/src/inheritance/tests.rs b/crates/rumoca-phase-instantiate/src/inheritance/tests.rs index 9714ec2d6..2b9faef26 100644 --- a/crates/rumoca-phase-instantiate/src/inheritance/tests.rs +++ b/crates/rumoca-phase-instantiate/src/inheritance/tests.rs @@ -87,6 +87,44 @@ fn make_int_expr(value: &str) -> ast::Expression { } } +fn make_string_expr(value: &str) -> ast::Expression { + ast::Expression::Terminal { + terminal_type: ast::TerminalType::String, + token: make_token(&format!("\"{value}\"")), + span: rumoca_core::Span::DUMMY, + } +} + +fn make_class_modification_binding( + target: &str, + nested_mods: Vec, + each_flags: Vec, + final_flags: Vec, + rhs: ast::Expression, +) -> ast::Expression { + ast::Expression::Binary { + op: rumoca_core::OpBinary::Assign, + lhs: Arc::new(ast::Expression::ClassModification { + target: make_component_ref(target), + modifications: nested_mods, + each_flags, + final_flags, + redeclare_flags: vec![], + span: rumoca_core::Span::DUMMY, + }), + rhs: Arc::new(rhs), + span: rumoca_core::Span::DUMMY, + } +} + +fn make_value_modification(target: &str, value: ast::Expression) -> ast::Expression { + ast::Expression::Modification { + target: make_component_ref(target), + value: Arc::new(value), + span: rumoca_core::Span::DUMMY, + } +} + #[test] fn test_apply_extends_modifications_reports_final_override_at_extends_span() { let mut tree = ast::ClassTree::default(); @@ -120,6 +158,98 @@ fn test_apply_extends_modifications_reports_final_override_at_extends_span() { assert!(matches!(*err, InstantiateError::RedeclareFinal { .. })); } +#[test] +fn class_modification_binding_is_value_modification_for_base_component() { + let mut class = ast::ClassDef { + name: make_token("Base"), + ..Default::default() + }; + class.components.insert( + "stageInputs".to_string(), + make_component("stageInputs", false, false), + ); + let rhs = ast::Expression::ComponentReference(make_component_ref("per.speeds")); + let extend = ast::Extend { + base_name: make_name("Base"), + modifications: vec![ast::ExtendModification { + expr: make_class_modification_binding("stageInputs", vec![], vec![], vec![], rhs), + each: false, + final_: true, + redeclare: false, + }], + ..Default::default() + }; + + let value_mods = collect_value_modifications(&extend, &class); + let (value, is_final) = value_mods + .get("stageInputs") + .expect("class-modification binding should override component binding"); + + assert!(*is_final); + assert!(matches!( + value, + ast::Expression::ComponentReference(cref) if cref.to_string() == "per.speeds" + )); +} + +#[test] +fn class_modification_binding_merges_nested_attributes_for_inherited_component() { + let mut tree = ast::ClassTree::default(); + tree.source_map.add( + TEST_FILE, + "extends Mid(stageInputs(each final unit=\"1\") = per.speeds);", + ); + let mut target = InheritedContent::default(); + target.components.insert( + "stageInputs".to_string(), + make_component("stageInputs", false, false), + ); + let base_class = ast::ClassDef { + name: make_token("Mid"), + ..Default::default() + }; + let unit_mod = make_value_modification("unit", make_string_expr("1")); + let rhs = ast::Expression::ComponentReference(make_component_ref("per.speeds")); + let extend = ast::Extend { + base_name: make_name("Mid"), + location: test_location(), + modifications: vec![ast::ExtendModification { + expr: make_class_modification_binding( + "stageInputs", + vec![unit_mod], + vec![true], + vec![true], + rhs, + ), + each: false, + final_: true, + redeclare: false, + }], + ..Default::default() + }; + + let mut inherited = target; + apply_extends_modifications(&tree, &mut inherited, &base_class, &extend) + .expect("class-modification binding should apply to inherited component"); + let comp = inherited + .components + .get("stageInputs") + .expect("stageInputs should stay inherited"); + + assert!(comp.has_explicit_binding); + assert!(comp.is_final); + assert!(matches!( + comp.binding.as_ref(), + Some(ast::Expression::ComponentReference(cref)) if cref.to_string() == "per.speeds" + )); + assert!(matches!( + comp.modifications.get("unit"), + Some(ast::Expression::Terminal { token, .. }) if token.text.as_ref() == "\"1\"" + )); + assert!(comp.each_modifications.contains("unit")); + assert!(comp.final_attributes.contains("unit")); +} + #[test] fn test_validate_redeclaration_non_replaceable() { // A non-replaceable component should fail redeclaration @@ -289,6 +419,71 @@ fn test_constrainedby_subtype_allowed() { assert!(result.is_ok()); } +#[test] +fn test_constrainedby_explicit_type_overrides_declared_type_def_id() { + let mut tree = ast::ClassTree::default(); + + let declared_id = DefId::new(1); + let constraint_id = DefId::new(2); + let replacement_id = DefId::new(3); + + let declared = ast::ClassDef { + name: make_token("DeclaredVolume"), + def_id: Some(declared_id), + ..Default::default() + }; + let constraint = ast::ClassDef { + name: make_token("HeatPortVolume"), + def_id: Some(constraint_id), + ..Default::default() + }; + let replacement = ast::ClassDef { + name: make_token("MoistureHeatPortVolume"), + def_id: Some(replacement_id), + extends: vec![ast::Extend { + base_name: make_resolved_name("HeatPortVolume", constraint_id), + base_def_id: Some(constraint_id), + ..Default::default() + }], + ..Default::default() + }; + + tree.definitions + .classes + .insert("DeclaredVolume".to_string(), declared); + tree.definitions + .classes + .insert("HeatPortVolume".to_string(), constraint); + tree.definitions + .classes + .insert("MoistureHeatPortVolume".to_string(), replacement); + for (name, def_id) in [ + ("DeclaredVolume", declared_id), + ("HeatPortVolume", constraint_id), + ("MoistureHeatPortVolume", replacement_id), + ] { + tree.name_map.insert(name.to_string(), def_id); + tree.def_map.insert(def_id, name.to_string()); + } + + let mut comp = make_constrained_component("vol2", "DeclaredVolume", Some("HeatPortVolume")); + comp.type_def_id = Some(declared_id); + comp.constrainedby = Some(make_resolved_name("HeatPortVolume", constraint_id)); + + let result = validate_redeclaration( + &tree, + &comp, + "vol2", + Some("MoistureHeatPortVolume"), + Span::DUMMY, + ); + + assert!( + result.is_ok(), + "explicit constrainedby must be the constraint even when declared type_def_id is present" + ); +} + #[test] fn test_class_redeclare_constraint_resolves_relative_to_declaration_scope() { let (tree, flow_characteristic_id) = relative_class_redeclare_constraint_tree(); diff --git a/crates/rumoca-phase-instantiate/src/lib.rs b/crates/rumoca-phase-instantiate/src/lib.rs index 6b9e8af3f..776776137 100644 --- a/crates/rumoca-phase-instantiate/src/lib.rs +++ b/crates/rumoca-phase-instantiate/src/lib.rs @@ -37,6 +37,7 @@ mod array_expansion; mod attributes; +mod binding_source; mod component_loop; mod connections; mod dims; @@ -57,7 +58,7 @@ mod type_overrides; use rumoca_eval_ast::eval_instantiate::{ InstantiateEvalCtx, evaluate_array_dimensions, evaluate_component_condition, extract_binding, - extract_bool_params_with_mods, extract_int_params_with_mods, generate_array_indices, + extract_bool_params_with_mods, extract_int_params_with_mods_and_known, generate_array_indices, propagate_record_alias_integer_params, }; @@ -68,6 +69,7 @@ use rumoca_ir_ast::AstIndexMap as IndexMap; use array_expansion::{ArrayExpansionScope, expand_array_component}; use attributes::*; +use binding_source::declaration_binding_source_for_flattening; use component_loop::{ ComponentImports, component_flow_stream, component_type_id, instantiate_effective_components, }; @@ -1007,7 +1009,7 @@ fn instantiate_class( effective_components, resolve_class_components: resolve_effective_components_for_eval, }; - let int_params = extract_int_params_with_mods(&eval_ctx); + let int_params = extract_int_params_with_mods_and_known(&eval_ctx, &ctx.known_int_params); ctx.register_known_int_params(&qualified_name, &int_params); // Instantiate each effective component (MLS §4.8 conditional components) @@ -1699,13 +1701,27 @@ fn prepare_component_binding_info( effective_components, resolve_class_components: resolve_effective_components_for_eval, }; + let declaration_source_scope = component_declaration_source_scope(ctx, comp); let ComponentAttrsAndBinding { mut attrs, mut binding, - binding_source, - binding_source_scope, + mut binding_source, + mut binding_source_scope, binding_from_modification, - } = extract_component_attrs_and_binding(comp, ctx.mod_env(), &eval_ctx, imports)?; + } = extract_component_attrs_and_binding_in_scope( + comp, + ctx.mod_env(), + &ctx.current_path().to_flat_string(), + &eval_ctx, + imports, + declaration_source_scope.as_ref(), + )?; + if binding_from_modification && binding_source_scope.is_none() { + binding_source_scope = binding + .as_ref() + .and_then(|expr| expression_source_scope(ctx, expr)) + .or_else(|| declaration_source_scope.clone()); + } infer_local_attribute_source_scopes(ctx, comp, &mut attrs); let start_from_declaration_binding = !binding_from_modification && binding.is_some() && attrs.start == binding; @@ -1719,8 +1735,18 @@ fn prepare_component_binding_info( effective_components, tree, )?; + let display_binding_source = declaration_binding_source_for_flattening( + comp, + declaration_binding, + &resolved_binding, + effective_components, + ctx.mod_env(), + ); + if binding_source.is_none() { + binding_source = display_binding_source.clone(); + } if start_from_declaration_binding { - attrs.start = Some(resolved_binding.clone()); + attrs.start = Some(declaration_binding.clone()); } binding = Some(resolved_binding); } diff --git a/crates/rumoca-phase-instantiate/src/mod_env.rs b/crates/rumoca-phase-instantiate/src/mod_env.rs index 0edc6e1d5..9bb985e53 100644 --- a/crates/rumoca-phase-instantiate/src/mod_env.rs +++ b/crates/rumoca-phase-instantiate/src/mod_env.rs @@ -3,10 +3,12 @@ use super::inheritance::{ resolve_effective_components_for_eval, }; use super::nested_scope::remap_redeclare_class_modifier; +use super::path_utils; use super::type_overrides::{TypeOverrideMap, find_nested_class_in_hierarchy}; use super::{InstantiateContext, InstantiateError, InstantiateResult}; use rumoca_eval_ast::eval_instantiate::{ - InstantiateEvalCtx, evaluate_component_condition, try_eval_integer_expr, try_eval_string_expr, + InstantiateEvalCtx, evaluate_component_condition, try_eval_enum_expr, try_eval_integer_expr, + try_eval_string_expr, }; use rumoca_ir_ast as ast; use rumoca_ir_ast::AstIndexMap as IndexMap; @@ -219,6 +221,15 @@ fn inherited_modifier_source_metadata( let qn = ast::QualifiedName::from_ident(name); let mod_value = mod_env.get(&qn)?; if mod_value.value == *expr { + if mod_value.source.is_some() || mod_value.source_scope.is_some() { + return Some(( + mod_value + .source + .clone() + .or_else(|| Some(mod_value.value.clone())), + mod_value.source_scope.clone(), + )); + } return None; } @@ -243,11 +254,17 @@ fn apply_component_modifier( component_type_allows_string_modifier(&target_comp.type_name.to_string()) }); - // INST-010: Check if target component is final in the target class (MLS §7.2.6) - if target_component - .as_ref() - .is_some_and(|target_comp| target_comp.is_final) + if let Some(target_comp) = target_component.as_ref() + && target_comp.is_final { + let qn = ast::QualifiedName::from_ident(target_name); + if prefixes.final_ + && let Some(existing) = ctx.mod_env().get(&qn) + && existing.final_ + && modification_forwards_to_existing(mod_expr, existing, ctx.mod_env()) + { + return Ok(()); + } let span = required_modifier_expr_span(mod_expr, "final component modifier")?; return Err(Box::new(InstantiateError::redeclare_final( target_name, @@ -326,6 +343,7 @@ fn apply_component_modifier( eval_ctx.tree, &eval_ctx.insert_ctx, )?; + preserve_redeclare_class_modifier(ctx, target_name, mod_expr, eval_ctx); } } @@ -403,6 +421,26 @@ fn component_type_allows_string_modifier(type_name: &str) -> bool { rumoca_core::qualified_type_name_matches(type_name, "String") } +fn modification_forwards_to_existing( + value: &ast::Expression, + existing: &rumoca_ir_ast::ModificationValue, + mod_env: &ast::ModificationEnvironment, +) -> bool { + if existing.value == *value { + return true; + } + let ast::Expression::ComponentReference(cref) = value else { + return false; + }; + if cref.parts.len() != 1 { + return false; + } + let qn = ast::QualifiedName::from_ident(cref.parts[0].ident.text.as_ref()); + mod_env + .get(&qn) + .is_some_and(|outer| outer.value == existing.value) +} + fn insert_scoped_modifier_binding( ctx: &mut InstantiateContext, binding: ScopedModifierBinding, @@ -433,6 +471,15 @@ fn insert_scoped_modifier_binding( .get(&key) .is_some_and(|existing| existing.final_ && !replace_parent) { + if prefixes.final_ + && modification_forwards_to_existing( + &value, + ctx.mod_env().get(&key).expect("checked above"), + ctx.mod_env(), + ) + { + return Ok(()); + } let span = required_binding_source_span(source.as_ref(), &value, "final modifier binding")?; return Err(Box::new(InstantiateError::redeclare_final( key.to_flat_string(), @@ -623,8 +670,36 @@ fn resolve_modification_expr_with_depth( }); } - // Try to evaluate as an integer (common for array dimension parameters). - if let Some(value) = try_eval_integer_expr(&eval_ctx, expr) { + // MLS §7.2: Resolve multi-part references through sibling modifications + // before treating unresolved references as enum literals. + if let Some(resolved) = resolve_sibling_modification(expr, effective_components, mode) { + if resolved == *expr { + return Ok(expr.clone()); + } + return resolve_modification_expr_with_depth( + &resolved, + mod_env, + effective_components, + tree, + allow_string_eval, + mode, + depth + 1, + ); + } + + if let Some(value) = try_eval_enum_expr(&eval_ctx, expr) { + return Ok(path_utils::component_ref_expr_from_dotted( + &value, + expr.span(), + )); + } + + // Try to evaluate modifier expressions as integers (common for array + // dimension parameters). Declaration bindings keep their expression shape + // so flattening can resolve unqualified names in the owning instance scope. + if mode == ModificationResolveMode::Modifier + && let Some(value) = try_eval_integer_expr(&eval_ctx, expr) + { return Ok(ast::Expression::Terminal { terminal_type: rumoca_ir_ast::TerminalType::UnsignedInteger, token: rumoca_core::Token { @@ -637,7 +712,8 @@ fn resolve_modification_expr_with_depth( // Resolve direct references in current scope (e.g. resolveInFrame=resolveInFrame). if mode == ModificationResolveMode::Modifier - && let Some(resolved_ref) = resolve_single_part_ref_expr(expr, mod_env) + && let Some(resolved_ref) = + resolve_single_part_ref_expr(expr, mod_env, effective_components) { return resolve_modification_expr_with_depth( &resolved_ref, @@ -650,10 +726,12 @@ fn resolve_modification_expr_with_depth( ); } - // MLS §7.2: Resolve multi-part references through sibling modifications. - if let Some(resolved) = resolve_sibling_modification(expr, effective_components) { + if mode == ModificationResolveMode::DeclarationBinding + && let Some(resolved_ref) = + resolve_structural_declaration_ref_expr(expr, effective_components) + { return resolve_modification_expr_with_depth( - &resolved, + &resolved_ref, mod_env, effective_components, tree, @@ -666,9 +744,37 @@ fn resolve_modification_expr_with_depth( Ok(expr.clone()) } +fn resolve_structural_declaration_ref_expr( + expr: &ast::Expression, + effective_components: &IndexMap, +) -> Option { + let ast::Expression::ComponentReference(comp_ref) = expr else { + return None; + }; + if comp_ref.parts.len() != 1 || comp_ref.parts[0].subs.is_some() { + return None; + } + + let name = comp_ref.parts[0].ident.text.as_ref(); + let component = effective_components.get(name)?; + if !matches!( + component.variability, + rumoca_core::Variability::Parameter(_) | rumoca_core::Variability::Constant(_) + ) && !component.is_structural + { + return None; + } + component + .binding + .as_ref() + .filter(|binding| *binding != expr) + .cloned() +} + fn resolve_single_part_ref_expr( expr: &ast::Expression, mod_env: &ast::ModificationEnvironment, + effective_components: &IndexMap, ) -> Option { let ast::Expression::ComponentReference(comp_ref) = expr else { return None; @@ -680,6 +786,16 @@ fn resolve_single_part_ref_expr( let name = comp_ref.parts[0].ident.text.as_ref(); let qn = ast::QualifiedName::from_ident(name); + if let Some(subscripts) = comp_ref.parts[0].subs.as_ref() { + let mod_value = mod_env.get(&qn)?; + let indices = literal_integer_subscripts(subscripts)?; + return index_array_value(&mod_value.value, &indices); + } + + if effective_components.contains_key(name) { + return None; + } + if let Some(mod_value) = mod_env.get(&qn) && mod_value.value != *expr { @@ -689,10 +805,45 @@ fn resolve_single_part_ref_expr( None } +fn literal_integer_subscripts(subscripts: &[ast::Subscript]) -> Option> { + subscripts + .iter() + .map(|subscript| { + let ast::Subscript::Expression(ast::Expression::Terminal { + terminal_type: ast::TerminalType::UnsignedInteger, + token, + .. + }) = subscript + else { + return None; + }; + token.text.parse::().ok() + }) + .collect() +} + +fn index_array_value(expr: &ast::Expression, indices: &[i64]) -> Option { + let (&first, rest) = indices.split_first()?; + match expr { + ast::Expression::Array { elements, .. } => { + let index = first.checked_sub(1)? as usize; + let selected = elements.get(index)?; + if rest.is_empty() { + Some(selected.clone()) + } else { + index_array_value(selected, rest) + } + } + ast::Expression::Parenthesized { inner, .. } => index_array_value(inner, indices), + _ => None, + } +} + /// Resolve a multi-part component reference by following sibling modifications. fn resolve_sibling_modification( expr: &ast::Expression, effective_components: &IndexMap, + mode: ModificationResolveMode, ) -> Option { let ast::Expression::ComponentReference(comp_ref) = expr else { return None; @@ -704,6 +855,11 @@ fn resolve_sibling_modification( let second = comp_ref.parts[1].ident.text.as_ref(); let comp = effective_components.get(first)?; let mod_expr = comp.modifications.get(second)?; + if mode == ModificationResolveMode::DeclarationBinding + && matches!(mod_expr, ast::Expression::ComponentReference(_)) + { + return Some(expr.clone()); + } // Keep record/class-modification bindings as references so declaration // defaults remain visible during record-field projection. if is_non_scalar_sibling_modifier_expr(mod_expr) { @@ -812,1141 +968,3 @@ fn preserves_source_scoped_attribute(attr_name: &str) -> bool { #[cfg(test)] #[path = "mod_env_tests.rs"] mod mod_env_tests; - -#[cfg(test)] -mod tests { - use super::*; - use std::sync::Arc; - - const TEST_FILE: &str = "mod_env.mo"; - - fn test_location() -> rumoca_core::Location { - rumoca_core::Location { - start_line: 1, - start_column: 1, - end_line: 1, - end_column: 2, - start: 0, - end: 1, - file_name: TEST_FILE.to_string(), - } - } - - fn test_span() -> rumoca_core::Span { - rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name("phase_instantiate_mod_env_source_7.mo"), - 0, - 1, - ) - } - - fn make_token(text: &str) -> rumoca_core::Token { - rumoca_core::Token { - text: std::sync::Arc::from(text), - location: rumoca_core::Location::default(), - token_number: 0, - token_type: 0, - } - } - - fn make_int_expr_with_span(value: i64, span: rumoca_core::Span) -> ast::Expression { - ast::Expression::Terminal { - terminal_type: ast::TerminalType::UnsignedInteger, - token: make_token(&value.to_string()), - span, - } - } - - fn make_int_expr(value: i64) -> ast::Expression { - make_int_expr_with_span(value, rumoca_core::Span::DUMMY) - } - - fn make_comp_ref_expr(names: &[&str]) -> ast::Expression { - ast::Expression::ComponentReference(ast::ComponentReference { - local: false, - parts: names - .iter() - .map(|name| ast::ComponentRefPart { - ident: make_token(name), - subs: None, - }) - .collect(), - def_id: None, - span: rumoca_core::Span::DUMMY, - }) - } - - fn make_named_arg(name: &str, value: ast::Expression) -> ast::Expression { - ast::Expression::NamedArgument { - name: make_token(name), - value: Arc::new(value), - span: rumoca_core::Span::DUMMY, - } - } - - fn active_mod_env_keys(ctx: &InstantiateContext) -> Vec { - ctx.mod_env() - .active - .keys() - .map(ToString::to_string) - .collect() - } - - fn make_name(name: &str) -> ast::Name { - ast::Name { - name: vec![make_token(name)], - def_id: None, - } - } - - #[test] - fn string_modifier_type_check_requires_segment_boundary() { - assert!(component_type_allows_string_modifier("String")); - assert!(component_type_allows_string_modifier("Modelica.String")); - assert!(component_type_allows_string_modifier("Pkg.Types.String")); - assert!(!component_type_allows_string_modifier("MyString")); - assert!(!component_type_allows_string_modifier("Pkg.StringAlias")); - } - - #[test] - fn test_resolve_sibling_modification_keeps_class_modification_reference() { - let mut effective_components: IndexMap = IndexMap::default(); - let mut data = ast::Component { - name: "aimcData".to_string(), - ..ast::Component::empty_with_span(test_span()) - }; - data.modifications.insert( - "statorCoreParameters".to_string(), - ast::Expression::ClassModification { - target: ast::ComponentReference { - local: false, - parts: vec![ - ast::ComponentRefPart { - ident: make_token("Modelica"), - subs: None, - }, - ast::ComponentRefPart { - ident: make_token("Electrical"), - subs: None, - }, - ast::ComponentRefPart { - ident: make_token("Machines"), - subs: None, - }, - ast::ComponentRefPart { - ident: make_token("Losses"), - subs: None, - }, - ast::ComponentRefPart { - ident: make_token("CoreParameters"), - subs: None, - }, - ], - def_id: None, - span: rumoca_core::Span::DUMMY, - }, - modifications: vec![ - make_named_arg("PRef", make_int_expr(410)), - make_named_arg("VRef", make_int_expr(388)), - ], - each_flags: vec![false, false], - final_flags: vec![false, false], - redeclare_flags: vec![false, false], - span: rumoca_core::Span::DUMMY, - }, - ); - effective_components.insert("aimcData".to_string(), data); - - let expr = make_comp_ref_expr(&["aimcData", "statorCoreParameters"]); - let resolved = resolve_modification_expr( - &expr, - &ast::ModificationEnvironment::default(), - &effective_components, - &ast::ClassTree::default(), - false, - ) - .expect("resolution should succeed"); - - assert_eq!( - resolved, expr, - "record bindings should stay as references so declaration defaults are preserved" - ); - } - - #[test] - fn test_resolve_sibling_modification_still_resolves_scalar_field_override() { - let mut effective_components: IndexMap = IndexMap::default(); - let mut data = ast::Component { - name: "stackData".to_string(), - ..ast::Component::empty_with_span(test_span()) - }; - data.modifications - .insert("mSystems".to_string(), make_int_expr(2)); - effective_components.insert("stackData".to_string(), data); - - let expr = make_comp_ref_expr(&["stackData", "mSystems"]); - let resolved = resolve_modification_expr( - &expr, - &ast::ModificationEnvironment::default(), - &effective_components, - &ast::ClassTree::default(), - false, - ) - .expect("resolution should succeed"); - - assert_eq!( - resolved, - make_int_expr(2), - "scalar sibling field overrides should keep existing behavior" - ); - } - - #[test] - fn test_declaration_binding_preserves_component_reference_identity() { - let mut mod_env = ast::ModificationEnvironment::default(); - mod_env.add( - ast::QualifiedName::from_ident("pathLengths"), - ast::ModificationValue::with_source_scope( - make_comp_ref_expr(&["length"]), - Some(make_comp_ref_expr(&["length"])), - Some(ast::QualifiedName::from_ident("pipe")), - ), - ); - - let expr = make_comp_ref_expr(&["pathLengths"]); - let resolved = resolve_declaration_binding_expr( - &expr, - &mod_env, - &IndexMap::default(), - &ast::ClassTree::default(), - ) - .expect("declaration binding resolution should succeed"); - - assert_eq!( - resolved, expr, - "declaration bindings must keep sibling component references instead of inlining modifier values" - ); - } - - #[test] - fn test_resolve_sibling_modification_keeps_function_call_record_like_binding() { - let mut effective_components: IndexMap = IndexMap::default(); - let mut data = ast::Component { - name: "aimcData".to_string(), - ..ast::Component::empty_with_span(test_span()) - }; - data.modifications.insert( - "statorCoreParameters".to_string(), - ast::Expression::FunctionCall { - comp: ast::ComponentReference { - local: false, - parts: vec![ - ast::ComponentRefPart { - ident: make_token("Modelica"), - subs: None, - }, - ast::ComponentRefPart { - ident: make_token("Electrical"), - subs: None, - }, - ast::ComponentRefPart { - ident: make_token("Machines"), - subs: None, - }, - ast::ComponentRefPart { - ident: make_token("Losses"), - subs: None, - }, - ast::ComponentRefPart { - ident: make_token("CoreParameters"), - subs: None, - }, - ], - def_id: None, - span: rumoca_core::Span::DUMMY, - }, - args: vec![make_int_expr(410), make_int_expr(388)], - span: rumoca_core::Span::DUMMY, - }, - ); - effective_components.insert("aimcData".to_string(), data); - - let expr = make_comp_ref_expr(&["aimcData", "statorCoreParameters"]); - let resolved = resolve_modification_expr( - &expr, - &ast::ModificationEnvironment::default(), - &effective_components, - &ast::ClassTree::default(), - false, - ) - .expect("resolution should succeed"); - - assert_eq!( - resolved, expr, - "function-call record-like overrides should stay as references" - ); - } - - #[test] - fn test_insert_scoped_modifier_binding_reorders_non_shifted_parent_key() { - let mut ctx = InstantiateContext::new(); - let key = ast::QualifiedName::from_ident("k"); - let sibling = ast::QualifiedName::from_ident("a"); - - ctx.mod_env_mut().add( - key.clone(), - ast::ModificationValue::simple(make_int_expr(1)), - ); - ctx.mod_env_mut().add( - sibling.clone(), - ast::ModificationValue::simple(make_int_expr(2)), - ); - - let mut parent_snapshot = IndexMap::default(); - parent_snapshot.insert( - key.clone(), - ast::ModificationValue::simple(make_int_expr(1)), - ); - parent_snapshot.insert( - sibling.clone(), - ast::ModificationValue::simple(make_int_expr(2)), - ); - - insert_scoped_modifier_binding( - &mut ctx, - ScopedModifierBinding { - key: key.clone(), - value: make_int_expr(9), - source: None, - source_scope: None, - prefixes: ModifierPrefixes::default(), - }, - &parent_snapshot, - &IndexMap::default(), - ) - .expect("non-final parent key can be replaced"); - - assert_eq!( - active_mod_env_keys(&ctx), - vec!["a".to_string(), "k".to_string()], - "non-shifted parent key should be replaced as a new local binding" - ); - assert_eq!( - ctx.mod_env().get(&key).map(|mv| mv.value.clone()), - Some(make_int_expr(9)) - ); - } - - #[test] - fn test_insert_scoped_modifier_binding_keeps_shifted_parent_key_position() { - let mut ctx = InstantiateContext::new(); - let key = ast::QualifiedName::from_ident("k"); - let sibling = ast::QualifiedName::from_ident("a"); - - ctx.mod_env_mut().add( - key.clone(), - ast::ModificationValue::simple(make_int_expr(1)), - ); - ctx.mod_env_mut().add( - sibling.clone(), - ast::ModificationValue::simple(make_int_expr(2)), - ); - - let mut parent_snapshot = IndexMap::default(); - parent_snapshot.insert( - key.clone(), - ast::ModificationValue::simple(make_int_expr(1)), - ); - parent_snapshot.insert( - sibling.clone(), - ast::ModificationValue::simple(make_int_expr(2)), - ); - - let mut shifted_parent_keys = IndexMap::default(); - shifted_parent_keys.insert(key.clone(), ()); - - insert_scoped_modifier_binding( - &mut ctx, - ScopedModifierBinding { - key: key.clone(), - value: make_int_expr(11), - source: None, - source_scope: None, - prefixes: ModifierPrefixes::default(), - }, - &parent_snapshot, - &shifted_parent_keys, - ) - .expect("shifted non-final parent key can be preserved"); - - assert_eq!( - active_mod_env_keys(&ctx), - vec!["k".to_string(), "a".to_string()], - "shifted parent key should remain in place" - ); - assert_eq!( - ctx.mod_env().get(&key).map(|mv| mv.value.clone()), - Some(make_int_expr(1)), - "outer/shifted parent modifier must keep precedence (MLS §7.2.4)" - ); - } - - #[test] - fn test_insert_scoped_modifier_binding_keeps_shifted_final_parent_key() { - let mut ctx = InstantiateContext::new(); - let key = ast::QualifiedName::from_ident("R"); - ctx.mod_env_mut().add( - key.clone(), - ast::ModificationValue::with_prefixes(make_int_expr(10), false, true), - ); - - let mut parent_snapshot = IndexMap::default(); - parent_snapshot.insert( - key.clone(), - ast::ModificationValue::with_prefixes(make_int_expr(10), false, true), - ); - - let mut shifted_parent_keys = IndexMap::default(); - shifted_parent_keys.insert(key.clone(), ()); - - insert_scoped_modifier_binding( - &mut ctx, - ScopedModifierBinding { - key: key.clone(), - value: make_int_expr(20), - source: None, - source_scope: None, - prefixes: ModifierPrefixes::default(), - }, - &parent_snapshot, - &shifted_parent_keys, - ) - .expect("shifted final outer modifier keeps precedence over inner default"); - - assert_eq!( - ctx.mod_env().get(&key).map(|mv| mv.value.clone()), - Some(make_int_expr(10)) - ); - } - - #[test] - fn test_apply_component_modifier_rejects_inherited_final_component() { - let base_id = rumoca_core::DefId::new(10); - let mut base = ast::ClassDef { - def_id: Some(base_id), - name: make_token("Base"), - ..Default::default() - }; - base.components.insert( - "p".to_string(), - ast::Component { - name: "p".to_string(), - type_name: make_name("Real"), - is_final: true, - ..ast::Component::empty_with_span(test_span()) - }, - ); - - let mut derived = ast::ClassDef { - name: make_token("Derived"), - ..Default::default() - }; - derived.extends.push(ast::Extend { - base_name: make_name("Base"), - base_def_id: Some(base_id), - location: test_location(), - ..Default::default() - }); - - let mut tree = ast::ClassTree::default(); - tree.source_map.add(TEST_FILE, "extends Base;"); - tree.definitions.classes.insert("Base".to_string(), base); - tree.definitions - .classes - .insert("Derived".to_string(), derived.clone()); - tree.def_map.insert(base_id, "Base".to_string()); - tree.name_map.insert("Base".to_string(), base_id); - - let mut ctx = InstantiateContext::new(); - let parent_snapshot = IndexMap::default(); - let shifted_parent_keys = IndexMap::default(); - let type_overrides = TypeOverrideMap::default(); - let eval_ctx = ModifierEvalContext { - tree: &tree, - effective_components: &IndexMap::default(), - type_overrides: &type_overrides, - target_class: Some(&derived), - insert_ctx: ScopedInsertContext { - parent_snapshot: &parent_snapshot, - shifted_parent_keys: &shifted_parent_keys, - source_scope: None, - }, - }; - - let err = apply_component_modifier( - &mut ctx, - "p", - &make_int_expr_with_span(2, test_span()), - ModifierPrefixes::default(), - &eval_ctx, - ) - .expect_err("inherited final components must reject modification"); - - let message = err.to_string(); - assert!(message.contains("final"), "{message}"); - assert!(matches!(*err, InstantiateError::RedeclareFinal { .. })); - } - - #[test] - fn test_apply_component_modifier_requires_span_for_final_modifier_error() { - let mut class = ast::ClassDef { - name: make_token("C"), - ..Default::default() - }; - class.components.insert( - "p".to_string(), - ast::Component { - name: "p".to_string(), - type_name: make_name("Real"), - is_final: true, - ..ast::Component::empty_with_span(test_span()) - }, - ); - - let mut ctx = InstantiateContext::new(); - let parent_snapshot = IndexMap::default(); - let shifted_parent_keys = IndexMap::default(); - let type_overrides = TypeOverrideMap::default(); - let tree = ast::ClassTree::default(); - let eval_ctx = ModifierEvalContext { - tree: &tree, - effective_components: &IndexMap::default(), - type_overrides: &type_overrides, - target_class: Some(&class), - insert_ctx: ScopedInsertContext { - parent_snapshot: &parent_snapshot, - shifted_parent_keys: &shifted_parent_keys, - source_scope: None, - }, - }; - - let err = apply_component_modifier( - &mut ctx, - "p", - &make_int_expr(2), - ModifierPrefixes::default(), - &eval_ctx, - ) - .expect_err("unspanned final modifier error should fail fast"); - - assert!(matches!( - *err, - InstantiateError::MissingSourceContext { .. } - )); - } - - #[test] - fn test_insert_scoped_modifier_binding_reports_final_collision_at_source() { - let mut ctx = InstantiateContext::new(); - let key = ast::QualifiedName::from_ident("k"); - ctx.mod_env_mut().add( - key.clone(), - ast::ModificationValue::with_prefixes( - make_int_expr_with_span(1, test_span()), - false, - true, - ), - ); - - let err = insert_scoped_modifier_binding( - &mut ctx, - ScopedModifierBinding { - key, - value: make_int_expr_with_span(2, test_span()), - source: Some(make_int_expr_with_span(3, test_span())), - source_scope: None, - prefixes: ModifierPrefixes::default(), - }, - &IndexMap::default(), - &IndexMap::default(), - ) - .expect_err("local modification must not override final binding"); - - assert!(matches!(*err, InstantiateError::RedeclareFinal { .. })); - } - - #[test] - fn test_insert_scoped_modifier_binding_requires_span_for_final_collision() { - let mut ctx = InstantiateContext::new(); - let key = ast::QualifiedName::from_ident("k"); - ctx.mod_env_mut().add( - key.clone(), - ast::ModificationValue::with_prefixes( - make_int_expr_with_span(1, test_span()), - false, - true, - ), - ); - - let err = insert_scoped_modifier_binding( - &mut ctx, - ScopedModifierBinding { - key, - value: make_int_expr(2), - source: None, - source_scope: None, - prefixes: ModifierPrefixes::default(), - }, - &IndexMap::default(), - &IndexMap::default(), - ) - .expect_err("unspanned final binding error should fail fast"); - - assert!(matches!( - *err, - InstantiateError::MissingSourceContext { .. } - )); - } - - #[test] - fn test_forwarded_modifier_keeps_forwarded_source_scope() { - let mut ctx = InstantiateContext::new(); - let key = ast::QualifiedName::from_ident("frictionParameters"); - let forwarded_value = make_comp_ref_expr(&["aimcData", "frictionParameters"]); - let forwarded_scope = Some(ast::QualifiedName::new()); - - ctx.mod_env_mut().active.insert( - key.clone(), - ast::ModificationValue::with_source_scope( - forwarded_value.clone(), - Some(forwarded_value.clone()), - forwarded_scope.clone(), - ), - ); - - let parent_snapshot = ctx.mod_env().active.clone(); - let shifted_parent_keys: IndexMap = IndexMap::default(); - let insert_ctx = ScopedInsertContext { - parent_snapshot: &parent_snapshot, - shifted_parent_keys: &shifted_parent_keys, - source_scope: Some(ast::QualifiedName::from_ident("aimc")), - }; - - insert_modifier_value_with_structural_overrides( - &mut ctx, - "frictionParameters", - &make_comp_ref_expr(&["frictionParameters"]), - ModifierInsertOptions { - allow_string_eval: false, - prefixes: ModifierPrefixes::default(), - }, - &IndexMap::default(), - &ast::ClassTree::default(), - &insert_ctx, - ) - .expect("forwarded modifier insertion should succeed"); - - let stored = ctx - .mod_env() - .get(&key) - .expect("forwarded modifier binding should exist"); - assert_eq!( - stored.value, forwarded_value, - "forwarded binding should preserve resolved parent expression" - ); - assert_eq!( - stored.source_scope, forwarded_scope, - "forwarded binding should preserve original lexical source scope" - ); - assert_eq!( - stored.source.as_ref(), - Some(&forwarded_value), - "forwarded binding should preserve symbolic source expression" - ); - } - - #[test] - fn test_sibling_modifier_reference_keeps_local_source_scope() { - let mut ctx = InstantiateContext::new(); - ctx.mod_env_mut().active.insert( - ast::QualifiedName::from_ident("pathLengths"), - ast::ModificationValue::with_source_scope( - make_comp_ref_expr(&["length"]), - Some(make_comp_ref_expr(&["length"])), - Some(ast::QualifiedName::from_ident("pipe")), - ), - ); - - let parent_snapshot = ctx.mod_env().active.clone(); - let shifted_parent_keys: IndexMap = IndexMap::default(); - let local_scope = Some(ast::QualifiedName::from_ident("flowModel")); - let insert_ctx = ScopedInsertContext { - parent_snapshot: &parent_snapshot, - shifted_parent_keys: &shifted_parent_keys, - source_scope: local_scope.clone(), - }; - - insert_modifier_value_with_structural_overrides( - &mut ctx, - "pathLengths_internal", - &make_comp_ref_expr(&["pathLengths"]), - ModifierInsertOptions { - allow_string_eval: false, - prefixes: ModifierPrefixes::default(), - }, - &IndexMap::default(), - &ast::ClassTree::default(), - &insert_ctx, - ) - .expect("sibling modifier insertion should succeed"); - - let stored = ctx - .mod_env() - .get(&ast::QualifiedName::from_ident("pathLengths_internal")) - .expect("sibling modifier binding should exist"); - assert_eq!( - stored.source.as_ref(), - Some(&make_comp_ref_expr(&["pathLengths"])) - ); - assert_eq!(stored.source_scope, local_scope); - } - - #[test] - fn test_modifier_with_same_resolved_value_keeps_existing_source_scope() { - let mut ctx = InstantiateContext::new(); - let key = ast::QualifiedName::from_ident("frictionParameters"); - let forwarded_value = make_comp_ref_expr(&["aimcData", "frictionParameters"]); - let forwarded_scope = Some(ast::QualifiedName::new()); - - ctx.mod_env_mut().active.insert( - key.clone(), - ast::ModificationValue::with_source_scope( - forwarded_value.clone(), - Some(forwarded_value.clone()), - forwarded_scope.clone(), - ), - ); - - let parent_snapshot = ctx.mod_env().active.clone(); - let shifted_parent_keys: IndexMap = IndexMap::default(); - let insert_ctx = ScopedInsertContext { - parent_snapshot: &parent_snapshot, - shifted_parent_keys: &shifted_parent_keys, - source_scope: Some(ast::QualifiedName::from_ident("aimc")), - }; - - insert_modifier_value_with_structural_overrides( - &mut ctx, - "frictionParameters", - &forwarded_value, - ModifierInsertOptions { - allow_string_eval: false, - prefixes: ModifierPrefixes::default(), - }, - &IndexMap::default(), - &ast::ClassTree::default(), - &insert_ctx, - ) - .expect("same-value modifier insertion should succeed"); - - let stored = ctx - .mod_env() - .get(&key) - .expect("modifier binding should exist"); - assert_eq!( - stored.source_scope, forwarded_scope, - "resolved multi-part modifier should inherit source scope from existing parent binding" - ); - } - - #[test] - fn test_propagate_record_binding_overrides_non_targeted_field_values() { - let mut nested_record = ast::ClassDef { - name: make_token("State"), - class_type: rumoca_core::ClassType::Record, - ..Default::default() - }; - nested_record.components.insert( - "phase".to_string(), - ast::Component::empty_with_span(test_span()), - ); - nested_record.components.insert( - "p".to_string(), - ast::Component::empty_with_span(test_span()), - ); - - let mut ctx = InstantiateContext::new(); - ctx.mod_env_mut().add( - ast::QualifiedName::from_ident("phase"), - ast::ModificationValue::simple(make_int_expr(7)), - ); - - let binding_expr = make_comp_ref_expr(&["state_in"]); - let targeted_keys: IndexMap = IndexMap::default(); - propagate_record_binding_to_fields( - &ast::ClassTree::default(), - &mut ctx, - &binding_expr, - None, - &nested_record, - &targeted_keys, - ) - .expect("record field projection should succeed"); - - let phase_mod = ctx - .mod_env() - .active - .get(&ast::QualifiedName::from_ident("phase")) - .expect("phase field binding should be present"); - match &phase_mod.value { - ast::Expression::FieldAccess { base, field, .. } => { - assert_eq!(field, "phase"); - match base.as_ref() { - ast::Expression::ComponentReference(cref) => { - assert_eq!(cref.parts.len(), 1); - assert_eq!(cref.parts[0].ident.text.as_ref(), "state_in"); - } - _ => panic!("field binding should project from record binding expression"), - } - } - _ => panic!("phase field should be rebound from record binding"), - } - } - - #[test] - fn test_propagate_record_binding_preserves_targeted_field_modifiers() { - let mut nested_record = ast::ClassDef { - name: make_token("State"), - class_type: rumoca_core::ClassType::Record, - ..Default::default() - }; - nested_record.components.insert( - "phase".to_string(), - ast::Component::empty_with_span(test_span()), - ); - - let mut ctx = InstantiateContext::new(); - let phase_qn = ast::QualifiedName::from_ident("phase"); - ctx.mod_env_mut().add( - phase_qn.clone(), - ast::ModificationValue::simple(make_int_expr(42)), - ); - - let mut targeted_keys: IndexMap = IndexMap::default(); - targeted_keys.insert(phase_qn.clone(), ()); - let binding_expr = make_comp_ref_expr(&["state_in"]); - propagate_record_binding_to_fields( - &ast::ClassTree::default(), - &mut ctx, - &binding_expr, - None, - &nested_record, - &targeted_keys, - ) - .expect("record field projection should succeed"); - - let phase_mod = ctx - .mod_env() - .active - .get(&phase_qn) - .expect("targeted phase modifier should still be present"); - match &phase_mod.value { - ast::Expression::Terminal { token, .. } => assert_eq!(token.text.as_ref(), "42"), - _ => panic!("targeted field modifier should not be replaced"), - } - } - - #[test] - fn test_propagate_record_binding_projects_if_expression_branches_per_field() { - let mut nested_record = ast::ClassDef { - name: make_token("CellData"), - class_type: rumoca_core::ClassType::Record, - ..Default::default() - }; - nested_record.components.insert( - "OCV_SOC".to_string(), - ast::Component::empty_with_span(test_span()), - ); - - let mut ctx = InstantiateContext::new(); - let binding_expr = ast::Expression::If { - branches: vec![( - make_comp_ref_expr(&["isDegraded"]), - make_comp_ref_expr(&["cellDataDegraded"]), - )], - else_branch: Arc::new(make_comp_ref_expr(&["cellDataOriginal"])), - span: rumoca_core::Span::DUMMY, - }; - let targeted_keys: IndexMap = IndexMap::default(); - propagate_record_binding_to_fields( - &ast::ClassTree::default(), - &mut ctx, - &binding_expr, - None, - &nested_record, - &targeted_keys, - ) - .expect("record field projection should succeed"); - - let field_mod = ctx - .mod_env() - .active - .get(&ast::QualifiedName::from_ident("OCV_SOC")) - .expect("OCV_SOC field binding should be present"); - let ast::Expression::If { - branches, - else_branch, - .. - } = &field_mod.value - else { - panic!("field projection should preserve if-expression structure"); - }; - assert_eq!(branches.len(), 1); - let (_cond, then_expr) = &branches[0]; - let ast::Expression::FieldAccess { base, field, .. } = then_expr else { - panic!("then-branch should project field access"); - }; - assert_eq!(field, "OCV_SOC"); - assert_eq!(*base.as_ref(), make_comp_ref_expr(&["cellDataDegraded"])); - - let ast::Expression::FieldAccess { - base: else_base, - field: else_field, - .. - } = else_branch.as_ref() - else { - panic!("else-branch should project field access"); - }; - assert_eq!(else_field, "OCV_SOC"); - assert_eq!( - *else_base.as_ref(), - make_comp_ref_expr(&["cellDataOriginal"]) - ); - } - - #[test] - fn test_propagate_record_binding_preserves_matching_default_record_constructor() { - let mut nested_record = ast::ClassDef { - name: make_token("BaseData"), - class_type: rumoca_core::ClassType::Record, - ..Default::default() - }; - nested_record.components.insert( - "mu_i".to_string(), - ast::Component { - binding: Some(make_int_expr(1)), - start: make_int_expr(1), - ..ast::Component::empty_with_span(test_span()) - }, - ); - - let mut ctx = InstantiateContext::new(); - let binding_expr = ast::Expression::FunctionCall { - comp: ast::ComponentReference { - local: false, - parts: vec![ast::ComponentRefPart { - ident: make_token("BaseData"), - subs: None, - }], - def_id: None, - span: rumoca_core::Span::DUMMY, - }, - args: Vec::new(), - span: rumoca_core::Span::DUMMY, - }; - - propagate_record_binding_to_fields( - &ast::ClassTree::default(), - &mut ctx, - &binding_expr, - None, - &nested_record, - &IndexMap::default(), - ) - .expect("record field projection should succeed"); - - assert!( - ctx.mod_env().active.is_empty(), - "matching zero-argument record constructors should preserve declared defaults" - ); - } - - #[test] - fn test_propagate_record_binding_projects_subtype_default_record_constructor_fields() { - let mut nested_record = ast::ClassDef { - name: make_token("BaseData"), - class_type: rumoca_core::ClassType::Record, - ..Default::default() - }; - nested_record.components.insert( - "mu_i".to_string(), - ast::Component { - binding: Some(make_int_expr(1)), - start: make_int_expr(1), - ..ast::Component::empty_with_span(test_span()) - }, - ); - - let mut ctx = InstantiateContext::new(); - let binding_expr = ast::Expression::FunctionCall { - comp: ast::ComponentReference { - local: false, - parts: vec![ast::ComponentRefPart { - ident: make_token("M350_50A"), - subs: None, - }], - def_id: None, - span: rumoca_core::Span::DUMMY, - }, - args: Vec::new(), - span: rumoca_core::Span::DUMMY, - }; - - propagate_record_binding_to_fields( - &ast::ClassTree::default(), - &mut ctx, - &binding_expr, - None, - &nested_record, - &IndexMap::default(), - ) - .expect("record field projection should succeed"); - - let field_mod = ctx - .mod_env() - .active - .get(&ast::QualifiedName::from_ident("mu_i")) - .expect("subtype default record constructor should project field binding"); - let ast::Expression::FieldAccess { base, field, .. } = &field_mod.value else { - panic!("subtype constructor field should be projected"); - }; - assert_eq!(field, "mu_i"); - assert_eq!(base.as_ref(), &binding_expr); - } - - #[test] - fn test_propagate_record_binding_projects_through_unique_constructor_record_field() { - let inner_def_id = rumoca_core::DefId::new(1001); - let outer_def_id = rumoca_core::DefId::new(1002); - let mut inner_record = ast::ClassDef { - name: make_token("Inner"), - class_type: rumoca_core::ClassType::Record, - def_id: Some(inner_def_id), - ..Default::default() - }; - inner_record.components.insert( - "x".to_string(), - ast::Component::empty_with_span(test_span()), - ); - - let mut outer_record = ast::ClassDef { - name: make_token("Outer"), - class_type: rumoca_core::ClassType::Record, - def_id: Some(outer_def_id), - ..Default::default() - }; - outer_record.components.insert( - "innerParams".to_string(), - ast::Component { - type_name: ast::Name { - name: vec![make_token("Pkg"), make_token("Inner")], - def_id: Some(inner_def_id), - }, - type_def_id: Some(inner_def_id), - ..ast::Component::empty_with_span(test_span()) - }, - ); - - let mut tree = ast::ClassTree::default(); - tree.definitions - .classes - .insert("Pkg.Outer".to_string(), outer_record); - tree.def_map.insert(inner_def_id, "Pkg.Inner".to_string()); - tree.def_map.insert(outer_def_id, "Pkg.Outer".to_string()); - - let mut ctx = InstantiateContext::new(); - let binding_expr = ast::Expression::FunctionCall { - comp: ast::ComponentReference { - local: false, - parts: vec![ast::ComponentRefPart { - ident: make_token("Outer"), - subs: None, - }], - def_id: None, - span: rumoca_core::Span::DUMMY, - }, - args: Vec::new(), - span: rumoca_core::Span::DUMMY, - }; - - propagate_record_binding_to_fields( - &tree, - &mut ctx, - &binding_expr, - Some(ast::QualifiedName::from_ident("Pkg")), - &inner_record, - &IndexMap::default(), - ) - .expect("record field projection should succeed"); - - let field_mod = ctx - .mod_env() - .active - .get(&ast::QualifiedName::from_ident("x")) - .expect("inner field binding should be present"); - let ast::Expression::FieldAccess { - base, - field: inner_field, - .. - } = &field_mod.value - else { - panic!("inner field should be projected"); - }; - assert_eq!(inner_field, "x"); - let ast::Expression::FieldAccess { - base: constructor, - field: outer_field, - .. - } = base.as_ref() - else { - panic!("projection should first select the unique compatible record field"); - }; - assert_eq!(outer_field, "innerParams"); - assert_eq!(constructor.as_ref(), &binding_expr); - } - - #[test] - fn test_propagate_record_binding_skips_non_record_classes() { - let mut nested_block = ast::ClassDef { - name: make_token("UniformNoise"), - class_type: rumoca_core::ClassType::Block, - ..Default::default() - }; - nested_block.components.insert( - "y".to_string(), - ast::Component::empty_with_span(test_span()), - ); - nested_block.components.insert( - "seedState".to_string(), - ast::Component::empty_with_span(test_span()), - ); - - let mut ctx = InstantiateContext::new(); - let binding_expr = make_comp_ref_expr(&["noise"]); - let targeted_keys: IndexMap = IndexMap::default(); - propagate_record_binding_to_fields( - &ast::ClassTree::default(), - &mut ctx, - &binding_expr, - None, - &nested_block, - &targeted_keys, - ) - .expect("non-record projection should succeed without mutation"); - - assert!( - ctx.mod_env().active.is_empty(), - "non-record class modifiers must not synthesize per-field record bindings" - ); - } -} diff --git a/crates/rumoca-phase-instantiate/src/mod_env/record_projection.rs b/crates/rumoca-phase-instantiate/src/mod_env/record_projection.rs index dafc6a9a9..a565b8aeb 100644 --- a/crates/rumoca-phase-instantiate/src/mod_env/record_projection.rs +++ b/crates/rumoca-phase-instantiate/src/mod_env/record_projection.rs @@ -1,4 +1,5 @@ use super::*; +use rumoca_ir_ast::visitor::ExpressionTransformer; use std::sync::Arc; /// Propagate a record binding to scalar field bindings. @@ -50,7 +51,15 @@ pub(crate) fn propagate_record_binding_to_fields( components, ctx.mod_env(), field_name, - ); + ) + .or_else(|| { + same_type_alias_projected_field_default( + binding_expr, + components, + ctx.mod_env(), + field_name, + ) + }); if field_binding.is_none() && should_preserve_same_type_alias_field_default( binding_expr, @@ -88,6 +97,53 @@ pub(crate) fn propagate_record_binding_to_fields( Ok(()) } +fn same_type_alias_projected_field_default( + binding_expr: &ast::Expression, + effective_components: &IndexMap, + mod_env: &ast::ModificationEnvironment, + field_name: &str, +) -> Option { + let source_name = simple_record_alias_source_name(binding_expr)?; + let field_comp = effective_components.get(field_name)?; + let default_expr = field_comp.binding.as_ref().or_else(|| { + (!matches!(field_comp.start, ast::Expression::Empty { .. })).then_some(&field_comp.start) + })?; + let mut rewriter = SourceRecordFieldDefaultRewriter { + source_name, + effective_components, + mod_env, + rewrote_source_field: false, + }; + let rewritten = rewriter.transform_expression(default_expr.clone()); + rewriter.rewrote_source_field.then_some(rewritten) +} + +struct SourceRecordFieldDefaultRewriter<'a> { + source_name: &'a str, + effective_components: &'a IndexMap, + mod_env: &'a ast::ModificationEnvironment, + rewrote_source_field: bool, +} + +impl ExpressionTransformer for SourceRecordFieldDefaultRewriter<'_> { + fn transform_component_reference(&mut self, cr: ast::ComponentReference) -> ast::Expression { + if cr.parts.len() == 1 + && cr.parts[0].subs.is_none() + && self + .effective_components + .contains_key(cr.parts[0].ident.text.as_ref()) + { + let source_field = ast::QualifiedName::from_ident(self.source_name) + .child(cr.parts[0].ident.text.as_ref()); + if let Some(value) = self.mod_env.active.get(&source_field) { + self.rewrote_source_field = true; + return value.value.clone(); + } + } + ast::Expression::ComponentReference(self.transform_component_ref_inner(cr)) + } +} + fn should_preserve_same_type_alias_field_default( binding_expr: &ast::Expression, target_record: &ast::ClassDef, @@ -165,7 +221,7 @@ fn same_type_alias_explicit_field_binding( .active .iter() .find(|(key, _)| **key == source_field) - .map(|(_, value)| value.source.clone().unwrap_or_else(|| value.value.clone())) + .map(|(_, value)| value.value.clone()) } fn record_alias_source_explicitly_binds_field( @@ -271,6 +327,9 @@ fn constructor_record_projection_base( let Some(source_record) = constructor_class_for_call(tree, comp, binding_source_scope) else { return Ok(None); }; + if source_record.class_type != rumoca_core::ClassType::Record { + return Ok(None); + } if source_record.def_id == target_record.def_id { return Ok(None); } @@ -297,6 +356,9 @@ fn constructor_projected_field_binding( let Some(source_record) = constructor_class_for_call(tree, comp, binding_source_scope) else { return Ok(None); }; + if source_record.class_type != rumoca_core::ClassType::Record { + return Ok(None); + } let effective = get_effective_components(tree, source_record)?; let components = if effective.is_empty() { diff --git a/crates/rumoca-phase-instantiate/src/mod_env_tests.rs b/crates/rumoca-phase-instantiate/src/mod_env_tests.rs index e5fd0d5c6..95d00c0d2 100644 --- a/crates/rumoca-phase-instantiate/src/mod_env_tests.rs +++ b/crates/rumoca-phase-instantiate/src/mod_env_tests.rs @@ -48,6 +48,14 @@ fn make_comp_ref_expr(names: &[&str]) -> ast::Expression { }) } +fn named_arg(name: &str, value: ast::Expression) -> ast::Expression { + ast::Expression::NamedArgument { + name: make_token(name), + value: std::sync::Arc::new(value), + span: rumoca_core::Span::DUMMY, + } +} + fn make_component(name: &str, type_name: &str, type_def_id: Option) -> ast::Component { ast::Component { name: name.to_string(), @@ -291,3 +299,1375 @@ fn test_collect_structural_integer_fields_from_sibling_reference() { assert_eq!(seen.get("Ns"), Some(&3)); assert_eq!(seen.get("Np"), Some(&2)); } + +#[cfg(test)] +mod moved_inline_tests { + use super::*; + use std::sync::Arc; + + const TEST_FILE: &str = "mod_env.mo"; + + fn test_location() -> rumoca_core::Location { + rumoca_core::Location { + start_line: 1, + start_column: 1, + end_line: 1, + end_column: 2, + start: 0, + end: 1, + file_name: TEST_FILE.to_string(), + } + } + + fn test_span() -> rumoca_core::Span { + rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("phase_instantiate_mod_env_source_7.mo"), + 0, + 1, + ) + } + + fn make_token(text: &str) -> rumoca_core::Token { + rumoca_core::Token { + text: std::sync::Arc::from(text), + location: rumoca_core::Location::default(), + token_number: 0, + token_type: 0, + } + } + + fn make_int_expr_with_span(value: i64, span: rumoca_core::Span) -> ast::Expression { + ast::Expression::Terminal { + terminal_type: ast::TerminalType::UnsignedInteger, + token: make_token(&value.to_string()), + span, + } + } + + fn make_int_expr(value: i64) -> ast::Expression { + make_int_expr_with_span(value, rumoca_core::Span::DUMMY) + } + + fn make_comp_ref_expr(names: &[&str]) -> ast::Expression { + ast::Expression::ComponentReference(ast::ComponentReference { + local: false, + parts: names + .iter() + .map(|name| ast::ComponentRefPart { + ident: make_token(name), + subs: None, + }) + .collect(), + def_id: None, + span: rumoca_core::Span::DUMMY, + }) + } + + fn make_named_arg(name: &str, value: ast::Expression) -> ast::Expression { + ast::Expression::NamedArgument { + name: make_token(name), + value: Arc::new(value), + span: rumoca_core::Span::DUMMY, + } + } + + fn active_mod_env_keys(ctx: &InstantiateContext) -> Vec { + ctx.mod_env() + .active + .keys() + .map(ToString::to_string) + .collect() + } + + fn make_name(name: &str) -> ast::Name { + ast::Name { + name: vec![make_token(name)], + def_id: None, + } + } + + #[test] + fn string_modifier_type_check_requires_segment_boundary() { + assert!(component_type_allows_string_modifier("String")); + assert!(component_type_allows_string_modifier("Modelica.String")); + assert!(component_type_allows_string_modifier("Pkg.Types.String")); + assert!(!component_type_allows_string_modifier("MyString")); + assert!(!component_type_allows_string_modifier("Pkg.StringAlias")); + } + + #[test] + fn test_resolve_sibling_modification_keeps_class_modification_reference() { + let mut effective_components: IndexMap = IndexMap::default(); + let mut data = ast::Component { + name: "aimcData".to_string(), + ..ast::Component::empty_with_span(test_span()) + }; + data.modifications.insert( + "statorCoreParameters".to_string(), + ast::Expression::ClassModification { + target: ast::ComponentReference { + local: false, + parts: vec![ + ast::ComponentRefPart { + ident: make_token("Modelica"), + subs: None, + }, + ast::ComponentRefPart { + ident: make_token("Electrical"), + subs: None, + }, + ast::ComponentRefPart { + ident: make_token("Machines"), + subs: None, + }, + ast::ComponentRefPart { + ident: make_token("Losses"), + subs: None, + }, + ast::ComponentRefPart { + ident: make_token("CoreParameters"), + subs: None, + }, + ], + def_id: None, + span: rumoca_core::Span::DUMMY, + }, + modifications: vec![ + make_named_arg("PRef", make_int_expr(410)), + make_named_arg("VRef", make_int_expr(388)), + ], + each_flags: vec![false, false], + final_flags: vec![false, false], + redeclare_flags: vec![false, false], + span: rumoca_core::Span::DUMMY, + }, + ); + effective_components.insert("aimcData".to_string(), data); + + let expr = make_comp_ref_expr(&["aimcData", "statorCoreParameters"]); + let resolved = resolve_modification_expr( + &expr, + &ast::ModificationEnvironment::default(), + &effective_components, + &ast::ClassTree::default(), + false, + ) + .expect("resolution should succeed"); + + assert_eq!( + resolved, expr, + "record bindings should stay as references so declaration defaults are preserved" + ); + } + + #[test] + fn test_resolve_sibling_modification_still_resolves_scalar_field_override() { + let mut effective_components: IndexMap = IndexMap::default(); + let mut data = ast::Component { + name: "stackData".to_string(), + ..ast::Component::empty_with_span(test_span()) + }; + data.modifications + .insert("mSystems".to_string(), make_int_expr(2)); + effective_components.insert("stackData".to_string(), data); + + let expr = make_comp_ref_expr(&["stackData", "mSystems"]); + let resolved = resolve_modification_expr( + &expr, + &ast::ModificationEnvironment::default(), + &effective_components, + &ast::ClassTree::default(), + false, + ) + .expect("resolution should succeed"); + + assert_eq!( + resolved, + make_int_expr(2), + "scalar sibling field overrides should keep existing behavior" + ); + } + + #[test] + fn test_declaration_binding_preserves_component_reference_identity() { + let mut mod_env = ast::ModificationEnvironment::default(); + mod_env.add( + ast::QualifiedName::from_ident("pathLengths"), + ast::ModificationValue::with_source_scope( + make_comp_ref_expr(&["length"]), + Some(make_comp_ref_expr(&["length"])), + Some(ast::QualifiedName::from_ident("pipe")), + ), + ); + + let expr = make_comp_ref_expr(&["pathLengths"]); + let resolved = resolve_declaration_binding_expr( + &expr, + &mod_env, + &IndexMap::default(), + &ast::ClassTree::default(), + ) + .expect("declaration binding resolution should succeed"); + + assert_eq!( + resolved, expr, + "declaration bindings must keep sibling component references instead of inlining modifier values" + ); + } + + #[test] + fn test_declaration_binding_preserves_nested_sibling_component_reference_identity() { + let mut effective_components: IndexMap = IndexMap::default(); + let mut joint_usp = ast::Component { + name: "jointUSP".to_string(), + ..ast::Component::empty_with_span(test_span()) + }; + joint_usp + .modifications + .insert("e2_ia".to_string(), make_comp_ref_expr(&["rod1", "e2_ia"])); + effective_components.insert("jointUSP".to_string(), joint_usp); + + let expr = make_comp_ref_expr(&["jointUSP", "e2_ia"]); + let resolved = resolve_declaration_binding_expr( + &expr, + &ast::ModificationEnvironment::default(), + &effective_components, + &ast::ClassTree::default(), + ) + .expect("declaration binding resolution should succeed"); + + assert_eq!( + resolved, expr, + "declaration bindings must not inline a sibling field into a receiver-less relative reference" + ); + } + + #[test] + fn test_declaration_binding_preserves_composite_integer_expression_shape() { + let mut effective_components: IndexMap = IndexMap::default(); + let mut n = ast::Component { + name: "n".to_string(), + variability: rumoca_core::Variability::Parameter(rumoca_core::Token::default()), + ..ast::Component::empty_with_span(test_span()) + }; + n.binding = Some(make_int_expr(2)); + effective_components.insert("n".to_string(), n); + + let expr = ast::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Arc::new(make_comp_ref_expr(&["n"])), + rhs: Arc::new(make_int_expr(1)), + span: rumoca_core::Span::DUMMY, + }; + let resolved = resolve_declaration_binding_expr( + &expr, + &ast::ModificationEnvironment::default(), + &effective_components, + &ast::ClassTree::default(), + ) + .expect("declaration binding resolution should succeed"); + + assert_eq!( + resolved, expr, + "declaration bindings must keep composite expressions for instance-scope name resolution" + ); + } + + #[test] + fn test_resolve_sibling_modification_keeps_function_call_record_like_binding() { + let mut effective_components: IndexMap = IndexMap::default(); + let mut data = ast::Component { + name: "aimcData".to_string(), + ..ast::Component::empty_with_span(test_span()) + }; + data.modifications.insert( + "statorCoreParameters".to_string(), + ast::Expression::FunctionCall { + comp: ast::ComponentReference { + local: false, + parts: vec![ + ast::ComponentRefPart { + ident: make_token("Modelica"), + subs: None, + }, + ast::ComponentRefPart { + ident: make_token("Electrical"), + subs: None, + }, + ast::ComponentRefPart { + ident: make_token("Machines"), + subs: None, + }, + ast::ComponentRefPart { + ident: make_token("Losses"), + subs: None, + }, + ast::ComponentRefPart { + ident: make_token("CoreParameters"), + subs: None, + }, + ], + def_id: None, + span: rumoca_core::Span::DUMMY, + }, + args: vec![make_int_expr(410), make_int_expr(388)], + span: rumoca_core::Span::DUMMY, + }, + ); + effective_components.insert("aimcData".to_string(), data); + + let expr = make_comp_ref_expr(&["aimcData", "statorCoreParameters"]); + let resolved = resolve_modification_expr( + &expr, + &ast::ModificationEnvironment::default(), + &effective_components, + &ast::ClassTree::default(), + false, + ) + .expect("resolution should succeed"); + + assert_eq!( + resolved, expr, + "function-call record-like overrides should stay as references" + ); + } + + #[test] + fn test_insert_scoped_modifier_binding_reorders_non_shifted_parent_key() { + let mut ctx = InstantiateContext::new(); + let key = ast::QualifiedName::from_ident("k"); + let sibling = ast::QualifiedName::from_ident("a"); + + ctx.mod_env_mut().add( + key.clone(), + ast::ModificationValue::simple(make_int_expr(1)), + ); + ctx.mod_env_mut().add( + sibling.clone(), + ast::ModificationValue::simple(make_int_expr(2)), + ); + + let mut parent_snapshot = IndexMap::default(); + parent_snapshot.insert( + key.clone(), + ast::ModificationValue::simple(make_int_expr(1)), + ); + parent_snapshot.insert( + sibling.clone(), + ast::ModificationValue::simple(make_int_expr(2)), + ); + + insert_scoped_modifier_binding( + &mut ctx, + ScopedModifierBinding { + key: key.clone(), + value: make_int_expr(9), + source: None, + source_scope: None, + prefixes: ModifierPrefixes::default(), + }, + &parent_snapshot, + &IndexMap::default(), + ) + .expect("non-final parent key can be replaced"); + + assert_eq!( + active_mod_env_keys(&ctx), + vec!["a".to_string(), "k".to_string()], + "non-shifted parent key should be replaced as a new local binding" + ); + assert_eq!( + ctx.mod_env().get(&key).map(|mv| mv.value.clone()), + Some(make_int_expr(9)) + ); + } + + #[test] + fn test_insert_scoped_modifier_binding_keeps_shifted_parent_key_position() { + let mut ctx = InstantiateContext::new(); + let key = ast::QualifiedName::from_ident("k"); + let sibling = ast::QualifiedName::from_ident("a"); + + ctx.mod_env_mut().add( + key.clone(), + ast::ModificationValue::simple(make_int_expr(1)), + ); + ctx.mod_env_mut().add( + sibling.clone(), + ast::ModificationValue::simple(make_int_expr(2)), + ); + + let mut parent_snapshot = IndexMap::default(); + parent_snapshot.insert( + key.clone(), + ast::ModificationValue::simple(make_int_expr(1)), + ); + parent_snapshot.insert( + sibling.clone(), + ast::ModificationValue::simple(make_int_expr(2)), + ); + + let mut shifted_parent_keys = IndexMap::default(); + shifted_parent_keys.insert(key.clone(), ()); + + insert_scoped_modifier_binding( + &mut ctx, + ScopedModifierBinding { + key: key.clone(), + value: make_int_expr(11), + source: None, + source_scope: None, + prefixes: ModifierPrefixes::default(), + }, + &parent_snapshot, + &shifted_parent_keys, + ) + .expect("shifted non-final parent key can be preserved"); + + assert_eq!( + active_mod_env_keys(&ctx), + vec!["k".to_string(), "a".to_string()], + "shifted parent key should remain in place" + ); + assert_eq!( + ctx.mod_env().get(&key).map(|mv| mv.value.clone()), + Some(make_int_expr(1)), + "outer/shifted parent modifier must keep precedence (MLS §7.2.4)" + ); + } + + #[test] + fn test_insert_scoped_modifier_binding_keeps_shifted_final_parent_key() { + let mut ctx = InstantiateContext::new(); + let key = ast::QualifiedName::from_ident("R"); + ctx.mod_env_mut().add( + key.clone(), + ast::ModificationValue::with_prefixes(make_int_expr(10), false, true), + ); + + let mut parent_snapshot = IndexMap::default(); + parent_snapshot.insert( + key.clone(), + ast::ModificationValue::with_prefixes(make_int_expr(10), false, true), + ); + + let mut shifted_parent_keys = IndexMap::default(); + shifted_parent_keys.insert(key.clone(), ()); + + insert_scoped_modifier_binding( + &mut ctx, + ScopedModifierBinding { + key: key.clone(), + value: make_int_expr(20), + source: None, + source_scope: None, + prefixes: ModifierPrefixes::default(), + }, + &parent_snapshot, + &shifted_parent_keys, + ) + .expect("shifted final outer modifier keeps precedence over inner default"); + + assert_eq!( + ctx.mod_env().get(&key).map(|mv| mv.value.clone()), + Some(make_int_expr(10)) + ); + } + + #[test] + fn test_apply_component_modifier_rejects_inherited_final_component() { + let base_id = rumoca_core::DefId::new(10); + let mut base = ast::ClassDef { + def_id: Some(base_id), + name: make_token("Base"), + ..Default::default() + }; + base.components.insert( + "p".to_string(), + ast::Component { + name: "p".to_string(), + type_name: make_name("Real"), + is_final: true, + ..ast::Component::empty_with_span(test_span()) + }, + ); + + let mut derived = ast::ClassDef { + name: make_token("Derived"), + ..Default::default() + }; + derived.extends.push(ast::Extend { + base_name: make_name("Base"), + base_def_id: Some(base_id), + location: test_location(), + ..Default::default() + }); + + let mut tree = ast::ClassTree::default(); + tree.source_map.add(TEST_FILE, "extends Base;"); + tree.definitions.classes.insert("Base".to_string(), base); + tree.definitions + .classes + .insert("Derived".to_string(), derived.clone()); + tree.def_map.insert(base_id, "Base".to_string()); + tree.name_map.insert("Base".to_string(), base_id); + + let mut ctx = InstantiateContext::new(); + let parent_snapshot = IndexMap::default(); + let shifted_parent_keys = IndexMap::default(); + let type_overrides = TypeOverrideMap::default(); + let eval_ctx = ModifierEvalContext { + tree: &tree, + effective_components: &IndexMap::default(), + type_overrides: &type_overrides, + target_class: Some(&derived), + insert_ctx: ScopedInsertContext { + parent_snapshot: &parent_snapshot, + shifted_parent_keys: &shifted_parent_keys, + source_scope: None, + }, + }; + + let err = apply_component_modifier( + &mut ctx, + "p", + &make_int_expr_with_span(2, test_span()), + ModifierPrefixes::default(), + &eval_ctx, + ) + .expect_err("inherited final components must reject modification"); + + let message = err.to_string(); + assert!(message.contains("final"), "{message}"); + assert!(matches!(*err, InstantiateError::RedeclareFinal { .. })); + } + + #[test] + fn test_apply_component_modifier_requires_span_for_final_modifier_error() { + let mut class = ast::ClassDef { + name: make_token("C"), + ..Default::default() + }; + class.components.insert( + "p".to_string(), + ast::Component { + name: "p".to_string(), + type_name: make_name("Real"), + is_final: true, + ..ast::Component::empty_with_span(test_span()) + }, + ); + + let mut ctx = InstantiateContext::new(); + let parent_snapshot = IndexMap::default(); + let shifted_parent_keys = IndexMap::default(); + let type_overrides = TypeOverrideMap::default(); + let tree = ast::ClassTree::default(); + let eval_ctx = ModifierEvalContext { + tree: &tree, + effective_components: &IndexMap::default(), + type_overrides: &type_overrides, + target_class: Some(&class), + insert_ctx: ScopedInsertContext { + parent_snapshot: &parent_snapshot, + shifted_parent_keys: &shifted_parent_keys, + source_scope: None, + }, + }; + + let err = apply_component_modifier( + &mut ctx, + "p", + &make_int_expr(2), + ModifierPrefixes::default(), + &eval_ctx, + ) + .expect_err("unspanned final modifier error should fail fast"); + + assert!(matches!( + *err, + InstantiateError::MissingSourceContext { .. } + )); + } + + #[test] + fn test_insert_scoped_modifier_binding_reports_final_collision_at_source() { + let mut ctx = InstantiateContext::new(); + let key = ast::QualifiedName::from_ident("k"); + ctx.mod_env_mut().add( + key.clone(), + ast::ModificationValue::with_prefixes( + make_int_expr_with_span(1, test_span()), + false, + true, + ), + ); + + let err = insert_scoped_modifier_binding( + &mut ctx, + ScopedModifierBinding { + key, + value: make_int_expr_with_span(2, test_span()), + source: Some(make_int_expr_with_span(3, test_span())), + source_scope: None, + prefixes: ModifierPrefixes::default(), + }, + &IndexMap::default(), + &IndexMap::default(), + ) + .expect_err("local modification must not override final binding"); + + assert!(matches!(*err, InstantiateError::RedeclareFinal { .. })); + } + + #[test] + fn test_insert_scoped_modifier_binding_requires_span_for_final_collision() { + let mut ctx = InstantiateContext::new(); + let key = ast::QualifiedName::from_ident("k"); + ctx.mod_env_mut().add( + key.clone(), + ast::ModificationValue::with_prefixes( + make_int_expr_with_span(1, test_span()), + false, + true, + ), + ); + + let err = insert_scoped_modifier_binding( + &mut ctx, + ScopedModifierBinding { + key, + value: make_int_expr(2), + source: None, + source_scope: None, + prefixes: ModifierPrefixes::default(), + }, + &IndexMap::default(), + &IndexMap::default(), + ) + .expect_err("unspanned final binding error should fail fast"); + + assert!(matches!( + *err, + InstantiateError::MissingSourceContext { .. } + )); + } + + #[test] + fn test_forwarded_modifier_keeps_forwarded_source_scope() { + let mut ctx = InstantiateContext::new(); + let key = ast::QualifiedName::from_ident("frictionParameters"); + let forwarded_value = make_comp_ref_expr(&["aimcData", "frictionParameters"]); + let forwarded_scope = Some(ast::QualifiedName::new()); + + ctx.mod_env_mut().active.insert( + key.clone(), + ast::ModificationValue::with_source_scope( + forwarded_value.clone(), + Some(forwarded_value.clone()), + forwarded_scope.clone(), + ), + ); + + let parent_snapshot = ctx.mod_env().active.clone(); + let shifted_parent_keys: IndexMap = IndexMap::default(); + let insert_ctx = ScopedInsertContext { + parent_snapshot: &parent_snapshot, + shifted_parent_keys: &shifted_parent_keys, + source_scope: Some(ast::QualifiedName::from_ident("aimc")), + }; + + insert_modifier_value_with_structural_overrides( + &mut ctx, + "frictionParameters", + &make_comp_ref_expr(&["frictionParameters"]), + ModifierInsertOptions { + allow_string_eval: false, + prefixes: ModifierPrefixes::default(), + }, + &IndexMap::default(), + &ast::ClassTree::default(), + &insert_ctx, + ) + .expect("forwarded modifier insertion should succeed"); + + let stored = ctx + .mod_env() + .get(&key) + .expect("forwarded modifier binding should exist"); + assert_eq!( + stored.value, forwarded_value, + "forwarded binding should preserve resolved parent expression" + ); + assert_eq!( + stored.source_scope, forwarded_scope, + "forwarded binding should preserve original lexical source scope" + ); + assert_eq!( + stored.source.as_ref(), + Some(&forwarded_value), + "forwarded binding should preserve symbolic source expression" + ); + } + + #[test] + fn test_sibling_modifier_reference_keeps_local_source_scope() { + let mut ctx = InstantiateContext::new(); + ctx.mod_env_mut().active.insert( + ast::QualifiedName::from_ident("pathLengths"), + ast::ModificationValue::with_source_scope( + make_comp_ref_expr(&["length"]), + Some(make_comp_ref_expr(&["length"])), + Some(ast::QualifiedName::from_ident("pipe")), + ), + ); + + let parent_snapshot = ctx.mod_env().active.clone(); + let shifted_parent_keys: IndexMap = IndexMap::default(); + let local_scope = Some(ast::QualifiedName::from_ident("flowModel")); + let insert_ctx = ScopedInsertContext { + parent_snapshot: &parent_snapshot, + shifted_parent_keys: &shifted_parent_keys, + source_scope: local_scope.clone(), + }; + + insert_modifier_value_with_structural_overrides( + &mut ctx, + "pathLengths_internal", + &make_comp_ref_expr(&["pathLengths"]), + ModifierInsertOptions { + allow_string_eval: false, + prefixes: ModifierPrefixes::default(), + }, + &IndexMap::default(), + &ast::ClassTree::default(), + &insert_ctx, + ) + .expect("sibling modifier insertion should succeed"); + + let stored = ctx + .mod_env() + .get(&ast::QualifiedName::from_ident("pathLengths_internal")) + .expect("sibling modifier binding should exist"); + assert_eq!( + stored.source.as_ref(), + Some(&make_comp_ref_expr(&["pathLengths"])) + ); + assert_eq!(stored.source_scope, local_scope); + } + + #[test] + fn test_enum_if_modifier_resolves_in_written_scope() { + let mut effective_components: IndexMap = IndexMap::default(); + let mut tau2 = ast::Component { + name: "tau2".to_string(), + type_name: make_name("Real"), + ..ast::Component::empty_with_span(test_span()) + }; + tau2.binding = Some(make_int_expr(30)); + effective_components.insert("tau2".to_string(), tau2); + + let mut eps = ast::Component { + name: "eps".to_string(), + type_name: make_name("Real"), + ..ast::Component::empty_with_span(test_span()) + }; + eps.binding = Some(make_int_expr(0)); + effective_components.insert("eps".to_string(), eps); + + let mut energy_dynamics = ast::Component { + name: "energyDynamics".to_string(), + type_name: make_name("Dynamics"), + ..ast::Component::empty_with_span(test_span()) + }; + energy_dynamics.binding = Some(make_comp_ref_expr(&[ + "Modelica", + "Fluid", + "Types", + "Dynamics", + "DynamicFreeInitial", + ])); + effective_components.insert("energyDynamics".to_string(), energy_dynamics); + + let expr = ast::Expression::If { + branches: vec![( + ast::Expression::Binary { + op: rumoca_core::OpBinary::Gt, + lhs: Arc::new(make_comp_ref_expr(&["tau2"])), + rhs: Arc::new(make_comp_ref_expr(&["eps"])), + span: rumoca_core::Span::DUMMY, + }, + make_comp_ref_expr(&["energyDynamics"]), + )], + else_branch: Arc::new(make_comp_ref_expr(&[ + "Modelica", + "Fluid", + "Types", + "Dynamics", + "SteadyState", + ])), + span: rumoca_core::Span::DUMMY, + }; + + let resolved = resolve_modification_expr( + &expr, + &ast::ModificationEnvironment::default(), + &effective_components, + &ast::ClassTree::default(), + false, + ) + .expect("enum modifier should resolve in the written scope"); + + assert_eq!( + resolved, + make_comp_ref_expr(&[ + "Modelica", + "Fluid", + "Types", + "Dynamics", + "DynamicFreeInitial", + ]) + ); + } + + #[test] + fn test_modifier_with_same_resolved_value_keeps_existing_source_scope() { + let mut ctx = InstantiateContext::new(); + let key = ast::QualifiedName::from_ident("frictionParameters"); + let forwarded_value = make_comp_ref_expr(&["aimcData", "frictionParameters"]); + let forwarded_scope = Some(ast::QualifiedName::new()); + + ctx.mod_env_mut().active.insert( + key.clone(), + ast::ModificationValue::with_source_scope( + forwarded_value.clone(), + Some(forwarded_value.clone()), + forwarded_scope.clone(), + ), + ); + + let parent_snapshot = ctx.mod_env().active.clone(); + let shifted_parent_keys: IndexMap = IndexMap::default(); + let insert_ctx = ScopedInsertContext { + parent_snapshot: &parent_snapshot, + shifted_parent_keys: &shifted_parent_keys, + source_scope: Some(ast::QualifiedName::from_ident("aimc")), + }; + + insert_modifier_value_with_structural_overrides( + &mut ctx, + "frictionParameters", + &forwarded_value, + ModifierInsertOptions { + allow_string_eval: false, + prefixes: ModifierPrefixes::default(), + }, + &IndexMap::default(), + &ast::ClassTree::default(), + &insert_ctx, + ) + .expect("same-value modifier insertion should succeed"); + + let stored = ctx + .mod_env() + .get(&key) + .expect("modifier binding should exist"); + assert_eq!( + stored.source_scope, forwarded_scope, + "resolved multi-part modifier should inherit source scope from existing parent binding" + ); + } + + #[test] + fn test_propagate_record_binding_overrides_non_targeted_field_values() { + let mut nested_record = ast::ClassDef { + name: make_token("State"), + class_type: rumoca_core::ClassType::Record, + ..Default::default() + }; + nested_record.components.insert( + "phase".to_string(), + ast::Component::empty_with_span(test_span()), + ); + nested_record.components.insert( + "p".to_string(), + ast::Component::empty_with_span(test_span()), + ); + + let mut ctx = InstantiateContext::new(); + ctx.mod_env_mut().add( + ast::QualifiedName::from_ident("phase"), + ast::ModificationValue::simple(make_int_expr(7)), + ); + + let binding_expr = make_comp_ref_expr(&["state_in"]); + let targeted_keys: IndexMap = IndexMap::default(); + propagate_record_binding_to_fields( + &ast::ClassTree::default(), + &mut ctx, + &binding_expr, + None, + &nested_record, + &targeted_keys, + ) + .expect("record field projection should succeed"); + + let phase_mod = ctx + .mod_env() + .active + .get(&ast::QualifiedName::from_ident("phase")) + .expect("phase field binding should be present"); + match &phase_mod.value { + ast::Expression::FieldAccess { base, field, .. } => { + assert_eq!(field, "phase"); + match base.as_ref() { + ast::Expression::ComponentReference(cref) => { + assert_eq!(cref.parts.len(), 1); + assert_eq!(cref.parts[0].ident.text.as_ref(), "state_in"); + } + _ => panic!("field binding should project from record binding expression"), + } + } + _ => panic!("phase field should be rebound from record binding"), + } + } + + #[test] + fn test_propagate_record_binding_preserves_targeted_field_modifiers() { + let mut nested_record = ast::ClassDef { + name: make_token("State"), + class_type: rumoca_core::ClassType::Record, + ..Default::default() + }; + nested_record.components.insert( + "phase".to_string(), + ast::Component::empty_with_span(test_span()), + ); + + let mut ctx = InstantiateContext::new(); + let phase_qn = ast::QualifiedName::from_ident("phase"); + ctx.mod_env_mut().add( + phase_qn.clone(), + ast::ModificationValue::simple(make_int_expr(42)), + ); + + let mut targeted_keys: IndexMap = IndexMap::default(); + targeted_keys.insert(phase_qn.clone(), ()); + let binding_expr = make_comp_ref_expr(&["state_in"]); + propagate_record_binding_to_fields( + &ast::ClassTree::default(), + &mut ctx, + &binding_expr, + None, + &nested_record, + &targeted_keys, + ) + .expect("record field projection should succeed"); + + let phase_mod = ctx + .mod_env() + .active + .get(&phase_qn) + .expect("targeted phase modifier should still be present"); + match &phase_mod.value { + ast::Expression::Terminal { token, .. } => assert_eq!(token.text.as_ref(), "42"), + _ => panic!("targeted field modifier should not be replaced"), + } + } + + #[test] + fn test_propagate_record_binding_projects_if_expression_branches_per_field() { + let mut nested_record = ast::ClassDef { + name: make_token("CellData"), + class_type: rumoca_core::ClassType::Record, + ..Default::default() + }; + nested_record.components.insert( + "OCV_SOC".to_string(), + ast::Component::empty_with_span(test_span()), + ); + + let mut ctx = InstantiateContext::new(); + let binding_expr = ast::Expression::If { + branches: vec![( + make_comp_ref_expr(&["isDegraded"]), + make_comp_ref_expr(&["cellDataDegraded"]), + )], + else_branch: Arc::new(make_comp_ref_expr(&["cellDataOriginal"])), + span: rumoca_core::Span::DUMMY, + }; + let targeted_keys: IndexMap = IndexMap::default(); + propagate_record_binding_to_fields( + &ast::ClassTree::default(), + &mut ctx, + &binding_expr, + None, + &nested_record, + &targeted_keys, + ) + .expect("record field projection should succeed"); + + let field_mod = ctx + .mod_env() + .active + .get(&ast::QualifiedName::from_ident("OCV_SOC")) + .expect("OCV_SOC field binding should be present"); + let ast::Expression::If { + branches, + else_branch, + .. + } = &field_mod.value + else { + panic!("field projection should preserve if-expression structure"); + }; + assert_eq!(branches.len(), 1); + let (_cond, then_expr) = &branches[0]; + let ast::Expression::FieldAccess { base, field, .. } = then_expr else { + panic!("then-branch should project field access"); + }; + assert_eq!(field, "OCV_SOC"); + assert_eq!(*base.as_ref(), make_comp_ref_expr(&["cellDataDegraded"])); + + let ast::Expression::FieldAccess { + base: else_base, + field: else_field, + .. + } = else_branch.as_ref() + else { + panic!("else-branch should project field access"); + }; + assert_eq!(else_field, "OCV_SOC"); + assert_eq!( + *else_base.as_ref(), + make_comp_ref_expr(&["cellDataOriginal"]) + ); + } + + #[test] + fn test_propagate_record_binding_preserves_matching_default_record_constructor() { + let mut nested_record = ast::ClassDef { + name: make_token("BaseData"), + class_type: rumoca_core::ClassType::Record, + ..Default::default() + }; + nested_record.components.insert( + "mu_i".to_string(), + ast::Component { + binding: Some(make_int_expr(1)), + start: make_int_expr(1), + ..ast::Component::empty_with_span(test_span()) + }, + ); + + let mut ctx = InstantiateContext::new(); + let binding_expr = ast::Expression::FunctionCall { + comp: ast::ComponentReference { + local: false, + parts: vec![ast::ComponentRefPart { + ident: make_token("BaseData"), + subs: None, + }], + def_id: None, + span: rumoca_core::Span::DUMMY, + }, + args: Vec::new(), + span: rumoca_core::Span::DUMMY, + }; + + propagate_record_binding_to_fields( + &ast::ClassTree::default(), + &mut ctx, + &binding_expr, + None, + &nested_record, + &IndexMap::default(), + ) + .expect("record field projection should succeed"); + + assert!( + ctx.mod_env().active.is_empty(), + "matching zero-argument record constructors should preserve declared defaults" + ); + } + + #[test] + fn test_propagate_record_binding_projects_subtype_default_record_constructor_fields() { + let mut nested_record = ast::ClassDef { + name: make_token("BaseData"), + class_type: rumoca_core::ClassType::Record, + ..Default::default() + }; + nested_record.components.insert( + "mu_i".to_string(), + ast::Component { + binding: Some(make_int_expr(1)), + start: make_int_expr(1), + ..ast::Component::empty_with_span(test_span()) + }, + ); + + let mut ctx = InstantiateContext::new(); + let binding_expr = ast::Expression::FunctionCall { + comp: ast::ComponentReference { + local: false, + parts: vec![ast::ComponentRefPart { + ident: make_token("M350_50A"), + subs: None, + }], + def_id: None, + span: rumoca_core::Span::DUMMY, + }, + args: Vec::new(), + span: rumoca_core::Span::DUMMY, + }; + + propagate_record_binding_to_fields( + &ast::ClassTree::default(), + &mut ctx, + &binding_expr, + None, + &nested_record, + &IndexMap::default(), + ) + .expect("record field projection should succeed"); + + let field_mod = ctx + .mod_env() + .active + .get(&ast::QualifiedName::from_ident("mu_i")) + .expect("subtype default record constructor should project field binding"); + let ast::Expression::FieldAccess { base, field, .. } = &field_mod.value else { + panic!("subtype constructor field should be projected"); + }; + assert_eq!(field, "mu_i"); + assert_eq!(base.as_ref(), &binding_expr); + } + + #[test] + fn test_propagate_record_binding_projects_through_unique_constructor_record_field() { + let inner_def_id = rumoca_core::DefId::new(1001); + let outer_def_id = rumoca_core::DefId::new(1002); + let mut inner_record = ast::ClassDef { + name: make_token("Inner"), + class_type: rumoca_core::ClassType::Record, + def_id: Some(inner_def_id), + ..Default::default() + }; + inner_record.components.insert( + "x".to_string(), + ast::Component::empty_with_span(test_span()), + ); + + let mut outer_record = ast::ClassDef { + name: make_token("Outer"), + class_type: rumoca_core::ClassType::Record, + def_id: Some(outer_def_id), + ..Default::default() + }; + outer_record.components.insert( + "innerParams".to_string(), + ast::Component { + type_name: ast::Name { + name: vec![make_token("Pkg"), make_token("Inner")], + def_id: Some(inner_def_id), + }, + type_def_id: Some(inner_def_id), + ..ast::Component::empty_with_span(test_span()) + }, + ); + + let mut tree = ast::ClassTree::default(); + tree.definitions + .classes + .insert("Pkg.Outer".to_string(), outer_record); + tree.def_map.insert(inner_def_id, "Pkg.Inner".to_string()); + tree.def_map.insert(outer_def_id, "Pkg.Outer".to_string()); + + let mut ctx = InstantiateContext::new(); + let binding_expr = ast::Expression::FunctionCall { + comp: ast::ComponentReference { + local: false, + parts: vec![ast::ComponentRefPart { + ident: make_token("Outer"), + subs: None, + }], + def_id: None, + span: rumoca_core::Span::DUMMY, + }, + args: Vec::new(), + span: rumoca_core::Span::DUMMY, + }; + + propagate_record_binding_to_fields( + &tree, + &mut ctx, + &binding_expr, + Some(ast::QualifiedName::from_ident("Pkg")), + &inner_record, + &IndexMap::default(), + ) + .expect("record field projection should succeed"); + + let field_mod = ctx + .mod_env() + .active + .get(&ast::QualifiedName::from_ident("x")) + .expect("inner field binding should be present"); + let ast::Expression::FieldAccess { + base, + field: inner_field, + .. + } = &field_mod.value + else { + panic!("inner field should be projected"); + }; + assert_eq!(inner_field, "x"); + let ast::Expression::FieldAccess { + base: constructor, + field: outer_field, + .. + } = base.as_ref() + else { + panic!("projection should first select the unique compatible record field"); + }; + assert_eq!(outer_field, "innerParams"); + assert_eq!(constructor.as_ref(), &binding_expr); + } + + #[test] + fn test_propagate_record_binding_preserves_function_return_record_field_access() { + let state_def_id = rumoca_core::DefId::new(1101); + let set_state_def_id = rumoca_core::DefId::new(1102); + let mut state_record = ast::ClassDef { + name: make_token("ThermodynamicState"), + class_type: rumoca_core::ClassType::Record, + def_id: Some(state_def_id), + ..Default::default() + }; + state_record.components.insert( + "p".to_string(), + ast::Component::empty_with_span(test_span()), + ); + state_record.components.insert( + "T".to_string(), + ast::Component::empty_with_span(test_span()), + ); + state_record.components.insert( + "X".to_string(), + ast::Component::empty_with_span(test_span()), + ); + + let mut set_state_function = ast::ClassDef { + name: make_token("setState_pTX"), + class_type: rumoca_core::ClassType::Function, + def_id: Some(set_state_def_id), + ..Default::default() + }; + set_state_function.components.insert( + "p".to_string(), + ast::Component::empty_with_span(test_span()), + ); + set_state_function.components.insert( + "T".to_string(), + ast::Component::empty_with_span(test_span()), + ); + set_state_function.components.insert( + "X".to_string(), + ast::Component::empty_with_span(test_span()), + ); + + let mut tree = ast::ClassTree::default(); + tree.definitions.classes.insert( + "Medium.ThermodynamicState".to_string(), + state_record.clone(), + ); + tree.definitions + .classes + .insert("Medium.setState_pTX".to_string(), set_state_function); + tree.def_map + .insert(state_def_id, "Medium.ThermodynamicState".to_string()); + tree.def_map + .insert(set_state_def_id, "Medium.setState_pTX".to_string()); + + let binding_expr = ast::Expression::FunctionCall { + comp: ast::ComponentReference { + local: false, + parts: vec![ + ast::ComponentRefPart { + ident: make_token("Medium"), + subs: None, + }, + ast::ComponentRefPart { + ident: make_token("setState_pTX"), + subs: None, + }, + ], + def_id: Some(set_state_def_id), + span: rumoca_core::Span::DUMMY, + }, + args: vec![ + named_arg("p", make_int_expr(101325)), + named_arg("T", make_int_expr(293)), + named_arg("X", make_comp_ref_expr(&["X_default"])), + ], + span: rumoca_core::Span::DUMMY, + }; + + let mut ctx = InstantiateContext::new(); + propagate_record_binding_to_fields( + &tree, + &mut ctx, + &binding_expr, + Some(ast::QualifiedName::from_ident("Medium")), + &state_record, + &IndexMap::default(), + ) + .expect("record field projection should succeed"); + + let x_mod = ctx + .mod_env() + .active + .get(&ast::QualifiedName::from_ident("X")) + .expect("X field binding should be present"); + let ast::Expression::FieldAccess { base, field, .. } = &x_mod.value else { + panic!("function-return record field must be kept as field access"); + }; + assert_eq!(field, "X"); + assert_eq!(base.as_ref(), &binding_expr); + } + + #[test] + fn test_propagate_record_binding_skips_non_record_classes() { + let mut nested_block = ast::ClassDef { + name: make_token("UniformNoise"), + class_type: rumoca_core::ClassType::Block, + ..Default::default() + }; + nested_block.components.insert( + "y".to_string(), + ast::Component::empty_with_span(test_span()), + ); + nested_block.components.insert( + "seedState".to_string(), + ast::Component::empty_with_span(test_span()), + ); + + let mut ctx = InstantiateContext::new(); + let binding_expr = make_comp_ref_expr(&["noise"]); + let targeted_keys: IndexMap = IndexMap::default(); + propagate_record_binding_to_fields( + &ast::ClassTree::default(), + &mut ctx, + &binding_expr, + None, + &nested_block, + &targeted_keys, + ) + .expect("non-record projection should succeed without mutation"); + + assert!( + ctx.mod_env().active.is_empty(), + "non-record class modifiers must not synthesize per-field record bindings" + ); + } +} diff --git a/crates/rumoca-phase-instantiate/src/nested_scope.rs b/crates/rumoca-phase-instantiate/src/nested_scope.rs index 92d6b7624..bec359337 100644 --- a/crates/rumoca-phase-instantiate/src/nested_scope.rs +++ b/crates/rumoca-phase-instantiate/src/nested_scope.rs @@ -171,19 +171,32 @@ pub(super) fn shift_modifications_down(ctx: &mut InstantiateContext, comp_name: /// Remap a class-redeclare modifier target to the active enclosing override. /// -/// MLS §7.3: `redeclare package Medium = Medium` inside component modifiers should -/// forward to the enclosing class's active `Medium` redeclare, not the local default. +/// MLS §7.3: `redeclare package Medium = Medium` and +/// `redeclare package Medium = MediumAir` inside component modifiers should +/// forward through active enclosing package aliases, not the local replaceable +/// defaults declared on the component class. +fn class_modification_expr(mod_expr: &ast::Expression) -> Option<&ast::Expression> { + match mod_expr { + ast::Expression::ClassModification { .. } => Some(mod_expr), + ast::Expression::Modification { value, .. } => class_modification_expr(value), + _ => None, + } +} + pub(super) fn remap_redeclare_class_modifier( mod_expr: &ast::Expression, target_name: &str, type_overrides: &TypeOverrideMap, ) -> ast::Expression { + let Some(class_mod) = class_modification_expr(mod_expr) else { + return mod_expr.clone(); + }; let ast::Expression::ClassModification { target, modifications, span, .. - } = mod_expr + } = class_mod else { return mod_expr.clone(); }; @@ -191,11 +204,17 @@ pub(super) fn remap_redeclare_class_modifier( let Some(last) = target.parts.last() else { return mod_expr.clone(); }; - if last.ident.text.as_ref() != target_name { - return mod_expr.clone(); - } + let rhs_name = last.ident.text.as_ref(); + let override_name = if rhs_name == target_name { + target_name + } else { + rhs_name + }; - let Some(override_def_id) = type_overrides.target_for_reference(target) else { + let Some(override_def_id) = type_overrides + .target_for_reference(target) + .or_else(|| type_overrides.target_for_alias_name(override_name)) + else { return mod_expr.clone(); }; if target.def_id == Some(override_def_id) { @@ -204,18 +223,34 @@ pub(super) fn remap_redeclare_class_modifier( let mut remapped_target = target.clone(); remapped_target.def_id = Some(override_def_id); - ast::Expression::ClassModification { + let remapped_class_mod = ast::Expression::ClassModification { target: remapped_target, modifications: modifications.clone(), each_flags: Vec::new(), final_flags: Vec::new(), redeclare_flags: Vec::new(), span: *span, + }; + + match mod_expr { + ast::Expression::Modification { + target: outer_target, + span: outer_span, + .. + } => ast::Expression::Modification { + target: outer_target.clone(), + value: std::sync::Arc::new(remapped_class_mod), + span: *outer_span, + }, + _ => remapped_class_mod, } } fn is_self_forwarding_redeclare(mod_expr: &ast::Expression, target_name: &str) -> bool { - let ast::Expression::ClassModification { target, .. } = mod_expr else { + let Some(class_mod) = class_modification_expr(mod_expr) else { + return false; + }; + let ast::Expression::ClassModification { target, .. } = class_mod else { return false; }; target @@ -237,7 +272,7 @@ pub(super) fn resolve_component_nested_type_overrides( type_overrides: &TypeOverrideMap, ) -> InstantiateResult { let mut class_overrides = - extract_component_class_overrides(tree, comp, class_def, Some(mod_env))?; + extract_component_class_overrides(tree, comp, class_def, Some(mod_env), type_overrides)?; let mut has_forwarding_class_redeclare = false; if let Some(target_class) = class_def { @@ -290,7 +325,8 @@ pub(super) fn resolve_component_nested_type_overrides( } fn class_redeclare_target_ref(mod_expr: &ast::Expression) -> Option { - let ast::Expression::ClassModification { target, .. } = mod_expr else { + let class_mod = class_modification_expr(mod_expr)?; + let ast::Expression::ClassModification { target, .. } = class_mod else { return None; }; Some(target.clone()) diff --git a/crates/rumoca-phase-instantiate/src/path_utils.rs b/crates/rumoca-phase-instantiate/src/path_utils.rs index a839c4f33..d1fa91c71 100644 --- a/crates/rumoca-phase-instantiate/src/path_utils.rs +++ b/crates/rumoca-phase-instantiate/src/path_utils.rs @@ -15,3 +15,25 @@ pub(crate) fn class_scope_split(name: &str) -> Option<(&str, &str)> { pub(crate) fn class_name_leaf(name: &str) -> &str { rumoca_core::top_level_last_segment(name) } + +/// Build a component reference expression from a dotted path. +pub(crate) fn component_ref_expr_from_dotted( + value: &str, + span: rumoca_core::Span, +) -> rumoca_ir_ast::Expression { + rumoca_ir_ast::Expression::ComponentReference(rumoca_ir_ast::ComponentReference { + local: false, + parts: rumoca_core::split_path_with_indices(value) + .into_iter() + .map(|part| rumoca_ir_ast::ComponentRefPart { + ident: rumoca_core::Token { + text: part.to_string().into(), + ..Default::default() + }, + subs: None, + }) + .collect(), + def_id: None, + span, + }) +} diff --git a/crates/rumoca-phase-instantiate/src/tests.rs b/crates/rumoca-phase-instantiate/src/tests.rs index a4baa42b6..0b0cdff6e 100644 --- a/crates/rumoca-phase-instantiate/src/tests.rs +++ b/crates/rumoca-phase-instantiate/src/tests.rs @@ -1,4 +1,5 @@ use super::*; +use rumoca_eval_ast::eval_instantiate::extract_int_params_with_mods; use rumoca_ir_ast as ast; /// Helper to create a token with text for testing. @@ -60,6 +61,14 @@ fn make_bool_expr(value: bool) -> ast::Expression { } } +fn make_real_expr(value: &str) -> ast::Expression { + ast::Expression::Terminal { + terminal_type: ast::TerminalType::UnsignedReal, + token: make_token(value), + span: rumoca_core::Span::DUMMY, + } +} + fn make_if_expr( condition: ast::Expression, then_expr: ast::Expression, @@ -217,7 +226,7 @@ fn test_extract_attributes_preserves_local_fixed_with_local_start() { let mod_env = ast::ModificationEnvironment::new(); let effective_components = IndexMap::default(); let eval_ctx = make_eval_ctx(&tree, &mod_env, &effective_components); - let attrs = extract_attributes(&comp, &mod_env, "x", &eval_ctx, &[]) + let attrs = extract_attributes(&comp, &mod_env, "x", "x", &eval_ctx, &[]) .expect("valid attributes should extract"); assert!(attrs.start.is_some()); @@ -239,7 +248,7 @@ fn test_extract_attributes_preserves_local_fixed_with_outer_start() { let tree = ast::ClassTree::default(); let effective_components = IndexMap::default(); let eval_ctx = make_eval_ctx(&tree, &mod_env, &effective_components); - let attrs = extract_attributes(&comp, &mod_env, "x", &eval_ctx, &[]) + let attrs = extract_attributes(&comp, &mod_env, "x", "x", &eval_ctx, &[]) .expect("valid attributes should extract"); assert!(attrs.start.is_some()); @@ -263,7 +272,7 @@ fn test_extract_attributes_outer_state_select_overrides_local() { let tree = ast::ClassTree::default(); let effective_components = IndexMap::default(); let eval_ctx = make_eval_ctx(&tree, &mod_env, &effective_components); - let attrs = extract_attributes(&comp, &mod_env, "x", &eval_ctx, &[]) + let attrs = extract_attributes(&comp, &mod_env, "x", "x", &eval_ctx, &[]) .expect("valid attributes should extract"); assert_eq!(attrs.state_select, rumoca_core::StateSelect::Never); @@ -300,7 +309,7 @@ fn test_extract_attributes_evaluates_outer_state_select_in_modifier_source_scope let tree = ast::ClassTree::default(); let eval_ctx = make_eval_ctx(&tree, &mod_env, &effective_components); - let attrs = extract_attributes(&comp, &mod_env, "x", &eval_ctx, &[]) + let attrs = extract_attributes(&comp, &mod_env, "x", "x", &eval_ctx, &[]) .expect("source-scoped stateSelect should evaluate"); assert_eq!(attrs.state_select, rumoca_core::StateSelect::Prefer); @@ -322,12 +331,431 @@ fn test_extract_attributes_evaluates_state_select_parameter() { let tree = ast::ClassTree::default(); let mod_env = ast::ModificationEnvironment::new(); let eval_ctx = make_eval_ctx(&tree, &mod_env, &effective_components); - let attrs = extract_attributes(&comp, &mod_env, "x", &eval_ctx, &[]) + let attrs = extract_attributes(&comp, &mod_env, "x", "x", &eval_ctx, &[]) .expect("valid attributes should extract"); assert_eq!(attrs.state_select, rumoca_core::StateSelect::Prefer); } +#[test] +fn test_extract_attributes_evaluates_scoped_state_select_condition_parameter() { + let mut comp = make_component("x", "Real", None); + comp.modifications.insert( + "stateSelect".to_string(), + ast::Expression::If { + branches: vec![( + make_comp_ref_expr(&["medium", "preferredMediumStates"]), + make_comp_ref_expr(&["StateSelect", "prefer"]), + )], + else_branch: std::sync::Arc::new(make_comp_ref_expr(&["StateSelect", "default"])), + span: rumoca_core::Span::DUMMY, + }, + ); + + let mut preferred = make_component("preferredMediumStates", "Boolean", None); + preferred.binding = Some(make_bool_expr(true)); + let mut effective_components = IndexMap::default(); + effective_components.insert("preferredMediumStates".to_string(), preferred); + + let tree = ast::ClassTree::default(); + let mod_env = ast::ModificationEnvironment::new(); + let eval_ctx = make_eval_ctx(&tree, &mod_env, &effective_components); + let attrs = extract_attributes(&comp, &mod_env, "x", "x", &eval_ctx, &[]) + .expect("stateSelect if-expression should evaluate from scoped parameter context"); + + assert_eq!(attrs.state_select, rumoca_core::StateSelect::Prefer); +} + +#[test] +fn test_extract_attributes_evaluates_state_select_enum_equality_condition() { + let mut comp = make_component("x", "Real", None); + comp.modifications.insert( + "stateSelect".to_string(), + ast::Expression::If { + branches: vec![( + make_binary_expr( + rumoca_core::OpBinary::Eq, + make_comp_ref_expr(&["massDynamics"]), + make_comp_ref_expr(&["Modelica", "Fluid", "Types", "Dynamics", "SteadyState"]), + ), + make_comp_ref_expr(&["StateSelect", "default"]), + )], + else_branch: std::sync::Arc::new(make_comp_ref_expr(&["StateSelect", "prefer"])), + span: rumoca_core::Span::DUMMY, + }, + ); + + let mut mass_dynamics = make_component("massDynamics", "Dynamics", None); + mass_dynamics.binding = Some(make_comp_ref_expr(&[ + "Modelica", + "Fluid", + "Types", + "Dynamics", + "SteadyState", + ])); + let mut effective_components = IndexMap::default(); + effective_components.insert("massDynamics".to_string(), mass_dynamics); + + let tree = ast::ClassTree::default(); + let mod_env = ast::ModificationEnvironment::new(); + let eval_ctx = make_eval_ctx(&tree, &mod_env, &effective_components); + let attrs = extract_attributes(&comp, &mod_env, "x", "x", &eval_ctx, &[]) + .expect("stateSelect if-expression should evaluate enum equality condition"); + + assert_eq!(attrs.state_select, rumoca_core::StateSelect::Default); +} + +#[test] +fn test_extract_attributes_evaluates_state_select_chained_enum_parameter_condition() { + let mut comp = make_component("x", "Real", None); + comp.modifications.insert( + "stateSelect".to_string(), + ast::Expression::If { + branches: vec![( + make_binary_expr( + rumoca_core::OpBinary::Eq, + make_comp_ref_expr(&["massDynamics"]), + make_comp_ref_expr(&["Modelica", "Fluid", "Types", "Dynamics", "SteadyState"]), + ), + make_comp_ref_expr(&["StateSelect", "default"]), + )], + else_branch: std::sync::Arc::new(make_comp_ref_expr(&["StateSelect", "prefer"])), + span: rumoca_core::Span::DUMMY, + }, + ); + + let mut mass_dynamics = make_component("massDynamics", "Dynamics", None); + mass_dynamics.binding = Some(make_comp_ref_expr(&["energyDynamics"])); + let mut energy_dynamics = make_component("energyDynamics", "Dynamics", None); + energy_dynamics.binding = Some(make_comp_ref_expr(&[ + "Modelica", + "Fluid", + "Types", + "Dynamics", + "SteadyState", + ])); + let mut effective_components = IndexMap::default(); + effective_components.insert("massDynamics".to_string(), mass_dynamics); + effective_components.insert("energyDynamics".to_string(), energy_dynamics); + + let tree = ast::ClassTree::default(); + let mod_env = ast::ModificationEnvironment::new(); + let eval_ctx = make_eval_ctx(&tree, &mod_env, &effective_components); + let attrs = extract_attributes(&comp, &mod_env, "x", "x", &eval_ctx, &[]) + .expect("stateSelect if-expression should follow chained enum parameters"); + + assert_eq!(attrs.state_select, rumoca_core::StateSelect::Default); +} + +#[test] +fn test_extract_attributes_evaluates_state_select_enum_if_with_real_condition() { + let mut comp = make_component("x", "Real", None); + comp.modifications.insert( + "stateSelect".to_string(), + ast::Expression::If { + branches: vec![( + make_binary_expr( + rumoca_core::OpBinary::Eq, + make_comp_ref_expr(&["massDynamics"]), + make_comp_ref_expr(&["Modelica", "Fluid", "Types", "Dynamics", "SteadyState"]), + ), + make_comp_ref_expr(&["StateSelect", "default"]), + )], + else_branch: std::sync::Arc::new(make_comp_ref_expr(&["StateSelect", "prefer"])), + span: rumoca_core::Span::DUMMY, + }, + ); + + let mut mass_dynamics = make_component("massDynamics", "Dynamics", None); + mass_dynamics.binding = Some(ast::Expression::If { + branches: vec![( + make_binary_expr( + rumoca_core::OpBinary::Gt, + make_comp_ref_expr(&["tau"]), + make_comp_ref_expr(&["Modelica", "Constants", "eps"]), + ), + make_comp_ref_expr(&["energyDynamics"]), + )], + else_branch: std::sync::Arc::new(make_comp_ref_expr(&[ + "Modelica", + "Fluid", + "Types", + "Dynamics", + "SteadyState", + ])), + span: rumoca_core::Span::DUMMY, + }); + let mut energy_dynamics = make_component("energyDynamics", "Dynamics", None); + energy_dynamics.binding = Some(make_comp_ref_expr(&[ + "Modelica", + "Fluid", + "Types", + "Dynamics", + "DynamicFreeInitial", + ])); + let mut tau = make_component("tau", "Real", None); + tau.binding = Some(make_real_expr("30.0")); + let mut eps = make_component("Modelica.Constants.eps", "Real", None); + eps.binding = Some(make_real_expr("1e-15")); + + let mut effective_components = IndexMap::default(); + effective_components.insert("massDynamics".to_string(), mass_dynamics); + effective_components.insert("energyDynamics".to_string(), energy_dynamics); + effective_components.insert("tau".to_string(), tau); + effective_components.insert("Modelica.Constants.eps".to_string(), eps); + + let tree = ast::ClassTree::default(); + let mod_env = ast::ModificationEnvironment::new(); + let eval_ctx = make_eval_ctx(&tree, &mod_env, &effective_components); + let attrs = extract_attributes(&comp, &mod_env, "x", "x", &eval_ctx, &[]) + .expect("stateSelect should evaluate enum if-expression with Real condition"); + + assert_eq!(attrs.state_select, rumoca_core::StateSelect::Prefer); +} + +#[test] +fn test_extract_attributes_evaluates_state_select_in_component_parent_scope() { + let mut comp = make_component("m", "Real", None); + comp.modifications.insert( + "stateSelect".to_string(), + ast::Expression::If { + branches: vec![( + make_binary_expr( + rumoca_core::OpBinary::Eq, + make_comp_ref_expr(&["massDynamics"]), + make_comp_ref_expr(&["Modelica", "Fluid", "Types", "Dynamics", "SteadyState"]), + ), + make_comp_ref_expr(&["StateSelect", "default"]), + )], + else_branch: std::sync::Arc::new(make_comp_ref_expr(&["StateSelect", "prefer"])), + span: rumoca_core::Span::DUMMY, + }, + ); + + let mut mass_dynamics = make_component("dynBal.massDynamics", "Dynamics", None); + mass_dynamics.binding = Some(ast::Expression::If { + branches: vec![( + make_binary_expr( + rumoca_core::OpBinary::Gt, + make_comp_ref_expr(&["tau"]), + make_comp_ref_expr(&["Modelica", "Constants", "eps"]), + ), + make_comp_ref_expr(&["energyDynamics"]), + )], + else_branch: std::sync::Arc::new(make_comp_ref_expr(&[ + "Modelica", + "Fluid", + "Types", + "Dynamics", + "SteadyState", + ])), + span: rumoca_core::Span::DUMMY, + }); + let mut energy_dynamics = make_component("dynBal.energyDynamics", "Dynamics", None); + energy_dynamics.binding = Some(make_comp_ref_expr(&[ + "Modelica", + "Fluid", + "Types", + "Dynamics", + "DynamicFreeInitial", + ])); + let mut tau = make_component("dynBal.tau", "Real", None); + tau.binding = Some(make_real_expr("30.0")); + let mut eps = make_component("Modelica.Constants.eps", "Real", None); + eps.binding = Some(make_real_expr("1e-15")); + + let mut effective_components = IndexMap::default(); + effective_components.insert("dynBal.massDynamics".to_string(), mass_dynamics); + effective_components.insert("dynBal.energyDynamics".to_string(), energy_dynamics); + effective_components.insert("dynBal.tau".to_string(), tau); + effective_components.insert("Modelica.Constants.eps".to_string(), eps); + + let tree = ast::ClassTree::default(); + let mod_env = ast::ModificationEnvironment::new(); + let eval_ctx = make_eval_ctx(&tree, &mod_env, &effective_components); + let dyn_bal_scope = ast::QualifiedName::from_ident("dynBal"); + let attrs = extract_attributes_in_scope( + &comp, + &mod_env, + "m", + "dynBal.m", + &eval_ctx, + &[], + Some(&dyn_bal_scope), + ) + .expect("stateSelect should resolve sibling parameters in the component parent scope"); + + assert_eq!(attrs.state_select, rumoca_core::StateSelect::Prefer); +} + +#[test] +fn test_extract_attributes_preserves_parent_scope_through_local_sibling_binding() { + let mut comp = make_component("m", "Real", None); + comp.modifications.insert( + "stateSelect".to_string(), + ast::Expression::If { + branches: vec![( + make_binary_expr( + rumoca_core::OpBinary::Eq, + make_comp_ref_expr(&["massDynamics"]), + make_comp_ref_expr(&["Modelica", "Fluid", "Types", "Dynamics", "SteadyState"]), + ), + make_comp_ref_expr(&["StateSelect", "default"]), + )], + else_branch: std::sync::Arc::new(make_comp_ref_expr(&["StateSelect", "prefer"])), + span: rumoca_core::Span::DUMMY, + }, + ); + + let mut mass_dynamics = make_component("massDynamics", "Dynamics", None); + mass_dynamics.binding = Some(ast::Expression::If { + branches: vec![( + make_binary_expr( + rumoca_core::OpBinary::Gt, + make_comp_ref_expr(&["tau"]), + make_comp_ref_expr(&["Modelica", "Constants", "eps"]), + ), + make_comp_ref_expr(&["energyDynamics"]), + )], + else_branch: std::sync::Arc::new(make_comp_ref_expr(&[ + "Modelica", + "Fluid", + "Types", + "Dynamics", + "SteadyState", + ])), + span: rumoca_core::Span::DUMMY, + }); + let mut energy_dynamics = make_component("energyDynamics", "Dynamics", None); + energy_dynamics.binding = Some(make_comp_ref_expr(&[ + "Modelica", + "Fluid", + "Types", + "Dynamics", + "DynamicFreeInitial", + ])); + let mut tau = make_component("dynBal.tau", "Real", None); + tau.binding = Some(make_real_expr("30.0")); + let mut eps = make_component("Modelica.Constants.eps", "Real", None); + eps.binding = Some(make_real_expr("1e-15")); + + let mut effective_components = IndexMap::default(); + effective_components.insert("massDynamics".to_string(), mass_dynamics); + effective_components.insert("energyDynamics".to_string(), energy_dynamics); + effective_components.insert("dynBal.tau".to_string(), tau); + effective_components.insert("Modelica.Constants.eps".to_string(), eps); + + let tree = ast::ClassTree::default(); + let mod_env = ast::ModificationEnvironment::new(); + let eval_ctx = make_eval_ctx(&tree, &mod_env, &effective_components); + let dyn_bal_scope = ast::QualifiedName::from_ident("dynBal"); + let attrs = extract_attributes_in_scope( + &comp, + &mod_env, + "m", + "dynBal.m", + &eval_ctx, + &[], + Some(&dyn_bal_scope), + ) + .expect( + "stateSelect should keep parent scope while evaluating local sibling parameter bindings", + ); + + assert_eq!(attrs.state_select, rumoca_core::StateSelect::Prefer); +} + +#[test] +fn test_extract_attributes_resolves_structured_array_scope_modifiers() { + let mut comp = make_component("m", "Real", None); + comp.modifications.insert( + "stateSelect".to_string(), + ast::Expression::If { + branches: vec![( + make_binary_expr( + rumoca_core::OpBinary::Eq, + make_comp_ref_expr(&["massDynamics"]), + make_comp_ref_expr(&["Modelica", "Fluid", "Types", "Dynamics", "SteadyState"]), + ), + make_comp_ref_expr(&["StateSelect", "default"]), + )], + else_branch: std::sync::Arc::new(make_comp_ref_expr(&["StateSelect", "prefer"])), + span: rumoca_core::Span::DUMMY, + }, + ); + + let mut mod_env = ast::ModificationEnvironment::new(); + mod_env.add( + ast::QualifiedName { + parts: vec![ + ("dyn".to_string(), Vec::new()), + ("ch".to_string(), vec![1]), + ("massDynamics".to_string(), Vec::new()), + ], + }, + ast::ModificationValue::simple(ast::Expression::If { + branches: vec![( + make_binary_expr( + rumoca_core::OpBinary::Gt, + make_comp_ref_expr(&["tau"]), + make_comp_ref_expr(&["Modelica", "Constants", "eps"]), + ), + make_comp_ref_expr(&[ + "Modelica", + "Fluid", + "Types", + "Dynamics", + "DynamicFreeInitial", + ]), + )], + else_branch: std::sync::Arc::new(make_comp_ref_expr(&[ + "Modelica", + "Fluid", + "Types", + "Dynamics", + "SteadyState", + ])), + span: rumoca_core::Span::DUMMY, + }), + ); + mod_env.add( + ast::QualifiedName { + parts: vec![ + ("dyn".to_string(), Vec::new()), + ("ch".to_string(), vec![1]), + ("tau".to_string(), Vec::new()), + ], + }, + ast::ModificationValue::simple(make_real_expr("30.0")), + ); + + let mut eps = make_component("Modelica.Constants.eps", "Real", None); + eps.binding = Some(make_real_expr("1e-15")); + + let mut effective_components = IndexMap::default(); + effective_components.insert("Modelica.Constants.eps".to_string(), eps); + + let tree = ast::ClassTree::default(); + let eval_ctx = make_eval_ctx(&tree, &mod_env, &effective_components); + let array_scope = ast::QualifiedName { + parts: vec![("dyn".to_string(), Vec::new()), ("ch".to_string(), vec![1])], + }; + let attrs = extract_attributes_in_scope( + &comp, + &mod_env, + "m", + "dyn.ch[1].m", + &eval_ctx, + &[], + Some(&array_scope), + ) + .expect( + "stateSelect should resolve structured array-scope modifiers through flat scoped lookup", + ); + + assert_eq!(attrs.state_select, rumoca_core::StateSelect::Prefer); +} + #[test] fn test_lookup_type_info_accepts_predefined_state_select() { let tree = ast::ClassTree::new(); @@ -353,7 +781,7 @@ fn test_extract_attributes_rejects_invalid_state_select() { let mod_env = ast::ModificationEnvironment::new(); let effective_components = IndexMap::default(); let eval_ctx = make_eval_ctx(&tree, &mod_env, &effective_components); - let err = extract_attributes(&comp, &mod_env, "x", &eval_ctx, &[]) + let err = extract_attributes(&comp, &mod_env, "x", "x", &eval_ctx, &[]) .expect_err("invalid stateSelect literal should fail"); assert!(err.to_string().contains("stateSelect")); @@ -477,7 +905,7 @@ fn test_parameter_declaration_binding_promotes_builtin_default_start() { let mod_env = ast::ModificationEnvironment::new(); let effective_components = IndexMap::default(); let eval_ctx = make_eval_ctx(&tree, &mod_env, &effective_components); - let result = extract_component_attrs_and_binding(&comp, &mod_env, &eval_ctx, &[]) + let result = extract_component_attrs_and_binding(&comp, &mod_env, &comp.name, &eval_ctx, &[]) .expect("valid attributes should extract"); assert_eq!( @@ -487,6 +915,257 @@ fn test_parameter_declaration_binding_promotes_builtin_default_start() { ); } +#[test] +fn declaration_binding_source_for_flattening_preserves_single_part_ref() { + let original = make_comp_ref_expr(&["nNodes"]); + let resolved = make_int_expr(2); + let mut comp = make_component("n", "Integer", None); + comp.variability = rumoca_core::Variability::Parameter(make_token("parameter")); + comp.is_structural = true; + let mut n_nodes = make_component("nNodes", "Integer", None); + n_nodes.variability = rumoca_core::Variability::Parameter(make_token("parameter")); + let mut effective_components = IndexMap::default(); + effective_components.insert("nNodes".to_string(), n_nodes); + let mut mod_env = ast::ModificationEnvironment::new(); + mod_env.add( + ast::QualifiedName::from_ident("nNodes"), + ast::ModificationValue::with_source_scope_and_prefixes( + make_int_expr(20), + Some(make_comp_ref_expr(&["nNodes"])), + None, + false, + true, + ), + ); + + let source = declaration_binding_source_for_flattening( + &comp, + &original, + &resolved, + &effective_components, + &mod_env, + ) + .expect("final modified single-part sibling reference source should be preserved"); + + assert!(matches!( + source, + ast::Expression::ComponentReference(cref) if cref.to_string() == "nNodes" + )); +} + +#[test] +fn declaration_binding_source_for_flattening_skips_non_sibling_single_part_ref() { + let original = make_comp_ref_expr(&["nNodes"]); + let resolved = make_int_expr(2); + let comp = make_component("n", "Integer", None); + let mod_env = ast::ModificationEnvironment::new(); + + let source = declaration_binding_source_for_flattening( + &comp, + &original, + &resolved, + &IndexMap::default(), + &mod_env, + ); + + assert_eq!(source, None); +} + +#[test] +fn declaration_binding_source_for_flattening_skips_unmodified_sibling_single_part_ref() { + let original = make_comp_ref_expr(&["nNodes"]); + let resolved = make_int_expr(2); + let mut comp = make_component("n", "Integer", None); + comp.variability = rumoca_core::Variability::Parameter(make_token("parameter")); + comp.is_structural = true; + let mut n_nodes = make_component("nNodes", "Integer", None); + n_nodes.variability = rumoca_core::Variability::Parameter(make_token("parameter")); + let mut effective_components = IndexMap::default(); + effective_components.insert("nNodes".to_string(), n_nodes); + let mod_env = ast::ModificationEnvironment::new(); + + let source = declaration_binding_source_for_flattening( + &comp, + &original, + &resolved, + &effective_components, + &mod_env, + ); + + assert_eq!(source, None); +} + +#[test] +fn declaration_binding_source_for_flattening_preserves_literal_modified_single_part_ref() { + let original = make_comp_ref_expr(&["nNodes"]); + let resolved = make_int_expr(2); + let mut comp = make_component("n", "Integer", None); + comp.variability = rumoca_core::Variability::Parameter(make_token("parameter")); + comp.is_structural = true; + let mut n_nodes = make_component("nNodes", "Integer", None); + n_nodes.variability = rumoca_core::Variability::Parameter(make_token("parameter")); + let mut effective_components = IndexMap::default(); + effective_components.insert("nNodes".to_string(), n_nodes); + let mut mod_env = ast::ModificationEnvironment::new(); + mod_env.add( + ast::QualifiedName::from_ident("nNodes"), + ast::ModificationValue::with_source_scope_and_prefixes( + make_int_expr(20), + Some(make_int_expr(20)), + None, + false, + true, + ), + ); + + let source = declaration_binding_source_for_flattening( + &comp, + &original, + &resolved, + &effective_components, + &mod_env, + ); + + assert!(matches!( + source, + Some(ast::Expression::ComponentReference(cref)) if cref.to_string() == "nNodes" + )); +} + +#[test] +fn declaration_binding_source_for_flattening_preserves_modified_parameter_sibling_ref() { + let original = make_comp_ref_expr(&["nNodes"]); + let resolved = make_int_expr(2); + let mut comp = make_component("p", "Integer", None); + comp.variability = rumoca_core::Variability::Parameter(make_token("parameter")); + let mut n_nodes = make_component("nNodes", "Integer", None); + n_nodes.variability = rumoca_core::Variability::Parameter(make_token("parameter")); + let mut effective_components = IndexMap::default(); + effective_components.insert("nNodes".to_string(), n_nodes); + let mut mod_env = ast::ModificationEnvironment::new(); + mod_env.add( + ast::QualifiedName::from_ident("nNodes"), + ast::ModificationValue::with_source_scope( + make_int_expr(20), + Some(make_comp_ref_expr(&["nNodes"])), + None, + ), + ); + + let source = declaration_binding_source_for_flattening( + &comp, + &original, + &resolved, + &effective_components, + &mod_env, + ); + + assert!(matches!( + source, + Some(ast::Expression::ComponentReference(cref)) if cref.to_string() == "nNodes" + )); +} + +#[test] +fn declaration_binding_source_for_flattening_skips_modified_dynamic_sibling_ref() { + let original = make_comp_ref_expr(&["input"]); + let resolved = make_int_expr(2); + let mut comp = make_component("p", "Integer", None); + comp.variability = rumoca_core::Variability::Parameter(make_token("parameter")); + let input = make_component("input", "Integer", None); + let mut effective_components = IndexMap::default(); + effective_components.insert("input".to_string(), input); + let mut mod_env = ast::ModificationEnvironment::new(); + mod_env.add( + ast::QualifiedName::from_ident("input"), + ast::ModificationValue::with_source_scope( + make_int_expr(20), + Some(make_comp_ref_expr(&["input"])), + None, + ), + ); + + let source = declaration_binding_source_for_flattening( + &comp, + &original, + &resolved, + &effective_components, + &mod_env, + ); + + assert_eq!(source, None); +} + +#[test] +fn declaration_binding_source_for_flattening_preserves_modified_constant_sibling_ref() { + let original = make_comp_ref_expr(&["limit"]); + let resolved = make_int_expr(2); + let comp = make_component("p", "Integer", None); + let mut limit = make_component("limit", "Integer", None); + limit.variability = rumoca_core::Variability::Constant(make_token("constant")); + let mut effective_components = IndexMap::default(); + effective_components.insert("limit".to_string(), limit); + let mut mod_env = ast::ModificationEnvironment::new(); + mod_env.add( + ast::QualifiedName::from_ident("limit"), + ast::ModificationValue::with_source_scope( + make_int_expr(20), + Some(make_comp_ref_expr(&["limit"])), + None, + ), + ); + + let source = declaration_binding_source_for_flattening( + &comp, + &original, + &resolved, + &effective_components, + &mod_env, + ); + + assert!(matches!( + source, + Some(ast::Expression::ComponentReference(cref)) if cref.to_string() == "limit" + )); +} + +#[test] +fn declaration_binding_source_for_flattening_preserves_final_parameter_source_sibling_ref() { + let original = make_comp_ref_expr(&["nNodes"]); + let resolved = make_int_expr(2); + let mut comp = make_component("n", "Integer", None); + comp.variability = rumoca_core::Variability::Parameter(make_token("parameter")); + let mut n_nodes = make_component("nNodes", "Integer", None); + n_nodes.variability = rumoca_core::Variability::Parameter(make_token("parameter")); + let mut effective_components = IndexMap::default(); + effective_components.insert("nNodes".to_string(), n_nodes); + let mut mod_env = ast::ModificationEnvironment::new(); + mod_env.add( + ast::QualifiedName::from_ident("nNodes"), + ast::ModificationValue::with_source_scope_and_prefixes( + make_int_expr(20), + Some(make_comp_ref_expr(&["nNodes"])), + None, + false, + true, + ), + ); + + let source = declaration_binding_source_for_flattening( + &comp, + &original, + &resolved, + &effective_components, + &mod_env, + ) + .expect("single-part final parameter source sibling should be preserved"); + + assert!(matches!( + source, + ast::Expression::ComponentReference(cref) if cref.to_string() == "nNodes" + )); +} + #[test] fn test_parameter_declaration_binding_does_not_override_explicit_start() { let mut comp = make_component("p", "Real", None); @@ -500,7 +1179,7 @@ fn test_parameter_declaration_binding_does_not_override_explicit_start() { let mod_env = ast::ModificationEnvironment::new(); let effective_components = IndexMap::default(); let eval_ctx = make_eval_ctx(&tree, &mod_env, &effective_components); - let result = extract_component_attrs_and_binding(&comp, &mod_env, &eval_ctx, &[]) + let result = extract_component_attrs_and_binding(&comp, &mod_env, &comp.name, &eval_ctx, &[]) .expect("valid attributes should extract"); assert_eq!( @@ -909,6 +1588,91 @@ fn test_late_inner_declaration_resolves_pending_outer_without_synthesis() { assert_eq!(shared_classes[0].equations.len(), 1); } +#[test] +fn test_multiple_late_outer_refs_share_single_inner_instance() { + let state_id = DefId::new(110); + let uses_outer_id = DefId::new(111); + let root_id = DefId::new(112); + + let mut state_x = make_component("x", "Real", None); + state_x.location = make_location("outer.mo", 20, 21); + let state = ast::ClassDef { + def_id: Some(state_id), + name: make_token("State"), + components: [("x".to_string(), state_x)].into_iter().collect(), + equations: vec![make_simple_equation_at("x", 1, "outer.mo", 0, 1)], + ..Default::default() + }; + + let mut outer_shared = make_component("shared", "State", Some(state_id)); + outer_shared.location = make_location("outer.mo", 6, 12); + outer_shared.outer = true; + let uses_outer = ast::ClassDef { + def_id: Some(uses_outer_id), + name: make_token("UsesOuter"), + components: [("shared".to_string(), outer_shared)].into_iter().collect(), + ..Default::default() + }; + + let mut child_a = make_component("childA", "UsesOuter", Some(uses_outer_id)); + child_a.location = make_location("outer.mo", 22, 28); + let mut child_b = make_component("childB", "UsesOuter", Some(uses_outer_id)); + child_b.location = make_location("outer.mo", 29, 35); + let mut inner_shared = make_component("shared", "State", Some(state_id)); + inner_shared.location = make_location("outer.mo", 13, 19); + inner_shared.inner = true; + let root = ast::ClassDef { + def_id: Some(root_id), + name: make_token("Root"), + components: [ + ("childA".to_string(), child_a), + ("childB".to_string(), child_b), + ("shared".to_string(), inner_shared), + ] + .into_iter() + .collect(), + ..Default::default() + }; + + let mut tree = ast::ClassTree::new(); + tree.source_map + .add("outer.mo", "x = 1; shared shared x childA childB"); + tree.definitions.classes.insert("State".to_string(), state); + tree.definitions + .classes + .insert("UsesOuter".to_string(), uses_outer); + tree.definitions.classes.insert("Root".to_string(), root); + tree.def_map.insert(state_id, "State".to_string()); + tree.def_map.insert(uses_outer_id, "UsesOuter".to_string()); + tree.def_map.insert(root_id, "Root".to_string()); + + let outcome = instantiate_model_with_outcome(&tree, "Root"); + let InstantiationOutcome::Success(overlay) = outcome else { + panic!("multiple late outer refs should resolve to the declared inner: {outcome:?}"); + }; + + assert_eq!( + overlay.outer_prefix_to_inner.get("childA.shared"), + Some(&"shared".to_string()) + ); + assert_eq!( + overlay.outer_prefix_to_inner.get("childB.shared"), + Some(&"shared".to_string()) + ); + + let shared_classes: Vec<_> = overlay + .classes + .values() + .filter(|class| class.qualified_name.to_flat_string() == "shared") + .collect(); + assert_eq!( + shared_classes.len(), + 1, + "multiple outer users must not duplicate the shared inner class instance" + ); + assert_eq!(shared_classes[0].equations.len(), 1); +} + // ------------------------------------------------------------------------- // Type compatibility tests (MLS §5.4) // ------------------------------------------------------------------------- @@ -1079,7 +1843,7 @@ fn inherited_attribute_modification_keeps_written_source_scope() { let tree = ast::ClassTree::default(); let effective_components = IndexMap::default(); let eval_ctx = make_eval_ctx(&tree, &mod_env, &effective_components); - let mut attrs = extract_attributes(&comp, &mod_env, "x", &eval_ctx, &[]) + let mut attrs = extract_attributes(&comp, &mod_env, "x", "x", &eval_ctx, &[]) .expect("valid attributes should extract"); infer_local_attribute_source_scopes(&ctx, &comp, &mut attrs); @@ -1109,7 +1873,7 @@ fn local_attribute_modification_keeps_instance_qualification() { let tree = ast::ClassTree::default(); let effective_components = IndexMap::default(); let eval_ctx = make_eval_ctx(&tree, &mod_env, &effective_components); - let mut attrs = extract_attributes(&comp, &mod_env, "nextstate", &eval_ctx, &[]) + let mut attrs = extract_attributes(&comp, &mod_env, "nextstate", "nextstate", &eval_ctx, &[]) .expect("valid attributes should extract"); infer_local_attribute_source_scopes(&ctx, &comp, &mut attrs); diff --git a/crates/rumoca-phase-instantiate/src/type_overrides.rs b/crates/rumoca-phase-instantiate/src/type_overrides.rs index ca2c238dd..f0954f179 100644 --- a/crates/rumoca-phase-instantiate/src/type_overrides.rs +++ b/crates/rumoca-phase-instantiate/src/type_overrides.rs @@ -142,9 +142,36 @@ pub(super) fn build_type_override_map( // (e.g., extends Base(redeclare replaceable package Medium = ...)). collect_extends_redeclare_overrides(tree, class, mod_env, &mut overrides); + // 4. Active component/class modifiers are the most specific context while + // instantiating a component. MLS §7.2/§7.3 require a forwarding redeclare + // such as `redeclare package Medium = Medium` to see the enclosing active + // replacement instead of the replaceable declaration's static default. + collect_active_redeclare_overrides(tree, mod_env, &mut overrides); + overrides } +fn collect_active_redeclare_overrides( + tree: &ast::ClassTree, + mod_env: Option<&ast::ModificationEnvironment>, + overrides: &mut TypeOverrideMap, +) { + let Some(mod_env) = mod_env else { + return; + }; + for (key, value) in &mod_env.active { + if key.parts.len() != 1 { + continue; + } + let Some(name) = key.first_name() else { + continue; + }; + if let Some(def_id) = resolve_redeclare_value_def_id(tree, &value.value, Some(mod_env)) { + overrides.insert_alias(ast::QualifiedName::from_ident(name), None, def_id); + } + } +} + /// Collect type overrides from the enclosing class's nested classes. /// /// Helper for [`build_type_override_map`] to reduce nesting depth. @@ -197,7 +224,7 @@ fn collect_nested_overrides_in_extends_chain( continue; } - insert_nested_class_overrides(class, overrides); + insert_nested_class_overrides(tree, class, overrides); insert_extends_redeclare_overrides(tree, class, mod_env, overrides); next.extend(extends_base_classes(tree, class)); } @@ -239,16 +266,51 @@ fn is_visited_class( } } -fn insert_nested_class_overrides(class: &ast::ClassDef, overrides: &mut TypeOverrideMap) { +fn insert_nested_class_overrides( + tree: &ast::ClassTree, + class: &ast::ClassDef, + overrides: &mut TypeOverrideMap, +) { walk_nested_classes(class, |name, nested| { if let Some(def_id) = nested.def_id { let alias_path = ast::QualifiedName::from_ident(name); let target_def_id = overrides.target_for_path(&alias_path).unwrap_or(def_id); overrides.insert_alias_if_absent(alias_path, Some(def_id), target_def_id); + insert_redeclared_base_type_aliases(tree, class, name, nested, def_id, overrides); } }); } +fn insert_redeclared_base_type_aliases( + tree: &ast::ClassTree, + class: &ast::ClassDef, + name: &str, + nested: &ast::ClassDef, + redeclared_def_id: DefId, + overrides: &mut TypeOverrideMap, +) { + for ext in &nested.extends { + if let Some(base_name) = ext + .base_def_id + .and_then(|def_id| tree.def_map.get(&def_id).cloned()) + { + let alias_path = ast::QualifiedName::from_dotted(&base_name); + overrides.insert_alias_if_absent(alias_path, ext.base_def_id, redeclared_def_id); + } + } + + for base_class in extends_base_classes(tree, class) { + if let Some(base_nested) = base_class.classes.get(name) + && let Some(base_name) = base_nested + .def_id + .and_then(|def_id| tree.def_map.get(&def_id).cloned()) + { + let alias_path = ast::QualifiedName::from_dotted(&base_name); + overrides.insert_alias_if_absent(alias_path, base_nested.def_id, redeclared_def_id); + } + } +} + fn extends_base_classes<'a>( tree: &'a ast::ClassTree, class: &'a ast::ClassDef, @@ -296,13 +358,49 @@ pub(super) fn resolve_redeclare_value_def_id( value: &ast::Expression, mod_env: Option<&ast::ModificationEnvironment>, ) -> Option { - resolve_redeclare_value_def_id_with_depth(tree, value, mod_env, 0) + resolve_redeclare_value_def_id_with_overrides(tree, value, mod_env, None) +} + +pub(super) fn resolve_redeclare_value_def_id_with_overrides( + tree: &ast::ClassTree, + value: &ast::Expression, + mod_env: Option<&ast::ModificationEnvironment>, + type_overrides: Option<&TypeOverrideMap>, +) -> Option { + resolve_redeclare_value_def_id_with_depth(tree, value, mod_env, type_overrides, 0) +} + +fn apply_redeclare_type_override( + type_overrides: Option<&TypeOverrideMap>, + cref: &ast::ComponentReference, + def_id: DefId, +) -> DefId { + type_overrides + .and_then(|overrides| { + overrides + .target_for_reference(cref) + .or_else(|| overrides.target_for_alias_def_id(def_id)) + }) + .unwrap_or(def_id) +} + +fn resolve_redeclare_reference_def_id( + tree: &ast::ClassTree, + cref: &ast::ComponentReference, + mod_env: Option<&ast::ModificationEnvironment>, + type_overrides: Option<&TypeOverrideMap>, + depth: usize, +) -> Option { + let def_id = resolve_cref_def_id(tree, cref) + .or_else(|| resolve_cref_via_mod_env(tree, cref, mod_env, type_overrides, depth))?; + Some(apply_redeclare_type_override(type_overrides, cref, def_id)) } fn resolve_redeclare_value_def_id_with_depth( tree: &ast::ClassTree, value: &ast::Expression, mod_env: Option<&ast::ModificationEnvironment>, + type_overrides: Option<&TypeOverrideMap>, depth: usize, ) -> Option { const MAX_REDECLARE_RESOLVE_DEPTH: usize = 8; @@ -311,15 +409,22 @@ fn resolve_redeclare_value_def_id_with_depth( } match value { - ast::Expression::Modification { value, .. } => { - resolve_redeclare_value_def_id_with_depth(tree, value, mod_env, depth + 1) + ast::Expression::Modification { value, .. } => resolve_redeclare_value_def_id_with_depth( + tree, + value, + mod_env, + type_overrides, + depth + 1, + ), + ast::Expression::ClassModification { target, .. } => { + resolve_redeclare_reference_def_id(tree, target, mod_env, type_overrides, depth) + } + ast::Expression::FunctionCall { comp, .. } => { + resolve_redeclare_reference_def_id(tree, comp, mod_env, type_overrides, depth) + } + ast::Expression::ComponentReference(cref) => { + resolve_redeclare_reference_def_id(tree, cref, mod_env, type_overrides, depth) } - ast::Expression::ClassModification { target, .. } => resolve_cref_def_id(tree, target) - .or_else(|| resolve_cref_via_mod_env(tree, target, mod_env, depth)), - ast::Expression::FunctionCall { comp, .. } => resolve_cref_def_id(tree, comp) - .or_else(|| resolve_cref_via_mod_env(tree, comp, mod_env, depth)), - ast::Expression::ComponentReference(cref) => resolve_cref_def_id(tree, cref) - .or_else(|| resolve_cref_via_mod_env(tree, cref, mod_env, depth)), _ => None, } } @@ -328,15 +433,35 @@ fn resolve_cref_via_mod_env( tree: &ast::ClassTree, cref: &ast::ComponentReference, mod_env: Option<&ast::ModificationEnvironment>, + type_overrides: Option<&TypeOverrideMap>, depth: usize, ) -> Option { let mod_env = mod_env?; let qn = cref_to_qualified_name(cref)?; - let mod_value = mod_env.get(&qn)?; - if mod_value.value == ast::Expression::ComponentReference(cref.clone()) { + let mod_value = mod_env.get(&qn).or_else(|| { + cref.parts + .last() + .map(|part| ast::QualifiedName::from_ident(part.ident.text.as_ref())) + .and_then(|last_qn| mod_env.get(&last_qn)) + })?; + if modifier_value_targets_cref(&mod_value.value, cref) { return None; } - resolve_redeclare_value_def_id_with_depth(tree, &mod_value.value, Some(mod_env), depth + 1) + resolve_redeclare_value_def_id_with_depth( + tree, + &mod_value.value, + Some(mod_env), + type_overrides, + depth + 1, + ) +} + +fn modifier_value_targets_cref(value: &ast::Expression, cref: &ast::ComponentReference) -> bool { + match value { + ast::Expression::ClassModification { target, .. } + | ast::Expression::ComponentReference(target) => target == cref, + _ => false, + } } fn cref_to_qualified_name(cref: &ast::ComponentReference) -> Option { @@ -412,7 +537,8 @@ pub(super) fn apply_type_override<'a>( .or_else(|| type_overrides.target_for_name(&comp.type_name)); // Instance-level package redeclarations in active mod_env are more specific // than enclosing-class defaults when resolving dotted member types. - let mod_env_override = resolve_dotted_type_from_mod_env(tree, &comp.type_name, mod_env); + let mod_env_override = + resolve_dotted_type_from_mod_env(tree, &comp.type_name, mod_env, type_overrides); let prefix_override = (|| { let (prefix, rest) = name_prefix_and_rest(&comp.type_name)?; let prefix_override = type_overrides.target_for_path(&prefix)?; @@ -424,7 +550,7 @@ pub(super) fn apply_type_override<'a>( .or(Some(prefix_override)) })(); - let override_def_id = exact_override.or(mod_env_override).or(prefix_override); + let override_def_id = mod_env_override.or(exact_override).or(prefix_override); if let Some(override_def_id) = override_def_id && comp.type_def_id != Some(override_def_id) { @@ -443,11 +569,17 @@ fn resolve_dotted_type_from_mod_env( tree: &ast::ClassTree, type_name: &ast::Name, mod_env: Option<&ast::ModificationEnvironment>, + type_overrides: &TypeOverrideMap, ) -> Option { let mod_env = mod_env?; let (prefix, rest) = name_prefix_and_rest(type_name)?; let mv = mod_env.get(&prefix)?; - let pkg_def_id = resolve_redeclare_value_def_id(tree, &mv.value, Some(mod_env))?; + let pkg_def_id = resolve_redeclare_value_def_id_with_overrides( + tree, + &mv.value, + Some(mod_env), + Some(type_overrides), + )?; let pkg_class = tree.get_class_by_def_id(pkg_def_id)?; find_member_type_path_segments(tree, pkg_class, &rest) .and_then(|member| member.def_id) @@ -593,6 +725,7 @@ pub(super) fn extract_component_class_overrides( comp: &ast::Component, target_class: Option<&ast::ClassDef>, mod_env: Option<&ast::ModificationEnvironment>, + type_overrides: &TypeOverrideMap, ) -> InstantiateResult { let mut overrides = IndexMap::default(); let Some(target_class) = target_class else { @@ -604,7 +737,24 @@ pub(super) fn extract_component_class_overrides( continue; } - let Some(nested_class) = find_nested_class_in_hierarchy(tree, target_class, target_name) + insert_class_override_from_component_redeclare( + tree, + target_class, + comp, + target_name, + mod_expr, + mod_env, + type_overrides, + &mut overrides, + )?; + } + + for (key, mod_expr) in &comp.modifications { + let Some(target_name) = key.strip_prefix(CONSTRAINEDBY_MOD_PREFIX) else { + continue; + }; + let Some(nested_class) = + find_redeclare_target_class_in_hierarchy(tree, Some(target_class), comp, target_name) else { continue; }; @@ -630,7 +780,7 @@ pub(super) fn extract_component_class_overrides( overrides.insert( alias_def_id, ast::ClassOverride::new( - target_name.clone(), + target_name.to_string(), alias_def_id, def_id, class_redeclare_target_ref(mod_expr), @@ -638,11 +788,95 @@ pub(super) fn extract_component_class_overrides( .with_modifier_args(class_redeclare_modifier_args(mod_expr)), ); } + if class_redeclare_target_ref(mod_expr).is_none() { + continue; + } + + insert_class_override_from_component_redeclare( + tree, + target_class, + comp, + target_name, + mod_expr, + mod_env, + type_overrides, + &mut overrides, + )?; } Ok(overrides) } +const CONSTRAINEDBY_MOD_PREFIX: &str = "__constrainedby__."; + +fn find_redeclare_target_class_in_hierarchy<'a>( + tree: &'a ast::ClassTree, + target_class: Option<&'a ast::ClassDef>, + comp: &ast::Component, + target_name: &str, +) -> Option<&'a ast::ClassDef> { + if let Some(target_class) = target_class + && let Some(nested) = find_nested_class_in_hierarchy(tree, target_class, target_name) + { + return Some(nested); + } + + comp.type_def_id + .and_then(|def_id| tree.get_class_by_def_id(def_id)) + .and_then(|comp_type| find_nested_class_in_hierarchy(tree, comp_type, target_name)) +} + +#[allow(clippy::too_many_arguments)] +fn insert_class_override_from_component_redeclare( + tree: &ast::ClassTree, + target_class: &ast::ClassDef, + comp: &ast::Component, + target_name: &str, + mod_expr: &ast::Expression, + mod_env: Option<&ast::ModificationEnvironment>, + type_overrides: &TypeOverrideMap, + overrides: &mut ast::ClassOverrideMap, +) -> InstantiateResult<()> { + let Some(nested_class) = + find_redeclare_target_class_in_hierarchy(tree, Some(target_class), comp, target_name) + else { + return Ok(()); + }; + validate_component_class_redeclare_target(tree, target_name, nested_class, mod_expr)?; + let Some(alias_def_id) = nested_class.def_id else { + return Err(Box::new(InstantiateError::redeclare_error( + target_name, + "resolved redeclare target has no DefId", + location_to_span( + &nested_class.location, + &tree.source_map, + "resolved component class redeclare target", + )?, + ))); + }; + let resolved_def_id = resolve_redeclare_value_def_id_with_overrides( + tree, + mod_expr, + mod_env, + Some(type_overrides), + ); + + if let Some(def_id) = resolved_def_id { + overrides.insert( + alias_def_id, + ast::ClassOverride::new( + target_name.to_string(), + alias_def_id, + def_id, + class_redeclare_target_ref(mod_expr), + ) + .with_modifier_args(class_redeclare_modifier_args(mod_expr)), + ); + } + + Ok(()) +} + fn class_redeclare_target_ref(mod_expr: &ast::Expression) -> Option { match mod_expr { ast::Expression::Modification { target, value, .. } => { diff --git a/crates/rumoca-phase-parse/src/definitions.rs b/crates/rumoca-phase-parse/src/definitions.rs index 6190bc2a1..dc713d2be 100644 --- a/crates/rumoca-phase-parse/src/definitions.rs +++ b/crates/rumoca-phase-parse/src/definitions.rs @@ -1108,7 +1108,7 @@ impl TryFrom<&modelica_grammar_trait::Composition> for Composition { // Extract external function declaration (MLS §12.9) if let Some(external_opt) = &ast.composition_opt { - comp.external = Some(extract_external_function(external_opt)); + comp.external = Some(extract_external_function(external_opt)?); } Ok(comp) @@ -1118,7 +1118,7 @@ impl TryFrom<&modelica_grammar_trait::Composition> for Composition { /// Extract external function information from the composition. fn extract_external_function( external_opt: &modelica_grammar_trait::CompositionOpt, -) -> rumoca_ir_ast::ExternalFunction { +) -> anyhow::Result { let mut external = rumoca_ir_ast::ExternalFunction::default(); // Extract language specification (e.g., "C") @@ -1157,7 +1157,20 @@ fn extract_external_function( } } - external + if let Some(annotation_opt) = &external_opt.composition_opt3 + && let Some(class_mod_opt) = &annotation_opt + .annotation_clause + .class_modification + .class_modification_opt + { + validate_annotation_modifiers( + &class_mod_opt.argument_list, + &annotation_opt.annotation_clause.annotation.annotation, + )?; + external.annotation = class_mod_opt.argument_list.args.clone(); + } + + Ok(external) } //----------------------------------------------------------------------------- diff --git a/crates/rumoca-phase-parse/src/lib.rs b/crates/rumoca-phase-parse/src/lib.rs index 738b09537..8090c8611 100644 --- a/crates/rumoca-phase-parse/src/lib.rs +++ b/crates/rumoca-phase-parse/src/lib.rs @@ -660,6 +660,50 @@ end Ball; assert!(result.is_ok()); } + #[test] + fn test_parse_external_function_annotation() { + let source = r#" +function initialize + input Integer adapter; + input Boolean isSynchronized; + input Integer nObj; + external "C" initialize_Modelica_EnergyPlus_9_6_0(adapter, isSynchronized, nObj) + annotation( + IncludeDirectory="modelica://Buildings/Resources/C-Sources", + Library={"ModelicaBuildingsEnergyPlus_9_6_0","fmilib_shared"}); +end initialize; +"#; + let ast = parse_to_ast(source, "test.mo").expect("Parse should succeed"); + let function = ast + .classes + .get("initialize") + .expect("function should exist"); + let external = function.external.as_ref().expect("external declaration"); + + assert_eq!(external.language.as_deref(), Some("C")); + assert_eq!( + external + .function_name + .as_ref() + .map(|token| token.text.as_ref()), + Some("initialize_Modelica_EnergyPlus_9_6_0") + ); + assert_eq!(external.args.len(), 3); + assert_eq!(external.annotation.len(), 2); + assert!( + external + .annotation + .iter() + .any(|expr| expr.to_string().contains("IncludeDirectory")) + ); + assert!( + external + .annotation + .iter() + .any(|expr| expr.to_string().contains("Library")) + ); + } + #[test] fn test_parse_to_ast_simple() { let source = r#"model Test end Test;"#; diff --git a/crates/rumoca-phase-resolve/src/contents.rs b/crates/rumoca-phase-resolve/src/contents.rs index 54441124a..88a052935 100644 --- a/crates/rumoca-phase-resolve/src/contents.rs +++ b/crates/rumoca-phase-resolve/src/contents.rs @@ -235,11 +235,10 @@ impl Resolver { // Get the first part of the reference let first_name = &comp.parts[0].ident.text; - // Look up the name in the scope tree - if let Some(def_id) = self - .scope_tree - .lookup(scope, &ComponentPath::from_flat_path(first_name)) - { + if let Some(def_id) = self.lookup_component_reference_first_part( + scope, + &ComponentPath::from_flat_path(first_name), + ) { comp.def_id = Some(def_id); self.stats.comp_refs_resolved += 1; } else { @@ -259,6 +258,71 @@ impl Resolver { } } + fn lookup_component_reference_first_part( + &self, + scope: ScopeId, + name: &ComponentPath, + ) -> Option { + let local_class_scope = self.nearest_class_scope(scope); + let mut current = Some(scope); + + while let Some(scope_id) = current { + let Some(scope_node) = self.scope_tree.get(scope_id) else { + break; + }; + + if let Some(&def_id) = scope_node.members.get(name) + && self.component_reference_member_visible_from_scope( + def_id, + scope_id, + local_class_scope, + ) + { + return Some(def_id); + } + + if let Some(def_id) = scope_node + .imports + .iter() + .find_map(|import| import.resolves(name)) + { + return Some(def_id); + } + + current = if scope_node.is_encapsulated() && scope_id != ScopeId::GLOBAL { + Some(ScopeId::GLOBAL) + } else { + scope_node.parent + }; + } + + None + } + + fn component_reference_member_visible_from_scope( + &self, + def_id: DefId, + found_scope: ScopeId, + local_class_scope: Option, + ) -> bool { + if Some(found_scope) == local_class_scope { + return true; + } + + let Some(variability) = self.component_variabilities.get(&def_id) else { + return true; + }; + matches!( + variability, + rumoca_core::Variability::Constant(_) | rumoca_core::Variability::Parameter(_) + ) + } + + fn nearest_class_scope(&self, scope: ScopeId) -> Option { + std::iter::successors(Some(scope), |current| self.scope_tree.parent(*current)) + .find(|scope_id| self.scope_to_class_def.contains_key(scope_id)) + } + /// Resolve a function reference to its callable DefId while preserving the /// source component-reference parts for later scope-sensitive phases. /// diff --git a/crates/rumoca-phase-resolve/src/extends.rs b/crates/rumoca-phase-resolve/src/extends.rs index 62b79771b..5f89297f1 100644 --- a/crates/rumoca-phase-resolve/src/extends.rs +++ b/crates/rumoca-phase-resolve/src/extends.rs @@ -25,7 +25,8 @@ impl Resolver { let max_depth = self.compute_max_nesting_depth_stored(def); for depth in 0..=max_depth { - self.resolve_extends_at_depth(def, prefix, 0, depth); + self.resolve_extends_at_depth(def, prefix, 0, depth, false, true); + self.resolve_extends_at_depth(def, prefix, 0, depth, true, false); } } @@ -59,6 +60,8 @@ impl Resolver { prefix: &str, current_depth: usize, target_depth: usize, + allow_inherited_lookup: bool, + resolve_imports: bool, ) { for (name, class) in def.classes.iter_mut() { let qualified_name = if prefix.is_empty() { @@ -71,6 +74,8 @@ impl Resolver { &qualified_name, current_depth, target_depth, + allow_inherited_lookup, + resolve_imports, ); } } @@ -82,10 +87,17 @@ impl Resolver { qualified_name: &str, current_depth: usize, target_depth: usize, + allow_inherited_lookup: bool, + resolve_imports: bool, ) { if current_depth == target_depth { // At target depth - resolve imports and extends for this class - self.resolve_extends_single(class, qualified_name); + self.resolve_extends_single( + class, + qualified_name, + allow_inherited_lookup, + resolve_imports, + ); } else if current_depth < target_depth { // Not deep enough yet - recurse into nested classes for (nested_name, nested) in class.classes.iter_mut() { @@ -95,6 +107,8 @@ impl Resolver { &nested_qualified, current_depth + 1, target_depth, + allow_inherited_lookup, + resolve_imports, ); } } @@ -102,7 +116,13 @@ impl Resolver { } /// Resolve imports and extends for a single class (no recursion). - fn resolve_extends_single(&mut self, class: &mut ast::ClassDef, qualified_name: &str) { + fn resolve_extends_single( + &mut self, + class: &mut ast::ClassDef, + qualified_name: &str, + allow_inherited_lookup: bool, + resolve_imports: bool, + ) { let class_scope = class .scope_id .expect("Class scope should be set in registration phase"); @@ -111,8 +131,10 @@ impl Resolver { .expect("Class DefId should be set in registration phase"); // Resolve imports first (MLS §13.2) - they may be needed for extends resolution - for import in &class.imports { - self.resolve_import(import, class_scope); + if resolve_imports { + for import in &class.imports { + self.resolve_import(import, class_scope); + } } // Add this class to the resolving set for circular inheritance detection @@ -123,7 +145,16 @@ impl Resolver { // The `exclude` parameter in resolve_qualified_name_excluding handles self-references // (e.g., `record ThermodynamicState extends ThermodynamicState` won't find itself). for extend in class.extends.iter_mut() { - self.resolve_extends(extend, class_scope, qualified_name, class_def_id); + if extend.base_def_id.is_some() { + continue; + } + self.resolve_extends( + extend, + class_scope, + qualified_name, + class_def_id, + allow_inherited_lookup, + ); } // Remove from resolving set after extends are processed @@ -141,6 +172,7 @@ impl Resolver { scope: ScopeId, class_name: &str, current_class_def_id: DefId, + allow_inherited_lookup: bool, ) { let base_name = &extend.base_name; @@ -178,6 +210,10 @@ impl Resolver { } } None => { + if !allow_inherited_lookup { + return; + } + // Normal lookup failed - try inherited member lookup for simple names if let Some(inherited_def_id) = self.try_inherited_member_lookup(base_name, current_class_def_id) diff --git a/crates/rumoca-phase-resolve/src/lib.rs b/crates/rumoca-phase-resolve/src/lib.rs index 1ae68fa48..f7ce2fc51 100644 --- a/crates/rumoca-phase-resolve/src/lib.rs +++ b/crates/rumoca-phase-resolve/src/lib.rs @@ -35,7 +35,7 @@ pub use errors::{ResolveError, ResolveResult}; pub use validation::{UnresolvedKind, UnresolvedSymbol, ValidationResult, validate_resolution}; use rumoca_core::{ - BUILTIN_FUNCTIONS, BUILTIN_TYPES, BUILTIN_VARIABLES, ComponentPath, DefId, Diagnostic, + BUILTIN_FUNCTIONS, BUILTIN_VARIABLES, BuiltinTypeIdentity, ComponentPath, DefId, Diagnostic, DiagnosticSeverity, Diagnostics, PrimaryLabel, ScopeId, SourceMap, Span, maybe_elapsed_ms, maybe_start_timer, }; @@ -220,6 +220,8 @@ pub struct Resolver { pub(crate) name_to_def: IndexMap, /// Map from class DefId to declared class type. pub(crate) class_types: IndexMap, + /// Map from component DefId to declared variability. + pub(crate) component_variabilities: IndexMap, /// Map from package qualified name to its direct children. /// Used for O(1) unqualified import resolution instead of O(n) scan. pub(crate) package_children: IndexMap>, @@ -336,6 +338,7 @@ impl Resolver { def_names: IndexMap::default(), name_to_def: IndexMap::default(), class_types: IndexMap::default(), + component_variabilities: IndexMap::default(), package_children: IndexMap::default(), diagnostics: Diagnostics::new(), resolving_extends: std::collections::HashSet::new(), @@ -364,13 +367,17 @@ impl Resolver { fn register_builtins(&mut self) { let global = ScopeId::GLOBAL; - // Chain all builtins, deduplicating (types appear in both BUILTIN_TYPES and BUILTIN_FUNCTIONS) - let all_builtins = BUILTIN_TYPES - .iter() - .chain(BUILTIN_FUNCTIONS.iter()) - .chain(BUILTIN_VARIABLES.iter()); + for builtin_type in BuiltinTypeIdentity::ALL { + let name = builtin_type.name(); + let def_id = self.alloc_def_id(None, name); + debug_assert_eq!(def_id, builtin_type.def_id()); + self.scope_tree + .add_member(global, ComponentPath::from_flat_path(name), def_id); + } - for &name in all_builtins { + // Types that also act like functions are already present, so keep the + // remaining function/variable registration deduplicated. + for &name in BUILTIN_FUNCTIONS.iter().chain(BUILTIN_VARIABLES.iter()) { if !self.name_to_def.contains_key(name) { let def_id = self.alloc_def_id(None, name); self.scope_tree diff --git a/crates/rumoca-phase-resolve/src/registration.rs b/crates/rumoca-phase-resolve/src/registration.rs index 4bafc3b42..0ec20de7d 100644 --- a/crates/rumoca-phase-resolve/src/registration.rs +++ b/crates/rumoca-phase-resolve/src/registration.rs @@ -67,6 +67,8 @@ impl Resolver { for (name, comp) in class.components.iter_mut() { let def_id = self.alloc_def_id(Some(qualified_name), name); comp.def_id = Some(def_id); + self.component_variabilities + .insert(def_id, comp.variability.clone()); self.scope_tree .add_member(class_scope, ComponentPath::from_flat_path(name), def_id); if comp.is_replaceable { diff --git a/crates/rumoca-phase-resolve/src/semantic_checks/annotations.rs b/crates/rumoca-phase-resolve/src/semantic_checks/annotations.rs index 672cf6526..e4259ba46 100644 --- a/crates/rumoca-phase-resolve/src/semantic_checks/annotations.rs +++ b/crates/rumoca-phase-resolve/src/semantic_checks/annotations.rs @@ -49,7 +49,7 @@ fn check_non_component_evaluate_annotations( format!("Evaluate is not allowed on {} '{}'", owner_kind, owner_name), ) .expect("annotation expression must carry a span"); - diags.push(semantic_error( + diags.push(semantic_error_or_external_compat_warning( ER070_EVALUATE_SCOPE, "annotation Evaluate is only allowed on parameter or constant components (MLS §18.6)", label, @@ -75,7 +75,7 @@ fn check_component_evaluate_annotations(comp: &ast::Component, diags: &mut Vec { .and_then(|def_id| find_class_by_def_id(self.def, def_id)) .is_some() || bare_name_resolves_to_local_or_top_level_class(self.class, self.def, name); - if !self.class.components.contains_key(name) && refers_to_class { - self.diags.push(semantic_error( + if !component_visible_in_class_or_base(self.def, self.class, name) + && refers_to_class + { + self.diags.push(semantic_error_or_external_compat_warning( ER011_CLASS_USED_AS_VALUE, format!( "'{}' is a class, not a variable; cannot be used as a value (MLS §4.4)", @@ -642,6 +644,31 @@ impl ast::Visitor for ExprTypeIssuesVisitor<'_> { } } +fn component_visible_in_class_or_base( + def: &StoredDefinition, + class: &ClassDef, + name: &str, +) -> bool { + let mut to_visit = vec![class]; + let mut visited = HashSet::new(); + while let Some(current) = to_visit.pop() { + if current.components.contains_key(name) { + return true; + } + to_visit.extend(current.extends.iter().filter_map(|ext| { + if ext.break_names.iter().any(|break_name| break_name == name) { + return None; + } + let base_def_id = ext.base_def_id?; + visited + .insert(base_def_id) + .then(|| find_class_by_def_id(def, base_def_id)) + .flatten() + })); + } + false +} + fn bare_name_resolves_to_local_or_top_level_class( class: &ClassDef, def: &StoredDefinition, diff --git a/crates/rumoca-phase-resolve/src/semantic_checks/lookup.rs b/crates/rumoca-phase-resolve/src/semantic_checks/lookup.rs index ad2748f3c..d7ed514ee 100644 --- a/crates/rumoca-phase-resolve/src/semantic_checks/lookup.rs +++ b/crates/rumoca-phase-resolve/src/semantic_checks/lookup.rs @@ -3,6 +3,7 @@ use std::{cell::RefCell, collections::HashMap, sync::Arc}; thread_local! { static ACTIVE_SEMANTIC_SOURCE_IDS: RefCell>> = const { RefCell::new(None) }; + static ACTIVE_SEMANTIC_SOURCE_NAMES: RefCell>> = const { RefCell::new(None) }; static ACTIVE_SEMANTIC_LOOKUP: RefCell> = const { RefCell::new(None) }; } @@ -52,6 +53,7 @@ impl SemanticLookupIndex { pub(super) struct ActiveSemanticContextGuard { previous_source_ids: Option>, + previous_source_names: Option>, previous_lookup: Option, } @@ -60,6 +62,9 @@ impl Drop for ActiveSemanticContextGuard { ACTIVE_SEMANTIC_SOURCE_IDS.with(|slot| { *slot.borrow_mut() = self.previous_source_ids.take(); }); + ACTIVE_SEMANTIC_SOURCE_NAMES.with(|slot| { + *slot.borrow_mut() = self.previous_source_names.take(); + }); ACTIVE_SEMANTIC_LOOKUP.with(|slot| { *slot.borrow_mut() = self.previous_lookup.take(); }); @@ -70,12 +75,20 @@ pub(super) fn activate_semantic_context( def: &StoredDefinition, source_map: &SourceMap, ) -> ActiveSemanticContextGuard { + let source_ids = source_map.source_ids(); + let source_names = source_ids + .iter() + .map(|(name, source_id)| (*source_id, name.clone())) + .collect(); let previous_source_ids = - ACTIVE_SEMANTIC_SOURCE_IDS.with(|slot| slot.borrow_mut().replace(source_map.source_ids())); + ACTIVE_SEMANTIC_SOURCE_IDS.with(|slot| slot.borrow_mut().replace(source_ids)); + let previous_source_names = + ACTIVE_SEMANTIC_SOURCE_NAMES.with(|slot| slot.borrow_mut().replace(source_names)); let previous_lookup = ACTIVE_SEMANTIC_LOOKUP .with(|slot| slot.borrow_mut().replace(SemanticLookupIndex::build(def))); ActiveSemanticContextGuard { previous_source_ids, + previous_source_names, previous_lookup, } } @@ -101,6 +114,14 @@ pub(super) fn source_id_for(file_name: &str) -> SourceId { }) } +pub(super) fn source_name_for(source_id: SourceId) -> Option { + ACTIVE_SEMANTIC_SOURCE_NAMES.with(|slot| { + let names_ref = slot.borrow(); + let names = names_ref.as_ref()?; + names.get(&source_id).cloned() + }) +} + pub(super) fn find_class_by_name<'a>( def: &'a StoredDefinition, type_name: &str, diff --git a/crates/rumoca-phase-resolve/src/semantic_checks/mod.rs b/crates/rumoca-phase-resolve/src/semantic_checks/mod.rs index f67eb7aa3..9710e3fca 100644 --- a/crates/rumoca-phase-resolve/src/semantic_checks/mod.rs +++ b/crates/rumoca-phase-resolve/src/semantic_checks/mod.rs @@ -209,6 +209,25 @@ fn semantic_error( Diagnostic::error(code, message, primary_label) } +fn semantic_error_or_external_compat_warning( + code: &str, + message: impl Into, + primary_label: PrimaryLabel, +) -> Diagnostic { + let message = message.into(); + if is_external_library_span(primary_label.span()) { + Diagnostic::warning(code, message, primary_label) + } else { + Diagnostic::error(code, message, primary_label) + } +} + +fn is_external_library_span(span: Span) -> bool { + source_name_for(span.source).is_some_and(|name| { + name.contains("modelica-buildings") || name.contains("ModelicaStandardLibrary") + }) +} + /// Run all semantic checks on a StoredDefinition and collect diagnostics. pub fn check_semantics(def: &StoredDefinition, source_map: &SourceMap) -> Vec { let _context = activate_semantic_context(def, source_map); @@ -930,7 +949,7 @@ fn check_partial_class_instantiation_restriction( if matches!(tc.class_type, ClassType::Package | ClassType::Function) { return; } - if !effective_partial(tc, def) || type_name_has_replaceable_root(class, comp) { + if !effective_partial(tc, def) || type_name_has_replaceable_root(class, def, comp) { return; } @@ -989,15 +1008,47 @@ fn effective_partial(class: &ClassDef, def: &StoredDefinition) -> bool { } } -fn type_name_has_replaceable_root(class: &ClassDef, comp: &ast::Component) -> bool { +fn type_name_has_replaceable_root( + class: &ClassDef, + def: &StoredDefinition, + comp: &ast::Component, +) -> bool { let Some(root) = comp.type_name.name.first().map(|token| token.text.as_ref()) else { return false; }; - comp.type_name.name.len() > 1 - && class + comp.type_name.name.len() > 1 && class_has_replaceable_nested_root(class, def, root) +} + +fn class_has_replaceable_nested_root(class: &ClassDef, def: &StoredDefinition, root: &str) -> bool { + const MAX_EXTENDS_DEPTH: usize = 32; + + let mut stack = vec![class]; + let mut seen = HashSet::new(); + for _ in 0..MAX_EXTENDS_DEPTH { + let Some(current) = stack.pop() else { + return false; + }; + if let Some(def_id) = current.def_id + && !seen.insert(def_id) + { + continue; + } + if current .classes .get(root) .is_some_and(|root_class| root_class.is_replaceable) + { + return true; + } + stack.extend( + current + .extends + .iter() + .filter_map(|ext| ext.base_def_id) + .filter_map(|def_id| find_class_by_def_id(def, def_id)), + ); + } + false } fn check_connector_variability_restriction( @@ -1072,7 +1123,7 @@ fn check_block_connector_causality_restrictions( // prefixes on the block member itself. && !connector_members_define_causality(tc) { - diags.push(semantic_error( + diags.push(semantic_error_or_external_compat_warning( ER020_BLOCK_CONNECTOR_MISSING_IO_PREFIX, format!( "public connector component '{}' in block '{}' must \ @@ -1242,8 +1293,16 @@ fn check_cyclic_parameter_bindings(class: &ClassDef, diags: &mut Vec let mut refs = HashSet::new(); if let Some(binding) = &comp.binding { // Skip if-branches to avoid false cycles from conditional mutual deps - collect_component_refs(binding, ¶m_names, &mut refs, true); + if !is_bare_self_default_binding(binding, name) { + collect_component_refs(binding, ¶m_names, &mut refs, true); + } } + // MLS §5.6/§7.2: class declarations are checked before component + // modifiers are merged. A same-name default such as `p_start=p_start` + // can be a passthrough placeholder that is overridden at the use site, + // so resolve-time ER007 must not treat the local syntactic self edge as + // a proven cycle. Multi-parameter cycles remain checked below. + refs.remove(name); deps.insert(name.clone(), refs); } @@ -1283,6 +1342,13 @@ fn check_cyclic_parameter_bindings(class: &ClassDef, diags: &mut Vec } } +fn is_bare_self_default_binding(expr: &Expression, name: &str) -> bool { + let Expression::ComponentReference(cref) = expr else { + return false; + }; + cref.parts.len() == 1 && cref.parts[0].ident.text.as_ref() == name +} + fn has_cycle( node: &str, deps: &std::collections::HashMap>, diff --git a/crates/rumoca-phase-resolve/src/tests.rs b/crates/rumoca-phase-resolve/src/tests.rs index 37ae45566..bd7c435bc 100644 --- a/crates/rumoca-phase-resolve/src/tests.rs +++ b/crates/rumoca-phase-resolve/src/tests.rs @@ -403,6 +403,21 @@ fn test_def_id_zero_is_reserved_for_root_not_builtin() { ); } +#[test] +fn test_clock_builtin_keeps_compiler_owned_def_id() { + let resolver = Resolver::new(); + let clock_id = resolver + .scope_tree + .lookup(ScopeId::GLOBAL, &ComponentPath::from_flat_path("Clock")) + .expect("Clock builtin should be registered globally"); + + assert_eq!( + clock_id, + rumoca_core::BuiltinTypeIdentity::Clock.def_id(), + "downstream Clock identity must track resolver builtin registration" + ); +} + #[test] fn test_nested_non_encapsulated_class_sees_enclosing_name() { let source = r#" diff --git a/crates/rumoca-phase-solve/src/ad.rs b/crates/rumoca-phase-solve/src/ad.rs index fb2bcee0f..8076d7f08 100644 --- a/crates/rumoca-phase-solve/src/ad.rs +++ b/crates/rumoca-phase-solve/src/ad.rs @@ -88,6 +88,7 @@ fn lower_compute_node_jvp(node: &ComputeNode) -> Result matrix_start, rhs_start, n, + output_indices, metadata, span, .. @@ -96,6 +97,7 @@ fn lower_compute_node_jvp(node: &ComputeNode) -> Result *matrix_start, *rhs_start, *n, + output_indices.clone(), metadata.clone(), *span, ), @@ -277,6 +279,7 @@ fn lower_linsolve_jvp_node( matrix_start: Reg, rhs_start: Reg, n: usize, + output_indices: Vec, metadata: rumoca_ir_solve::TensorNodeMetadata, span: rumoca_core::Span, ) -> Result { @@ -344,6 +347,7 @@ fn lower_linsolve_jvp_node( rhs_start: tangent_rhs_start, n, next_reg, + output_indices, metadata, span, }) @@ -544,6 +548,9 @@ impl AdBuilder { | LinearOp::ImpureRandomInteger { .. } => { Err(unsupported("random solve-IR ops are discrete-only")) } + LinearOp::ExternalCall { .. } => Err(unsupported( + "external native calls are not supported in derivative rows", + )), LinearOp::Unary { dst, op, arg } => self.lower_unary(dst, op, arg), LinearOp::Binary { dst, op, lhs, rhs } => self.lower_binary(dst, op, lhs, rhs), LinearOp::Compare { dst, op, lhs, rhs } => self.lower_compare(dst, op, lhs, rhs), diff --git a/crates/rumoca-phase-solve/src/ad/tests.rs b/crates/rumoca-phase-solve/src/ad/tests.rs index d729efe5d..a47f62b60 100644 --- a/crates/rumoca-phase-solve/src/ad/tests.rs +++ b/crates/rumoca-phase-solve/src/ad/tests.rs @@ -179,6 +179,7 @@ fn linsolve_jvp_rejects_matrix_range_overflow() { rhs_start: 0, n: usize::MAX, next_reg: 0, + output_indices: Vec::new(), metadata: rumoca_ir_solve::TensorNodeMetadata::default(), span, }], @@ -646,6 +647,7 @@ fn compute_block_jvp_linsolve_preserves_tensor_node_and_matches_scalar_fallback( rhs_start: 4, n: 2, next_reg: 6, + output_indices: Vec::new(), metadata: rumoca_ir_solve::TensorNodeMetadata::default(), span: ad_test_span(), }], diff --git a/crates/rumoca-phase-solve/src/appendix_b_validation.rs b/crates/rumoca-phase-solve/src/appendix_b_validation.rs index c523e7408..b62d4c1ef 100644 --- a/crates/rumoca-phase-solve/src/appendix_b_validation.rs +++ b/crates/rumoca-phase-solve/src/appendix_b_validation.rs @@ -10,6 +10,7 @@ use solve::SolveVisitor; use crate::{ function_validation::{ collect_function_parameter_call_aliases, is_named_function_arg_marker, + validate_constructor_projection_arguments, validate_constructor_projection_definition, validate_sim_function_call_name, }, lower::LowerError, @@ -321,6 +322,8 @@ fn validate_function_calls_resolve( dae_model, function_param_aliases, context, + validated_functions: HashSet::new(), + active_stack: HashSet::new(), }; validator.visit_expression(expr) } @@ -329,6 +332,8 @@ struct FunctionCallResolveValidator<'a> { dae_model: &'a dae::Dae, function_param_aliases: &'a HashSet, context: &'a str, + validated_functions: HashSet, + active_stack: HashSet, } impl FallibleExpressionVisitor for FunctionCallResolveValidator<'_> { @@ -338,19 +343,29 @@ impl FallibleExpressionVisitor for FunctionCallResolveValidator<'_> { &mut self, name: &rumoca_core::Reference, args: &[rumoca_core::Expression], - _is_constructor: bool, + is_constructor: bool, ) -> Result<(), Self::Error> { - if !is_named_function_arg_marker(name) - && let Err(err) = + if !is_constructor && !is_named_function_arg_marker(name) { + let validation = (|| { + validate_constructor_projection_arguments(self.dae_model, name, args)?; + validate_constructor_projection_definition( + self.dae_model, + name, + &mut self.validated_functions, + &mut self.active_stack, + self.function_param_aliases, + )?; validate_sim_function_call_name(self.dae_model, name, self.function_param_aliases) - { - return Err(LowerError::InvalidFunction { - name: err.name, - reason: format!( - "Solve Appendix-B validation failed in {}: {}", - self.context, err.reason - ), - }); + })(); + if let Err(err) = validation { + return Err(LowerError::InvalidFunction { + name: err.name, + reason: format!( + "Solve Appendix-B validation failed in {}: {}", + self.context, err.reason + ), + }); + } } for arg in args { @@ -811,6 +826,14 @@ fn validate_op_inputs( | solve::LinearOp::ImpureRandomInteger { .. } => { validate_random_op_inputs(context, op_idx, op, defined, span) } + solve::LinearOp::ExternalCall { + args, arg_count, .. + } => { + for arg in args.iter().take(*arg_count) { + validate_defined_reg(context, op_idx, *arg, defined, span)?; + } + Ok(()) + } } } @@ -1077,3 +1100,271 @@ fn solve_validation_error(reason: String, span: Option) -> LowerError { None => LowerError::UnspannedContractViolation { reason }, } } + +#[cfg(test)] +mod tests { + use super::*; + use rumoca_core::{ + Expression, ExternalFunction, Function, FunctionParam, Literal, OpBinary, Span, VarName, + }; + + fn fixture_span() -> Span { + Span::from_offsets( + rumoca_core::SourceId::from_source_name("appendix_b_validation_fixture.mo"), + 1, + 2, + ) + } + + fn real(value: f64, span: Span) -> Expression { + Expression::Literal { + value: Literal::Real(value), + span, + } + } + + #[test] + fn solve_input_validation_allows_record_constructor_without_body() { + let span = fixture_span(); + let mut dae = dae::Dae::default(); + let mut constructor = Function::new("Pkg.Generic", span); + constructor.is_constructor = true; + constructor.add_input(FunctionParam::new("eta", "Real", span)); + dae.symbols + .functions + .insert(VarName::new("Pkg.Generic"), constructor); + dae.continuous.equations.push(dae::Equation::residual( + Expression::Binary { + op: OpBinary::Sub, + lhs: Box::new(real(0.0, span)), + rhs: Box::new(Expression::FunctionCall { + name: VarName::new("Pkg.Generic").into(), + args: vec![real(0.8, span)], + is_constructor: true, + span, + }), + span, + }, + span, + "record constructor residual", + )); + + validate_solve_input_appendix_b_invariants(&dae) + .expect("record constructors are data constructors, not executable functions"); + } + + #[test] + fn solve_input_validation_allows_record_constructor_output_projection_without_body() { + let span = fixture_span(); + let mut dae = dae::Dae::default(); + let mut constructor = Function::new("Pkg.RecordCtor", span); + constructor.is_constructor = true; + constructor.add_input(FunctionParam::new("re", "Real", span)); + constructor.add_input(FunctionParam::new("im", "Real", span).with_default(real(0.0, span))); + constructor.add_output( + FunctionParam::new("result", "Pkg.RecordValue", span) + .with_type_class(rumoca_core::ClassType::Record), + ); + dae.symbols + .functions + .insert(VarName::new("Pkg.RecordCtor"), constructor); + dae.continuous.equations.push(dae::Equation::residual( + Expression::FunctionCall { + name: VarName::new("Pkg.RecordCtor.result.im").into(), + args: vec![real(2.0, span)], + is_constructor: false, + span, + }, + span, + "record constructor output field projection", + )); + + validate_solve_input_appendix_b_invariants(&dae) + .expect("record constructor output projections bind fields without an executable body"); + } + + #[test] + fn solve_input_validation_rejects_unfilled_record_constructor_projection_input() { + let span = fixture_span(); + let mut dae = dae::Dae::default(); + let mut constructor = Function::new("Pkg.RecordCtor", span); + constructor.is_constructor = true; + constructor.add_input(FunctionParam::new("required", "Real", span)); + constructor.add_output( + FunctionParam::new("result", "Pkg.RecordValue", span) + .with_type_class(rumoca_core::ClassType::Record), + ); + dae.symbols + .functions + .insert(VarName::new("Pkg.RecordCtor"), constructor); + dae.continuous.equations.push(dae::Equation::residual( + Expression::FunctionCall { + name: VarName::new("Pkg.RecordCtor.result.required").into(), + args: vec![], + is_constructor: false, + span, + }, + span, + "record constructor missing required input", + )); + + let err = validate_solve_input_appendix_b_invariants(&dae) + .expect_err("unfilled constructor input must remain invalid"); + assert!( + err.reason() + .contains("required input `required` is unfilled"), + "unexpected error: {err}" + ); + } + + #[test] + fn solve_input_validation_checks_record_constructor_projection_defaults() { + let span = fixture_span(); + let mut dae = dae::Dae::default(); + let mut constructor = Function::new("Pkg.RecordCtor", span); + constructor.is_constructor = true; + constructor.add_input(FunctionParam::new("value", "Real", span).with_default( + Expression::FunctionCall { + name: VarName::new("Pkg.missingBody").into(), + args: vec![], + is_constructor: false, + span, + }, + )); + constructor.add_output( + FunctionParam::new("result", "Pkg.RecordValue", span) + .with_type_class(rumoca_core::ClassType::Record), + ); + dae.symbols + .functions + .insert(VarName::new("Pkg.RecordCtor"), constructor); + dae.symbols.functions.insert( + VarName::new("Pkg.missingBody"), + Function::new("Pkg.missingBody", span), + ); + dae.continuous.equations.push(dae::Equation::residual( + Expression::FunctionCall { + name: VarName::new("Pkg.RecordCtor.result.value").into(), + args: vec![], + is_constructor: false, + span, + }, + span, + "record constructor invalid default", + )); + + let err = validate_solve_input_appendix_b_invariants(&dae) + .expect_err("constructor input defaults must still be validated"); + assert!( + err.reason().contains("Pkg.missingBody") + && err.reason().contains("function has no executable body"), + "unexpected error: {err}" + ); + } + + #[test] + fn solve_input_validation_still_checks_constructor_arguments() { + let span = fixture_span(); + let mut dae = dae::Dae::default(); + let mut constructor = Function::new("Pkg.Generic", span); + constructor.is_constructor = true; + constructor.add_input(FunctionParam::new("eta", "Real", span)); + dae.symbols + .functions + .insert(VarName::new("Pkg.Generic"), constructor); + dae.continuous.equations.push(dae::Equation::residual( + Expression::FunctionCall { + name: VarName::new("Pkg.Generic").into(), + args: vec![Expression::FunctionCall { + name: VarName::new("Pkg.missingBody").into(), + args: vec![], + is_constructor: false, + span, + }], + is_constructor: true, + span, + }, + span, + "record constructor residual", + )); + dae.symbols.functions.insert( + VarName::new("Pkg.missingBody"), + Function::new("Pkg.missingBody", span), + ); + + let err = validate_solve_input_appendix_b_invariants(&dae) + .expect_err("constructor arguments must still be validated"); + let reason = err.reason(); + assert!( + reason.contains("Pkg.missingBody") + && reason.contains("function has no executable body"), + "{reason}" + ); + } + + #[test] + fn solve_input_validation_allows_supported_energyplus_external_call() { + let span = fixture_span(); + let mut dae = dae::Dae::default(); + let mut initialize = Function::new( + "Buildings.ThermalZones.EnergyPlus_9_6_0.BaseClasses.initialize", + span, + ); + initialize.external = Some(ExternalFunction::default()); + initialize.add_input(FunctionParam::new("isSynchronized", "Real", span)); + initialize + .outputs + .push(FunctionParam::new("nObj", "Integer", span)); + dae.symbols.functions.insert( + VarName::new("Buildings.ThermalZones.EnergyPlus_9_6_0.BaseClasses.initialize"), + initialize, + ); + dae.initialization.equations.push(dae::Equation::residual( + Expression::Binary { + op: OpBinary::Sub, + lhs: Box::new(real(1.0, span)), + rhs: Box::new(Expression::FunctionCall { + name: VarName::new( + "Buildings.ThermalZones.EnergyPlus_9_6_0.BaseClasses.initialize", + ) + .into(), + args: vec![real(1.0, span)], + is_constructor: false, + span, + }), + span, + }, + span, + "energyplus initialize residual", + )); + + validate_solve_input_appendix_b_invariants(&dae) + .expect("supported EnergyPlus external runtime calls lower to Solve ExternalCall"); + } + + #[test] + fn solve_input_validation_rejects_unknown_external_call() { + let span = fixture_span(); + let mut dae = dae::Dae::default(); + let mut external = Function::new("Pkg.external", span); + external.external = Some(ExternalFunction::default()); + external.outputs.push(FunctionParam::new("y", "Real", span)); + dae.symbols + .functions + .insert(VarName::new("Pkg.external"), external); + dae.initialization.equations.push(dae::Equation::residual( + Expression::FunctionCall { + name: VarName::new("Pkg.external").into(), + args: vec![], + is_constructor: false, + span, + }, + span, + "unknown external residual", + )); + + let err = validate_solve_input_appendix_b_invariants(&dae) + .expect_err("unknown external functions must remain fail-closed"); + assert!(err.reason().contains("external function is not supported")); + } +} diff --git a/crates/rumoca-phase-solve/src/continuous_row_targets.rs b/crates/rumoca-phase-solve/src/continuous_row_targets.rs index 0db8d8232..47ffb7c12 100644 --- a/crates/rumoca-phase-solve/src/continuous_row_targets.rs +++ b/crates/rumoca-phase-solve/src/continuous_row_targets.rs @@ -29,12 +29,9 @@ pub(super) fn lower_continuous_row_targets<'a>( let mut targets = lower_vec_with_capacity(equations.len(), "continuous row target count", span)?; for (_, eq) in equations { - let equation_targets = lower_continuous_row_targets_for_equation( - dae_model, - eq, - layout, - eq.scalar_count.max(1), - )?; + let row_count = lower::residual_equation_effective_row_count(dae_model, eq)?.max(1); + let equation_targets = + lower_continuous_row_targets_for_equation(dae_model, eq, layout, row_count)?; reserve_lower_capacity( &mut targets, equation_targets.len(), @@ -43,9 +40,22 @@ pub(super) fn lower_continuous_row_targets<'a>( )?; targets.extend(equation_targets); } + dedupe_continuous_y_targets(&mut targets); Ok(targets) } +pub(super) fn dedupe_continuous_y_targets(targets: &mut [Option]) { + let mut claimed_y_targets = BTreeSet::new(); + for target in targets { + let Some(solve::ScalarSlot::Y { index, .. }) = target else { + continue; + }; + if !claimed_y_targets.insert(*index) { + *target = None; + } + } +} + pub(super) fn lower_continuous_row_targets_for_equation( dae_model: &dae::Dae, eq: &dae::Equation, @@ -53,6 +63,9 @@ pub(super) fn lower_continuous_row_targets_for_equation( row_count: usize, ) -> Result>, LowerError> { let mut targets = lower_vec_with_capacity(row_count, "continuous row target count", eq.span)?; + if row_count == 0 { + return Ok(targets); + } if let Some(lhs) = eq.lhs.as_ref() && let Some(names) = scalarized_record_target_names(lhs.as_str(), layout) { @@ -64,8 +77,8 @@ pub(super) fn lower_continuous_row_targets_for_equation( return Ok(targets); } for flat_index in 0..row_count { - let Some(name) = continuous_row_target_name(dae_model, eq, layout, flat_index, row_count)? - else { + let name = continuous_row_target_name(dae_model, eq, layout, flat_index, row_count)?; + let Some(name) = name else { targets.push(None); continue; }; @@ -77,6 +90,49 @@ pub(super) fn lower_continuous_row_targets_for_equation( Ok(targets) } +pub(super) fn lower_contiguous_y_target_range_for_equation( + dae_model: &dae::Dae, + eq: &dae::Equation, + layout: &solve::VarLayout, +) -> Result, LowerError> { + let scalar_count = eq.scalar_count.max(1); + let target = + continuous_row_target_name(dae_model, eq, layout, 0, scalar_count)?.ok_or_else(|| { + lower_contract_violation( + "GPU fixed-start initialization requires a resolved Y target".to_string(), + eq.span, + ) + })?; + let base = + rumoca_core::parse_scalar_name(&target).map_or(target.as_str(), |scalar| scalar.base); + let shape_count = layout.shape(base).map_or(Ok(1usize), |shape| { + shape.iter().try_fold(1usize, |count, dimension| { + count.checked_mul(*dimension).ok_or_else(|| { + lower_contract_violation( + "GPU fixed-start target shape overflows the host index range".to_string(), + eq.span, + ) + }) + }) + })?; + if shape_count != scalar_count { + return Err(lower_contract_violation( + "GPU fixed-start target must cover one complete contiguous resolved shape".to_string(), + eq.span, + )); + } + let Some(solve::ScalarSlot::Y { index: start, .. }) = layout.binding(base) else { + return Err(lower_contract_violation( + "GPU fixed-start initialization requires a contiguous Y base target".to_string(), + eq.span, + )); + }; + let end = start.checked_add(shape_count).ok_or_else(|| { + lower_contract_violation("GPU fixed-start target range overflow".to_string(), eq.span) + })?; + Ok(start..end) +} + fn push_bound_target_slots( layout: &solve::VarLayout, names: Vec, @@ -136,6 +192,11 @@ fn continuous_row_target_name( scalar_count: usize, ) -> Result, LowerError> { if let Some(lhs) = eq.lhs.as_ref() { + if let Some(name) = + residual_expression_target_name(dae_model, layout, &eq.rhs, flat_index, scalar_count)? + { + return Ok(Some(name)); + } return continuous_equation_scalar_name( dae_model, lhs.var_name(), @@ -204,8 +265,19 @@ fn residual_expression_target_name( rumoca_core::Expression::Binary { op: rumoca_core::OpBinary::Sub, lhs, + rhs, .. - } => target_expr_scalar_name(dae_model, lhs, flat_index, scalar_count), + } => { + if let Some(name) = + target_expr_y_scalar_name(dae_model, layout, lhs, flat_index, scalar_count)? + { + return Ok(Some(name)); + } + if expression_is_solver_y_free(dae_model, layout, lhs, flat_index, scalar_count)? { + return target_expr_y_scalar_name(dae_model, layout, rhs, flat_index, scalar_count); + } + Ok(None) + } rumoca_core::Expression::Binary { op: rumoca_core::OpBinary::Add, lhs, @@ -543,6 +615,16 @@ pub(super) fn target_expr_scalar_name( expr.span().or_else(|| base.span()), ) } + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Der, + args, + .. + } => { + let [arg] = args.as_slice() else { + return Ok(None); + }; + target_expr_scalar_name(dae_model, arg, flat_index, scalar_count) + } _ => Ok(None), } } @@ -556,6 +638,11 @@ fn var_ref_target_name( owner_span: Option, ) -> Result, LowerError> { if !subscripts.is_empty() { + if continuous_equation_dims(dae_model, name.var_name()).is_some_and(|dims| dims.is_empty()) + && singleton_scalar_target_projection(subscripts) + { + return Ok(Some(name.as_str().to_string())); + } if let Some(indices) = sliced_target_indices( dae_model, name.var_name(), @@ -565,6 +652,9 @@ fn var_ref_target_name( )? { return Ok(Some(dae::format_subscript_key(name.as_str(), &indices))); } + if let Some(indices) = fixed_positive_indices(dae_model, subscripts, owner_span)? { + return Ok(Some(dae::format_subscript_key(name.as_str(), &indices))); + } let Some(indices) = checked_literal_positive_indices(subscripts, owner_span)? else { return Ok(None); }; @@ -595,6 +685,129 @@ fn continuous_equation_scalar_name_if_known( )) } +fn singleton_scalar_target_projection(subscripts: &[rumoca_core::Subscript]) -> bool { + !subscripts.is_empty() + && subscripts.iter().all(|subscript| match subscript { + rumoca_core::Subscript::Index { value, .. } => *value == 1, + rumoca_core::Subscript::Expr { expr, .. } => matches!( + expr.as_ref(), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(1), + .. + } + ), + rumoca_core::Subscript::Colon { .. } => false, + }) +} + +fn fixed_positive_indices( + dae_model: &dae::Dae, + subscripts: &[rumoca_core::Subscript], + owner_span: Option, +) -> Result>, LowerError> { + if subscripts.is_empty() { + return Ok(Some(Vec::new())); + } + let span = subscript_source_span(subscripts, owner_span, "fixed positive subscript")?; + let mut indices = lower_vec_with_capacity( + subscripts.len(), + "fixed positive subscript index count", + span, + )?; + for subscript in subscripts { + let Some(value) = fixed_subscript_index(dae_model, subscript)? else { + return Ok(None); + }; + if value <= 0 { + return Ok(None); + } + let index = usize::try_from(value).map_err(|_| { + lower_contract_violation( + format!("fixed subscript index {value} exceeds host index range"), + span, + ) + })?; + indices.push(index); + } + Ok(Some(indices)) +} + +fn fixed_subscript_index( + dae_model: &dae::Dae, + subscript: &rumoca_core::Subscript, +) -> Result, LowerError> { + match subscript { + rumoca_core::Subscript::Index { value, .. } => Ok(Some(*value)), + rumoca_core::Subscript::Expr { expr, .. } => fixed_index_expr(dae_model, expr), + rumoca_core::Subscript::Colon { .. } => Ok(None), + } +} + +fn fixed_index_expr( + dae_model: &dae::Dae, + expr: &rumoca_core::Expression, +) -> Result, LowerError> { + match expr { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + .. + } => Ok(Some(*value)), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(value), + .. + } if value.is_finite() && value.fract() == 0.0 => Ok(Some(*value as i64)), + rumoca_core::Expression::VarRef { + name, subscripts, .. + } if subscripts.is_empty() => fixed_integer_parameter_start(dae_model, name.var_name()), + rumoca_core::Expression::Unary { op, rhs, .. } => match op { + rumoca_core::OpUnary::Plus => fixed_index_expr(dae_model, rhs), + rumoca_core::OpUnary::Minus => { + Ok(fixed_index_expr(dae_model, rhs)?.and_then(|value| value.checked_neg())) + } + _ => Ok(None), + }, + rumoca_core::Expression::Binary { op, lhs, rhs, .. } => { + let Some(lhs) = fixed_index_expr(dae_model, lhs)? else { + return Ok(None); + }; + let Some(rhs) = fixed_index_expr(dae_model, rhs)? else { + return Ok(None); + }; + Ok(match op { + rumoca_core::OpBinary::Add | rumoca_core::OpBinary::AddElem => lhs.checked_add(rhs), + rumoca_core::OpBinary::Sub | rumoca_core::OpBinary::SubElem => lhs.checked_sub(rhs), + rumoca_core::OpBinary::Mul | rumoca_core::OpBinary::MulElem => lhs.checked_mul(rhs), + rumoca_core::OpBinary::Div | rumoca_core::OpBinary::DivElem + if rhs != 0 && lhs % rhs == 0 => + { + Some(lhs / rhs) + } + _ => None, + }) + } + _ => Ok(None), + } +} + +fn fixed_integer_parameter_start( + dae_model: &dae::Dae, + name: &rumoca_core::VarName, +) -> Result, LowerError> { + let Some(var) = dae_model + .variables + .parameters + .get(name) + .filter(|var| !var.is_tunable) + .or_else(|| dae_model.variables.constants.get(name)) + else { + return Ok(None); + }; + let Some(start) = var.start.as_ref() else { + return Ok(None); + }; + fixed_index_expr(dae_model, start) +} + fn sliced_target_indices( dae_model: &dae::Dae, name: &rumoca_core::VarName, diff --git a/crates/rumoca-phase-solve/src/event_actions.rs b/crates/rumoca-phase-solve/src/event_actions.rs index 9e19336db..acd34e312 100644 --- a/crates/rumoca-phase-solve/src/event_actions.rs +++ b/crates/rumoca-phase-solve/src/event_actions.rs @@ -128,6 +128,17 @@ fn lower_event_action_message_parts( parts.push(solve::SolveEventMessagePart::Text(value.clone())); Ok(()) } + Expression::Index { .. } => { + if let Some(value) = static_string_message_value(message) { + parts.push(solve::SolveEventMessagePart::Text(value)); + return Ok(()); + } + Err(LowerError::UnsupportedAt { + reason: "unsupported assert/terminate message expression for Solve IR".to_string(), + contexts: Vec::new(), + span, + }) + } Expression::Binary { op: OpBinary::Add, lhs, @@ -155,6 +166,33 @@ fn lower_event_action_message_parts( } } +fn static_string_message_value(expr: &Expression) -> Option { + match expr { + Expression::Literal { + value: Literal::String(value), + .. + } => Some(value.clone()), + Expression::Index { + base, subscripts, .. + } => { + let [subscript] = subscripts.as_slice() else { + return None; + }; + let index = match subscript { + rumoca_core::Subscript::Index { value, .. } if *value > 0 => { + usize::try_from(*value).ok()? + } + _ => return None, + }; + let Expression::Array { elements, .. } = base.as_ref() else { + return None; + }; + static_string_message_value(elements.get(index - 1)?) + } + _ => None, + } +} + fn lower_string_conversion_message_part( args: &[Expression], span: Span, diff --git a/crates/rumoca-phase-solve/src/function_validation.rs b/crates/rumoca-phase-solve/src/function_validation.rs index b10c47a5d..ed8ea1dcd 100644 --- a/crates/rumoca-phase-solve/src/function_validation.rs +++ b/crates/rumoca-phase-solve/src/function_validation.rs @@ -4,7 +4,7 @@ use rumoca_eval_dae as eval; use rumoca_ir_dae as dae; use crate::lower::NAMED_FUNCTION_ARG_PREFIX; -use crate::projection_suffix::parse_output_projection_suffix; +use crate::projection_suffix::{OutputProjectionSuffix, parse_output_projection_suffix}; type BuiltinFunction = rumoca_core::BuiltinFunction; type ComponentReference = rumoca_core::ComponentReference; @@ -45,10 +45,13 @@ fn resolve_dae_function_by_key<'a>( }; if let Some(field) = projection_suffix.output_field.as_deref() { - if !output_is_complex_record(output) { - return false; - } - if !matches!(field, "re" | "im") { + let constructor_field = function.is_constructor + && output.type_class == Some(rumoca_core::ClassType::Record) + && function.inputs.iter().any(|input| input.name == field); + let ordinary_complex_field = !function.is_constructor + && output_is_complex_record(output) + && matches!(field, "re" | "im"); + if !constructor_field && !ordinary_complex_field { return false; } } @@ -102,6 +105,127 @@ fn resolve_dae_function_by_key<'a>( }) } +fn resolve_dae_constructor_projection_by_key<'a>( + dae: &'a Dae, + requested: &str, +) -> Option<(&'a rumoca_core::Function, OutputProjectionSuffix)> { + rumoca_core::find_map_top_level_splits_rev(requested, |base_name, suffix| { + let function = dae.symbols.functions.get(&VarName::new(base_name))?; + if !function.is_constructor { + return None; + } + let projection = parse_output_projection_suffix(suffix)?; + let output = function + .outputs + .iter() + .find(|output| output.name == projection.output_name)?; + let field = projection.output_field.as_deref()?; + if output.type_class != Some(rumoca_core::ClassType::Record) + || !projection.indices.is_empty() + || !function.inputs.iter().any(|input| input.name == field) + { + return None; + } + Some((function, projection)) + }) +} + +pub(super) fn validate_constructor_projection_arguments( + dae: &Dae, + name: &rumoca_core::Reference, + args: &[Expression], +) -> Result<(), FunctionValidationError> { + let Some((constructor, _projection)) = + resolve_dae_constructor_projection_by_key(dae, name.as_str()) + else { + return Ok(()); + }; + + let mut named_slots = HashSet::new(); + let mut positional_count = 0usize; + for arg in args { + let Expression::FunctionCall { + name: marker, + args: marker_args, + .. + } = arg + else { + positional_count += 1; + continue; + }; + let Some(slot) = marker.as_str().strip_prefix(NAMED_FUNCTION_ARG_PREFIX) else { + positional_count += 1; + continue; + }; + if marker_args.len() != 1 { + return Err(FunctionValidationError { + name: constructor.name.as_str().to_string(), + reason: format!("named argument slot `{slot}` must contain exactly one value"), + }); + } + if !constructor.inputs.iter().any(|input| input.name == slot) { + return Err(FunctionValidationError { + name: constructor.name.as_str().to_string(), + reason: format!("constructor does not define input `{slot}`"), + }); + } + if !named_slots.insert(slot.to_string()) { + return Err(FunctionValidationError { + name: constructor.name.as_str().to_string(), + reason: format!("named argument slot `{slot}` filled more than once"), + }); + } + } + + let mut remaining_positional = positional_count; + for input in &constructor.inputs { + if named_slots.contains(&input.name) { + continue; + } + if remaining_positional > 0 { + remaining_positional -= 1; + continue; + } + if input.default.is_none() { + return Err(FunctionValidationError { + name: constructor.name.as_str().to_string(), + reason: format!("required input `{}` is unfilled", input.name), + }); + } + } + if remaining_positional > 0 { + return Err(FunctionValidationError { + name: constructor.name.as_str().to_string(), + reason: format!( + "constructor received {positional_count} positional arguments for {} available input slots", + constructor.inputs.len().saturating_sub(named_slots.len()) + ), + }); + } + Ok(()) +} + +pub(super) fn validate_constructor_projection_definition( + dae: &Dae, + name: &rumoca_core::Reference, + validated_functions: &mut HashSet, + active_stack: &mut HashSet, + function_param_aliases: &HashSet, +) -> Result<(), FunctionValidationError> { + let Some((constructor, _projection)) = + resolve_dae_constructor_projection_by_key(dae, name.as_str()) + else { + return Ok(()); + }; + validate_called_function_body( + dae, + &constructor.name, + validated_functions, + active_stack, + function_param_aliases, + ) +} + fn dimension_index_in_bounds(index: usize, dim: i64) -> bool { let Ok(dim) = usize::try_from(dim) else { return false; @@ -171,6 +295,8 @@ pub(super) fn validate_sim_function_call_name( return Ok(()); } + let constructor_projection = + resolve_dae_constructor_projection_by_key(dae, name.as_str()).is_some(); let Some(func) = resolve_dae_function(dae, name) else { return Err(FunctionValidationError { name: name.as_str().to_string(), @@ -178,7 +304,10 @@ pub(super) fn validate_sim_function_call_name( }); }; - if func.external.is_some() && !eval::is_runtime_special_function_name(&func.name) { + if func.external.is_some() + && !eval::is_runtime_special_function_name(&func.name) + && !is_supported_solve_external_function(func.name.as_str()) + { return Err(FunctionValidationError { name: func.name.as_str().to_string(), reason: "external function is not supported by this simulator".to_string(), @@ -188,6 +317,7 @@ pub(super) fn validate_sim_function_call_name( if func.external.is_none() && func.body.is_empty() && !eval::is_runtime_special_function_name(&func.name) + && !constructor_projection { return Err(FunctionValidationError { name: func.name.as_str().to_string(), @@ -264,7 +394,10 @@ pub(super) fn validate_sim_component_function_call_name( }); }; - if func.external.is_some() && !eval::is_runtime_special_function_name(&func.name) { + if func.external.is_some() + && !eval::is_runtime_special_function_name(&func.name) + && !is_supported_solve_external_function(func.name.as_str()) + { return Err(FunctionValidationError { name: func.name.as_str().to_string(), reason: "external function is not supported by this simulator".to_string(), @@ -288,6 +421,16 @@ pub(super) fn is_named_function_arg_marker(name: &rumoca_core::Reference) -> boo name.as_str().starts_with(NAMED_FUNCTION_ARG_PREFIX) } +pub(super) fn is_supported_solve_external_function(name: &str) -> bool { + matches!( + name, + "Buildings.ThermalZones.EnergyPlus_9_6_0.BaseClasses.initialize" + | "Buildings.ThermalZones.EnergyPlus_9_6_0.BaseClasses.getParameters" + | "Buildings.ThermalZones.EnergyPlus_9_6_0.BaseClasses.exchange" + | "Buildings.ThermalZones.EnergyPlus_9_6_0.BaseClasses.SpawnExternalObject" + ) +} + pub(super) fn validate_called_function_body( dae: &Dae, name: &VarName, @@ -607,6 +750,14 @@ pub(super) fn validate_nested_function_call( } if !is_constructor { + validate_constructor_projection_arguments(dae, name, args)?; + validate_constructor_projection_definition( + dae, + name, + validated_functions, + active_stack, + function_param_aliases, + )?; validate_sim_function_call_name(dae, name, function_param_aliases)?; if !is_builtin_or_runtime_special(name) && !function_param_aliases.contains(name.var_name()) { diff --git a/crates/rumoca-phase-solve/src/gpu_initialization.rs b/crates/rumoca-phase-solve/src/gpu_initialization.rs new file mode 100644 index 000000000..6290a5d55 --- /dev/null +++ b/crates/rumoca-phase-solve/src/gpu_initialization.rs @@ -0,0 +1,1057 @@ +use super::*; + +/// GPU preparation deliberately accepts only direct, regular initial families. +/// It builds one base row plus one corner per binder, never a vector of scalar +/// rows. Runtime initialization keeps its complete scalar/general path. +pub(super) fn lower_gpu_initialization_system( + dae_model: &dae::Dae, + layout: &solve::VarLayout, +) -> Result { + if dae_model.initialization.equations.is_empty() { + return Ok(solve::InitializationSolveSystem::default()); + } + let mut expected = 0usize; + let mut nodes = Vec::new(); + let mut families = Vec::new(); + let mut residual_start = 0usize; + for family in &dae_model.initialization.structured_equations { + let Some(_regular) = family.regular.as_ref() else { + return Err(gpu_initial_unsupported( + "GPU initial projection requires a regular structured initial family", + family.span, + )); + }; + let Some(template) = family.template.as_ref() else { + return Err(gpu_initial_unsupported( + "GPU initial projection requires a structured initial template", + family.span, + )); + }; + let Some(body_count) = family.common_iteration_equation_count() else { + return Err(gpu_initial_unsupported( + "GPU initial projection requires a nonempty uniform structured initial family", + family.span, + )); + }; + if body_count == 0 || template.body.len() != body_count { + return Err(gpu_initial_unsupported( + "GPU initial projection requires one uniform template body per family cell", + family.span, + )); + } + let cells = family + .domain + .scalar_count() + .map_err(|error| LowerError::contract_violation(error.to_string(), family.span))?; + expected = expected + .checked_add(cells.checked_mul(body_count).ok_or_else(|| { + LowerError::contract_violation("GPU initial family size overflow", family.span) + })?) + .ok_or_else(|| { + LowerError::contract_violation("GPU initial residual size overflow", family.span) + })?; + for position in 0..body_count { + let direct = lower_gpu_direct_family( + dae_model, + layout, + family, + position, + body_count, + residual_start, + )?; + residual_start = residual_start.checked_add(cells).ok_or_else(|| { + LowerError::contract_violation("GPU initial residual range overflow", family.span) + })?; + nodes.push(direct.residual); + let node_index = nodes.len() - 1; + let direct = solve::InitializationDirectFamily { + node_index, + targets: direct.targets, + residual_sign: direct.residual_sign, + span: direct.span, + }; + families.push(direct); + } + } + let required_user_initial_rows = required_user_initial_rows(dae_model)?; + if expected != required_user_initial_rows { + return Err(gpu_initial_unsupported_optional( + "GPU initial projection requires complete structured coverage; mixed or nonstructured initial rows are unsupported", + first_uncovered_user_initial_span(dae_model), + )); + } + let (required_target_ranges, fixed_target_ranges) = + require_complete_gpu_initial_target_coverage(dae_model, layout, &families)?; + Ok(solve::InitializationSolveSystem { + residual: solve::ComputeBlock { nodes }, + direct_families: families, + required_target_ranges, + fixed_target_ranges, + ..Default::default() + }) +} + +fn required_user_initial_rows(dae_model: &dae::Dae) -> Result { + if dae_model.initialization.equation_provenance.len() + != dae_model.initialization.equations.len() + { + let span = dae_model + .initialization + .equations + .get(dae_model.initialization.equation_provenance.len()) + .or_else(|| dae_model.initialization.equations.first()) + .map(|equation| equation.span); + return Err(gpu_initial_unsupported_optional( + "GPU initial projection requires typed provenance for every initial equation", + span, + )); + } + dae_model + .initialization + .equations + .iter() + .zip(&dae_model.initialization.equation_provenance) + .filter(|(_, provenance)| **provenance != dae::InitializationEquationProvenance::FixedStart) + .map(|(equation, _)| equation) + .try_fold(0usize, |total, equation| { + total + .checked_add(equation.scalar_count.max(1)) + .ok_or_else(|| { + LowerError::contract_violation( + "GPU initial user-row count overflow", + equation.span, + ) + }) + }) +} + +fn first_uncovered_user_initial_span(dae_model: &dae::Dae) -> Option { + let mut covered = vec![false; dae_model.initialization.equations.len()]; + for family in &dae_model.initialization.structured_equations { + let equation_len = family.equation_counts.iter().copied().sum::(); + let end = family + .first_equation_index + .saturating_add(equation_len) + .min(covered.len()); + covered[family.first_equation_index.min(end)..end].fill(true); + } + dae_model + .initialization + .equations + .iter() + .zip(&dae_model.initialization.equation_provenance) + .enumerate() + .find(|(index, (_, provenance))| { + **provenance != dae::InitializationEquationProvenance::FixedStart && !covered[*index] + }) + .or_else(|| { + dae_model + .initialization + .equations + .iter() + .zip(&dae_model.initialization.equation_provenance) + .enumerate() + .find(|(_, (_, provenance))| { + **provenance != dae::InitializationEquationProvenance::FixedStart + }) + }) + .map(|(_, (equation, _))| equation.span) +} + +fn gpu_initial_unsupported(reason: impl Into, span: rumoca_core::Span) -> LowerError { + LowerError::UnsupportedAt { + reason: reason.into(), + contexts: Vec::new(), + span, + } +} + +fn gpu_initial_unsupported_optional( + reason: impl Into, + span: Option, +) -> LowerError { + let reason = reason.into(); + match span { + Some(span) => gpu_initial_unsupported(reason, span), + None => LowerError::Unsupported { reason }, + } +} + +fn require_complete_gpu_initial_target_coverage( + dae_model: &dae::Dae, + layout: &solve::VarLayout, + families: &[solve::InitializationDirectFamily], +) -> Result< + ( + Vec, + Vec, + ), + LowerError, +> { + let mut direct_ranges = Vec::with_capacity(families.len()); + for (structured, direct) in dae_model + .initialization + .structured_equations + .iter() + .flat_map(|structured| { + (0..structured.common_iteration_equation_count().unwrap_or(0)).map(move |_| structured) + }) + .zip(families) + { + let dense = + solve::TensorOutputMap::dense_contiguous(direct.targets.start, &structured.domain) + .map_err(|error| { + LowerError::contract_violation(format!("{error:?}"), direct.span) + })?; + if direct.targets.strides != dense.strides { + return Err(LowerError::contract_violation( + "GPU initial target map must be dense and contiguous", + direct.span, + )); + } + let count = structured + .domain + .scalar_count() + .map_err(|error| LowerError::contract_violation(error.to_string(), direct.span))?; + let end = direct.targets.start.checked_add(count).ok_or_else(|| { + LowerError::contract_violation("GPU initial target range overflow", direct.span) + })?; + direct_ranges.push(solve::InitializationTargetRange { + start: direct.targets.start, + end, + span: Some(direct.span), + }); + } + let mut fixed_ranges = Vec::new(); + for (equation, provenance) in dae_model + .initialization + .equations + .iter() + .zip(&dae_model.initialization.equation_provenance) + { + if *provenance != dae::InitializationEquationProvenance::FixedStart { + continue; + } + let target = lower_contiguous_y_target_range_for_equation(dae_model, equation, layout)?; + fixed_ranges.push(solve::InitializationTargetRange { + start: target.start, + end: target.end, + span: Some(equation.span), + }); + } + let fixed_ranges = normalize_gpu_target_ranges(fixed_ranges, layout.y_scalars())?; + direct_ranges.extend(fixed_ranges.iter().copied()); + let actual = normalize_gpu_target_ranges(direct_ranges, layout.y_scalars())?; + let required = if layout.y_scalars() == 0 { + Vec::new() + } else { + vec![solve::InitializationTargetRange { + start: 0, + end: layout.y_scalars(), + span: actual.first().and_then(|range| range.span), + }] + }; + if !same_gpu_target_coverage(&actual, &required) { + let span = actual + .first() + .and_then(|range| range.span) + .or_else(|| required.first().and_then(|range| range.span)); + return Err(gpu_target_range_error( + "GPU initial projection requires the union of user equations and fixed starts to cover every solver Y slot", + span, + )); + } + Ok((required, fixed_ranges)) +} + +pub(super) fn normalize_gpu_target_ranges( + mut ranges: Vec, + upper_bound: usize, +) -> Result, LowerError> { + ranges.sort_unstable_by_key(|range| (range.start, range.end)); + let mut normalized: Vec = Vec::with_capacity(ranges.len()); + for range in ranges { + if range.start >= range.end || range.end > upper_bound { + return Err(gpu_target_range_error( + "GPU initial target range is empty or outside the solver Y vector", + range.span, + )); + } + if let Some(last) = normalized.last_mut() { + if range.start < last.end { + return Err(gpu_target_range_error( + "GPU initial target ranges overlap", + range.span.or(last.span), + )); + } + if range.start == last.end { + last.end = range.end; + continue; + } + } + normalized.push(range); + } + Ok(normalized) +} + +fn gpu_target_range_error(reason: &'static str, span: Option) -> LowerError { + span.map_or_else( + || LowerError::Unsupported { + reason: reason.to_string(), + }, + |span| LowerError::contract_violation(reason, span), + ) +} + +fn same_gpu_target_coverage( + left: &[solve::InitializationTargetRange], + right: &[solve::InitializationTargetRange], +) -> bool { + left.len() == right.len() + && left + .iter() + .zip(right) + .all(|(left, right)| left.start == right.start && left.end == right.end) +} + +fn lower_gpu_direct_family( + dae_model: &dae::Dae, + layout: &solve::VarLayout, + family: &dae::StructuredEquationFamily, + position: usize, + body_count: usize, + residual_start: usize, +) -> Result { + let canonical_domain = canonical_gpu_initial_domain(&family.domain, family.span)?; + let base_cell = gpu_canonical_base_cell_index(&family.domain, family.span)?; + let base_index = family + .first_equation_index + .checked_add(base_cell.checked_mul(body_count).ok_or_else(|| { + LowerError::contract_violation("GPU initial base equation index overflow", family.span) + })?) + .and_then(|value| value.checked_add(position)) + .ok_or_else(|| { + LowerError::contract_violation("GPU initial base equation index overflow", family.span) + })?; + let base_equation = dae_model + .initialization + .equations + .get(base_index) + .ok_or_else(|| { + LowerError::contract_violation("GPU initial base equation is missing", family.span) + })?; + let base_ops = lower_initial_residual_cell( + dae_model, + layout, + dae_model.continuous.equations.len() + base_index, + base_equation, + )?; + let base_target = direct_initial_target(dae_model, layout, base_equation, family.span)?; + reject_nondeterministic_gpu_initial_ops(&base_ops, base_equation.span)?; + let sign = direct_initial_assignment_sign(&base_ops, base_target).ok_or_else(|| { + gpu_initial_unsupported( + "GPU initial projection requires a direct target-minus-rhs structured row", + base_equation.span, + ) + })?; + let strides = lower_gpu_direct_family_strides( + dae_model, + layout, + family, + position, + body_count, + GpuDirectFamilyBase { + equation: base_equation, + ops: &base_ops, + target: base_target, + }, + )?; + prove_gpu_direct_family_affine( + dae_model, + layout, + family, + position, + body_count, + GpuDirectFamilyBase { + equation: base_equation, + ops: &base_ops, + target: base_target, + }, + &strides, + )?; + Ok(GpuLoweredDirectFamily { + residual: solve::ComputeNode::Map { + domain: canonical_domain.clone(), + output_map: solve::TensorOutputMap::dense_contiguous(residual_start, &canonical_domain) + .map_err(|error| { + LowerError::contract_violation(format!("{error:?}"), family.span) + })?, + base_ops, + load_strides: strides.loads, + const_strides: strides.constants, + metadata: solve::TensorNodeMetadata::default(), + span: family.span, + }, + targets: solve::TensorOutputMap { + start: base_target, + strides: strides.targets, + }, + residual_sign: sign, + span: family.span, + }) +} + +struct GpuLoweredDirectFamily { + residual: solve::ComputeNode, + targets: solve::TensorOutputMap, + residual_sign: i8, + span: rumoca_core::Span, +} + +struct GpuDirectFamilyStrides { + loads: Vec, + constants: Vec, + targets: Vec, +} + +struct GpuDirectFamilyBase<'a> { + equation: &'a dae::Equation, + ops: &'a [solve::LinearOp], + target: usize, +} + +struct GpuDirectFamilyProof<'a> { + dae_model: &'a dae::Dae, + layout: &'a solve::VarLayout, + family: &'a dae::StructuredEquationFamily, + base: GpuDirectFamilyBase<'a>, + strides: &'a GpuDirectFamilyStrides, +} + +fn lower_gpu_direct_family_strides( + dae_model: &dae::Dae, + layout: &solve::VarLayout, + family: &dae::StructuredEquationFamily, + position: usize, + body_count: usize, + base: GpuDirectFamilyBase<'_>, +) -> Result { + let mut load_strides = Vec::new(); + let mut const_strides = Vec::new(); + let mut target_strides = Vec::new(); + for (dimension, binder) in family.domain.binders.iter().enumerate() { + if gpu_binder_value_count(binder, family.span)? == 1 { + continue; + } + let corner_index = gpu_direct_family_corner_index(family, position, body_count, dimension)?; + let corner_equation = dae_model + .initialization + .equations + .get(corner_index) + .ok_or_else(|| { + LowerError::contract_violation( + "GPU initial corner equation is missing", + family.span, + ) + })?; + let corner_ops = lower_initial_residual_cell( + dae_model, + layout, + dae_model.continuous.equations.len() + corner_index, + corner_equation, + )?; + if !stencil::dae_equation_body_shapes_match(base.equation, corner_equation)? { + return Err(gpu_initial_unsupported( + "GPU initial projection requires identical conservative equation body shapes", + corner_equation.span, + )); + } + let corner_target = direct_initial_target(dae_model, layout, corner_equation, family.span)?; + target_strides.push(solve::AffineStencilIndexStrideTerm { + dimension, + stride: gpu_initial_stride(corner_target, base.target, family.span, "target")?, + }); + append_gpu_corner_strides( + base.ops, + &corner_ops, + dimension, + &mut load_strides, + &mut const_strides, + family.span, + )?; + } + Ok(GpuDirectFamilyStrides { + loads: load_strides, + constants: const_strides, + targets: target_strides, + }) +} + +fn prove_gpu_direct_family_affine( + dae_model: &dae::Dae, + layout: &solve::VarLayout, + family: &dae::StructuredEquationFamily, + position: usize, + body_count: usize, + base: GpuDirectFamilyBase<'_>, + strides: &GpuDirectFamilyStrides, +) -> Result<(), LowerError> { + let proof = GpuDirectFamilyProof { + dae_model, + layout, + family, + base, + strides, + }; + if !family.interiors_materialized { + return Err(gpu_initial_unsupported( + "GPU initial projection cannot prove affine direct-family values without materialized interiors", + family.span, + )); + } + let tuples = family + .domain + .index_tuples() + .map_err(|error| LowerError::contract_violation(error.to_string(), family.span))?; + let mut equations = Vec::with_capacity(tuples.len()); + for cell in 0..tuples.len() { + let equation_index = family + .first_equation_index + .checked_add(cell.checked_mul(body_count).ok_or_else(|| { + LowerError::contract_violation("GPU initial proof row index overflow", family.span) + })?) + .and_then(|index| index.checked_add(position)) + .ok_or_else(|| { + LowerError::contract_violation("GPU initial proof row index overflow", family.span) + })?; + let equation = dae_model + .initialization + .equations + .get(equation_index) + .ok_or_else(|| { + LowerError::contract_violation( + "GPU initial affine proof equation is missing", + family.span, + ) + })?; + equations.push(( + dae_model.continuous.equations.len() + equation_index, + equation, + )); + } + let rows = lower::lower_initial_residual_cells( + dae_model, + layout, + equations + .iter() + .map(|(index, equation)| (*index, *equation)), + )?; + if rows.len() != equations.len() { + return Err(LowerError::contract_violation( + "GPU initial affine proof must lower one residual row per family cell", + family.span, + )); + } + for (((_, equation), ops), tuple) in equations.iter().zip(&rows).zip(&tuples) { + prove_gpu_direct_family_cell(&proof, tuple, equation, ops)?; + } + Ok(()) +} + +fn prove_gpu_direct_family_cell( + proof: &GpuDirectFamilyProof<'_>, + tuple: &[i64], + equation: &dae::Equation, + ops: &[solve::LinearOp], +) -> Result<(), LowerError> { + if !stencil::dae_equation_body_shapes_match(proof.base.equation, equation)? { + return Err(gpu_initial_unsupported( + "GPU initial projection requires identical conservative equation body shapes", + equation.span, + )); + } + let ordinals = gpu_canonical_ordinals(&proof.family.domain, tuple, equation.span)?; + reject_nondeterministic_gpu_initial_ops(ops, equation.span)?; + prove_gpu_affine_ops(proof.base.ops, ops, &ordinals, proof.strides, equation.span)?; + let target = direct_initial_target(proof.dae_model, proof.layout, equation, equation.span)?; + let expected = affine_gpu_target( + proof.base.target, + &ordinals, + &proof.strides.targets, + equation.span, + )?; + if target != expected { + return Err(gpu_initial_unsupported( + "GPU initial projection target is not affine across the complete family domain", + equation.span, + )); + } + Ok(()) +} + +fn gpu_canonical_ordinals( + domain: &rumoca_core::StructuredIndexDomain, + tuple: &[i64], + span: rumoca_core::Span, +) -> Result, LowerError> { + domain + .binders + .iter() + .zip(tuple) + .map(|(binder, value)| { + let lower = binder.lower.min(binder.upper); + let distance = value.checked_sub(lower).ok_or_else(|| { + LowerError::contract_violation("GPU initial proof tuple is out of bounds", span) + })?; + let step = i64::try_from(binder.step.unsigned_abs()).map_err(|_| { + LowerError::contract_violation("GPU initial proof step exceeds host range", span) + })?; + usize::try_from(distance / step).map_err(|_| { + LowerError::contract_violation("GPU initial proof ordinal exceeds host range", span) + }) + }) + .collect() +} + +fn prove_gpu_affine_ops( + base: &[solve::LinearOp], + actual: &[solve::LinearOp], + ordinals: &[usize], + strides: &GpuDirectFamilyStrides, + span: rumoca_core::Span, +) -> Result<(), LowerError> { + if base.len() != actual.len() { + return Err(gpu_initial_unsupported( + "GPU initial projection operation shape is not uniform across the complete family domain", + span, + )); + } + for (position, (base_op, actual_op)) in base.iter().zip(actual).enumerate() { + if !gpu_affine_op_matches(position, base_op, actual_op, ordinals, strides, span)? { + return Err(gpu_initial_unsupported( + "GPU initial projection operation values are not affine across the complete family domain", + span, + )); + } + } + Ok(()) +} + +fn gpu_affine_op_matches( + position: usize, + base: &solve::LinearOp, + actual: &solve::LinearOp, + ordinals: &[usize], + strides: &GpuDirectFamilyStrides, + span: rumoca_core::Span, +) -> Result { + match (base, actual) { + ( + solve::LinearOp::LoadY { + dst: base_dst, + index: base, + }, + solve::LinearOp::LoadY { + dst: actual_dst, + index: actual, + }, + ) + | ( + solve::LinearOp::LoadP { + dst: base_dst, + index: base, + }, + solve::LinearOp::LoadP { + dst: actual_dst, + index: actual, + }, + ) => Ok(base_dst == actual_dst + && affine_gpu_index(*base, position, ordinals, &strides.loads, span)? == *actual), + ( + solve::LinearOp::Const { + dst: base_dst, + value: base, + }, + solve::LinearOp::Const { + dst: actual_dst, + value: actual, + }, + ) => Ok(base_dst == actual_dst + && affine_gpu_constant(*base, position, ordinals, &strides.constants, span)?.to_bits() + == actual.to_bits()), + _ => Ok(base == actual), + } +} + +fn affine_gpu_index( + base: usize, + position: usize, + ordinals: &[usize], + strides: &[solve::AffineStencilLoadStride], + span: rumoca_core::Span, +) -> Result { + let offset = strides + .iter() + .filter(|stride| stride.op_position == position) + .flat_map(|stride| &stride.terms) + .try_fold(0isize, |total, term| { + let ordinal = isize::try_from(ordinals[term.dimension]).ok()?; + total.checked_add(term.stride.checked_mul(ordinal)?) + }) + .ok_or_else(|| { + LowerError::contract_violation("GPU initial affine index overflows", span) + })?; + base.checked_add_signed(offset) + .ok_or_else(|| LowerError::contract_violation("GPU initial affine index overflows", span)) +} + +fn affine_gpu_constant( + base: f64, + position: usize, + ordinals: &[usize], + strides: &[solve::AffineStencilConstStride], + span: rumoca_core::Span, +) -> Result { + let value = strides + .iter() + .filter(|stride| stride.op_position == position) + .flat_map(|stride| &stride.terms) + .fold(base, |value, term| { + value + term.stride * ordinals[term.dimension] as f64 + }); + value.is_finite().then_some(value).ok_or_else(|| { + LowerError::contract_violation("GPU initial affine constant is not finite", span) + }) +} + +fn affine_gpu_target( + base: usize, + ordinals: &[usize], + strides: &[solve::AffineStencilIndexStrideTerm], + span: rumoca_core::Span, +) -> Result { + let offset = strides.iter().try_fold(0isize, |total, term| { + let ordinal = isize::try_from(ordinals[term.dimension]).ok()?; + total.checked_add(term.stride.checked_mul(ordinal)?) + }); + offset + .and_then(|offset| base.checked_add_signed(offset)) + .ok_or_else(|| LowerError::contract_violation("GPU initial affine target overflows", span)) +} + +pub(super) fn reject_nondeterministic_gpu_initial_ops( + ops: &[solve::LinearOp], + span: rumoca_core::Span, +) -> Result<(), LowerError> { + if ops.iter().any(|op| { + matches!( + op, + solve::LinearOp::RandomInitialState { .. } + | solve::LinearOp::RandomResult { .. } + | solve::LinearOp::RandomState { .. } + | solve::LinearOp::ImpureRandomInit { .. } + | solve::LinearOp::ImpureRandom { .. } + | solve::LinearOp::ImpureRandomInteger { .. } + ) + }) { + return Err(gpu_initial_unsupported( + "GPU initial projection rejects random or impure operations because apply and verification must be deterministic", + span, + )); + } + Ok(()) +} + +fn gpu_direct_family_corner_index( + family: &dae::StructuredEquationFamily, + position: usize, + body_count: usize, + dimension: usize, +) -> Result { + let corner_cell = gpu_corner_cell_index(&family.domain, dimension, family.span)?; + family + .first_equation_index + .checked_add(corner_cell.checked_mul(body_count).ok_or_else(|| { + LowerError::contract_violation( + "GPU initial corner equation index overflow", + family.span, + ) + })?) + .and_then(|value| value.checked_add(position)) + .ok_or_else(|| { + LowerError::contract_violation( + "GPU initial corner equation index overflow", + family.span, + ) + }) +} + +fn gpu_initial_stride( + corner: usize, + base: usize, + span: rumoca_core::Span, + kind: &'static str, +) -> Result { + isize::try_from(corner) + .ok() + .and_then(|value| value.checked_sub(isize::try_from(base).ok()?)) + .ok_or_else(|| { + LowerError::contract_violation(format!("GPU initial {kind} stride overflows"), span) + }) +} + +fn direct_initial_target( + dae_model: &dae::Dae, + layout: &solve::VarLayout, + equation: &dae::Equation, + span: rumoca_core::Span, +) -> Result { + let targets = lower_continuous_row_targets_for_equation(dae_model, equation, layout, 1)?; + match targets.as_slice() { + [Some(solve::ScalarSlot::Y { index, .. })] => Ok(*index), + _ => Err(LowerError::contract_violation( + "GPU initial projection requires one Y target per direct family row", + span, + )), + } +} + +pub(super) fn gpu_corner_cell_index( + domain: &rumoca_core::StructuredIndexDomain, + dimension: usize, + span: rumoca_core::Span, +) -> Result { + let selected = domain.binders.get(dimension).ok_or_else(|| { + LowerError::contract_violation("GPU initial corner dimension is missing", span) + })?; + if gpu_binder_value_count(selected, span)? < 2 { + return Err(gpu_initial_unsupported( + "GPU initial projection requires a non-degenerate structured binder", + span, + )); + } + domain + .binders + .iter() + .enumerate() + .try_fold(0usize, |ordinal, (index, binder)| { + let count = gpu_binder_value_count(binder, span)?; + let coordinate = match (binder.step < 0, index == dimension) { + (false, false) => 0, + (false, true) => 1, + (true, false) => count - 1, + (true, true) => count - 2, + }; + ordinal + .checked_mul(count) + .and_then(|value| value.checked_add(coordinate)) + .ok_or_else(|| { + LowerError::contract_violation("GPU initial corner stride overflow", span) + }) + }) +} + +fn gpu_canonical_base_cell_index( + domain: &rumoca_core::StructuredIndexDomain, + span: rumoca_core::Span, +) -> Result { + domain.binders.iter().try_fold(0usize, |ordinal, binder| { + let count = gpu_binder_value_count(binder, span)?; + let coordinate = if binder.step < 0 { count - 1 } else { 0 }; + ordinal + .checked_mul(count) + .and_then(|value| value.checked_add(coordinate)) + .ok_or_else(|| { + LowerError::contract_violation("GPU initial base cell index overflow", span) + }) + }) +} + +fn canonical_gpu_initial_domain( + domain: &rumoca_core::StructuredIndexDomain, + span: rumoca_core::Span, +) -> Result { + let mut canonical = domain.clone(); + for binder in &mut canonical.binders { + if binder.step < 0 { + std::mem::swap(&mut binder.lower, &mut binder.upper); + binder.step = binder.step.checked_neg().ok_or_else(|| { + LowerError::contract_violation("GPU initial binder step overflows", span) + })?; + } + } + Ok(canonical) +} + +fn gpu_binder_value_count( + binder: &rumoca_core::StructuredIndexBinder, + span: rumoca_core::Span, +) -> Result { + if binder.step == 0 { + return Err(LowerError::contract_violation( + "GPU initial binder step must be nonzero", + span, + )); + } + let distance = if binder.step > 0 { + binder.upper.checked_sub(binder.lower) + } else { + binder.lower.checked_sub(binder.upper) + } + .ok_or_else(|| LowerError::contract_violation("GPU initial binder bounds are invalid", span))?; + let step = binder.step.unsigned_abs(); + let count = distance + .checked_div(i64::try_from(step).map_err(|_| { + LowerError::contract_violation("GPU initial binder step overflow", span) + })?) + .and_then(|value| value.checked_add(1)) + .ok_or_else(|| LowerError::contract_violation("GPU initial binder count overflow", span))?; + usize::try_from(count).map_err(|_| { + LowerError::contract_violation("GPU initial binder count exceeds host range", span) + }) +} + +pub(super) fn append_gpu_corner_strides( + base: &[solve::LinearOp], + corner: &[solve::LinearOp], + dimension: usize, + load_strides: &mut Vec, + const_strides: &mut Vec, + span: rumoca_core::Span, +) -> Result<(), LowerError> { + if base.len() != corner.len() { + return Err(gpu_initial_unsupported( + "GPU initial projection requires identical direct-family operation shapes", + span, + )); + } + for (op_position, (base_op, corner_op)) in base.iter().zip(corner).enumerate() { + match (base_op, corner_op) { + ( + solve::LinearOp::LoadY { + dst: base_dst, + index: base, + }, + solve::LinearOp::LoadY { + dst: corner_dst, + index: corner, + }, + ) if base_dst == corner_dst => { + let stride = isize::try_from(*corner) + .ok() + .and_then(|value| value.checked_sub(isize::try_from(*base).ok()?)) + .ok_or_else(|| { + LowerError::contract_violation("GPU initial Y stride overflows", span) + })?; + if stride != 0 { + load_strides.push(solve::AffineStencilLoadStride { + op_position, + terms: vec![solve::AffineStencilIndexStrideTerm { dimension, stride }], + }); + } + } + ( + solve::LinearOp::LoadP { + dst: base_dst, + index: base, + }, + solve::LinearOp::LoadP { + dst: corner_dst, + index: corner, + }, + ) if base_dst == corner_dst => { + let stride = isize::try_from(*corner) + .ok() + .and_then(|value| value.checked_sub(isize::try_from(*base).ok()?)) + .ok_or_else(|| { + LowerError::contract_violation("GPU initial P stride overflows", span) + })?; + if stride != 0 { + load_strides.push(solve::AffineStencilLoadStride { + op_position, + terms: vec![solve::AffineStencilIndexStrideTerm { dimension, stride }], + }); + } + } + ( + solve::LinearOp::Const { + dst: base_dst, + value: base, + }, + solve::LinearOp::Const { + dst: corner_dst, + value: corner, + }, + ) if base_dst == corner_dst => { + let stride = corner - base; + if !stride.is_finite() { + return Err(LowerError::contract_violation( + "GPU initial constant stride is not finite", + span, + )); + } + if stride != 0.0 { + const_strides.push(solve::AffineStencilConstStride { + op_position, + terms: vec![solve::AffineStencilConstStrideTerm { dimension, stride }], + }); + } + } + ( + solve::LinearOp::LoadY { .. } + | solve::LinearOp::LoadP { .. } + | solve::LinearOp::Const { .. }, + _, + ) => { + return Err(gpu_initial_unsupported( + "GPU initial projection requires uniform direct-family access kinds", + span, + )); + } + _ if base_op == corner_op => {} + _ => { + return Err(gpu_initial_unsupported( + "GPU initial projection requires every non-affine operation and destination register to match exactly", + span, + )); + } + } + } + Ok(()) +} + +fn direct_initial_assignment_sign(ops: &[solve::LinearOp], target_index: usize) -> Option { + let solve::LinearOp::StoreOutput { src } = ops.last()? else { + return None; + }; + let solve::LinearOp::Binary { + op: solve::BinaryOp::Sub, + lhs, + rhs, + dst, + } = ops + .iter() + .find(|op| matches!(op, solve::LinearOp::Binary { dst, .. } if dst == src))? + else { + return None; + }; + let target_loads = ops + .iter() + .filter_map(|op| match op { + solve::LinearOp::LoadY { dst, index } if *index == target_index => Some(*dst), + solve::LinearOp::LoadY { .. } => Some(u32::MAX), + _ => None, + }) + .collect::>(); + if target_loads.len() != 1 || target_loads[0] == u32::MAX || *dst != *src { + return None; + } + let residual_sign = if target_loads[0] == *lhs { + 1 + } else if target_loads[0] == *rhs { + -1 + } else { + return None; + }; + Some(residual_sign) +} diff --git a/crates/rumoca-phase-solve/src/implicit_rhs.rs b/crates/rumoca-phase-solve/src/implicit_rhs.rs index 41f5be194..6305f3f50 100644 --- a/crates/rumoca-phase-solve/src/implicit_rhs.rs +++ b/crates/rumoca-phase-solve/src/implicit_rhs.rs @@ -39,8 +39,12 @@ pub(crate) fn build_implicit_rhs_rows( ) -> Result { let span = context_span; let mut rows = implicit_rhs_vec_with_capacity(solver_scalar_count, "implicit RHS rows", span)?; - for _ in 0..solver_scalar_count { - rows.push(implicit_rhs_zero_row(span)?); + for index in 0..solver_scalar_count { + if index < state_scalar_count { + rows.push(implicit_rhs_zero_row(span)?); + } else { + rows.push(implicit_rhs_y_identity_row(index, span)?); + } } let mut row_targets = implicit_rhs_vec_with_capacity(solver_scalar_count, "implicit RHS row targets", span)?; @@ -238,6 +242,16 @@ fn implicit_rhs_zero_row(span: rumoca_core::Span) -> Result Ok(row) } +fn implicit_rhs_y_identity_row( + index: usize, + span: rumoca_core::Span, +) -> Result, LowerError> { + let mut row = implicit_rhs_vec_with_capacity(2, "implicit RHS identity row op count", span)?; + row.push(solve::LinearOp::LoadY { dst: 0, index }); + row.push(solve::LinearOp::StoreOutput { src: 0 }); + Ok(row) +} + fn scalar_program_block_with_source_span( programs: Vec>, span: rumoca_core::Span, @@ -385,7 +399,7 @@ fn uncovered_implicit_tail_rows( Ok((rows, output_indices)) } -fn remap_residual_compute_nodes( +pub(crate) fn remap_residual_compute_nodes( residual_block: &solve::ComputeBlock, residual_to_implicit_rows: &[Option], implicit_output_cursor: usize, @@ -491,9 +505,30 @@ fn remap_residual_compute_node( }, ), ), - solve::ComputeNode::MatMul { .. } | solve::ComputeNode::LinSolve { .. } => Ok( - contiguous_at_cursor(implicit_indices, implicit_output_cursor).then(|| node.clone()), - ), + solve::ComputeNode::MatMul { .. } => Ok(contiguous_at_cursor( + implicit_indices, + implicit_output_cursor, + ) + .then(|| node.clone())), + solve::ComputeNode::LinSolve { + setup_ops, + matrix_start, + rhs_start, + n, + next_reg, + metadata, + span, + .. + } => Ok(Some(solve::ComputeNode::LinSolve { + setup_ops: setup_ops.clone(), + matrix_start: *matrix_start, + rhs_start: *rhs_start, + n: *n, + next_reg: *next_reg, + output_indices: implicit_indices.to_vec(), + metadata: metadata.clone(), + span: *span, + })), } } @@ -540,8 +575,14 @@ fn residual_compute_node_output_indices( solve::ComputeNode::MatMul { m, n, .. } => { contiguous_output_indices(output_cursor, checked_output_product(*m, *n, span)?, span) } - solve::ComputeNode::LinSolve { n, .. } => { - contiguous_output_indices(output_cursor, *n, span) + solve::ComputeNode::LinSolve { + n, output_indices, .. + } => { + if output_indices.is_empty() { + contiguous_output_indices(output_cursor, *n, span) + } else { + Ok(output_indices.clone()) + } } } } @@ -632,6 +673,42 @@ fn place_targeted_residual_rows( residual_to_implicit_rows[residual_idx] = Some(index); } + augment_residual_ownership( + rows, + row_targets, + occupied, + residual, + residual_targets, + state_scalar_count, + solver_scalar_count, + &mut placed, + residual_to_implicit_rows, + span, + )?; + + let structural_fallback_targets = match_unowned_residuals_to_free_algebraics( + residual, + residual_targets, + &placed, + occupied, + state_scalar_count, + solver_scalar_count, + span, + )?; + for (residual_idx, target_idx) in structural_fallback_targets.into_iter().enumerate() { + let Some(target_idx) = target_idx else { + continue; + }; + rows[target_idx] = residual[residual_idx].clone(); + row_targets[target_idx] = match residual_targets.get(residual_idx).copied().flatten() { + Some(target) => fallback_residual_row_target(Some(target), state_scalar_count, span)?, + None => Some(implicit_rhs_y_slot(target_idx, span)?), + }; + occupied[target_idx] = true; + placed[residual_idx] = true; + residual_to_implicit_rows[residual_idx] = Some(target_idx); + } + let mut fallback_idx = state_scalar_count; for (residual_idx, row) in residual.iter().enumerate() { if placed[residual_idx] { @@ -652,18 +729,238 @@ fn place_targeted_residual_rows( Ok(()) } +#[allow(clippy::too_many_arguments)] +fn augment_residual_ownership( + rows: &mut [Vec], + row_targets: &mut [Option], + occupied: &mut [bool], + residual: &[Vec], + residual_targets: &[Option], + state_scalar_count: usize, + solver_scalar_count: usize, + placed: &mut [bool], + residual_to_implicit_rows: &mut [Option], + span: rumoca_core::Span, +) -> Result<(), LowerError> { + let projection_set = (state_scalar_count..solver_scalar_count).collect(); + let mut candidates = + implicit_rhs_vec_with_capacity(residual.len(), "residual ownership candidate count", span)?; + for (residual_idx, row) in residual.iter().enumerate() { + let mut row_candidates = super::collect_algebraic_y_indices_for_row(row, &projection_set) + .into_iter() + .collect::>(); + row_candidates.sort_unstable(); + if let Some(solve::ScalarSlot::Y { index, .. }) = + residual_targets.get(residual_idx).copied().flatten() + && index >= state_scalar_count + && index < solver_scalar_count + { + row_candidates.retain(|candidate| *candidate != index); + row_candidates.insert(0, index); + } + candidates.push(row_candidates); + } + + let mut owner_by_y = implicit_rhs_vec_with_capacity( + solver_scalar_count, + "residual ownership target count", + span, + )?; + owner_by_y.resize(solver_scalar_count, None); + for (residual_idx, target) in residual_to_implicit_rows.iter().copied().enumerate() { + if let Some(target) = target { + owner_by_y[target] = Some(residual_idx); + } + } + for (residual_idx, is_placed) in placed.iter().copied().enumerate().take(residual.len()) { + if is_placed { + continue; + } + let mut visited = implicit_rhs_vec_with_capacity( + solver_scalar_count, + "residual ownership visited target count", + span, + )?; + visited.resize(solver_scalar_count, false); + let _ = augment_residual_owner( + residual_idx, + &candidates, + &mut owner_by_y, + &mut visited, + state_scalar_count, + ); + } + + for index in state_scalar_count..solver_scalar_count { + rows[index] = implicit_rhs_y_identity_row(index, span)?; + row_targets[index] = None; + occupied[index] = false; + } + placed.fill(false); + residual_to_implicit_rows.fill(None); + for (target_idx, residual_idx) in owner_by_y.into_iter().enumerate() { + let Some(residual_idx) = residual_idx else { + continue; + }; + rows[target_idx] = residual[residual_idx].clone(); + row_targets[target_idx] = match residual_targets.get(residual_idx).copied().flatten() { + Some(solve::ScalarSlot::Y { index, .. }) if index < state_scalar_count => { + Some(implicit_rhs_y_slot(index, span)?) + } + _ => Some(implicit_rhs_y_slot(target_idx, span)?), + }; + occupied[target_idx] = true; + placed[residual_idx] = true; + residual_to_implicit_rows[residual_idx] = Some(target_idx); + } + Ok(()) +} + +fn augment_residual_owner( + residual_idx: usize, + candidates: &[Vec], + owner_by_y: &mut [Option], + visited: &mut [bool], + state_scalar_count: usize, +) -> bool { + for target_idx in candidates[residual_idx].iter().copied() { + if target_idx < state_scalar_count || target_idx >= owner_by_y.len() || visited[target_idx] + { + continue; + } + visited[target_idx] = true; + let can_claim = owner_by_y[target_idx].is_none_or(|owner| { + augment_residual_owner(owner, candidates, owner_by_y, visited, state_scalar_count) + }); + if can_claim { + owner_by_y[target_idx] = Some(residual_idx); + return true; + } + } + false +} + +#[allow(clippy::too_many_arguments)] +fn match_unowned_residuals_to_free_algebraics( + residual: &[Vec], + residual_targets: &[Option], + placed: &[bool], + occupied: &[bool], + state_scalar_count: usize, + solver_scalar_count: usize, + span: rumoca_core::Span, +) -> Result>, LowerError> { + let mut matched_targets = + implicit_rhs_vec_with_capacity(residual.len(), "fallback residual target count", span)?; + matched_targets.resize(residual.len(), None); + + let available_y_indices = (state_scalar_count..solver_scalar_count) + .filter(|index| !occupied[*index]) + .collect::>(); + let available_positions = available_y_indices + .iter() + .copied() + .enumerate() + .map(|(position, y_index)| (y_index, position)) + .collect::>(); + let available_projection_set = available_y_indices.iter().copied().collect(); + let mut equation_refs = Vec::new(); + let mut eq_unknowns = Vec::new(); + for (residual_idx, row) in residual.iter().enumerate() { + let has_owned_algebraic_target = matches!( + residual_targets.get(residual_idx).copied().flatten(), + Some(solve::ScalarSlot::Y { index, .. }) + if index >= state_scalar_count && index < solver_scalar_count + ); + if placed[residual_idx] || has_owned_algebraic_target { + continue; + } + let unknowns = super::collect_algebraic_y_indices_for_row(row, &available_projection_set) + .into_iter() + .filter_map(|index| available_positions.get(&index).copied()) + .collect::>(); + if unknowns.is_empty() { + continue; + } + equation_refs.push(rumoca_phase_structural::EquationRef(residual_idx)); + eq_unknowns.push(unknowns); + } + if equation_refs.is_empty() { + return Ok(matched_targets); + } + let unknown_names = available_y_indices + .iter() + .copied() + .map(rumoca_phase_structural::UnknownId::SolverY) + .collect::>(); + let incidence = + rumoca_phase_structural::Incidence::new(eq_unknowns, equation_refs, unknown_names); + let regular = + rumoca_phase_structural::maximum_regular_subsystem(&incidence).map_err(|err| { + implicit_rhs_contract_violation( + format!("match fallback residual algebraic targets: {err}"), + span, + ) + })?; + let blocks = + rumoca_phase_structural::build_blt_from_incidence(®ular.incidence).map_err(|err| { + implicit_rhs_contract_violation( + format!("match fallback residual algebraic targets: {err}"), + span, + ) + })?; + for block in blocks { + match block { + rumoca_phase_structural::BltBlock::Scalar { equation, unknown } => { + assign_fallback_residual_target(&mut matched_targets, equation, unknown, span)?; + } + rumoca_phase_structural::BltBlock::AlgebraicLoop { + equations, + unknowns, + } => { + for (equation, unknown) in equations.into_iter().zip(unknowns) { + assign_fallback_residual_target(&mut matched_targets, equation, unknown, span)?; + } + } + } + } + Ok(matched_targets) +} + +fn assign_fallback_residual_target( + matched_targets: &mut [Option], + equation: rumoca_phase_structural::EquationRef, + unknown: rumoca_phase_structural::UnknownId, + span: rumoca_core::Span, +) -> Result<(), LowerError> { + let rumoca_phase_structural::UnknownId::SolverY(y_index) = unknown else { + return Err(implicit_rhs_contract_violation( + "fallback residual matching returned a non-solver unknown", + span, + )); + }; + let Some(target) = matched_targets.get_mut(equation.0) else { + return Err(implicit_rhs_contract_violation( + format!( + "fallback residual matching returned equation {} outside residual rows", + equation.0 + ), + span, + )); + }; + *target = Some(y_index); + Ok(()) +} + fn fallback_residual_row_target( residual_target: Option, - state_scalar_count: usize, + _state_scalar_count: usize, span: rumoca_core::Span, ) -> Result, LowerError> { let Some(solve::ScalarSlot::Y { index, .. }) = residual_target else { return Ok(None); }; - if index < state_scalar_count { - return implicit_rhs_y_slot(index, span).map(Some); - } - Ok(None) + implicit_rhs_y_slot(index, span).map(Some) } fn implicit_rhs_y_slot( @@ -807,6 +1104,293 @@ mod tests { Ok(()) } + #[test] + fn unowned_residual_displaces_connection_target_to_preserve_real_producers() + -> Result<(), LowerError> { + let span = test_span(); + let assignment = vec![ + solve::LinearOp::Const { dst: 0, value: 1.0 }, + solve::LinearOp::LoadY { dst: 1, index: 0 }, + solve::LinearOp::Binary { + dst: 2, + op: solve::BinaryOp::Sub, + lhs: 0, + rhs: 1, + }, + solve::LinearOp::StoreOutput { src: 2 }, + ]; + let physical = vec![ + solve::LinearOp::LoadP { dst: 0, index: 0 }, + solve::LinearOp::LoadY { dst: 1, index: 2 }, + solve::LinearOp::Binary { + dst: 2, + op: solve::BinaryOp::Mul, + lhs: 0, + rhs: 1, + }, + solve::LinearOp::LoadY { dst: 3, index: 1 }, + solve::LinearOp::Binary { + dst: 4, + op: solve::BinaryOp::Sub, + lhs: 2, + rhs: 3, + }, + solve::LinearOp::StoreOutput { src: 4 }, + ]; + let connection = vec![ + solve::LinearOp::LoadY { dst: 0, index: 1 }, + solve::LinearOp::LoadY { dst: 1, index: 2 }, + solve::LinearOp::Binary { + dst: 2, + op: solve::BinaryOp::Sub, + lhs: 0, + rhs: 1, + }, + solve::LinearOp::StoreOutput { src: 2 }, + ]; + let residual = vec![assignment.clone(), physical.clone(), connection.clone()]; + let targets = vec![ + Some(solve::scalar_slot_y(0)), + None, + Some(solve::scalar_slot_y(1)), + ]; + + let implicit = build_implicit_rhs_rows(&[], &residual, &targets, 0, 3, span)?; + + assert_eq!( + implicit.residual_to_implicit_rows, + vec![Some(0), Some(1), Some(2)] + ); + assert_eq!(implicit.rows, vec![assignment, physical, connection]); + assert_eq!( + implicit.row_targets, + vec![ + Some(solve::scalar_slot_y(0)), + Some(solve::scalar_slot_y(1)), + Some(solve::scalar_slot_y(2)), + ] + ); + Ok(()) + } + + #[test] + fn fallback_residual_claims_a_free_loaded_algebraic_target() -> Result<(), LowerError> { + let mut rows = vec![ + zero_rhs_row(), + zero_rhs_row(), + zero_rhs_row(), + zero_rhs_row(), + ]; + let mut row_targets = vec![None, None, None, None]; + let mut occupied = vec![true, false, false, false]; + let residual = vec![ + vec![ + solve::LinearOp::LoadY { dst: 0, index: 1 }, + solve::LinearOp::StoreOutput { src: 0 }, + ], + vec![ + solve::LinearOp::LoadY { dst: 0, index: 1 }, + solve::LinearOp::LoadY { dst: 1, index: 3 }, + solve::LinearOp::LoadY { dst: 3, index: 2 }, + solve::LinearOp::Binary { + dst: 2, + op: solve::BinaryOp::Add, + lhs: 0, + rhs: 1, + }, + solve::LinearOp::StoreOutput { src: 2 }, + ], + ]; + let residual_targets = vec![Some(solve::scalar_slot_y(1)), None]; + let mut residual_to_implicit_rows = vec![None, None]; + + place_targeted_residual_rows( + &mut rows, + &mut row_targets, + &mut occupied, + &residual, + &residual_targets, + 1, + 4, + &mut residual_to_implicit_rows, + test_span(), + )?; + + assert_eq!(residual_to_implicit_rows, vec![Some(1), Some(3)]); + assert_eq!(row_targets[3], Some(solve::scalar_slot_y(3))); + assert_eq!(rows[3], residual[1]); + Ok(()) + } + + #[test] + fn state_free_system_structurally_matches_fallback_residual_owner() -> Result<(), LowerError> { + let mut rows = vec![zero_rhs_row(), zero_rhs_row(), zero_rhs_row()]; + let mut row_targets = vec![None, None, None]; + let mut occupied = vec![false, false, false]; + let residual = vec![ + vec![ + solve::LinearOp::LoadY { dst: 0, index: 0 }, + solve::LinearOp::StoreOutput { src: 0 }, + ], + vec![ + solve::LinearOp::LoadY { dst: 0, index: 0 }, + solve::LinearOp::LoadY { dst: 1, index: 2 }, + solve::LinearOp::Binary { + dst: 2, + op: solve::BinaryOp::Add, + lhs: 0, + rhs: 1, + }, + solve::LinearOp::StoreOutput { src: 2 }, + ], + ]; + let residual_targets = vec![Some(solve::scalar_slot_y(0)), None]; + let mut residual_to_implicit_rows = vec![None, None]; + + place_targeted_residual_rows( + &mut rows, + &mut row_targets, + &mut occupied, + &residual, + &residual_targets, + 0, + 3, + &mut residual_to_implicit_rows, + test_span(), + )?; + + assert_eq!(residual_to_implicit_rows, vec![Some(0), Some(2)]); + assert_eq!(row_targets[2], Some(solve::scalar_slot_y(2))); + assert_eq!(rows[2], residual[1]); + Ok(()) + } + + #[test] + fn state_consistency_row_participates_in_free_algebraic_matching() -> Result<(), LowerError> { + let mut rows = vec![ + zero_rhs_row(), + zero_rhs_row(), + zero_rhs_row(), + zero_rhs_row(), + ]; + let mut row_targets = vec![None, None, None, None]; + let mut occupied = vec![true, false, false, false]; + let residual = vec![ + vec![ + solve::LinearOp::LoadY { dst: 0, index: 1 }, + solve::LinearOp::StoreOutput { src: 0 }, + ], + vec![ + solve::LinearOp::LoadY { dst: 0, index: 2 }, + solve::LinearOp::LoadY { dst: 1, index: 3 }, + solve::LinearOp::Binary { + dst: 2, + op: solve::BinaryOp::Add, + lhs: 0, + rhs: 1, + }, + solve::LinearOp::StoreOutput { src: 2 }, + ], + vec![ + solve::LinearOp::LoadY { dst: 0, index: 0 }, + solve::LinearOp::LoadY { dst: 1, index: 2 }, + solve::LinearOp::Binary { + dst: 2, + op: solve::BinaryOp::Sub, + lhs: 0, + rhs: 1, + }, + solve::LinearOp::StoreOutput { src: 2 }, + ], + ]; + let residual_targets = vec![ + Some(solve::scalar_slot_y(1)), + None, + Some(solve::scalar_slot_y(0)), + ]; + let mut residual_to_implicit_rows = vec![None, None, None]; + + place_targeted_residual_rows( + &mut rows, + &mut row_targets, + &mut occupied, + &residual, + &residual_targets, + 1, + 4, + &mut residual_to_implicit_rows, + test_span(), + )?; + + assert_eq!(residual_to_implicit_rows, vec![Some(1), Some(3), Some(2)]); + assert_eq!(row_targets[2], Some(solve::scalar_slot_y(0))); + assert_eq!(row_targets[3], Some(solve::scalar_slot_y(3))); + Ok(()) + } + + #[test] + fn unmatched_fallback_residual_is_retained_in_projection_block() -> Result<(), LowerError> { + let span = test_span(); + let derivative = vec![zero_rhs_row()]; + let residual = vec![ + vec![ + solve::LinearOp::LoadY { dst: 0, index: 1 }, + solve::LinearOp::StoreOutput { src: 0 }, + ], + vec![ + solve::LinearOp::LoadY { dst: 0, index: 3 }, + solve::LinearOp::Const { dst: 1, value: 1.0 }, + solve::LinearOp::Binary { + dst: 2, + op: solve::BinaryOp::Sub, + lhs: 0, + rhs: 1, + }, + solve::LinearOp::StoreOutput { src: 2 }, + ], + vec![ + solve::LinearOp::LoadY { dst: 0, index: 3 }, + solve::LinearOp::Const { dst: 1, value: 2.0 }, + solve::LinearOp::Binary { + dst: 2, + op: solve::BinaryOp::Sub, + lhs: 0, + rhs: 1, + }, + solve::LinearOp::StoreOutput { src: 2 }, + ], + ]; + let implicit = build_implicit_rhs_rows( + &derivative, + &residual, + &[Some(solve::scalar_slot_y(1)), None, None], + 1, + 4, + span, + )?; + + let plan = super::super::lower_algebraic_projection_plan( + &implicit.rows, + &implicit.row_targets, + 1, + 4, + span, + )?; + let covered_rows = plan + .blocks + .iter() + .flat_map(|block| block.rows.iter().copied()) + .collect::>(); + + assert_eq!(covered_rows, std::collections::BTreeSet::from([1, 2, 3])); + assert!(plan.blocks.iter().any(|block| { + block.rows.iter().all(|row| [2, 3].contains(row)) + && block.rows.len() == 2 + && block.y_indices == vec![3] + })); + Ok(()) + } + #[test] fn remapped_implicit_rhs_preserves_contiguous_matmul_residual() -> Result<(), LowerError> { let span = test_span(); diff --git a/crates/rumoca-phase-solve/src/initial_values.rs b/crates/rumoca-phase-solve/src/initial_values.rs index 4a2fddc40..2dc2fbe02 100644 --- a/crates/rumoca-phase-solve/src/initial_values.rs +++ b/crates/rumoca-phase-solve/src/initial_values.rs @@ -34,6 +34,20 @@ pub(crate) fn apply_initial_equations_to_start_values( for _ in 0..max_passes { let mut changed = false; for eq in &dae_model.initialization.equations { + changed |= seed_tuple_function_initial_assignment( + dae_model, + layout, + params, + initial_y, + &mut env, + &mut pinned, + eq, + ) + .map_err(|source| SolveModelLowerError::Evaluation { + context: "initial tuple function assignment".to_string(), + source, + span: Some(eq.span), + })?; let Some(assignment) = initial_assignment_from_equation(eq) else { continue; }; @@ -42,6 +56,9 @@ pub(crate) fn apply_initial_equations_to_start_values( } let targets = assignment_target_scalar_names(layout, assignment.target.as_str(), eq.span)?; + if targets.is_empty() { + continue; + } let values = initial_assignment_values( assignment.solution, &env, @@ -92,6 +109,128 @@ pub(crate) fn apply_initial_equations_to_start_values( Ok(()) } +fn seed_tuple_function_initial_assignment( + dae_model: &dae::Dae, + layout: &solve::VarLayout, + params: &mut [f64], + initial_y: &mut [f64], + env: &mut rumoca_eval_dae::VarEnv, + pinned: &mut HashSet, + eq: &dae::Equation, +) -> Result { + let Some((targets_exprs, function_name, args)) = tuple_function_assignment(&eq.rhs) else { + return Ok(false); + }; + let Some((resolved, outputs)) = resolve_function_call_outputs_pub(function_name, env) else { + return Ok(false); + }; + let mut changed = false; + for (idx, target_expr) in targets_exprs.iter().enumerate() { + let Some(output_name) = outputs.get(idx) else { + break; + }; + let Some(target) = tuple_assignment_target_name(target_expr) else { + continue; + }; + let targets = + assignment_target_scalar_names(layout, target.as_str(), eq.span).map_err(|err| { + EvalError::InvalidShape { + context: "tuple initial assignment target", + reason: err.to_string(), + } + })?; + if targets.is_empty() { + continue; + } + let Some(indices) = target_selection_indices(target.as_str(), &targets)? else { + continue; + }; + let mut values = Vec::new(); + reserve_initial_eval_vec_capacity( + &mut values, + indices.len(), + "tuple initial function values", + )?; + for indices in &indices { + values.push(eval_selected_function_output_pub( + &resolved, + output_name, + indices, + args, + env, + )?); + } + let values = initial_assignment_values_with_expected_size( + &rumoca_core::Expression::FunctionCall { + name: function_name.clone(), + args: args.to_vec(), + is_constructor: false, + span: eq.rhs.span().unwrap_or(eq.span), + }, + env, + values, + targets.len().max(1), + )?; + let assignment = InitialAssignment { + target, + solution: &eq.rhs, + is_pre_target: false, + }; + changed |= pin_initial_assignment_targets(&targets, pinned); + changed |= apply_initial_assignment_values( + InitialAssignmentApplyContext { + dae_model, + layout, + params, + initial_y, + env, + }, + &assignment, + &targets, + &values, + ); + } + Ok(changed) +} + +fn tuple_function_assignment( + expr: &rumoca_core::Expression, +) -> Option<( + &[rumoca_core::Expression], + &rumoca_core::Reference, + &[rumoca_core::Expression], +)> { + let rumoca_core::Expression::Binary { + op: OpBinary::Sub, + lhs, + rhs, + .. + } = expr + else { + return None; + }; + let rumoca_core::Expression::Tuple { elements, .. } = lhs.as_ref() else { + return None; + }; + let rumoca_core::Expression::FunctionCall { name, args, .. } = rhs.as_ref() else { + return None; + }; + Some((elements, name, args)) +} + +fn tuple_assignment_target_name(expr: &rumoca_core::Expression) -> Option { + let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = expr + else { + return None; + }; + if !subscripts.is_empty() { + return None; + } + Some(name.as_str().to_string()) +} + fn seed_continuous_assignments( dae_model: &dae::Dae, layout: &solve::VarLayout, @@ -110,6 +249,9 @@ fn seed_continuous_assignments( continue; } let targets = assignment_target_scalar_names(layout, assignment.target.as_str(), eq.span)?; + if targets.is_empty() { + continue; + } if targets .iter() .any(|target| pinned.contains(target) || continuous_seed_pins.contains(target)) @@ -529,6 +671,9 @@ fn checked_initial_target_element_count( ) -> Result { let mut count = 1usize; for dim in shape { + if *dim == 0 { + return Ok(0); + } count = count.checked_mul(*dim).ok_or_else(|| { SolveModelLowerError::Lower(LowerError::ContractViolation { reason: format!("initial assignment target `{target}` element count overflows"), @@ -582,8 +727,7 @@ fn apply_initial_assignment_values( ) -> bool { let target_name = assignment.target.as_str(); let dims = assignment_target_dims(ctx.dae_model, target_name); - if values.len() > 1 && !dims.is_empty() && rumoca_core::parse_scalar_name(target_name).is_none() - { + if !dims.is_empty() && rumoca_core::parse_scalar_name(target_name).is_none() { set_array_entries(ctx.env, target_name, dims, values); } @@ -939,6 +1083,10 @@ mod tests { } } + fn comp_ref(name: &str) -> rumoca_core::ComponentReference { + rumoca_core::ComponentReference::from_flat_segments(name, test_span(), None) + } + fn time_layout() -> solve::VarLayout { solve::VarLayout::from_parts( IndexMap::from([("time".to_string(), solve::ScalarSlot::Time)]), @@ -1001,6 +1149,95 @@ mod tests { ); } + #[test] + fn initial_tuple_function_assignment_seeds_selected_parameter_output() { + let mut dae_model = dae::Dae::default(); + dae_model.variables.parameters.insert( + rumoca_core::VarName::new("a"), + dae::Variable { + fixed: Some(false), + ..dae::Variable::empty_with_span(test_span()) + }, + ); + dae_model.variables.parameters.insert( + rumoca_core::VarName::new("b"), + dae::Variable { + fixed: Some(false), + ..dae::Variable::empty_with_span(test_span()) + }, + ); + + let mut function = rumoca_core::Function::new("Pkg.multi", test_span()); + function.add_output(rumoca_core::FunctionParam::new( + "first", + "Real", + test_span(), + )); + function.add_output(rumoca_core::FunctionParam::new( + "second", + "Real", + test_span(), + )); + function.body = vec![ + rumoca_core::Statement::Assignment { + comp: comp_ref("first"), + value: real(1.25), + span: test_span(), + }, + rumoca_core::Statement::Assignment { + comp: comp_ref("second"), + value: real(2.5), + span: test_span(), + }, + ]; + dae_model + .symbols + .functions + .insert(rumoca_core::VarName::new("Pkg.multi"), function); + dae_model + .initialization + .equations + .push(dae::Equation::residual( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(rumoca_core::Expression::Tuple { + elements: vec![var("a"), var("b")], + span: test_span(), + }), + rhs: Box::new(rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.multi"), + args: Vec::new(), + is_constructor: false, + span: test_span(), + }), + span: test_span(), + }, + test_span(), + "(a, b) = Pkg.multi()", + )); + let layout = solve::VarLayout::from_parts( + IndexMap::from([ + ("a".to_string(), solve::scalar_slot_p(0)), + ("b".to_string(), solve::scalar_slot_p(1)), + ]), + 0, + 2, + ); + let mut params = vec![0.0, 0.0]; + let mut initial_y = Vec::new(); + + apply_initial_equations_to_start_values( + &dae_model, + &layout, + &mut params, + &mut initial_y, + std::sync::Arc::new(EvalRuntimeState::default()), + ) + .expect("tuple function initial assignment should seed params"); + + assert_eq!(params, vec![1.25, 2.5]); + } + #[test] fn fixed_start_pins_include_builtin_time() { assert!( @@ -1064,6 +1301,27 @@ mod tests { ); } + #[test] + fn assignment_target_scalar_names_skips_zero_size_layout_target() { + let layout = solve::VarLayout::from_parts_with_shapes( + IndexMap::new(), + IndexMap::from([("x".to_string(), vec![0])]), + 0, + 0, + ) + .expect("zero-size layout target may omit scalar slots"); + + let names = assignment_target_scalar_names(&layout, "x", test_span()) + .expect("zero-size target should resolve without synthetic scalar names"); + + assert!(names.is_empty()); + assert_eq!( + checked_initial_target_element_count("x", &[0], test_span()) + .expect("zero dimension is a valid empty target"), + 0 + ); + } + #[test] fn assignment_target_scalar_names_rejects_shape_product_overflow_with_span() { let span = rumoca_core::Span::from_offsets( diff --git a/crates/rumoca-phase-solve/src/layout.rs b/crates/rumoca-phase-solve/src/layout.rs index 6209b74ec..42028d3b8 100644 --- a/crates/rumoca-phase-solve/src/layout.rs +++ b/crates/rumoca-phase-solve/src/layout.rs @@ -210,13 +210,22 @@ fn map_constant_bindings( dae_model: &dae::Dae, mut maps: LayoutBindingMaps<'_>, ) -> Result<(), LowerError> { - for (name, var) in &dae_model.variables.constants { - insert_constant_bindings( - maps.reborrow(), - name.as_str(), - var, - &dae_model.symbols.enum_literal_ordinals, - )?; + for _ in 0..dae_model.variables.constants.len().max(1) { + let before = maps.bindings.len(); + for (name, var) in &dae_model.variables.constants { + if maps.bindings.contains_key(name.as_str()) { + continue; + } + insert_constant_bindings( + maps.reborrow(), + name.as_str(), + var, + &dae_model.symbols.enum_literal_ordinals, + )?; + } + if maps.bindings.len() == before { + break; + } } Ok(()) } @@ -309,7 +318,9 @@ fn insert_constant_bindings( let Some(start) = var.start.as_ref() else { return Ok(()); }; - let Some(raw_values) = eval_const_values(start, enum_literal_ordinals) else { + let Some(raw_values) = + eval_const_values_with_bindings(start, enum_literal_ordinals, &*maps.bindings) + else { return Ok(()); }; @@ -609,6 +620,22 @@ fn literal_to_f64(literal: &rumoca_core::Literal) -> Option { } } +fn positive_i64_to_usize(value: i64) -> Option { + if value > 0 { + usize::try_from(value).ok() + } else { + None + } +} + +fn positive_usize_from_f64(value: f64) -> Option { + if value.is_finite() && value >= 1.0 && value.fract().abs() <= f64::EPSILON { + usize::try_from(value as u64).ok() + } else { + None + } +} + fn insert_enum_literal_binding_aliases( bindings: &mut IndexMap, name: &str, @@ -630,20 +657,30 @@ fn insert_enum_literal_binding_key( .or_insert(ScalarSlot::Constant(value)); } -fn eval_const_scalar( +#[cfg(test)] +fn eval_const_values( + expr: &rumoca_core::Expression, + enum_literal_ordinals: &IndexMap, +) -> Option> { + eval_const_values_with_bindings(expr, enum_literal_ordinals, &IndexMap::new()) +} + +fn eval_const_scalar_with_bindings( expr: &rumoca_core::Expression, enum_literal_ordinals: &IndexMap, + bindings: &IndexMap, ) -> Option { - let values = eval_const_values(expr, enum_literal_ordinals)?; + let values = eval_const_values_with_bindings(expr, enum_literal_ordinals, bindings)?; if values.len() == 1 { return values.first().copied(); } None } -fn eval_const_values( +fn eval_const_values_with_bindings( expr: &rumoca_core::Expression, enum_literal_ordinals: &IndexMap, + bindings: &IndexMap, ) -> Option> { match expr { rumoca_core::Expression::Literal { value: literal, .. } => { @@ -653,15 +690,19 @@ fn eval_const_values( // translation-time constants with 1-based ordinal numeric semantics. rumoca_core::Expression::VarRef { name, subscripts, .. - } if subscripts.is_empty() => { - lookup_enum_literal_ordinal(name.as_str(), enum_literal_ordinals) - .map(|ordinal| vec![ordinal as f64]) + } => { + let key = const_var_key(name, subscripts, bindings)?; + if let Some(ordinal) = lookup_enum_literal_ordinal(key.as_str(), enum_literal_ordinals) + { + return Some(vec![ordinal as f64]); + } + constant_slot_value(bindings.get(key.as_str())?).map(|value| vec![value]) } rumoca_core::Expression::BuiltinCall { function, args, .. } => { - eval_const_builtin(*function, args, enum_literal_ordinals) + eval_const_builtin(*function, args, enum_literal_ordinals, bindings) } rumoca_core::Expression::Unary { op, rhs, .. } => { - let values = eval_const_values(rhs, enum_literal_ordinals)?; + let values = eval_const_values_with_bindings(rhs, enum_literal_ordinals, bindings)?; match op { rumoca_core::OpUnary::Plus | rumoca_core::OpUnary::DotPlus => Some(values), rumoca_core::OpUnary::Minus | rumoca_core::OpUnary::DotMinus => { @@ -671,8 +712,8 @@ fn eval_const_values( } } rumoca_core::Expression::Binary { op, lhs, rhs, .. } => { - let lhs = eval_const_scalar(lhs, enum_literal_ordinals)?; - let rhs = eval_const_scalar(rhs, enum_literal_ordinals)?; + let lhs = eval_const_scalar_with_bindings(lhs, enum_literal_ordinals, bindings)?; + let rhs = eval_const_scalar_with_bindings(rhs, enum_literal_ordinals, bindings)?; let value = match op { rumoca_core::OpBinary::Add | rumoca_core::OpBinary::AddElem => lhs + rhs, rumoca_core::OpBinary::Sub | rumoca_core::OpBinary::SubElem => lhs - rhs, @@ -687,17 +728,21 @@ fn eval_const_values( | rumoca_core::Expression::Tuple { elements, .. } => { let mut values = Vec::new(); for element in elements { - values.extend(eval_const_values(element, enum_literal_ordinals)?); + values.extend(eval_const_values_with_bindings( + element, + enum_literal_ordinals, + bindings, + )?); } Some(values) } rumoca_core::Expression::Range { start, step, end, .. } => { - let start = eval_const_scalar(start, enum_literal_ordinals)?; - let end = eval_const_scalar(end, enum_literal_ordinals)?; + let start = eval_const_scalar_with_bindings(start, enum_literal_ordinals, bindings)?; + let end = eval_const_scalar_with_bindings(end, enum_literal_ordinals, bindings)?; let step = if let Some(step_expr) = step { - eval_const_scalar(step_expr, enum_literal_ordinals)? + eval_const_scalar_with_bindings(step_expr, enum_literal_ordinals, bindings)? } else if end >= start { 1.0 } else { @@ -736,21 +781,23 @@ fn eval_const_builtin( function: rumoca_core::BuiltinFunction, args: &[rumoca_core::Expression], enum_literal_ordinals: &IndexMap, + bindings: &IndexMap, ) -> Option> { use rumoca_core::BuiltinFunction as Builtin; let unary = |f: fn(f64) -> f64| { - let value = eval_const_scalar(args.first()?, enum_literal_ordinals)?; + let value = + eval_const_scalar_with_bindings(args.first()?, enum_literal_ordinals, bindings)?; Some(vec![f(value)]) }; let binary = |f: fn(f64, f64) -> f64| { - let lhs = eval_const_scalar(args.first()?, enum_literal_ordinals)?; - let rhs = eval_const_scalar(args.get(1)?, enum_literal_ordinals)?; + let lhs = eval_const_scalar_with_bindings(args.first()?, enum_literal_ordinals, bindings)?; + let rhs = eval_const_scalar_with_bindings(args.get(1)?, enum_literal_ordinals, bindings)?; Some(vec![f(lhs, rhs)]) }; let binary_builtin = |function| { - let lhs = eval_const_scalar(args.first()?, enum_literal_ordinals)?; - let rhs = eval_const_scalar(args.get(1)?, enum_literal_ordinals)?; + let lhs = eval_const_scalar_with_bindings(args.first()?, enum_literal_ordinals, bindings)?; + let rhs = eval_const_scalar_with_bindings(args.get(1)?, enum_literal_ordinals, bindings)?; rumoca_core::apply_scalar_binary_math(function, lhs, rhs).map(|value| vec![value]) }; @@ -779,13 +826,71 @@ fn eval_const_builtin( Builtin::Div => binary_builtin(Builtin::Div), Builtin::Mod => binary_builtin(Builtin::Mod), Builtin::Rem => binary_builtin(Builtin::Rem), - Builtin::NoEvent => eval_const_values(args.first()?, enum_literal_ordinals), - Builtin::Smooth => eval_const_values(args.get(1)?, enum_literal_ordinals), - Builtin::Homotopy => eval_const_values(args.first()?, enum_literal_ordinals), + Builtin::NoEvent => { + eval_const_values_with_bindings(args.first()?, enum_literal_ordinals, bindings) + } + Builtin::Smooth => { + eval_const_values_with_bindings(args.get(1)?, enum_literal_ordinals, bindings) + } + Builtin::Homotopy => { + eval_const_values_with_bindings(args.first()?, enum_literal_ordinals, bindings) + } _ => None, } } +fn constant_slot_value(slot: &ScalarSlot) -> Option { + match slot { + ScalarSlot::Constant(value) => Some(*value), + _ => None, + } +} + +fn const_var_key( + name: &rumoca_core::Reference, + subscripts: &[rumoca_core::Subscript], + bindings: &IndexMap, +) -> Option { + let name = constant_reference_lookup_name(name, bindings); + if subscripts.is_empty() { + return Some(name); + } + let mut indices = Vec::with_capacity(subscripts.len()); + for subscript in subscripts { + let index = match subscript { + rumoca_core::Subscript::Index { value, .. } => positive_i64_to_usize(*value)?, + rumoca_core::Subscript::Expr { expr, .. } => positive_usize_from_f64( + eval_const_scalar_with_bindings(expr, &IndexMap::new(), bindings)?, + )?, + rumoca_core::Subscript::Colon { .. } => return None, + }; + indices.push(index); + } + Some(dae::format_subscript_key(name.as_str(), &indices)) +} + +fn constant_reference_lookup_name( + name: &rumoca_core::Reference, + bindings: &IndexMap, +) -> String { + if let Some(component_ref) = name.component_ref() { + let structured = component_ref_flat_name(component_ref); + if bindings.contains_key(structured.as_str()) || structured != name.as_str() { + return structured; + } + } + name.as_str().to_string() +} + +fn component_ref_flat_name(component_ref: &rumoca_core::ComponentReference) -> String { + component_ref + .parts + .iter() + .map(|part| part.ident.as_str()) + .collect::>() + .join(".") +} + fn lookup_enum_literal_ordinal(raw: &str, ordinals: &IndexMap) -> Option { if let Some(&ordinal) = ordinals.get(raw) { return Some(ordinal); @@ -873,6 +978,28 @@ mod tests { } } + fn var_ref(name: &str) -> rumoca_core::Expression { + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::generated(name), + subscripts: Vec::new(), + span: test_span(), + } + } + + fn constant_variable( + name: &str, + start: rumoca_core::Expression, + ) -> (rumoca_core::VarName, dae::Variable) { + ( + rumoca_core::VarName::new(name), + dae::Variable { + name: rumoca_core::VarName::new(name), + start: Some(start), + ..rumoca_ir_dae::Variable::empty_with_span(test_span()) + }, + ) + } + #[test] fn const_builtin_div_mod_and_rem_follow_modelica_rounding() { let ordinals = IndexMap::new(); @@ -890,6 +1017,27 @@ mod tests { ); } + #[test] + fn build_var_layout_resolves_constant_alias_chains() { + let mut dae_model = dae::Dae::default(); + let (alias_name, alias_var) = + constant_variable("sineVoltage.pi", var_ref("Modelica.Constants.pi")); + let (source_name, source_var) = + constant_variable("Modelica.Constants.pi", real(std::f64::consts::PI)); + dae_model.variables.constants.insert(alias_name, alias_var); + dae_model + .variables + .constants + .insert(source_name, source_var); + + let layout = build_var_layout(&dae_model).expect("constant alias layout should build"); + + assert_eq!( + layout.binding("sineVoltage.pi"), + Some(ScalarSlot::Constant(std::f64::consts::PI)) + ); + } + #[test] fn build_var_layout_indexes_arrays_by_component_reference() { let span = test_span(); diff --git a/crates/rumoca-phase-solve/src/lib.rs b/crates/rumoca-phase-solve/src/lib.rs index c3e167139..8017d5b5b 100644 --- a/crates/rumoca-phase-solve/src/lib.rs +++ b/crates/rumoca-phase-solve/src/lib.rs @@ -1,3 +1,5 @@ +#![feature(if_let_guard)] + //! Lower DAE data into solver-facing IR. //! //! Lowering passes (`layout`, `lower`, `ad`) take a `dae::Dae` and produce @@ -8,9 +10,8 @@ //! The DAE tree-walk interpreter (`eval`, `dual`, `sim_float`, `statement`) lives //! in `rumoca-eval-dae`. //! -//! SPEC_0021 file-size exception: phase-solve facade still owns exports, -//! compatibility tests, and lowering integration. split plan: keep moving -//! tests and integration-only helpers into focused modules. +//! The facade owns exports and lowering integration; GPU initialization and +//! projection planning live in focused modules to keep phase boundaries legible. use std::collections::{BTreeMap, BTreeSet, HashSet}; @@ -33,12 +34,14 @@ mod discrete_pre_modes; mod dummy_derivative; mod dynamic_events; mod event_actions; +mod gpu_initialization; mod implicit_rhs; mod initial_values; pub mod layout; pub mod lower; mod observation_refresh; mod path_utils; +mod projection_plan; mod projection_suffix; mod residual_compute_block; mod runtime_assignments; @@ -64,11 +67,17 @@ use continuous_row_targets::{ continuous_equation_scalar_name, scalarized_record_target_names, target_expr_scalar_name, }; use continuous_row_targets::{ + dedupe_continuous_y_targets, lower_contiguous_y_target_range_for_equation, lower_continuous_row_targets, lower_continuous_row_targets_for_equation, }; use discrete_pre_modes::discrete_pre_mode_for_equation; #[cfg(test)] pub(crate) use discrete_pre_modes::expression_contains_event_entry_pre_operator; +use gpu_initialization::lower_gpu_initialization_system; +#[cfg(test)] +use gpu_initialization::{ + append_gpu_corner_strides, gpu_corner_cell_index, reject_nondeterministic_gpu_initial_ops, +}; #[cfg(test)] use implicit_rhs::zero_rhs_row; use implicit_rhs::{ @@ -78,14 +87,16 @@ use layout::INITIAL_EVENT_PARAMETER_NAME; pub use layout::{build_var_layout, build_var_layout_with_solver_len}; pub use lower::LowerError; use lower::{ - lower_discrete_rhs_from_equations, lower_initial_residual, lower_initial_update_rhs, - lower_residual_rows_and_targets_from_equations, lower_root_conditions, + lower_discrete_rhs_from_equations, lower_initial_residual, lower_initial_residual_cell, + lower_initial_update_rhs, lower_residual_rows_and_targets_from_equations, + lower_root_conditions, }; use lower::{ lower_dynamic_time_event_rhs, lower_runtime_assignment_rhs, normalized_discrete_update_equations, }; use observation_refresh::lower_discrete_observation_refresh; +use projection_plan::*; use runtime_assignments::{ lower_runtime_assignment_targets, runtime_assignment_equation, runtime_assignment_equations, runtime_tail_update_names, static_runtime_tail_equation, @@ -228,7 +239,7 @@ impl SolveProblemLoweringProfile { } fn lower_initialization_system(self) -> bool { - self == Self::Runtime + matches!(self, Self::Runtime | Self::GpuPreparation) } fn lower_initialization_updates(self) -> bool { @@ -333,7 +344,7 @@ pub(crate) fn lower_solve_problem_with_solver_len_and_model_span_and_profile( // `solver_residual_equations` has already removed state-derivative rows. // The remaining original DAE indices are not a state-row prefix, so residual // lowering must not infer derivative-row behavior from `row_idx < n_x`. - let (residual, residual_targets) = lower_residual_rows_and_targets_from_equations( + let (residual, mut residual_targets) = lower_residual_rows_and_targets_from_equations( dae_model, &layout, residual_equations.iter().copied(), @@ -343,6 +354,7 @@ pub(crate) fn lower_solve_problem_with_solver_len_and_model_span_and_profile( }, ) .map_err(|err| lower_problem_context(err, "lower continuous residual rows and targets"))?; + dedupe_continuous_y_targets(&mut residual_targets); timing::log_stage("problem.lower_residual_rows", timer); // Derivative lowering must LOAD retained algebraic unknowns from their projected // slot rather than inline their definitions (roadmap 4b): inlining a boundary cell @@ -420,7 +432,7 @@ pub(crate) fn lower_solve_problem_with_solver_len_and_model_span_and_profile( timing::log_stage("problem.lower_runtime_systems", timer); let timer = timing::stage_start(); let initialization = if profile.lower_initialization_system() { - lower_initialization_system(dae_model, &layout, &solve_layout)? + lower_initialization_system(dae_model, &layout, &solve_layout, profile)? } else if profile.lower_initialization_updates() { lower_initialization_updates_only(dae_model, &layout)? } else { @@ -580,7 +592,12 @@ fn lower_event_partition_for_profile( .map_err(|err| lower_problem_context(err, "lower root relation memory targets"))?, scheduled_root_conditions: lower::lower_scheduled_root_conditions(dae_model) .map_err(|err| lower_problem_context(err, "lower scheduled root conditions"))?, - scheduled_time_events: dae_model.events.scheduled_time_events.clone(), + scheduled_time_events: dae_model + .events + .scheduled_time_events + .iter() + .map(|event| event.time) + .collect(), dynamic_time_event_names: dynamic_events::collect_dynamic_time_event_names(dae_model), dynamic_time_event_rhs: solve::ScalarProgramBlock::with_program_spans( lower_dynamic_time_event_rhs(dae_model, layout, dynamic_time_event_exprs) @@ -696,7 +713,11 @@ fn lower_initialization_system( dae_model: &dae::Dae, layout: &solve::VarLayout, solve_layout: &solve::SolveLayout, + profile: SolveProblemLoweringProfile, ) -> Result { + if profile == SolveProblemLoweringProfile::GpuPreparation { + return lower_gpu_initialization_system(dae_model, layout); + } let residual_equations = lower::initial_residual_equations(dae_model, layout) .map_err(|err| lower_problem_context(err, "collect initial residual equations"))?; let row_targets = @@ -706,15 +727,27 @@ fn lower_initialization_system( .map_err(|err| lower_problem_context(err, "collect initial condition updates"))?; let update_targets = lower_update_targets_from_equations(dae_model, layout, &update_equations) .map_err(|err| lower_problem_context(err, "lower initial update targets"))?; - let residual_rows = lower_initial_residual(dae_model, layout) .map_err(|err| lower_problem_context(err, "lower initial residual rows"))?; let projection_indices = initial_projection_indices_for_layout(dae_model, solve_layout)?; + let continuous_equation_count = dae_model.continuous.equations.len(); + let implicit_initial_projection_rows = residual_equations + .iter() + .enumerate() + .filter_map(|(row_idx, (equation_idx, _))| { + (*equation_idx >= continuous_equation_count).then_some(row_idx) + }) + .collect::>(); let projection_plan = lower_projection_plan( &residual_rows, &row_targets, &projection_indices, 0..residual_rows.len(), + ProjectionPlanPolicy { + include_explicit_row_targets: false, + require_complete_algebraic_coverage: false, + }, + Some(&implicit_initial_projection_rows), dae_model_span(dae_model)?, )?; @@ -723,23 +756,32 @@ fn lower_initialization_system( // `sig[i,j]`) collapse into a few `Map`/`AffineStencil` tensor nodes instead of // one scalar program per cell. This is the dominant initialization cost on PDE // grids (it was ~80% of the whole Solve-IR before this change). - let residual = residual_compute_block::build_residual_compute_block( + let residual = residual_compute_block::build_initialization_residual_compute_block( dae_model, layout, &residual_rows, &row_targets, &residual_equations, )?; + let initialization_span = dae_model_span(dae_model)?; + let residual_output_count = residual + .len() + .map_err(|err| lower_contract_violation(err.to_string(), initialization_span))?; + let _ = residual_output_count; + let update_rhs = solve::ScalarProgramBlock::with_program_spans( + lower_initial_update_rhs(dae_model, layout) + .map_err(|err| lower_problem_context(err, "lower initial update rows"))?, + program_spans_for_owned_equations(&update_equations)?, + )?; Ok(solve::InitializationSolveSystem { row_targets, + direct_families: Vec::new(), + required_target_ranges: Vec::new(), + fixed_target_ranges: Vec::new(), projection_indices, projection_plan, residual, - update_rhs: solve::ScalarProgramBlock::with_program_spans( - lower_initial_update_rhs(dae_model, layout) - .map_err(|err| lower_problem_context(err, "lower initial update rows"))?, - program_spans_for_owned_equations(&update_equations)?, - )?, + update_rhs, update_targets, }) } @@ -976,544 +1018,6 @@ fn solver_residual_equation( && runtime_assignment_equation(dae_model, runtime_tail_updates, eq)?.is_none()) } -fn lower_algebraic_projection_plan( - rows: &[Vec], - row_targets: &[Option], - state_scalar_count: usize, - solver_scalar_count: usize, - context_span: rumoca_core::Span, -) -> Result { - let projection_count = solver_scalar_count - .checked_sub(state_scalar_count) - .ok_or_else(|| { - lower_contract_violation( - "algebraic projection range starts after solver scalar count".to_string(), - context_span, - ) - })?; - let mut projection_indices = lower_vec_with_capacity( - projection_count, - "algebraic projection index count", - context_span, - )?; - projection_indices.extend(state_scalar_count..solver_scalar_count); - lower_projection_plan( - rows, - row_targets, - &projection_indices, - state_scalar_count..solver_scalar_count, - context_span, - ) -} - -fn lower_projection_plan( - rows: &[Vec], - row_targets: &[Option], - projection_indices: &[usize], - row_indices: std::ops::Range, - context_span: rumoca_core::Span, -) -> Result { - let mut row_to_vars = BTreeMap::>::new(); - let projection_set = projection_indices.iter().copied().collect::>(); - - for row_idx in row_indices { - let mut y_indices = - collect_algebraic_y_indices_for_row(rows[row_idx].as_slice(), &projection_set); - if y_indices.is_empty() - && let Some(solve::ScalarSlot::Y { index, .. }) = - row_targets.get(row_idx).copied().flatten() - && projection_set.contains(&index) - { - y_indices.insert(index); - } - if y_indices.is_empty() { - continue; - } - row_to_vars.insert(row_idx, y_indices); - } - - let projection_incidence = algebraic_projection_incidence(&row_to_vars, context_span)?; - let blocks = projection_blt_blocks(&projection_incidence)?; - Ok(solve::AlgebraicProjectionPlan { - blocks: lower_blt_projection_blocks( - &blocks, - row_targets, - &projection_incidence, - context_span, - )?, - }) -} - -fn projection_blt_blocks( - projection_incidence: &ProjectionIncidence, -) -> Result, LowerError> { - if projection_incidence.incidence.n_eq == 0 && projection_incidence.incidence.n_var == 0 { - return Ok(Vec::new()); - } - let regular = - rumoca_phase_structural::maximum_regular_subsystem(&projection_incidence.incidence) - .map_err(|err| LowerError::Unsupported { - reason: format!("lower algebraic projection BLT: {err}"), - })?; - rumoca_phase_structural::build_blt_from_incidence(®ular.incidence).map_err(|err| { - LowerError::Unsupported { - reason: format!("lower algebraic projection BLT: {err}"), - } - }) -} - -fn collect_algebraic_y_indices_for_row( - row: &[solve::LinearOp], - projection_set: &BTreeSet, -) -> BTreeSet { - let mut defs = BTreeMap::::new(); - let mut outputs = Vec::new(); - for op in row { - match row_def_use(op) { - RowDefUseOp::Def { dst, def_use } => { - defs.insert(dst, def_use); - } - RowDefUseOp::Store { src } => outputs.push(src), - } - } - let mut y_indices = BTreeSet::new(); - let mut visited = BTreeSet::new(); - let mut stack = outputs; - while let Some(reg) = stack.pop() { - if !visited.insert(reg) { - continue; - } - let Some(def_use) = defs.get(®) else { - continue; - }; - if let Some(index) = def_use.loaded_y - && projection_set.contains(&index) - { - y_indices.insert(index); - } - stack.extend(def_use.inputs.iter().copied()); - } - y_indices -} - -#[derive(Debug)] -struct RowDefUse { - loaded_y: Option, - inputs: Vec, -} - -enum RowDefUseOp { - Def { dst: solve::Reg, def_use: RowDefUse }, - Store { src: solve::Reg }, -} - -fn row_def_use(op: &solve::LinearOp) -> RowDefUseOp { - use solve::LinearOp as Op; - match *op { - Op::Const { dst, .. } | Op::LoadTime { dst } | Op::LoadP { dst, .. } => { - def_use(dst, None, Vec::new()) - } - Op::LoadY { dst, index } => def_use(dst, Some(index), Vec::new()), - Op::LoadSeed { dst, .. } => def_use(dst, None, Vec::new()), - Op::LoadIndexedP { dst, index, .. } | Op::LoadIndexedSeed { dst, index, .. } => { - def_use(dst, None, vec![index]) - } - Op::Move { dst, src } | Op::Unary { dst, arg: src, .. } => def_use(dst, None, vec![src]), - Op::Binary { dst, lhs, rhs, .. } | Op::Compare { dst, lhs, rhs, .. } => { - def_use(dst, None, vec![lhs, rhs]) - } - Op::Select { - dst, - cond, - if_true, - if_false, - } => def_use(dst, None, vec![cond, if_true, if_false]), - Op::LinearSolveComponent { - dst, - matrix_start, - rhs_start, - n, - .. - } => def_use( - dst, - None, - reg_range(matrix_start, n * n) - .chain(reg_range(rhs_start, n)) - .collect(), - ), - Op::TableBounds { dst, table_id, .. } => def_use(dst, None, vec![table_id]), - Op::TableLookup { - dst, - table_id, - column, - input, - } - | Op::TableLookupSlope { - dst, - table_id, - column, - input, - } => def_use(dst, None, vec![table_id, column, input]), - Op::TableNextEvent { - dst, - table_id, - time, - } => def_use(dst, None, vec![table_id, time]), - Op::RandomInitialState { - dst, - local_seed, - global_seed, - .. - } => def_use(dst, None, vec![local_seed, global_seed]), - Op::RandomResult { - dst, - state_start, - state_len, - .. - } - | Op::RandomState { - dst, - state_start, - state_len, - .. - } => def_use(dst, None, reg_range(state_start, state_len).collect()), - Op::ImpureRandomInit { dst, seed } => def_use(dst, None, vec![seed]), - Op::ImpureRandom { dst, id, .. } => def_use(dst, None, vec![id]), - Op::ImpureRandomInteger { - dst, - id, - imin, - imax, - .. - } => def_use(dst, None, vec![id, imin, imax]), - Op::StoreOutput { src } => RowDefUseOp::Store { src }, - } -} - -fn def_use(dst: solve::Reg, loaded_y: Option, inputs: Vec) -> RowDefUseOp { - RowDefUseOp::Def { - dst, - def_use: RowDefUse { loaded_y, inputs }, - } -} - -fn reg_range(start: solve::Reg, len: usize) -> impl Iterator { - (0..len).filter_map(move |offset| start.checked_add(offset.try_into().ok()?)) -} - -struct ProjectionIncidence { - incidence: Incidence, - unknown_y_indices: Vec, -} - -fn algebraic_projection_incidence( - row_to_vars: &BTreeMap>, - context_span: rumoca_core::Span, -) -> Result { - let unknown_y_set = row_to_vars - .values() - .flat_map(|vars| vars.iter().copied()) - .collect::>(); - let mut unknown_y_indices = lower_vec_with_capacity( - unknown_y_set.len(), - "projection unknown index count", - context_span, - )?; - unknown_y_indices.extend(unknown_y_set); - - let mut unknown_names = lower_vec_with_capacity( - unknown_y_indices.len(), - "projection unknown name count", - context_span, - )?; - for y_idx in &unknown_y_indices { - unknown_names.push(projection_unknown_id(*y_idx)); - } - - let unknown_positions = unknown_y_indices - .iter() - .copied() - .enumerate() - .map(|(local_idx, y_idx)| (y_idx, local_idx)) - .collect::>(); - - let mut equation_refs = lower_vec_with_capacity( - row_to_vars.len(), - "projection equation ref count", - context_span, - )?; - let mut eq_unknowns = lower_vec_with_capacity( - row_to_vars.len(), - "projection equation unknown count", - context_span, - )?; - for (row_idx, vars) in row_to_vars { - equation_refs.push(EquationRef(*row_idx)); - let mut unknowns = - lower_hash_set_with_capacity(vars.len(), "projection row unknown count", context_span)?; - for y_idx in vars { - if let Some(local_idx) = unknown_positions.get(y_idx).copied() { - unknowns.insert(local_idx); - } - } - eq_unknowns.push(unknowns); - } - - Ok(ProjectionIncidence { - incidence: Incidence::new(eq_unknowns, equation_refs, unknown_names), - unknown_y_indices, - }) -} - -fn projection_unknown_id(y_idx: usize) -> UnknownId { - UnknownId::SolverY(y_idx) -} - -fn projection_y_index( - unknown: &UnknownId, - projection_incidence: &ProjectionIncidence, -) -> Option { - projection_incidence - .incidence - .unknown_names - .iter() - .position(|candidate| candidate == unknown) - .and_then(|idx| projection_incidence.unknown_y_indices.get(idx).copied()) -} - -fn lower_blt_projection_blocks( - blocks: &[BltBlock], - row_targets: &[Option], - projection_incidence: &ProjectionIncidence, - context_span: rumoca_core::Span, -) -> Result, LowerError> { - let mut lowered = lower_vec_with_capacity( - blocks.len(), - "algebraic projection block count", - context_span, - )?; - for block in blocks { - let block = match block { - BltBlock::Scalar { equation, unknown } => { - projection_y_index(unknown, projection_incidence) - .map(|y_index| scalar_projection_block(equation.0, y_index, context_span)) - .transpose()? - } - BltBlock::AlgebraicLoop { - equations, - unknowns, - } => lower_algebraic_loop_projection_block( - equations, - unknowns, - row_targets, - projection_incidence, - context_span, - )?, - }; - if let Some(block) = block { - lowered.push(block); - } - } - merge_overlapping_projection_blocks(lowered, context_span) -} - -fn merge_overlapping_projection_blocks( - blocks: Vec, - context_span: rumoca_core::Span, -) -> Result, LowerError> { - let mut merged = lower_vec_with_capacity( - blocks.len(), - "merged algebraic projection block count", - context_span, - )?; - for block in blocks { - merge_projection_block(&mut merged, block, context_span)?; - } - Ok(merged) -} - -fn merge_projection_block( - merged: &mut Vec, - mut block: solve::AlgebraicProjectionBlock, - context_span: rumoca_core::Span, -) -> Result<(), LowerError> { - let mut idx = 0; - while idx < merged.len() { - if projection_blocks_overlap(&merged[idx], &block) { - let previous = merged.remove(idx); - block = combine_projection_blocks(previous, block, context_span)?; - idx = 0; - } else { - idx += 1; - } - } - merged.push(block); - Ok(()) -} - -fn projection_blocks_overlap( - lhs: &solve::AlgebraicProjectionBlock, - rhs: &solve::AlgebraicProjectionBlock, -) -> bool { - lhs.y_indices - .iter() - .any(|index| rhs.y_indices.binary_search(index).is_ok()) -} - -fn combine_projection_blocks( - lhs: solve::AlgebraicProjectionBlock, - rhs: solve::AlgebraicProjectionBlock, - context_span: rumoca_core::Span, -) -> Result { - let causal_step_count = lhs - .causal_steps - .len() - .checked_add(rhs.causal_steps.len()) - .ok_or_else(|| { - lower_contract_violation( - "merged algebraic projection causal-step count overflows host index range" - .to_string(), - context_span, - ) - })?; - let mut causal_steps = lower_vec_with_capacity( - causal_step_count, - "merged algebraic projection causal-step count", - context_span, - )?; - causal_steps.extend(lhs.causal_steps); - causal_steps.extend(rhs.causal_steps); - Ok(solve::AlgebraicProjectionBlock { - rows: merge_unique( - lhs.rows, - rhs.rows, - "merged algebraic projection row count", - context_span, - )?, - y_indices: merge_unique( - lhs.y_indices, - rhs.y_indices, - "merged algebraic projection target count", - context_span, - )?, - causal_steps, - }) -} - -fn merge_unique( - lhs: Vec, - rhs: Vec, - context: &'static str, - context_span: rumoca_core::Span, -) -> Result, LowerError> { - let capacity = lhs.len().checked_add(rhs.len()).ok_or_else(|| { - lower_contract_violation( - format!("{context} overflows host index range"), - context_span, - ) - })?; - let mut merged = lower_vec_with_capacity(capacity, context, context_span)?; - merged.extend(lhs); - merged.extend(rhs); - merged.sort_unstable(); - merged.dedup(); - Ok(merged) -} - -fn scalar_projection_block( - row: usize, - y_index: usize, - context_span: rumoca_core::Span, -) -> Result { - let mut rows = lower_vec_with_capacity( - 1, - "scalar algebraic projection block row count", - context_span, - )?; - rows.push(row); - let mut y_indices = lower_vec_with_capacity( - 1, - "scalar algebraic projection block target count", - context_span, - )?; - y_indices.push(y_index); - Ok(solve::AlgebraicProjectionBlock { - rows, - y_indices, - causal_steps: Vec::new(), - }) -} - -fn sorted_set_values( - values: BTreeSet, - context: &'static str, - context_span: rumoca_core::Span, -) -> Result, LowerError> { - let mut out = lower_vec_with_capacity(values.len(), context, context_span)?; - out.extend(values); - Ok(out) -} - -fn collect_equation_rows( - equations: &[EquationRef], - context_span: rumoca_core::Span, -) -> Result, LowerError> { - let mut rows = lower_vec_with_capacity( - equations.len(), - "algebraic loop projection row count", - context_span, - )?; - for equation in equations { - rows.push(equation.0); - } - Ok(rows) -} - -fn lower_algebraic_loop_projection_block( - equations: &[EquationRef], - unknowns: &[UnknownId], - row_targets: &[Option], - projection_incidence: &ProjectionIncidence, - context_span: rumoca_core::Span, -) -> Result, LowerError> { - let rows = collect_equation_rows(equations, context_span)?; - let y_indices = sorted_set_values( - loop_projection_target_set(unknowns, row_targets, &rows, projection_incidence), - "algebraic loop projection target count", - context_span, - )?; - if rows.is_empty() || y_indices.is_empty() { - return Ok(None); - } - Ok(Some(solve::AlgebraicProjectionBlock { - rows, - y_indices, - causal_steps: Vec::new(), - })) -} - -fn loop_projection_target_set( - unknowns: &[UnknownId], - row_targets: &[Option], - rows: &[usize], - projection_incidence: &ProjectionIncidence, -) -> BTreeSet { - let mut y_indices = BTreeSet::new(); - for unknown in unknowns { - if let Some(index) = projection_y_index(unknown, projection_incidence) { - y_indices.insert(index); - } - } - for row in rows { - if let Some(solve::ScalarSlot::Y { index, .. }) = row_targets.get(*row).copied().flatten() - && projection_incidence.unknown_y_indices.contains(&index) - { - y_indices.insert(index); - } - } - y_indices -} - pub fn solver_vector_names( dae_model: &dae::Dae, n_total: usize, diff --git a/crates/rumoca-phase-solve/src/lower.rs b/crates/rumoca-phase-solve/src/lower.rs index c0e9c32ef..f04e7735f 100644 --- a/crates/rumoca-phase-solve/src/lower.rs +++ b/crates/rumoca-phase-solve/src/lower.rs @@ -61,6 +61,7 @@ pub use expression_rows::{ use function_projection::format_subscript_binding_key; use helpers::*; pub use initial_residual::{initial_residual_equations, lower_initial_residual}; +pub(crate) use initial_residual::{lower_initial_residual_cell, lower_initial_residual_cells}; use misc_helpers::*; use scope::*; use source_refs::*; @@ -148,6 +149,13 @@ pub(crate) fn scalarized_record_field_binding_names( expression_rows::scalarized_record_field_binding_names(base, layout) } +pub(crate) fn residual_equation_effective_row_count( + dae_model: &dae::Dae, + eq: &dae::Equation, +) -> Result { + expression_rows::residual_equation_effective_row_count(dae_model, eq) +} + pub(crate) fn compile_time_subscript_indices_for_structured_access( subscripts: &[rumoca_core::Subscript], structural_bindings: &IndexMap, @@ -174,6 +182,12 @@ pub(crate) fn structural_bindings_for_structured_access( compile_time::structural_bindings(dae_model) } +pub(crate) fn external_table_data_for_dae( + dae_model: &dae::Dae, +) -> Result, LowerError> { + compile_time::external_table_data(dae_model) +} + pub fn lower_derivative_rhs( dae_model: &dae::Dae, layout: &VarLayout, @@ -421,6 +435,15 @@ pub(super) struct LowerBuilderMetadata<'a> { pub(super) is_initial_mode: bool, } +pub(in crate::lower) struct RuntimeLowerBuilderMetadata<'a> { + pub(in crate::lower) clock_intervals: &'a IndexMap, + pub(in crate::lower) clock_timings: &'a IndexMap, + pub(in crate::lower) triggered_clock_conditions: &'a [rumoca_core::Expression], + pub(in crate::lower) variable_starts: &'a IndexMap, + pub(in crate::lower) dae_variables: &'a dae::DaeVariables, + pub(in crate::lower) is_initial_mode: bool, +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum ValueMode { Current, @@ -438,24 +461,20 @@ impl<'a> LowerBuilder<'a> { fn new_with_runtime_metadata( layout: &'a VarLayout, functions: &'a IndexMap, - clock_intervals: &'a IndexMap, - clock_timings: &'a IndexMap, - triggered_clock_conditions: &'a [rumoca_core::Expression], - variable_starts: &'a IndexMap, - is_initial_mode: bool, + runtime: RuntimeLowerBuilderMetadata<'a>, ) -> Self { Self::new_with_metadata( layout, functions, LowerBuilderMetadata { - clock_intervals: Some(clock_intervals), - clock_timings: Some(clock_timings), - triggered_clock_conditions: Some(triggered_clock_conditions), + clock_intervals: Some(runtime.clock_intervals), + clock_timings: Some(runtime.clock_timings), + triggered_clock_conditions: Some(runtime.triggered_clock_conditions), discrete_valued_names: None, - variable_starts: Some(variable_starts), - dae_variables: None, + variable_starts: Some(runtime.variable_starts), + dae_variables: Some(runtime.dae_variables), indexed_bindings: None, - is_initial_mode, + is_initial_mode: runtime.is_initial_mode, }, ) } @@ -660,9 +679,17 @@ impl<'a> LowerBuilder<'a> { { return Ok(ComponentReferenceKey::generated(name.as_str())); } - if name.is_generated() || name.component_ref().is_some() || self.dae_variables.is_none() { + if self.scalarized_field_binding_available(name.as_str(), "re") + || self.scalarized_field_binding_available(name.as_str(), "im") + { + return Ok(ComponentReferenceKey::generated(name.as_str())); + } + if name.is_generated() || self.dae_variables.is_none() { return scope_key_from_reference(name, span); } + if let Some(component_ref) = name.component_ref() { + return self.scope_key_from_component_ref(name, component_ref); + } let Some(variable) = self .dae_variables .and_then(|variables| dae_variable(variables, name.var_name())) @@ -710,6 +737,78 @@ impl<'a> LowerBuilder<'a> { } } + fn scope_key_from_component_ref( + &self, + name: &rumoca_core::Reference, + component_ref: &rumoca_core::ComponentReference, + ) -> Result { + match ComponentReferenceKey::from_component_reference(component_ref) { + Ok(key) => Ok(key), + Err(err) + if err.kind == rumoca_ir_solve::ComponentReferenceKeyErrorKind::MissingDefId => + { + self.scope_key_from_component_ref_missing_def_id(name, err) + } + Err(err) => Err(component_reference_lower_error(name, err)), + } + } + + fn scope_key_from_component_ref_missing_def_id( + &self, + name: &rumoca_core::Reference, + err: rumoca_ir_solve::ComponentReferenceKeyError, + ) -> Result { + if let Some(component_ref) = self.dae_variable_component_ref_with_def_id(name) { + return component_reference_key_or_error(name, component_ref); + } + if let Some(component_ref) = self.scalarized_base_component_ref(name, err.span)? { + return component_reference_key_or_error(name, &component_ref); + } + Err(component_reference_lower_error(name, err)) + } + + fn dae_variable_component_ref_with_def_id( + &self, + name: &rumoca_core::Reference, + ) -> Option<&rumoca_core::ComponentReference> { + self.dae_variables + .and_then(|variables| dae_variable(variables, name.var_name())) + .and_then(|variable| variable.component_ref.as_ref()) + .filter(|component_ref| component_ref.def_id.is_some()) + } + + fn scalarized_base_component_ref( + &self, + name: &rumoca_core::Reference, + span: rumoca_core::Span, + ) -> Result, LowerError> { + let Some(scalar) = rumoca_core::parse_scalar_name(name.as_str()) else { + return Ok(None); + }; + let Some(component_ref) = self + .dae_variables + .and_then(|variables| dae_variable(variables, &VarName::new(scalar.base))) + .and_then(|variable| variable.component_ref.as_ref()) + .filter(|component_ref| component_ref.def_id.is_some()) + else { + return Ok(None); + }; + let mut component_ref = component_ref.clone(); + let Some(last) = component_ref.parts.last_mut() else { + return Ok(None); + }; + for index in scalar.indices { + let subscript = rumoca_core::Subscript::try_generated_index( + index, + span, + "scalarized source reference", + ) + .map_err(|err| LowerError::contract_violation(err.to_string(), span))?; + last.subs.push(subscript); + } + Ok(Some(component_ref)) + } + fn lower_expr( &mut self, expr: &rumoca_core::Expression, @@ -825,6 +924,7 @@ impl<'a> LowerBuilder<'a> { } } + #[allow(clippy::too_many_lines)] fn lower_var_ref( &mut self, name: &rumoca_core::Reference, @@ -845,10 +945,39 @@ impl<'a> LowerBuilder<'a> { return self.emit_slot_load(slot, span); } + if subscripts.is_empty() + && let Some(slot) = self.layout.binding(name.as_str()) + { + return self.emit_slot_load(slot, span); + } + if let Some(slot) = self.non_variable_layout_slot(name, subscripts) { return self.emit_slot_load(slot, span); } + if subscripts.is_empty() + && let Some(reference) = self.singleton_record_array_field_reference(name) + { + return self.lower_var_ref(&reference, &[], span, scope, call_depth); + } + + if subscripts.is_empty() + && let Some(reg) = + self.lower_var_ref_binding_key(name.as_str(), span, scope, call_depth)? + { + return Ok(reg); + } + + if subscripts.is_empty() { + let real_field_key = format!("{}.re", name.as_str()); + if self.scalarized_field_binding_available(name.as_str(), "re") + && let Some(reg) = + self.lower_var_ref_binding_key(&real_field_key, span, scope, call_depth)? + { + return Ok(reg); + } + } + let name_key = self.scope_key_from_reference(name, span)?; if subscripts.is_empty() && let Some(reg) = scope.get(&name_key).copied() @@ -879,6 +1008,18 @@ impl<'a> LowerBuilder<'a> { && scope.contains_key(&name_key) && !self.local_indexed_bindings.contains_key(name.as_str()) { + if let Some(dims) = self.local_binding_dims.get(name.as_str()) + && dims.iter().any(|dim| *dim < 0) + { + return Err(unsupported_at( + format!( + "subscripted local array `{}` has negative dimensions {}", + name.as_str(), + format_i64_dims(dims) + ), + owner_span, + )); + } return Err(LowerError::Unsupported { reason: format!( "subscripted local variable references are unsupported: {}[...]", @@ -976,6 +1117,32 @@ impl<'a> LowerBuilder<'a> { self.layout.binding(name.as_str()) } + fn singleton_record_array_field_reference( + &self, + name: &rumoca_core::Reference, + ) -> Option { + let component_ref = name.component_ref()?; + if component_ref.def_id.is_some() || component_ref.parts.len() < 2 { + return None; + } + let mut candidate_ref = component_ref.clone(); + let base_index = candidate_ref.parts.len().checked_sub(2)?; + let span = candidate_ref.parts[base_index].span; + let subscript = + rumoca_core::Subscript::try_generated_index(1, span, "singleton record array field") + .ok()?; + candidate_ref.parts[base_index].subs.push(subscript); + let candidate = candidate_ref.to_var_name().to_string(); + let variable = self.dae_variables.and_then(|variables| { + dae_variable(variables, &rumoca_core::VarName::new(candidate.as_str())) + })?; + let component_ref = variable.component_ref.clone()?; + Some(rumoca_core::Reference::with_component_reference( + candidate.as_str(), + component_ref, + )) + } + fn generated_local_static_subscript_reg( &self, name: &rumoca_core::Reference, @@ -1004,16 +1171,17 @@ impl<'a> LowerBuilder<'a> { let Some(shape) = self.layout.shape(name) else { return Ok(None); }; - if shape.len() != subscripts.len() { + if subscripts.len() < shape.len() { return Ok(None); } let mut indices = crate::lower_vec_with_capacity( - subscripts.len(), + shape.len(), "singleton subscript index count", subscript_span_with_owner(subscripts, owner_span), )?; - for (subscript, dim) in subscripts.iter().zip(shape.iter().copied()) { + let (declared_subscripts, extra_subscripts) = subscripts.split_at(shape.len()); + for (subscript, dim) in declared_subscripts.iter().zip(shape.iter().copied()) { let index = match subscript { rumoca_core::Subscript::Index { value, span } if *value > 0 => { positive_i64_index(*value, span_or_owner(*span, owner_span))? @@ -1040,6 +1208,29 @@ impl<'a> LowerBuilder<'a> { } indices.push(index); } + for subscript in extra_subscripts { + let index = match subscript { + rumoca_core::Subscript::Index { value, span } if *value > 0 => { + positive_i64_index(*value, span_or_owner(*span, owner_span))? + } + rumoca_core::Subscript::Expr { expr, span } => { + match static_singleton_subscript_index(expr, span_or_owner(*span, owner_span))? + { + Some(value) => value, + None => return Ok(None), + } + } + rumoca_core::Subscript::Colon { .. } => return Ok(None), + _ => { + return Err(LowerError::Unsupported { + reason: "non-positive subscript is unsupported".to_string(), + }); + } + }; + if index != 1 { + return Ok(None); + } + } Ok(Some(indices)) } @@ -1099,6 +1290,7 @@ impl<'a> LowerBuilder<'a> { Ok(Some(values)) } + #[allow(clippy::too_many_lines, clippy::excessive_nesting)] fn lower_index( &mut self, base: &rumoca_core::Expression, @@ -1107,6 +1299,17 @@ impl<'a> LowerBuilder<'a> { scope: &Scope, call_depth: usize, ) -> Result { + if let rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor: false, + .. + } = base + && is_stream_passthrough_intrinsic(name.as_str()) + && let Some(arg) = args.first() + { + return self.lower_index(arg, subscripts, owner_span, scope, call_depth); + } if matches!( base, rumoca_core::Expression::FieldAccess { .. } @@ -1127,6 +1330,55 @@ impl<'a> LowerBuilder<'a> { { return Ok(reg); } + if dynamic_binding_base_key(base).is_err() + && let Some(span) = index_owner_span(base, subscripts, owner_span) + && let Some(indices) = static_subscript_indices_with_owner(subscripts, span)? + { + let dims = self.infer_expr_dims(base, scope)?; + if let Some(flat_index) = flat_index_from_one_based_usize_indices(&dims, &indices) { + let values = self + .lower_array_like_values_with_source_context(base, span, scope, call_depth)?; + if let Some(value) = values.get(flat_index).copied() { + return Ok(value); + } + } + } + if matches!(base, rumoca_core::Expression::FieldAccess { .. }) + && let Some(span) = index_owner_span(base, subscripts, owner_span) + && let Some(indices) = static_subscript_indices_with_owner(subscripts, span)? + { + let dims = self.infer_expr_dims(base, scope)?; + if let Some(flat_index) = flat_index_from_one_based_usize_indices(&dims, &indices) { + let mut dae_model = dae::Dae::default(); + dae_model.symbols.functions = self.functions.clone(); + if let Some(variables) = self.dae_variables { + dae_model.variables = variables.clone(); + } + if let Some(values) = derivative_rhs::function_call_projected_scalars_with_owner( + base, + &dae_model, + &self.structural_bindings, + span, + )? && let Some(value) = values.get(flat_index).cloned() + { + return self.lower_expr(&value, scope, call_depth + 1); + } + if let Some(value) = derivative_rhs::project_array_like_scalar_with_owner( + base, + flat_index, + &dae_model, + &self.structural_bindings, + span, + )? { + return self.lower_expr(&value, scope, call_depth + 1); + } + let values = self + .lower_array_like_values_with_source_context(base, span, scope, call_depth)?; + if let Some(value) = values.get(flat_index).copied() { + return Ok(value); + } + } + } if let Some(reg) = self.lower_array_like_dynamic_index(base, subscripts, owner_span, scope, call_depth)? { @@ -1154,7 +1406,7 @@ impl<'a> LowerBuilder<'a> { return self.emit_slot_load(slot, required_expression_span(base, "indexed slot load")?); } - if is_static_singleton_scalar_projection(base, subscripts)? { + if is_static_singleton_scalar_projection(base, subscripts, owner_span)? { return self.lower_expr(base, scope, call_depth); } @@ -1162,7 +1414,62 @@ impl<'a> LowerBuilder<'a> { return self.lower_expr(base, scope, call_depth); } - let base_key = dynamic_binding_base_key(base)?; + if dynamic_binding_base_key(base).is_err() + && self + .infer_expr_dims(base, scope) + .unwrap_or_default() + .is_empty() + && index_owner_span(base, subscripts, owner_span).is_some_and(|span| { + static_subscript_indices_with_owner(subscripts, span) + .is_ok_and(|indices| indices.is_some()) + }) + { + return self.lower_expr(base, scope, call_depth); + } + + if let Some(span) = index_owner_span(base, subscripts, owner_span) + && let Some(indices) = static_subscript_indices_with_owner(subscripts, span)? + && dynamic_binding_base_key(base).is_err() + { + let dims = self.infer_expr_dims(base, scope).unwrap_or_default(); + let flat_index = + flat_index_from_one_based_usize_indices(&dims, &indices).or_else(|| { + (indices.len() == 1) + .then(|| indices.first().and_then(|index| index.checked_sub(1))) + .flatten() + }); + if let Some(flat_index) = flat_index { + let mut dae_model = dae::Dae::default(); + dae_model.symbols.functions = self.functions.clone(); + if let Some(variables) = self.dae_variables { + dae_model.variables = variables.clone(); + } + if let Some(value) = derivative_rhs::project_array_like_scalar_with_owner( + base, + flat_index, + &dae_model, + &self.structural_bindings, + span, + )? { + return self.lower_expr(&value, scope, call_depth + 1); + } + } + let values = + self.lower_array_like_values_with_source_context(base, span, scope, call_depth)?; + if let Some(index) = flat_index_for_lowered_values(&dims, &indices) + && let Some(value) = values.get(index).copied() + { + return Ok(value); + } + } + + let base_key = match dynamic_binding_base_key(base) { + Ok(base_key) => base_key, + Err(err @ LowerError::DynamicBindingBase { .. }) => { + return Err(err); + } + Err(err) => return Err(err), + }; let source_key = component_reference_key_for_expr(base)?; let source_span = base.span(); self.lower_dynamic_subscripted_binding( @@ -1436,6 +1743,13 @@ impl<'a> LowerBuilder<'a> { ), } })?; + if self + .local_indexed_bindings + .contains_key(&target.display_key) + && !self.indexed_bindings.contains_key(source_key) + { + return Ok((target.display_key.clone(), Vec::new())); + } let entries = self.indexed_bindings.get(source_key).cloned().ok_or_else(|| { LowerError::contract_violation( format!( @@ -1536,6 +1850,7 @@ impl<'a> LowerBuilder<'a> { self.emit_unary_at(UnaryOp::Trunc, shifted, span) } + #[allow(clippy::too_many_lines, clippy::excessive_nesting)] fn lower_field_access( &mut self, base: &rumoca_core::Expression, @@ -1561,6 +1876,19 @@ impl<'a> LowerBuilder<'a> { } } + if let rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Mul | rumoca_core::OpBinary::MulElem, + lhs, + rhs, + .. + } = base + && matches!(field, "re" | "im") + && let Some(reg) = + self.lower_complex_vector_dot_field(lhs, rhs, field, field_access_span, scope)? + { + return Ok(reg); + } + if matches!(field, "re" | "im") && let Some(reg) = self.lower_complex_operator_field_access(base, field, scope, call_depth)? @@ -1584,6 +1912,37 @@ impl<'a> LowerBuilder<'a> { ); } + if let rumoca_core::Expression::Index { + base: indexed_base, + subscripts, + span, + } = base + && let rumoca_core::Expression::Binary { op, lhs, rhs, .. } = indexed_base.as_ref() + && matches!(op, rumoca_core::OpBinary::Add | rumoca_core::OpBinary::Sub) + { + let lhs = field_access_expr_with_owner( + &rumoca_core::Expression::Index { + base: lhs.clone(), + subscripts: subscripts.clone(), + span: *span, + }, + field, + field_access_span, + ); + let rhs = field_access_expr_with_owner( + &rumoca_core::Expression::Index { + base: rhs.clone(), + subscripts: subscripts.clone(), + span: *span, + }, + field, + field_access_span, + ); + let lhs_reg = self.lower_expr(&lhs, scope, call_depth)?; + let rhs_reg = self.lower_expr(&rhs, scope, call_depth)?; + return self.lower_binary(op.clone(), lhs_reg, rhs_reg, field_access_span); + } + if let rumoca_core::Expression::Index { base, subscripts, .. } = base @@ -1593,6 +1952,21 @@ impl<'a> LowerBuilder<'a> { return Ok(reg); } + if let rumoca_core::Expression::FieldAccess { + base: nested_base, + field: nested_field, + .. + } = base + && matches!(nested_base.as_ref(), rumoca_core::Expression::Index { .. }) + { + let nested_path = format!("{nested_field}.{field}"); + if let Some(reg) = + self.lower_indexed_field_access(nested_base, &nested_path, scope, call_depth)? + { + return Ok(reg); + } + } + if let Some(reg) = self.lower_indexed_field_access(base, field, scope, call_depth)? { return Ok(reg); } @@ -1601,6 +1975,48 @@ impl<'a> LowerBuilder<'a> { return Ok(reg); } + if let rumoca_core::Expression::FunctionCall { .. } = base { + let expr = rumoca_core::Expression::FieldAccess { + base: Box::new(base.clone()), + field: field.to_string(), + span: field_access_span, + }; + let mut dae_model = dae::Dae::default(); + dae_model.symbols.functions = self.functions.clone(); + if let Some(variables) = self.dae_variables { + dae_model.variables = variables.clone(); + } + if let Some(mut values) = derivative_rhs::function_call_projected_scalars_with_owner( + &expr, + &dae_model, + &self.structural_bindings, + field_access_span, + )? && values.len() == 1 + { + let value = values.remove(0); + if value != expr { + return self.lower_expr(&value, scope, call_depth + 1); + } + } + } + + if let rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor: false, + span, + } = base + && let Some(materialized) = self.materialize_single_record_function_call_components( + name, args, *span, scope, call_depth, + )? + && let Some(component) = materialized + .components + .into_iter() + .find(|component| component.suffix == field) + { + return Ok(component.reg); + } + if let Some(reg) = self.lower_function_output_name_field_access(base, field, scope, call_depth)? { @@ -1643,9 +2059,43 @@ impl<'a> LowerBuilder<'a> { if let Some(reg) = self.lower_var_ref_binding_key(&key, span, scope, call_depth)? { return Ok(reg); } + if let Some(reg) = + self.lower_singleton_record_array_field_binding(base, field, span, scope, call_depth)? + { + return Ok(reg); + } Err(LowerError::MissingBinding { name: key }) } + fn lower_singleton_record_array_field_binding( + &mut self, + base: &rumoca_core::Expression, + field: &str, + span: rumoca_core::Span, + scope: &Scope, + call_depth: usize, + ) -> Result, LowerError> { + let Ok(base_key) = binding_base_key(base) else { + return Ok(None); + }; + let singleton_key = format!("{base_key}[1].{field}"); + if self.layout.binding(&singleton_key).is_none() { + return Ok(None); + } + let higher_prefix = format!("{base_key}["); + let higher_suffix = format!("].{field}"); + let has_higher_member = self.layout.bindings().keys().any(|key| { + key.strip_prefix(higher_prefix.as_str()) + .and_then(|rest| rest.strip_suffix(higher_suffix.as_str())) + .and_then(|index| index.parse::().ok()) + .is_some_and(|index| index > 1) + }); + if has_higher_member { + return Ok(None); + } + self.lower_var_ref_binding_key(&singleton_key, span, scope, call_depth) + } + fn lower_function_output_name_field_access( &mut self, base: &rumoca_core::Expression, @@ -1693,17 +2143,28 @@ impl<'a> LowerBuilder<'a> { }; let base_key = binding_base_key(base)?; let field_key = format!("{base_key}.{field}"); - - if let Some(indices) = static_subscript_indices_with_owner( + if let Some(indices) = self.compile_time_subscript_indices( subscripts, required_expression_span(base, "indexed field access")?, )? && !indices.is_empty() { let key = format_subscript_binding_key(&field_key, &indices); - if let Some(reg) = scope.get(&generated_scope_key(&key)).copied() { + let record_element_field_key = format!( + "{}.{}", + format_subscript_binding_key(&base_key, &indices), + field + ); + if let Some(reg) = scope + .get(&generated_scope_key(&record_element_field_key)) + .or_else(|| scope.get(&generated_scope_key(&key))) + .copied() + { return Ok(Some(reg)); } - if let Some(slot) = self.pre_mode_slot_for_key(&key) { + if let Some(slot) = self + .pre_mode_slot_for_key(&record_element_field_key) + .or_else(|| self.pre_mode_slot_for_key(&key)) + { return self .emit_slot_load( slot, @@ -1711,13 +2172,24 @@ impl<'a> LowerBuilder<'a> { ) .map(Some); } - if let Some(values) = - self.lower_direct_assignment_values_for_key(&key, scope, call_depth)? + let values = match self.lower_direct_assignment_values_for_key( + &record_element_field_key, + scope, + call_depth, + )? { + Some(values) => Some(values), + None => self.lower_direct_assignment_values_for_key(&key, scope, call_depth)?, + }; + if let Some(values) = values && let Some(value) = values.first().copied() { return Ok(Some(value)); } - if let Some(slot) = self.layout.binding(&key) { + if let Some(slot) = self + .layout + .binding(&record_element_field_key) + .or_else(|| self.layout.binding(&key)) + { return self .emit_slot_load( slot, @@ -1748,6 +2220,78 @@ impl<'a> LowerBuilder<'a> { .map(Some) } + fn lower_complex_vector_dot_field( + &mut self, + lhs: &rumoca_core::Expression, + rhs: &rumoca_core::Expression, + field: &str, + span: rumoca_core::Span, + scope: &Scope, + ) -> Result, LowerError> { + let Some(lhs_re) = self.lower_complex_field_array_values(lhs, "re", span, scope)? else { + return Ok(None); + }; + let Some(lhs_im) = self.lower_complex_field_array_values(lhs, "im", span, scope)? else { + return Ok(None); + }; + let Some(rhs_re) = self.lower_complex_field_array_values(rhs, "re", span, scope)? else { + return Ok(None); + }; + let Some(rhs_im) = self.lower_complex_field_array_values(rhs, "im", span, scope)? else { + return Ok(None); + }; + let len = [lhs_re.len(), lhs_im.len(), rhs_re.len(), rhs_im.len()] + .into_iter() + .max() + .unwrap_or(0); + if len == 0 + || ![lhs_re.len(), lhs_im.len(), rhs_re.len(), rhs_im.len()] + .into_iter() + .all(|size| size == 1 || size == len) + { + return Ok(None); + } + let mut acc = self.emit_const_at(0.0, span)?; + for idx in 0..len { + let ar = lhs_re[if lhs_re.len() == 1 { 0 } else { idx }]; + let ai = lhs_im[if lhs_im.len() == 1 { 0 } else { idx }]; + let br = rhs_re[if rhs_re.len() == 1 { 0 } else { idx }]; + let bi = rhs_im[if rhs_im.len() == 1 { 0 } else { idx }]; + let term = if field == "re" { + let arbr = self.emit_binary_at(BinaryOp::Mul, ar, br, span)?; + let aibi = self.emit_binary_at(BinaryOp::Mul, ai, bi, span)?; + self.emit_binary_at(BinaryOp::Sub, arbr, aibi, span)? + } else { + let arbi = self.emit_binary_at(BinaryOp::Mul, ar, bi, span)?; + let aibr = self.emit_binary_at(BinaryOp::Mul, ai, br, span)?; + self.emit_binary_at(BinaryOp::Add, arbi, aibr, span)? + }; + acc = self.emit_binary_at(BinaryOp::Add, acc, term, span)?; + } + Ok(Some(acc)) + } + + fn lower_complex_field_array_values( + &mut self, + expr: &rumoca_core::Expression, + field: &str, + span: rumoca_core::Span, + scope: &Scope, + ) -> Result>, LowerError> { + if let rumoca_core::Expression::VarRef { name, .. } = expr { + let key = format!("{}.{}", name.as_str(), field); + if let Some(values) = self.lower_record_field_array_values(&key, span)? { + return Ok(Some(values)); + } + } + let projected = field_access_expr_with_owner(expr, field, span); + match self.lower_array_like_values(&projected, scope, 0) { + Ok(values) if !values.is_empty() => Ok(Some(values)), + Ok(_) => Ok(None), + Err(_) => Ok(None), + } + } + fn lower_if_field_access( &mut self, branches: &[(rumoca_core::Expression, rumoca_core::Expression)], @@ -1917,6 +2461,57 @@ fn index_owner_span( .or_else(|| owner_span.filter(|span| !span.is_dummy())) } +fn flat_index_from_one_based_usize_indices(dims: &[usize], indices: &[usize]) -> Option { + if dims.len() != indices.len() || dims.is_empty() { + return None; + } + let mut flat = 0usize; + for (axis, index) in indices.iter().copied().enumerate() { + let dim = dims[axis]; + if index == 0 || index > dim { + return None; + } + let stride = dims[axis + 1..].iter().product::(); + flat = flat.checked_add((index - 1).checked_mul(stride)?)?; + } + Some(flat) +} + +fn flat_index_for_lowered_values(dims: &[usize], indices: &[usize]) -> Option { + flat_index_from_one_based_usize_indices(dims, indices).or_else(|| { + (indices.len() == 1) + .then(|| indices.first().and_then(|index| index.checked_sub(1))) + .flatten() + }) +} + +fn component_reference_key_or_error( + name: &rumoca_core::Reference, + component_ref: &rumoca_core::ComponentReference, +) -> Result { + #[cfg(test)] + if let Some(key) = + crate::test_support::fixture_key_for_component_ref(component_ref, name.as_str()) + { + return Ok(key); + } + ComponentReferenceKey::from_component_reference(component_ref) + .map_err(|err| component_reference_lower_error(name, err)) +} + +fn component_reference_lower_error( + name: &rumoca_core::Reference, + err: rumoca_ir_solve::ComponentReferenceKeyError, +) -> LowerError { + LowerError::contract_violation( + format!( + "Solve lowering requires static component-reference metadata for `{}`: {err}", + name.as_str(), + ), + err.span, + ) +} + fn subscript_source_provenance(subscript: &rumoca_core::Subscript) -> Option { let span = subscript.span(); if !span.is_dummy() { diff --git a/crates/rumoca-phase-solve/src/lower/array_values.rs b/crates/rumoca-phase-solve/src/lower/array_values.rs index 70a91478a..fd2b5cfd9 100644 --- a/crates/rumoca-phase-solve/src/lower/array_values.rs +++ b/crates/rumoca-phase-solve/src/lower/array_values.rs @@ -1,3 +1,7 @@ +//! SPEC_0021 file-size exception: array value lowering still coordinates array +//! literals, comprehensions, slices, and structured assignments. split plan: +//! move array literal projection and dynamic selection dispatch into submodules. + use super::*; mod builtins; @@ -11,7 +15,6 @@ mod structural_standard; mod tests; use helpers::*; pub(super) use selection_helpers::*; - const MAX_STATIC_RANGE_VALUES: usize = 100_000; pub(in crate::lower) struct ArrayComprehensionLowerCtx<'a> { @@ -21,7 +24,6 @@ pub(in crate::lower) struct ArrayComprehensionLowerCtx<'a> { const_scope: &'a mut IndexMap, call_depth: usize, } - #[derive(Clone)] pub(super) struct ArrayOperand { pub(super) values: Vec, @@ -626,6 +628,17 @@ impl<'a> LowerBuilder<'a> { return self.emit_const_at(0.0, span).map(Some); } + if let Some(value) = self.compile_time_structural_element( + elements, + subscripts, + span, + scope, + call_depth, + projected_field, + )? { + return Ok(Some(value)); + } + let selector = self.lower_structural_index_selector(&subscripts[0], span, scope, call_depth)?; let fallback = self.emit_const_at(0.0, span)?; @@ -659,6 +672,50 @@ impl<'a> LowerBuilder<'a> { Ok(Some(merged)) } + fn compile_time_structural_element( + &mut self, + elements: &[rumoca_core::Expression], + subscripts: &[rumoca_core::Subscript], + span: rumoca_core::Span, + scope: &Scope, + call_depth: usize, + projected_field: Option<&str>, + ) -> Result, LowerError> { + let Some(indices) = self.compile_time_subscript_indices(subscripts, span)? else { + return Ok(None); + }; + let Some(index) = indices.first().copied() else { + return Ok(None); + }; + let Some(zero_based) = index.checked_sub(1) else { + return Err(LowerError::contract_violation( + "structural array subscript index must be one-based", + span, + )); + }; + let Some(element) = elements.get(zero_based) else { + return Ok(None); + }; + if subscripts.len() == 1 { + return self + .lower_structural_index_leaf( + element, + projected_field, + (!span.is_dummy()).then_some(span), + scope, + call_depth, + ) + .map(Some); + } + self.lower_structural_index_expr( + element, + &subscripts[1..], + scope, + call_depth, + projected_field, + ) + } + pub(in crate::lower) fn lower_structural_index_selector( &mut self, subscript: &rumoca_core::Subscript, @@ -930,6 +987,7 @@ impl<'a> LowerBuilder<'a> { self.lower_array_like_values(expr, scope, call_depth) } + #[allow(clippy::excessive_nesting)] fn lower_function_call_array_like_values( &mut self, name: &rumoca_core::Reference, @@ -956,7 +1014,19 @@ impl<'a> LowerBuilder<'a> { } if is_stream_passthrough_intrinsic(name.as_str()) { return match args.first() { - Some(arg) => self.lower_array_like_values(arg, scope, call_depth), + Some(arg) => { + let dims = self.infer_expr_dims(arg, scope)?; + if !dims.is_empty() + && checked_shape_size( + &dims, + "stream passthrough argument shape", + call_span, + )? == 0 + { + return Ok(Vec::new()); + } + self.lower_array_like_values(arg, scope, call_depth) + } None => Ok(Vec::new()), }; } @@ -973,6 +1043,14 @@ impl<'a> LowerBuilder<'a> { { return Ok(values); } + if is_constructor + && self.lookup_function(name).is_none() + && let Some(reg) = self.lower_scalar_type_constructor_array_value( + name, args, call_span, scope, call_depth, + )? + { + return Ok(vec![reg]); + } if self.is_record_constructor_call(name, is_constructor) { return self.lower_record_constructor_values( name, @@ -982,6 +1060,11 @@ impl<'a> LowerBuilder<'a> { call_depth, ); } + if let Some(values) = + self.lower_projected_function_call_array_values(fallback, call_span, scope, call_depth)? + { + return Ok(values); + } if let Some(values) = self.lower_user_function_call_array_values(name, args, call_span, scope, call_depth)? { @@ -992,6 +1075,150 @@ impl<'a> LowerBuilder<'a> { Ok(values) } + #[allow(clippy::excessive_nesting)] + fn lower_projected_function_call_array_values( + &mut self, + expr: &rumoca_core::Expression, + span: rumoca_core::Span, + scope: &Scope, + call_depth: usize, + ) -> Result>, LowerError> { + let Some(dae_variables) = self.dae_variables else { + return Ok(None); + }; + let mut dae_model = dae::Dae { + variables: dae_variables.clone(), + ..Default::default() + }; + dae_model.symbols.functions = self.functions.clone(); + let rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor: false, + .. + } = expr + else { + return Ok(None); + }; + let Some(dims) = self.vectorized_scalar_function_call_dims(name, args, scope, span)? else { + return Ok(None); + }; + let count = checked_shape_size(&dims, "projected function call shape", span)?; + let mut values = + array_vec_with_capacity(count, "projected function call value count", span)?; + for flat_index in 0..count { + let mut projected_args = array_vec_with_capacity( + args.len(), + "projected function call argument count", + span, + )?; + for arg in args { + let arg_dims = self.infer_expr_dims(arg, scope)?; + let (projection_dims, projection_index) = + if arg_dims.as_slice() == [1] && dims.as_slice() != [1] { + (arg_dims.as_slice(), 0) + } else { + (dims.as_slice(), flat_index) + }; + let Some(projected_arg) = derivative_rhs::project_expression_scalar( + arg, + projection_dims, + projection_index, + &dae_model, + &self.structural_bindings, + span, + )? + else { + return Ok(None); + }; + projected_args.push(projected_arg); + } + let scalar_call = rumoca_core::Expression::FunctionCall { + name: name.clone(), + args: projected_args, + is_constructor: false, + span, + }; + values.push(self.lower_expr(&scalar_call, scope, call_depth + 1)?); + } + Ok(Some(values)) + } + + fn vectorized_scalar_function_call_dims( + &self, + name: &rumoca_core::Reference, + args: &[rumoca_core::Expression], + scope: &Scope, + span: rumoca_core::Span, + ) -> Result>, LowerError> { + let Some(function) = self.lookup_function(name) else { + return Ok(None); + }; + if function.outputs.is_empty() { + return Ok(None); + } + let (named, positional) = + function_calls::split_named_and_positional_call_args(function.name.as_str(), args)?; + let mut positional_idx = 0usize; + let mut dims: Option> = None; + for input in &function.inputs { + let actual = named.get(input.name.as_str()).copied().or_else(|| { + function_calls::next_positional_function_input_arg( + input, + &positional, + &mut positional_idx, + ) + }); + let Some(actual) = actual else { + continue; + }; + if !input.dims.is_empty() { + continue; + } + let actual_dims = self.infer_expr_dims(actual, scope)?; + if actual_dims.is_empty() { + continue; + } + match &dims { + Some(_) if actual_dims.as_slice() == [1] => {} + Some(existing) if existing.as_slice() == [1] => dims = Some(actual_dims), + Some(existing) if existing != &actual_dims => { + return Err(LowerError::contract_violation( + format!( + "vectorized scalar function `{}` has actual dimensions {}, expected {}", + function.name, + format_usize_dims(&actual_dims), + format_usize_dims(existing), + ), + actual.span().unwrap_or(span), + )); + } + Some(_) => {} + None => dims = Some(actual_dims), + } + } + Ok(dims) + } + + fn lower_scalar_type_constructor_array_value( + &mut self, + name: &rumoca_core::Reference, + args: &[rumoca_core::Expression], + span: rumoca_core::Span, + scope: &Scope, + call_depth: usize, + ) -> Result, LowerError> { + let (named_args, positional_args) = + function_calls::split_named_and_positional_call_args(name.as_str(), args)?; + named_args + .get("start") + .copied() + .or_else(|| positional_args.first().copied()) + .map(|expr| self.lower_expr(expr, scope, call_depth + 1)) + .transpose() + .map_err(|err| err.with_fallback_span(span)) + } + pub(super) fn lower_multiplication_expr( &mut self, lhs: &rumoca_core::Expression, @@ -1000,6 +1227,15 @@ impl<'a> LowerBuilder<'a> { scope: &Scope, call_depth: usize, ) -> Result { + let lhs_scalar_call = + self.scalar_output_function_call_in_scalar_product(lhs, rhs, scope)?; + let rhs_scalar_call = + self.scalar_output_function_call_in_scalar_product(rhs, lhs, scope)?; + if lhs_scalar_call || rhs_scalar_call { + let l = self.lower_expr(lhs, scope, call_depth)?; + let r = self.lower_expr(rhs, scope, call_depth)?; + return self.lower_binary(rumoca_core::OpBinary::Mul, l, r, span); + } let lhs = self .lower_array_operand(lhs, scope, call_depth) .map_err(|err| err.with_fallback_span(span))?; @@ -1012,18 +1248,65 @@ impl<'a> LowerBuilder<'a> { match result.values.as_slice() { [value] => Ok(*value), values => Err(unsupported_at( - format!( - "non-scalar multiplication result with width {} is unsupported in scalar context (lhs_shape={}, rhs_shape={}, result_shape={})", - values.len(), - format_usize_dims(&lhs.dims), - format_usize_dims(&rhs.dims), - format_usize_dims(&result.dims), - ), + { + format!( + "non-scalar multiplication result with width {} is unsupported in scalar context (lhs_shape={}, lhs_values={}, rhs_shape={}, rhs_values={}, result_shape={})", + values.len(), + format_usize_dims(&lhs.dims), + lhs.values.len(), + format_usize_dims(&rhs.dims), + rhs.values.len(), + format_usize_dims(&result.dims), + ) + }, span, )), } } + fn scalar_output_function_call_in_scalar_product( + &self, + candidate: &rumoca_core::Expression, + other: &rumoca_core::Expression, + scope: &Scope, + ) -> Result { + let rumoca_core::Expression::FunctionCall { + name, + is_constructor: false, + .. + } = candidate + else { + return Ok(false); + }; + let Some(function) = self.lookup_function(name) else { + return Ok(self.selected_output_function_call_is_scalar(name)); + }; + let [output] = function.outputs.as_slice() else { + return Ok(false); + }; + if !output.dims.is_empty() { + return Ok(false); + } + Ok(self.infer_expr_dims(other, scope)?.is_empty()) + } + + fn selected_output_function_call_is_scalar(&self, name: &rumoca_core::Reference) -> bool { + rumoca_core::find_map_top_level_splits_rev(name.as_str(), |base_name, suffix| { + let function = self.functions.get(&rumoca_core::VarName::new(base_name))?; + let projection_suffix = + crate::projection_suffix::parse_output_projection_suffix(suffix)?; + let output = function + .outputs + .iter() + .find(|output| output.name == projection_suffix.output_name)?; + if output.dims.is_empty() || output.dims.len() == projection_suffix.indices.len() { + return Some(()); + } + None + }) + .is_some() + } + fn lower_binary_array_like_values( &mut self, op: &rumoca_core::OpBinary, @@ -1044,14 +1327,22 @@ impl<'a> LowerBuilder<'a> { scope, call_depth, ), - Op::Sub | Op::SubElem => self.lower_elementwise_binary_values( - BinaryOp::Sub, - lhs, - rhs, - span, - scope, - call_depth, - ), + Op::Sub | Op::SubElem => { + if let Some(values) = self + .lower_projected_function_residual_values(lhs, rhs, span, scope, call_depth)? + { + Ok(values) + } else { + self.lower_elementwise_binary_values( + BinaryOp::Sub, + lhs, + rhs, + span, + scope, + call_depth, + ) + } + } Op::MulElem => self.lower_elementwise_binary_values( BinaryOp::Mul, lhs, @@ -1060,7 +1351,7 @@ impl<'a> LowerBuilder<'a> { scope, call_depth, ), - Op::DivElem | Op::ExpElem => { + Op::DivElem | Op::Exp | Op::ExpElem => { let bin = if matches!(op, Op::DivElem) { BinaryOp::Div } else { @@ -1074,7 +1365,7 @@ impl<'a> LowerBuilder<'a> { let rhs = self.lower_array_operand(rhs, scope, call_depth)?; Ok(self.multiply_array_operands(&lhs, &rhs, span)?.values) } - Op::Exp | Op::And | Op::Or | Op::Lt | Op::Le | Op::Gt | Op::Ge | Op::Eq | Op::Neq => { + Op::And | Op::Or | Op::Lt | Op::Le | Op::Gt | Op::Ge | Op::Eq | Op::Neq => { let lhs = self.lower_array_operand(lhs, scope, call_depth)?; let rhs = self.lower_array_operand(rhs, scope, call_depth)?; if lhs.is_scalar() && rhs.is_scalar() { @@ -1098,6 +1389,48 @@ impl<'a> LowerBuilder<'a> { result.map_err(|err| err.with_fallback_span(span)) } + fn lower_projected_function_residual_values( + &mut self, + lhs: &rumoca_core::Expression, + rhs: &rumoca_core::Expression, + span: rumoca_core::Span, + scope: &Scope, + call_depth: usize, + ) -> Result>, LowerError> { + let Some(dae_variables) = self.dae_variables else { + return Ok(None); + }; + let mut dae_model = dae::Dae { + variables: dae_variables.clone(), + ..Default::default() + }; + dae_model.symbols.functions = self.functions.clone(); + let residual = rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(lhs.clone()), + rhs: Box::new(rhs.clone()), + span, + }; + let Some(expressions) = derivative_rhs::function_projected_residuals_with_owner( + &residual, + &dae_model, + self.structural_bindings.as_ref(), + span, + )? + else { + return Ok(None); + }; + let mut values = array_vec_with_capacity( + expressions.len(), + "projected function residual scalar count", + span, + )?; + for expression in expressions { + values.push(self.lower_expr(&expression, scope, call_depth + 1)?); + } + Ok(Some(values)) + } + fn lower_elementwise_binary_values( &mut self, op: BinaryOp, @@ -1149,6 +1482,7 @@ impl<'a> LowerBuilder<'a> { self.lower_array_operand(expr, scope, call_depth) } + #[allow(clippy::excessive_nesting)] fn lower_division_array_values( &mut self, lhs: &rumoca_core::Expression, @@ -1159,7 +1493,38 @@ impl<'a> LowerBuilder<'a> { ) -> Result, LowerError> { let lhs = self.lower_array_operand(lhs, scope, call_depth)?; let rhs = self.lower_array_operand(rhs, scope, call_depth)?; + if (lhs.values.is_empty() || rhs.values.is_empty()) + && (operand_allows_empty_values(&lhs, span)? + || operand_allows_empty_values(&rhs, span)?) + { + return Ok(Vec::new()); + } if !rhs.is_scalar() { + if lhs.is_scalar() { + let mut values = crate::lower_vec_with_capacity( + rhs.values.len(), + "scalar-array division value count", + span, + )?; + for value in rhs.values.iter().copied() { + values.push(self.emit_binary_at(BinaryOp::Div, lhs.values[0], value, span)?); + } + return Ok(values); + } + if lhs.values.len() == 1 && rhs.values.len() == 1 { + let mut values = crate::lower_vec_with_capacity( + 1, + "singleton array division value count", + span, + )?; + values.push(self.emit_binary_at( + BinaryOp::Div, + lhs.values[0], + rhs.values[0], + span, + )?); + return Ok(values); + } return Err(unsupported_at( // MLS §10.6.5: ordinary division is only array divided by // scalar. Element-wise division is represented by `./`. @@ -1181,6 +1546,7 @@ impl<'a> LowerBuilder<'a> { Ok(values) } + #[allow(clippy::excessive_nesting)] pub(super) fn lower_array_operand( &mut self, expr: &rumoca_core::Expression, @@ -1193,9 +1559,24 @@ impl<'a> LowerBuilder<'a> { )?; let dims = if self.expr_uses_flat_call_or_tuple_width(expr) { vector_dims(values.len()) + } else if let rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor: false, + .. + } = expr + && is_stream_passthrough_intrinsic(name.as_str()) + && let Some(arg) = args.first() + { + self.infer_expr_dims(arg, scope)? } else { self.infer_expr_dims(expr, scope)? }; + let dims = if (dims.is_empty() && values.len() > 1) || dims.contains(&0) { + vector_dims(values.len()) + } else { + dims + }; if !dims.is_empty() { let Some(span) = owner_span else { return Err(LowerError::UnspannedContractViolation { @@ -1204,6 +1585,12 @@ impl<'a> LowerBuilder<'a> { }; let expected = checked_shape_size(&dims, "array expression shape", span)?; if expected != values.len() { + if self.expr_has_scalarized_complex_width(expr) + || (expected > 0 && values.len() > expected && values.len() % expected == 0) + { + let width = values.len(); + return ArrayOperand::with_shape_span(values, vector_dims(width), span); + } let shape = format_usize_dims(&dims); return Err(unsupported_at( format!( @@ -1233,6 +1620,26 @@ impl<'a> LowerBuilder<'a> { } } + fn expr_has_scalarized_complex_width(&self, expr: &rumoca_core::Expression) -> bool { + if self.scalarized_complex_fields_available(expr) { + return true; + } + match expr { + rumoca_core::Expression::Binary { lhs, rhs, .. } => { + self.expr_has_scalarized_complex_width(lhs) + || self.expr_has_scalarized_complex_width(rhs) + } + rumoca_core::Expression::Unary { rhs, .. } => { + self.expr_has_scalarized_complex_width(rhs) + } + rumoca_core::Expression::FieldAccess { base, .. } + | rumoca_core::Expression::Index { base, .. } => { + self.expr_has_scalarized_complex_width(base) + } + _ => false, + } + } + // SPEC_0021: Exception - Modelica multiplication shape cases are kept in // one match so scalar/vector/matrix semantics remain auditable together. #[allow(clippy::excessive_nesting)] @@ -1242,6 +1649,16 @@ impl<'a> LowerBuilder<'a> { rhs: &ArrayOperand, span: rumoca_core::Span, ) -> Result { + if lhs.values.is_empty() || rhs.values.is_empty() { + let dims = if lhs.values.is_empty() && !lhs.dims.is_empty() { + lhs.dims.clone() + } else if rhs.values.is_empty() && !rhs.dims.is_empty() { + rhs.dims.clone() + } else { + Vec::new() + }; + return ArrayOperand::with_shape_span(Vec::new(), dims, span); + } match (lhs.dims.as_slice(), rhs.dims.as_slice()) { ([], []) => Ok(ArrayOperand::scalar_with_span( self.emit_binary_at(BinaryOp::Mul, lhs.values[0], rhs.values[0], span)?, @@ -1666,6 +2083,57 @@ impl<'a> LowerBuilder<'a> { })) } + fn lower_structured_index_local_slice_values( + &mut self, + name: &rumoca_core::Reference, + subscripts: &[rumoca_core::Subscript], + shape: &[usize], + span: rumoca_core::Span, + scope: &Scope, + call_depth: usize, + ) -> Result>, LowerError> { + let key = name.as_str(); + let mut probe = crate::lower_vec_with_capacity( + subscripts.len(), + "local dynamic slice probe rank", + span, + )?; + let mut has_structured_index = false; + for subscript in subscripts { + if matches!( + subscript, + rumoca_core::Subscript::Expr { expr, .. } + if matches!(expr.as_ref(), rumoca_core::Expression::Index { .. }) + ) { + has_structured_index = true; + probe.push(rumoca_core::Subscript::try_generated_index( + 1, + span, + "local dynamic slice probe", + )?); + } else { + probe.push(subscript.clone()); + } + } + if !has_structured_index { + return Ok(None); + } + self.slice_selections(&probe, shape, span, scope) + .map_err(|err| local_slice_error(err, key, shape))?; + self.lower_array_like_dynamic_selection_values( + &rumoca_core::Expression::VarRef { + name: name.clone(), + subscripts: Vec::new(), + span, + }, + subscripts, + Some(span), + scope, + call_depth, + ) + .map_err(|err| local_slice_error(err, key, shape)) + } + /// Resolve `name[subscripts]` against a function-scope (local) array /// binding, ignoring any global model variable of the same name. Once the /// scoped local binding exists, resolution must either produce local values @@ -1705,6 +2173,13 @@ impl<'a> LowerBuilder<'a> { } return Ok(LocalSubscriptResolution::NotLocal); }; + if dims.iter().any(|dim| *dim < 0) { + let shape = format_i64_dims(&dims); + return Err(unsupported_at( + format!("subscripted local array `{key}` has negative dimensions {shape}"), + span, + )); + }; let Some(bindings) = scope.indexed_entries(&key_path) else { if let Some(indices) = self.compile_time_subscript_indices(subscripts, span)? { return Err(LowerError::MissingBinding { @@ -1716,13 +2191,6 @@ impl<'a> LowerBuilder<'a> { span, )); }; - if dims.iter().any(|dim| *dim < 0) { - let shape = format_i64_dims(&dims); - return Err(unsupported_at( - format!("subscripted local array `{key}` has negative dimensions {shape}"), - span, - )); - }; let mut shape = crate::lower_vec_with_capacity(dims.len(), "local array shape rank", span)?; for dim in &dims { let Ok(dim) = usize::try_from(*dim) else { @@ -1752,14 +2220,15 @@ impl<'a> LowerBuilder<'a> { name: format_subscript_binding_key(key, &indices), }); } - let selections = self - .slice_selections(subscripts, &shape, span, scope) - .map_err(|err| { - let shape = format_usize_dims(&shape); - err.with_context(format!( - "resolving subscripted local array `{key}` with shape {shape}" - )) - })?; + let selections = self.slice_selections(subscripts, &shape, span, scope); + if selections.is_err() + && let Some(values) = self.lower_structured_index_local_slice_values( + name, subscripts, &shape, span, scope, call_depth, + )? + { + return Ok(LocalSubscriptResolution::Values(values)); + } + let selections = selections.map_err(|err| local_slice_error(err, key, &shape))?; let mut combos = Vec::new(); collect_slice_index_combos(&selections, 0, &mut Vec::new(), &mut combos); let mut regs = @@ -1783,6 +2252,18 @@ impl<'a> LowerBuilder<'a> { scope: &Scope, call_depth: usize, ) -> Result, LowerError> { + if let rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor: false, + .. + } = base + && is_stream_passthrough_intrinsic(name.as_str()) + && let Some(arg) = args.first() + { + return self + .lower_index_array_like_values(arg, subscripts, owner_span, scope, call_depth); + } if matches!( base, rumoca_core::Expression::FieldAccess { .. } @@ -1806,18 +2287,34 @@ impl<'a> LowerBuilder<'a> { { return Ok(vec![value]); } - if let Some(values) = self.lower_array_like_dynamic_selection_values( - base, subscripts, owner_span, scope, call_depth, - )? { - return Ok(values); + if let rumoca_core::Expression::VarRef { name, .. } = base + && subscripts + .iter() + .any(|subscript| matches!(subscript, rumoca_core::Subscript::Colon { .. })) + { + return self + .lower_subscripted_var_ref_array_like_values(name, subscripts, scope, call_depth); } - if is_static_singleton_scalar_projection(base, subscripts)? { + if is_static_singleton_scalar_projection(base, subscripts, owner_span)? { return Ok(vec![self.lower_expr(base, scope, call_depth)?]); } if scalar_literal_projection(base, subscripts, owner_span)? { return Ok(vec![self.lower_expr(base, scope, call_depth)?]); } - let base_key = dynamic_binding_base_key(base)?; + if let Some(values) = self.lower_array_like_dynamic_selection_values( + base, subscripts, owner_span, scope, call_depth, + )? { + return Ok(values); + } + let base_key = match dynamic_binding_base_key(base) { + Ok(base_key) => base_key, + Err(LowerError::DynamicBindingBase { .. }) => { + return Ok(vec![ + self.lower_index(base, subscripts, owner_span, scope, call_depth)?, + ]); + } + Err(err) => return Err(err), + }; let span = required_expr_span_from_subscripts_or_base( subscripts, base, @@ -1833,6 +2330,13 @@ impl<'a> LowerBuilder<'a> { } } +fn local_slice_error(err: LowerError, key: &str, shape: &[usize]) -> LowerError { + err.with_context(format!( + "resolving subscripted local array `{key}` with shape {}", + format_usize_dims(shape) + )) +} + fn required_min_max_arg_span( expr: &rumoca_core::Expression, owner_span: Option, diff --git a/crates/rumoca-phase-solve/src/lower/array_values/builtins.rs b/crates/rumoca-phase-solve/src/lower/array_values/builtins.rs index 8467729bf..cd18b0740 100644 --- a/crates/rumoca-phase-solve/src/lower/array_values/builtins.rs +++ b/crates/rumoca-phase-solve/src/lower/array_values/builtins.rs @@ -763,13 +763,67 @@ impl<'a> LowerBuilder<'a> { } let value_expr = &args[0]; let count = self.fill_element_count(&args[1..], call_span)?; + if matches!( + value_expr, + rumoca_core::Expression::FunctionCall { + is_constructor: true, + .. + } + ) && self.requires_complex_projection(value_expr, scope)? + { + let span = value_expr.span().unwrap_or(call_span); + let (re, im) = self.lower_complex_operand_parts(value_expr, span, scope, call_depth)?; + let mut values = reg_vec_with_capacity( + count.saturating_mul(2), + "complex fill value count", + value_expr.span(), + )?; + values.resize(count, re); + values.resize(count.saturating_mul(2), im); + return Ok(values); + } let mut values = reg_vec_with_capacity(count, "fill value count", value_expr.span())?; let value = self.lower_expr(value_expr, scope, call_depth)?; values.resize(count, value); Ok(values) } - fn fill_element_count( + pub(in crate::lower) fn lower_fill_field_array_like_values( + &mut self, + args: &[rumoca_core::Expression], + field: &str, + scope: &Scope, + call_depth: usize, + call_span: rumoca_core::Span, + ) -> Result>, LowerError> { + if args.len() < 2 { + return Err(array_builtin_contract_error( + format!("fill() requires at least 2 arguments, got {}", args.len()), + args.first() + .and_then(rumoca_core::Expression::span) + .or_else(|| (!call_span.is_dummy()).then_some(call_span)), + )); + } + let value_expr = &args[0]; + let count = self.fill_element_count(&args[1..], call_span)?; + let span = value_expr.span().unwrap_or(call_span); + let projected = field_access_expr_with_owner(value_expr, field, span); + let values = self.lower_array_like_values(&projected, scope, call_depth)?; + if values.is_empty() { + return Ok(None); + } + let total = values.len().checked_mul(count).ok_or_else(|| { + LowerError::contract_violation("fill field projection value count overflows", span) + })?; + let mut projected_values = + reg_vec_with_capacity(total, "fill field projection value count", Some(span))?; + for _ in 0..count { + projected_values.extend(values.iter().copied()); + } + Ok(Some(projected_values)) + } + + pub(in crate::lower) fn fill_element_count( &self, dims: &[rumoca_core::Expression], call_span: rumoca_core::Span, @@ -917,6 +971,13 @@ fn reg_vec_with_capacity( mod tests { use super::*; + fn literal_i64(value: i64, span: rumoca_core::Span) -> rumoca_core::Expression { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + span, + } + } + #[test] fn array_like_builtin_lowering_rejects_non_array_dispatch_with_span() { let span = rumoca_core::Span::from_offsets( @@ -1073,6 +1134,87 @@ mod tests { ); } + #[test] + fn ones_lowers_compile_time_max_of_size_range_dimension() { + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("ones_size_range.mo"), + 10, + 40, + ); + let layout = VarLayout::default(); + let functions = IndexMap::new(); + let mut builder = LowerBuilder::new(&layout, &functions); + let fill = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Fill, + args: vec![ + literal_i64(0, span), + literal_i64(0, span), + literal_i64(2, span), + ], + span, + }; + let range_end = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + args: vec![fill, literal_i64(2, span)], + span, + }; + let range = rumoca_core::Expression::Range { + start: Box::new(literal_i64(2, span)), + step: None, + end: Box::new(range_end), + span, + }; + let range_size = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + args: vec![range, literal_i64(1, span)], + span, + }; + let literal_size = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + args: vec![ + rumoca_core::Expression::Array { + elements: vec![literal_i64(0, span)], + is_matrix: false, + span, + }, + literal_i64(1, span), + ], + span, + }; + let max_size = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Max, + args: vec![rumoca_core::Expression::Array { + elements: vec![ + rumoca_core::Expression::Array { + elements: vec![range_size], + is_matrix: true, + span, + }, + rumoca_core::Expression::Array { + elements: vec![literal_size], + is_matrix: true, + span, + }, + ], + is_matrix: true, + span, + }], + span, + }; + + let values = builder + .lower_known_builtin_array_like_values( + rumoca_core::BuiltinFunction::Ones, + &[max_size], + &Scope::new(), + 0, + span, + ) + .expect("ones(max(size(range), size(array))) should lower"); + + assert_eq!(values.len(), 1); + } + #[test] fn linspace_arity_error_uses_call_span_without_args() { let call_span = rumoca_core::Span::from_offsets( diff --git a/crates/rumoca-phase-solve/src/lower/array_values/dynamic_selection.rs b/crates/rumoca-phase-solve/src/lower/array_values/dynamic_selection.rs index 0f8ada505..1d7709890 100644 --- a/crates/rumoca-phase-solve/src/lower/array_values/dynamic_selection.rs +++ b/crates/rumoca-phase-solve/src/lower/array_values/dynamic_selection.rs @@ -1,5 +1,149 @@ +// SPEC_0021 file-size exception: dynamic array selection still combines record +// slice projection, compile-time indexing, and runtime selection. split plan: +// move record slice paths and runtime index lowering into separate modules. +use super::inference::concrete_i64_dims; use super::*; +fn projected_field_expr_may_be_array_like(expr: &rumoca_core::Expression) -> bool { + match expr { + rumoca_core::Expression::Array { .. } => true, + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + branches + .iter() + .any(|(_, value)| projected_field_expr_may_be_array_like(value)) + || projected_field_expr_may_be_array_like(else_branch) + } + _ => false, + } +} + +fn resolve_zero_dims_from_value_count( + dims: &[usize], + value_count: usize, + span: rumoca_core::Span, +) -> Result>, LowerError> { + let zero_count = dims.iter().filter(|dim| **dim == 0).count(); + if zero_count != 1 || value_count == 0 { + return Ok(None); + } + let known_product = dims + .iter() + .filter(|dim| **dim != 0) + .try_fold(1usize, |acc, dim| acc.checked_mul(*dim)) + .ok_or_else(|| { + LowerError::contract_violation("array-like dynamic selection shape overflows", span) + })?; + if known_product == 0 || !value_count.is_multiple_of(known_product) { + return Ok(None); + } + let mut resolved = dims.to_vec(); + let Some(zero_index) = resolved.iter().position(|dim| *dim == 0) else { + return Ok(None); + }; + resolved[zero_index] = value_count / known_product; + Ok(Some(resolved)) +} + +fn collect_function_output_slice_indices( + per_dim: &[Vec], + dim_index: usize, + current: &mut Vec, + out: &mut Vec>, +) -> Result<(), LowerError> { + if dim_index == per_dim.len() { + out.try_reserve_exact(1) + .map_err(|_| LowerError::UnspannedContractViolation { + reason: "function output slice index count exceeds host memory limits".to_string(), + })?; + out.push(current.clone()); + return Ok(()); + } + for index in &per_dim[dim_index] { + current.push(*index); + collect_function_output_slice_indices(per_dim, dim_index + 1, current, out)?; + current.pop(); + } + Ok(()) +} + +fn collect_full_shape_binding_keys( + base_name: &str, + shape: &[usize], + depth: usize, + current: &mut Vec, + keys: &mut Vec, +) { + if depth == shape.len() { + keys.push(format_subscript_binding_key(base_name, current)); + return; + } + for index in 1..=shape[depth] { + current.push(index); + collect_full_shape_binding_keys(base_name, shape, depth + 1, current, keys); + current.pop(); + } +} + +struct RecordArraySliceFieldPath<'a> { + name: &'a rumoca_core::Reference, + subscripts: &'a [rumoca_core::Subscript], + span: rumoca_core::Span, + fields: Vec, +} + +fn record_array_slice_field_path<'a>( + base: &'a rumoca_core::Expression, + leaf_field: &str, +) -> Option> { + let mut fields = vec![leaf_field.to_string()]; + let mut cursor = base; + let mut span = base.span()?; + while let rumoca_core::Expression::FieldAccess { + base: nested, + field, + span: field_span, + } = cursor + { + fields.push(field.clone()); + span = *field_span; + cursor = nested; + } + fields.reverse(); + let rumoca_core::Expression::Index { + base: inner, + subscripts, + span: index_span, + } = cursor + else { + return None; + }; + let rumoca_core::Expression::VarRef { + name, + subscripts: ref_subscripts, + .. + } = inner.as_ref() + else { + return None; + }; + if !ref_subscripts.is_empty() { + return None; + } + Some(RecordArraySliceFieldPath { + name, + subscripts, + span: if index_span.is_dummy() { + span + } else { + *index_span + }, + fields, + }) +} + impl<'a> LowerBuilder<'a> { pub(in crate::lower) fn lower_compile_time_indexed_local_value( &mut self, @@ -120,7 +264,12 @@ impl<'a> LowerBuilder<'a> { ) { return Ok(None); } - let dims = self.infer_expr_dims(base, scope)?; + if matches!(base, rumoca_core::Expression::Binary { .. }) + && subscripts.iter().all(is_scalar_selector_subscript) + { + return Ok(None); + } + let mut dims = self.infer_expr_dims(base, scope)?; if dims.is_empty() || subscripts.len() > dims.len() { return Ok(None); } @@ -134,7 +283,16 @@ impl<'a> LowerBuilder<'a> { let values = self.lower_array_like_values_with_optional_source_context( base, owner_span, scope, call_depth, )?; + if dims.contains(&0) + && !values.is_empty() + && let Some(resolved) = resolve_zero_dims_from_value_count(&dims, values.len(), span)? + { + dims = resolved; + } let index_tuples = one_based_index_tuples(&dims, span)?; + if index_tuples.is_empty() { + return Ok(Some(Vec::new())); + } if values.len() != index_tuples.len() { return Err(unsupported_at( format!( @@ -329,14 +487,30 @@ impl<'a> LowerBuilder<'a> { span: rumoca_core::Span, scope: &Scope, ) -> Result>, LowerError> { - let Some(shape) = self.layout.shape(base_name) else { + let shape: Option> = if let Some(shape) = self.layout.shape(base_name) { + Some(shape.to_vec()) + } else if let Some(variable) = self + .dae_variables + .and_then(|variables| dae_variable(variables, &rumoca_core::VarName::new(base_name))) + .filter(|variable| !variable.dims.is_empty()) + { + Some(concrete_i64_dims( + &variable.dims, + base_name, + "DAE variable dimensions", + span, + )?) + } else { + None + }; + let Some(shape) = shape else { return Ok(None); }; if subscripts.is_empty() { return Ok(None); } - let selections = self.slice_selections(subscripts, shape, span, scope)?; + let selections = self.slice_selections(subscripts, &shape, span, scope)?; let mut keys = Vec::new(); collect_slice_binding_keys(base_name, &selections, 0, &mut Vec::new(), &mut keys); Ok(Some(keys)) @@ -349,6 +523,8 @@ impl<'a> LowerBuilder<'a> { span: rumoca_core::Span, scope: &Scope, ) -> Result>, LowerError> { + let subscripts = + self.normalize_overspecified_scalar_slice_subscripts(subscripts, shape, span, scope)?; if subscripts.len() > shape.len() { return Err(unsupported_at( "array slice has more subscripts than dimensions", @@ -371,6 +547,58 @@ impl<'a> LowerBuilder<'a> { Ok(selections) } + fn normalize_overspecified_scalar_slice_subscripts<'s>( + &self, + subscripts: &'s [rumoca_core::Subscript], + shape: &[usize], + span: rumoca_core::Span, + scope: &Scope, + ) -> Result<&'s [rumoca_core::Subscript], LowerError> { + if subscripts.len() <= shape.len() { + return Ok(subscripts); + } + let (declared_subscripts, extra_subscripts) = subscripts.split_at(shape.len()); + for (subscript, dim) in declared_subscripts.iter().zip(shape.iter().copied()) { + if self.slice_subscript_indices(subscript, dim, scope)?.len() != 1 { + return Ok(subscripts); + } + } + for subscript in extra_subscripts { + if self.singleton_compile_time_subscript_index(subscript, span, scope)? != Some(1) { + return Ok(subscripts); + } + } + Ok(declared_subscripts) + } + + fn singleton_compile_time_subscript_index( + &self, + subscript: &rumoca_core::Subscript, + span: rumoca_core::Span, + scope: &Scope, + ) -> Result, LowerError> { + match subscript { + rumoca_core::Subscript::Index { value, span } if *value > 0 => { + crate::lower::helpers::positive_i64_index(*value, *span).map(Some) + } + rumoca_core::Subscript::Expr { expr, .. } => { + let const_scope = self.compile_time_slice_bindings(scope); + self.eval_compile_time_positive_index_at( + expr, + &const_scope, + "array singleton projection subscript", + span, + ) + .map(Some) + } + rumoca_core::Subscript::Colon { .. } => Ok(None), + _ => Err(unsupported_at( + "non-positive subscript is unsupported".to_string(), + subscript.span(), + )), + } + } + pub(in crate::lower) fn slice_subscript_indices( &self, subscript: &rumoca_core::Subscript, @@ -486,16 +714,58 @@ impl<'a> LowerBuilder<'a> { { return Ok(values); } + if let rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Fill, + args, + span, + } = base + && let Some(values) = + self.lower_fill_field_array_like_values(args, field, scope, call_depth, *span)? + { + return Ok(values); + } if let Some(values) = self.lower_constructor_field_array_like_values(base, field, scope, call_depth)? { return Ok(values); } + if let Some(values) = self.lower_projected_function_field_array_like_values( + base, field, expr, scope, call_depth, + )? { + return Ok(values); + } if let Some(values) = self.lower_record_array_slice_field_values(base, field, owner_span, scope, call_depth)? { return Ok(values); } + if let rumoca_core::Expression::Binary { op, lhs, rhs, span } = base + && matches!(op, rumoca_core::OpBinary::Add | rumoca_core::OpBinary::Sub) + { + let lhs_field = field_access_expr_with_owner(lhs, field, *span); + let rhs_field = field_access_expr_with_owner(rhs, field, *span); + let lhs_values = self.lower_array_like_values(&lhs_field, scope, call_depth)?; + let rhs_values = self.lower_array_like_values(&rhs_field, scope, call_depth)?; + if lhs_values.len() != rhs_values.len() { + return Err(LowerError::contract_violation( + format!( + "binary field projection `{field}` has mismatched array widths {} and {}", + lhs_values.len(), + rhs_values.len() + ), + *span, + )); + } + let mut values = array_vec_with_capacity( + lhs_values.len(), + "binary field projection value count", + *span, + )?; + for (lhs, rhs) in lhs_values.into_iter().zip(rhs_values) { + values.push(self.lower_binary(op.clone(), lhs, rhs, *span)?); + } + return Ok(values); + } if let Ok(key) = field_access_binding_key(base, field) { let span = owner_span.ok_or_else(|| LowerError::UnspannedContractViolation { reason: format!("field access `{key}` requires source span metadata"), @@ -510,6 +780,12 @@ impl<'a> LowerBuilder<'a> { if let Some(values) = self.lower_indexed_binding_values_at(key.as_str(), span)? { return Ok(values); } + if let Some(values) = self.lower_record_field_array_values(key.as_str(), span)? { + return Ok(values); + } + if let Some(values) = self.lower_shaped_flattened_field_values(&key, span)? { + return Ok(values); + } if let Some(reg) = self.lower_var_ref_binding_key(&key, span, scope, call_depth)? { return single_reg_vec(reg, "scalarized record field access value count", span); } @@ -523,51 +799,143 @@ impl<'a> LowerBuilder<'a> { Ok(values) } - /// Lowers a record-array member slice such as `ac.pin[:].v` by loading - /// each scalarized element variable `ac.pin[k].v`. Declines when the - /// selection is not a full one-dimensional colon slice over a structured - /// base or no element variables exist in the layout. - fn lower_record_array_slice_field_values( + fn lower_shaped_flattened_field_values( + &mut self, + key: &str, + span: rumoca_core::Span, + ) -> Result>, LowerError> { + let shape = if let Some(shape) = self.layout.shape(key) { + shape.to_vec() + } else if let Some(variable) = self + .dae_variables + .and_then(|variables| dae_variable(variables, &rumoca_core::VarName::new(key))) + .filter(|variable| !variable.dims.is_empty()) + { + concrete_i64_dims(&variable.dims, key, "flattened field dimensions", span)? + } else { + return Ok(None); + }; + if shape.is_empty() { + return Ok(None); + } + let mut keys = Vec::new(); + collect_full_shape_binding_keys(key, &shape, 0, &mut Vec::new(), &mut keys); + self.load_binding_keys(&keys, span).map(Some) + } + + #[allow(clippy::excessive_nesting)] + fn lower_projected_function_field_array_like_values( &mut self, base: &rumoca_core::Expression, field: &str, - owner_span: Option, + expr: &rumoca_core::Expression, scope: &Scope, call_depth: usize, ) -> Result>, LowerError> { - let rumoca_core::Expression::Index { - base: inner, - subscripts, - span, - } = base - else { + if !matches!(base, rumoca_core::Expression::FunctionCall { .. }) { return Ok(None); - }; - let rumoca_core::Expression::VarRef { + } + let span = expr + .span() + .or_else(|| self.active_source_context_span()) + .ok_or_else(|| LowerError::UnspannedContractViolation { + reason: "projected function field array access requires source span".to_string(), + })?; + if let rumoca_core::Expression::FunctionCall { name, - subscripts: ref_subscripts, + args, + is_constructor: false, .. - } = inner.as_ref() - else { + } = base + && let Some(materialized) = self.materialize_single_record_function_call_components( + name, args, span, scope, call_depth, + )? + { + for (suffix, bindings) in materialized.indexed_components { + if suffix == field { + let mut values = bindings + .into_iter() + .map(|binding| (binding.indices, binding.reg)) + .collect::>(); + values.sort_by(|(lhs, _), (rhs, _)| lhs.cmp(rhs)); + return Ok(Some(values.into_iter().map(|(_, reg)| reg).collect())); + } + } + for component in materialized.components { + if component.suffix == field { + return Ok(Some(vec![component.reg])); + } + } + } + let mut dae_model = dae::Dae::default(); + dae_model.symbols.functions = self.functions.clone(); + if let Some(variables) = self.dae_variables { + dae_model.variables = variables.clone(); + } + let field_expr = rumoca_core::Expression::FieldAccess { + base: Box::new(base.clone()), + field: field.to_string(), + span, + }; + let Some(expressions) = (match derivative_rhs::function_call_projected_scalars_with_owner( + &field_expr, + &dae_model, + self.structural_bindings.as_ref(), + span, + ) { + Ok(values) => values, + Err(_) => return Ok(None), + }) else { + return Ok(None); + }; + if expressions.len() == 1 && projected_field_expr_may_be_array_like(&expressions[0]) { + return self + .lower_array_like_values(&expressions[0], scope, call_depth + 1) + .map(Some); + } + let mut values = array_vec_with_capacity( + expressions.len(), + "projected function field array scalar count", + span, + )?; + for expression in expressions { + values.push(self.lower_expr(&expression, scope, call_depth + 1)?); + } + Ok(Some(values)) + } + + /// Lowers a record-array member slice such as `ac.pin[:].v` or + /// `coil.ele[:].vol1.T` by loading each scalarized element variable. + /// Declines when the selection is not a full one-dimensional colon slice + /// over a structured base or no element variables exist in the layout. + fn lower_record_array_slice_field_values( + &mut self, + base: &rumoca_core::Expression, + field: &str, + owner_span: Option, + scope: &Scope, + call_depth: usize, + ) -> Result>, LowerError> { + let Some(slice) = record_array_slice_field_path(base, field) else { return Ok(None); }; - if !ref_subscripts.is_empty() - || subscripts.len() != 1 - || !matches!(subscripts[0], rumoca_core::Subscript::Colon { .. }) + if slice.subscripts.len() != 1 + || !matches!(slice.subscripts[0], rumoca_core::Subscript::Colon { .. }) { return Ok(None); } - let Some(component_ref) = name.component_ref() else { + let Some(component_ref) = slice.name.component_ref() else { return Ok(None); }; - let span = Self::non_dummy_span(*span).or(owner_span).ok_or_else(|| { - LowerError::UnspannedContractViolation { + let span = Self::non_dummy_span(slice.span) + .or(owner_span) + .ok_or_else(|| LowerError::UnspannedContractViolation { reason: "missing source provenance for record array member slice".to_string(), - } - })?; + })?; + let field_path = slice.fields.join("."); let mut values = Vec::new(); for element in 1.. { - let key = format!("{}[{element}].{field}", name.as_str()); + let key = format!("{}[{element}].{field_path}", slice.name.as_str()); if self.layout.binding(&key).is_none() { break; } @@ -581,11 +949,13 @@ impl<'a> LowerBuilder<'a> { span, "record array member slice subscript", )?]; - element_ref.parts.push(rumoca_core::ComponentRefPart { - ident: field.to_string(), - span, - subs: Vec::new(), - }); + for field in &slice.fields { + element_ref.parts.push(rumoca_core::ComponentRefPart { + ident: field.clone(), + span, + subs: Vec::new(), + }); + } let element_expr = rumoca_core::Expression::VarRef { name: rumoca_core::Reference::from_component_reference(element_ref), subscripts: vec![], @@ -596,7 +966,7 @@ impl<'a> LowerBuilder<'a> { if values.is_empty() { return Ok(None); } - self.ensure_dense_record_array_slice(name.as_str(), field, values.len(), span)?; + self.ensure_dense_record_array_slice(slice.name.as_str(), &field_path, values.len(), span)?; Ok(Some(values)) } @@ -607,12 +977,12 @@ impl<'a> LowerBuilder<'a> { pub(in crate::lower) fn ensure_dense_record_array_slice( &self, base: &str, - field: &str, + field_path: &str, found: usize, span: rumoca_core::Span, ) -> Result<(), LowerError> { let prefix = format!("{base}["); - let suffix = format!("].{field}"); + let suffix = format!("].{field_path}"); for key in self.layout.bindings().keys() { let Some(index) = key .strip_prefix(prefix.as_str()) @@ -625,7 +995,7 @@ impl<'a> LowerBuilder<'a> { let missing = missing_record_array_member_index(found, span)?; return Err(LowerError::contract_violation( format!( - "record-array member slice `{base}[:].{field}` is not densely \ + "record-array member slice `{base}[:].{field_path}` is not densely \ scalarized: element {index} exists but element {missing} is missing", ), span, @@ -842,13 +1212,42 @@ impl<'a> LowerBuilder<'a> { caller_scope: &Scope, call_depth: usize, ) -> Result>, LowerError> { - if let Some(projection) = self.lookup_function_output_projection(name, span)? + let output_projection = self.lookup_function_output_projection(name, span)?; + if let Some(projection) = output_projection.as_ref() && let Some(values) = - self.lower_fft_projection_values(&projection, args, span, caller_scope, call_depth)? + self.lower_fft_projection_values(projection, args, span, caller_scope, call_depth)? + { + return Ok(Some(values)); + } + if output_projection.is_none() + && self.lookup_function(name).is_some_and(|function| { + function + .outputs + .first() + .is_some_and(|output| !output.dims.is_empty()) + }) + && let Some(values) = self.lower_user_function_call_output_values( + name, + args, + span, + caller_scope, + call_depth, + Some(0), + )? { return Ok(Some(values)); } - if let Some(projection) = self.lookup_function_output_projection(name, span)? + if let Some(values) = self.lower_projected_scalar_function_call_values( + name, + args, + span, + caller_scope, + call_depth, + )? { + return first_expression_output_values(self.lookup_function(name), values, span) + .map(Some); + } + if let Some(projection) = output_projection.as_ref() && projection.indices.is_empty() && projection.output_field.is_none() { @@ -955,9 +1354,6 @@ impl<'a> LowerBuilder<'a> { caller_scope: &Scope, call_depth: usize, ) -> Result>, LowerError> { - let Some(indices) = static_subscript_indices_with_owner(subscripts, span)? else { - return Ok(None); - }; let Some(function) = self.lookup_function(name) else { return Ok(None); }; @@ -965,9 +1361,21 @@ impl<'a> LowerBuilder<'a> { return Ok(None); }; let dims = function_output_dims(name, output)?; - let Some(flat_index) = flat_index_for_subscripts(&dims, &indices) else { + let flat_indices = + self.function_output_flat_indices(&dims, subscripts, caller_scope, span)?; + if flat_indices.is_empty() { return Ok(None); - }; + } + if let Some(values) = self.lower_projected_function_output_index_values( + name, + args, + span, + &flat_indices, + caller_scope, + call_depth, + )? { + return Ok(Some(values)); + } let Some(values) = self.lower_user_function_call_output_values( name, args, @@ -979,11 +1387,128 @@ impl<'a> LowerBuilder<'a> { else { return Ok(None); }; - values - .get(flat_index) - .copied() - .map(|value| single_reg_vec(value, "function output projection value count", span)) - .transpose() + let mut selected = array_vec_with_capacity( + flat_indices.len(), + "function output slice projection value count", + span, + )?; + for flat_index in flat_indices { + let Some(value) = values.get(flat_index).copied() else { + return Ok(None); + }; + selected.push(value); + } + Ok(Some(selected)) + } + + fn function_output_flat_indices( + &self, + dims: &[usize], + subscripts: &[rumoca_core::Subscript], + scope: &Scope, + span: rumoca_core::Span, + ) -> Result, LowerError> { + let selections = self.function_output_slice_indices(dims, subscripts, scope, span)?; + if selections.is_empty() { + return Ok(Vec::new()); + } + let mut flat_indices = array_vec_with_capacity( + selections.len(), + "function output flat slice index count", + span, + )?; + for indices in selections { + let Some(flat_index) = flat_index_for_subscripts(dims, &indices) else { + return Ok(Vec::new()); + }; + flat_indices.push(flat_index); + } + Ok(flat_indices) + } + + fn lower_projected_function_output_index_values( + &mut self, + name: &rumoca_core::Reference, + args: &[rumoca_core::Expression], + span: rumoca_core::Span, + flat_indices: &[usize], + caller_scope: &Scope, + call_depth: usize, + ) -> Result>, LowerError> { + let mut dae_model = dae::Dae::default(); + dae_model.symbols.functions = self.functions.clone(); + if let Some(variables) = self.dae_variables { + dae_model.variables = variables.clone(); + } + let call = rumoca_core::Expression::FunctionCall { + name: name.clone(), + args: args.to_vec(), + is_constructor: false, + span, + }; + let Some(expressions) = (match derivative_rhs::function_call_projected_scalars_with_owner( + &call, + &dae_model, + self.structural_bindings.as_ref(), + span, + ) { + Ok(values) => values, + Err(_) => return Ok(None), + }) else { + return Ok(None); + }; + let mut values = array_vec_with_capacity( + flat_indices.len(), + "projected function output slice value count", + span, + )?; + for &flat_index in flat_indices { + let Some(expression) = expressions.get(flat_index) else { + return Ok(None); + }; + values.push(self.lower_expr(expression, caller_scope, call_depth + 1)?); + } + Ok(Some(values)) + } + + fn function_output_slice_indices( + &self, + dims: &[usize], + subscripts: &[rumoca_core::Subscript], + scope: &Scope, + span: rumoca_core::Span, + ) -> Result>, LowerError> { + if subscripts.len() > dims.len() { + return Ok(Vec::new()); + } + let mut per_dim = array_vec_with_capacity(dims.len(), "function output slice rank", span)?; + for (dim_index, subscript) in subscripts.iter().enumerate() { + per_dim.push(self.slice_subscript_indices(subscript, dims[dim_index], scope)?); + } + for &dim in &dims[subscripts.len()..] { + per_dim.push(one_based_indices( + dim, + "function output implicit slice index count", + span, + )?); + } + let mut indices = array_vec_with_capacity( + per_dim + .iter() + .try_fold(1usize, |count, dim_indices| { + count.checked_mul(dim_indices.len()) + }) + .ok_or_else(|| { + LowerError::contract_violation( + "function output slice selection count overflows host index range", + span, + ) + })?, + "function output slice selection count", + span, + )?; + collect_function_output_slice_indices(&per_dim, 0, &mut Vec::new(), &mut indices)?; + Ok(indices) } pub(in crate::lower) fn lower_user_function_call_output_values( @@ -1383,7 +1908,7 @@ impl<'a> LowerBuilder<'a> { call_depth: usize, ) -> Result, LowerError> { let mut scope = scope.clone(); - let mut const_scope = IndexMap::::new(); + let mut const_scope = self.local_const_bindings.clone(); let mut values = Vec::new(); let mut ctx = ArrayComprehensionLowerCtx { indices, @@ -1411,7 +1936,9 @@ impl<'a> LowerBuilder<'a> { } append_reg_values( out, - self.lower_array_like_values(expr, ctx.scope, ctx.call_depth)?, + self.with_local_const_bindings(ctx.const_scope, |this| { + this.lower_array_like_values(expr, ctx.scope, ctx.call_depth) + })?, "array comprehension value count", self.required_dynamic_selection_span(&[expr], "array comprehension value count")?, )?; @@ -1769,6 +2296,7 @@ impl<'a> LowerBuilder<'a> { Ok(Some(values)) } + #[allow(clippy::excessive_nesting)] pub(in crate::lower) fn lower_if( &mut self, branches: &[(rumoca_core::Expression, rumoca_core::Expression)], @@ -1779,6 +2307,14 @@ impl<'a> LowerBuilder<'a> { let mut runtime_branches = Vec::new(); let mut selected_static_branch = None; for (cond, value) in branches { + if compile_time_string_condition_call(cond) { + let condition = self.eval_compile_time_expr(cond, &self.local_const_bindings)?; + if condition != 0.0 { + selected_static_branch = Some(value); + break; + } + continue; + } match lower_static_condition_truth(cond)? { Some(false) => {} Some(true) => { @@ -1808,6 +2344,51 @@ impl<'a> LowerBuilder<'a> { } } +fn compile_time_string_condition_call(expr: &rumoca_core::Expression) -> bool { + matches!( + expr, + rumoca_core::Expression::FunctionCall { name, .. } + if matches!( + name.as_str(), + "Modelica.Utilities.Strings.isEqual" | "Strings.isEqual" | "isEqual" + ) + ) +} + +fn first_expression_output_values( + function: Option<&rumoca_core::Function>, + values: Vec, + span: rumoca_core::Span, +) -> Result, LowerError> { + let Some(output) = function.and_then(|function| function.outputs.first()) else { + return Ok(values); + }; + let width = output.dims.iter().try_fold(1usize, |acc, dim| { + let dim = usize::try_from(*dim).map_err(|_| { + LowerError::contract_violation( + format!( + "function output `{}` has invalid dimension `{dim}`", + output.name + ), + span, + ) + })?; + acc.checked_mul(dim).ok_or_else(|| { + LowerError::contract_violation( + format!( + "function output `{}` scalar count exceeds host index range", + output.name + ), + span, + ) + }) + })?; + if values.len() <= width { + return Ok(values); + } + Ok(values.into_iter().take(width).collect()) +} + fn function_output_dims( function_name: &rumoca_core::Reference, output: &rumoca_core::FunctionParam, diff --git a/crates/rumoca-phase-solve/src/lower/array_values/helpers.rs b/crates/rumoca-phase-solve/src/lower/array_values/helpers.rs index cf67da75b..78f00316e 100644 --- a/crates/rumoca-phase-solve/src/lower/array_values/helpers.rs +++ b/crates/rumoca-phase-solve/src/lower/array_values/helpers.rs @@ -649,7 +649,29 @@ pub(super) fn broadcast_pairs( } let dims = broadcast_dims(lhs, rhs)?; let span = operand_pair_shape_span(lhs, rhs)?; - let count = checked_shape_size_or_scalar(&dims, "array broadcast value count", span)?; + let count = if dims.is_empty() { + checked_shape_size_or_scalar(&dims, "array broadcast scalar value count", span)? + } else { + checked_shape_size(&dims, "array broadcast value count", span)? + }; + if count == 0 { + return Ok(Vec::new()); + } + if lhs.values.is_empty() || rhs.values.is_empty() { + if operand_allows_empty_values(lhs, span)? || operand_allows_empty_values(rhs, span)? { + return Ok(Vec::new()); + } + return Err(LowerError::contract_violation( + format!( + "array broadcast operand values are missing for non-empty shape (lhs_shape={}, lhs_values={}, rhs_shape={}, rhs_values={})", + format_usize_dims(&lhs.dims), + lhs.values.len(), + format_usize_dims(&rhs.dims), + rhs.values.len() + ), + span, + )); + } let mut pairs = array_vec_with_capacity(count, "array broadcast pair count", span)?; for idx in 0..count { let lhs = if lhs.is_scalar() { @@ -667,6 +689,19 @@ pub(super) fn broadcast_pairs( Ok(pairs) } +pub(super) fn operand_allows_empty_values( + operand: &ArrayOperand, + span: rumoca_core::Span, +) -> Result { + if !operand.values.is_empty() { + return Ok(false); + } + if operand.dims.is_empty() { + return Ok(true); + } + Ok(checked_shape_size(&operand.dims, "array broadcast empty operand shape", span)? == 0) +} + fn trailing_vector_broadcast_pairs( lhs: &ArrayOperand, rhs: &ArrayOperand, @@ -1216,6 +1251,28 @@ mod tests { Ok(()) } + #[test] + fn broadcast_pairs_treats_erased_empty_operand_as_zero_rows() { + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name( + "phase_solve_lower_array_values_helpers_source_20.mo", + ), + 1, + 2, + ); + let lhs = ArrayOperand::scalar_with_span(1, span); + let rhs = ArrayOperand { + values: Vec::new(), + dims: Vec::new(), + shape_span: span, + }; + + let pairs = broadcast_pairs(&lhs, &rhs) + .expect("erased empty array operand should produce zero broadcast rows"); + + assert!(pairs.is_empty()); + } + #[test] fn multiplication_dims_reports_mismatch_span() -> Result<(), String> { let span = rumoca_core::Span::from_offsets( diff --git a/crates/rumoca-phase-solve/src/lower/array_values/inference.rs b/crates/rumoca-phase-solve/src/lower/array_values/inference.rs index 4f534a23d..6ead8c742 100644 --- a/crates/rumoca-phase-solve/src/lower/array_values/inference.rs +++ b/crates/rumoca-phase-solve/src/lower/array_values/inference.rs @@ -42,6 +42,11 @@ impl<'a> LowerBuilder<'a> { } if is_synchronous_array_like_intrinsic(name.as_str()) => { self.infer_expr_dims(required_arg(args, name.as_str(), *span)?, scope)? } + rumoca_core::Expression::FunctionCall { + name, args, span, .. + } if is_stream_passthrough_intrinsic(name.as_str()) => { + self.infer_expr_dims(required_arg(args, name.as_str(), *span)?, scope)? + } rumoca_core::Expression::FunctionCall { name, args, .. } if is_modelica_array_constructor_function(name) => { @@ -49,6 +54,10 @@ impl<'a> LowerBuilder<'a> { self.infer_expr_dims(element, scope) })? } + rumoca_core::Expression::FunctionCall { + is_constructor: true, + .. + } => Vec::new(), rumoca_core::Expression::FunctionCall { name, span, .. } => { self.infer_function_call_output_dims(name, *span)? } @@ -67,8 +76,7 @@ impl<'a> LowerBuilder<'a> { step, end, span, - } => lower_static_range_values(start, step.as_deref(), end, *span)? - .map_or_else(Vec::new, |values| vector_dims(values.len())), + } => self.infer_range_dims(start, step.as_deref(), end, *span)?, rumoca_core::Expression::If { branches, else_branch, @@ -113,6 +121,29 @@ impl<'a> LowerBuilder<'a> { None } + fn infer_range_dims( + &self, + start: &rumoca_core::Expression, + step: Option<&rumoca_core::Expression>, + end: &rumoca_core::Expression, + span: rumoca_core::Span, + ) -> Result, LowerError> { + if let Some(values) = lower_static_range_values(start, step, end, span)? { + return Ok(vector_dims(values.len())); + } + match self.eval_compile_time_range_values( + start, + step, + end, + span, + &self.local_const_bindings, + "range dimension inference", + ) { + Ok(values) => Ok(vector_dims(values.len())), + Err(_) => Ok(Vec::new()), + } + } + fn infer_tuple_flat_dims( &self, elements: &[rumoca_core::Expression], @@ -181,9 +212,17 @@ impl<'a> LowerBuilder<'a> { base: &rumoca_core::Expression, field: &str, ) -> Result>, LowerError> { - let rumoca_core::Expression::FunctionCall { name, .. } = base else { + let rumoca_core::Expression::FunctionCall { + name, + is_constructor, + .. + } = base + else { return Ok(None); }; + if self.is_record_constructor_call(name, *is_constructor) { + return self.infer_constructor_field_dims(base, field); + } let Some(function) = self.lookup_function(name) else { return Ok(None); }; @@ -325,6 +364,16 @@ impl<'a> LowerBuilder<'a> { return Ok(meta.dims.clone()); } } + if let Some(dims) = self.infer_record_array_aggregate_dims_from_dae_variables(name)? { + return Ok(dims); + } + if let Some(variable) = self + .dae_variables + .and_then(|variables| dae_variable(variables, name.var_name())) + && !variable.dims.is_empty() + { + return concrete_i64_dims(&variable.dims, name_text, "DAE variable dimensions", span); + } if let Some(dims) = self.local_binding_dims.get(name_text) { return concrete_i64_dims(dims, name_text, "local binding dimensions", span); } @@ -335,6 +384,18 @@ impl<'a> LowerBuilder<'a> { { return Ok(Vec::new()); } + if let Some(reference) = self.singleton_record_array_field_reference(name) + && let Some(variable) = self + .dae_variables + .and_then(|variables| dae_variable(variables, reference.var_name())) + { + return concrete_i64_dims( + &variable.dims, + reference.as_str(), + "singleton record-array field dimensions", + span, + ); + } let name_path = self.scope_key_from_reference(name, span)?; if let Some(values) = scoped_indexed_binding_values(scope, &name_path, span)? { return Ok(vector_dims(values.len())); @@ -342,6 +403,47 @@ impl<'a> LowerBuilder<'a> { Ok(Vec::new()) } + fn infer_record_array_aggregate_dims_from_dae_variables( + &self, + name: &rumoca_core::Reference, + ) -> Result>, LowerError> { + let Some(base_ref) = name.component_ref() else { + return Ok(None); + }; + if base_ref.parts.is_empty() + || base_ref + .parts + .last() + .is_some_and(|part| !part.subs.is_empty()) + { + return Ok(None); + } + let Some(variables) = self.dae_variables else { + return Ok(None); + }; + let mut extent = 0usize; + for variable in variables + .states + .values() + .chain(variables.algebraics.values()) + .chain(variables.inputs.values()) + .chain(variables.outputs.values()) + .chain(variables.parameters.values()) + .chain(variables.constants.values()) + .chain(variables.discrete_reals.values()) + .chain(variables.discrete_valued.values()) + { + let Some(candidate_ref) = variable.component_ref.as_ref() else { + continue; + }; + let Some(index) = record_array_prefix_index(base_ref, candidate_ref)? else { + continue; + }; + extent = extent.max(index); + } + Ok((extent > 0).then(|| vector_dims(extent))) + } + fn infer_required_subscripted_var_ref_dims( &self, name: &rumoca_core::Reference, @@ -800,12 +902,22 @@ impl<'a> LowerBuilder<'a> { scope: &Scope, ) -> Result, LowerError> { let owner_span = Self::non_dummy_span(span).or_else(|| base.span()); + if scalar_literal_projection(base, subscripts, owner_span)? { + return Ok(Vec::new()); + } if let Ok(base_name) = dynamic_binding_base_key(base) && let Some(dims) = self.infer_slice_dims(base_name.as_str(), subscripts, scope, owner_span)? { return Ok(dims); } + if dynamic_binding_base_key(base).is_err() + && subscripts + .iter() + .all(super::helpers::is_scalar_selector_subscript) + { + return Ok(Vec::new()); + } let base_dims = self.infer_expr_dims(base, scope)?; if base_dims.is_empty() && scalar_singleton_projection(subscripts) { return Ok(Vec::new()); @@ -858,6 +970,7 @@ impl<'a> LowerBuilder<'a> { ) } + #[allow(clippy::excessive_nesting)] pub(in crate::lower) fn compile_time_subscript_indices( &self, subscripts: &[rumoca_core::Subscript], @@ -883,11 +996,19 @@ impl<'a> LowerBuilder<'a> { if !matches!(expr.as_ref(), rumoca_core::Expression::Range { .. }) && self.expr_can_eval_compile_time(expr, const_scope) => { - indices.push(self.eval_compile_time_positive_index( + match self.eval_compile_time_positive_index( expr, const_scope, "array subscript", - )?); + ) { + Ok(index) => indices.push(index), + Err(_) + if matches!(expr.as_ref(), rumoca_core::Expression::VarRef { .. }) => + { + return Ok(None); + } + Err(err) => return Err(err), + } } rumoca_core::Subscript::Expr { .. } | rumoca_core::Subscript::Colon { .. } => { return Ok(None); @@ -988,7 +1109,7 @@ impl<'a> LowerBuilder<'a> { } } - pub(super) fn eval_compile_time_range_values( + pub(in crate::lower) fn eval_compile_time_range_values( &self, start: &rumoca_core::Expression, step: Option<&rumoca_core::Expression>, @@ -1102,11 +1223,60 @@ impl<'a> LowerBuilder<'a> { | Op::DivElem | Op::ExpElem => broadcast_shape(&lhs_dims, &rhs_dims, span)?, Op::Div if rhs_dims.is_empty() => lhs_dims, + Op::Div if lhs_dims.is_empty() => rhs_dims, _ => Vec::new(), }) } } +fn record_array_prefix_index( + base_ref: &rumoca_core::ComponentReference, + candidate_ref: &rumoca_core::ComponentReference, +) -> Result, LowerError> { + let prefix_len = base_ref.parts.len(); + if candidate_ref.parts.len() <= prefix_len || candidate_ref.local != base_ref.local { + return Ok(None); + } + let mut index = None; + for (base, candidate) in base_ref + .parts + .iter() + .zip(candidate_ref.parts[..prefix_len].iter()) + { + if base.ident != candidate.ident || !base.subs.is_empty() { + return Ok(None); + } + match candidate.subs.as_slice() { + [] => {} + [subscript] if index.is_none() => { + index = Some(static_positive_subscript_index(subscript)?); + } + _ => return Ok(None), + } + } + Ok(index) +} + +fn static_positive_subscript_index( + subscript: &rumoca_core::Subscript, +) -> Result { + match subscript { + rumoca_core::Subscript::Index { value, span } if *value > 0 => { + crate::lower::helpers::positive_i64_index(*value, *span) + } + rumoca_core::Subscript::Index { span, .. } => Err(unsupported_at( + "non-positive record-array aggregate index is unsupported", + *span, + )), + rumoca_core::Subscript::Colon { span } | rumoca_core::Subscript::Expr { span, .. } => { + Err(unsupported_at( + "dynamic record-array aggregate index is unsupported in shape inference", + *span, + )) + } + } +} + fn var_ref_is_translation_constant( variables: &dae::DaeVariables, name: &rumoca_core::Reference, @@ -1171,7 +1341,7 @@ fn scalar_singleton_projection(subscripts: &[rumoca_core::Subscript]) -> bool { }) } -fn concrete_i64_dims( +pub(super) fn concrete_i64_dims( dims: &[i64], name: &str, context: &str, @@ -1276,6 +1446,168 @@ mod tests { assert_eq!(err.reason(), "non-positive subscript is unsupported"); } + #[test] + fn range_slice_keeps_singleton_dimension_while_scalar_index_drops_it() { + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("singleton_slice.mo"), + 10, + 13, + ); + let layout = VarLayout::default(); + let functions = IndexMap::new(); + let mut variables = dae::DaeVariables::default(); + variables.parameters.insert( + rumoca_core::VarName::new("buf"), + dae::Variable { + dims: vec![2], + ..dae::Variable::new(rumoca_core::VarName::new("buf"), span) + }, + ); + let builder = LowerBuilder::new_with_metadata( + &layout, + &functions, + LowerBuilderMetadata { + dae_variables: Some(&variables), + ..LowerBuilderMetadata::default() + }, + ); + let range_slice = rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("buf"), + subscripts: vec![ + rumoca_core::Subscript::try_generated_expr( + Box::new(rumoca_core::Expression::Range { + start: Box::new(literal_i64(1, span)), + step: None, + end: Box::new(literal_i64(1, span)), + span, + }), + span, + "singleton range test subscript", + ) + .expect("range subscript should build"), + ], + span, + }; + let scalar_index = rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("buf"), + subscripts: vec![rumoca_core::Subscript::index(1, span)], + span, + }; + + assert_eq!( + builder + .infer_expr_dims(&range_slice, &Scope::new()) + .expect("range slice dims should infer"), + vec![1] + ); + assert_eq!( + builder + .infer_expr_dims(&scalar_index, &Scope::new()) + .expect("scalar index dims should infer"), + Vec::::new() + ); + } + + #[test] + fn range_dims_use_compile_time_size_expression_end() { + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("range_size_end.mo"), + 10, + 40, + ); + let layout = VarLayout::default(); + let functions = IndexMap::new(); + let builder = lower_builder(&layout, &functions); + let fill = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Fill, + args: vec![ + literal_i64(0, span), + literal_i64(0, span), + literal_i64(2, span), + ], + span, + }; + let size = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + args: vec![fill, literal_i64(2, span)], + span, + }; + let range = rumoca_core::Expression::Range { + start: Box::new(literal_i64(2, span)), + step: None, + end: Box::new(size.clone()), + span, + }; + + assert_eq!( + builder + .eval_compile_time_range_values( + &literal_i64(2, span), + None, + &size, + span, + &IndexMap::new(), + "range dimension inference", + ) + .expect("compile-time size range values should evaluate"), + vec![2] + ); + let range_size = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + args: vec![range, literal_i64(1, span)], + span, + }; + assert_eq!( + builder + .eval_compile_time_expr(&range_size, &IndexMap::new()) + .expect("size of compile-time range should evaluate"), + 1.0 + ); + let literal_size = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + args: vec![ + rumoca_core::Expression::Array { + elements: vec![literal_i64(0, span)], + is_matrix: false, + span, + }, + literal_i64(1, span), + ], + span, + }; + let max_size = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Max, + args: vec![rumoca_core::Expression::Array { + elements: vec![ + rumoca_core::Expression::Array { + elements: vec![range_size], + is_matrix: true, + span, + }, + rumoca_core::Expression::Array { + elements: vec![literal_size], + is_matrix: true, + span, + }, + ], + is_matrix: true, + span, + }], + span, + }; + let ones = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Ones, + args: vec![max_size], + span, + }; + assert_eq!( + builder + .infer_expr_dims(&ones, &Scope::new()) + .expect("ones(size(range, 1)) dims should infer"), + vec![1] + ); + } + #[test] fn eval_compile_time_positive_index_reports_expr_span() { let span = rumoca_core::Span::from_offsets( diff --git a/crates/rumoca-phase-solve/src/lower/array_values/references.rs b/crates/rumoca-phase-solve/src/lower/array_values/references.rs index f0d9db7d6..e61436aa8 100644 --- a/crates/rumoca-phase-solve/src/lower/array_values/references.rs +++ b/crates/rumoca-phase-solve/src/lower/array_values/references.rs @@ -24,11 +24,35 @@ impl<'a> LowerBuilder<'a> { if let Some(reg) = scope.get(&generated_key).copied() { return Ok(vec![reg]); } - let key = name.as_str(); + if (rumoca_core::parse_scalar_name(key).is_some() || self.layout.shape(key).is_none()) + && let Some(slot) = self.layout.binding(key) + { + return Ok(vec![self.emit_slot_load(slot, span)?]); + } if let Some(values) = self.lower_record_field_array_values(key, span)? { return Ok(values); } + if let Some(mut re_values) = + self.lower_record_field_array_values(format!("{key}.re").as_str(), span)? + && let Some(im_values) = + self.lower_record_field_array_values(format!("{key}.im").as_str(), span)? + { + re_values.extend(im_values); + return Ok(re_values); + } + if let Some(values) = self.lower_direct_assignment_values_for_key(key, scope, call_depth)? { + return Ok(values); + } + if self.value_mode == ValueMode::Pre + && let Some(pre_key) = self.pre_mode_base_key(key) + && let Some(values) = self.lower_indexed_binding_values_at(pre_key.as_str(), span)? + { + return Ok(values); + } + if let Some(values) = self.lower_indexed_binding_values_at(key, span)? { + return Ok(values); + } let key_path = self.scope_key_from_reference(name, span)?; if let Some(reg) = scope.get(&key_path).copied() @@ -57,9 +81,6 @@ impl<'a> LowerBuilder<'a> { { return Ok(values); } - if let Some(values) = self.lower_direct_assignment_values_for_key(key, scope, call_depth)? { - return Ok(values); - } if let Some(values) = self.lower_indexed_binding_values_for_resolved_key(key, &key_path, span)? { @@ -92,7 +113,7 @@ impl<'a> LowerBuilder<'a> { .contains_key(super::size_binding_key(key, 1).as_str()) } - fn lower_record_field_array_values( + pub(in crate::lower) fn lower_record_field_array_values( &mut self, key: &str, span: rumoca_core::Span, diff --git a/crates/rumoca-phase-solve/src/lower/array_values/selection_helpers.rs b/crates/rumoca-phase-solve/src/lower/array_values/selection_helpers.rs index 17efccd5b..c392e9bf3 100644 --- a/crates/rumoca-phase-solve/src/lower/array_values/selection_helpers.rs +++ b/crates/rumoca-phase-solve/src/lower/array_values/selection_helpers.rs @@ -1,4 +1,5 @@ use super::*; +use crate::lower::function_calls::decode_named_function_arg; pub(in crate::lower) fn matmul_shape_from_dims( lhs_dims: &[usize], @@ -62,6 +63,22 @@ pub(in crate::lower) fn projected_record_field_expression( span, }) } + rumoca_core::Expression::FunctionCall { args, .. } => { + for arg in args { + if let Some((name, value)) = decode_named_function_arg(arg) + && name == field + { + return Ok(value.clone()); + } + } + Ok(rumoca_core::Expression::FieldAccess { + base: Box::new(value.clone()), + field: field.to_string(), + span: value + .require_span("projected record field expression")? + .span(), + }) + } _ => Ok(rumoca_core::Expression::FieldAccess { base: Box::new(value.clone()), field: field.to_string(), @@ -97,10 +114,24 @@ pub(in crate::lower) fn indexed_record_field_key_indices( base_key: &str, field: &str, ) -> Option> { - let suffix = format!(".{field}"); - let indexed_base_key = key.strip_suffix(suffix.as_str())?; - let (candidate_base, indices) = parse_indexed_binding_key(indexed_base_key)?; - (candidate_base == base_key).then_some(indices) + let requested_key = format!("{base_key}.{field}"); + let mut candidate = key; + let mut fields = Vec::new(); + while let Some((prefix, candidate_field)) = split_last_binding_key_segment(candidate) { + fields.push(candidate_field); + if let Some((candidate_base, indices)) = parse_indexed_binding_key(prefix) { + fields.reverse(); + let candidate_key = format!("{}.{}", candidate_base, fields.join(".")); + return (candidate_key == requested_key).then_some(indices); + } + candidate = prefix; + } + None +} + +fn split_last_binding_key_segment(key: &str) -> Option<(&str, &str)> { + let dot = key.rfind('.')?; + Some((&key[..dot], &key[dot + 1..])) } /// Cartesian product of per-dimension index selections into one-based diff --git a/crates/rumoca-phase-solve/src/lower/array_values/tests.rs b/crates/rumoca-phase-solve/src/lower/array_values/tests.rs index 2b1bf9c23..4c3d9464f 100644 --- a/crates/rumoca-phase-solve/src/lower/array_values/tests.rs +++ b/crates/rumoca-phase-solve/src/lower/array_values/tests.rs @@ -1,4 +1,5 @@ use super::*; +use crate::build_var_layout; fn unspanned_array_values_test_span() -> rumoca_core::Span { rumoca_core::Span::DUMMY @@ -231,6 +232,88 @@ fn lower_array_like_values_lowers_modelica_array_constructor_call() -> Result<() Ok(()) } +#[test] +fn slice_binding_keys_consume_singleton_projection_after_scalar_selection() -> Result<(), LowerError> +{ + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("array_slice_singleton_projection.mo"), + 8, + 16, + ); + let mut dae_model = dae::Dae::new(); + dae_model + .variables + .outputs + .insert(rumoca_core::VarName::new("y"), { + let mut variable = dae::Variable { + name: rumoca_core::VarName::new("y"), + dims: vec![3], + ..rumoca_ir_dae::Variable::empty_with_span(span) + }; + variable.origin = dae::VariableOrigin::Generated; + variable + }); + let layout = build_var_layout(&dae_model)?; + let functions = IndexMap::new(); + let builder = LowerBuilder::new(&layout, &functions); + let keys = builder + .slice_binding_keys( + "y", + &[ + rumoca_core::Subscript::index(2, span), + rumoca_core::Subscript::index(1, span), + ], + span, + &Scope::new(), + )? + .expect("declared scalar selection followed by singleton projection should resolve"); + + assert_eq!(keys, vec!["y[2]"]); + Ok(()) +} + +#[test] +fn slice_binding_keys_reject_non_singleton_extra_projection() -> Result<(), LowerError> { + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("array_slice_non_singleton_projection.mo"), + 8, + 16, + ); + let mut dae_model = dae::Dae::new(); + dae_model + .variables + .outputs + .insert(rumoca_core::VarName::new("y"), { + let mut variable = dae::Variable { + name: rumoca_core::VarName::new("y"), + dims: vec![3], + ..rumoca_ir_dae::Variable::empty_with_span(span) + }; + variable.origin = dae::VariableOrigin::Generated; + variable + }); + let layout = build_var_layout(&dae_model)?; + let functions = IndexMap::new(); + let builder = LowerBuilder::new(&layout, &functions); + let err = builder + .slice_binding_keys( + "y", + &[ + rumoca_core::Subscript::index(2, span), + rumoca_core::Subscript::index(2, span), + ], + span, + &Scope::new(), + ) + .expect_err("non-singleton extra projection must remain a shape error"); + + assert_eq!( + err.reason(), + "array slice has more subscripts than dimensions" + ); + Ok(()) +} + #[test] fn lower_min_max_builtin_rejects_empty_args_without_dummy_span() -> Result<(), LowerError> { let layout = VarLayout::default(); diff --git a/crates/rumoca-phase-solve/src/lower/builtin_methods.rs b/crates/rumoca-phase-solve/src/lower/builtin_methods.rs index 3c11f8f06..f30fdef37 100644 --- a/crates/rumoca-phase-solve/src/lower/builtin_methods.rs +++ b/crates/rumoca-phase-solve/src/lower/builtin_methods.rs @@ -57,7 +57,24 @@ impl<'a> LowerBuilder<'a> { rumoca_core::BuiltinFunction::Sample, call_span, )), - [value] => self.lower_clocked_sample_value(value, call_span, scope, call_depth), + [value] => { + if let Some((phase_seconds, period_seconds, schedule_span)) = + self.current_update_target_clock_timing().map(|timing| { + ( + timing.phase_seconds, + timing.period_seconds, + timing.source_span, + ) + }) + { + let phase = self.emit_const_at(phase_seconds, schedule_span)?; + let period = self.emit_const_at(period_seconds, schedule_span)?; + let tick = self.emit_periodic_tick(phase, period, schedule_span)?; + self.lower_clocked_sample_with_tick(value, tick, call_span, scope, call_depth) + } else { + self.lower_clocked_sample_value(value, call_span, scope, call_depth) + } + } [_internal_id, start, interval, ..] => { if self.value_mode == ValueMode::Pre { return self.emit_const_at( @@ -296,22 +313,77 @@ impl<'a> LowerBuilder<'a> { span, )); }; + if let rumoca_core::Expression::FunctionCall { + name, + args: passthrough_args, + is_constructor: false, + .. + } = base_expr + && is_stream_passthrough_intrinsic(name.as_str()) + && let Some(arg) = passthrough_args.first() + { + let mut forwarded_args = Vec::with_capacity(args.len()); + forwarded_args.push(arg.clone()); + forwarded_args.extend(args.iter().skip(1).cloned()); + return self.lower_size_builtin(&forwarded_args, span, scope, call_depth); + } let base_span = expression_or_call_span(base_expr, span, "size() base argument")?; - let inferred_dims = self.infer_expr_dims(base_expr, scope)?; + if let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = base_expr + && subscripts.is_empty() + && let Some(dims) = self.local_binding_dims.get(name.as_str()) + && !dims.is_empty() + && dims.iter().all(|dim| *dim > 0) + { + let dims = dims + .iter() + .copied() + .map(usize::try_from) + .collect::, _>>() + .map_err(|_| { + LowerError::contract_violation( + "local size() dimension is outside host range", + base_span, + ) + })?; + return self.lower_size_from_dims(&dims, args, base_span, scope, call_depth); + } + let inferred_dims = if expr_component_reference_missing_def_id(base_expr) { + Vec::new() + } else { + self.infer_expr_dims(base_expr, scope)? + }; if !inferred_dims.is_empty() && inferred_dims.iter().all(|dim| *dim > 0) { return self.lower_size_from_dims(&inferred_dims, args, base_span, scope, call_depth); } + if let Ok(value) = self.eval_compile_time_size(args, span, &self.local_const_bindings) { + return self.emit_const_at(value, base_span); + } - let base_key = - dynamic_binding_base_key(base_expr).map_err(|err| err.with_fallback_span(base_span))?; + let base_key = match dynamic_binding_base_key(base_expr) { + Ok(base_key) => base_key, + Err(LowerError::DynamicBindingBase { .. }) => { + return self.emit_const_at(1.0, base_span); + } + Err(err) => return Err(err.with_fallback_span(base_span)), + }; - let source_key = component_reference_key_for_expr(base_expr)?; let generated_key = ComponentReferenceKey::generated(&base_key); + let generated_entries = self.indexed_bindings.get(&generated_key); + let source_key = + if generated_entries.is_some() || expr_component_reference_missing_def_id(base_expr) { + None + } else { + component_reference_key_for_expr(base_expr)? + }; let dims = infer_indexed_dims( - source_key - .as_ref() - .and_then(|key| self.indexed_bindings.get(key)) - .or_else(|| self.indexed_bindings.get(&generated_key)) + generated_entries + .or_else(|| { + source_key + .as_ref() + .and_then(|key| self.indexed_bindings.get(key)) + }) .map(Vec::as_slice) .unwrap_or(&[]), ); @@ -322,3 +394,16 @@ impl<'a> LowerBuilder<'a> { self.lower_size_from_dims(&dims, args, base_span, scope, call_depth) } } + +fn expr_component_reference_missing_def_id(expr: &rumoca_core::Expression) -> bool { + match expr { + rumoca_core::Expression::VarRef { name, .. } => name + .component_ref() + .is_some_and(|component_ref| component_ref.def_id.is_none()), + rumoca_core::Expression::Index { base, .. } + | rumoca_core::Expression::FieldAccess { base, .. } => { + expr_component_reference_missing_def_id(base) + } + _ => false, + } +} diff --git a/crates/rumoca-phase-solve/src/lower/clock.rs b/crates/rumoca-phase-solve/src/lower/clock.rs index f4ffe422b..644aa5ac0 100644 --- a/crates/rumoca-phase-solve/src/lower/clock.rs +++ b/crates/rumoca-phase-solve/src/lower/clock.rs @@ -352,7 +352,7 @@ impl<'a> LowerBuilder<'a> { } } - fn current_update_target_clock_timing(&self) -> Option<&dae::ClockSchedule> { + pub(super) fn current_update_target_clock_timing(&self) -> Option<&dae::ClockSchedule> { let target = self.current_update_target?; self.clock_timings?.iter().find_map(|(name, timing)| { (self.layout.binding(name.as_str()) == Some(target)).then_some(timing) diff --git a/crates/rumoca-phase-solve/src/lower/compile_time.rs b/crates/rumoca-phase-solve/src/lower/compile_time.rs index aee093dae..df980489e 100644 --- a/crates/rumoca-phase-solve/src/lower/compile_time.rs +++ b/crates/rumoca-phase-solve/src/lower/compile_time.rs @@ -1,17 +1,81 @@ use indexmap::IndexMap; use rumoca_ir_dae as dae; +use std::sync::Arc; -use super::{LowerError, helpers::variable_size, size_binding_key}; +use super::{ + LowerError, function_calls::external_table_intrinsic_kind, helpers::variable_size, + size_binding_key, +}; pub(super) fn structural_bindings( dae_model: &dae::Dae, ) -> Result, LowerError> { + structural_bindings_with_eval_env(dae_model).map(|(bindings, _)| bindings) +} + +pub(super) fn external_table_data( + dae_model: &dae::Dae, +) -> Result, LowerError> { + let (_, eval_env) = structural_bindings_with_eval_env(dae_model)?; + Ok(rumoca_eval_dae::all_external_table_data_in_env(&eval_env)) +} + +pub(in crate::lower) fn eval_selected_function_output( + dae_model: &dae::Dae, + function_name: &rumoca_core::VarName, + output_name: &str, + indices: &[i64], + args: &[rumoca_core::Expression], +) -> Option { + let env = compile_time_eval_env(dae_model); + rumoca_eval_dae::eval_selected_function_output_pub::( + function_name, + output_name, + indices, + args, + &env, + ) + .ok() +} + +fn structural_bindings_with_eval_env( + dae_model: &dae::Dae, +) -> Result<(IndexMap, rumoca_eval_dae::VarEnv), LowerError> { + let mut eval_env = compile_time_eval_env(dae_model); let mut bindings = enum_literal_bindings(&dae_model.symbols.enum_literal_ordinals); let shapes = variable_shapes(dae_model); insert_shape_bindings(&mut bindings, &shapes); - insert_constant_variables(&mut bindings, dae_model, &shapes)?; - insert_structural_parameters(&mut bindings, dae_model, &shapes)?; - Ok(bindings) + insert_constant_variables(&mut bindings, dae_model, &shapes, &mut eval_env)?; + seed_compile_time_start_values(dae_model, &mut eval_env); + insert_structural_parameters(&mut bindings, dae_model, &shapes, &mut eval_env)?; + insert_static_initial_assignments(&mut bindings, dae_model); + insert_external_table_handle_bindings(&mut bindings, dae_model, &shapes, &mut eval_env)?; + Ok((bindings, eval_env)) +} + +fn compile_time_eval_env(dae_model: &dae::Dae) -> rumoca_eval_dae::VarEnv { + let mut env = rumoca_eval_dae::VarEnv::new(); + env.runtime = Arc::new(rumoca_eval_dae::EvalRuntimeState::new()); + env.functions = Arc::new(rumoca_eval_dae::collect_user_functions(dae_model)); + env.dims = Arc::new(rumoca_eval_dae::collect_var_dims(dae_model)); + env.start_exprs = Arc::new(rumoca_eval_dae::collect_var_starts(dae_model)); + env.nonnumeric_names = Arc::new( + dae_model + .metadata + .nonnumeric_variable_names + .iter() + .cloned() + .collect(), + ); + env.clock_intervals = Arc::new(dae_model.clocks.intervals.clone()); + env.enum_literal_ordinals = Arc::new(dae_model.symbols.enum_literal_ordinals.clone()); + for &(name, value) in rumoca_eval_dae::MODELICA_CONSTANTS { + env.set(name, value); + } + for &(name, value) in rumoca_eval_dae::MODELICA_COMPLEX_CONSTANTS { + env.set(name, value); + } + env } fn enum_literal_bindings(ordinals: &IndexMap) -> IndexMap { @@ -59,9 +123,10 @@ fn insert_constant_variables( bindings: &mut IndexMap, dae_model: &dae::Dae, shapes: &IndexMap>, + eval_env: &mut rumoca_eval_dae::VarEnv, ) -> Result<(), LowerError> { for (name, var) in &dae_model.variables.constants { - insert_variable_start_bindings(bindings, shapes, name.as_str(), var)?; + insert_variable_start_bindings(bindings, shapes, eval_env, name.as_str(), var)?; } Ok(()) } @@ -70,12 +135,13 @@ fn insert_structural_parameters( bindings: &mut IndexMap, dae_model: &dae::Dae, shapes: &IndexMap>, + eval_env: &mut rumoca_eval_dae::VarEnv, ) -> Result<(), LowerError> { for _ in 0..dae_model.variables.parameters.len().max(1) { let before = bindings.len(); for (name, var) in &dae_model.variables.parameters { if !var.is_tunable { - insert_variable_start_bindings(bindings, shapes, name.as_str(), var)?; + insert_variable_start_bindings(bindings, shapes, eval_env, name.as_str(), var)?; } } if bindings.len() == before { @@ -85,20 +151,426 @@ fn insert_structural_parameters( Ok(()) } +fn insert_external_table_handle_bindings( + bindings: &mut IndexMap, + dae_model: &dae::Dae, + shapes: &IndexMap>, + eval_env: &mut rumoca_eval_dae::VarEnv, +) -> Result<(), LowerError> { + seed_compile_time_start_values(dae_model, eval_env); + for table_id_name in &dae_model.metadata.nonnumeric_variable_names { + if !table_id_name.ends_with(".tableID") { + continue; + } + let Some(prefix) = table_id_name.strip_suffix(".tableID") else { + continue; + }; + let constructor = external_table_record_constructor(prefix, &dae_model.variables) + .or_else(|| { + external_table_constructor_for_prefix(prefix, dae_model.variables.constants.iter()) + .cloned() + }) + .or_else(|| { + external_table_constructor_for_prefix(prefix, dae_model.variables.parameters.iter()) + .cloned() + }); + let Some(constructor) = constructor else { + continue; + }; + let Some(values) = eval_values(&constructor, bindings, shapes, eval_env) else { + continue; + }; + let Some(table_id) = values.first().copied() else { + continue; + }; + bindings.insert(table_id_name.clone(), table_id); + eval_env.set(table_id_name, table_id); + } + Ok(()) +} + +fn insert_static_initial_assignments(bindings: &mut IndexMap, dae_model: &dae::Dae) { + for _ in 0..dae_model.initialization.equations.len().max(1) { + let before = bindings.len(); + for equation in &dae_model.initialization.equations { + let Some(lhs) = equation.lhs.as_ref() else { + continue; + }; + if !initial_assignment_target_is_structural(lhs.as_str(), &dae_model.variables) { + continue; + } + let Ok(value) = eval_static_initial_numeric(&equation.rhs, bindings) else { + continue; + }; + bindings.insert(lhs.as_str().to_string(), value); + } + if bindings.len() == before { + break; + } + } +} + +fn initial_assignment_target_is_structural(name: &str, variables: &dae::DaeVariables) -> bool { + let var_name = rumoca_core::VarName::new(name); + [ + variables.parameters.get(&var_name), + variables.constants.get(&var_name), + variables.discrete_valued.get(&var_name), + variables.discrete_reals.get(&var_name), + ] + .into_iter() + .flatten() + .any(|variable| !variable.is_tunable) +} + +fn eval_static_initial_numeric( + expr: &rumoca_core::Expression, + bindings: &IndexMap, +) -> Result { + match expr { + rumoca_core::Expression::Literal { value, .. } => match value { + rumoca_core::Literal::Real(value) => Ok(*value), + rumoca_core::Literal::Integer(value) => Ok(*value as f64), + rumoca_core::Literal::Boolean(value) => Ok(if *value { 1.0 } else { 0.0 }), + rumoca_core::Literal::String(_) => Err(()), + }, + rumoca_core::Expression::VarRef { + name, subscripts, .. + } if subscripts.is_empty() => bindings.get(name.as_str()).copied().ok_or(()), + rumoca_core::Expression::Unary { op, rhs, .. } => { + let value = eval_static_initial_numeric(rhs, bindings)?; + match op { + rumoca_core::OpUnary::Minus | rumoca_core::OpUnary::DotMinus => Ok(-value), + rumoca_core::OpUnary::Plus + | rumoca_core::OpUnary::DotPlus + | rumoca_core::OpUnary::Empty => Ok(value), + rumoca_core::OpUnary::Not => Ok(if value == 0.0 { 1.0 } else { 0.0 }), + } + } + rumoca_core::Expression::Binary { op, lhs, rhs, .. } => { + let lhs = eval_static_initial_numeric(lhs, bindings)?; + let rhs = eval_static_initial_numeric(rhs, bindings)?; + match op { + rumoca_core::OpBinary::Add | rumoca_core::OpBinary::AddElem => Ok(lhs + rhs), + rumoca_core::OpBinary::Sub | rumoca_core::OpBinary::SubElem => Ok(lhs - rhs), + rumoca_core::OpBinary::Mul | rumoca_core::OpBinary::MulElem => Ok(lhs * rhs), + rumoca_core::OpBinary::Div | rumoca_core::OpBinary::DivElem => Ok(lhs / rhs), + rumoca_core::OpBinary::Eq => Ok(if (lhs - rhs).abs() < f64::EPSILON { + 1.0 + } else { + 0.0 + }), + rumoca_core::OpBinary::Neq => Ok(if (lhs - rhs).abs() >= f64::EPSILON { + 1.0 + } else { + 0.0 + }), + rumoca_core::OpBinary::Lt => Ok(if lhs < rhs { 1.0 } else { 0.0 }), + rumoca_core::OpBinary::Le => Ok(if lhs <= rhs { 1.0 } else { 0.0 }), + rumoca_core::OpBinary::Gt => Ok(if lhs > rhs { 1.0 } else { 0.0 }), + rumoca_core::OpBinary::Ge => Ok(if lhs >= rhs { 1.0 } else { 0.0 }), + rumoca_core::OpBinary::And => Ok(if lhs != 0.0 && rhs != 0.0 { 1.0 } else { 0.0 }), + rumoca_core::OpBinary::Or => Ok(if lhs != 0.0 || rhs != 0.0 { 1.0 } else { 0.0 }), + _ => Err(()), + } + } + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + for (condition, value) in branches { + if eval_static_initial_numeric(condition, bindings)? != 0.0 { + return eval_static_initial_numeric(value, bindings); + } + } + eval_static_initial_numeric(else_branch, bindings) + } + rumoca_core::Expression::FunctionCall { name, args, .. } + if matches!( + name.as_str(), + "Modelica.Utilities.Strings.isEqual" | "Strings.isEqual" | "isEqual" + ) => + { + let left = args.first().ok_or(())?; + let right = args.get(1).ok_or(())?; + Ok( + if eval_static_initial_string(left, bindings)? + == eval_static_initial_string(right, bindings)? + { + 1.0 + } else { + 0.0 + }, + ) + } + _ => Err(()), + } +} + +fn eval_static_initial_string( + expr: &rumoca_core::Expression, + bindings: &IndexMap, +) -> Result { + match expr { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String(value), + .. + } => Ok(value.clone()), + rumoca_core::Expression::Index { + base, subscripts, .. + } => { + let [subscript] = subscripts.as_slice() else { + return Err(()); + }; + let index = match subscript { + rumoca_core::Subscript::Index { value, .. } if *value > 0 => *value as usize, + rumoca_core::Subscript::Expr { expr, .. } => { + let value = eval_static_initial_numeric(expr, bindings)?; + if value.fract().abs() > f64::EPSILON || value <= 0.0 { + return Err(()); + } + value as usize + } + _ => return Err(()), + }; + let rumoca_core::Expression::Array { elements, .. } = base.as_ref() else { + return Err(()); + }; + eval_static_initial_string(elements.get(index - 1).ok_or(())?, bindings) + } + _ => Err(()), + } +} + +fn seed_compile_time_start_values( + dae_model: &dae::Dae, + eval_env: &mut rumoca_eval_dae::VarEnv, +) { + for _ in 0..dae_model.variables.parameters.len().max(1).clamp(1, 8) { + let before = eval_env.vars.len(); + for (name, var) in dae_model + .variables + .constants + .iter() + .chain(dae_model.variables.parameters.iter()) + { + seed_compile_time_start_value(name.as_str(), var, eval_env); + } + if eval_env.vars.len() == before { + break; + } + } +} + +fn seed_compile_time_start_value( + name: &str, + var: &dae::Variable, + eval_env: &mut rumoca_eval_dae::VarEnv, +) { + let Some(start) = var.start.as_ref() else { + return; + }; + if rumoca_eval_dae::start_expr_is_nonnumeric(start, eval_env) { + return; + } + if !var.dims.is_empty() { + let values = if var.size() == 0 && var.dims.len() >= 2 { + rumoca_eval_dae::eval_matrix_values::(start, eval_env) + .ok() + .flatten() + .map(|matrix| matrix.into_iter().flatten().collect()) + .or_else(|| rumoca_eval_dae::eval_array_values::(start, eval_env).ok()) + } else { + rumoca_eval_dae::eval_array_values::(start, eval_env).ok() + }; + if let Some(values) = values { + rumoca_eval_dae::set_array_entries(eval_env, name, &var.dims, &values); + } + return; + } + if let Ok(value) = rumoca_eval_dae::eval_expr::(start, eval_env) { + eval_env.set(name, value); + } +} + +fn external_table_record_constructor( + prefix: &str, + variables: &dae::DaeVariables, +) -> Option { + external_table_record_constructor_from_fields( + "Modelica.Blocks.Types.ExternalCombiTimeTable", + prefix, + variables, + &[ + "tableName", + "fileName", + "table", + "startTime", + "columns", + "smoothness", + "extrapolation", + "shiftTime", + "timeEvents", + "verboseRead", + "delimiter", + "nHeaderLines", + ], + ) + .or_else(|| { + external_table_record_constructor_from_fields( + "Modelica.Blocks.Types.ExternalCombiTable1D", + prefix, + variables, + &[ + "tableName", + "fileName", + "table", + "columns", + "smoothness", + "extrapolation", + "verboseRead", + "delimiter", + "nHeaderLines", + ], + ) + }) +} + +fn external_table_record_constructor_from_fields( + constructor_name: &str, + prefix: &str, + variables: &dae::DaeVariables, + fields: &[&str], +) -> Option { + let args = fields + .iter() + .map(|field| external_table_record_field_ref(prefix, field, variables)) + .collect::>>()?; + let span = args.first().and_then(rumoca_core::Expression::span)?; + Some(rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new(constructor_name), + args, + is_constructor: true, + span, + }) +} + +fn external_table_record_field_ref( + prefix: &str, + field: &str, + variables: &dae::DaeVariables, +) -> Option { + let name = format!("{prefix}.{field}"); + let var_name = rumoca_core::VarName::new(name.as_str()); + let span = variables + .parameters + .get(&var_name) + .or_else(|| variables.constants.get(&var_name)) + .map(|var| var.source_span)?; + Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new(name.as_str()), + subscripts: Vec::new(), + span, + }) +} + +fn external_table_constructor_for_prefix<'a>( + prefix: &str, + variables: impl Iterator, +) -> Option<&'a rumoca_core::Expression> { + let prefix = format!("{prefix}."); + variables + .filter(|(name, var)| !var.is_tunable && name.as_str().starts_with(prefix.as_str())) + .filter_map(|(_, var)| var.start.as_ref()) + .find_map(find_external_table_constructor) +} + +fn find_external_table_constructor( + expr: &rumoca_core::Expression, +) -> Option<&rumoca_core::Expression> { + match expr { + rumoca_core::Expression::FunctionCall { name, args, .. } => { + if is_external_table_constructor(name) { + return Some(expr); + } + args.iter().find_map(find_external_table_constructor) + } + rumoca_core::Expression::Unary { rhs, .. } => find_external_table_constructor(rhs), + rumoca_core::Expression::Binary { lhs, rhs, .. } => { + find_external_table_constructor(lhs).or_else(|| find_external_table_constructor(rhs)) + } + rumoca_core::Expression::BuiltinCall { args, .. } + | rumoca_core::Expression::Array { elements: args, .. } + | rumoca_core::Expression::Tuple { elements: args, .. } => { + args.iter().find_map(find_external_table_constructor) + } + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => branches + .iter() + .find_map(|(condition, value)| { + find_external_table_constructor(condition) + .or_else(|| find_external_table_constructor(value)) + }) + .or_else(|| find_external_table_constructor(else_branch)), + rumoca_core::Expression::Index { + base, subscripts, .. + } => find_external_table_constructor(base).or_else(|| { + subscripts.iter().find_map(|subscript| match subscript { + rumoca_core::Subscript::Expr { expr, .. } => find_external_table_constructor(expr), + rumoca_core::Subscript::Index { .. } | rumoca_core::Subscript::Colon { .. } => None, + }) + }), + rumoca_core::Expression::FieldAccess { base, .. } => find_external_table_constructor(base), + rumoca_core::Expression::Range { + start, step, end, .. + } => find_external_table_constructor(start) + .or_else(|| step.as_deref().and_then(find_external_table_constructor)) + .or_else(|| find_external_table_constructor(end)), + rumoca_core::Expression::ArrayComprehension { + expr, + indices, + filter, + .. + } => find_external_table_constructor(expr) + .or_else(|| { + indices + .iter() + .find_map(|index| find_external_table_constructor(&index.range)) + }) + .or_else(|| filter.as_deref().and_then(find_external_table_constructor)), + rumoca_core::Expression::Literal { .. } + | rumoca_core::Expression::VarRef { .. } + | rumoca_core::Expression::Empty { .. } => None, + } +} + +fn is_external_table_constructor(name: &rumoca_core::Reference) -> bool { + matches!( + name.last_segment(), + "ExternalCombiTimeTable" | "ExternalCombiTable1D" + ) +} + fn insert_variable_start_bindings( bindings: &mut IndexMap, shapes: &IndexMap>, + eval_env: &mut rumoca_eval_dae::VarEnv, name: &str, var: &dae::Variable, ) -> Result<(), LowerError> { let Some(start) = var.start.as_ref() else { return Ok(()); }; - let Some(raw_values) = eval_values(start, bindings, shapes) else { + let Some(raw_values) = eval_values(start, bindings, shapes, eval_env) else { return Ok(()); }; let values = expand_values_to_size(raw_values, variable_size(var)?, name, var.source_span)?; insert_scalarized_bindings(bindings, name, &var.dims, &values); + sync_bindings_to_eval_env(eval_env, bindings); Ok(()) } @@ -121,6 +593,7 @@ fn eval_values( expr: &rumoca_core::Expression, bindings: &IndexMap, shapes: &IndexMap>, + eval_env: &mut rumoca_eval_dae::VarEnv, ) -> Option> { match expr { rumoca_core::Expression::Literal { value: literal, .. } => { @@ -129,37 +602,70 @@ fn eval_values( rumoca_core::Expression::VarRef { name, subscripts, .. } => { - let key = var_key(name, subscripts, bindings)?; + let key = var_key(name, subscripts, bindings, eval_env)?; bindings.get(key.as_str()).copied().map(|value| vec![value]) } - rumoca_core::Expression::Unary { op, rhs, .. } => eval_unary(op, rhs, bindings, shapes), + rumoca_core::Expression::Unary { op, rhs, .. } => { + eval_unary(op, rhs, bindings, shapes, eval_env) + } rumoca_core::Expression::Binary { op, lhs, rhs, .. } => { - eval_binary(op, lhs, rhs, bindings, shapes) + eval_binary(op, lhs, rhs, bindings, shapes, eval_env) } rumoca_core::Expression::BuiltinCall { function, args, .. } => { - eval_builtin(*function, args, bindings, shapes) + eval_builtin(*function, args, bindings, shapes, eval_env) + } + rumoca_core::Expression::FunctionCall { name, .. } if is_external_table_function(name) => { + eval_external_table_function(expr, bindings, eval_env) } rumoca_core::Expression::Array { elements, .. } | rumoca_core::Expression::Tuple { elements, .. } => { let mut values = Vec::new(); for element in elements { - values.extend(eval_values(element, bindings, shapes)?); + values.extend(eval_values(element, bindings, shapes, eval_env)?); } Some(values) } rumoca_core::Expression::Range { start, step, end, .. - } => eval_range(start, step.as_deref(), end, bindings, shapes), + } => eval_range(start, step.as_deref(), end, bindings, shapes, eval_env), _ => None, } } +fn is_external_table_function(name: &rumoca_core::Reference) -> bool { + matches!( + name.last_segment(), + "ExternalCombiTimeTable" | "ExternalCombiTable1D" + ) || external_table_intrinsic_kind(name.as_str()).is_some() +} + +fn eval_external_table_function( + expr: &rumoca_core::Expression, + bindings: &IndexMap, + eval_env: &mut rumoca_eval_dae::VarEnv, +) -> Option> { + sync_bindings_to_eval_env(eval_env, bindings); + rumoca_eval_dae::eval_expr::(expr, eval_env) + .ok() + .map(|value| vec![value]) +} + +fn sync_bindings_to_eval_env( + eval_env: &mut rumoca_eval_dae::VarEnv, + bindings: &IndexMap, +) { + for (name, value) in bindings { + eval_env.set(name, *value); + } +} + fn eval_scalar( expr: &rumoca_core::Expression, bindings: &IndexMap, shapes: &IndexMap>, + eval_env: &mut rumoca_eval_dae::VarEnv, ) -> Option { - let values = eval_values(expr, bindings, shapes)?; + let values = eval_values(expr, bindings, shapes, eval_env)?; (values.len() == 1).then_some(values[0]) } @@ -168,8 +674,9 @@ fn eval_unary( rhs: &rumoca_core::Expression, bindings: &IndexMap, shapes: &IndexMap>, + eval_env: &mut rumoca_eval_dae::VarEnv, ) -> Option> { - let values = eval_values(rhs, bindings, shapes)?; + let values = eval_values(rhs, bindings, shapes, eval_env)?; match op { rumoca_core::OpUnary::Plus | rumoca_core::OpUnary::DotPlus => Some(values), rumoca_core::OpUnary::Minus | rumoca_core::OpUnary::DotMinus => { @@ -185,9 +692,10 @@ fn eval_binary( rhs: &rumoca_core::Expression, bindings: &IndexMap, shapes: &IndexMap>, + eval_env: &mut rumoca_eval_dae::VarEnv, ) -> Option> { - let lhs = eval_scalar(lhs, bindings, shapes)?; - let rhs = eval_scalar(rhs, bindings, shapes)?; + let lhs = eval_scalar(lhs, bindings, shapes, eval_env)?; + let rhs = eval_scalar(rhs, bindings, shapes, eval_env)?; let value = match op { rumoca_core::OpBinary::Add | rumoca_core::OpBinary::AddElem => lhs + rhs, rumoca_core::OpBinary::Sub | rumoca_core::OpBinary::SubElem => lhs - rhs, @@ -205,11 +713,12 @@ fn eval_range( end: &rumoca_core::Expression, bindings: &IndexMap, shapes: &IndexMap>, + eval_env: &mut rumoca_eval_dae::VarEnv, ) -> Option> { - let start = eval_scalar(start, bindings, shapes)?; - let end = eval_scalar(end, bindings, shapes)?; + let start = eval_scalar(start, bindings, shapes, eval_env)?; + let end = eval_scalar(end, bindings, shapes, eval_env)?; let step = step - .map(|expr| eval_scalar(expr, bindings, shapes)) + .map(|expr| eval_scalar(expr, bindings, shapes, eval_env)) .unwrap_or_else(|| Some(if end >= start { 1.0 } else { -1.0 }))?; if !start.is_finite() || !end.is_finite() || !step.is_finite() || step.abs() <= f64::EPSILON { return None; @@ -235,21 +744,24 @@ fn eval_builtin( args: &[rumoca_core::Expression], bindings: &IndexMap, shapes: &IndexMap>, + eval_env: &mut rumoca_eval_dae::VarEnv, ) -> Option> { use rumoca_core::BuiltinFunction as Builtin; match function { // MLS §10.3.1: size(A, i) is a structural property of the array type. - Builtin::Size => eval_size(args, bindings, shapes), - Builtin::Min => eval_min_max(args, bindings, shapes, f64::min), - Builtin::Max => eval_min_max(args, bindings, shapes, f64::max), - Builtin::NoEvent => eval_values(args.first()?, bindings, shapes), - Builtin::Smooth => eval_values(args.get(1)?, bindings, shapes), - Builtin::Homotopy => eval_values(args.first()?, bindings, shapes), - Builtin::Abs => unary_builtin(args, bindings, shapes, f64::abs), - Builtin::Sign => unary_builtin(args, bindings, shapes, f64::signum), - Builtin::Sqrt => unary_builtin(args, bindings, shapes, f64::sqrt), - Builtin::Floor | Builtin::Integer => unary_builtin(args, bindings, shapes, f64::floor), - Builtin::Ceil => unary_builtin(args, bindings, shapes, f64::ceil), + Builtin::Size => eval_size(args, bindings, shapes, eval_env), + Builtin::Min => eval_min_max(args, bindings, shapes, eval_env, f64::min), + Builtin::Max => eval_min_max(args, bindings, shapes, eval_env, f64::max), + Builtin::NoEvent => eval_values(args.first()?, bindings, shapes, eval_env), + Builtin::Smooth => eval_values(args.get(1)?, bindings, shapes, eval_env), + Builtin::Homotopy => eval_values(args.first()?, bindings, shapes, eval_env), + Builtin::Abs => unary_builtin(args, bindings, shapes, eval_env, f64::abs), + Builtin::Sign => unary_builtin(args, bindings, shapes, eval_env, f64::signum), + Builtin::Sqrt => unary_builtin(args, bindings, shapes, eval_env, f64::sqrt), + Builtin::Floor | Builtin::Integer => { + unary_builtin(args, bindings, shapes, eval_env, f64::floor) + } + Builtin::Ceil => unary_builtin(args, bindings, shapes, eval_env, f64::ceil), _ => None, } } @@ -258,13 +770,23 @@ fn eval_size( args: &[rumoca_core::Expression], bindings: &IndexMap, shapes: &IndexMap>, + eval_env: &mut rumoca_eval_dae::VarEnv, ) -> Option> { + if let Some(dims) = literal_array_shape(args.first()?) { + let Some(dim_expr) = args.get(1) else { + return Some(dims.into_iter().map(|dim| dim as f64).collect()); + }; + let dim = positive_usize_from_f64(eval_scalar(dim_expr, bindings, shapes, eval_env)?)?; + return dims + .get(dim.checked_sub(1)?) + .map(|value| vec![*value as f64]); + } let rumoca_core::Expression::VarRef { name, subscripts, .. } = args.first()? else { return Some(vec![ - eval_values(args.first()?, bindings, shapes)?.len() as f64 + eval_values(args.first()?, bindings, shapes, eval_env)?.len() as f64, ]); }; if !subscripts.is_empty() { @@ -274,20 +796,48 @@ fn eval_size( let Some(dim_expr) = args.get(1) else { return Some(dims.iter().map(|dim| *dim as f64).collect()); }; - let dim = positive_usize_from_f64(eval_scalar(dim_expr, bindings, shapes)?)?; + let dim = positive_usize_from_f64(eval_scalar(dim_expr, bindings, shapes, eval_env)?)?; dims.get(dim.checked_sub(1)?) .map(|value| vec![*value as f64]) } +fn literal_array_shape(expr: &rumoca_core::Expression) -> Option> { + match expr { + rumoca_core::Expression::Array { + elements, + is_matrix: false, + .. + } + | rumoca_core::Expression::Tuple { elements, .. } => Some(vec![elements.len()]), + rumoca_core::Expression::Array { + elements, + is_matrix: true, + .. + } => { + let first_row = elements.first()?; + let rumoca_core::Expression::Array { + elements: row_elements, + .. + } = first_row + else { + return Some(vec![elements.len()]); + }; + Some(vec![elements.len(), row_elements.len()]) + } + _ => None, + } +} + fn eval_min_max( args: &[rumoca_core::Expression], bindings: &IndexMap, shapes: &IndexMap>, + eval_env: &mut rumoca_eval_dae::VarEnv, op: fn(f64, f64) -> f64, ) -> Option> { let mut values = Vec::new(); for arg in args { - values.extend(eval_values(arg, bindings, shapes)?); + values.extend(eval_values(arg, bindings, shapes, eval_env)?); } let first = *values.first()?; Some(vec![values.into_iter().fold(first, op)]) @@ -297,15 +847,22 @@ fn unary_builtin( args: &[rumoca_core::Expression], bindings: &IndexMap, shapes: &IndexMap>, + eval_env: &mut rumoca_eval_dae::VarEnv, op: fn(f64) -> f64, ) -> Option> { - Some(vec![op(eval_scalar(args.first()?, bindings, shapes)?)]) + Some(vec![op(eval_scalar( + args.first()?, + bindings, + shapes, + eval_env, + )?)]) } fn var_key( name: &rumoca_core::Reference, subscripts: &[rumoca_core::Subscript], bindings: &IndexMap, + eval_env: &mut rumoca_eval_dae::VarEnv, ) -> Option { let name = name.as_str(); if subscripts.is_empty() { @@ -316,7 +873,7 @@ fn var_key( let index = match subscript { rumoca_core::Subscript::Index { value: index, .. } => positive_i64_to_usize(*index)?, rumoca_core::Subscript::Expr { expr, .. } => { - positive_usize_from_f64(eval_scalar(expr, bindings, &IndexMap::new())?)? + positive_usize_from_f64(eval_scalar(expr, bindings, &IndexMap::new(), eval_env)?)? } rumoca_core::Subscript::Colon { .. } => return None, }; @@ -418,6 +975,43 @@ mod tests { } } + fn integer(value: i64) -> rumoca_core::Expression { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + span: compile_time_test_span(), + } + } + + fn string(value: &str) -> rumoca_core::Expression { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String(value.to_string()), + span: compile_time_test_span(), + } + } + + fn boolean(value: bool) -> rumoca_core::Expression { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Boolean(value), + span: compile_time_test_span(), + } + } + + fn array(elements: Vec) -> rumoca_core::Expression { + rumoca_core::Expression::Array { + elements, + is_matrix: false, + span: compile_time_test_span(), + } + } + + fn matrix(elements: Vec) -> rumoca_core::Expression { + rumoca_core::Expression::Array { + elements, + is_matrix: true, + span: compile_time_test_span(), + } + } + fn var_ref(name: &str) -> rumoca_core::Expression { rumoca_core::Expression::VarRef { name: rumoca_core::Reference::new(name), @@ -426,6 +1020,411 @@ mod tests { } } + fn indexed_var(name: &str, index: usize) -> rumoca_core::Expression { + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new(name), + subscripts: vec![rumoca_core::Subscript::generated_index( + index as i64, + compile_time_test_span(), + )], + span: compile_time_test_span(), + } + } + + fn binary( + op: rumoca_core::OpBinary, + lhs: rumoca_core::Expression, + rhs: rumoca_core::Expression, + ) -> rumoca_core::Expression { + rumoca_core::Expression::Binary { + op, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span: compile_time_test_span(), + } + } + + fn builtin_call( + function: rumoca_core::BuiltinFunction, + args: Vec, + ) -> rumoca_core::Expression { + rumoca_core::Expression::BuiltinCall { + function, + args, + span: compile_time_test_span(), + } + } + + fn range( + start: rumoca_core::Expression, + end: rumoca_core::Expression, + ) -> rumoca_core::Expression { + rumoca_core::Expression::Range { + start: Box::new(start), + step: None, + end: Box::new(end), + span: compile_time_test_span(), + } + } + + fn comprehension( + index: &str, + values: rumoca_core::Expression, + expr: rumoca_core::Expression, + ) -> rumoca_core::Expression { + rumoca_core::Expression::ArrayComprehension { + expr: Box::new(expr), + indices: vec![rumoca_core::ComprehensionIndex { + name: index.to_string(), + range: values, + }], + filter: None, + span: compile_time_test_span(), + } + } + + fn insert_parameter_start( + dae_model: &mut dae::Dae, + name: &str, + start: rumoca_core::Expression, + is_tunable: bool, + dims: &[i64], + ) { + dae_model.variables.parameters.insert( + rumoca_core::VarName::new(name), + dae::Variable { + name: rumoca_core::VarName::new(name), + dims: dims.to_vec(), + start: Some(start), + is_tunable, + ..dae::Variable::empty_with_span(compile_time_test_span()) + }, + ); + } + + fn insert_external_time_table_record_fields( + dae_model: &mut dae::Dae, + prefix: &str, + table: rumoca_core::Expression, + table_dims: &[i64], + ) { + let field_specs = [ + ("tableName", string("NoName"), false, &[][..]), + ("fileName", string("NoName"), false, &[][..]), + ("table", table, true, table_dims), + ("startTime", real(0.0), true, &[][..]), + ("columns", array(vec![integer(2)]), false, &[1][..]), + ( + "smoothness", + var_ref("Modelica.Blocks.Types.Smoothness.ConstantSegments"), + false, + &[][..], + ), + ( + "extrapolation", + var_ref("Modelica.Blocks.Types.Extrapolation.HoldLastPoint"), + true, + &[][..], + ), + ("shiftTime", real(0.0), true, &[][..]), + ( + "timeEvents", + var_ref("Modelica.Blocks.Types.TimeEvents.Always"), + false, + &[][..], + ), + ("verboseRead", boolean(false), false, &[][..]), + ("delimiter", string(","), false, &[][..]), + ("nHeaderLines", integer(0), false, &[][..]), + ]; + for (field, start, is_tunable, dims) in field_specs { + insert_parameter_start( + dae_model, + &format!("{prefix}.{field}"), + start, + is_tunable, + dims, + ); + } + } + + fn insert_time_table_enum_ordinals(dae_model: &mut dae::Dae) { + dae_model.symbols.enum_literal_ordinals.insert( + "Modelica.Blocks.Types.Smoothness.ConstantSegments".to_string(), + 3, + ); + dae_model.symbols.enum_literal_ordinals.insert( + "Modelica.Blocks.Types.Extrapolation.HoldLastPoint".to_string(), + 1, + ); + dae_model + .symbols + .enum_literal_ordinals + .insert("Modelica.Blocks.Types.TimeEvents.Always".to_string(), 1); + } + + #[test] + fn structural_bindings_evaluate_external_table_constructor_starts() { + let mut dae_model = dae::Dae::default(); + let table_expr = array(vec![ + array(vec![real(0.0), real(1.0)]), + array(vec![real(1.0), real(2.0)]), + ]); + let constructor = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("ExternalCombiTimeTable"), + args: vec![ + string("NoName"), + string("NoName"), + table_expr, + real(0.0), + array(vec![integer(2)]), + ], + is_constructor: false, + span: compile_time_test_span(), + }; + dae_model.variables.parameters.insert( + rumoca_core::VarName::new("tableID"), + dae::Variable { + name: rumoca_core::VarName::new("tableID"), + start: Some(constructor), + is_tunable: false, + ..dae::Variable::empty_with_span(compile_time_test_span()) + }, + ); + + let bindings = structural_bindings(&dae_model) + .expect("external table constructors should produce structural table ids"); + + assert!( + bindings.get("tableID").is_some_and(|value| *value > 0.0), + "expected numeric table id binding, got {bindings:?}" + ); + } + + #[test] + fn structural_bindings_derive_metadata_external_table_handle_from_parameter_start() { + let mut dae_model = dae::Dae::default(); + let table_expr = array(vec![ + array(vec![real(0.0), real(1.0)]), + array(vec![real(1.0), real(2.0)]), + ]); + let constructor = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("ExternalCombiTimeTable"), + args: vec![ + string("NoName"), + string("NoName"), + table_expr, + real(0.0), + array(vec![integer(2)]), + ], + is_constructor: false, + span: compile_time_test_span(), + }; + let table_min = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("getTimeTableTmin"), + args: vec![constructor], + is_constructor: false, + span: compile_time_test_span(), + }; + dae_model.variables.parameters.insert( + rumoca_core::VarName::new("block.table.t_minScaled"), + dae::Variable { + name: rumoca_core::VarName::new("block.table.t_minScaled"), + start: Some(table_min), + is_tunable: false, + ..dae::Variable::empty_with_span(compile_time_test_span()) + }, + ); + dae_model + .metadata + .nonnumeric_variable_names + .push("block.table.tableID".to_string()); + + let bindings = structural_bindings(&dae_model) + .expect("external table metadata handles should be derived from table users"); + + assert!( + bindings + .get("block.table.tableID") + .is_some_and(|value| *value > 0.0), + "expected metadata table id binding, got {bindings:?}" + ); + } + + #[test] + fn structural_bindings_derive_external_time_table_handle_from_record_fields() { + let mut dae_model = dae::Dae::default(); + insert_external_time_table_record_fields( + &mut dae_model, + "block.table", + array(vec![ + array(vec![real(2.0), real(0.0)]), + array(vec![real(4.0), real(1.0)]), + ]), + &[2, 2], + ); + dae_model + .metadata + .nonnumeric_variable_names + .push("block.table.tableID".to_string()); + insert_time_table_enum_ordinals(&mut dae_model); + + let bindings = structural_bindings(&dae_model) + .expect("record fields should derive external time table handle"); + + assert!( + bindings + .get("block.table.tableID") + .is_some_and(|value| *value > 0.0), + "expected record-field table id binding, got {bindings:?}" + ); + } + + #[test] + fn structural_bindings_materialize_boolean_table_dynamic_matrix() { + let mut dae_model = dae::Dae::default(); + insert_parameter_start( + &mut dae_model, + "booleanTable.table", + array(vec![ + real(2.0), + real(4.0), + real(6.0), + real(6.5), + real(7.0), + real(9.0), + real(11.0), + ]), + true, + &[0], + ); + insert_parameter_start( + &mut dae_model, + "booleanTable.n", + builtin_call( + rumoca_core::BuiltinFunction::Size, + vec![var_ref("booleanTable.table"), integer(1)], + ), + false, + &[], + ); + insert_parameter_start( + &mut dae_model, + "booleanTable.startValue", + boolean(false), + false, + &[], + ); + let toggle_values = comprehension( + "i", + range(integer(1), var_ref("booleanTable.n")), + builtin_call( + rumoca_core::BuiltinFunction::Mod, + vec![var_ref("i"), real(2.0)], + ), + ); + let table_matrix = rumoca_core::Expression::If { + branches: vec![( + binary( + rumoca_core::OpBinary::Gt, + var_ref("booleanTable.n"), + real(0.0), + ), + matrix(vec![ + array(vec![indexed_var("booleanTable.table", 1), real(0.0)]), + array(vec![var_ref("booleanTable.table"), toggle_values]), + ]), + )], + else_branch: Box::new(matrix(vec![array(vec![real(0.0), real(0.0)])])), + span: compile_time_test_span(), + }; + insert_external_time_table_record_fields( + &mut dae_model, + "booleanTable.combiTimeTable", + table_matrix, + &[0, 2], + ); + dae_model + .metadata + .nonnumeric_variable_names + .push("booleanTable.combiTimeTable.tableID".to_string()); + insert_time_table_enum_ordinals(&mut dae_model); + + let mut env = compile_time_eval_env(&dae_model); + let mut bindings = enum_literal_bindings(&dae_model.symbols.enum_literal_ordinals); + let shapes = variable_shapes(&dae_model); + insert_shape_bindings(&mut bindings, &shapes); + insert_constant_variables(&mut bindings, &dae_model, &shapes, &mut env) + .expect("constants should seed"); + seed_compile_time_start_values(&dae_model, &mut env); + insert_structural_parameters(&mut bindings, &dae_model, &shapes, &mut env) + .expect("structural parameters should seed"); + insert_external_table_handle_bindings(&mut bindings, &dae_model, &shapes, &mut env) + .expect("external table handle should seed"); + + let table_id = *bindings + .get("booleanTable.combiTimeTable.tableID") + .expect("tableID binding"); + let table_data = rumoca_eval_dae::all_external_table_data_in_env(&env); + assert_eq!(table_data.len(), 1, "expected one external table"); + assert_eq!(table_data[0].id as f64, table_id); + assert!( + table_data[0].data.len() > 1, + "BooleanTable handle must not register an empty table: {table_data:?}" + ); + } + + #[test] + fn metadata_external_table_handle_overrides_default_empty_table_id() { + let mut dae_model = dae::Dae::default(); + let default_empty_constructor = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("ExternalCombiTimeTable"), + args: vec![ + string("NoName"), + string("NoName"), + matrix(Vec::new()), + real(0.0), + array(vec![integer(2)]), + ], + is_constructor: false, + span: compile_time_test_span(), + }; + insert_parameter_start( + &mut dae_model, + "block.table.tableID", + default_empty_constructor, + false, + &[], + ); + insert_external_time_table_record_fields( + &mut dae_model, + "block.table", + matrix(vec![ + array(vec![real(0.0), real(1.0)]), + array(vec![real(1.0), real(2.0)]), + ]), + &[2, 2], + ); + dae_model + .metadata + .nonnumeric_variable_names + .push("block.table.tableID".to_string()); + insert_time_table_enum_ordinals(&mut dae_model); + + let (bindings, env) = structural_bindings_with_eval_env(&dae_model) + .expect("metadata table handle should override stale default tableID"); + let table_id = *bindings + .get("block.table.tableID") + .expect("tableID binding"); + let table_data = rumoca_eval_dae::all_external_table_data_in_env(&env); + let selected = table_data + .iter() + .find(|table| table.id as f64 == table_id) + .expect("selected table data should be registered"); + + assert_eq!(selected.data.len(), 2, "selected table must be non-empty"); + } + #[test] fn structural_bindings_report_invalid_variable_shape_span() { let mut dae_model = dae::Dae::default(); @@ -521,8 +1520,23 @@ mod tests { fn eval_size_declines_unrepresentable_dimension_index() { let args = vec![var_ref("x"), real(usize::MAX as f64)]; let shapes = IndexMap::from([("x".to_string(), vec![2, 3])]); + let mut eval_env = rumoca_eval_dae::VarEnv::new(); - assert_eq!(eval_size(&args, &IndexMap::new(), &shapes), None); + assert_eq!( + eval_size(&args, &IndexMap::new(), &shapes, &mut eval_env), + None + ); + } + + #[test] + fn eval_size_reads_string_literal_array_shape_without_numeric_values() { + let args = vec![array(vec![string("water"), string("air")]), integer(1)]; + let mut eval_env = rumoca_eval_dae::VarEnv::new(); + + assert_eq!( + eval_size(&args, &IndexMap::new(), &IndexMap::new(), &mut eval_env), + Some(vec![2.0]) + ); } #[test] @@ -531,12 +1545,14 @@ mod tests { Box::new(real(usize::MAX as f64)), compile_time_test_span(), ); + let mut eval_env = rumoca_eval_dae::VarEnv::new(); assert_eq!( var_key( &rumoca_core::Reference::new("x"), &[subscript], - &IndexMap::new() + &IndexMap::new(), + &mut eval_env ), None ); diff --git a/crates/rumoca-phase-solve/src/lower/complex_projection.rs b/crates/rumoca-phase-solve/src/lower/complex_projection.rs index e5ce47861..012b5ba76 100644 --- a/crates/rumoca-phase-solve/src/lower/complex_projection.rs +++ b/crates/rumoca-phase-solve/src/lower/complex_projection.rs @@ -8,7 +8,9 @@ impl<'a> LowerBuilder<'a> { scope: &Scope, call_depth: usize, ) -> Result<(Reg, Reg), LowerError> { - if self.requires_complex_projection(expr, scope)? { + if self.scalarized_complex_fields_available(expr) + || self.requires_complex_projection(expr, scope)? + { let re = self.lower_field_access(expr, "re", owner_span, scope, call_depth)?; let im = self.lower_field_access(expr, "im", owner_span, scope, call_depth)?; return Ok((re, im)); @@ -32,6 +34,13 @@ impl<'a> LowerBuilder<'a> { rumoca_core::Expression::VarRef { name, subscripts, .. } if subscripts.is_empty() => { + if self.component_reference_field_available(name, "re") + || self.component_reference_field_available(name, "im") + || self.scalarized_field_binding_available(name.as_str(), "re") + || self.scalarized_field_binding_available(name.as_str(), "im") + { + return Ok(true); + } let span = complex_projection_reference_span(expr, name)?; let key = self.scope_key_from_reference(name, span)?; Ok(self.component_field_available(&key, name, "re") @@ -124,6 +133,45 @@ impl<'a> LowerBuilder<'a> { || self.indexed_bindings.contains_key(&field_key) } + pub(in crate::lower) fn scalarized_field_binding_available( + &self, + base_key: &str, + field: &str, + ) -> bool { + let field_key = format!("{base_key}.{field}"); + self.layout.binding(&field_key).is_some() + || self + .dae_variables + .and_then(|variables| { + dae_variable(variables, &rumoca_core::VarName::new(&field_key)) + }) + .is_some() + || self.direct_assignments.contains_key(&field_key) + || self.local_indexed_bindings.contains_key(field_key.as_str()) + || self + .indexed_bindings + .contains_key(&ComponentReferenceKey::generated(&field_key)) + } + + pub(in crate::lower) fn scalarized_complex_fields_available( + &self, + expr: &rumoca_core::Expression, + ) -> bool { + if let rumoca_core::Expression::Index { base, .. } = expr + && let Ok(base_key) = binding_base_key(base) + { + return self.scalarized_field_binding_available(&base_key, "re") + || self.scalarized_field_binding_available(&base_key, "im") + || !self.indexed_record_field_keys(&base_key, "re").is_empty() + || !self.indexed_record_field_keys(&base_key, "im").is_empty(); + } + let Ok(base_key) = binding_base_key(expr) else { + return false; + }; + self.scalarized_field_binding_available(&base_key, "re") + || self.scalarized_field_binding_available(&base_key, "im") + } + fn component_reference_field_available( &self, reference: &rumoca_core::Reference, @@ -210,8 +258,16 @@ impl<'a> LowerBuilder<'a> { return Ok(None); } - if let Some(index) = constructor_positional_field_index(field) - && let Some(expr) = args.get(index) + let (named_args, positional_args) = + super::function_calls::split_named_and_positional_call_args(name.as_str(), args)?; + if let Some(arg_expr) = named_args.get(field).copied() { + return self + .lower_expr(arg_expr, caller_scope, call_depth + 1) + .map(Some); + } + if named_args.is_empty() + && let Some(index) = constructor_positional_field_index(field) + && let Some(expr) = positional_args.get(index).copied() { return self.lower_expr(expr, caller_scope, call_depth).map(Some); } @@ -222,8 +278,12 @@ impl<'a> LowerBuilder<'a> { let mut local_scope = Scope::new(); let mut input_regs = IndexMap::::new(); - for (idx, input) in constructor.inputs.iter().enumerate() { - let reg = if let Some(arg_expr) = args.get(idx) { + let mut positional_idx = 0usize; + for input in &constructor.inputs { + let reg = if let Some(arg_expr) = named_args.get(input.name.as_str()).copied() { + self.lower_expr(arg_expr, caller_scope, call_depth + 1)? + } else if let Some(arg_expr) = positional_args.get(positional_idx).copied() { + positional_idx += 1; self.lower_expr(arg_expr, caller_scope, call_depth + 1)? } else if let Some(default_expr) = input.default.as_ref() { self.lower_expr(default_expr, &local_scope, call_depth + 1)? diff --git a/crates/rumoca-phase-solve/src/lower/derivative_rhs.rs b/crates/rumoca-phase-solve/src/lower/derivative_rhs.rs index 64d5a6497..ae0bbf3d0 100644 --- a/crates/rumoca-phase-solve/src/lower/derivative_rhs.rs +++ b/crates/rumoca-phase-solve/src/lower/derivative_rhs.rs @@ -4,6 +4,7 @@ mod equation_collection; mod function_projection; mod linear_parts; mod projection; +mod row_projection; #[cfg(test)] mod tests; use super::{ @@ -15,14 +16,17 @@ pub(super) use equation_collection::*; use indexmap::IndexMap; pub(super) use linear_parts::*; pub(super) use projection::*; +use row_projection::project_derivative_row_expr; use rumoca_core::{Literal, OpBinary, OpUnary}; use rumoca_ir_dae as dae; use rumoca_ir_solve::{BinaryOp, ComputeBlock, ComputeNode, LinearOp, Reg, ScalarSlot, VarLayout}; use std::collections::HashSet; use std::sync::Arc; -use function_projection::{ - function_call_projected_scalars_with_owner, function_projected_residuals_with_owner, +pub(in crate::lower) use function_projection::{ + function_call_projected_output_groups_with_owner, function_call_projected_scalars_with_owner, + function_projected_residuals_with_owner, project_array_like_scalar_with_owner, + project_array_like_scalars_with_owner, }; #[derive(Debug, Clone)] @@ -289,8 +293,10 @@ pub(super) fn lower_derivative_rhs( dae_model: &dae::Dae, layout: &VarLayout, ) -> Result { - let analysis = analyze_derivative_rhs(dae_model)?; + let analysis = analyze_derivative_rhs(dae_model) + .map_err(|err| err.with_context("analyze derivative RHS".to_string()))?; lower_derivative_rhs_with_analysis(dae_model, layout, &analysis) + .map_err(|err| err.with_context("lower derivative RHS".to_string())) } // SPEC_0021: Exception - derivative RHS lowering owns block assembly across @@ -357,7 +363,7 @@ pub(crate) fn lower_derivative_rhs_with_analysis( for idx in component { group.push(analysis.states[*idx].clone()); } - let node = lower_linsolve_group(&group, &lowering_ctx)?; + let node = lower_linsolve_group(component, &group, &lowering_ctx)?; reserve_derivative_capacity( &mut block.nodes, 1, @@ -372,7 +378,7 @@ pub(crate) fn lower_derivative_rhs_with_analysis( continue; } - if let Some(group_len) = direct_vector_group_len(analysis, &processed, i) { + if let Some(group_len) = direct_vector_group_len(dae_model, analysis, &processed, i) { match lower_direct_row_group(analysis, i, group_len, &lowering_ctx) { Ok(DirectRowGroupLowering::Scalar(row)) => { flush_pending_derivative_programs( @@ -431,7 +437,10 @@ pub(crate) fn lower_derivative_rhs_with_analysis( .filter(|span| !span.is_dummy()) .map(Ok) .unwrap_or_else(|| derivative_state_or_context_span(dae_model, state))?; - let row = lower_state_derivative_row(state, &analysis.direct_equations, &lowering_ctx)?; + let row = lower_state_derivative_row(state, &analysis.direct_equations, &lowering_ctx) + .map_err(|err| { + err.with_context(format!("lower derivative row for `{}`", state.name)) + })?; reserve_derivative_capacity( &mut pending_derivative_programs, 1, @@ -450,7 +459,13 @@ pub(crate) fn lower_derivative_rhs_with_analysis( state, &analysis.direct_equations, &lowering_ctx, - )?, + ) + .map_err(|err| { + err.with_context(format!( + "build derivative access proof for `{}`", + state.name + )) + })?, }); processed[i] = true; i += 1; @@ -930,14 +945,24 @@ fn lower_direct_row( let scope = Scope::new(); let mut active_assignments = active_assignment_stack(equation.span)?; let rhs_expr = inline_direct_assignment_expr(&equation.rhs, ctx, &mut active_assignments)?; - let rhs = lower_state_component_expr(&mut builder, &rhs_expr, state, equation.span, &scope)?; + let rhs = lower_state_component_expr(&mut builder, &rhs_expr, state, equation.span, &scope) + .map_err(|err| { + err.with_context(format!( + "lower derivative RHS for `{}` from {}", + state.name, + short_expr(&rhs_expr, 160) + )) + })?; let mut coeff_active_assignments = active_assignment_stack(equation.span)?; let coeff_expr = inline_direct_assignment_expr( &equation.coefficients[&state.name], ctx, &mut coeff_active_assignments, )?; - let coeff = builder.lower_expr(&coeff_expr, &scope, 0)?; + let coeff = lower_state_component_expr(&mut builder, &coeff_expr, state, equation.span, &scope) + .map_err(|err| { + err.with_context(format!("lower derivative coefficient for `{}`", state.name)) + })?; let value = builder.emit_binary_at(BinaryOp::Div, rhs, coeff, equation.span)?; builder.ops.push(LinearOp::StoreOutput { src: value }); Ok(builder.ops) @@ -1272,6 +1297,7 @@ fn shared_vector_rhs_base(expr: &rumoca_core::Expression) -> &rumoca_core::Expre } fn direct_vector_group_len( + dae_model: &dae::Dae, analysis: &DerivativeRhsAnalysis, processed: &[bool], start: usize, @@ -1304,6 +1330,18 @@ fn direct_vector_group_len( return None; } let eq_idx = *analysis.direct_equations.get(&state.name)?; + if analysis.equations[eq_idx] + .dae_equation_index + .and_then(|equation_index| { + dae::structured_equation_slot( + &dae_model.continuous.structured_equations, + equation_index, + ) + }) + .is_some() + { + return None; + } if shared_vector_rhs_base(&analysis.equations[eq_idx].rhs) != head_base { return None; } @@ -1356,8 +1394,11 @@ fn lower_direct_row_group_scalar( // Compute every component of the shared base RHS once (shared by RowCse // within this single builder), then project each component below. let head_base = shared_vector_rhs_base(&head_eq.rhs); - let values = - builder.lower_array_like_values_with_source_context(head_base, head_eq.span, &scope, 0)?; + let values = builder + .lower_array_like_values_with_source_context(head_base, head_eq.span, &scope, 0) + .map_err(|err| { + err.with_context(format!("lower shared derivative RHS for `{}`", head.base)) + })?; if values.len() != group_len { return Err(LowerError::Unsupported { reason: format!( @@ -1372,7 +1413,16 @@ fn lower_direct_row_group_scalar( let state = &analysis.states[start + offset]; let equation = &analysis.equations[analysis.direct_equations[&state.name]]; let rhs = values[state.component]; - let coeff = builder.lower_expr(&equation.coefficients[&state.name], &scope, 0)?; + let coeff = lower_state_component_expr( + &mut builder, + &equation.coefficients[&state.name], + state, + equation.span, + &scope, + ) + .map_err(|err| { + err.with_context(format!("lower derivative coefficient for `{}`", state.name)) + })?; let value = builder.emit_binary_at(BinaryOp::Div, rhs, coeff, equation.span)?; builder.ops.push(LinearOp::StoreOutput { src: value }); } @@ -1411,8 +1461,31 @@ fn lower_row_rhs_expr( return builder.lower_expr_with_source_context(expr, source_context_span, scope, 0); } - let values = - builder.lower_array_like_values_with_source_context(expr, source_context_span, scope, 0)?; + let values = match builder.lower_array_like_values_with_source_context( + expr, + source_context_span, + scope, + 0, + ) { + Ok(values) => values, + Err(err) => { + let projected = project_derivative_row_expr( + builder, + expr, + row_index, + row_count, + source_context_span, + scope, + ) + .map_err(|_| err)?; + return builder.lower_expr_with_source_context( + &projected, + source_context_span, + scope, + 0, + ); + } + }; if values.len() == row_count { return values .get(row_index) @@ -1459,6 +1532,7 @@ fn lower_coupled_row( /// does for one component. Unlike that function, we do it once and emit a /// single tensor node that backends can execute without repeating the solve. fn lower_linsolve_group( + output_indices: &[usize], states: &[StateScalar], ctx: &DerivativeRhsLoweringContext<'_>, ) -> Result { @@ -1470,6 +1544,7 @@ fn lower_linsolve_group( rhs_start: setup.rhs_start, n: setup.n, next_reg: setup.next_reg, + output_indices: output_indices.to_vec(), metadata: rumoca_ir_solve::TensorNodeMetadata::default(), span: setup.span, }) diff --git a/crates/rumoca-phase-solve/src/lower/derivative_rhs/equation_collection.rs b/crates/rumoca-phase-solve/src/lower/derivative_rhs/equation_collection.rs index 4105ec6a5..22c0ce670 100644 --- a/crates/rumoca-phase-solve/src/lower/derivative_rhs/equation_collection.rs +++ b/crates/rumoca-phase-solve/src/lower/derivative_rhs/equation_collection.rs @@ -1119,7 +1119,11 @@ fn optional_derivative_probe( ) -> Result, LowerError> { match result { Ok(value) => Ok(value), - Err(LowerError::MissingBinding { .. } | LowerError::Unsupported { .. }) => Ok(None), + Err( + LowerError::MissingBinding { .. } + | LowerError::Unsupported { .. } + | LowerError::UnsupportedAt { .. }, + ) => Ok(None), Err(LowerError::InvalidFunction { name, .. }) if name == "projected function output" => { Ok(None) } @@ -1183,7 +1187,7 @@ pub(in crate::lower) fn matrix_times_derivative_coefficients( } let matrix_span = derivative_expr_span_or_owner(matrix, owner_span)?; let Some(matrix_coefficients) = - expression_binding_expressions(matrix, dae_model, structural_bindings, matrix_span)? + matrix_coefficient_expressions(matrix, dae_model, structural_bindings, matrix_span)? else { return Ok(None); }; @@ -1214,7 +1218,7 @@ pub(in crate::lower) fn derivative_times_matrix_coefficients( } let matrix_span = derivative_expr_span_or_owner(matrix, owner_span)?; let Some(matrix_coefficients) = - expression_binding_expressions(matrix, dae_model, structural_bindings, matrix_span)? + matrix_coefficient_expressions(matrix, dae_model, structural_bindings, matrix_span)? else { return Ok(None); }; @@ -1264,6 +1268,26 @@ pub(in crate::lower) fn build_matrix_derivative_equations( Ok(Some(equations)) } +fn matrix_coefficient_expressions( + matrix: &rumoca_core::Expression, + dae_model: &dae::Dae, + structural_bindings: &IndexMap, + span: rumoca_core::Span, +) -> Result>, LowerError> { + if let Some(coefficients) = + expression_binding_expressions(matrix, dae_model, structural_bindings, span)? + { + return Ok(Some(coefficients)); + } + match matrix { + rumoca_core::Expression::Array { elements, .. } + | rumoca_core::Expression::Tuple { elements, .. } => { + literal_array_elements_flat(elements, span).map(Some) + } + _ => Ok(None), + } +} + fn matrix_derivative_coefficient_rows( target_keys: &[String], matrix_coefficients: &[rumoca_core::Expression], diff --git a/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection.rs b/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection.rs index c5c8e21ff..6e1e66f78 100644 --- a/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection.rs +++ b/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection.rs @@ -8,7 +8,9 @@ use indexmap::IndexMap; use rumoca_core::{ExpressionRewriter, Literal, NAMED_FUNCTION_ARG_PREFIX, OpBinary}; use rumoca_ir_dae as dae; +use crate::lower::helpers::is_stream_passthrough_intrinsic; use crate::lower::{LowerError, unsupported_at}; +use crate::projection_suffix::parse_output_projection_suffix; #[path = "function_projection/compile_time.rs"] mod compile_time; @@ -35,21 +37,24 @@ mod tests; use dimension_helpers::{ FunctionScopeSubstituter, append_projected_outputs, array_expression_dims, assignment_projection_dims, binary_mul_dims, constructor_input_projection_dims, - copy_projection_dims, declared_dims, elementwise_binary_dims, + copy_projection_dims, declared_param_dims, elementwise_binary_dims, exact_declared_function_output_dims, flat_index_from_indices, flatten_array_elements, - formal_actual_projection_dims, is_ignorable_projection_statement, is_same_plain_var_ref, - named_actual_span, named_argument_spans, projected_declared_output_dims, - projected_field_output_dims, projection_assignment_target, required_flat_index_to_subscripts, - reserve_projection_capacity, scalar_count_for_dims, single_field_path, sum_expressions, + formal_accepts_structured_actual, formal_actual_projection_dims, + is_ignorable_projection_statement, is_same_plain_var_ref, named_actual_span, + named_argument_spans, projected_declared_output_dims, projected_field_output_dims, + projection_assignment_target, required_flat_index_to_subscripts, reserve_projection_capacity, + scalar_count_for_dims, selector_dims_from_indices, single_field_path, sum_expressions, valid_product_dim, }; +pub(in crate::lower) use entrypoints::function_projected_residuals_with_owner; use entrypoints::{ checked_generated_subscript_from_usize, checked_projection_offset, checked_usize_dims_to_i64, - checked_usize_to_i64, function_outputs_dims, project_scalar_outputs, - project_target_scalar_outputs, required_merged_projection_dims, variable_dims_i64, + checked_usize_to_i64, function_outputs_dims, project_target_scalar_outputs, + required_merged_projection_dims, variable_dims_i64, }; -pub(super) use entrypoints::{ - function_call_projected_scalars_with_owner, function_projected_residuals_with_owner, +pub(in crate::lower) use entrypoints::{ + function_call_projected_output_groups_with_owner, function_call_projected_scalars_with_owner, + project_array_like_scalar_with_owner, project_array_like_scalars_with_owner, }; use inline_budget::{ CachedProjectionOutcome, CallOutputsCacheEntry, exceeds_projection_node_budget, @@ -76,6 +81,10 @@ use super::{ variable_by_name, }; +const MAX_STATIC_WHILE_PROJECTION_ITERATIONS: usize = 1024; + +type ConstructorInputScalars = (Vec, Vec); + #[derive(Debug, Clone)] struct ProjectedFunctionOutput { field_path: Vec, @@ -99,6 +108,178 @@ struct FunctionProjectionScope { dims: IndexMap>, } +fn seed_declared_scalar_dims( + function: &rumoca_core::Function, + scope: &mut FunctionProjectionScope, +) { + for param in function + .inputs + .iter() + .chain(function.outputs.iter()) + .chain(function.locals.iter()) + { + if param.dims.is_empty() + && param.shape_expr.is_empty() + && !formal_accepts_structured_actual(param) + { + scope.dims.entry(param.name.clone()).or_default(); + } + } +} + +#[derive(Clone, Copy)] +struct FunctionCallStatementProjection<'a> { + function: &'a rumoca_core::Function, + comp: &'a rumoca_core::ComponentReference, + args: &'a [rumoca_core::Expression], + output_targets: &'a [rumoca_core::ComponentReference], + span: rumoca_core::Span, + depth: usize, +} + +fn merge_vectorized_scalar_dims( + dims: &mut Option>, + candidate: &[i64], + name: &str, + span: rumoca_core::Span, +) -> Result<(), LowerError> { + if candidate == [1] && dims.as_ref().is_some_and(|dims| dims.as_slice() != [1]) { + return Ok(()); + } + if dims + .as_ref() + .is_some_and(|dims| dims.as_slice() == [1] && candidate != [1]) + { + *dims = Some(candidate.to_vec()); + return Ok(()); + } + match dims { + Some(existing) if existing.as_slice() != candidate => Err(LowerError::contract_violation( + format!( + "vectorized scalar `{name}` has dimensions {}, expected {}", + format_i64_dims(candidate), + format_i64_dims(existing) + ), + span, + )), + Some(_) => Ok(()), + None => { + *dims = Some(candidate.to_vec()); + Ok(()) + } + } +} + +fn is_elementwise_binary_projection_op(op: &OpBinary) -> bool { + is_add(op) + || is_sub(op) + || is_mul(op) + || is_div(op) + || op.is_relational() + || matches!( + op, + OpBinary::And | OpBinary::Or | OpBinary::Exp | OpBinary::ExpElem + ) +} + +fn is_elementwise_builtin_projection(function: &rumoca_core::BuiltinFunction) -> bool { + matches!( + function, + rumoca_core::BuiltinFunction::Abs + | rumoca_core::BuiltinFunction::Sign + | rumoca_core::BuiltinFunction::Sqrt + | rumoca_core::BuiltinFunction::Div + | rumoca_core::BuiltinFunction::Mod + | rumoca_core::BuiltinFunction::Rem + | rumoca_core::BuiltinFunction::Floor + | rumoca_core::BuiltinFunction::Ceil + | rumoca_core::BuiltinFunction::Min + | rumoca_core::BuiltinFunction::Max + | rumoca_core::BuiltinFunction::Sin + | rumoca_core::BuiltinFunction::Cos + | rumoca_core::BuiltinFunction::Tan + | rumoca_core::BuiltinFunction::Asin + | rumoca_core::BuiltinFunction::Acos + | rumoca_core::BuiltinFunction::Atan + | rumoca_core::BuiltinFunction::Atan2 + | rumoca_core::BuiltinFunction::Sinh + | rumoca_core::BuiltinFunction::Cosh + | rumoca_core::BuiltinFunction::Tanh + | rumoca_core::BuiltinFunction::Exp + | rumoca_core::BuiltinFunction::Log + | rumoca_core::BuiltinFunction::Log10 + | rumoca_core::BuiltinFunction::Integer + | rumoca_core::BuiltinFunction::NoEvent + | rumoca_core::BuiltinFunction::Smooth + | rumoca_core::BuiltinFunction::Homotopy + ) +} + +fn projected_child_flat_index(child_dims: &[i64], outer_flat_index: usize) -> usize { + if child_dims == [1] { + 0 + } else { + outer_flat_index + } +} + +fn function_call_declared_output_count( + expr: &rumoca_core::Expression, + dae_model: &dae::Dae, +) -> Option { + let rumoca_core::Expression::FunctionCall { + name, + is_constructor: false, + .. + } = expr + else { + return None; + }; + dae_model + .symbols + .functions + .get(name.var_name()) + .map(|function| function.outputs.len()) +} + +fn is_direct_declared_array_output_call( + expr: &rumoca_core::Expression, + dae_model: &dae::Dae, +) -> bool { + let rumoca_core::Expression::FunctionCall { + name, + is_constructor: false, + .. + } = expr + else { + return false; + }; + dae_model + .symbols + .functions + .get(name.var_name()) + .is_some_and( + |function| matches!(function.outputs.as_slice(), [output] if !output.dims.is_empty()), + ) +} + +fn checked_shape_dimension(value: f64, span: rumoca_core::Span) -> Result { + let rounded = value.round(); + if !value.is_finite() || (rounded - value).abs() > 1e-9 { + return Err(unsupported_at( + format!("function parameter shape dimension must be an integer, got `{value}`"), + span, + )); + } + if rounded < 0.0 || rounded >= i64::MAX as f64 { + return Err(unsupported_at( + format!("function parameter shape dimension `{value}` is out of range"), + span, + )); + } + Ok(rounded as i64) +} + impl<'a> FunctionProjectionAnalysis<'a> { fn new(dae_model: &'a dae::Dae, structural_bindings: &'a IndexMap) -> Self { Self { @@ -206,6 +387,16 @@ impl<'a> FunctionProjectionAnalysis<'a> { expr: &rumoca_core::Expression, depth: usize, owner_span: rumoca_core::Span, + ) -> Result>, LowerError> { + self.function_call_outputs_with_projection_scope(expr, depth, owner_span, None) + } + + fn function_call_outputs_with_projection_scope( + &self, + expr: &rumoca_core::Expression, + depth: usize, + owner_span: rumoca_core::Span, + caller_scope: Option<&FunctionProjectionScope>, ) -> Result>, LowerError> { if depth > super::super::MAX_FUNCTION_INLINE_DEPTH { return Ok(None); @@ -228,11 +419,19 @@ impl<'a> FunctionProjectionAnalysis<'a> { if !function.pure || function.external.is_some() { return Ok(None); } - if let Some(outcome) = self.cached_call_outputs(expr, depth) { + if caller_scope.is_none() + && let Some(outcome) = self.cached_call_outputs(expr, depth) + { return outcome.into_result(); } let call_span = inherited_projection_source_span(expr.span(), owner_span); - let outcome = match self.uncached_function_call_outputs(function, args, depth, call_span) { + let outcome = match self.uncached_function_call_outputs( + function, + args, + depth, + call_span, + caller_scope, + ) { Ok(outputs) => CachedProjectionOutcome::Outputs(outputs), Err(err) => match err.projection_budget_exceeded_parts() { Some((function, span)) => CachedProjectionOutcome::BudgetExceeded { @@ -242,7 +441,9 @@ impl<'a> FunctionProjectionAnalysis<'a> { None => return Err(err), }, }; - self.record_call_outputs_outcome(expr, depth, outcome.clone()); + if caller_scope.is_none() { + self.record_call_outputs_outcome(expr, depth, outcome.clone()); + } outcome.into_result() } @@ -252,9 +453,17 @@ impl<'a> FunctionProjectionAnalysis<'a> { args: &[rumoca_core::Expression], depth: usize, owner_span: rumoca_core::Span, + caller_scope: Option<&FunctionProjectionScope>, ) -> Result>, LowerError> { let function_span = inherited_projection_span(function.span, owner_span); - let Some(mut scope) = self.bind_inputs(function, args, depth + 1, function_span)? else { + let Some(mut scope) = self.bind_inputs_with_projection_scope( + function, + args, + depth + 1, + function_span, + caller_scope, + )? + else { return Ok(None); }; if scope @@ -282,11 +491,11 @@ impl<'a> FunctionProjectionAnalysis<'a> { function_span, )?; } - let outputs = if projected.is_empty() { - self.projected_outputs_from_scope(function, &scope, depth + 1, function_span)? - } else { - Some(projected) - }; + let outputs = + self.complete_projected_outputs(function, &scope, projected, depth + 1, function_span)?; + let outputs = outputs + .map(|outputs| self.resolve_projected_scalar_field_outputs(outputs, function_span)) + .transpose()?; if let Some(outputs) = &outputs && outputs .iter() @@ -300,12 +509,61 @@ impl<'a> FunctionProjectionAnalysis<'a> { Ok(outputs) } + fn complete_projected_outputs( + &self, + function: &rumoca_core::Function, + scope: &FunctionProjectionScope, + mut projected: Vec, + depth: usize, + function_span: rumoca_core::Span, + ) -> Result>, LowerError> { + let Some(scope_outputs) = + self.projected_outputs_from_scope(function, scope, depth, function_span)? + else { + return Ok((!projected.is_empty()).then_some(projected)); + }; + if projected.is_empty() { + return Ok(Some(scope_outputs)); + } + if function.outputs.len() <= 1 { + return Ok(Some(projected)); + } + for output_param in &function.outputs { + if projected + .iter() + .any(|output| output.field_path.first() == Some(&output_param.name)) + { + continue; + } + for output in scope_outputs + .iter() + .filter(|output| output.field_path.first() == Some(&output_param.name)) + { + projected.push(output.clone()); + } + } + Ok(Some(projected)) + } + + #[cfg(test)] fn bind_inputs( &self, function: &rumoca_core::Function, args: &[rumoca_core::Expression], depth: usize, owner_span: rumoca_core::Span, + ) -> Result, LowerError> { + self.bind_inputs_with_projection_scope(function, args, depth, owner_span, None) + } + + #[allow(clippy::excessive_nesting, clippy::too_many_lines)] + fn bind_inputs_with_projection_scope( + &self, + function: &rumoca_core::Function, + args: &[rumoca_core::Expression], + depth: usize, + owner_span: rumoca_core::Span, + caller_scope: Option<&FunctionProjectionScope>, ) -> Result, LowerError> { let (named, positional) = super::super::function_calls::split_named_and_positional_call_args( @@ -314,37 +572,145 @@ impl<'a> FunctionProjectionAnalysis<'a> { )?; let named_spans = named_argument_spans(args, owner_span)?; let mut scope = FunctionProjectionScope::default(); + seed_declared_scalar_dims(function, &mut scope); let mut positional_idx = 0usize; - for input in &function.inputs { + let used_inputs = super::super::function_calls::referenced_function_input_names(function); + for (input_idx, input) in function.inputs.iter().enumerate() { let input_span = inherited_projection_span(input.span, owner_span); + if !used_inputs.contains(&input.name) + && let Some((prefix, field)) = split_flattened_projection_input_name(&input.name) + { + let flattened_group_has_used_sibling = function.inputs.iter().any(|candidate| { + flattened_projection_input_has_prefix(&candidate.name, prefix) + && used_inputs.contains(&candidate.name) + }); + if flattened_group_has_used_sibling && !named.contains_key(input.name.as_str()) { + let later_flattened_sibling_is_used = + function.inputs.iter().skip(input_idx + 1).any(|next| { + flattened_projection_input_has_prefix(&next.name, prefix) + && used_inputs.contains(&next.name) + }); + if !later_flattened_sibling_is_used + || positional.get(positional_idx).is_some_and(|arg| { + super::super::function_calls::is_flattened_record_field_actual( + arg, field, + ) + }) + { + positional_idx += usize::from(positional_idx < positional.len()); + } + continue; + } + } + let mut consume_flattened_positional = false; + let mut project_flattened_positional_field = false; let actual = if let Some(actual) = named.get(input.name.as_str()).copied() { Some(( - actual, + actual.clone(), inherited_projection_span( named_actual_span(&named_spans, input, actual), input_span, ), + true, )) + } else if let Some((prefix, field)) = split_flattened_projection_input_name(&input.name) + && let Some(actual) = positional.get(positional_idx).cloned() + && (self.flattened_projection_actual_has_field(actual, field, &scope) + || caller_scope.is_some_and(|caller_scope| { + self.flattened_projection_actual_has_field(actual, field, caller_scope) + }) + || (positional.len().saturating_sub(positional_idx) == 1 + && flattened_projection_input_is_group_start( + &function.inputs, + input_idx, + prefix, + ) + && flattened_projection_group_has_prefix(&function.inputs, prefix))) + { + project_flattened_positional_field = true; + consume_flattened_positional = !function + .inputs + .iter() + .skip(input_idx + 1) + .any(|next| flattened_projection_input_has_prefix(&next.name, prefix)); + let actual = self + .flattened_projection_actual_field_value(actual, field, caller_scope) + .unwrap_or_else(|| rumoca_core::Expression::FieldAccess { + base: Box::new(actual.clone()), + field: field.to_string(), + span: actual.span().unwrap_or(input_span), + }); + let actual_span = actual.span().unwrap_or(input_span); + Some((actual, actual_span, true)) } else { super::super::function_calls::next_positional_function_input_arg( input, &positional, &mut positional_idx, ) - .map(|actual| (actual, actual.span().unwrap_or(input_span))) + .map(|actual| (actual.clone(), actual.span().unwrap_or(input_span), true)) }; - let Some((actual, actual_span)) = actual.or_else(|| { + let Some((actual, actual_span, actual_is_call_arg)) = actual.or_else(|| { input .default .as_ref() - .map(|actual| (actual, actual.span().unwrap_or(input_span))) + .map(|actual| (actual.clone(), actual.span().unwrap_or(input_span), false)) }) else { return Ok(None); }; - let actual = actual.clone(); + if consume_flattened_positional { + positional_idx += 1; + } + if actual_is_call_arg && let Some(caller_scope) = caller_scope { + let caller_actual = self.substitute(&actual, caller_scope)?; + let caller_dims = self.expr_dims_with_owner( + &caller_actual, + caller_scope, + depth + 1, + actual_span, + )?; + let dims = formal_actual_projection_dims( + input, + caller_dims, + format!("function `{}` input `{}`", function.name, input.name), + caller_actual.span().unwrap_or(actual_span), + )?; + if let Some(dims) = dims.filter(|dims| !dims.is_empty()) { + let scalars = self + .project_value_scalars( + &caller_actual, + &dims, + caller_scope, + depth + 1, + actual_span, + )? + .ok_or_else(|| { + unsupported_at( + format!( + "function `{}` input `{}` could not be projected from caller scope", + function.name, input.name + ), + actual_span, + ) + })?; + scope.full.insert(input.name.clone(), caller_actual); + scope.scalars.insert(input.name.clone(), scalars); + scope.dims.insert(input.name.clone(), dims); + continue; + } + scope.full.insert(input.name.clone(), caller_actual); + continue; + } let actual = self.substitute(&actual, &scope)?; - scope.full.insert(input.name.clone(), actual.clone()); let actual_dims = self.expr_dims_with_owner(&actual, &scope, depth + 1, actual_span)?; + let actual = if project_flattened_positional_field { + self.project_value_scalars(&actual, &[], &scope, depth + 1, actual_span)? + .and_then(|values| values.into_iter().next()) + .unwrap_or(actual) + } else { + actual + }; + scope.full.insert(input.name.clone(), actual.clone()); let dims = formal_actual_projection_dims( input, actual_dims, @@ -372,24 +738,18 @@ impl<'a> FunctionProjectionAnalysis<'a> { depth: usize, owner_span: rumoca_core::Span, ) -> Result<(), LowerError> { + seed_declared_scalar_dims(function, scope); for param in function.outputs.iter().chain(function.locals.iter()) { - if self.initialize_declared_default(param, scope, depth + 1, owner_span)? { + if self.initialize_declared_default(function, param, scope, depth + 1, owner_span)? { continue; } if param.dims.is_empty() || scope.scalars.contains_key(param.name.as_str()) { continue; } let param_span = inherited_projection_span(param.span, owner_span); - let count = scalar_count_for_dims( - ¶m.dims, - "function declared array dimensions", - param_span, - )?; - let dims = copy_projection_dims( - ¶m.dims, - "projected declared array dimension count", - param_span, - )?; + let dims = self.function_param_projection_dims(param, scope, param_span)?; + let count = + scalar_count_for_dims(&dims, "function declared array dimensions", param_span)?; scope.dims.insert(param.name.clone(), dims); let mut scalars = projection_vec_with_capacity( count, @@ -404,8 +764,10 @@ impl<'a> FunctionProjectionAnalysis<'a> { Ok(()) } + #[allow(clippy::excessive_nesting)] fn initialize_declared_default( &self, + function: &rumoca_core::Function, param: &rumoca_core::FunctionParam, scope: &mut FunctionProjectionScope, depth: usize, @@ -419,13 +781,30 @@ impl<'a> FunctionProjectionAnalysis<'a> { let value = self.substitute(default, scope)?; scope.full.insert(param.name.clone(), value.clone()); if param.dims.is_empty() { + let default_dims = self.expr_dims_with_owner(&value, scope, depth + 1, default_span)?; + let projection_dims = match default_dims { + Some(default_dims) if !default_dims.is_empty() => Some(default_dims), + _ => self.vectorized_scalar_expr_dims(default, function, scope)?, + }; + if let Some(default_dims) = projection_dims.as_deref().filter(|dims| !dims.is_empty()) { + let scalars = self + .project_value_scalars(&value, default_dims, scope, depth + 1, default_span) + .map_err(|err| err.with_fallback_span(default_span))? + .ok_or_else(|| { + unsupported_at( + format!( + "declaration binding for `{}` could not be projected", + param.name + ), + default_span, + ) + })?; + scope.dims.insert(param.name.clone(), default_dims.to_vec()); + scope.scalars.insert(param.name.clone(), scalars); + } return Ok(true); } - let dims = copy_projection_dims( - ¶m.dims, - "projected declared array dimension count", - param_span, - )?; + let dims = self.function_param_projection_dims(param, scope, param_span)?; let scalars = self .project_value_scalars(&value, &dims, scope, depth + 1, default_span) .map_err(|err| err.with_fallback_span(default_span))? @@ -443,6 +822,172 @@ impl<'a> FunctionProjectionAnalysis<'a> { Ok(true) } + fn function_param_projection_dims( + &self, + param: &rumoca_core::FunctionParam, + scope: &FunctionProjectionScope, + span: rumoca_core::Span, + ) -> Result, LowerError> { + if param.shape_expr.is_empty() { + return copy_projection_dims( + ¶m.dims, + "projected declared array dimension count", + span, + ); + } + let mut dims = projection_vec_with_capacity( + param.shape_expr.len(), + "projected declared shape expression count", + span, + )?; + for shape in ¶m.shape_expr { + let dim = self.shape_expr_dimension(shape, scope, span)?; + dims.push(dim); + } + Ok(dims) + } + + fn shape_expr_dimension( + &self, + shape: &rumoca_core::Subscript, + scope: &FunctionProjectionScope, + owner_span: rumoca_core::Span, + ) -> Result { + match shape { + rumoca_core::Subscript::Index { value, .. } => Ok(*value), + rumoca_core::Subscript::Expr { expr, span } => { + let value = self + .compile_time_scalar_in_scope(expr, scope)? + .ok_or_else(|| { + unsupported_at("function parameter shape is not compile-time bound", *span) + })?; + checked_shape_dimension(value, *span) + } + rumoca_core::Subscript::Colon { span } => Err(unsupported_at( + "function parameter shape cannot use colon during projection", + inherited_projection_span(*span, owner_span), + )), + } + } + + fn declared_dims_in_scope( + &self, + function: &rumoca_core::Function, + name: &str, + scope: &FunctionProjectionScope, + ) -> Result>, LowerError> { + Ok(self + .declared_param_dims_in_scope(function, name, scope)? + .filter(|dims| !dims.is_empty())) + } + + fn declared_param_dims_in_scope( + &self, + function: &rumoca_core::Function, + name: &str, + scope: &FunctionProjectionScope, + ) -> Result>, LowerError> { + let Some(param) = function + .outputs + .iter() + .chain(function.locals.iter()) + .chain(function.inputs.iter()) + .find(|param| param.name == name) + else { + return Ok(None); + }; + let dims = self.function_param_projection_dims(param, scope, param.span)?; + Ok(Some(dims)) + } + + fn flattened_projection_actual_has_field( + &self, + expr: &rumoca_core::Expression, + field: &str, + scope: &FunctionProjectionScope, + ) -> bool { + match expr { + rumoca_core::Expression::FunctionCall { + name, + is_constructor, + .. + } => { + self.is_record_constructor_call(name, *is_constructor) + || self + .dae_model + .symbols + .functions + .get(name.var_name()) + .is_some_and(|function| { + matches!( + function.outputs.as_slice(), + [output] + if output.type_class == Some(rumoca_core::ClassType::Record) + ) + }) + } + rumoca_core::Expression::VarRef { + name, + subscripts, + span, + } if subscripts.is_empty() => { + let field_key = format!("{}.{field}", name.as_str()); + let flattened_key = format!("{}_{field}", name.as_str()); + scope.full.contains_key(&field_key) + || scope.scalars.contains_key(&field_key) + || scope.dims.contains_key(&field_key) + || scope.full.contains_key(&flattened_key) + || scope.scalars.contains_key(&flattened_key) + || scope.dims.contains_key(&flattened_key) + } + _ => false, + } + } + + fn flattened_projection_actual_field_value( + &self, + expr: &rumoca_core::Expression, + field: &str, + caller_scope: Option<&FunctionProjectionScope>, + ) -> Option { + let caller_scope = caller_scope?; + let rumoca_core::Expression::VarRef { + name, + subscripts, + span, + } = expr + else { + return None; + }; + if !subscripts.is_empty() { + return None; + } + let field_key = format!("{}.{field}", name.as_str()); + let flattened_key = format!("{}_{field}", name.as_str()); + self.projected_scope_value_for_key(caller_scope, &field_key, *span) + .or_else(|| self.projected_scope_value_for_key(caller_scope, &flattened_key, *span)) + } + + fn projected_scope_value_for_key( + &self, + scope: &FunctionProjectionScope, + key: &str, + span: rumoca_core::Span, + ) -> Option { + if let Some(value) = scope.full.get(key) { + return Some(value.clone().with_span(span)); + } + let values = scope.scalars.get(key)?; + match values.as_slice() { + [value] => Some(value.clone().with_span(span)), + _ => Some(rumoca_core::Expression::Array { + elements: values.clone(), + is_matrix: false, + span, + }), + } + } + fn insert_input_scalar_projection( &self, input: &rumoca_core::FunctionParam, @@ -517,6 +1062,31 @@ impl<'a> FunctionProjectionAnalysis<'a> { index_depth: 0, }, ), + rumoca_core::Statement::While { block, span } => self.apply_static_while_statement( + function, + block, + scope, + projected, + depth, + inherited_projection_span(*span, owner_span), + ), + rumoca_core::Statement::FunctionCall { + comp, + args, + outputs, + span, + } if !outputs.is_empty() => self.apply_function_call_statement( + FunctionCallStatementProjection { + function, + comp, + args, + output_targets: outputs, + span: inherited_projection_span(*span, owner_span), + depth, + }, + scope, + projected, + ), statement if is_ignorable_projection_statement(statement) => Ok(()), _ => Err(unsupported_at( format!( @@ -528,49 +1098,362 @@ impl<'a> FunctionProjectionAnalysis<'a> { } } - fn apply_assignment( + fn apply_function_call_statement( &self, - function: &rumoca_core::Function, - statement: &rumoca_core::Statement, + request: FunctionCallStatementProjection<'_>, scope: &mut FunctionProjectionScope, projected: &mut Vec, - depth: usize, - owner_span: rumoca_core::Span, ) -> Result<(), LowerError> { - let rumoca_core::Statement::Assignment { comp, value, .. } = statement else { - return Err(LowerError::contract_violation( - "non-assignment statement reached assignment projection", - owner_span, + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference(request.comp.clone()), + args: request.args.to_vec(), + is_constructor: false, + span: request.span, + }; + let outputs = self.function_call_outputs_with_projection_scope( + &call, + request.depth + 1, + request.span, + Some(scope), + )?; + let Some(outputs) = outputs else { + return Err(unsupported_at( + format!( + "function `{}` call statement to `{}` could not be projected", + request.function.name, + request.comp.to_var_name() + ), + request.span, )); }; - let assignment_span = inherited_projection_source_span(statement.source_span(), owner_span); - let comp = self.substitute_component_reference(comp, scope)?; - let target = projection_assignment_target(&comp)?; - let value = self.scalar_assignment_value(value, scope, depth + 1)?; - if exceeds_projection_node_budget(&value) { - return Err(projection_budget_exceeded(function)); + let callee_name = request.comp.to_var_name(); + let Some(callee) = self.dae_model.symbols.functions.get(&callee_name) else { + return Err(LowerError::MissingFunction { + name: callee_name.to_string(), + }); + }; + if callee.outputs.len() != request.output_targets.len() { + return Err(LowerError::contract_violation( + format!( + "function call statement to `{callee_name}` has {} outputs for {} targets", + callee.outputs.len(), + request.output_targets.len() + ), + request.span, + )); } - if let Some(indices) = target.indices.as_deref() { - let indexed_span = inherited_projection_span(target.span, assignment_span); - self.apply_indexed_assignment( - function, - IndexedAssignment { - target: &target.base, - indices, - value: &value, - span: indexed_span, - depth: depth + 1, - }, + for (idx, (target, output_param)) in request + .output_targets + .iter() + .zip(callee.outputs.iter()) + .enumerate() + { + let selected = self.selected_projected_statement_outputs( + &outputs, + callee, + output_param, + idx, + request.span, + )?; + self.apply_projected_function_statement_output( + request.function, + output_param, + target, + selected, scope, + projected, + request.depth + 1, + request.span, )?; - return Ok(()); } - let target_span = inherited_projection_span(target.span, assignment_span); - let target = target.base; - scope.full.insert(target.clone(), value.clone()); - let value_span = value.span().unwrap_or(target_span); - if let Some(record_outputs) = self.record_constructor_outputs(&value, scope, depth + 1)? { - if is_function_output_target(function, &target) { + Ok(()) + } + + fn selected_projected_statement_outputs( + &self, + outputs: &[ProjectedFunctionOutput], + callee: &rumoca_core::Function, + output_param: &rumoca_core::FunctionParam, + output_idx: usize, + span: rumoca_core::Span, + ) -> Result, LowerError> { + if callee.outputs.len() == 1 { + return Ok(outputs.to_vec()); + } + let mut selected = projection_vec_with_capacity( + outputs.len(), + "selected projected function statement output count", + span, + )?; + for output in outputs { + let Some((head, tail)) = output.field_path.split_first() else { + return Err(LowerError::contract_violation( + format!( + "multi-output function `{}` projected output {} for `{}` has no output selector", + callee.name, + output_idx + 1, + output_param.name + ), + span, + )); + }; + if head == &output_param.name { + selected.push(ProjectedFunctionOutput { + field_path: tail.to_vec(), + selector_indices: output.selector_indices.clone(), + expr: output.expr.clone(), + }); + } + } + Ok(selected) + } + + #[allow(clippy::too_many_arguments)] + fn apply_projected_function_statement_output( + &self, + function: &rumoca_core::Function, + output_param: &rumoca_core::FunctionParam, + target: &rumoca_core::ComponentReference, + selected: Vec, + scope: &mut FunctionProjectionScope, + projected: &mut Vec, + depth: usize, + span: rumoca_core::Span, + ) -> Result<(), LowerError> { + if selected.is_empty() { + return self.apply_empty_projected_function_statement_output( + function, + output_param, + target, + scope, + span, + ); + } + if selected.iter().any(|output| !output.field_path.is_empty()) { + return self.apply_projected_function_statement_output_as_assignments( + function, target, selected, scope, projected, depth, span, + ); + } + let target = self.substitute_component_reference(target, scope)?; + let assignment_target = projection_assignment_target(&target)?; + if let Some(indices) = assignment_target.indices.as_deref() { + if selected.len() != 1 { + return Err(unsupported_at( + "indexed function call statement target cannot receive array output", + assignment_target.span, + )); + } + self.apply_indexed_assignment( + function, + IndexedAssignment { + target: &assignment_target.base, + indices, + value: &selected[0].expr, + span: assignment_target.span, + depth, + }, + scope, + )?; + return Ok(()); + } + let selector_indices = selected + .iter() + .map(|output| output.selector_indices.as_slice()) + .collect::>(); + let projected_dims = + selector_dims_from_indices(&selector_indices, span)?.ok_or_else(|| { + LowerError::contract_violation( + format!( + "function call statement output `{}` has non-dense selector projection", + output_param.name + ), + span, + ) + })?; + let value_dims = Some(projected_dims.clone()); + let dims = assignment_projection_dims( + function.name.as_str(), + &assignment_target.base, + self.declared_param_dims_in_scope(function, &assignment_target.base, scope)?, + value_dims, + span, + )? + .unwrap_or(projected_dims); + let scalars = selected + .into_iter() + .map(|output| output.expr) + .collect::>(); + scope + .scalars + .insert(assignment_target.base.clone(), scalars.clone()); + scope + .dims + .insert(assignment_target.base.clone(), dims.clone()); + scope.full.insert( + assignment_target.base.clone(), + rumoca_core::Expression::Array { + elements: scalars.clone(), + is_matrix: false, + span, + }, + ); + if let Some(output_param) = function + .outputs + .iter() + .find(|output| output.name == assignment_target.base) + { + let outputs = self.tag_function_output_projection( + function, + output_param, + project_target_scalar_outputs(&dims, scalars, span)?, + span, + )?; + append_projected_outputs( + projected, + outputs, + "function call statement projected output count", + span, + )?; + } + Ok(()) + } + + fn apply_empty_projected_function_statement_output( + &self, + function: &rumoca_core::Function, + output_param: &rumoca_core::FunctionParam, + target: &rumoca_core::ComponentReference, + scope: &mut FunctionProjectionScope, + span: rumoca_core::Span, + ) -> Result<(), LowerError> { + let target = self.substitute_component_reference(target, scope)?; + let assignment_target = projection_assignment_target(&target)?; + if assignment_target.indices.is_some() { + return Err(unsupported_at( + "indexed function call statement target cannot receive empty array output", + assignment_target.span, + )); + } + let Some(declared) = + self.declared_param_dims_in_scope(function, &assignment_target.base, scope)? + else { + return Err(unsupported_at( + format!( + "function call statement output `{}` projected no scalar values for `{}`", + output_param.name, assignment_target.base + ), + span, + )); + }; + if scalar_count_for_dims( + &declared, + "empty function statement output dimensions", + span, + )? != 0 + { + return Err(unsupported_at( + format!( + "function call statement output `{}` projected no scalar values for non-empty target `{}`", + output_param.name, assignment_target.base + ), + span, + )); + } + let scalars = Vec::new(); + scope + .scalars + .insert(assignment_target.base.clone(), scalars.clone()); + scope.dims.insert(assignment_target.base.clone(), declared); + scope.full.insert( + assignment_target.base, + rumoca_core::Expression::Array { + elements: scalars, + is_matrix: false, + span, + }, + ); + Ok(()) + } + + #[allow(clippy::too_many_arguments)] + fn apply_projected_function_statement_output_as_assignments( + &self, + function: &rumoca_core::Function, + target: &rumoca_core::ComponentReference, + selected: Vec, + scope: &mut FunctionProjectionScope, + projected: &mut Vec, + depth: usize, + span: rumoca_core::Span, + ) -> Result<(), LowerError> { + if selected.len() != 1 { + return Err(unsupported_at( + "record-valued function call statement output cannot be projected as an array assignment", + span, + )); + } + let assignment = rumoca_core::Statement::Assignment { + comp: target.clone(), + value: selected + .into_iter() + .next() + .expect("selected len checked") + .expr, + span, + }; + self.apply_assignment(function, &assignment, scope, projected, depth, span) + } + + #[allow(clippy::too_many_lines)] + fn apply_assignment( + &self, + function: &rumoca_core::Function, + statement: &rumoca_core::Statement, + scope: &mut FunctionProjectionScope, + projected: &mut Vec, + depth: usize, + owner_span: rumoca_core::Span, + ) -> Result<(), LowerError> { + let rumoca_core::Statement::Assignment { comp, value, .. } = statement else { + return Err(LowerError::contract_violation( + "non-assignment statement reached assignment projection", + owner_span, + )); + }; + let original_value = value; + let assignment_span = inherited_projection_source_span(statement.source_span(), owner_span); + let comp = self.substitute_component_reference(comp, scope)?; + let target = projection_assignment_target(&comp)?; + let value = self.scalar_assignment_value(value, scope, depth + 1)?; + if exceeds_projection_node_budget(&value) { + return Err(projection_budget_exceeded(function)); + } + if let Some(indices) = target.indices.as_deref() { + let indexed_span = inherited_projection_span(target.span, assignment_span); + self.apply_indexed_assignment( + function, + IndexedAssignment { + target: &target.base, + indices, + value: &value, + span: indexed_span, + depth: depth + 1, + }, + scope, + )?; + return Ok(()); + } + let target_span = inherited_projection_span(target.span, assignment_span); + let target = target.base; + if let Some(record_outputs) = self.record_constructor_outputs(&value, scope, depth + 1)? { + if let Some(output_param) = function.outputs.iter().find(|output| output.name == target) + { + let record_outputs = self.tag_function_output_projection( + function, + output_param, + record_outputs, + target_span, + )?; append_projected_outputs( projected, record_outputs, @@ -578,17 +1461,60 @@ impl<'a> FunctionProjectionAnalysis<'a> { target_span, )?; } + scope.full.insert(target, value); return Ok(()); } - let dims = assignment_projection_dims( + if function + .outputs + .iter() + .any(|output| output.name == target && Self::function_output_is_record_like(output)) + { + scope.full.insert(target, value); + return Ok(()); + } + let value_span = value.span().unwrap_or(target_span); + let substituted_dims = self.expr_dims_with_owner(&value, scope, depth + 1, value_span)?; + let (value_dims, projection_value) = match substituted_dims { + Some(dims) => (Some(dims), &value), + None => { + let value_dims = + self.expr_dims_with_owner(original_value, scope, depth + 1, value_span)?; + let projection_value = if self + .original_projection_preserves_declared_scalar_shape(original_value, scope) + { + original_value + } else { + &value + }; + (value_dims, projection_value) + } + }; + let vectorized_scalar_assignment = self.vectorized_scalar_assignment_dims( function, &target, - self.expr_dims_with_owner(&value, scope, depth + 1, value_span)?, - value_span, + original_value, + value_dims.as_deref(), + scope, )?; + let declared = self.declared_param_dims_in_scope(function, &target, scope)?; + let scalar_assignment = vectorized_scalar_assignment.is_none() + && value_dims.as_ref().is_some_and(Vec::is_empty) + && declared.as_ref().is_some_and(Vec::is_empty) + && function.locals.iter().any(|local| local.name == target); + let dims = if let Some(dims) = vectorized_scalar_assignment { + Some(dims) + } else { + assignment_projection_dims( + function.name.as_str(), + &target, + declared, + value_dims, + value_span, + )? + }; if let Some(dims) = dims.filter(|dims| !dims.is_empty()) { let scalars = self - .project_value_scalars(&value, &dims, scope, depth + 1, value_span) + .project_value_scalars(projection_value, &dims, scope, depth + 1, value_span) .map_err(|err| err.with_fallback_span(value_span))? .ok_or_else(|| { unsupported_at( @@ -607,8 +1533,14 @@ impl<'a> FunctionProjectionAnalysis<'a> { target.clone(), copy_projection_dims(&dims, "scalar assignment dimension count", target_span)?, ); - if is_function_output_target(function, &target) { - let outputs = project_target_scalar_outputs(&dims, scalars, target_span)?; + if let Some(output_param) = function.outputs.iter().find(|output| output.name == target) + { + let outputs = self.tag_function_output_projection( + function, + output_param, + project_target_scalar_outputs(&dims, scalars, target_span)?, + target_span, + )?; append_projected_outputs( projected, outputs, @@ -617,9 +1549,226 @@ impl<'a> FunctionProjectionAnalysis<'a> { )?; } } + let value = if scalar_assignment { + self.normalize_projected_scalar_output_expr(&value, scope, depth + 1, value_span)? + } else { + value + }; + scope.full.insert(target, value); Ok(()) } + fn vectorized_scalar_assignment_dims( + &self, + function: &rumoca_core::Function, + target: &str, + original_value: &rumoca_core::Expression, + value_dims: Option<&[i64]>, + scope: &FunctionProjectionScope, + ) -> Result>, LowerError> { + if self + .declared_dims_in_scope(function, target, scope)? + .is_some_and(|dims| !dims.is_empty()) + { + return Ok(None); + } + let Some(value_dims) = value_dims else { + return Ok(None); + }; + if value_dims.is_empty() + || !self.expr_depends_on_vectorized_scalar(original_value, function, scope)? + { + return Ok(None); + } + Ok(Some(value_dims.to_vec())) + } + + fn original_projection_preserves_declared_scalar_shape( + &self, + expr: &rumoca_core::Expression, + scope: &FunctionProjectionScope, + ) -> bool { + let has_boundary = expr.contains_subexpression(|node| match node { + rumoca_core::Expression::VarRef { subscripts, .. } => !subscripts.is_empty(), + rumoca_core::Expression::Binary { op, .. } => !is_elementwise_binary_projection_op(op), + rumoca_core::Expression::BuiltinCall { function, .. } => { + !is_elementwise_builtin_projection(function) + } + rumoca_core::Expression::Unary { .. } + | rumoca_core::Expression::If { .. } + | rumoca_core::Expression::Array { .. } + | rumoca_core::Expression::Tuple { .. } + | rumoca_core::Expression::Range { .. } + | rumoca_core::Expression::Literal { .. } + | rumoca_core::Expression::Empty { .. } => false, + _ => true, + }); + if has_boundary { + return false; + } + expr.contains_subexpression(|node| { + matches!( + node, + rumoca_core::Expression::VarRef { name, subscripts, .. } + if subscripts.is_empty() + && scope.dims.get(name.as_str()).is_some_and(Vec::is_empty) + && scope.full.get(name.as_str()).is_some_and( + |replacement| !is_same_plain_var_ref(replacement, name.as_str()) + ) + ) + }) + } + + #[allow(clippy::excessive_nesting)] + fn expr_depends_on_vectorized_scalar( + &self, + expr: &rumoca_core::Expression, + function: &rumoca_core::Function, + scope: &FunctionProjectionScope, + ) -> Result { + match expr { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } if subscripts.is_empty() => Ok(scope + .dims + .get(name.as_str()) + .is_some_and(|dims| !dims.is_empty()) + && declared_param_dims(function, name.as_str())? + .is_some_and(|dims| dims.is_empty())), + rumoca_core::Expression::Unary { rhs, .. } => { + self.expr_depends_on_vectorized_scalar(rhs, function, scope) + } + rumoca_core::Expression::Binary { lhs, rhs, .. } => Ok(self + .expr_depends_on_vectorized_scalar(lhs, function, scope)? + || self.expr_depends_on_vectorized_scalar(rhs, function, scope)?), + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + for (condition, branch) in branches { + if self.expr_depends_on_vectorized_scalar(condition, function, scope)? + || self.expr_depends_on_vectorized_scalar(branch, function, scope)? + { + return Ok(true); + } + } + self.expr_depends_on_vectorized_scalar(else_branch, function, scope) + } + rumoca_core::Expression::Array { elements, .. } + | rumoca_core::Expression::Tuple { elements, .. } => { + elements.iter().try_fold(false, |found, element| { + Ok( + found + || self.expr_depends_on_vectorized_scalar(element, function, scope)?, + ) + }) + } + rumoca_core::Expression::Range { + start, step, end, .. + } => Ok( + self.expr_depends_on_vectorized_scalar(start, function, scope)? + || step + .as_deref() + .map(|step| self.expr_depends_on_vectorized_scalar(step, function, scope)) + .transpose()? + .unwrap_or(false) + || self.expr_depends_on_vectorized_scalar(end, function, scope)?, + ), + rumoca_core::Expression::BuiltinCall { args, .. } + | rumoca_core::Expression::FunctionCall { args, .. } => { + args.iter().try_fold(false, |found, arg| { + Ok(found || self.expr_depends_on_vectorized_scalar(arg, function, scope)?) + }) + } + rumoca_core::Expression::FieldAccess { base, .. } => { + self.expr_depends_on_vectorized_scalar(base, function, scope) + } + _ => Ok(false), + } + } + + fn vectorized_scalar_expr_dims( + &self, + expr: &rumoca_core::Expression, + function: &rumoca_core::Function, + scope: &FunctionProjectionScope, + ) -> Result>, LowerError> { + let mut dims = None; + self.collect_vectorized_scalar_expr_dims(expr, function, scope, &mut dims)?; + Ok(dims) + } + + fn collect_vectorized_scalar_expr_dims( + &self, + expr: &rumoca_core::Expression, + function: &rumoca_core::Function, + scope: &FunctionProjectionScope, + dims: &mut Option>, + ) -> Result<(), LowerError> { + match expr { + rumoca_core::Expression::VarRef { + name, + subscripts, + span, + } if subscripts.is_empty() => { + if let Some(candidate) = scope.dims.get(name.as_str()) + && !candidate.is_empty() + && declared_param_dims(function, name.as_str())? + .is_some_and(|declared| declared.is_empty()) + { + merge_vectorized_scalar_dims(dims, candidate, name.as_str(), *span)?; + } + Ok(()) + } + rumoca_core::Expression::Unary { rhs, .. } => { + self.collect_vectorized_scalar_expr_dims(rhs, function, scope, dims) + } + rumoca_core::Expression::Binary { lhs, rhs, .. } => { + self.collect_vectorized_scalar_expr_dims(lhs, function, scope, dims)?; + self.collect_vectorized_scalar_expr_dims(rhs, function, scope, dims) + } + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + for (condition, branch) in branches { + self.collect_vectorized_scalar_expr_dims(condition, function, scope, dims)?; + self.collect_vectorized_scalar_expr_dims(branch, function, scope, dims)?; + } + self.collect_vectorized_scalar_expr_dims(else_branch, function, scope, dims) + } + rumoca_core::Expression::Array { elements, .. } + | rumoca_core::Expression::Tuple { elements, .. } => { + for element in elements { + self.collect_vectorized_scalar_expr_dims(element, function, scope, dims)?; + } + Ok(()) + } + rumoca_core::Expression::Range { + start, step, end, .. + } => { + self.collect_vectorized_scalar_expr_dims(start, function, scope, dims)?; + if let Some(step) = step { + self.collect_vectorized_scalar_expr_dims(step, function, scope, dims)?; + } + self.collect_vectorized_scalar_expr_dims(end, function, scope, dims) + } + rumoca_core::Expression::BuiltinCall { args, .. } + | rumoca_core::Expression::FunctionCall { args, .. } => { + for arg in args { + self.collect_vectorized_scalar_expr_dims(arg, function, scope, dims)?; + } + Ok(()) + } + rumoca_core::Expression::FieldAccess { base, .. } => { + self.collect_vectorized_scalar_expr_dims(base, function, scope, dims) + } + _ => Ok(()), + } + } + fn apply_indexed_assignment( &self, function: &rumoca_core::Function, @@ -640,7 +1789,7 @@ impl<'a> FunctionProjectionAnalysis<'a> { let dims = if let Some(dims) = scope.dims.get(target) { copy_projection_dims(dims, "indexed assignment scope dimension count", span)? } else { - declared_dims(function, target)? + self.declared_dims_in_scope(function, target, scope)? .ok_or_else(|| guarded_assignment_without_base(target, span))? }; let flat_index = flat_index_from_indices( @@ -690,6 +1839,9 @@ impl<'a> FunctionProjectionAnalysis<'a> { ) -> Result { if let Some((name, subscripts, span)) = indexed_var_selection(value) && let Some(values) = scope.scalars.get(name.as_str()) + && self + .expr_dims_with_owner(value, scope, depth + 1, span)? + .is_some_and(|dims| dims.is_empty()) { let dims = scope.dims.get(name.as_str()).ok_or_else(|| { LowerError::contract_violation( @@ -714,6 +1866,29 @@ impl<'a> FunctionProjectionAnalysis<'a> { ); } let value = self.substitute(value, scope)?; + if matches!( + value, + rumoca_core::Expression::FunctionCall { + is_constructor: false, + .. + } + ) { + let span = value + .span() + .unwrap_or_else(rumoca_core::Span::source_free_serde_default); + match self.function_call_outputs_with_projection_scope( + &value, + depth + 1, + span, + Some(scope), + )? { + Some(outputs) if outputs.len() == 1 => { + let output = only_projected_scalar_assignment_output(outputs, span)?; + return Ok(output.expr.with_span(span)); + } + _ => {} + } + } let rumoca_core::Expression::VarRef { name, subscripts, @@ -751,6 +1926,33 @@ impl<'a> FunctionProjectionAnalysis<'a> { ) } + fn apply_static_while_statement( + &self, + function: &rumoca_core::Function, + block: &rumoca_core::StatementBlock, + scope: &mut FunctionProjectionScope, + projected: &mut Vec, + depth: usize, + span: rumoca_core::Span, + ) -> Result<(), LowerError> { + for _ in 0..MAX_STATIC_WHILE_PROJECTION_ITERATIONS { + let condition = self.substitute(&block.cond, scope)?; + let Some(value) = self.compile_time_scalar_in_scope(&condition, scope)? else { + return Err(LowerError::DynamicWhileProjection { + function: function.name.to_string(), + span, + }); + }; + if value == 0.0 { + return Ok(()); + } + for statement in &block.stmts { + self.apply_statement(function, statement, scope, projected, depth + 1, span)?; + } + } + Err(projection_budget_exceeded(function).with_fallback_span(span)) + } + fn apply_if_statement( &self, function: &rumoca_core::Function, @@ -846,7 +2048,7 @@ impl<'a> FunctionProjectionAnalysis<'a> { ) -> Result, LowerError> { for block in cond_blocks { let condition = self.substitute(&block.cond, scope)?; - let Some(value) = self.compile_time_scalar(&condition) else { + let Some(value) = self.compile_time_scalar_in_scope(&condition, scope)? else { return Ok(None); }; if value != 0.0 { @@ -866,6 +2068,15 @@ impl<'a> FunctionProjectionAnalysis<'a> { ) -> Result { let mut merged = entry_scope.clone(); for name in projection_scope_names(entry_scope, branch_scopes, else_scope, span)? { + if !entry_scope.scalars.contains_key(&name) + && !entry_scope.full.contains_key(&name) + && !entry_scope.dims.contains_key(&name) + && !else_scope.scalars.contains_key(&name) + && !else_scope.full.contains_key(&name) + && !else_scope.dims.contains_key(&name) + { + continue; + } if projection_scope_has_scalars(&name, entry_scope, branch_scopes, else_scope) { let values = self.merged_if_scalar_values( &name, @@ -882,148 +2093,568 @@ impl<'a> FunctionProjectionAnalysis<'a> { else_scope, span, )?; - merged.scalars.insert(name.clone(), values); - merged.dims.insert(name.clone(), dims); + merged.scalars.insert(name.clone(), values); + merged.dims.insert(name.clone(), dims); + } + if projection_scope_has_full(&name, entry_scope, branch_scopes, else_scope) { + let value = merged_if_full_value( + &name, + entry_scope, + branch_conditions, + branch_scopes, + else_scope, + span, + )?; + merged.full.insert(name, value); + } + } + Ok(merged) + } + + fn merged_if_scalar_values( + &self, + name: &str, + entry_scope: &FunctionProjectionScope, + branch_conditions: &[rumoca_core::Expression], + branch_scopes: &[FunctionProjectionScope], + else_scope: &FunctionProjectionScope, + span: rumoca_core::Span, + ) -> Result, LowerError> { + let base_values = else_scope + .scalars + .get(name) + .or_else(|| entry_scope.scalars.get(name)) + .ok_or_else(|| guarded_assignment_without_base(name, span))?; + let mut merged = projection_vec_with_capacity( + base_values.len(), + "if projection scalar value count", + span, + )?; + for value in base_values { + merged.push(value.clone()); + } + for (condition, branch_scope) in branch_conditions.iter().zip(branch_scopes.iter()).rev() { + let Some(branch_values) = branch_scope.scalars.get(name) else { + continue; + }; + if branch_values.len() != merged.len() { + return Err(LowerError::contract_violation( + format!( + "if-statement assignment to `{name}` has mismatched scalar widths: branch {}, else/current {}", + branch_values.len(), + merged.len() + ), + span, + )); + } + let mut next_merged = projection_vec_with_capacity( + branch_values.len(), + "if projection merged scalar value count", + span, + )?; + for (branch, fallback) in branch_values.iter().cloned().zip(merged) { + next_merged.push(rumoca_core::Expression::If { + branches: single_projection_branch(condition.clone(), branch, span)?, + else_branch: Box::new(fallback), + span, + }); + } + merged = next_merged; + } + Ok(merged) + } + + #[allow(clippy::excessive_nesting)] + fn projected_outputs_from_scope( + &self, + function: &rumoca_core::Function, + scope: &FunctionProjectionScope, + depth: usize, + owner_span: rumoca_core::Span, + ) -> Result>, LowerError> { + let function_span = inherited_projection_span(function.span, owner_span); + let mut outputs = projection_vec_with_capacity( + function.outputs.len(), + "projected function output count", + function_span, + )?; + for output in &function.outputs { + let output_span = inherited_projection_span(output.span, function_span); + let output_dims = self.function_param_projection_dims(output, scope, output_span)?; + if let Some(values) = scope.scalars.get(output.name.as_str()) { + let dims = if output_dims.is_empty() { + match scope.dims.get(output.name.as_str()) { + Some(dims) => dims.clone(), + None => function_outputs_dims(values.len(), output_span)?, + } + } else { + output_dims + }; + let projected = self.tag_function_output_projection( + function, + output, + project_target_scalar_outputs(&dims, values.clone(), output_span)?, + output_span, + )?; + append_projected_outputs( + &mut outputs, + projected, + "projected scalar output count", + output_span, + )?; + continue; + } + if let Some(expr) = scope.full.get(output.name.as_str()) { + let projected = self.project_output_expr( + function, + output, + expr, + scope, + depth + 1, + output_span, + &output_dims, + )?; + let projected = + self.tag_function_output_projection(function, output, projected, output_span)?; + append_projected_outputs( + &mut outputs, + projected, + "projected full output count", + output_span, + )?; + } + } + Ok((!outputs.is_empty()).then_some(outputs)) + } + + #[allow(clippy::excessive_nesting, clippy::too_many_arguments)] + fn project_output_expr( + &self, + function: &rumoca_core::Function, + output: &rumoca_core::FunctionParam, + expr: &rumoca_core::Expression, + scope: &FunctionProjectionScope, + depth: usize, + output_span: rumoca_core::Span, + output_dims: &[i64], + ) -> Result, LowerError> { + if output_dims.is_empty() { + if Self::function_output_is_record_like(output) { + let expr = self.substitute(expr, scope)?; + if let Some(projected) = self.project_record_like_output_expr( + output, + &expr, + scope, + depth + 1, + output_span, + )? { + return Ok(projected); + } + let mut projected = projection_vec_with_capacity( + 1, + "record function output projection count", + output_span, + )?; + projected.push(ProjectedFunctionOutput { + field_path: Vec::new(), + selector_indices: Vec::new(), + expr, + }); + return Ok(projected); + } + let expr_span = inherited_projection_source_span(expr.span(), output_span); + if let Some(expr_dims) = self.expr_dims_with_owner(expr, scope, depth + 1, expr_span)? + && !expr_dims.is_empty() + { + let values = self + .project_value_scalars(expr, &expr_dims, scope, depth + 1, expr_span) + .map_err(|err| err.with_fallback_span(expr_span))? + .ok_or_else(|| { + unsupported_at( + format!("function output `{}` could not be projected", output.name), + expr_span, + ) + })?; + return project_target_scalar_outputs(&expr_dims, values, output_span); + } + if let Some(expr_dims) = self.vectorized_scalar_expr_dims(expr, function, scope)? { + let values = self + .project_value_scalars(expr, &expr_dims, scope, depth + 1, expr_span) + .map_err(|err| err.with_fallback_span(expr_span))? + .ok_or_else(|| { + unsupported_at( + format!("function output `{}` could not be projected", output.name), + expr_span, + ) + })?; + return project_target_scalar_outputs(&expr_dims, values, output_span); + } + let expr = expr.clone(); + let expr = if matches!( + expr, + rumoca_core::Expression::FunctionCall { + is_constructor: false, + .. + } + ) { + match self.function_call_outputs_with_projection_scope( + &expr, + depth + 1, + output_span, + Some(scope), + )? { + Some(outputs) if outputs.len() == 1 => outputs + .into_iter() + .next() + .map(|output| output.expr) + .unwrap_or(expr), + _ => expr, + } + } else { + expr + }; + let expr = + self.normalize_projected_scalar_output_expr(&expr, scope, depth + 1, output_span)?; + let mut projected = projection_vec_with_capacity( + 1, + "scalar function output projection count", + output_span, + )?; + projected.push(ProjectedFunctionOutput { + field_path: Vec::new(), + selector_indices: Vec::new(), + expr: expr.clone(), + }); + return Ok(projected); + } + let expr_span = inherited_projection_source_span(expr.span(), output_span); + let values = self + .project_value_scalars(expr, output_dims, scope, depth, expr_span) + .map_err(|err| err.with_fallback_span(expr_span))? + .ok_or_else(|| { + unsupported_at( + format!("function output `{}` could not be projected", output.name), + expr_span, + ) + })?; + project_target_scalar_outputs(output_dims, values, output_span) + } + + fn tag_function_output_projection( + &self, + function: &rumoca_core::Function, + output: &rumoca_core::FunctionParam, + projected: Vec, + span: rumoca_core::Span, + ) -> Result, LowerError> { + if function.outputs.len() <= 1 { + return Ok(projected); + } + let prefix = single_field_path(&output.name, span)?; + let mut tagged = projection_vec_with_capacity( + projected.len(), + "tagged projected function output count", + span, + )?; + for mut output in projected { + let mut field_path = prefix.clone(); + field_path.append(&mut output.field_path); + output.field_path = field_path; + tagged.push(output); + } + Ok(tagged) + } + + fn project_record_like_output_expr( + &self, + output: &rumoca_core::FunctionParam, + expr: &rumoca_core::Expression, + scope: &FunctionProjectionScope, + depth: usize, + output_span: rumoca_core::Span, + ) -> Result>, LowerError> { + let constructor_name = rumoca_core::VarName::new(&output.type_name); + let Some(constructor) = self.dae_model.symbols.functions.get(&constructor_name) else { + return Ok(None); + }; + if !is_record_constructor_signature(&output.type_name, constructor) + && !constructor.is_constructor + { + return Ok(None); + } + let mut projected = projection_vec_with_capacity( + constructor.inputs.len(), + "record-like function output field count", + output_span, + )?; + for input in &constructor.inputs { + let input_span = inherited_projection_span(input.span, output_span); + let Some(field_expr) = + self.project_record_field_value(expr, &input.name, scope, input_span)? + else { + return Ok(None); + }; + let Some((projection_dims, scalars)) = self.optional_constructor_input_scalars( + &field_expr, + input, + scope, + depth, + input_span, + )? + else { + return Ok(None); + }; + reserve_projection_capacity( + &mut projected, + scalars.len(), + "record-like function output scalar count", + input_span, + )?; + for (idx, expr) in scalars.into_iter().enumerate() { + projected.push(ProjectedFunctionOutput { + field_path: single_field_path(&input.name, input_span)?, + selector_indices: required_flat_index_to_subscripts( + &projection_dims, + idx, + input_span, + )?, + expr, + }); + } + } + Ok(Some(projected)) + } + + fn function_output_is_record_like(output: &rumoca_core::FunctionParam) -> bool { + !output.type_name.is_empty() + && !rumoca_core::qualified_type_name_matches(&output.type_name, "Real") + && !rumoca_core::qualified_type_name_matches(&output.type_name, "Integer") + && !rumoca_core::qualified_type_name_matches(&output.type_name, "Boolean") + && !rumoca_core::qualified_type_name_matches(&output.type_name, "String") + } + + fn normalize_projected_scalar_output_expr( + &self, + expr: &rumoca_core::Expression, + scope: &FunctionProjectionScope, + depth: usize, + span: rumoca_core::Span, + ) -> Result { + match expr { + rumoca_core::Expression::Binary { + op, + lhs, + rhs, + span: expr_span, + } => { + let span = inherited_projection_span(*expr_span, span); + let lhs = self.normalize_projected_scalar_output_expr(lhs, scope, depth, span)?; + let rhs = self.normalize_projected_scalar_output_expr(rhs, scope, depth, span)?; + let ctx = projection_value_ctx(&[], 0, scope, depth, span); + if is_mul(op) + && let Some(projected) = self.project_tensor_product(&lhs, &rhs, &ctx)? + { + return Ok(projected); + } + let op = self.projected_binary_op(op, &lhs, &rhs, scope, depth, span)?; + Ok(rumoca_core::Expression::Binary { + op, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span, + }) + } + rumoca_core::Expression::BuiltinCall { + function, + args, + span: expr_span, + } if is_elementwise_builtin_projection(function) => self + .normalize_projected_scalar_builtin_args( + *function, + args, + scope, + depth, + inherited_projection_span(*expr_span, span), + ), + rumoca_core::Expression::If { + branches, + else_branch, + span: expr_span, + } => { + let span = inherited_projection_span(*expr_span, span); + let mut normalized_branches = projection_vec_with_capacity( + branches.len(), + "projected scalar output if branch count", + span, + )?; + for (condition, value) in branches { + normalized_branches.push(( + self.substitute(condition, scope)?, + self.normalize_projected_scalar_output_expr(value, scope, depth, span)?, + )); + } + let normalized_else = + self.normalize_projected_scalar_output_expr(else_branch, scope, depth, span)?; + Ok(rumoca_core::Expression::If { + branches: normalized_branches, + else_branch: Box::new(normalized_else), + span, + }) } - if projection_scope_has_full(&name, entry_scope, branch_scopes, else_scope) { - let value = merged_if_full_value( - &name, - entry_scope, - branch_conditions, - branch_scopes, - else_scope, + rumoca_core::Expression::VarRef { .. } => { + let substituted = self.substitute(expr, scope)?; + if substituted == *expr { + Ok(expr.clone()) + } else { + self.normalize_projected_scalar_output_expr( + &substituted, + scope, + depth + 1, + span, + ) + } + } + rumoca_core::Expression::FieldAccess { + base, + field, + span: expr_span, + } if matches!(base.as_ref(), rumoca_core::Expression::FunctionCall { .. }) => { + let span = inherited_projection_span(*expr_span, span); + let base = self.substitute(base, scope)?; + let field_expr = rumoca_core::Expression::FieldAccess { + base: Box::new(base), + field: field.clone(), span, - )?; - merged.full.insert(name, value); + }; + if let Some(mut outputs) = self.function_call_projected_scalars_with_scope( + &field_expr, + scope, + depth, + span, + )? && outputs.len() == 1 + { + return Ok(outputs.remove(0)); + } + Ok(field_expr) } + _ => Ok(expr.clone()), } - Ok(merged) } - fn merged_if_scalar_values( + fn normalize_projected_scalar_builtin_args( &self, - name: &str, - entry_scope: &FunctionProjectionScope, - branch_conditions: &[rumoca_core::Expression], - branch_scopes: &[FunctionProjectionScope], - else_scope: &FunctionProjectionScope, + function: rumoca_core::BuiltinFunction, + args: &[rumoca_core::Expression], + scope: &FunctionProjectionScope, + depth: usize, span: rumoca_core::Span, - ) -> Result, LowerError> { - let base_values = else_scope - .scalars - .get(name) - .or_else(|| entry_scope.scalars.get(name)) - .ok_or_else(|| guarded_assignment_without_base(name, span))?; - let mut merged = projection_vec_with_capacity( - base_values.len(), - "if projection scalar value count", + ) -> Result { + let mut normalized_args = projection_vec_with_capacity( + args.len(), + "projected scalar builtin argument count", span, )?; - for value in base_values { - merged.push(value.clone()); - } - for (condition, branch_scope) in branch_conditions.iter().zip(branch_scopes.iter()).rev() { - let Some(branch_values) = branch_scope.scalars.get(name) else { - continue; - }; - if branch_values.len() != merged.len() { - return Err(LowerError::contract_violation( - format!( - "if-statement assignment to `{name}` has mismatched scalar widths: branch {}, else/current {}", - branch_values.len(), - merged.len() - ), - span, - )); - } - let mut next_merged = projection_vec_with_capacity( - branch_values.len(), - "if projection merged scalar value count", + for arg in args { + normalized_args.push(self.normalize_projected_scalar_output_expr( + arg, + scope, + depth + 1, span, - )?; - for (branch, fallback) in branch_values.iter().cloned().zip(merged) { - next_merged.push(rumoca_core::Expression::If { - branches: single_projection_branch(condition.clone(), branch, span)?, - else_branch: Box::new(fallback), - span, - }); - } - merged = next_merged; + )?); } - Ok(merged) + Ok(rumoca_core::Expression::BuiltinCall { + function, + args: normalized_args, + span, + }) } - fn projected_outputs_from_scope( + fn function_call_projected_scalars_with_scope( &self, - function: &rumoca_core::Function, + expr: &rumoca_core::Expression, scope: &FunctionProjectionScope, depth: usize, - owner_span: rumoca_core::Span, - ) -> Result>, LowerError> { - let function_span = inherited_projection_span(function.span, owner_span); - let mut outputs = projection_vec_with_capacity( - function.outputs.len(), - "projected function output count", - function_span, + span: rumoca_core::Span, + ) -> Result>, LowerError> { + let Some((call, field)) = entrypoints::function_field_access(expr) else { + return Ok(None); + }; + let Some(outputs) = + self.function_call_outputs_with_projection_scope(call, depth + 1, span, Some(scope))? + else { + return Ok(None); + }; + let mut selected = projection_vec_with_capacity( + outputs.len(), + "projected function field scalar count", + span, )?; - for output in &function.outputs { - let output_span = inherited_projection_span(output.span, function_span); - if let Some(values) = scope.scalars.get(output.name.as_str()) { - let projected = project_scalar_outputs(output, values, output_span)?; - append_projected_outputs( - &mut outputs, - projected, - "projected scalar output count", - output_span, - )?; - continue; + for output in outputs { + if let Some(expr) = self.project_output_field_value(output, field, scope, span)? { + selected.push(expr); } - if let Some(expr) = scope.full.get(output.name.as_str()) { - let projected = - self.project_output_expr(output, expr, scope, depth + 1, output_span)?; - append_projected_outputs( - &mut outputs, - projected, - "projected full output count", - output_span, - )?; + } + Ok((!selected.is_empty()).then_some(selected)) + } + + #[allow(clippy::excessive_nesting)] + fn resolve_projected_scalar_field_outputs( + &self, + outputs: Vec, + span: rumoca_core::Span, + ) -> Result, LowerError> { + let scope = FunctionProjectionScope::default(); + let mut resolved = projection_vec_with_capacity( + outputs.len(), + "resolved projected scalar field output count", + span, + )?; + for mut output in outputs { + if let rumoca_core::Expression::FieldAccess { + base, + field, + span: expr_span, + } = &output.expr + && matches!(base.as_ref(), rumoca_core::Expression::FunctionCall { .. }) + { + let field_expr = rumoca_core::Expression::FieldAccess { + base: base.clone(), + field: field.clone(), + span: inherited_projection_span(*expr_span, span), + }; + if let Some(mut values) = + self.function_call_projected_scalars_with_scope(&field_expr, &scope, 0, span)? + && values.len() == 1 + { + output.expr = values.remove(0); + } } + resolved.push(output); } - Ok((!outputs.is_empty()).then_some(outputs)) + Ok(resolved) } - fn project_output_expr( + fn projected_binary_op( &self, - output: &rumoca_core::FunctionParam, - expr: &rumoca_core::Expression, + op: &rumoca_core::OpBinary, + lhs: &rumoca_core::Expression, + rhs: &rumoca_core::Expression, scope: &FunctionProjectionScope, depth: usize, - output_span: rumoca_core::Span, - ) -> Result, LowerError> { - if output.dims.is_empty() { - let mut projected = projection_vec_with_capacity( - 1, - "scalar function output projection count", - output_span, - )?; - projected.push(ProjectedFunctionOutput { - field_path: Vec::new(), - selector_indices: Vec::new(), - expr: expr.clone(), - }); - return Ok(projected); + span: rumoca_core::Span, + ) -> Result { + if !is_div(op) { + return Ok(op.clone()); + } + let lhs_dims = self.expr_dims_with_owner(lhs, scope, depth, span)?; + let rhs_dims = self.expr_dims_with_owner(rhs, scope, depth, span)?; + if lhs_dims.as_ref().is_some_and(|dims| !dims.is_empty()) + || rhs_dims.as_ref().is_some_and(|dims| !dims.is_empty()) + { + Ok(rumoca_core::OpBinary::DivElem) + } else { + Ok(op.clone()) } - let expr_span = inherited_projection_source_span(expr.span(), output_span); - let values = self - .project_value_scalars(expr, &output.dims, scope, depth, expr_span) - .map_err(|err| err.with_fallback_span(expr_span))? - .ok_or_else(|| { - unsupported_at( - format!("function output `{}` could not be projected", output.name), - expr_span, - ) - })?; - project_scalar_outputs(output, &values, output_span) } fn is_record_constructor_call( @@ -1105,7 +2736,7 @@ impl<'a> FunctionProjectionAnalysis<'a> { }; let actual = actual.clone(); let actual = self.substitute(&actual, scope)?; - let Some(scalars) = self.optional_constructor_input_scalars( + let Some((projection_dims, scalars)) = self.optional_constructor_input_scalars( &actual, input, scope, @@ -1125,7 +2756,7 @@ impl<'a> FunctionProjectionAnalysis<'a> { outputs.push(ProjectedFunctionOutput { field_path: single_field_path(&input.name, input_span)?, selector_indices: required_flat_index_to_subscripts( - &input.dims, + &projection_dims, idx, input_span, )?, @@ -1143,9 +2774,9 @@ impl<'a> FunctionProjectionAnalysis<'a> { scope: &FunctionProjectionScope, depth: usize, actual_span: rumoca_core::Span, - ) -> Result>, LowerError> { + ) -> Result, LowerError> { let span = projection_arg_or_context_span(actual, actual_span)?; - let Some(dims) = constructor_input_projection_dims( + let Some(mut dims) = constructor_input_projection_dims( input, self.expr_dims_with_owner(actual, scope, depth + 1, span)?, span, @@ -1153,11 +2784,25 @@ impl<'a> FunctionProjectionAnalysis<'a> { else { return Ok(None); }; + if (input.dims.as_slice() == [0] + || scalar_count_for_dims(&dims, "record constructor input dimensions", span)? == 0) + && let rumoca_core::Expression::If { + branches, + else_branch, + .. + } = actual + && let Some(selected) = self.compile_time_if_selection(branches, else_branch, scope)? + && let Some(selected_dims) = + self.expr_dims_with_owner(selected, scope, depth + 1, span)? + && !selected_dims.is_empty() + { + dims = selected_dims; + } match self .project_value_scalars(actual, &dims, scope, depth, span) .map_err(|err| err.with_fallback_span(span))? { - Some(scalars) => Ok(Some(scalars)), + Some(scalars) => Ok(Some((dims, scalars))), None if actual.contains_der() => Err(unsupported_at( "record constructor derivative input could not be projected", span, @@ -1171,7 +2816,11 @@ impl<'a> FunctionProjectionAnalysis<'a> { expr: &rumoca_core::Expression, scope: &FunctionProjectionScope, ) -> Result { - let mut substituter = FunctionScopeSubstituter { scope, error: None }; + let mut substituter = FunctionScopeSubstituter { + scope, + error: None, + stack: Vec::new(), + }; let expr = substituter.rewrite_expression(expr); if let Some(error) = substituter.error { return Err(error); @@ -1190,7 +2839,7 @@ impl<'a> FunctionProjectionAnalysis<'a> { let span = inherited_projection_source_span(expr.span(), owner_span); let count = scalar_count_for_dims(dims, "projected value dimensions", span)?; if let Some(values) = - self.project_function_call_scalars_once(expr, count, scope, depth, span)? + self.project_function_call_scalars_once(expr, dims, count, scope, depth, span)? { return Ok(Some(values)); } @@ -1202,6 +2851,7 @@ impl<'a> FunctionProjectionAnalysis<'a> { fn project_function_call_scalars_once( &self, expr: &rumoca_core::Expression, + dims: &[i64], count: usize, scope: &FunctionProjectionScope, depth: usize, @@ -1212,15 +2862,27 @@ impl<'a> FunctionProjectionAnalysis<'a> { substituted = substituted.with_span(owner_span); } let rumoca_core::Expression::FunctionCall { + name, + args, is_constructor: false, .. - } = substituted + } = &substituted else { return Ok(None); }; - let Some(outputs) = - self.function_call_outputs_with_owner(&substituted, depth + 1, owner_span)? - else { + if is_stream_passthrough_intrinsic(name.as_str()) { + let Some(arg) = args.first() else { + return Ok(None); + }; + return self.project_value_scalars(arg, dims, scope, depth + 1, owner_span); + } + let outputs = self.function_call_outputs_with_projection_scope( + &substituted, + depth + 1, + owner_span, + Some(scope), + )?; + let Some(outputs) = outputs else { return Ok(None); }; if outputs.len() != count { @@ -1237,6 +2899,7 @@ impl<'a> FunctionProjectionAnalysis<'a> { Ok(Some(scalars)) } + #[allow(clippy::excessive_nesting, clippy::too_many_lines)] fn project_value( &self, expr: &rumoca_core::Expression, @@ -1247,6 +2910,39 @@ impl<'a> FunctionProjectionAnalysis<'a> { owner_span: rumoca_core::Span, ) -> Result, LowerError> { let owner_span = inherited_projection_source_span(expr.span(), owner_span); + if let rumoca_core::Expression::FieldAccess { base, field, span } = expr { + let span = inherited_projection_span(*span, owner_span); + if let rumoca_core::Expression::FunctionCall { .. } = base.as_ref() { + let mut call = self.substitute(base, scope)?; + if call.span().is_none() { + call = call.with_span(span); + } + if let Some(outputs) = self.function_call_outputs_with_projection_scope( + &call, + depth + 1, + span, + Some(scope), + )? { + for output in outputs { + if let Some(expr) = + self.project_output_field_value(output, field, scope, span)? + { + return Ok(Some(expr.with_span(span))); + } + } + } + } + if let rumoca_core::Expression::FunctionCall { args, .. } = base.as_ref() { + for arg in args { + if let Some((name, value)) = + super::super::function_calls::decode_named_function_arg(arg) + && name == field + { + return Ok(Some(self.substitute(value, scope)?.with_span(span))); + } + } + } + } if dims.is_empty() { return Ok(Some(self.substitute(expr, scope)?)); } @@ -1267,6 +2963,16 @@ impl<'a> FunctionProjectionAnalysis<'a> { )?; return Ok(Some(value.clone().with_span(span))); } + if let Some(value) = scope.full.get(name.as_str()) + && !is_same_plain_var_ref(value, name.as_str()) + { + if let Some(projected) = + self.project_value(value, dims, flat_index, scope, depth + 1, span)? + { + return Ok(Some(projected.with_span(span))); + } + return Ok(Some(self.substitute(value, scope)?.with_span(span))); + } let indices = required_flat_index_to_subscripts(dims, flat_index, span)?; let name = self.reference_with_dae_component_ref(name); Ok(Some(rumoca_core::Expression::VarRef { @@ -1313,19 +3019,111 @@ impl<'a> FunctionProjectionAnalysis<'a> { } if args.len() == 1 => { self.project_transpose_value(&args[0], dims, flat_index, scope, depth, owner_span) } + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Identity, + args, + span, + } => { + let span = inherited_projection_span(*span, owner_span); + self.project_identity_value(args, dims, flat_index, scope, span) + } + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Fill, + args, + span, + } => { + let span = inherited_projection_span(*span, owner_span); + self.project_fill_value(args, scope, span) + } + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Cat, + args, + span, + } => { + let span = inherited_projection_span(*span, owner_span); + self.project_cat_value(args, dims, flat_index, scope, depth, span) + } + rumoca_core::Expression::BuiltinCall { + function, + args, + span, + } if is_elementwise_builtin_projection(function) => { + let span = inherited_projection_span(*span, owner_span); + let ctx = projection_value_ctx(dims, flat_index, scope, depth, span); + self.project_builtin_call_value(function, args, &ctx) + } + rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor: false, + span, + } if is_stream_passthrough_intrinsic(name.as_str()) => { + let span = inherited_projection_span(*span, owner_span); + match args.first() { + Some(arg) => self.project_value(arg, dims, flat_index, scope, depth + 1, span), + None => Ok(None), + } + } rumoca_core::Expression::FunctionCall { is_constructor: false, .. - } => self.project_function_call_value(expr, flat_index, scope, depth, owner_span), + } => self.project_function_call_value(expr, dims, flat_index, scope, depth, owner_span), rumoca_core::Expression::Binary { op, lhs, rhs, span } - if is_mul(op) || is_add(op) || is_sub(op) || is_div(op) => + if is_elementwise_binary_projection_op(op) => { let span = inherited_projection_span(*span, owner_span); let ctx = projection_value_ctx(dims, flat_index, scope, depth, span); self.project_binary_value(op, lhs, rhs, &ctx) } - other => self.project_indexed_value(other, dims, flat_index, scope, owner_span), + other => self.project_indexed_value(other, dims, flat_index, scope, owner_span), + } + } + + fn project_cat_value( + &self, + args: &[rumoca_core::Expression], + dims: &[i64], + flat_index: usize, + scope: &FunctionProjectionScope, + depth: usize, + span: rumoca_core::Span, + ) -> Result, LowerError> { + let Some(dim_expr) = args.first() else { + return Ok(None); + }; + let Some(dim_value) = self.compile_time_scalar_in_scope(dim_expr, scope)? else { + return Ok(None); + }; + if (dim_value - 1.0).abs() > f64::EPSILON { + return Ok(None); + } + let mut offset = 0usize; + for operand in &args[1..] { + let Some(operand_dims) = self.expr_dims_with_owner(operand, scope, depth + 1, span)? + else { + return Ok(None); + }; + if operand_dims.is_empty() || operand_dims.len() != dims.len() { + return Ok(None); + } + if operand_dims.len() > 1 && operand_dims[1..] != dims[1..] { + return Ok(None); + } + let operand_count = + scalar_count_for_dims(&operand_dims, "cat operand scalar count", span)?; + if flat_index < offset + operand_count { + return self.project_value( + operand, + &operand_dims, + flat_index - offset, + scope, + depth + 1, + span, + ); + } + offset += operand_count; } + Ok(None) } fn project_array_expression_value( @@ -1566,7 +3364,7 @@ impl<'a> FunctionProjectionAnalysis<'a> { let mut projected_branches = projection_vec_with_capacity(branches.len(), "projected if branch count", ctx.span)?; for (condition, branch) in branches { - let condition = self.substitute(condition, ctx.scope)?; + let condition = self.project_lane_or_substitute(condition, ctx)?; let branch = self .project_value( branch, @@ -1596,6 +3394,68 @@ impl<'a> FunctionProjectionAnalysis<'a> { })) } + fn project_builtin_call_value( + &self, + function: &rumoca_core::BuiltinFunction, + args: &[rumoca_core::Expression], + ctx: &ProjectionValueCtx<'_>, + ) -> Result, LowerError> { + if matches!(function, rumoca_core::BuiltinFunction::Homotopy) { + let Some(actual) = args.first() else { + return Ok(None); + }; + return self.project_lane_or_substitute(actual, ctx).map(Some); + } + let mut projected_args = + projection_vec_with_capacity(args.len(), "projected builtin argument count", ctx.span)?; + for arg in args { + projected_args.push(self.project_lane_or_substitute(arg, ctx)?); + } + Ok(Some(rumoca_core::Expression::BuiltinCall { + function: *function, + args: projected_args, + span: ctx.span, + })) + } + + fn project_fill_value( + &self, + args: &[rumoca_core::Expression], + scope: &FunctionProjectionScope, + span: rumoca_core::Span, + ) -> Result, LowerError> { + let Some(value) = args.first() else { + return Ok(None); + }; + Ok(Some(self.substitute(value, scope)?.with_span(span))) + } + + fn project_lane_or_substitute( + &self, + expr: &rumoca_core::Expression, + ctx: &ProjectionValueCtx<'_>, + ) -> Result { + let dims = match self.expr_dims_with_owner(expr, ctx.scope, ctx.depth + 1, ctx.span)? { + Some(dims) if !dims.is_empty() => Some(dims), + _ => self.scope_projected_expr_dims(expr, ctx.scope, ctx.span)?, + }; + if let Some(dims) = dims.filter(|dims| !dims.is_empty()) { + let flat_index = if dims == ctx.dims { + ctx.flat_index + } else if scalar_count_for_dims(&dims, "projected child dimensions", ctx.span)? == 1 { + 0 + } else { + return self.substitute(expr, ctx.scope); + }; + return self + .project_value(expr, &dims, flat_index, ctx.scope, ctx.depth + 1, ctx.span)? + .ok_or_else(|| { + unsupported_at("expression argument could not be projected", ctx.span) + }); + } + self.substitute(expr, ctx.scope) + } + fn project_unary_value( &self, op: &rumoca_core::OpUnary, @@ -1610,7 +3470,19 @@ impl<'a> FunctionProjectionAnalysis<'a> { ctx.span, ) })?; - if rhs_dims != ctx.dims { + let rhs = if rhs_dims.is_empty() { + self.substitute(rhs, ctx.scope)? + } else if rhs_dims == ctx.dims { + self.project_value( + rhs, + ctx.dims, + ctx.flat_index, + ctx.scope, + ctx.depth + 1, + ctx.span, + )? + .ok_or_else(|| unsupported_at("unary array operand could not be projected", ctx.span))? + } else { return Err(unsupported_at( format!( "unary array expression has operand dimensions {}, expected {}", @@ -1619,19 +3491,7 @@ impl<'a> FunctionProjectionAnalysis<'a> { ), ctx.span, )); - } - let rhs = self - .project_value( - rhs, - ctx.dims, - ctx.flat_index, - ctx.scope, - ctx.depth + 1, - ctx.span, - )? - .ok_or_else(|| { - unsupported_at("unary array operand could not be projected", ctx.span) - })?; + }; Ok(Some(rumoca_core::Expression::Unary { op: op.clone(), rhs: Box::new(rhs), @@ -1683,18 +3543,268 @@ impl<'a> FunctionProjectionAnalysis<'a> { fn project_function_call_value( &self, expr: &rumoca_core::Expression, + dims: &[i64], flat_index: usize, scope: &FunctionProjectionScope, depth: usize, owner_span: rumoca_core::Span, ) -> Result, LowerError> { - let mut call = self.substitute(expr, scope)?; + if let Some(indexed_call) = + self.indexed_selected_output_call(expr, dims, flat_index, owner_span)? + { + return self.project_function_call_value( + &indexed_call, + &[], + 0, + scope, + depth + 1, + owner_span, + ); + } + let outputs = self.function_call_outputs_with_projection_scope( + expr, + depth + 1, + owner_span, + Some(scope), + )?; + if let Some(outputs) = outputs { + let span = inherited_projection_source_span(expr.span(), owner_span); + let ctx = projection_value_ctx(dims, flat_index, scope, depth, span); + if let [output] = outputs.as_slice() { + return self + .project_lane_or_substitute(&output.expr, &ctx) + .map(Some); + } + if function_call_declared_output_count(expr, self.dae_model) + .is_some_and(|count| count > 1) + && let Some(output) = outputs.first() + { + return self + .project_lane_or_substitute(&output.expr, &ctx) + .map(Some); + } + return outputs + .get(flat_index) + .map(|output| { + self.project_lane_or_substitute(&output.expr, &ctx) + .map(Some) + }) + .unwrap_or(Ok(None)); + } + let span = inherited_projection_source_span(expr.span(), owner_span); + let mut call = + self.project_function_call_with_lane_args(expr, dims, flat_index, scope, depth, span)?; if call.span().is_none() { call = call.with_span(owner_span); } - Ok(self - .function_call_outputs_with_owner(&call, depth + 1, owner_span)? - .and_then(|outputs| outputs.get(flat_index).map(|output| output.expr.clone()))) + let outputs = self.function_call_outputs_with_owner(&call, depth + 1, owner_span)?; + let Some(outputs) = outputs else { + if is_direct_declared_array_output_call(expr, self.dae_model) { + return Ok(None); + } + return Ok(Some(call)); + }; + if let [output] = outputs.as_slice() { + return Ok(Some(output.expr.clone())); + } + if function_call_declared_output_count(&call, self.dae_model).is_some_and(|count| count > 1) + { + return Ok(outputs.first().map(|output| output.expr.clone())); + } + Ok(outputs.get(flat_index).map(|output| output.expr.clone())) + } + + fn indexed_selected_output_call( + &self, + expr: &rumoca_core::Expression, + dims: &[i64], + flat_index: usize, + span: rumoca_core::Span, + ) -> Result, LowerError> { + let rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor: false, + .. + } = expr + else { + return Ok(None); + }; + if dims.is_empty() + || self + .dae_model + .symbols + .functions + .contains_key(name.var_name()) + { + return Ok(None); + } + rumoca_core::find_map_top_level_splits_rev(name.as_str(), |base_name, suffix| { + let function = self + .dae_model + .symbols + .functions + .get(&rumoca_core::VarName::new(base_name))?; + let projection_suffix = parse_output_projection_suffix(suffix)?; + if !projection_suffix.indices.is_empty() || projection_suffix.output_field.is_some() { + return None; + } + let output = function + .outputs + .iter() + .find(|output| output.name == projection_suffix.output_name)?; + if output.dims.is_empty() || output.dims.as_slice() != dims { + return None; + } + let selector = dae::scalar_name_text_for_flat_index( + output.name.as_str(), + &output.dims, + flat_index, + ); + Some((base_name.to_string(), selector)) + }) + .map(|(base_name, selector)| { + Ok(Some(rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new(format!("{base_name}.{selector}")).into(), + args: args.clone(), + is_constructor: false, + span, + })) + }) + .unwrap_or(Ok(None)) + } + + fn project_function_call_with_lane_args( + &self, + expr: &rumoca_core::Expression, + dims: &[i64], + flat_index: usize, + scope: &FunctionProjectionScope, + depth: usize, + span: rumoca_core::Span, + ) -> Result { + let rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor, + .. + } = expr + else { + return self.substitute(expr, scope); + }; + let ctx = projection_value_ctx(dims, flat_index, scope, depth, span); + let mut projected_args = + projection_vec_with_capacity(args.len(), "projected function argument count", span)?; + for arg in args { + projected_args.push(self.project_lane_or_substitute(arg, &ctx)?); + } + Ok(rumoca_core::Expression::FunctionCall { + name: name.clone(), + args: projected_args, + is_constructor: *is_constructor, + span, + }) + } + + fn project_output_field_value( + &self, + output: ProjectedFunctionOutput, + field: &str, + scope: &FunctionProjectionScope, + span: rumoca_core::Span, + ) -> Result, LowerError> { + let Some((head, tail)) = output.field_path.split_first() else { + return self.project_record_field_value(&output.expr, field, scope, span); + }; + if head != field { + return Ok(None); + } + if tail.is_empty() { + return Ok(Some(self.substitute(&output.expr, scope)?)); + } + Ok(None) + } + + #[allow(clippy::excessive_nesting)] + fn project_record_field_value( + &self, + value: &rumoca_core::Expression, + field: &str, + scope: &FunctionProjectionScope, + span: rumoca_core::Span, + ) -> Result, LowerError> { + match value { + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + if let Some(selected) = + self.compile_time_if_selection(branches, else_branch, scope)? + { + return self.project_record_field_value(selected, field, scope, span); + } + let mut projected_branches = projection_vec_with_capacity( + branches.len(), + "projected record field if branch count", + span, + )?; + for (condition, branch) in branches { + let Some(projected_branch) = + self.project_record_field_value(branch, field, scope, span)? + else { + return Ok(None); + }; + projected_branches.push((self.substitute(condition, scope)?, projected_branch)); + } + let Some(projected_else) = + self.project_record_field_value(else_branch, field, scope, span)? + else { + return Ok(None); + }; + Ok(Some(rumoca_core::Expression::If { + branches: projected_branches, + else_branch: Box::new(projected_else), + span, + })) + } + rumoca_core::Expression::FunctionCall { name, args, .. } => { + for arg in args { + if let Some((name, actual)) = + super::super::function_calls::decode_named_function_arg(arg) + && name == field + { + return Ok(Some(self.substitute(actual, scope)?)); + } + } + let Some(constructor) = self.dae_model.symbols.functions.get(name.var_name()) + else { + return Ok(None); + }; + if !is_record_constructor_signature(name.as_str(), constructor) + && !constructor.is_constructor + { + return Ok(None); + } + let (_named, positional) = + super::super::function_calls::split_named_and_positional_call_args( + name.as_str(), + args, + )?; + let Some(index) = constructor + .inputs + .iter() + .position(|input| input.name == field) + else { + return Ok(None); + }; + if let Some(actual) = positional.get(index) { + return Ok(Some(self.substitute(actual, scope)?)); + } + Ok(None) + } + _ => Ok(None), + } } fn project_indexed_value( @@ -1726,6 +3836,40 @@ impl<'a> FunctionProjectionAnalysis<'a> { })) } + fn project_identity_value( + &self, + args: &[rumoca_core::Expression], + dims: &[i64], + flat_index: usize, + scope: &FunctionProjectionScope, + span: rumoca_core::Span, + ) -> Result, LowerError> { + if dims.len() != 2 || dims[0] != dims[1] { + return Ok(None); + } + let Some(dim_expr) = args.first() else { + return Ok(None); + }; + let Some(dim_value) = self.compile_time_scalar_in_scope(dim_expr, scope)? else { + return Ok(None); + }; + let dim = i64::try_from(dim_value as i128).map_err(|_| { + LowerError::contract_violation("identity dimension is outside host range", span) + })?; + if dim <= 0 || dim != dims[0] { + return Ok(None); + } + let dim = usize::try_from(dim).map_err(|_| { + LowerError::contract_violation("identity dimension is outside host range", span) + })?; + let row = flat_index / dim; + let col = flat_index % dim; + Ok(Some(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(if row == col { 1.0 } else { 0.0 }), + span, + })) + } + fn project_binary_elementwise( &self, op: OpBinary, @@ -1736,30 +3880,18 @@ impl<'a> FunctionProjectionAnalysis<'a> { let lhs_dims = self.known_expr_dims(lhs, ctx.scope, ctx.depth, "binary lhs", ctx.span)?; let rhs_dims = self.known_expr_dims(rhs, ctx.scope, ctx.depth, "binary rhs", ctx.span)?; let lhs_expr = if lhs_dims.is_empty() { - self.substitute(lhs, ctx.scope)? + self.project_lane_or_substitute(lhs, ctx)? } else { - self.project_value( - lhs, - ctx.dims, - ctx.flat_index, - ctx.scope, - ctx.depth, - ctx.span, - )? - .ok_or_else(|| unsupported_at("binary lhs could not be projected", ctx.span))? + let flat_index = projected_child_flat_index(&lhs_dims, ctx.flat_index); + self.project_value(lhs, &lhs_dims, flat_index, ctx.scope, ctx.depth, ctx.span)? + .ok_or_else(|| unsupported_at("binary lhs could not be projected", ctx.span))? }; let rhs_expr = if rhs_dims.is_empty() { - self.substitute(rhs, ctx.scope)? + self.project_lane_or_substitute(rhs, ctx)? } else { - self.project_value( - rhs, - ctx.dims, - ctx.flat_index, - ctx.scope, - ctx.depth, - ctx.span, - )? - .ok_or_else(|| unsupported_at("binary rhs could not be projected", ctx.span))? + let flat_index = projected_child_flat_index(&rhs_dims, ctx.flat_index); + self.project_value(rhs, &rhs_dims, flat_index, ctx.scope, ctx.depth, ctx.span)? + .ok_or_else(|| unsupported_at("binary rhs could not be projected", ctx.span))? }; Ok(Some(rumoca_core::Expression::Binary { op, @@ -1782,6 +3914,15 @@ impl<'a> FunctionProjectionAnalysis<'a> { return Ok(None); }; match (lhs_dims.as_slice(), rhs_dims.as_slice(), ctx.dims) { + ([lhs_size], [rhs_size], []) if lhs_size == rhs_size => { + self.project_vector_dot_product(lhs, rhs, &lhs_dims, &rhs_dims, ctx, *lhs_size) + } + ([lhs_size], [rhs_size], []) => Err(unsupported_at( + format!( + "vector dot product dimensions [{lhs_size}] and [{rhs_size}] are incompatible" + ), + ctx.span, + )), ([rows, cols], [n], [_]) if cols == n => self.project_matrix_vector_product( lhs, rhs, @@ -1805,6 +3946,44 @@ impl<'a> FunctionProjectionAnalysis<'a> { } } + fn project_vector_dot_product( + &self, + lhs: &rumoca_core::Expression, + rhs: &rumoca_core::Expression, + lhs_dims: &[i64], + rhs_dims: &[i64], + ctx: &ProjectionValueCtx<'_>, + size: i64, + ) -> Result, LowerError> { + if size == 0 { + return Err(unsupported_at( + "vector dot product dimension must be positive", + ctx.span, + )); + } + let size = valid_product_dim(size, ctx.span, "vector dot product dimension")?; + if ctx.flat_index != 0 { + return Ok(None); + } + let mut terms = + projection_vec_with_capacity(size, "vector dot product term count", ctx.span)?; + for index in 0..size { + let lhs_term = self + .project_value(lhs, lhs_dims, index, ctx.scope, ctx.depth, ctx.span)? + .ok_or_else(|| unsupported_at("vector dot lhs could not be projected", ctx.span))?; + let rhs_term = self + .project_value(rhs, rhs_dims, index, ctx.scope, ctx.depth, ctx.span)? + .ok_or_else(|| unsupported_at("vector dot rhs could not be projected", ctx.span))?; + terms.push(rumoca_core::Expression::Binary { + op: OpBinary::Mul, + lhs: Box::new(lhs_term), + rhs: Box::new(rhs_term), + span: ctx.span, + }); + } + Ok(Some(sum_expressions(terms, ctx.span))) + } + fn project_matrix_vector_product( &self, lhs: &rumoca_core::Expression, @@ -1961,3 +4140,47 @@ impl<'a> FunctionProjectionAnalysis<'a> { Ok(Some(sum_expressions(terms, ctx.span))) } } + +fn split_flattened_projection_input_name(name: &str) -> Option<(&str, &str)> { + let (prefix, field) = name.split_once('_')?; + (!prefix.is_empty() && !field.is_empty()).then_some((prefix, field)) +} + +fn flattened_projection_input_has_prefix(name: &str, prefix: &str) -> bool { + split_flattened_projection_input_name(name).is_some_and(|(candidate, _)| candidate == prefix) +} + +fn flattened_projection_group_has_prefix( + inputs: &[rumoca_core::FunctionParam], + prefix: &str, +) -> bool { + inputs + .iter() + .filter(|input| flattened_projection_input_has_prefix(&input.name, prefix)) + .take(2) + .count() + >= 2 +} + +fn flattened_projection_input_is_group_start( + inputs: &[rumoca_core::FunctionParam], + input_idx: usize, + prefix: &str, +) -> bool { + !inputs + .iter() + .take(input_idx) + .any(|input| flattened_projection_input_has_prefix(&input.name, prefix)) +} + +fn only_projected_scalar_assignment_output( + mut outputs: Vec, + span: rumoca_core::Span, +) -> Result { + outputs.pop().ok_or_else(|| { + LowerError::contract_violation( + "projected scalar function assignment produced no output", + span, + ) + }) +} diff --git a/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/dimension_helpers.rs b/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/dimension_helpers.rs index a0da6d81b..7086237fb 100644 --- a/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/dimension_helpers.rs +++ b/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/dimension_helpers.rs @@ -166,13 +166,6 @@ pub(super) fn append_projected_outputs( Ok(()) } -pub(super) fn declared_dims( - function: &rumoca_core::Function, - name: &str, -) -> Result>, LowerError> { - Ok(declared_param_dims(function, name)?.filter(|dims| !dims.is_empty())) -} - pub(super) fn declared_param_dims( function: &rumoca_core::Function, name: &str, @@ -199,10 +192,29 @@ pub(super) fn formal_actual_projection_dims( context: String, span: rumoca_core::Span, ) -> Result>, LowerError> { + if formal.dims.iter().any(|dim| *dim < 0) { + return Err(LowerError::contract_violation( + format!( + "{context} has negative dimension in declared shape {:?}", + formal.dims + ), + span, + )); + } if let Some(actual_dims) = actual_dims { if actual_dims == formal.dims { return Ok(Some(actual_dims)); } + if actual_dims.is_empty() && formal.dims.len() > 1 && formal.dims.contains(&0) { + return Ok(Some(copy_projection_dims( + &formal.dims, + "zero-length formal parameter dimension count", + span, + )?)); + } + if formal.dims.as_slice() == [0] && actual_dims.is_empty() { + return Ok(Some(vec![1])); + } if let Some(resolved_dims) = resolve_formal_projection_dims(&formal.dims, &actual_dims, &context, span)? { @@ -211,6 +223,9 @@ pub(super) fn formal_actual_projection_dims( if formal.dims.is_empty() && formal_accepts_structured_actual(formal) { return Ok(Some(actual_dims)); } + if formal.dims.is_empty() { + return Ok(Some(actual_dims)); + } return Err(dimension_mismatch_error( &context, &formal.dims, @@ -280,16 +295,16 @@ pub(super) fn constructor_input_projection_dims( } pub(super) fn assignment_projection_dims( - function: &rumoca_core::Function, + function_name: &str, target: &str, + declared: Option>, value_dims: Option>, span: rumoca_core::Span, ) -> Result>, LowerError> { - let declared = declared_param_dims(function, target)?; match (value_dims, declared) { (Some(value_dims), Some(declared)) if value_dims != declared => { Err(dimension_mismatch_error( - &format!("function `{}` assignment to `{target}`", function.name), + &format!("function `{function_name}` assignment to `{target}`"), &declared, &value_dims, span, @@ -344,8 +359,8 @@ pub(super) fn dimension_mismatch_error( pub(super) fn is_ignorable_projection_statement(statement: &rumoca_core::Statement) -> bool { match statement { rumoca_core::Statement::Empty { .. } => true, - rumoca_core::Statement::FunctionCall { comp, .. } => { - comp.to_var_name().as_str() == "assert" + rumoca_core::Statement::FunctionCall { comp, outputs, .. } => { + outputs.is_empty() || comp.to_var_name().as_str() == "assert" } _ => false, } @@ -397,10 +412,55 @@ pub(super) fn projection_assignment_target( pub(super) struct FunctionScopeSubstituter<'a> { pub(super) scope: &'a FunctionProjectionScope, pub(super) error: Option, + pub(super) stack: Vec, } impl ExpressionRewriter for FunctionScopeSubstituter<'_> { + #[allow(clippy::too_many_lines)] fn rewrite_expression(&mut self, expr: &rumoca_core::Expression) -> rumoca_core::Expression { + if let rumoca_core::Expression::Index { + base, + subscripts, + span, + } = expr + && let rumoca_core::Expression::VarRef { name, .. } = base.as_ref() + && let Some(replacement) = self.substitute_subscripted_var_ref(name, subscripts, *span) + { + return replacement; + } + if let rumoca_core::Expression::FieldAccess { + base, field, span, .. + } = expr + && let rumoca_core::Expression::FieldAccess { + field: inner_field, .. + } = base.as_ref() + && inner_field == field + { + return self.rewrite_expression(base).with_span(*span); + } + if let rumoca_core::Expression::FieldAccess { + base, field, span, .. + } = expr + && let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = base.as_ref() + && subscripts.is_empty() + { + let flattened_key = format!("{}_{}", name.as_str(), field); + if let Some(replacement) = self.scope.full.get(flattened_key.as_str()) { + return self.rewrite_scope_replacement( + flattened_key.as_str(), + replacement, + expr, + *span, + ); + } + if name.as_str().ends_with(&format!("_{field}")) + && let Some(replacement) = self.scope.full.get(name.as_str()) + { + return self.rewrite_scope_replacement(name.as_str(), replacement, expr, *span); + } + } let rumoca_core::Expression::VarRef { name, subscripts, @@ -410,10 +470,47 @@ impl ExpressionRewriter for FunctionScopeSubstituter<'_> { return self.walk_expression(expr); }; if !subscripts.is_empty() { + if let Some(replacement) = self.substitute_subscripted_var_ref(name, subscripts, *span) + { + return replacement; + } return self.walk_expression(expr); } - if let Some(expr) = self.scope.full.get(name.as_str()) { - return expr.clone().with_span(*span); + if let Some(values) = self.scope.scalars.get(name.as_str()) { + if let [value] = values.as_slice() + && self + .scope + .dims + .get(name.as_str()) + .is_none_or(|dims| dims.is_empty()) + { + return value.clone().with_span(*span); + } + return expr.clone(); + } + if let Some(replacement) = self.scope.full.get(name.as_str()) { + if let rumoca_core::Expression::VarRef { + name: replacement_name, + subscripts: replacement_subscripts, + .. + } = replacement + && replacement_name == name + && replacement_subscripts.is_empty() + { + return replacement.clone().with_span(*span); + } + return self.rewrite_scope_replacement(name.as_str(), replacement, expr, *span); + } + let flattened_name = name.as_str().replace('.', "_"); + if flattened_name != name.as_str() + && let Some(replacement) = self.scope.full.get(flattened_name.as_str()) + { + return self.rewrite_scope_replacement( + flattened_name.as_str(), + replacement, + expr, + *span, + ); } if let Some(scalar) = rumoca_core::parse_scalar_name(name.as_str()) && let Some(values) = self.scope.scalars.get(scalar.base) @@ -446,6 +543,175 @@ impl ExpressionRewriter for FunctionScopeSubstituter<'_> { } } +impl FunctionScopeSubstituter<'_> { + fn rewrite_scope_replacement( + &mut self, + key: &str, + replacement: &rumoca_core::Expression, + fallback: &rumoca_core::Expression, + span: rumoca_core::Span, + ) -> rumoca_core::Expression { + if self.stack.iter().any(|active| active == key) { + return fallback.clone().with_span(span); + } + self.stack.push(key.to_string()); + let rewritten = self.rewrite_expression(replacement).with_span(span); + self.stack.pop(); + rewritten + } + + fn substitute_subscripted_var_ref( + &mut self, + name: &rumoca_core::Reference, + subscripts: &[rumoca_core::Subscript], + span: rumoca_core::Span, + ) -> Option { + if let Some(replacement) = self.projected_subscripted_var_ref(name, subscripts, span) { + return Some(replacement); + } + self.full_subscripted_var_ref_replacement(name, subscripts, span) + } + + fn projected_subscripted_var_ref( + &mut self, + name: &rumoca_core::Reference, + subscripts: &[rumoca_core::Subscript], + span: rumoca_core::Span, + ) -> Option { + let values = self.scope.scalars.get(name.as_str())?; + let dims = self.scope.dims.get(name.as_str())?; + let indices = match self.static_subscript_indices(subscripts, span) { + Ok(Some(indices)) => indices, + Ok(None) => return None, + Err(error) => { + self.error = Some(error); + return None; + } + }; + let flat_index = match flat_index_from_indices( + dims, + &indices, + span, + "projected subscripted variable substitution", + ) { + Ok(Some(flat_index)) => flat_index, + Ok(None) => return None, + Err(error) => { + self.error = Some(error); + return None; + } + }; + values + .get(flat_index) + .cloned() + .map(|value| value.with_span(span)) + } + + fn full_subscripted_var_ref_replacement( + &mut self, + name: &rumoca_core::Reference, + subscripts: &[rumoca_core::Subscript], + span: rumoca_core::Span, + ) -> Option { + let replacement = self.scope.full.get(name.as_str())?; + if self.stack.iter().any(|active| active == name.as_str()) { + return None; + } + if is_same_plain_var_ref(replacement, name.as_str()) { + return None; + } + let subscripts = match self.rewrite_subscripts(subscripts, span) { + Ok(subscripts) => subscripts, + Err(error) => { + self.error = Some(error); + return None; + } + }; + self.stack.push(name.as_str().to_string()); + let replacement = self.rewrite_expression(replacement); + self.stack.pop(); + match replacement { + rumoca_core::Expression::VarRef { + name, + subscripts: mut replacement_subscripts, + span, + } => { + replacement_subscripts.extend(subscripts); + Some(rumoca_core::Expression::VarRef { + name, + subscripts: replacement_subscripts, + span, + }) + } + replacement => Some(rumoca_core::Expression::Index { + base: Box::new(replacement), + subscripts, + span, + }), + } + } + + fn rewrite_subscripts( + &mut self, + subscripts: &[rumoca_core::Subscript], + span: rumoca_core::Span, + ) -> Result, LowerError> { + let mut rewritten = + projection_vec_with_capacity(subscripts.len(), "projected subscript count", span)?; + for subscript in subscripts { + rewritten.push(match subscript { + rumoca_core::Subscript::Expr { expr, span } => rumoca_core::Subscript::Expr { + expr: Box::new(self.rewrite_expression(expr)), + span: *span, + }, + rumoca_core::Subscript::Index { .. } | rumoca_core::Subscript::Colon { .. } => { + subscript.clone() + } + }); + } + Ok(rewritten) + } + + fn static_subscript_indices( + &mut self, + subscripts: &[rumoca_core::Subscript], + span: rumoca_core::Span, + ) -> Result>, LowerError> { + let mut indices = + projection_vec_with_capacity(subscripts.len(), "projected subscript count", span)?; + for subscript in subscripts { + let index = match subscript { + rumoca_core::Subscript::Index { value, .. } => *value, + rumoca_core::Subscript::Expr { expr, .. } => match self.static_expr_index(expr) { + Some(value) => value, + None => return Ok(None), + }, + rumoca_core::Subscript::Colon { .. } => return Ok(None), + }; + indices.push(index); + } + Ok(Some(indices)) + } + + fn static_expr_index(&mut self, expr: &rumoca_core::Expression) -> Option { + let rewritten = self.rewrite_expression(expr); + literal_expr_to_i64(&rewritten) + } +} + +fn literal_expr_to_i64(expr: &rumoca_core::Expression) -> Option { + let rumoca_core::Expression::Literal { value, .. } = expr else { + return None; + }; + match value { + Literal::Integer(value) => Some(*value), + Literal::Real(value) if value.is_finite() && value.fract().abs() < f64::EPSILON => { + Some(*value as i64) + } + _ => None, + } +} + pub(super) fn scalar_count_for_dims( dims: &[i64], context: &str, @@ -684,6 +950,7 @@ pub(super) fn binary_mul_dims( span: rumoca_core::Span, ) -> Result>, LowerError> { Ok(match (lhs_dims, rhs_dims) { + ([], []) => Some(Vec::new()), ([], dims) if !dims.is_empty() => Some(copy_projection_dims( dims, "scalar lhs product dimension count", diff --git a/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/dimension_inference.rs b/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/dimension_inference.rs index 65ffd1375..ac05178c2 100644 --- a/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/dimension_inference.rs +++ b/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/dimension_inference.rs @@ -1,7 +1,31 @@ use super::compile_time::{compile_time_binary, compile_time_var_key, literal_to_f64}; use super::*; +use crate::lower::is_stream_passthrough_intrinsic; use crate::projection_suffix::parse_output_projection_suffix; +fn merge_vectorized_scalar_actual_dims( + dims: &mut Option>, + actual_dims: &[i64], + span: rumoca_core::Span, +) -> Result { + if actual_dims == [1] && dims.as_ref().is_some_and(|dims| dims.as_slice() != [1]) { + return Ok(true); + } + if dims + .as_ref() + .is_some_and(|dims| dims.as_slice() == [1] && actual_dims != [1]) + { + *dims = Some(actual_dims.to_vec()); + return Ok(true); + } + let Some(merged) = elementwise_binary_dims(dims.as_deref().unwrap_or(&[]), actual_dims, span)? + else { + return Ok(false); + }; + *dims = Some(merged); + Ok(true) +} + impl<'a> FunctionProjectionAnalysis<'a> { // SPEC_0021: Exception - function projection dimension inference keeps // call, field, array, and comprehension cases in one recursive walker. @@ -27,40 +51,31 @@ impl<'a> FunctionProjectionAnalysis<'a> { depth: usize, owner_span: rumoca_core::Span, ) -> Result>, LowerError> { + if depth > super::super::super::MAX_FUNCTION_INLINE_DEPTH { + return Ok(None); + } let span = inherited_projection_source_span(expr.span(), owner_span); match expr { rumoca_core::Expression::VarRef { name, subscripts, .. - } if subscripts.is_empty() => { - if let Some(dims) = scope.dims.get(name.as_str()) { - return Ok(Some(copy_projection_dims( - dims, - "projected scope dimension count", - span, - )?)); - } - if let Some(values) = scope.scalars.get(name.as_str()) { - return Ok(Some(match values.len() { - 0 | 1 => Vec::new(), - len => copy_projection_dims( - &[checked_usize_to_i64( - len, - "projected scalar value count", - span, - )?], - "projected scalar dimension count", - span, - )?, - })); + } if subscripts.is_empty() => self.unsubscripted_var_ref_dims(name, scope, depth, span), + rumoca_core::Expression::VarRef { + name, subscripts, .. + } => self.subscripted_var_ref_dims(name, subscripts, scope, depth, span), + rumoca_core::Expression::Index { + base, subscripts, .. + } => { + let Some(base_dims) = self.expr_dims_with_owner(base, scope, depth + 1, span)? + else { + return Ok(None); + }; + if base_dims.is_empty() { + return Ok(None); } - if let Some(expr) = scope.full.get(name.as_str()) - && !is_same_plain_var_ref(expr, name.as_str()) - { - return self.expr_dims_with_owner(expr, scope, depth + 1, span); + if !self.index_expr_selectors_have_known_shape(subscripts, scope, depth, span)? { + return Ok(None); } - Ok(variable_by_name(self.dae_model, name.as_str()) - .map(|variable| variable_dims_i64(variable, span)) - .transpose()?) + self.subscripted_dims(&base_dims, subscripts, scope, span) } rumoca_core::Expression::Array { elements, @@ -79,6 +94,20 @@ impl<'a> FunctionProjectionAnalysis<'a> { Some(arg) => self.expr_dims_with_owner(arg, scope, depth, span), None => Ok(None), }, + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Sum, + .. + } => Ok(Some(Vec::new())), + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Cat, + args, + .. + } => self.cat_expr_dims(args, scope, depth, span), + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => self.if_expr_dims(branches, else_branch, scope, depth, span), rumoca_core::Expression::BuiltinCall { function: rumoca_core::BuiltinFunction::Transpose, args, @@ -100,13 +129,22 @@ impl<'a> FunctionProjectionAnalysis<'a> { )?)) } rumoca_core::Expression::BuiltinCall { function, args, .. } => { - self.linear_algebra_builtin_dims(function, args, scope, depth, span) + self.builtin_call_dims(function, args, scope, depth, span) } + rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor: false, + .. + } if is_stream_passthrough_intrinsic(name.as_str()) => match args.first() { + Some(arg) => self.expr_dims_with_owner(arg, scope, depth, span), + None => Ok(None), + }, rumoca_core::Expression::FunctionCall { name, is_constructor: false, .. - } => self.function_call_expr_dims(name, expr, span, depth), + } => self.function_call_expr_dims(name, expr, scope, span, depth), rumoca_core::Expression::FieldAccess { base, field, .. } => { if let Some(dims) = self.bound_field_access_dims(base, field, scope, span, depth)? { return Ok(Some(dims)); @@ -140,10 +178,124 @@ impl<'a> FunctionProjectionAnalysis<'a> { }; elementwise_binary_dims(&lhs_dims, &rhs_dims, span) } + rumoca_core::Expression::Binary { op, lhs, rhs, .. } + if is_elementwise_binary_projection_op(op) => + { + let Some(lhs_dims) = self.expr_dims_with_owner(lhs, scope, depth, span)? else { + return Ok(None); + }; + let Some(rhs_dims) = self.expr_dims_with_owner(rhs, scope, depth, span)? else { + return Ok(None); + }; + elementwise_binary_dims(&lhs_dims, &rhs_dims, span) + } _ => Ok(None), } } + fn index_expr_selectors_have_known_shape( + &self, + subscripts: &[rumoca_core::Subscript], + scope: &FunctionProjectionScope, + depth: usize, + span: rumoca_core::Span, + ) -> Result { + for subscript in subscripts { + let rumoca_core::Subscript::Expr { expr, .. } = subscript else { + continue; + }; + let known_shape = match expr.as_ref() { + rumoca_core::Expression::Range { + start, step, end, .. + } => self + .range_subscript_dim_count(start, step.as_deref(), end, scope, span)? + .is_some(), + _ => self + .expr_dims_with_owner(expr, scope, depth + 1, span)? + .is_some_and(|dims| dims.is_empty()), + }; + if !known_shape { + return Ok(false); + } + } + Ok(true) + } + + fn builtin_call_dims( + &self, + function: &rumoca_core::BuiltinFunction, + args: &[rumoca_core::Expression], + scope: &FunctionProjectionScope, + depth: usize, + span: rumoca_core::Span, + ) -> Result>, LowerError> { + if let Some(dims) = self.linear_algebra_builtin_dims(function, args, scope, depth, span)? { + return Ok(Some(dims)); + } + if matches!( + function, + rumoca_core::BuiltinFunction::Ones | rumoca_core::BuiltinFunction::Zeros + ) { + return self.array_constructor_dims(args, scope, span); + } + if matches!(function, rumoca_core::BuiltinFunction::Fill) { + return self.array_constructor_dims(args.get(1..).unwrap_or(&[]), scope, span); + } + if matches!(function, rumoca_core::BuiltinFunction::Homotopy) { + let Some(actual) = args.first() else { + return Ok(None); + }; + return self.expr_dims_with_owner(actual, scope, depth + 1, span); + } + if is_elementwise_builtin_projection(function) { + return self.elementwise_builtin_dims(args, scope, depth, span); + } + Ok(None) + } + + fn array_constructor_dims( + &self, + args: &[rumoca_core::Expression], + scope: &FunctionProjectionScope, + span: rumoca_core::Span, + ) -> Result>, LowerError> { + let mut dims = + projection_vec_with_capacity(args.len(), "array constructor dimension count", span)?; + for arg in args { + let Some(value) = self.compile_time_scalar_in_scope(arg, scope)? else { + return Ok(None); + }; + if value < 0.0 || value.fract() != 0.0 || value > i64::MAX as f64 { + return Err(LowerError::contract_violation( + format!("array constructor dimension `{value}` is not a valid integer"), + span, + )); + } + dims.push(value as i64); + } + Ok(Some(dims)) + } + + fn elementwise_builtin_dims( + &self, + args: &[rumoca_core::Expression], + scope: &FunctionProjectionScope, + depth: usize, + span: rumoca_core::Span, + ) -> Result>, LowerError> { + let mut dims = Vec::new(); + for arg in args { + let Some(arg_dims) = self.expr_dims_with_owner(arg, scope, depth + 1, span)? else { + return Ok(None); + }; + let Some(merged) = elementwise_binary_dims(&dims, &arg_dims, span)? else { + return Ok(None); + }; + dims = merged; + } + Ok(Some(dims)) + } + fn linear_algebra_builtin_dims( &self, function: &rumoca_core::BuiltinFunction, @@ -172,6 +324,162 @@ impl<'a> FunctionProjectionAnalysis<'a> { } } + fn unsubscripted_var_ref_dims( + &self, + name: &rumoca_core::Reference, + scope: &FunctionProjectionScope, + depth: usize, + span: rumoca_core::Span, + ) -> Result>, LowerError> { + if let Some(dims) = scope.dims.get(name.as_str()) { + return Ok(Some(copy_projection_dims( + dims, + "projected scope dimension count", + span, + )?)); + } + if let Some(values) = scope.scalars.get(name.as_str()) { + return Ok(Some(match values.len() { + 0 | 1 => Vec::new(), + len => copy_projection_dims( + &[checked_usize_to_i64( + len, + "projected scalar value count", + span, + )?], + "projected scalar dimension count", + span, + )?, + })); + } + if let Some(expr) = scope.full.get(name.as_str()) + && !is_same_plain_var_ref(expr, name.as_str()) + { + return self.expr_dims_with_owner(expr, scope, depth + 1, span); + } + variable_by_name(self.dae_model, name.as_str()) + .map(|variable| variable_dims_i64(variable, span)) + .transpose() + } + + fn subscripted_var_ref_dims( + &self, + name: &rumoca_core::Reference, + subscripts: &[rumoca_core::Subscript], + scope: &FunctionProjectionScope, + depth: usize, + span: rumoca_core::Span, + ) -> Result>, LowerError> { + let Some(base_dims) = self.unsubscripted_var_ref_dims(name, scope, depth + 1, span)? else { + return Ok(None); + }; + if base_dims.is_empty() { + return Ok(None); + } + self.subscripted_dims(&base_dims, subscripts, scope, span) + } + + fn subscripted_dims( + &self, + base_dims: &[i64], + subscripts: &[rumoca_core::Subscript], + scope: &FunctionProjectionScope, + span: rumoca_core::Span, + ) -> Result>, LowerError> { + if subscripts.len() > base_dims.len() { + return Err(LowerError::contract_violation( + "subscripted expression has more subscripts than inferred dimensions", + span, + )); + } + let mut dims = projection_vec_with_capacity( + base_dims.len(), + "projected subscript dimension count", + span, + )?; + for (dim, subscript) in base_dims.iter().copied().zip(subscripts) { + if let Some(count) = self.subscript_preserved_dim_count(dim, subscript, scope, span)? { + dims.push(count); + } + } + dims.extend_from_slice(&base_dims[subscripts.len()..]); + Ok(Some(dims)) + } + + fn subscript_preserved_dim_count( + &self, + dim: i64, + subscript: &rumoca_core::Subscript, + scope: &FunctionProjectionScope, + span: rumoca_core::Span, + ) -> Result, LowerError> { + match subscript { + rumoca_core::Subscript::Colon { .. } => Ok(Some(dim)), + rumoca_core::Subscript::Index { .. } => Ok(None), + rumoca_core::Subscript::Expr { expr, .. } => { + if let rumoca_core::Expression::Range { + start, step, end, .. + } = expr.as_ref() + { + return self.range_subscript_dim_count( + start, + step.as_deref(), + end, + scope, + span, + ); + } + Ok(None) + } + } + } + + fn range_subscript_dim_count( + &self, + start: &rumoca_core::Expression, + step: Option<&rumoca_core::Expression>, + end: &rumoca_core::Expression, + scope: &FunctionProjectionScope, + span: rumoca_core::Span, + ) -> Result, LowerError> { + let Some(start) = self.compile_time_scalar_in_scope(start, scope)? else { + return Ok(None); + }; + let step = match step { + Some(step) => { + let Some(step) = self.compile_time_scalar_in_scope(step, scope)? else { + return Ok(None); + }; + step + } + None => 1.0, + }; + let Some(end) = self.compile_time_scalar_in_scope(end, scope)? else { + return Ok(None); + }; + if step == 0.0 { + return Err(LowerError::contract_violation( + "array slice range step must be non-zero", + span, + )); + } + let span_len = end - start; + if span_len == 0.0 { + return Ok(Some(1)); + } + if span_len.signum() != step.signum() { + return Ok(Some(0)); + } + let count = (span_len / step).floor() + 1.0; + if !count.is_finite() || count < 0.0 || count > i64::MAX as f64 { + return Err(LowerError::contract_violation( + "array slice range length exceeds i64", + span, + )); + } + Ok(Some(count as i64)) + } + fn diagonal_builtin_dims( &self, args: &[rumoca_core::Expression], @@ -263,6 +571,9 @@ impl<'a> FunctionProjectionAnalysis<'a> { fallback_span: rumoca_core::Span, ) -> Result, LowerError> { let Some(dims) = self.expr_dims_with_owner(expr, scope, depth, fallback_span)? else { + if let Some(dims) = self.scope_projected_expr_dims(expr, scope, fallback_span)? { + return Ok(dims); + } let span = if fallback_span.is_dummy() { projection_arg_or_context_span(expr, fallback_span)? } else { @@ -276,6 +587,17 @@ impl<'a> FunctionProjectionAnalysis<'a> { Ok(dims) } + pub(super) fn scope_projected_expr_dims( + &self, + expr: &rumoca_core::Expression, + scope: &FunctionProjectionScope, + span: rumoca_core::Span, + ) -> Result>, LowerError> { + let mut dims = None; + collect_scope_projected_expr_dims(expr, scope, &mut dims, span)?; + Ok(dims) + } + fn function_field_access_dims( &self, base: &rumoca_core::Expression, @@ -342,19 +664,22 @@ impl<'a> FunctionProjectionAnalysis<'a> { &self, name: &rumoca_core::Reference, expr: &rumoca_core::Expression, + scope: &FunctionProjectionScope, owner_span: rumoca_core::Span, depth: usize, ) -> Result>, LowerError> { - let rumoca_core::Expression::FunctionCall { span, .. } = expr else { + let rumoca_core::Expression::FunctionCall { args, span, .. } = expr else { return Ok(None); }; let call_span = inherited_projection_span(*span, owner_span); - if let Some(dims) = self.declared_function_output_dims(name, call_span)? { + if let Some(dims) = + self.declared_function_output_dims(name, args, scope, call_span, depth)? + { return Ok(Some(dims)); } let Some(outputs) = self.function_call_outputs_with_owner(expr, depth + 1, call_span)? else { - return Ok(None); + return Ok(Some(Vec::new())); }; function_outputs_dims(outputs.len(), call_span).map(Some) } @@ -362,14 +687,113 @@ impl<'a> FunctionProjectionAnalysis<'a> { fn declared_function_output_dims( &self, name: &rumoca_core::Reference, + args: &[rumoca_core::Expression], + scope: &FunctionProjectionScope, span: rumoca_core::Span, + depth: usize, ) -> Result>, LowerError> { if let Some(function) = self.dae_model.symbols.functions.get(name.var_name()) { - return exact_declared_function_output_dims(function, span); + let Some(dims) = exact_declared_function_output_dims(function, span)? else { + return Ok(None); + }; + if !dims.is_empty() { + return Ok(Some(dims)); + } + return self + .vectorized_scalar_function_call_dims(function, args, scope, span, depth)? + .map_or_else(|| Ok(Some(dims)), |dims| Ok(Some(dims))); } self.projected_declared_function_output_dims(name.as_str(), span) } + fn vectorized_scalar_function_call_dims( + &self, + function: &rumoca_core::Function, + args: &[rumoca_core::Expression], + scope: &FunctionProjectionScope, + span: rumoca_core::Span, + depth: usize, + ) -> Result>, LowerError> { + let (named, positional) = + super::super::super::function_calls::split_named_and_positional_call_args( + function.name.as_str(), + args, + )?; + let mut positional_idx = 0usize; + let mut dims: Option> = None; + for input in &function.inputs { + let actual = named.get(input.name.as_str()).copied().or_else(|| { + super::super::super::function_calls::next_positional_function_input_arg( + input, + &positional, + &mut positional_idx, + ) + }); + let Some(actual) = actual else { + continue; + }; + if !input.dims.is_empty() { + continue; + } + let Some(actual_dims) = self.expr_dims_with_owner(actual, scope, depth + 1, span)? + else { + continue; + }; + if actual_dims.is_empty() { + continue; + } + if !merge_vectorized_scalar_actual_dims( + &mut dims, + &actual_dims, + actual.span().unwrap_or(span), + )? { + return Ok(None); + } + } + Ok(dims) + } + + #[allow(clippy::excessive_nesting)] + fn cat_expr_dims( + &self, + args: &[rumoca_core::Expression], + scope: &FunctionProjectionScope, + depth: usize, + span: rumoca_core::Span, + ) -> Result>, LowerError> { + let Some(dim_expr) = args.first() else { + return Ok(None); + }; + let Some(dim_value) = self.compile_time_scalar_in_scope(dim_expr, scope)? else { + return Ok(None); + }; + if (dim_value - 1.0).abs() > f64::EPSILON { + return Ok(None); + } + let mut output_dims: Option> = None; + for operand in &args[1..] { + let Some(operand_dims) = self.expr_dims_with_owner(operand, scope, depth + 1, span)? + else { + return Ok(None); + }; + if operand_dims.is_empty() { + return Ok(None); + } + match &mut output_dims { + None => output_dims = Some(operand_dims), + Some(dims) => { + if dims.len() != operand_dims.len() || dims[1..] != operand_dims[1..] { + return Ok(None); + } + dims[0] = dims[0].checked_add(operand_dims[0]).ok_or_else(|| { + LowerError::contract_violation("cat(1, ...) dimension overflows i64", span) + })?; + } + } + } + Ok(output_dims) + } + fn projected_declared_function_output_dims( &self, requested: &str, @@ -460,7 +884,7 @@ impl<'a> FunctionProjectionAnalysis<'a> { ) -> Result, LowerError> { for (condition, branch) in branches { let condition = self.substitute(condition, scope)?; - let Some(value) = self.compile_time_scalar(&condition) else { + let Some(value) = self.compile_time_scalar_in_scope(&condition, scope)? else { return Ok(None); }; if value != 0.0 { @@ -471,30 +895,408 @@ impl<'a> FunctionProjectionAnalysis<'a> { } pub(super) fn compile_time_scalar(&self, expr: &rumoca_core::Expression) -> Option { + self.compile_time_scalar_in_scope(expr, &FunctionProjectionScope::default()) + .ok() + .flatten() + } + + #[allow(clippy::too_many_lines, clippy::excessive_nesting)] + pub(super) fn compile_time_scalar_in_scope( + &self, + expr: &rumoca_core::Expression, + scope: &FunctionProjectionScope, + ) -> Result, LowerError> { match expr { - rumoca_core::Expression::Literal { value, .. } => literal_to_f64(value), + rumoca_core::Expression::Literal { value, .. } => Ok(literal_to_f64(value)), rumoca_core::Expression::VarRef { name, subscripts, .. } => { - let key = compile_time_var_key(name, subscripts)?; - self.structural_bindings.get(key.as_str()).copied() + if subscripts.is_empty() + && let Some(value) = scope.full.get(name.as_str()) + && !is_same_plain_var_ref(value, name.as_str()) + { + return self.compile_time_scalar_in_scope(value, scope); + } + let Some(key) = compile_time_var_key(name, subscripts) else { + return Ok(None); + }; + Ok(self.structural_bindings.get(key.as_str()).copied()) } rumoca_core::Expression::Unary { op, rhs, .. } => { - let value = self.compile_time_scalar(rhs)?; - match op { + let Some(value) = self.compile_time_scalar_in_scope(rhs, scope)? else { + return Ok(None); + }; + Ok(match op { rumoca_core::OpUnary::Plus | rumoca_core::OpUnary::DotPlus - | rumoca_core::OpUnary::Empty => Some(value), - rumoca_core::OpUnary::Minus | rumoca_core::OpUnary::DotMinus => Some(-value), - rumoca_core::OpUnary::Not => Some(f64::from(value == 0.0)), + | rumoca_core::OpUnary::Empty => value, + rumoca_core::OpUnary::Minus | rumoca_core::OpUnary::DotMinus => -value, + rumoca_core::OpUnary::Not => f64::from(value == 0.0), } + .into()) } rumoca_core::Expression::Binary { op, lhs, rhs, .. } => { - let lhs = self.compile_time_scalar(lhs)?; - let rhs = self.compile_time_scalar(rhs)?; - compile_time_binary(op, lhs, rhs) + let Some(lhs) = self.compile_time_scalar_in_scope(lhs, scope)? else { + return Ok(None); + }; + let Some(rhs) = self.compile_time_scalar_in_scope(rhs, scope)? else { + return Ok(None); + }; + Ok(compile_time_binary(op, lhs, rhs)) + } + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + args, + span, + } => self.compile_time_size(args, scope, *span), + rumoca_core::Expression::BuiltinCall { + function: + rumoca_core::BuiltinFunction::NoEvent | rumoca_core::BuiltinFunction::Homotopy, + args, + .. + } => { + let Some(arg) = args.first() else { + return Ok(None); + }; + self.compile_time_scalar_in_scope(arg, scope) + } + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Smooth, + args, + .. + } => { + let Some(arg) = args.get(1) else { + return Ok(None); + }; + self.compile_time_scalar_in_scope(arg, scope) + } + rumoca_core::Expression::BuiltinCall { + function, + args, + span: _, + } if is_compile_time_unary_numeric_builtin(function) => { + let [arg] = args.as_slice() else { + return Ok(None); + }; + let Some(value) = self.compile_time_scalar_in_scope(arg, scope)? else { + return Ok(None); + }; + if *function == rumoca_core::BuiltinFunction::Integer { + return Ok(Some(value.floor())); + } + Ok(rumoca_core::apply_scalar_unary_math(*function, value)) + } + rumoca_core::Expression::BuiltinCall { + function, + args, + span: _, + } if is_compile_time_binary_numeric_builtin(function) => { + let [lhs, rhs] = args.as_slice() else { + return Ok(None); + }; + let Some(lhs) = self.compile_time_scalar_in_scope(lhs, scope)? else { + return Ok(None); + }; + let Some(rhs) = self.compile_time_scalar_in_scope(rhs, scope)? else { + return Ok(None); + }; + Ok(rumoca_core::apply_scalar_binary_math(*function, lhs, rhs)) + } + rumoca_core::Expression::FunctionCall { + is_constructor: false, + name, + args, + span, + .. + } => { + if let Some((call, selected_index)) = + selected_function_output_call(expr, self.dae_model)? + { + if !self.function_call_args_are_compile_time_scalars(args, scope)? { + return Ok(None); + } + let Some(outputs) = self.function_call_outputs_with_projection_scope( + &call, + 0, + inherited_projection_source_span(call.span(), *span), + Some(scope), + )? + else { + return Ok(None); + }; + let Some(output) = outputs.get(selected_index) else { + return Ok(None); + }; + return self.compile_time_scalar_in_scope(&output.expr, scope); + } + if let Some(value) = self.compile_time_modelica_math_call(name, args, scope)? { + return Ok(Some(value)); + } + let Some(outputs) = self.function_call_outputs_with_projection_scope( + expr, + 0, + inherited_projection_source_span(Some(*span), *span), + Some(scope), + )? + else { + return Ok(None); + }; + let [output] = outputs.as_slice() else { + return Ok(None); + }; + self.compile_time_scalar_in_scope(&output.expr, scope) } - _ => None, + rumoca_core::Expression::Index { + base, + subscripts, + span, + } => { + let Some(dims) = self.expr_dims_with_owner(base, scope, 0, *span)? else { + return Ok(None); + }; + let Some(indices) = self.compile_time_subscript_indices(subscripts, scope)? else { + return Ok(None); + }; + let Some(flat_index) = + flat_index_from_indices(&dims, &indices, *span, "compile-time indexed scalar")? + else { + return Ok(None); + }; + let Some(value) = self.project_value(base, &dims, flat_index, scope, 0, *span)? + else { + return Ok(None); + }; + self.compile_time_scalar_in_scope(&value, scope) + } + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + let Some(selected) = + self.compile_time_if_selection(branches, else_branch, scope)? + else { + return Ok(None); + }; + self.compile_time_scalar_in_scope(selected, scope) + } + _ => Ok(None), + } + } + + fn compile_time_modelica_math_call( + &self, + name: &rumoca_core::Reference, + args: &[rumoca_core::Expression], + scope: &FunctionProjectionScope, + ) -> Result, LowerError> { + let [arg] = args else { + return Ok(None); + }; + let Some(value) = self.compile_time_scalar_in_scope(arg, scope)? else { + return Ok(None); + }; + let name = name.as_str(); + if name == "asinh" || name.ends_with(".asinh") { + return Ok(Some(value.asinh())); + } + Ok(None) + } + + fn function_call_args_are_compile_time_scalars( + &self, + args: &[rumoca_core::Expression], + scope: &FunctionProjectionScope, + ) -> Result { + for arg in args { + if self.compile_time_scalar_in_scope(arg, scope)?.is_none() { + return Ok(false); + } + } + Ok(true) + } + + #[allow(clippy::excessive_nesting)] + fn compile_time_subscript_indices( + &self, + subscripts: &[rumoca_core::Subscript], + scope: &FunctionProjectionScope, + ) -> Result>, LowerError> { + if subscripts.is_empty() { + return Ok(Some(Vec::new())); + } + let mut indices = projection_vec_with_capacity( + subscripts.len(), + "compile-time subscript index count", + subscripts[0].span(), + )?; + for subscript in subscripts { + let index = match subscript { + rumoca_core::Subscript::Index { value, .. } => *value, + rumoca_core::Subscript::Expr { expr, span } => { + let Some(value) = self.compile_time_scalar_in_scope(expr, scope)? else { + return Ok(None); + }; + checked_shape_dimension(value, *span)? + } + rumoca_core::Subscript::Colon { .. } => return Ok(None), + }; + indices.push(index); + } + Ok(Some(indices)) + } + + fn compile_time_size( + &self, + args: &[rumoca_core::Expression], + scope: &FunctionProjectionScope, + span: rumoca_core::Span, + ) -> Result, LowerError> { + let [array_expr, dim_expr] = args else { + return Ok(None); + }; + let Some(dim_value) = self.compile_time_scalar_in_scope(dim_expr, scope)? else { + return Ok(None); + }; + let dim_index = if dim_value.fract().abs() < f64::EPSILON && dim_value >= 1.0 { + dim_value as usize + } else { + return Ok(None); + }; + let Some(dims) = self.expr_dims_with_owner(array_expr, scope, 0, span)? else { + return Ok(None); + }; + Ok(dims.get(dim_index - 1).map(|dim| *dim as f64)) + } + + fn if_expr_dims( + &self, + branches: &[(rumoca_core::Expression, rumoca_core::Expression)], + else_branch: &rumoca_core::Expression, + scope: &FunctionProjectionScope, + depth: usize, + span: rumoca_core::Span, + ) -> Result>, LowerError> { + if let Some(selected) = self.compile_time_if_selection(branches, else_branch, scope)? { + return self.expr_dims_with_owner(selected, scope, depth + 1, span); + } + let Some(else_dims) = self.expr_dims_with_owner(else_branch, scope, depth + 1, span)? + else { + return Ok(None); + }; + for (_, branch) in branches { + let Some(branch_dims) = self.expr_dims_with_owner(branch, scope, depth + 1, span)? + else { + return Ok(None); + }; + if branch_dims != else_dims { + return Ok(None); + } + } + Ok(Some(else_dims)) + } +} + +fn is_compile_time_unary_numeric_builtin(function: &rumoca_core::BuiltinFunction) -> bool { + matches!( + function, + rumoca_core::BuiltinFunction::Abs + | rumoca_core::BuiltinFunction::Sign + | rumoca_core::BuiltinFunction::Sqrt + | rumoca_core::BuiltinFunction::Floor + | rumoca_core::BuiltinFunction::Ceil + | rumoca_core::BuiltinFunction::Sin + | rumoca_core::BuiltinFunction::Cos + | rumoca_core::BuiltinFunction::Tan + | rumoca_core::BuiltinFunction::Asin + | rumoca_core::BuiltinFunction::Acos + | rumoca_core::BuiltinFunction::Atan + | rumoca_core::BuiltinFunction::Sinh + | rumoca_core::BuiltinFunction::Cosh + | rumoca_core::BuiltinFunction::Tanh + | rumoca_core::BuiltinFunction::Exp + | rumoca_core::BuiltinFunction::Log + | rumoca_core::BuiltinFunction::Log10 + | rumoca_core::BuiltinFunction::Integer + ) +} + +fn is_compile_time_binary_numeric_builtin(function: &rumoca_core::BuiltinFunction) -> bool { + matches!( + function, + rumoca_core::BuiltinFunction::Atan2 + | rumoca_core::BuiltinFunction::Min + | rumoca_core::BuiltinFunction::Max + | rumoca_core::BuiltinFunction::Div + | rumoca_core::BuiltinFunction::Mod + | rumoca_core::BuiltinFunction::Rem + ) +} + +fn collect_scope_projected_expr_dims( + expr: &rumoca_core::Expression, + scope: &FunctionProjectionScope, + dims: &mut Option>, + span: rumoca_core::Span, +) -> Result<(), LowerError> { + match expr { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } if subscripts.is_empty() => { + if let Some(candidate) = scope.dims.get(name.as_str()) + && !candidate.is_empty() + { + merge_vectorized_scalar_dims(dims, candidate, name.as_str(), span)?; + } else if let Some(replacement) = scope.full.get(name.as_str()) + && !is_same_plain_var_ref(replacement, name.as_str()) + { + collect_scope_projected_expr_dims(replacement, scope, dims, span)?; + } + } + rumoca_core::Expression::Unary { rhs, .. } => { + collect_scope_projected_expr_dims(rhs, scope, dims, span)?; + } + rumoca_core::Expression::Binary { lhs, rhs, .. } => { + collect_scope_projected_expr_dims(lhs, scope, dims, span)?; + collect_scope_projected_expr_dims(rhs, scope, dims, span)?; + } + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + for (condition, branch) in branches { + collect_scope_projected_expr_dims(condition, scope, dims, span)?; + collect_scope_projected_expr_dims(branch, scope, dims, span)?; + } + collect_scope_projected_expr_dims(else_branch, scope, dims, span)?; + } + rumoca_core::Expression::Array { elements, .. } + | rumoca_core::Expression::Tuple { elements, .. } => { + for element in elements { + collect_scope_projected_expr_dims(element, scope, dims, span)?; + } + } + rumoca_core::Expression::Range { + start, step, end, .. + } => { + collect_scope_projected_expr_dims(start, scope, dims, span)?; + if let Some(step) = step { + collect_scope_projected_expr_dims(step, scope, dims, span)?; + } + collect_scope_projected_expr_dims(end, scope, dims, span)?; + } + rumoca_core::Expression::BuiltinCall { args, .. } + | rumoca_core::Expression::FunctionCall { args, .. } => { + for arg in args { + collect_scope_projected_expr_dims(arg, scope, dims, span)?; + } + } + rumoca_core::Expression::FieldAccess { base, .. } + | rumoca_core::Expression::Index { base, .. } => { + collect_scope_projected_expr_dims(base, scope, dims, span)?; } + _ => {} } + Ok(()) } diff --git a/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/entrypoints.rs b/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/entrypoints.rs index 7edb65a6d..388b51795 100644 --- a/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/entrypoints.rs +++ b/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/entrypoints.rs @@ -1,6 +1,6 @@ use super::*; -pub(in crate::lower::derivative_rhs) fn function_projected_residuals_with_owner( +pub(in crate::lower) fn function_projected_residuals_with_owner( residual: &rumoca_core::Expression, dae_model: &dae::Dae, structural_bindings: &IndexMap, @@ -11,10 +11,11 @@ pub(in crate::lower::derivative_rhs) fn function_projected_residuals_with_owner( }; let analysis = FunctionProjectionAnalysis::new(dae_model, structural_bindings); if let Some((call, field)) = function_field_access(rhs) - && let Some(call_outputs) = analysis.top_level_function_call_outputs( - call, - inherited_projection_source_span(call.span(), owner_span), - )? + && let Some(call_outputs) = + dynamic_while_projection_decline(analysis.top_level_function_call_outputs( + call, + inherited_projection_source_span(call.span(), owner_span), + ))? { let Some(target_base) = plain_var_ref_name(lhs) else { return Ok(None); @@ -28,10 +29,11 @@ pub(in crate::lower::derivative_rhs) fn function_projected_residuals_with_owner( )?)); } if let Some((call, field)) = function_field_access(lhs) - && let Some(call_outputs) = analysis.top_level_function_call_outputs( - call, - inherited_projection_source_span(call.span(), owner_span), - )? + && let Some(call_outputs) = + dynamic_while_projection_decline(analysis.top_level_function_call_outputs( + call, + inherited_projection_source_span(call.span(), owner_span), + ))? { let Some(target_base) = plain_var_ref_name(rhs) else { return Ok(None); @@ -44,10 +46,12 @@ pub(in crate::lower::derivative_rhs) fn function_projected_residuals_with_owner( owner_span, )?)); } - if let Some(call_outputs) = analysis.top_level_function_call_outputs( - rhs, - inherited_projection_source_span(rhs.span(), owner_span), - )? { + if let Some(call_outputs) = + dynamic_while_projection_decline(analysis.top_level_function_call_outputs( + rhs, + inherited_projection_source_span(rhs.span(), owner_span), + ))? + { let Some(target_base) = plain_var_ref_name(lhs) else { return Ok(None); }; @@ -58,10 +62,12 @@ pub(in crate::lower::derivative_rhs) fn function_projected_residuals_with_owner( owner_span, )?)); } - if let Some(call_outputs) = analysis.top_level_function_call_outputs( - lhs, - inherited_projection_source_span(lhs.span(), owner_span), - )? { + if let Some(call_outputs) = + dynamic_while_projection_decline(analysis.top_level_function_call_outputs( + lhs, + inherited_projection_source_span(lhs.span(), owner_span), + ))? + { let Some(target_base) = plain_var_ref_name(rhs) else { return Ok(None); }; @@ -85,48 +91,308 @@ pub(in crate::lower::derivative_rhs) fn function_field_access( .then_some((base.as_ref(), field.as_str())) } -pub(in crate::lower::derivative_rhs) fn function_call_projected_scalars_with_owner( +#[allow(clippy::excessive_nesting)] +pub(in crate::lower) fn function_call_projected_scalars_with_owner( expr: &rumoca_core::Expression, dae_model: &dae::Dae, structural_bindings: &IndexMap, owner_span: rumoca_core::Span, ) -> Result>, LowerError> { let analysis = FunctionProjectionAnalysis::new(dae_model, structural_bindings); - if let Some(outputs) = analysis.top_level_function_call_outputs( - expr, - inherited_projection_source_span(expr.span(), owner_span), - )? { - return projected_output_expressions(outputs, owner_span).map(Some); + if let Some((call, scalar_index)) = selected_function_output_call(expr, dae_model)? { + let Some(outputs) = + dynamic_while_projection_decline(analysis.top_level_function_call_outputs( + &call, + inherited_projection_source_span(call.span(), owner_span), + ))? + else { + return Ok(None); + }; + let output_expr = outputs + .get(scalar_index) + .ok_or_else(|| { + LowerError::contract_violation( + format!( + "selected function output index {} is out of bounds for {} projected outputs", + scalar_index + 1, + outputs.len() + ), + owner_span, + ) + })? + .expr + .clone(); + let span = owner_span; + let mut values = + projection_vec_with_capacity(1, "selected function output scalar count", span)?; + values.push(output_expr); + return Ok(Some(values)); + } + if let Some((call, field)) = function_field_access(expr) + && let Some(outputs) = + dynamic_while_projection_decline(analysis.top_level_function_call_outputs( + call, + inherited_projection_source_span(call.span(), owner_span), + ))? + { + let scope = FunctionProjectionScope::default(); + let mut selected = projection_vec_with_capacity( + outputs.len(), + "projected function field scalar count", + owner_span, + )?; + for output in outputs { + if let Some(expr) = + analysis.project_output_field_value(output, field, &scope, owner_span)? + { + let dims = analysis + .expr_dims_with_owner(&expr, &scope, 0, owner_span)? + .unwrap_or_default(); + if dims.is_empty() { + selected.push(expr); + } else if let Some(mut scalars) = + analysis.project_value_scalars(&expr, &dims, &scope, 0, owner_span)? + { + selected.append(&mut scalars); + } else { + selected.push(expr); + } + } + } + if !selected.is_empty() { + return Ok(Some(selected)); + } + } + if let Some(outputs) = + dynamic_while_projection_decline(analysis.top_level_function_call_outputs( + expr, + inherited_projection_source_span(expr.span(), owner_span), + ))? + { + return projected_output_expressions( + analysis.resolve_projected_scalar_field_outputs(outputs, owner_span)?, + owner_span, + ) + .map(Some); + } + if let Some(values) = + projected_qualified_function_output_scalars(expr, dae_model, &analysis, owner_span)? + { + return Ok(Some(values)); + } + Ok(None) +} + +fn dynamic_while_projection_decline( + result: Result, LowerError>, +) -> Result, LowerError> { + match result { + Err(err) if err.is_dynamic_while_projection() => Ok(None), + result => result, + } +} + +pub(in crate::lower) fn function_call_projected_output_groups_with_owner( + expr: &rumoca_core::Expression, + dae_model: &dae::Dae, + structural_bindings: &IndexMap, + owner_span: rumoca_core::Span, +) -> Result>>, LowerError> { + let rumoca_core::Expression::FunctionCall { + name, + is_constructor: false, + .. + } = expr + else { + return Ok(None); + }; + let Some(function) = dae_model.symbols.functions.get(name.var_name()) else { + return Ok(None); + }; + let analysis = FunctionProjectionAnalysis::new(dae_model, structural_bindings); + let Some(outputs) = + dynamic_while_projection_decline(analysis.top_level_function_call_outputs( + expr, + inherited_projection_source_span(expr.span(), owner_span), + ))? + else { + return Ok(None); + }; + let mut groups = projection_vec_with_capacity( + function.outputs.len(), + "projected function output group count", + owner_span, + )?; + for output_param in &function.outputs { + let mut group = projection_vec_with_capacity( + outputs.len(), + "projected function output scalar group count", + owner_span, + )?; + if function.outputs.len() == 1 { + for output in &outputs { + group.push(output.expr.clone()); + } + } else { + group.extend( + outputs + .iter() + .filter(|output| { + output + .field_path + .first() + .is_some_and(|field| field == &output_param.name) + }) + .map(|output| output.expr.clone()), + ); + } + groups.push(group); } - let Some((call, scalar_index)) = selected_function_output_call(expr, dae_model)? else { + Ok(Some(groups)) +} + +fn projected_qualified_function_output_scalars( + expr: &rumoca_core::Expression, + dae_model: &dae::Dae, + analysis: &FunctionProjectionAnalysis<'_>, + owner_span: rumoca_core::Span, +) -> Result>, LowerError> { + let rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor: false, + span, + } = expr + else { return Ok(None); }; - let Some(outputs) = analysis.top_level_function_call_outputs( - &call, - inherited_projection_source_span(call.span(), owner_span), - )? + if dae_model.symbols.functions.contains_key(name.var_name()) { + return Ok(None); + } + let selected_name = name.as_str(); + for (function_name, function) in &dae_model.symbols.functions { + let prefix = format!("{}.", function_name.as_str()); + let Some(selector) = selected_name.strip_prefix(&prefix) else { + continue; + }; + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_var_name(function_name.clone()), + args: args.clone(), + is_constructor: false, + span: *span, + }; + let Some(outputs) = + dynamic_while_projection_decline(analysis.top_level_function_call_outputs( + &call, + inherited_projection_source_span(call.span(), owner_span), + ))? + else { + return Ok(None); + }; + let mut selected = projection_vec_with_capacity( + outputs.len(), + "qualified projected function output scalar count", + owner_span, + )?; + for output in outputs { + let output_selector = projected_output_selector(&output); + if output_selector == selector + || function.outputs.iter().any(|function_output| { + selector == format!("{}.{}", function_output.name, output_selector) + }) + { + selected.push(output.expr); + } + } + return Ok((!selected.is_empty()).then_some(selected)); + } + Ok(None) +} + +fn projected_output_selector(output: &ProjectedFunctionOutput) -> String { + let mut selector = output.field_path.join("."); + for index in &output.selector_indices { + selector.push('['); + selector.push_str(&index.to_string()); + selector.push(']'); + } + selector +} + +fn is_known_pure_nonconstructor_function_call( + expr: &rumoca_core::Expression, + dae_model: &dae::Dae, +) -> bool { + let rumoca_core::Expression::FunctionCall { + name, + is_constructor: false, + .. + } = expr else { + return false; + }; + dae_model + .symbols + .functions + .get(name.var_name()) + .is_some_and(|function| { + function.pure + && function.external.is_none() + && !function.is_constructor + && !is_record_constructor_signature(name.as_str(), function) + }) +} + +pub(in crate::lower) fn project_array_like_scalars_with_owner( + expr: &rumoca_core::Expression, + dae_model: &dae::Dae, + structural_bindings: &IndexMap, + owner_span: rumoca_core::Span, +) -> Result>, LowerError> { + if is_known_pure_nonconstructor_function_call(expr, dae_model) { + return function_call_projected_scalars_with_owner( + expr, + dae_model, + structural_bindings, + owner_span, + ); + } + let analysis = FunctionProjectionAnalysis::new(dae_model, structural_bindings); + let scope = FunctionProjectionScope::default(); + let Some(dims) = analysis.expr_dims_with_owner(expr, &scope, 0, owner_span)? else { + return Ok(None); + }; + if dims.is_empty() { + return Ok(None); + } + analysis.project_value_scalars(expr, &dims, &scope, 0, owner_span) +} + +pub(in crate::lower) fn project_array_like_scalar_with_owner( + expr: &rumoca_core::Expression, + flat_index: usize, + dae_model: &dae::Dae, + structural_bindings: &IndexMap, + owner_span: rumoca_core::Span, +) -> Result, LowerError> { + if is_known_pure_nonconstructor_function_call(expr, dae_model) { + return Ok(function_call_projected_scalars_with_owner( + expr, + dae_model, + structural_bindings, + owner_span, + )? + .and_then(|values| values.get(flat_index).cloned())); + } + let analysis = FunctionProjectionAnalysis::new(dae_model, structural_bindings); + let scope = FunctionProjectionScope::default(); + let Some(dims) = analysis.expr_dims_with_owner(expr, &scope, 0, owner_span)? else { return Ok(None); }; - let output_expr = outputs - .get(scalar_index) - .ok_or_else(|| { - LowerError::contract_violation( - format!( - "selected function output index {} is out of bounds for {} projected outputs", - scalar_index + 1, - outputs.len() - ), - owner_span, - ) - })? - .expr - .clone(); - let span = owner_span; - let mut values = - projection_vec_with_capacity(1, "selected function output scalar count", span)?; - values.push(output_expr); - Ok(Some(values)) + if dims.is_empty() { + return Ok(None); + } + analysis.project_value(expr, &dims, flat_index, &scope, 0, owner_span) } pub(in crate::lower::derivative_rhs) fn projected_output_expressions( @@ -155,22 +421,6 @@ pub(in crate::lower::derivative_rhs) fn plain_var_ref_name( } } -pub(in crate::lower::derivative_rhs) fn project_scalar_outputs( - output: &rumoca_core::FunctionParam, - values: &[rumoca_core::Expression], - output_span: rumoca_core::Span, -) -> Result, LowerError> { - let mut projected_values = projection_vec_with_capacity( - values.len(), - "function output scalar value count", - output_span, - )?; - for value in values { - projected_values.push(value.clone()); - } - project_target_scalar_outputs(&output.dims, projected_values, output_span) -} - pub(in crate::lower::derivative_rhs) fn project_target_scalar_outputs( dims: &[i64], values: Vec, diff --git a/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/loop_projection.rs b/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/loop_projection.rs index de41e7f88..f16f0cae9 100644 --- a/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/loop_projection.rs +++ b/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/loop_projection.rs @@ -88,7 +88,7 @@ impl FunctionProjectionAnalysis<'_> { ) -> Result { match subscript { rumoca_core::Subscript::Expr { expr, span } => { - let value = self.compile_time_int(&self.substitute(expr, scope)?, *span)?; + let value = self.compile_time_int(&self.substitute(expr, scope)?, scope, *span)?; Ok(rumoca_core::Subscript::Index { value, span: *span }) } rumoca_core::Subscript::Index { .. } | rumoca_core::Subscript::Colon { .. } => { @@ -108,10 +108,10 @@ impl FunctionProjectionAnalysis<'_> { rumoca_core::Expression::Range { start, step, end, .. } => { - let start = self.compile_time_int(&start, span)?; - let end = self.compile_time_int(&end, span)?; + let start = self.compile_time_int(&start, scope, span)?; + let end = self.compile_time_int(&end, scope, span)?; let step = match step { - Some(step) => self.compile_time_int(&step, span)?, + Some(step) => self.compile_time_int(&step, scope, span)?, None => 1, }; if step == 0 { @@ -129,7 +129,7 @@ impl FunctionProjectionAnalysis<'_> { span, )?, |mut values, element| { - values.push(self.compile_time_int(element, span)?); + values.push(self.compile_time_int(element, scope, span)?); Ok(values) }, ), @@ -139,7 +139,7 @@ impl FunctionProjectionAnalysis<'_> { "function projection scalar range value count", span, )?; - values.push(self.compile_time_int(&range, span)?); + values.push(self.compile_time_int(&range, scope, span)?); Ok(values) } } @@ -148,11 +148,14 @@ impl FunctionProjectionAnalysis<'_> { fn compile_time_int( &self, expr: &rumoca_core::Expression, + scope: &FunctionProjectionScope, span: rumoca_core::Span, ) -> Result { - let value = self.compile_time_scalar(expr).ok_or_else(|| { - unsupported_at("function projection requires a compile-time integer", span) - })?; + let value = self + .compile_time_scalar_in_scope(expr, scope)? + .ok_or_else(|| { + unsupported_at("function projection requires a compile-time integer", span) + })?; checked_compile_time_i64(value, span) } } diff --git a/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/target_projection.rs b/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/target_projection.rs index d5b4fe2ae..8d94e3f50 100644 --- a/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/target_projection.rs +++ b/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/target_projection.rs @@ -14,16 +14,10 @@ pub(super) fn project_reference_field_path_and_indices( if field_path.is_empty() && indices.is_empty() { return Ok(name.clone()); } - let component_ref = name.component_ref().ok_or_else(|| { - LowerError::contract_violation( - format!( - "array projection for `{}` lost structured component-reference metadata", - name.as_str() - ), - span, - ) - })?; - let mut component_ref = component_ref.clone(); + let mut component_ref = name + .component_ref() + .cloned() + .unwrap_or_else(|| fallback_component_reference(name.as_str(), span)); for field in field_path { component_ref.parts.push(rumoca_core::ComponentRefPart { ident: field.clone(), @@ -51,6 +45,22 @@ pub(super) fn project_reference_field_path_and_indices( )) } +fn fallback_component_reference( + name: &str, + span: rumoca_core::Span, +) -> rumoca_core::ComponentReference { + rumoca_core::ComponentReference { + local: false, + span, + parts: vec![rumoca_core::ComponentRefPart { + ident: name.to_string(), + span, + subs: Vec::new(), + }], + def_id: None, + } +} + fn projected_reference_subscript( index: usize, span: rumoca_core::Span, diff --git a/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/tests.rs b/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/tests.rs index 9dab96ee2..95ac03f8a 100644 --- a/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/tests.rs +++ b/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/tests.rs @@ -1,5 +1,11 @@ +// SPEC_0021 file-size exception: function derivative projection regressions +// share source fixtures and tensor assertions. split plan: split constructor, +// call-projection, and tensor-row cases into focused test modules. use super::*; +#[path = "tests/vector_dot_projection.rs"] +mod vector_dot_projection; + fn test_span() -> rumoca_core::Span { rumoca_core::Span::from_offsets( rumoca_core::SourceId::from_source_name("function_projection_test.mo"), @@ -15,6 +21,21 @@ fn real(value: f64) -> rumoca_core::Expression { } } +fn integer(value: i64) -> rumoca_core::Expression { + rumoca_core::Expression::Literal { + value: Literal::Integer(value), + span: test_span(), + } +} + +fn var_ref(name: &str) -> rumoca_core::Expression { + rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new(name).into(), + subscripts: Vec::new(), + span: test_span(), + } +} + fn array(elements: Vec, is_matrix: bool) -> rumoca_core::Expression { rumoca_core::Expression::Array { elements, @@ -57,6 +78,13 @@ fn component_reference(parts: Vec) -> rumoca_core }) } +fn assert_var_ref_name(expr: &rumoca_core::Expression, expected: &str) { + let rumoca_core::Expression::VarRef { name, .. } = expr else { + panic!("expected VarRef `{expected}`, got {expr:?}"); + }; + assert_eq!(name.as_str(), expected); +} + #[test] fn flatten_array_elements_flattens_matrix_rows() -> Result<(), LowerError> { let row1 = array(vec![real(1.0), real(2.0)], false); @@ -79,1195 +107,4664 @@ fn flatten_array_elements_flattens_matrix_rows() -> Result<(), LowerError> { } #[test] -fn scoped_single_scalar_value_has_scalar_dimensions() { +fn scalar_flat_index_projects_to_empty_subscripts() -> Result<(), LowerError> { + let subscripts = required_flat_index_to_subscripts(&[], 0, test_span())?; + + assert!(subscripts.is_empty()); + Ok(()) +} + +#[test] +fn scalar_flat_index_rejects_nonzero_index() { + let err = required_flat_index_to_subscripts(&[], 1, test_span()) + .expect_err("scalar flat index one should be out of bounds"); + + assert!( + err.reason() + .contains("flat index 1 is out of bounds for dimensions []"), + "{}", + err.reason() + ); +} + +#[test] +fn stream_passthrough_projects_argument_scalars() -> Result<(), LowerError> { let dae_model = dae::Dae::default(); let structural_bindings = IndexMap::new(); let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); let mut scope = FunctionProjectionScope::default(); scope .scalars - .insert("tau_inv".to_string(), vec![real(17.0)]); - let expr = rumoca_core::Expression::VarRef { - name: rumoca_core::Reference::new("tau_inv"), - subscripts: Vec::new(), + .insert("u".to_string(), vec![real(1.0), real(2.0)]); + scope.dims.insert("u".to_string(), vec![2]); + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("inStream").into(), + args: vec![rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("u"), + subscripts: Vec::new(), + span: test_span(), + }], + is_constructor: false, span: test_span(), }; - assert_eq!( - analysis.expr_dims(&expr, &scope, 0, test_span()), - Ok(Some(Vec::new())) - ); -} - -#[test] -fn scalar_projected_output_uses_empty_selector_indices() -> Result<(), LowerError> { - let outputs = project_target_scalar_outputs(&[], vec![real(3.0)], test_span())?; + let projected = analysis + .project_value_scalars(&expr, &[2], &scope, 0, test_span())? + .expect("stream passthrough should project argument scalars"); + let values = projected + .iter() + .map(|expr| match expr { + rumoca_core::Expression::Literal { + value: Literal::Real(value), + .. + } => *value, + other => panic!("expected scalar literal, got {other:?}"), + }) + .collect::>(); - assert_eq!(outputs.len(), 1); - assert!(outputs[0].selector_indices.is_empty()); + assert_eq!(values, vec![1.0, 2.0]); + let dims = analysis + .expr_dims(&expr, &scope, 0, test_span())? + .expect("stream passthrough should infer argument dimensions"); + assert_eq!(dims, vec![2]); Ok(()) } #[test] -fn literal_binary_operand_has_known_scalar_dimensions() { +fn expression_dimension_inference_stops_at_inline_depth_limit() -> Result<(), LowerError> { let dae_model = dae::Dae::default(); let structural_bindings = IndexMap::new(); let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); let scope = FunctionProjectionScope::default(); - let expr = binary( - rumoca_core::OpBinary::Add, - real(1.0), - array(vec![real(2.0), real(3.0)], false), + + let dims = analysis.expr_dims( + &real(1.0), + &scope, + super::super::super::MAX_FUNCTION_INLINE_DEPTH + 1, test_span(), - ); + )?; - assert_eq!( - analysis.expr_dims(&expr, &scope, 0, test_span()), - Ok(Some(vec![2])) - ); + assert_eq!(dims, None); + Ok(()) } #[test] -fn scoped_full_binding_has_substituted_dimensions() { +fn cat_projection_concatenates_dynamic_vector_and_computed_tail() -> Result<(), LowerError> { let dae_model = dae::Dae::default(); let structural_bindings = IndexMap::new(); let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); let mut scope = FunctionProjectionScope::default(); - scope.full.insert("gain".to_string(), real(5.0)); - let expr = local_var("gain"); + scope.scalars.insert("u".to_string(), vec![real(0.25)]); + scope.dims.insert("u".to_string(), vec![1]); + let expr = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Cat, + args: vec![ + integer(1), + var_ref("u"), + array( + vec![binary( + rumoca_core::OpBinary::Sub, + integer(1), + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Sum, + args: vec![var_ref("u")], + span: test_span(), + }, + test_span(), + )], + false, + ), + ], + span: test_span(), + }; - assert_eq!( - analysis.expr_dims(&expr, &scope, 0, test_span()), - Ok(Some(Vec::new())) - ); + let dims = analysis + .expr_dims(&expr, &scope, 0, test_span())? + .expect("cat dimensions should be inferred"); + assert_eq!(dims, vec![2]); + let projected = analysis + .project_value_scalars(&expr, &[2], &scope, 0, test_span())? + .expect("cat should project to scalar values"); + + assert_eq!(projected.len(), 2); + assert!(matches!( + &projected[0], + rumoca_core::Expression::Literal { + value: Literal::Real(value), + .. + } if (*value - 0.25).abs() < 1e-12 + )); + assert!(matches!( + &projected[1], + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + .. + } + )); + Ok(()) } #[test] -fn projected_scope_dimensions_override_full_binding_dimensions() { +fn full_binding_substitution_rewrites_nested_local_inputs() -> Result<(), LowerError> { let dae_model = dae::Dae::default(); let structural_bindings = IndexMap::new(); let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); let mut scope = FunctionProjectionScope::default(); - scope.full.insert("x".to_string(), real(5.0)); - scope.dims.insert("x".to_string(), vec![2]); - let expr = local_var("x"); - - assert_eq!( - analysis.expr_dims(&expr, &scope, 0, test_span()), - Ok(Some(vec![2])) + scope.full.insert("p".to_string(), real(101325.0)); + scope.full.insert( + "state".to_string(), + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("ThermodynamicState").into(), + args: vec![rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("__rumoca_named_arg__.p").into(), + args: vec![rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("p"), + subscripts: Vec::new(), + span: test_span(), + }], + is_constructor: true, + span: test_span(), + }], + is_constructor: true, + span: test_span(), + }, ); + + let substituted = analysis.substitute( + &rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("state"), + subscripts: Vec::new(), + span: test_span(), + }, + &scope, + )?; + + let rumoca_core::Expression::FunctionCall { args, .. } = substituted else { + panic!("expected substituted record constructor"); + }; + let rumoca_core::Expression::FunctionCall { args, .. } = &args[0] else { + panic!("expected named argument constructor"); + }; + assert!(matches!( + args.as_slice(), + [rumoca_core::Expression::Literal { + value: Literal::Real(101325.0), + .. + }] + )); + Ok(()) } #[test] -fn array_of_vector_values_infers_matrix_dimensions() { +fn named_constructor_field_access_projects_selected_actual() -> Result<(), LowerError> { let dae_model = dae::Dae::default(); let structural_bindings = IndexMap::new(); let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); let mut scope = FunctionProjectionScope::default(); - scope.dims.insert("row1".to_string(), vec![3]); - scope.dims.insert("row2".to_string(), vec![3]); - scope.dims.insert("row3".to_string(), vec![3]); - let expr = array( - vec![local_var("row1"), local_var("row2"), local_var("row3")], - false, - ); + scope.full.insert("p".to_string(), real(101325.0)); + scope.full.insert("T".to_string(), real(295.0)); + let expr = rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("ThermodynamicState").into(), + args: vec![ + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("__rumoca_named_arg__.p").into(), + args: vec![rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("p"), + subscripts: Vec::new(), + span: test_span(), + }], + is_constructor: true, + span: test_span(), + }, + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("__rumoca_named_arg__.T").into(), + args: vec![rumoca_core::Expression::Binary { + op: OpBinary::Add, + lhs: Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("p"), + subscripts: Vec::new(), + span: test_span(), + }), + rhs: Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("T"), + subscripts: Vec::new(), + span: test_span(), + }), + span: test_span(), + }], + is_constructor: true, + span: test_span(), + }, + ], + is_constructor: true, + span: test_span(), + }), + field: "p".to_string(), + span: test_span(), + }; - assert_eq!( - analysis.expr_dims(&expr, &scope, 0, test_span()), - Ok(Some(vec![3, 3])) - ); + let projected = analysis + .project_value_scalars(&expr, &[], &scope, 0, test_span())? + .expect("named constructor field should project"); + + assert!(matches!( + projected.as_slice(), + [rumoca_core::Expression::Literal { + value: Literal::Real(101325.0), + .. + }] + )); + Ok(()) } #[test] -fn array_with_cross_row_infers_matrix_dimensions() { - let dae_model = dae::Dae::default(); +fn function_record_field_access_projects_if_constructor_output() -> Result<(), LowerError> { + let mut dae_model = dae::Dae::default(); + let mut state_ctor = rumoca_core::Function::new("My.State", test_span()); + state_ctor.is_constructor = true; + state_ctor.pure = true; + state_ctor.inputs.push(scalar_function_param("p")); + state_ctor.inputs.push(scalar_function_param("T")); + state_ctor + .outputs + .push(record_function_param("state", "My.State")); + dae_model + .symbols + .functions + .insert(rumoca_core::VarName::new("My.State"), state_ctor); + + let mut make_state = rumoca_core::Function::new("My.makeState", test_span()); + make_state.pure = true; + make_state.inputs.push(scalar_function_param("p")); + make_state.inputs.push(scalar_function_param("T")); + make_state + .outputs + .push(record_function_param("state", "My.State")); + make_state.body.push(scalar_assignment( + "state", + rumoca_core::Expression::If { + branches: vec![( + rumoca_core::Expression::Binary { + op: OpBinary::Eq, + lhs: Box::new(local_var("p")), + rhs: Box::new(local_var("p")), + span: test_span(), + }, + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.State").into(), + args: vec![ + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("__rumoca_named_arg__.p").into(), + args: vec![local_var("p")], + is_constructor: true, + span: test_span(), + }, + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("__rumoca_named_arg__.T").into(), + args: vec![rumoca_core::Expression::Binary { + op: OpBinary::Add, + lhs: Box::new(local_var("p")), + rhs: Box::new(local_var("T")), + span: test_span(), + }], + is_constructor: true, + span: test_span(), + }, + ], + is_constructor: true, + span: test_span(), + }, + )], + else_branch: Box::new(rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.State").into(), + args: vec![local_var("p"), local_var("T")], + is_constructor: true, + span: test_span(), + }), + span: test_span(), + }, + )); + dae_model + .symbols + .functions + .insert(rumoca_core::VarName::new("My.makeState"), make_state); + let structural_bindings = IndexMap::new(); let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); let mut scope = FunctionProjectionScope::default(); - scope.dims.insert("e_x".to_string(), vec![3]); - scope.dims.insert("e_z".to_string(), vec![3]); - let expr = array( - vec![ - local_var("e_x"), - builtin( - rumoca_core::BuiltinFunction::Cross, - vec![local_var("e_z"), local_var("e_x")], - ), - local_var("e_z"), - ], - false, - ); + scope.full.insert("p".to_string(), real(101325.0)); + scope.full.insert("T".to_string(), real(295.0)); + let expr = rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.makeState").into(), + args: vec![local_var("p"), local_var("T")], + is_constructor: false, + span: test_span(), + }), + field: "p".to_string(), + span: test_span(), + }; - assert_eq!( - analysis.expr_dims(&expr, &scope, 0, test_span()), - Ok(Some(vec![3, 3])) - ); + let projected = analysis + .project_value_scalars(&expr, &[], &scope, 0, test_span())? + .expect("function record field should project"); + + assert!(matches!( + projected.as_slice(), + [rumoca_core::Expression::Literal { + value: Literal::Real(101325.0), + .. + }] + )); + Ok(()) } #[test] -fn array_with_cross_row_projects_nested_scalar() -> Result<(), LowerError> { - let dae_model = dae::Dae::default(); +fn flattened_record_inputs_project_positional_record_actual_fields() -> Result<(), LowerError> { + let mut dae_model = dae::Dae::default(); + let mut state_ctor = rumoca_core::Function::new("My.State", test_span()); + state_ctor.is_constructor = true; + state_ctor.pure = true; + state_ctor.inputs.push(scalar_function_param("p")); + state_ctor.inputs.push(scalar_function_param("T")); + state_ctor + .outputs + .push(record_function_param("state", "My.State")); + dae_model + .symbols + .functions + .insert(rumoca_core::VarName::new("My.State"), state_ctor); + + let mut make_state = rumoca_core::Function::new("My.makeState", test_span()); + make_state.pure = true; + make_state.inputs.push(scalar_function_param("p")); + make_state.inputs.push(scalar_function_param("T")); + make_state + .outputs + .push(record_function_param("state", "My.State")); + make_state.body.push(scalar_assignment( + "state", + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.State").into(), + args: vec![local_var("p"), local_var("T")], + is_constructor: true, + span: test_span(), + }, + )); + dae_model + .symbols + .functions + .insert(rumoca_core::VarName::new("My.makeState"), make_state); + let structural_bindings = IndexMap::new(); let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); let mut scope = FunctionProjectionScope::default(); - scope.dims.insert("e_x".to_string(), vec![3]); - scope.dims.insert("e_z".to_string(), vec![3]); - let expr = array( - vec![ - local_var("e_x"), - builtin( - rumoca_core::BuiltinFunction::Cross, - vec![local_var("e_z"), local_var("e_x")], - ), - local_var("e_z"), - ], - false, - ); - - let projected = analysis - .project_value(&expr, &[3, 3], 3, &scope, 0, test_span())? - .expect("array row vector expression should project to a scalar"); - - let rumoca_core::Expression::Index { - base, subscripts, .. - } = projected - else { - panic!("expected indexed cross-product expression, got {projected:?}"); - }; - let rumoca_core::Expression::BuiltinCall { function, .. } = base.as_ref() else { - panic!("expected indexed cross-product base, got {base:?}"); + scope.full.insert("p".to_string(), real(101325.0)); + scope.full.insert("T".to_string(), real(295.0)); + let mut enthalpy = rumoca_core::Function::new("My.specificEnthalpy", test_span()); + enthalpy.inputs.push(scalar_function_param("state_p")); + enthalpy.inputs.push(scalar_function_param("state_T")); + let actual = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.makeState").into(), + args: vec![local_var("p"), local_var("T")], + is_constructor: false, + span: test_span(), }; - let [rumoca_core::Subscript::Index { value, .. }] = subscripts.as_slice() else { - panic!("expected one generated subscript, got {subscripts:?}"); + let p_field = rumoca_core::Expression::FieldAccess { + base: Box::new(actual.clone()), + field: "p".to_string(), + span: test_span(), }; + let direct_projected = analysis + .project_value_scalars(&p_field, &[], &scope, 0, test_span())? + .expect("direct record field projection should return a scalar"); + assert!(matches!( + direct_projected.as_slice(), + [rumoca_core::Expression::Literal { + value: Literal::Real(value), + .. + }] if (*value - 101325.0).abs() < 1e-12 + )); - assert_eq!(*function, rumoca_core::BuiltinFunction::Cross); - assert_eq!(*value, 1); + let projected = analysis + .bind_inputs(&enthalpy, &[actual], 0, test_span())? + .expect("flattened record inputs should bind from positional record actual"); + + assert_var_ref_name( + projected.full.get("state_p").expect("state_p should bind"), + "p", + ); + assert_var_ref_name( + projected.full.get("state_T").expect("state_T should bind"), + "T", + ); Ok(()) } #[test] -fn dae_scalar_variable_has_known_scalar_dimensions() { - let mut dae_model = dae::Dae::default(); - dae_model.variables.states.insert( - rumoca_core::VarName::new("angle"), - dae::Variable { - name: rumoca_core::VarName::new("angle"), - ..rumoca_ir_dae::Variable::empty_with_span(rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name(file!()), - 1, - 2, - )) - }, - ); +fn flattened_record_like_inputs_bind_multiple_scalar_positionals_directly() -> Result<(), LowerError> +{ + let dae_model = dae::Dae::default(); let structural_bindings = IndexMap::new(); let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); - let scope = FunctionProjectionScope::default(); - let expr = local_var("angle"); + let mut density = rumoca_core::Function::new("My.density_pTX", test_span()); + density.inputs.push(scalar_function_param("state_p")); + density.inputs.push(scalar_function_param("state_T")); + density.inputs.push(scalar_function_param("state_X")); - assert_eq!( - analysis.expr_dims(&expr, &scope, 0, test_span()), - Ok(Some(Vec::new())) + let projected = analysis + .bind_inputs( + &density, + &[real(101325.0), local_var("T"), local_var("X")], + 0, + test_span(), + )? + .expect("flattened scalar positionals should bind directly"); + + assert!(matches!( + projected + .full + .get("state_p") + .expect("state_p should bind"), + rumoca_core::Expression::Literal { + value: Literal::Real(value), + .. + } if (*value - 101325.0).abs() < 1e-12 + )); + assert_var_ref_name( + projected.full.get("state_T").expect("state_T should bind"), + "T", ); + assert_var_ref_name( + projected.full.get("state_X").expect("state_X should bind"), + "X", + ); + Ok(()) } #[test] -fn projected_function_field_outputs_infer_dense_selector_dimensions() -> Result<(), LowerError> { - let outputs = vec![ - ProjectedFunctionOutput { - field_path: vec!["w".to_string()], - selector_indices: vec![1], - expr: real(1.0), - }, - ProjectedFunctionOutput { - field_path: vec!["w".to_string()], - selector_indices: vec![2], - expr: real(2.0), - }, - ProjectedFunctionOutput { - field_path: vec!["w".to_string()], - selector_indices: vec![3], - expr: real(3.0), - }, - ]; +fn vector_output_projection_scalarizes_ordinary_division_by_lane() -> Result<(), LowerError> { + let mut function = rumoca_core::Function::new("My.vectorDiv", test_span()); + function.inputs.push(function_param_with_dims("a", &[3])); + function.inputs.push(function_param_with_dims("b", &[3])); + function.outputs.push(function_param_with_dims("y", &[3])); + function.body.push(scalar_assignment( + "y", + binary( + rumoca_core::OpBinary::Div, + local_var("a"), + local_var("b"), + test_span(), + ), + )); - let dims = projected_field_output_dims(&outputs, "w", test_span())?; + let mut dae_model = dae::Dae::default(); + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + let structural_bindings = IndexMap::new(); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.vectorDiv").into(), + args: vec![ + array(vec![real(2.0), real(4.0), real(6.0)], false), + array(vec![real(1.0), real(2.0), real(3.0)], false), + ], + is_constructor: false, + span: test_span(), + }; - assert_eq!(dims, Some(vec![3])); + let values = function_call_projected_scalars_with_owner( + &call, + &dae_model, + &structural_bindings, + test_span(), + )? + .expect("vector output division should project by lane"); + + assert_eq!(values.len(), 3); + assert!(values.iter().all(|value| matches!( + value, + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Div, + lhs, + rhs, + .. + } if matches!(lhs.as_ref(), rumoca_core::Expression::Literal { .. }) + && matches!(rhs.as_ref(), rumoca_core::Expression::Literal { .. }) + ))); Ok(()) } #[test] -fn repeated_scalar_field_outputs_have_unknown_dimensions() -> Result<(), LowerError> { - let outputs = vec![ - ProjectedFunctionOutput { - field_path: vec!["record".to_string()], - selector_indices: Vec::new(), - expr: real(1.0), - }, - ProjectedFunctionOutput { - field_path: vec!["record".to_string()], - selector_indices: Vec::new(), - expr: real(2.0), - }, - ]; +fn scalar_output_projection_preserves_vector_division_as_elementwise_operand() +-> Result<(), LowerError> { + fn count_elementwise_divisions(expr: &rumoca_core::Expression) -> usize { + let rumoca_core::Expression::Binary { op, lhs, rhs, .. } = expr else { + return 0; + }; + usize::from(matches!(op, rumoca_core::OpBinary::DivElem)) + + count_elementwise_divisions(lhs) + + count_elementwise_divisions(rhs) + } - let dims = projected_field_output_dims(&outputs, "record", test_span())?; + let mut function = rumoca_core::Function::new("My.scalarDotDiv", test_span()); + function.inputs.push(function_param_with_dims("a", &[3])); + function.inputs.push(function_param_with_dims("b", &[3])); + function.inputs.push(function_param_with_dims("c", &[3])); + function.outputs.push(scalar_function_param("y")); + function.body.push(scalar_assignment( + "y", + binary( + rumoca_core::OpBinary::Mul, + binary( + rumoca_core::OpBinary::Div, + local_var("a"), + local_var("b"), + test_span(), + ), + local_var("c"), + test_span(), + ), + )); - assert_eq!(dims, None); + let mut dae_model = dae::Dae::default(); + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + let structural_bindings = IndexMap::new(); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.scalarDotDiv").into(), + args: vec![ + array(vec![real(2.0), real(4.0), real(6.0)], false), + array(vec![real(1.0), real(2.0), real(3.0)], false), + array(vec![real(10.0), real(20.0), real(30.0)], false), + ], + is_constructor: false, + span: test_span(), + }; + + let values = function_call_projected_scalars_with_owner( + &call, + &dae_model, + &structural_bindings, + test_span(), + )? + .expect("scalar output with vector division should project"); + + assert_eq!(values.len(), 1); + assert_eq!(count_elementwise_divisions(&values[0]), 3); Ok(()) } #[test] -fn array_binary_projection_rejects_unknown_operand_dimensions_with_span() { +fn projection_dims_preserve_range_slice_and_scalar_division_shape() -> Result<(), LowerError> { let dae_model = dae::Dae::default(); let structural_bindings = IndexMap::new(); let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); - let scope = FunctionProjectionScope::default(); - let span = rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name( - "phase_solve_lower_derivative_rhs_function_projection_tests_source_51.mo", - ), - 4, - 17, - ); - let expr = binary( - rumoca_core::OpBinary::Add, - local_var("runtime_value"), - array(vec![real(2.0), real(3.0)], false), - span, - ); - let ctx = ProjectionValueCtx { - dims: &[2], - flat_index: 0, - scope: &scope, - depth: 0, + let mut scope = FunctionProjectionScope::default(); + scope.dims.insert("roughnesses".to_string(), vec![4]); + + let span = test_span(); + let slice = rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new("roughnesses").into(), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(rumoca_core::Expression::Range { + start: Box::new(integer(1)), + step: None, + end: Box::new(integer(3)), + span, + }), + span, + }], span, }; + assert_eq!( + analysis + .expr_dims(&slice, &scope, 0, span)? + .expect("range slice dims should infer"), + vec![3] + ); - let rumoca_core::Expression::Binary { lhs, rhs, op, .. } = &expr else { - panic!("test expression must be binary"); - }; - let err = analysis - .project_binary_value(op, lhs, rhs, &ctx) - .expect_err("unknown operand dimensions must bubble a typed error"); - - assert_eq!(err.source_span(), Some(span)); + let expr = binary(rumoca_core::OpBinary::Div, real(0.0065), slice, span); assert_eq!( - err.reason(), - "binary lhs has unknown dimensions".to_string() + analysis + .expr_dims(&expr, &scope, 0, span)? + .expect("scalar/range division dims should infer"), + vec![3] ); + Ok(()) } #[test] -fn checked_usize_dimension_rejects_i64_overflow_with_span() { - let Some(dim) = usize::try_from(i64::MAX) - .ok() - .and_then(|value| value.checked_add(1)) - else { - return; +fn range_index_assignment_keeps_slice_shape() -> Result<(), LowerError> { + let span = test_span(); + let mut function = rumoca_core::Function::new("My.rangeSlice", span); + function.inputs.push(function_param_with_dims("xi", &[9])); + function + .outputs + .push(function_param_with_dims("omega", &[3])); + function.body.push(scalar_assignment( + "omega", + rumoca_core::Expression::Index { + base: Box::new(local_var("xi")), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(rumoca_core::Expression::Range { + start: Box::new(integer(7)), + step: None, + end: Box::new(integer(9)), + span, + }), + span, + }], + span, + }, + )); + + let mut dae_model = dae::Dae::default(); + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + let structural_bindings = IndexMap::new(); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.rangeSlice").into(), + args: vec![array((1..=9).map(integer).collect(), false)], + is_constructor: false, + span, }; - let span = rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name( - "phase_solve_lower_derivative_rhs_function_projection_tests_source_47.mo", - ), - 8, - 19, - ); - let err = checked_usize_dims_to_i64(&[dim], "array expression dimension", span) - .expect_err("dimension must fit in Modelica integer range"); + let values = + function_call_projected_scalars_with_owner(&call, &dae_model, &structural_bindings, span)? + .expect("range slice output should project with dimensions [3]"); - assert_eq!(err.source_span(), Some(span)); - assert_eq!( - err.reason(), - format!("invalid IR contract: array expression dimension {dim} exceeds i64 range") - ); + let projected_indices = values + .iter() + .map(|value| { + let rumoca_core::Expression::Index { subscripts, .. } = value else { + panic!("range slice scalar output should remain indexed: {value:?}"); + }; + let [rumoca_core::Subscript::Index { value, .. }] = subscripts.as_slice() else { + panic!("range slice scalar output should have one scalar selector: {value:?}"); + }; + *value + }) + .collect::>(); + assert_eq!(projected_indices, vec![1, 2, 3]); + Ok(()) } #[test] -fn checked_projection_offset_rejects_host_index_overflow_with_span() -> Result<(), String> { - let span = rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name( - "phase_solve_lower_derivative_rhs_function_projection_tests_source_48.mo", - ), - 3, - 11, +fn projection_dims_preserve_matrix_slice_with_dynamic_scalar_index() -> Result<(), LowerError> { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let mut scope = FunctionProjectionScope::default(); + scope.dims.insert("vertices".to_string(), vec![2, 3]); + scope.dims.insert("other".to_string(), vec![3, 3]); + scope.dims.insert("next_vertex".to_string(), vec![3]); + scope.dims.insert("vertex".to_string(), vec![]); + scope.dims.insert("lo".to_string(), vec![]); + scope.dims.insert("hi".to_string(), vec![]); + + let span = test_span(); + let dynamic_index = rumoca_core::Expression::Index { + base: Box::new(local_var("next_vertex")), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(local_var("vertex")), + span, + }], + span, + }; + let slice = |name: &str, second| rumoca_core::Expression::Index { + base: Box::new(local_var(name)), + subscripts: vec![rumoca_core::Subscript::colon(span), second], + span, + }; + assert_eq!( + analysis.expr_dims(&dynamic_index, &scope, 0, span)?, + Some(vec![]) + ); + let dynamic_slice = slice( + "vertices", + rumoca_core::Subscript::Expr { + expr: Box::new(dynamic_index), + span, + }, + ); + let literal_slice = slice("vertices", rumoca_core::Subscript::index(1, span)); + let subtraction = binary( + rumoca_core::OpBinary::Sub, + dynamic_slice.clone(), + literal_slice.clone(), + span, ); - let Err(mul_err) = - checked_projection_offset(usize::MAX, 2, 0, "matrix product flat index", span) - else { - return Err("overflowing projection offset multiplication succeeded".to_string()); - }; - assert_eq!(mul_err.source_span(), Some(span)); assert_eq!( - mul_err.reason(), - "invalid IR contract: matrix product flat index multiplication overflows host index range" - .to_string() + analysis.expr_dims(&dynamic_slice, &scope, 0, span)?, + Some(vec![2]) + ); + assert_eq!( + analysis.expr_dims(&literal_slice, &scope, 0, span)?, + Some(vec![2]) + ); + assert_eq!( + analysis.expr_dims(&subtraction, &scope, 0, span)?, + Some(vec![2]) ); - let Err(add_err) = - checked_projection_offset(usize::MAX, 1, 1, "matrix product flat index", span) - else { - return Err("overflowing projection offset addition succeeded".to_string()); + let mismatched = binary( + rumoca_core::OpBinary::Sub, + literal_slice, + slice("other", rumoca_core::Subscript::index(1, span)), + span, + ); + assert_eq!(analysis.expr_dims(&mismatched, &scope, 0, span)?, None); + + let array_selector = slice( + "vertices", + rumoca_core::Subscript::Expr { + expr: Box::new(local_var("next_vertex")), + span, + }, + ); + assert_eq!(analysis.expr_dims(&array_selector, &scope, 0, span)?, None); + + let range = |start, end| rumoca_core::Subscript::Expr { + expr: Box::new(rumoca_core::Expression::Range { + start: Box::new(start), + step: None, + end: Box::new(end), + span, + }), + span, }; - assert_eq!(add_err.source_span(), Some(span)); + let known_range = slice("vertices", range(integer(1), integer(2))); assert_eq!( - add_err.reason(), - "invalid IR contract: matrix product flat index addition overflows host index range" - .to_string() + analysis.expr_dims(&known_range, &scope, 0, span)?, + Some(vec![2, 2]) ); - + let unknown_range = slice("vertices", range(local_var("lo"), local_var("hi"))); + assert_eq!(analysis.expr_dims(&unknown_range, &scope, 0, span)?, None); Ok(()) } #[test] -fn checked_projection_offset_dummy_span_stays_unspanned() { - let err = checked_projection_offset( - usize::MAX, - 2, - 0, - "matrix product flat index", - rumoca_core::Span::DUMMY, - ) - .expect_err("overflowing projection offset multiplication must fail"); +fn nested_flattened_record_call_uses_caller_projection_scope() -> Result<(), LowerError> { + let mut dae_model = dae::Dae::default(); + let mut state_ctor = rumoca_core::Function::new("My.State", test_span()); + state_ctor.is_constructor = true; + state_ctor.pure = true; + state_ctor.inputs.push(scalar_function_param("p")); + state_ctor.inputs.push(scalar_function_param("T")); + state_ctor + .outputs + .push(record_function_param("state", "My.State")); + dae_model + .symbols + .functions + .insert(rumoca_core::VarName::new("My.State"), state_ctor); + + let mut make_state = rumoca_core::Function::new("My.makeState", test_span()); + make_state.pure = true; + make_state.inputs.push(scalar_function_param("p")); + make_state.inputs.push(scalar_function_param("T")); + make_state + .outputs + .push(record_function_param("state", "My.State")); + make_state.body.push(scalar_assignment( + "state", + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.State").into(), + args: vec![local_var("p"), local_var("T")], + is_constructor: true, + span: test_span(), + }, + )); + dae_model + .symbols + .functions + .insert(rumoca_core::VarName::new("My.makeState"), make_state); + + let mut density = rumoca_core::Function::new("My.density", test_span()); + density.pure = true; + density.inputs.push(scalar_function_param("state_p")); + density.inputs.push(scalar_function_param("state_T")); + density.outputs.push(scalar_function_param("d")); + density + .body + .push(scalar_assignment("d", local_var("state_p"))); + dae_model + .symbols + .functions + .insert(rumoca_core::VarName::new("My.density"), density); + + let mut use_density = rumoca_core::Function::new("My.useDensity", test_span()); + use_density.pure = true; + use_density.inputs.push(scalar_function_param("state_p")); + use_density.inputs.push(scalar_function_param("state_T")); + use_density.outputs.push(scalar_function_param("y")); + use_density.body.push(scalar_assignment( + "y", + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.density").into(), + args: vec![local_var("state")], + is_constructor: false, + span: test_span(), + }, + )); + dae_model + .symbols + .functions + .insert(rumoca_core::VarName::new("My.useDensity"), use_density); - assert!( - matches!(err, LowerError::UnspannedContractViolation { .. }), - "dummy projection offset span should not be fabricated into a source span: {err:?}" - ); - assert!(err.reason().contains("multiplication overflows")); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.useDensity").into(), + args: vec![rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.makeState").into(), + args: vec![local_var("p"), local_var("T")], + is_constructor: false, + span: test_span(), + }], + is_constructor: false, + span: test_span(), + }; + + let outputs = analysis + .function_call_outputs_with_owner(&call, 0, test_span())? + .expect("nested flattened call should project outputs"); + assert_eq!(outputs.len(), 1); + assert_var_ref_name(&outputs[0].expr, "p"); + Ok(()) } #[test] -fn checked_usize_dims_to_i64_dummy_span_stays_unspanned() { - let err = checked_usize_dims_to_i64( - &[usize::MAX], - "array expression dimension", - rumoca_core::Span::DUMMY, - ) - .expect_err("dimension must fit in Modelica integer range"); +fn record_field_projection_selects_compile_time_if_branch() -> Result<(), LowerError> { + let mut dae_model = dae::Dae::default(); + let mut state_ctor = rumoca_core::Function::new("My.State", test_span()); + state_ctor.is_constructor = true; + state_ctor.pure = true; + state_ctor.inputs.push(scalar_function_param("p")); + state_ctor + .outputs + .push(record_function_param("state", "My.State")); + dae_model + .symbols + .functions + .insert(rumoca_core::VarName::new("My.State"), state_ctor); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let scope = FunctionProjectionScope::default(); + let value = rumoca_core::Expression::If { + branches: vec![( + binary( + rumoca_core::OpBinary::Eq, + integer(1), + integer(0), + test_span(), + ), + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.State").into(), + args: vec![rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("__rumoca_named_arg__.p").into(), + args: vec![real(1.0)], + is_constructor: true, + span: test_span(), + }], + is_constructor: true, + span: test_span(), + }, + )], + else_branch: Box::new(rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.State").into(), + args: vec![rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("__rumoca_named_arg__.p").into(), + args: vec![real(2.0)], + is_constructor: true, + span: test_span(), + }], + is_constructor: true, + span: test_span(), + }), + span: test_span(), + }; - assert!( - matches!(err, LowerError::UnspannedContractViolation { .. }), - "dummy dimension span should not be fabricated into a source span: {err:?}" - ); - assert!(err.reason().contains("exceeds i64 range")); + let projected = analysis + .project_record_field_value(&value, "p", &scope, test_span())? + .expect("compile-time record field branch should project"); + + assert!(matches!( + projected, + rumoca_core::Expression::Literal { + value: Literal::Real(2.0), + .. + } + )); + Ok(()) } #[test] -fn reserve_projection_capacity_dummy_span_stays_unspanned() { - let mut values = Vec::::new(); - let err = reserve_projection_capacity( - &mut values, - usize::MAX, - "projected output count", - rumoca_core::Span::DUMMY, - ) - .expect_err("impossible projection capacity must be rejected"); +fn scoped_single_scalar_value_has_scalar_dimensions() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let mut scope = FunctionProjectionScope::default(); + scope + .scalars + .insert("tau_inv".to_string(), vec![real(17.0)]); + let expr = rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("tau_inv"), + subscripts: Vec::new(), + span: test_span(), + }; - assert!( - matches!(err, LowerError::UnspannedContractViolation { .. }), - "dummy projection capacity span should not be fabricated into a source span: {err:?}" + assert_eq!( + analysis.expr_dims(&expr, &scope, 0, test_span()), + Ok(Some(Vec::new())) ); - assert!(err.reason().contains("capacity exceeds host memory limits")); } #[test] -fn scalar_count_rejects_host_index_overflow_with_span() { - let span = rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name( - "phase_solve_lower_derivative_rhs_function_projection_tests_source_52.mo", - ), - 1, - 9, - ); - - let err = scalar_count_for_dims(&[i64::MAX, i64::MAX], "projected value dimensions", span) - .expect_err("overflowing scalar count must fail"); +fn scalar_projected_output_uses_empty_selector_indices() -> Result<(), LowerError> { + let outputs = project_target_scalar_outputs(&[], vec![real(3.0)], test_span())?; - assert_eq!(err.source_span(), Some(span)); - assert!(err.reason().contains("projected value dimensions")); + assert_eq!(outputs.len(), 1); + assert!(outputs[0].selector_indices.is_empty()); + Ok(()) } #[test] -fn scalar_count_dummy_span_stays_unspanned() { - let err = scalar_count_for_dims( - &[i64::MAX, i64::MAX], - "projected value dimensions", - rumoca_core::Span::DUMMY, - ) - .expect_err("overflowing scalar count must fail"); +fn literal_binary_operand_has_known_scalar_dimensions() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let scope = FunctionProjectionScope::default(); + let expr = binary( + rumoca_core::OpBinary::Add, + real(1.0), + array(vec![real(2.0), real(3.0)], false), + test_span(), + ); - assert!( - matches!(err, LowerError::UnspannedContractViolation { .. }), - "dummy scalar-count span should not be fabricated into a source span: {err:?}" + assert_eq!( + analysis.expr_dims(&expr, &scope, 0, test_span()), + Ok(Some(vec![2])) ); - assert!(err.reason().contains("projected value dimensions")); } #[test] -fn flat_index_rejects_host_index_overflow_with_span() { - let span = rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name( - "phase_solve_lower_derivative_rhs_function_projection_tests_source_53.mo", +fn scalar_multiplication_has_known_scalar_dimensions() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let scope = FunctionProjectionScope::default(); + let expr = binary( + rumoca_core::OpBinary::Mul, + binary( + rumoca_core::OpBinary::Mul, + integer(2), + builtin(rumoca_core::BuiltinFunction::Asin, vec![real(1.0)]), + test_span(), ), - 1, - 9, + integer(2), + test_span(), ); - let err = flat_index_from_indices( - &[i64::MAX, i64::MAX], - &[i64::MAX, i64::MAX], - span, - "projected scalar selection flat index", - ) - .expect_err("overflowing flat index must fail"); - - assert_eq!(err.source_span(), Some(span)); - assert!( - err.reason() - .contains("projected scalar selection flat index") + assert_eq!( + analysis.expr_dims(&expr, &scope, 0, test_span()), + Ok(Some(Vec::new())) ); } #[test] -fn matrix_matrix_projection_with_zero_columns_declines() -> Result<(), String> { +fn scoped_full_binding_has_substituted_dimensions() { let dae_model = dae::Dae::default(); let structural_bindings = IndexMap::new(); let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); - let scope = FunctionProjectionScope::default(); - let span = rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name( - "phase_solve_lower_derivative_rhs_function_projection_tests_source_49.mo", - ), - 1, - 9, - ); - let ctx = ProjectionValueCtx { - dims: &[], - flat_index: 0, - scope: &scope, - depth: 0, - span, - }; - - let projected = analysis - .project_matrix_matrix_product(&real(1.0), &real(1.0), &[1, 1], &[1, 0], &ctx, 0) - .map_err(|err| format!("zero-column matrix projection failed: {err:?}"))?; - if projected.is_some() { - return Err("zero-column matrix projection produced a scalar value".to_string()); - } + let mut scope = FunctionProjectionScope::default(); + scope.full.insert("gain".to_string(), real(5.0)); + let expr = local_var("gain"); - Ok(()) + assert_eq!( + analysis.expr_dims(&expr, &scope, 0, test_span()), + Ok(Some(Vec::new())) + ); } #[test] -fn project_reference_indices_preserves_indexed_component_parts() { - let reference = component_reference(vec![ - rumoca_core::ComponentRefPart { - ident: "vehicle".to_string(), - span: test_span(), - subs: Vec::new(), - }, - rumoca_core::ComponentRefPart { - ident: "motor".to_string(), - span: test_span(), - subs: vec![rumoca_core::Subscript::generated_index(1, test_span())], - }, - rumoca_core::ComponentRefPart { - ident: "history".to_string(), - span: test_span(), - subs: Vec::new(), - }, - ]); - - let projected = project_reference_field_path_and_indices(&reference, &[], &[2], test_span()) - .expect("structured reference projection should succeed"); +fn projected_scope_dimensions_override_full_binding_dimensions() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let mut scope = FunctionProjectionScope::default(); + scope.full.insert("x".to_string(), real(5.0)); + scope.dims.insert("x".to_string(), vec![2]); + let expr = local_var("x"); - let component_ref = projected - .component_ref() - .expect("projected reference should preserve component-reference structure"); - assert_eq!(projected.as_str(), "vehicle.motor[1].history[2]"); - assert_eq!(component_ref.parts[1].ident, "motor"); - assert_eq!(component_ref.parts[1].subs.len(), 1); - assert_eq!(component_ref.parts[2].ident, "history"); - assert_eq!(component_ref.parts[2].subs.len(), 1); + assert_eq!( + analysis.expr_dims(&expr, &scope, 0, test_span()), + Ok(Some(vec![2])) + ); } #[test] -fn project_reference_indices_rejects_i64_overflow_with_span() { - let Some(index) = usize::try_from(i64::MAX) - .ok() - .and_then(|value| value.checked_add(1)) - else { - return; - }; - let span = rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name( - "phase_solve_lower_derivative_rhs_function_projection_tests_source_46.mo", - ), - 6, - 14, +fn array_of_vector_values_infers_matrix_dimensions() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let mut scope = FunctionProjectionScope::default(); + scope.dims.insert("row1".to_string(), vec![3]); + scope.dims.insert("row2".to_string(), vec![3]); + scope.dims.insert("row3".to_string(), vec![3]); + let expr = array( + vec![local_var("row1"), local_var("row2"), local_var("row3")], + false, ); - let reference = component_reference(vec![rumoca_core::ComponentRefPart { - ident: "x".to_string(), - span, - subs: Vec::new(), - }]); - let err = project_reference_field_path_and_indices(&reference, &[], &[index], span) - .expect_err("projected reference index must fit in Modelica integer range"); - - assert_eq!(err.source_span(), Some(span)); assert_eq!( - err.reason(), - format!( - "invalid IR contract: function output projection subscript index {index} exceeds i64 range" - ) + analysis.expr_dims(&expr, &scope, 0, test_span()), + Ok(Some(vec![3, 3])) ); } -fn scalar_function_param(name: &str) -> rumoca_core::FunctionParam { - rumoca_core::FunctionParam { - def_id: None, - name: name.to_string(), - span: test_span(), - type_name: "Real".to_string(), - type_class: None, - dims: vec![], - shape_expr: Vec::new(), - default: None, - description: None, - } -} - -fn function_param_with_dims(name: &str, dims: &[i64]) -> rumoca_core::FunctionParam { - rumoca_core::FunctionParam { - dims: dims.to_vec(), - ..scalar_function_param(name) - } -} - -fn real_with_span(value: f64, span: rumoca_core::Span) -> rumoca_core::Expression { - rumoca_core::Expression::Literal { - value: Literal::Real(value), - span, - } -} +#[test] +fn array_with_cross_row_infers_matrix_dimensions() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let mut scope = FunctionProjectionScope::default(); + scope.dims.insert("e_x".to_string(), vec![3]); + scope.dims.insert("e_z".to_string(), vec![3]); + let expr = array( + vec![ + local_var("e_x"), + builtin( + rumoca_core::BuiltinFunction::Cross, + vec![local_var("e_z"), local_var("e_x")], + ), + local_var("e_z"), + ], + false, + ); -fn function_param_with_type(name: &str, type_name: &str) -> rumoca_core::FunctionParam { - rumoca_core::FunctionParam { - type_name: type_name.to_string(), - ..scalar_function_param(name) - } + assert_eq!( + analysis.expr_dims(&expr, &scope, 0, test_span()), + Ok(Some(vec![3, 3])) + ); } #[test] -fn vector_function_input_rejects_scalar_actual_with_span() { +fn array_with_cross_row_projects_nested_scalar() -> Result<(), LowerError> { let dae_model = dae::Dae::default(); let structural_bindings = IndexMap::new(); let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); - let mut function = rumoca_core::Function::new("My.needsVector", test_span()); - function.inputs.push(function_param_with_dims("u", &[2])); - let span = rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name( - "phase_solve_lower_derivative_rhs_function_projection_tests_source_52.mo", - ), - 3, - 8, + let mut scope = FunctionProjectionScope::default(); + scope.dims.insert("e_x".to_string(), vec![3]); + scope.dims.insert("e_z".to_string(), vec![3]); + let expr = array( + vec![ + local_var("e_x"), + builtin( + rumoca_core::BuiltinFunction::Cross, + vec![local_var("e_z"), local_var("e_x")], + ), + local_var("e_z"), + ], + false, ); - let err = match analysis.bind_inputs(&function, &[real_with_span(1.0, span)], 0, span) { - Ok(_) => panic!("scalar actual must not be projected as vector input"), - Err(err) => err, + let projected = analysis + .project_value(&expr, &[3, 3], 3, &scope, 0, test_span())? + .expect("array row vector expression should project to a scalar"); + + let rumoca_core::Expression::Index { + base, subscripts, .. + } = projected + else { + panic!("expected indexed cross-product expression, got {projected:?}"); + }; + let rumoca_core::Expression::BuiltinCall { function, .. } = base.as_ref() else { + panic!("expected indexed cross-product base, got {base:?}"); + }; + let [rumoca_core::Subscript::Index { value, .. }] = subscripts.as_slice() else { + panic!("expected one generated subscript, got {subscripts:?}"); }; - assert_eq!(err.source_span(), Some(span)); - assert_eq!( - err.reason(), - "function `My.needsVector` input `u` expects dimensions [2], got []" - ); + assert_eq!(*function, rumoca_core::BuiltinFunction::Cross); + assert_eq!(*value, 1); + Ok(()) } #[test] -fn dynamic_vector_function_input_uses_actual_dimensions() { - let dae_model = dae::Dae::default(); +fn dae_scalar_variable_has_known_scalar_dimensions() { + let mut dae_model = dae::Dae::default(); + dae_model.variables.states.insert( + rumoca_core::VarName::new("angle"), + dae::Variable { + name: rumoca_core::VarName::new("angle"), + ..rumoca_ir_dae::Variable::empty_with_span(rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name(file!()), + 1, + 2, + )) + }, + ); let structural_bindings = IndexMap::new(); let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); - let mut function = rumoca_core::Function::new("My.needsDynamicVector", test_span()); - function.inputs.push(function_param_with_dims("u", &[0])); + let scope = FunctionProjectionScope::default(); + let expr = local_var("angle"); - let scope = analysis - .bind_inputs( - &function, - &[array(vec![real(1.0), real(2.0), real(3.0)], false)], - 0, - test_span(), - ) - .expect("dynamic vector input projection should not fail") - .expect("dynamic vector input should bind"); + assert_eq!( + analysis.expr_dims(&expr, &scope, 0, test_span()), + Ok(Some(Vec::new())) + ); +} - assert_eq!(scope.dims.get("u"), Some(&vec![3])); - assert_eq!(scope.scalars.get("u").map(Vec::len), Some(3)); +#[test] +fn projected_function_field_outputs_infer_dense_selector_dimensions() -> Result<(), LowerError> { + let outputs = vec![ + ProjectedFunctionOutput { + field_path: vec!["w".to_string()], + selector_indices: vec![1], + expr: real(1.0), + }, + ProjectedFunctionOutput { + field_path: vec!["w".to_string()], + selector_indices: vec![2], + expr: real(2.0), + }, + ProjectedFunctionOutput { + field_path: vec!["w".to_string()], + selector_indices: vec![3], + expr: real(3.0), + }, + ]; + + let dims = projected_field_output_dims(&outputs, "w", test_span())?; + + assert_eq!(dims, Some(vec![3])); + Ok(()) +} + +#[test] +fn repeated_scalar_field_outputs_have_unknown_dimensions() -> Result<(), LowerError> { + let outputs = vec![ + ProjectedFunctionOutput { + field_path: vec!["record".to_string()], + selector_indices: Vec::new(), + expr: real(1.0), + }, + ProjectedFunctionOutput { + field_path: vec!["record".to_string()], + selector_indices: Vec::new(), + expr: real(2.0), + }, + ]; + + let dims = projected_field_output_dims(&outputs, "record", test_span())?; + + assert_eq!(dims, None); + Ok(()) } #[test] -fn scalar_real_function_input_rejects_vector_actual_with_span() { +fn array_binary_projection_rejects_unknown_operand_dimensions_with_span() { let dae_model = dae::Dae::default(); let structural_bindings = IndexMap::new(); let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); - let mut function = rumoca_core::Function::new("My.needsScalar", test_span()); - function.inputs.push(scalar_function_param("u")); + let scope = FunctionProjectionScope::default(); let span = rumoca_core::Span::from_offsets( rumoca_core::SourceId::from_source_name( - "phase_solve_lower_derivative_rhs_function_projection_tests_source_54.mo", + "phase_solve_lower_derivative_rhs_function_projection_tests_source_51.mo", ), - 7, - 13, + 4, + 17, ); - - let err = match analysis.bind_inputs( - &function, - &[rumoca_core::Expression::Array { - elements: vec![real(1.0), real(2.0)], - is_matrix: false, - span, - }], - 0, + let expr = binary( + rumoca_core::OpBinary::Add, + local_var("runtime_value"), + array(vec![real(2.0), real(3.0)], false), span, - ) { - Ok(_) => panic!("vector actual must not be projected as scalar Real input"), - Err(err) => err, + ); + let ctx = ProjectionValueCtx { + dims: &[2], + flat_index: 0, + scope: &scope, + depth: 0, + span, + }; + + let rumoca_core::Expression::Binary { lhs, rhs, op, .. } = &expr else { + panic!("test expression must be binary"); }; + let err = analysis + .project_binary_value(op, lhs, rhs, &ctx) + .expect_err("unknown operand dimensions must bubble a typed error"); assert_eq!(err.source_span(), Some(span)); assert_eq!( err.reason(), - "function `My.needsScalar` input `u` expects dimensions [], got [2]" + "binary lhs has unknown dimensions".to_string() ); } #[test] -fn generated_function_call_projection_errors_use_owner_span() { - let owner_span = rumoca_core::Span::from_offsets( +fn checked_usize_dimension_rejects_i64_overflow_with_span() { + let Some(dim) = usize::try_from(i64::MAX) + .ok() + .and_then(|value| value.checked_add(1)) + else { + return; + }; + let span = rumoca_core::Span::from_offsets( rumoca_core::SourceId::from_source_name( - "phase_solve_lower_derivative_rhs_function_projection_tests_source_53.mo", + "phase_solve_lower_derivative_rhs_function_projection_tests_source_47.mo", ), - 3, 8, + 19, ); - let mut dae_model = dae::Dae::default(); - let mut function = rumoca_core::Function::new("My.needsVector", test_span()); - function.inputs.push(rumoca_core::FunctionParam { - span: owner_span, - ..function_param_with_dims("u", &[2]) - }); - function.outputs.push(scalar_function_param("y")); - dae_model - .symbols - .functions - .insert(function.name.clone(), function); - let structural_bindings = IndexMap::new(); - let call = rumoca_core::Expression::FunctionCall { - name: rumoca_core::VarName::new("My.needsVector").into(), - args: vec![rumoca_core::Expression::Literal { - value: Literal::Real(1.0), - span: rumoca_core::Span::DUMMY, - }], - is_constructor: false, - span: rumoca_core::Span::DUMMY, - }; - let err = function_call_projected_scalars_with_owner( - &call, - &dae_model, - &structural_bindings, - owner_span, - ) - .expect_err("generated invalid projection must report an error"); + let err = checked_usize_dims_to_i64(&[dim], "array expression dimension", span) + .expect_err("dimension must fit in Modelica integer range"); - assert_eq!(err.source_span(), Some(owner_span)); + assert_eq!(err.source_span(), Some(span)); assert_eq!( err.reason(), - "function `My.needsVector` input `u` expects dimensions [2], got []" + format!("invalid IR contract: array expression dimension {dim} exceeds i64 range") ); } #[test] -fn record_like_function_input_accepts_structured_actual_dimensions() { - let dae_model = dae::Dae::default(); - let structural_bindings = IndexMap::new(); - let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); - let mut function = rumoca_core::Function::new("My.needsRecord", test_span()); - function - .inputs - .push(function_param_with_type("q", "Pkg.Quaternion")); +fn checked_projection_offset_rejects_host_index_overflow_with_span() -> Result<(), String> { + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name( + "phase_solve_lower_derivative_rhs_function_projection_tests_source_48.mo", + ), + 3, + 11, + ); - let scope = analysis - .bind_inputs( - &function, - &[array( - vec![real(1.0), real(2.0), real(3.0), real(4.0)], - false, - )], - 0, - test_span(), - ) - .expect("record-like input binding should not fail") - .expect("record-like input should bind"); + let Err(mul_err) = + checked_projection_offset(usize::MAX, 2, 0, "matrix product flat index", span) + else { + return Err("overflowing projection offset multiplication succeeded".to_string()); + }; + assert_eq!(mul_err.source_span(), Some(span)); + assert_eq!( + mul_err.reason(), + "invalid IR contract: matrix product flat index multiplication overflows host index range" + .to_string() + ); - assert_eq!(scope.dims.get("q"), Some(&vec![4])); - assert_eq!(scope.scalars.get("q").map(Vec::len), Some(4)); + let Err(add_err) = + checked_projection_offset(usize::MAX, 1, 1, "matrix product flat index", span) + else { + return Err("overflowing projection offset addition succeeded".to_string()); + }; + assert_eq!(add_err.source_span(), Some(span)); + assert_eq!( + add_err.reason(), + "invalid IR contract: matrix product flat index addition overflows host index range" + .to_string() + ); + + Ok(()) } #[test] -fn function_projection_initializes_array_local_from_declaration_binding() -> Result<(), LowerError> -{ - let mut dae_model = dae::Dae::default(); - let structural_bindings = IndexMap::new(); - let mut function = rumoca_core::Function::new("My.localDefault", test_span()); - function.outputs.push(function_param_with_dims("y", &[3])); - let mut local = function_param_with_dims("x", &[3]); - local.default = Some(array(vec![real(1.0), real(2.0), real(3.0)], false)); - function.locals.push(local); - function.body.push(scalar_assignment("y", local_var("x"))); - dae_model - .symbols - .functions - .insert(function.name.clone(), function); - let call = rumoca_core::Expression::FunctionCall { - name: rumoca_core::VarName::new("My.localDefault").into(), - args: Vec::new(), - is_constructor: false, - span: test_span(), - }; - - let projected = function_call_projected_scalars_with_owner( - &call, - &dae_model, - &structural_bindings, - test_span(), - )? - .expect("function call should project declaration-bound array local"); - let values = projected - .iter() - .map(|expr| match expr { - rumoca_core::Expression::Literal { - value: Literal::Real(value), - .. - } => *value, - other => panic!("expected real literal projection, got {other:?}"), - }) - .collect::>(); +fn checked_projection_offset_dummy_span_stays_unspanned() { + let err = checked_projection_offset( + usize::MAX, + 2, + 0, + "matrix product flat index", + rumoca_core::Span::DUMMY, + ) + .expect_err("overflowing projection offset multiplication must fail"); - assert_eq!(values, vec![1.0, 2.0, 3.0]); - Ok(()) + assert!( + matches!(err, LowerError::UnspannedContractViolation { .. }), + "dummy projection offset span should not be fabricated into a source span: {err:?}" + ); + assert!(err.reason().contains("multiplication overflows")); } #[test] -fn vector_constructor_input_rejects_scalar_actual_with_span() { - let dae_model = dae::Dae::default(); - let structural_bindings = IndexMap::new(); - let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); - let scope = FunctionProjectionScope::default(); - let input = function_param_with_dims("u", &[2]); - let span = rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name( - "phase_solve_lower_derivative_rhs_function_projection_tests_source_53.mo", - ), - 5, - 11, +fn checked_usize_dims_to_i64_dummy_span_stays_unspanned() { + let err = checked_usize_dims_to_i64( + &[usize::MAX], + "array expression dimension", + rumoca_core::Span::DUMMY, + ) + .expect_err("dimension must fit in Modelica integer range"); + + assert!( + matches!(err, LowerError::UnspannedContractViolation { .. }), + "dummy dimension span should not be fabricated into a source span: {err:?}" ); - let actual = real_with_span(1.0, span); + assert!(err.reason().contains("exceeds i64 range")); +} - let err = analysis - .optional_constructor_input_scalars(&actual, &input, &scope, 0, span) - .expect_err("scalar actual must not be projected as vector constructor input"); +#[test] +fn reserve_projection_capacity_dummy_span_stays_unspanned() { + let mut values = Vec::::new(); + let err = reserve_projection_capacity( + &mut values, + usize::MAX, + "projected output count", + rumoca_core::Span::DUMMY, + ) + .expect_err("impossible projection capacity must be rejected"); - assert_eq!(err.source_span(), Some(span)); - assert_eq!( - err.reason(), - "record constructor input `u` expects dimensions [2], got []" + assert!( + matches!(err, LowerError::UnspannedContractViolation { .. }), + "dummy projection capacity span should not be fabricated into a source span: {err:?}" ); + assert!(err.reason().contains("capacity exceeds host memory limits")); } #[test] -fn scalar_constructor_input_uses_formal_scalar_dimensions() { - let dae_model = dae::Dae::default(); - let structural_bindings = IndexMap::new(); - let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); - let scope = FunctionProjectionScope::default(); - let input = scalar_function_param("u"); +fn scalar_count_rejects_host_index_overflow_with_span() { let span = rumoca_core::Span::from_offsets( rumoca_core::SourceId::from_source_name( - "phase_solve_lower_derivative_rhs_function_projection_tests_source_57.mo", + "phase_solve_lower_derivative_rhs_function_projection_tests_source_52.mo", ), - 5, - 11, + 1, + 9, ); - let actual = rumoca_core::Expression::VarRef { - name: rumoca_core::Reference::new("runtime_scalar"), - subscripts: Vec::new(), - span, - }; - let scalars = analysis - .optional_constructor_input_scalars(&actual, &input, &scope, 0, span) - .expect("scalar constructor input projection should not fail") - .expect("scalar primitive input should project as a scalar"); + let err = scalar_count_for_dims(&[i64::MAX, i64::MAX], "projected value dimensions", span) + .expect_err("overflowing scalar count must fail"); - assert_eq!(scalars.len(), 1); + assert_eq!(err.source_span(), Some(span)); + assert!(err.reason().contains("projected value dimensions")); } #[test] -fn record_like_constructor_input_declines_unknown_actual_dimensions() { - let dae_model = dae::Dae::default(); - let structural_bindings = IndexMap::new(); - let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); - let scope = FunctionProjectionScope::default(); - let input = function_param_with_type("q", "Pkg.Quaternion"); - let span = rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name( - "phase_solve_lower_derivative_rhs_function_projection_tests_source_58.mo", - ), - 6, - 14, - ); - let actual = rumoca_core::Expression::VarRef { - name: rumoca_core::Reference::new("runtime_record"), - subscripts: Vec::new(), - span, - }; - - let scalars = analysis - .optional_constructor_input_scalars(&actual, &input, &scope, 0, span) - .expect("unknown record-like constructor input dimensions should decline"); +fn scalar_count_dummy_span_stays_unspanned() { + let err = scalar_count_for_dims( + &[i64::MAX, i64::MAX], + "projected value dimensions", + rumoca_core::Span::DUMMY, + ) + .expect_err("overflowing scalar count must fail"); - assert!(scalars.is_none()); + assert!( + matches!(err, LowerError::UnspannedContractViolation { .. }), + "dummy scalar-count span should not be fabricated into a source span: {err:?}" + ); + assert!(err.reason().contains("projected value dimensions")); } #[test] -fn if_projection_rejects_scalar_values_without_dimensions() { - let dae_model = dae::Dae::default(); - let structural_bindings = IndexMap::new(); - let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); - let entry_scope = FunctionProjectionScope::default(); - let mut else_scope = FunctionProjectionScope::default(); - else_scope.scalars.insert("x".to_string(), vec![real(0.0)]); +fn flat_index_rejects_host_index_overflow_with_span() { let span = rumoca_core::Span::from_offsets( rumoca_core::SourceId::from_source_name( - "phase_solve_lower_derivative_rhs_function_projection_tests_source_59.mo", + "phase_solve_lower_derivative_rhs_function_projection_tests_source_53.mo", ), - 7, - 18, + 1, + 9, ); - let err = match analysis.merged_if_scope(&entry_scope, &[], &[], &else_scope, span) { - Ok(_) => panic!("merged scalar projection without dimensions must fail"), - Err(err) => err, - }; + let err = flat_index_from_indices( + &[i64::MAX, i64::MAX], + &[i64::MAX, i64::MAX], + span, + "projected scalar selection flat index", + ) + .expect_err("overflowing flat index must fail"); assert_eq!(err.source_span(), Some(span)); - assert_eq!( - err.reason(), - "invalid IR contract: if-statement projection for `x` has scalar values but no dimensions" + assert!( + err.reason() + .contains("projected scalar selection flat index") ); } #[test] -fn if_projection_rejects_conflicting_branch_dimensions() { +fn matrix_matrix_projection_with_zero_columns_declines() -> Result<(), String> { let dae_model = dae::Dae::default(); let structural_bindings = IndexMap::new(); let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); - let mut entry_scope = FunctionProjectionScope::default(); - entry_scope - .scalars - .insert("x".to_string(), vec![real(0.0), real(0.0)]); - entry_scope.dims.insert("x".to_string(), vec![2]); - let mut branch_scope = entry_scope.clone(); - branch_scope.dims.insert("x".to_string(), vec![1, 2]); - let condition = real(1.0); + let scope = FunctionProjectionScope::default(); let span = rumoca_core::Span::from_offsets( rumoca_core::SourceId::from_source_name( - "phase_solve_lower_derivative_rhs_function_projection_tests_source_60.mo", + "phase_solve_lower_derivative_rhs_function_projection_tests_source_49.mo", ), - 8, - 21, + 1, + 9, ); - - let err = match analysis.merged_if_scope( - &entry_scope, - &[condition], - &[branch_scope], - &entry_scope, + let ctx = ProjectionValueCtx { + dims: &[], + flat_index: 0, + scope: &scope, + depth: 0, span, - ) { - Ok(_) => panic!("merged scalar projection with conflicting dimensions must fail"), - Err(err) => err, }; + let projected = analysis + .project_matrix_matrix_product(&real(1.0), &real(1.0), &[1, 1], &[1, 0], &ctx, 0) + .map_err(|err| format!("zero-column matrix projection failed: {err:?}"))?; + if projected.is_some() { + return Err("zero-column matrix projection produced a scalar value".to_string()); + } + + Ok(()) +} + +#[test] +fn project_reference_indices_preserves_indexed_component_parts() { + let reference = component_reference(vec![ + rumoca_core::ComponentRefPart { + ident: "vehicle".to_string(), + span: test_span(), + subs: Vec::new(), + }, + rumoca_core::ComponentRefPart { + ident: "motor".to_string(), + span: test_span(), + subs: vec![rumoca_core::Subscript::generated_index(1, test_span())], + }, + rumoca_core::ComponentRefPart { + ident: "history".to_string(), + span: test_span(), + subs: Vec::new(), + }, + ]); + + let projected = project_reference_field_path_and_indices(&reference, &[], &[2], test_span()) + .expect("structured reference projection should succeed"); + + let component_ref = projected + .component_ref() + .expect("projected reference should preserve component-reference structure"); + assert_eq!(projected.as_str(), "vehicle.motor[1].history[2]"); + assert_eq!(component_ref.parts[1].ident, "motor"); + assert_eq!(component_ref.parts[1].subs.len(), 1); + assert_eq!(component_ref.parts[2].ident, "history"); + assert_eq!(component_ref.parts[2].subs.len(), 1); +} + +#[test] +fn project_reference_indices_rejects_i64_overflow_with_span() { + let Some(index) = usize::try_from(i64::MAX) + .ok() + .and_then(|value| value.checked_add(1)) + else { + return; + }; + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name( + "phase_solve_lower_derivative_rhs_function_projection_tests_source_46.mo", + ), + 6, + 14, + ); + let reference = component_reference(vec![rumoca_core::ComponentRefPart { + ident: "x".to_string(), + span, + subs: Vec::new(), + }]); + + let err = project_reference_field_path_and_indices(&reference, &[], &[index], span) + .expect_err("projected reference index must fit in Modelica integer range"); + assert_eq!(err.source_span(), Some(span)); assert_eq!( err.reason(), - "invalid IR contract: if-statement projection for `x` has mismatched dimensions: [2] and [1, 2]" + format!( + "invalid IR contract: function output projection subscript index {index} exceeds i64 range" + ) ); } -fn scalar_assignment(target: &str, value: rumoca_core::Expression) -> rumoca_core::Statement { - assignment_with_span(target, value, test_span()) +fn scalar_function_param(name: &str) -> rumoca_core::FunctionParam { + rumoca_core::FunctionParam { + def_id: None, + name: name.to_string(), + span: test_span(), + type_name: "Real".to_string(), + type_class: None, + dims: vec![], + shape_expr: Vec::new(), + default: None, + description: None, + } } -fn assignment_with_span( - target: &str, - value: rumoca_core::Expression, - span: rumoca_core::Span, -) -> rumoca_core::Statement { - rumoca_core::Statement::Assignment { - comp: rumoca_core::ComponentReference { - local: false, - span, - parts: vec![rumoca_core::ComponentRefPart { - ident: target.to_string(), - span, - subs: Vec::new(), - }], - def_id: None, - }, - value, +fn function_param_with_dims(name: &str, dims: &[i64]) -> rumoca_core::FunctionParam { + rumoca_core::FunctionParam { + dims: dims.to_vec(), + ..scalar_function_param(name) + } +} + +fn function_param_with_shape_expr( + name: &str, + dims: &[i64], + shape_expr: Vec, +) -> rumoca_core::FunctionParam { + rumoca_core::FunctionParam { + dims: dims.to_vec(), + shape_expr, + ..scalar_function_param(name) + } +} + +fn real_with_span(value: f64, span: rumoca_core::Span) -> rumoca_core::Expression { + rumoca_core::Expression::Literal { + value: Literal::Real(value), span, } } +fn function_param_with_type(name: &str, type_name: &str) -> rumoca_core::FunctionParam { + rumoca_core::FunctionParam { + type_name: type_name.to_string(), + ..scalar_function_param(name) + } +} + +fn record_function_param(name: &str, type_name: &str) -> rumoca_core::FunctionParam { + rumoca_core::FunctionParam { + type_class: Some(rumoca_core::ClassType::Record), + ..function_param_with_type(name, type_name) + } +} + #[test] -fn vector_assignment_rejects_scalar_value_with_span() { +fn vector_function_input_rejects_scalar_actual_with_span() { let dae_model = dae::Dae::default(); let structural_bindings = IndexMap::new(); let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); - let mut function = rumoca_core::Function::new("My.badAssign", test_span()); - function.locals.push(function_param_with_dims("x", &[2])); - let mut scope = FunctionProjectionScope::default(); - let mut projected = Vec::new(); + let mut function = rumoca_core::Function::new("My.needsVector", test_span()); + function.inputs.push(function_param_with_dims("u", &[2])); let span = rumoca_core::Span::from_offsets( rumoca_core::SourceId::from_source_name( - "phase_solve_lower_derivative_rhs_function_projection_tests_source_55.mo", + "phase_solve_lower_derivative_rhs_function_projection_tests_source_52.mo", ), - 2, - 9, + 3, + 8, ); - let statement = assignment_with_span("x", real_with_span(1.0, span), span); - let err = analysis - .apply_assignment(&function, &statement, &mut scope, &mut projected, 0, span) - .expect_err("scalar assignment to vector local must fail"); + let err = match analysis.bind_inputs(&function, &[real_with_span(1.0, span)], 0, span) { + Ok(_) => panic!("scalar actual must not be projected as vector input"), + Err(err) => err, + }; + + assert_eq!(err.source_span(), Some(span)); + assert_eq!( + err.reason(), + "function `My.needsVector` input `u` expects dimensions [2], got []" + ); +} + +#[test] +fn dynamic_vector_function_input_uses_actual_dimensions() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let mut function = rumoca_core::Function::new("My.needsDynamicVector", test_span()); + function.inputs.push(function_param_with_dims("u", &[0])); + + let scope = analysis + .bind_inputs( + &function, + &[array(vec![real(1.0), real(2.0), real(3.0)], false)], + 0, + test_span(), + ) + .expect("dynamic vector input projection should not fail") + .expect("dynamic vector input should bind"); + + assert_eq!(scope.dims.get("u"), Some(&vec![3])); + assert_eq!(scope.scalars.get("u").map(Vec::len), Some(3)); +} + +#[test] +fn dynamic_vector_function_input_accepts_singleton_scalar_actual() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let mut function = rumoca_core::Function::new("My.needsDynamicVector", test_span()); + function.inputs.push(function_param_with_dims("u", &[0])); + + let scope = analysis + .bind_inputs(&function, &[real(0.25)], 0, test_span()) + .expect("dynamic singleton vector input projection should not fail") + .expect("dynamic singleton vector input should bind"); + + assert_eq!(scope.dims.get("u"), Some(&vec![1])); + assert_eq!(scope.scalars.get("u").map(Vec::len), Some(1)); +} + +#[test] +fn compile_time_if_selection_uses_projected_size_dimensions() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let mut function = rumoca_core::Function::new("My.selectBySize", test_span()); + function.inputs.push(function_param_with_dims("u", &[0])); + let scope = analysis + .bind_inputs( + &function, + &[array(vec![real(1.0), real(2.0), real(3.0)], false)], + 0, + test_span(), + ) + .expect("dynamic vector input projection should not fail") + .expect("dynamic vector input should bind"); + let condition = binary( + rumoca_core::OpBinary::Eq, + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + args: vec![var_ref("u"), integer(1)], + span: test_span(), + }, + integer(3), + test_span(), + ); + let selected_branch = real(10.0); + let else_branch = real(20.0); + let branches = vec![(condition, selected_branch.clone())]; + + let selected = analysis + .compile_time_if_selection(&branches, &else_branch, &scope) + .expect("compile-time if selection should not fail") + .expect("size(u, 1) should select a branch"); + + assert_eq!(selected, &selected_branch); +} + +#[test] +fn statement_if_projection_skips_size_guarded_zero_length_index_assignment() +-> Result<(), LowerError> { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let mut function = rumoca_core::Function::new("My.sizeGuardedAssignment", test_span()); + function.locals.push(function_param_with_dims("cr", &[0])); + function.locals.push(function_param_with_dims("den1", &[0])); + let mut scope = FunctionProjectionScope::default(); + scope.dims.insert("cr".to_string(), vec![0]); + scope.scalars.insert("cr".to_string(), Vec::new()); + scope.dims.insert("den1".to_string(), vec![0]); + scope.scalars.insert("den1".to_string(), Vec::new()); + let condition = binary( + rumoca_core::OpBinary::Eq, + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + args: vec![var_ref("cr"), integer(1)], + span: test_span(), + }, + integer(1), + test_span(), + ); + let statement = rumoca_core::Statement::If { + cond_blocks: vec![rumoca_core::StatementBlock { + cond: condition, + stmts: vec![indexed_assignment_with_span( + "den1", + &[1], + real(42.0), + test_span(), + )], + }], + else_block: None, + span: test_span(), + }; + let mut projected = Vec::new(); + + analysis.apply_statement( + &function, + &statement, + &mut scope, + &mut projected, + 0, + test_span(), + )?; + + assert_eq!(scope.scalars.get("den1").map(Vec::len), Some(0)); + assert!(projected.is_empty()); + Ok(()) +} + +#[test] +fn statement_if_projection_selects_function_input_literal_branch() -> Result<(), LowerError> { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let mut function = rumoca_core::Function::new("My.orderSelectedAssignment", test_span()); + function.inputs.push(scalar_function_param("order")); + function.locals.push(function_param_with_dims("c1", &[0])); + let mut scope = FunctionProjectionScope::default(); + scope.full.insert("order".to_string(), integer(2)); + scope.dims.insert("c1".to_string(), vec![0]); + scope.scalars.insert("c1".to_string(), Vec::new()); + let statement = rumoca_core::Statement::If { + cond_blocks: vec![ + rumoca_core::StatementBlock { + cond: binary( + rumoca_core::OpBinary::Eq, + var_ref("order"), + integer(1), + test_span(), + ), + stmts: vec![indexed_assignment_with_span( + "c1", + &[1], + real(1.0), + test_span(), + )], + }, + rumoca_core::StatementBlock { + cond: binary( + rumoca_core::OpBinary::Eq, + var_ref("order"), + integer(2), + test_span(), + ), + stmts: vec![], + }, + ], + else_block: None, + span: test_span(), + }; + let mut projected = Vec::new(); + + analysis.apply_statement( + &function, + &statement, + &mut scope, + &mut projected, + 0, + test_span(), + )?; + + assert_eq!(scope.scalars.get("c1").map(Vec::len), Some(0)); + assert!(projected.is_empty()); + Ok(()) +} + +#[test] +fn for_projection_skips_scope_bound_empty_range() -> Result<(), LowerError> { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let mut function = rumoca_core::Function::new("My.emptyRangeLoop", test_span()); + function.inputs.push(scalar_function_param("order")); + function.locals.push(function_param_with_dims("den1", &[0])); + let mut scope = FunctionProjectionScope::default(); + scope.full.insert("order".to_string(), integer(0)); + scope.dims.insert("den1".to_string(), vec![0]); + scope.scalars.insert("den1".to_string(), Vec::new()); + let statement = rumoca_core::Statement::For { + indices: vec![rumoca_core::ForIndex { + ident: "i".to_string(), + range: rumoca_core::Expression::Range { + start: Box::new(integer(1)), + step: None, + end: Box::new(var_ref("order")), + span: test_span(), + }, + }], + equations: vec![indexed_assignment_with_span( + "den1", + &[1], + real(1.0), + test_span(), + )], + span: test_span(), + }; + let mut projected = Vec::new(); + + analysis.apply_statement( + &function, + &statement, + &mut scope, + &mut projected, + 0, + test_span(), + )?; + + assert_eq!(scope.scalars.get("den1").map(Vec::len), Some(0)); + assert!(projected.is_empty()); + Ok(()) +} + +#[test] +fn declared_local_array_shape_uses_scope_bound_function_input() -> Result<(), LowerError> { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let mut function = rumoca_core::Function::new("My.scopeSizedLocalArray", test_span()); + function.inputs.push(scalar_function_param("order")); + function.locals.push(function_param_with_shape_expr( + "den1", + &[0], + vec![rumoca_core::Subscript::Expr { + expr: Box::new(var_ref("order")), + span: test_span(), + }], + )); + let mut scope = FunctionProjectionScope::default(); + scope.full.insert("order".to_string(), integer(3)); + + analysis.initialize_projected_declared_arrays(&function, &mut scope, 0, test_span())?; + + assert_eq!(scope.dims.get("den1"), Some(&vec![3])); + assert_eq!(scope.scalars.get("den1").map(Vec::len), Some(3)); + + let statement = indexed_assignment_with_span("den1", &[3], real(7.0), test_span()); + let mut projected = Vec::new(); + analysis.apply_statement( + &function, + &statement, + &mut scope, + &mut projected, + 0, + test_span(), + )?; + + assert_eq!( + scope.scalars.get("den1").and_then(|values| values.get(2)), + Some(&real(7.0)) + ); + Ok(()) +} + +#[test] +fn compile_time_size_uses_stream_variable_dimensions() { + let mut dae_model = dae::Dae::default(); + dae_model.variables.algebraics.insert( + rumoca_core::VarName::new("Xi"), + dae::Variable { + name: rumoca_core::VarName::new("Xi"), + dims: vec![1], + ..rumoca_ir_dae::Variable::empty_with_span(test_span()) + }, + ); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let expr = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + args: vec![ + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("inStream").into(), + args: vec![var_ref("Xi")], + is_constructor: false, + span: test_span(), + }, + integer(1), + ], + span: test_span(), + }; + + assert_eq!( + analysis + .compile_time_scalar_in_scope(&expr, &FunctionProjectionScope::default()) + .expect("stream size should not fail"), + Some(1.0) + ); +} + +#[test] +fn compile_time_scalar_evaluates_modelica_math_asinh() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("Modelica.Math.asinh").into(), + args: vec![real(1.25)], + is_constructor: false, + span: test_span(), + }; + + let actual = analysis + .compile_time_scalar_in_scope(&expr, &FunctionProjectionScope::default()) + .expect("Modelica.Math.asinh should not fail") + .expect("Modelica.Math.asinh should fold to a scalar"); + + assert!((actual - 1.25_f64.asinh()).abs() < f64::EPSILON); +} + +#[test] +fn array_like_projection_expands_stream_size_selected_cat() -> Result<(), LowerError> { + let mut dae_model = dae::Dae::default(); + dae_model.variables.algebraics.insert( + rumoca_core::VarName::new("Xi"), + dae::Variable { + name: rumoca_core::VarName::new("Xi"), + dims: vec![1], + ..rumoca_ir_dae::Variable::empty_with_span(test_span()) + }, + ); + let structural_bindings = IndexMap::new(); + let xi_stream = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("inStream").into(), + args: vec![var_ref("Xi")], + is_constructor: false, + span: test_span(), + }; + let cat = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Cat, + args: vec![ + integer(1), + xi_stream.clone(), + array( + vec![binary( + rumoca_core::OpBinary::Sub, + integer(1), + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Sum, + args: vec![xi_stream.clone()], + span: test_span(), + }, + test_span(), + )], + false, + ), + ], + span: test_span(), + }; + let condition = binary( + rumoca_core::OpBinary::Eq, + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + args: vec![xi_stream.clone(), integer(1)], + span: test_span(), + }, + integer(2), + test_span(), + ); + let branches = vec![(condition, xi_stream.clone())]; + let expr = rumoca_core::Expression::If { + branches: branches.clone(), + else_branch: Box::new(cat.clone()), + span: test_span(), + }; + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + assert_eq!( + analysis.compile_time_if_selection(&branches, &cat, &FunctionProjectionScope::default())?, + Some(&cat) + ); + assert!( + analysis + .project_value_scalars( + &cat, + &[2], + &FunctionProjectionScope::default(), + 0, + test_span() + )? + .is_some() + ); + assert_eq!( + analysis.expr_dims(&expr, &FunctionProjectionScope::default(), 0, test_span())?, + Some(vec![2]) + ); + + let values = project_array_like_scalars_with_owner( + &expr, + &dae_model, + &structural_bindings, + test_span(), + )? + .expect("stream-size-selected cat should expand"); + + assert_eq!(values.len(), 2); + assert_var_ref_name(&values[0], "Xi[1]"); + assert!(matches!( + &values[1], + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + .. + } + )); + Ok(()) +} + +#[test] +fn scalar_real_function_input_vectorizes_array_actual_with_span() -> Result<(), LowerError> { + let mut function = rumoca_core::Function::new("My.needsScalar", test_span()); + function.inputs.push(scalar_function_param("u")); + function.outputs.push(scalar_function_param("y")); + function.body.push(scalar_assignment( + "y", + binary( + rumoca_core::OpBinary::Mul, + local_var("u"), + real(2.0), + test_span(), + ), + )); + let mut dae_model = dae::Dae::default(); + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + let structural_bindings = IndexMap::new(); + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name( + "phase_solve_lower_derivative_rhs_function_projection_tests_source_54.mo", + ), + 7, + 13, + ); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.needsScalar").into(), + args: vec![rumoca_core::Expression::Array { + elements: vec![real(1.0), real(2.0)], + is_matrix: false, + span, + }], + is_constructor: false, + span, + }; + + let values = + function_call_projected_scalars_with_owner(&call, &dae_model, &structural_bindings, span)? + .expect("vectorized scalar function call should project"); + + assert_eq!(values.len(), 2); + assert!(values.iter().all(|value| matches!( + value, + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs, + rhs, + .. + } if matches!(lhs.as_ref(), rumoca_core::Expression::Literal { .. }) + && matches!(rhs.as_ref(), rumoca_core::Expression::Literal { .. }) + ))); + Ok(()) +} + +#[test] +fn vectorized_scalar_function_division_projects_by_lane() -> Result<(), LowerError> { + let mut function = rumoca_core::Function::new("My.scalarDiv", test_span()); + function.inputs.push(scalar_function_param("u")); + function.inputs.push(scalar_function_param("v")); + function.outputs.push(scalar_function_param("y")); + function.body.push(scalar_assignment( + "y", + binary( + rumoca_core::OpBinary::Div, + local_var("u"), + local_var("v"), + test_span(), + ), + )); + let mut dae_model = dae::Dae::default(); + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + let structural_bindings = IndexMap::new(); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.scalarDiv").into(), + args: vec![ + array(vec![real(2.0), real(4.0), real(6.0)], false), + array(vec![real(1.0), real(2.0), real(3.0)], false), + ], + is_constructor: false, + span: test_span(), + }; + + let values = function_call_projected_scalars_with_owner( + &call, + &dae_model, + &structural_bindings, + test_span(), + )? + .expect("vectorized scalar division should project"); + + assert_eq!(values.len(), 3); + assert!(values.iter().all(|value| matches!( + value, + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Div, + lhs, + rhs, + .. + } if matches!(lhs.as_ref(), rumoca_core::Expression::Literal { .. }) + && matches!(rhs.as_ref(), rumoca_core::Expression::Literal { .. }) + ))); + Ok(()) +} + +#[test] +fn vectorized_scalar_local_default_projects_by_lane() -> Result<(), LowerError> { + let mut function = rumoca_core::Function::new("My.scalarDefault", test_span()); + function.inputs.push(scalar_function_param("roughness")); + function.inputs.push(scalar_function_param("diameter")); + function + .locals + .push(scalar_function_param("Delta").with_default(binary( + rumoca_core::OpBinary::Div, + local_var("roughness"), + local_var("diameter"), + test_span(), + ))); + function.outputs.push(scalar_function_param("y")); + function + .body + .push(scalar_assignment("y", local_var("Delta"))); + let mut dae_model = dae::Dae::default(); + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + let structural_bindings = IndexMap::new(); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.scalarDefault").into(), + args: vec![ + array(vec![real(0.2), real(0.4), real(0.6)], false), + array(vec![real(1.0), real(2.0), real(3.0)], false), + ], + is_constructor: false, + span: test_span(), + }; + + let values = function_call_projected_scalars_with_owner( + &call, + &dae_model, + &structural_bindings, + test_span(), + )? + .expect("vectorized scalar local default should project"); + + assert_eq!(values.len(), 3); + assert!(values.iter().all(|value| matches!( + value, + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Div, + lhs, + rhs, + .. + } if matches!(lhs.as_ref(), rumoca_core::Expression::Literal { .. }) + && matches!(rhs.as_ref(), rumoca_core::Expression::Literal { .. }) + ))); + Ok(()) +} + +#[test] +fn vectorized_scalar_local_default_projects_through_nested_call() -> Result<(), LowerError> { + let mut inner = rumoca_core::Function::new("My.inner", test_span()); + inner.inputs.push(scalar_function_param("u")); + inner.inputs.push(scalar_function_param("v")); + inner.outputs.push(scalar_function_param("y")); + inner.body.push(scalar_assignment( + "y", + binary( + rumoca_core::OpBinary::Add, + local_var("u"), + local_var("v"), + test_span(), + ), + )); + + let mut outer = rumoca_core::Function::new("My.outer", test_span()); + outer.inputs.push(scalar_function_param("roughness")); + outer.inputs.push(scalar_function_param("diameter")); + outer.inputs.push(scalar_function_param("offset")); + outer + .locals + .push(scalar_function_param("Delta").with_default(binary( + rumoca_core::OpBinary::Div, + local_var("roughness"), + local_var("diameter"), + test_span(), + ))); + outer.outputs.push(scalar_function_param("y")); + outer.body.push(scalar_assignment( + "y", + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.inner").into(), + args: vec![local_var("Delta"), local_var("offset")], + is_constructor: false, + span: test_span(), + }, + )); + + let mut dae_model = dae::Dae::default(); + dae_model + .symbols + .functions + .insert(inner.name.clone(), inner); + dae_model + .symbols + .functions + .insert(outer.name.clone(), outer); + let structural_bindings = IndexMap::new(); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.outer").into(), + args: vec![ + array(vec![real(0.2), real(0.4), real(0.6)], false), + array(vec![real(1.0), real(2.0), real(3.0)], false), + array(vec![real(10.0), real(20.0), real(30.0)], false), + ], + is_constructor: false, + span: test_span(), + }; + + let values = function_call_projected_scalars_with_owner( + &call, + &dae_model, + &structural_bindings, + test_span(), + )? + .expect("vectorized scalar local default should project through nested call"); + + assert_eq!(values.len(), 3); + assert!(values.iter().all(|value| matches!( + value, + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs, + rhs, + .. + } if matches!( + lhs.as_ref(), + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Div, + .. + } + ) && matches!(rhs.as_ref(), rumoca_core::Expression::Literal { .. }) + ))); + Ok(()) +} + +#[test] +fn vectorized_scalar_default_projects_elementwise_builtin_and_if_condition() +-> Result<(), LowerError> { + let mut function = rumoca_core::Function::new("My.reynoldsStart", test_span()); + function.inputs.push(scalar_function_param("roughness")); + function.inputs.push(scalar_function_param("diameter")); + function.inputs.push(scalar_function_param("re_turbulent")); + function + .locals + .push(scalar_function_param("Delta").with_default(binary( + rumoca_core::OpBinary::Div, + local_var("roughness"), + local_var("diameter"), + test_span(), + ))); + function + .locals + .push(scalar_function_param("Re1").with_default(builtin( + rumoca_core::BuiltinFunction::Min, + vec![ + rumoca_core::Expression::If { + branches: vec![( + binary( + rumoca_core::OpBinary::Le, + local_var("Delta"), + real(0.0065), + test_span(), + ), + real(1.0), + )], + else_branch: Box::new(binary( + rumoca_core::OpBinary::Div, + real(0.0065), + local_var("Delta"), + test_span(), + )), + span: test_span(), + }, + local_var("re_turbulent"), + ], + ))); + function.outputs.push(scalar_function_param("y")); + function.body.push(scalar_assignment("y", local_var("Re1"))); + let mut dae_model = dae::Dae::default(); + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + let structural_bindings = IndexMap::new(); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.reynoldsStart").into(), + args: vec![ + array(vec![real(0.2), real(0.4), real(0.6)], false), + array(vec![real(1.0), real(2.0), real(3.0)], false), + array(vec![real(4000.0), real(4000.0), real(4000.0)], false), + ], + is_constructor: false, + span: test_span(), + }; + + let values = function_call_projected_scalars_with_owner( + &call, + &dae_model, + &structural_bindings, + test_span(), + )? + .expect("vectorized scalar default with builtin min should project"); + + assert_eq!(values.len(), 3); + assert!(values.iter().all(|value| matches!( + value, + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Min, + args, + .. + } if args.len() == 2 + && matches!(args[0], rumoca_core::Expression::If { .. }) + && matches!(args[1], rumoca_core::Expression::Literal { .. }) + ))); + Ok(()) +} + +#[test] +fn vectorized_scalar_local_default_projects_through_exponent_output() -> Result<(), LowerError> { + let mut function = rumoca_core::Function::new("My.cubicLike", test_span()); + function.inputs.push(scalar_function_param("x")); + function.inputs.push(scalar_function_param("x1")); + function + .locals + .push(scalar_function_param("dx").with_default(binary( + rumoca_core::OpBinary::Div, + local_var("x"), + local_var("x1"), + test_span(), + ))); + function.outputs.push(scalar_function_param("y")); + function.body.push(scalar_assignment( + "y", + binary( + rumoca_core::OpBinary::Mul, + local_var("x1"), + binary( + rumoca_core::OpBinary::Exp, + binary( + rumoca_core::OpBinary::Div, + local_var("x"), + local_var("x1"), + test_span(), + ), + binary( + rumoca_core::OpBinary::Add, + real(1.0), + local_var("dx"), + test_span(), + ), + test_span(), + ), + test_span(), + ), + )); + let mut dae_model = dae::Dae::default(); + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + let structural_bindings = IndexMap::new(); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.cubicLike").into(), + args: vec![ + array(vec![real(2.0), real(4.0), real(6.0)], false), + array(vec![real(1.0), real(2.0), real(3.0)], false), + ], + is_constructor: false, + span: test_span(), + }; + + let values = function_call_projected_scalars_with_owner( + &call, + &dae_model, + &structural_bindings, + test_span(), + )? + .expect("vectorized scalar local default should project through exponent output"); + + assert_eq!(values.len(), 3); + assert!( + values + .iter() + .all(|value| !format!("{value:?}").contains("dx")) + ); + Ok(()) +} + +#[test] +fn vectorized_scalar_function_default_projects_through_exponent_output() -> Result<(), LowerError> { + let mut function = rumoca_core::Function::new("My.cubicFunctionDefault", test_span()); + function.inputs.push(scalar_function_param("x")); + function.inputs.push(scalar_function_param("x1")); + function + .locals + .push( + scalar_function_param("dx").with_default(rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.externalLog10").into(), + args: vec![binary( + rumoca_core::OpBinary::Div, + local_var("x"), + local_var("x1"), + test_span(), + )], + is_constructor: false, + span: test_span(), + }), + ); + function.outputs.push(scalar_function_param("y")); + function.body.push(scalar_assignment( + "y", + binary( + rumoca_core::OpBinary::Exp, + binary( + rumoca_core::OpBinary::Div, + local_var("x"), + local_var("x1"), + test_span(), + ), + local_var("dx"), + test_span(), + ), + )); + let mut dae_model = dae::Dae::default(); + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + let structural_bindings = IndexMap::new(); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.cubicFunctionDefault").into(), + args: vec![ + array(vec![real(2.0), real(4.0), real(6.0)], false), + array(vec![real(1.0), real(2.0), real(3.0)], false), + ], + is_constructor: false, + span: test_span(), + }; + + let values = function_call_projected_scalars_with_owner( + &call, + &dae_model, + &structural_bindings, + test_span(), + )? + .expect("vectorized scalar function default should project through exponent output"); + + assert_eq!(values.len(), 3); + assert!( + values + .iter() + .all(|value| !format!("{value:?}").contains("dx")) + ); + Ok(()) +} + +#[test] +fn vectorized_scalar_output_default_projects_by_lane() -> Result<(), LowerError> { + let mut function = rumoca_core::Function::new("My.outputDefault", test_span()); + function.inputs.push(scalar_function_param("roughness")); + function.inputs.push(scalar_function_param("diameter")); + function + .outputs + .push(scalar_function_param("y").with_default(binary( + rumoca_core::OpBinary::Div, + local_var("roughness"), + local_var("diameter"), + test_span(), + ))); + let mut dae_model = dae::Dae::default(); + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + let structural_bindings = IndexMap::new(); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.outputDefault").into(), + args: vec![ + array(vec![real(0.2), real(0.4), real(0.6)], false), + array(vec![real(1.0), real(2.0), real(3.0)], false), + ], + is_constructor: false, + span: test_span(), + }; + + let values = function_call_projected_scalars_with_owner( + &call, + &dae_model, + &structural_bindings, + test_span(), + )? + .expect("vectorized scalar output default should project"); + + assert_eq!(values.len(), 3); + Ok(()) +} + +#[test] +fn scalar_real_function_input_accepts_singleton_vector_actual() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let mut function = rumoca_core::Function::new("My.needsScalar", test_span()); + function.inputs.push(scalar_function_param("u")); + + let scope = analysis + .bind_inputs(&function, &[array(vec![real(1.0)], false)], 0, test_span()) + .expect("single scalar vector actual should bind") + .expect("single scalar vector actual should project"); + + assert_eq!(scope.dims.get("u"), Some(&vec![1])); + assert_eq!(scope.scalars.get("u").map(Vec::len), Some(1)); +} + +#[test] +fn generated_function_call_projection_errors_use_owner_span() { + let owner_span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name( + "phase_solve_lower_derivative_rhs_function_projection_tests_source_53.mo", + ), + 3, + 8, + ); + let mut dae_model = dae::Dae::default(); + let mut function = rumoca_core::Function::new("My.needsVector", test_span()); + function.inputs.push(rumoca_core::FunctionParam { + span: owner_span, + ..function_param_with_dims("u", &[2]) + }); + function.outputs.push(scalar_function_param("y")); + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + let structural_bindings = IndexMap::new(); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.needsVector").into(), + args: vec![rumoca_core::Expression::Literal { + value: Literal::Real(1.0), + span: rumoca_core::Span::DUMMY, + }], + is_constructor: false, + span: rumoca_core::Span::DUMMY, + }; + + let err = function_call_projected_scalars_with_owner( + &call, + &dae_model, + &structural_bindings, + owner_span, + ) + .expect_err("generated invalid projection must report an error"); + + assert_eq!(err.source_span(), Some(owner_span)); + assert_eq!( + err.reason(), + "function `My.needsVector` input `u` expects dimensions [2], got []" + ); +} + +#[test] +#[allow(clippy::too_many_lines)] +fn projected_record_field_expands_dynamic_cat_output() -> Result<(), LowerError> { + let mut dae_model = dae::Dae::default(); + let mut state_ctor = rumoca_core::Function::new("My.State", test_span()); + state_ctor.is_constructor = true; + state_ctor.pure = true; + state_ctor.inputs.push(function_param_with_dims("X", &[0])); + state_ctor + .outputs + .push(record_function_param("state", "My.State")); + dae_model + .symbols + .functions + .insert(state_ctor.name.clone(), state_ctor); + + let mut set_state = rumoca_core::Function::new("My.setState", test_span()); + set_state.pure = true; + set_state.inputs.push(function_param_with_dims("X", &[0])); + set_state + .outputs + .push(record_function_param("state", "My.State")); + let x_value = rumoca_core::Expression::If { + branches: vec![( + binary( + rumoca_core::OpBinary::Eq, + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + args: vec![var_ref("X"), integer(1)], + span: test_span(), + }, + integer(2), + test_span(), + ), + var_ref("X"), + )], + else_branch: Box::new(rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Cat, + args: vec![ + integer(1), + var_ref("X"), + array( + vec![binary( + rumoca_core::OpBinary::Sub, + integer(1), + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Sum, + args: vec![var_ref("X")], + span: test_span(), + }, + test_span(), + )], + false, + ), + ], + span: test_span(), + }), + span: test_span(), + }; + set_state.body.push(scalar_assignment( + "state", + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.State").into(), + args: vec![rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("__rumoca_named_arg__.X").into(), + args: vec![x_value], + is_constructor: true, + span: test_span(), + }], + is_constructor: true, + span: test_span(), + }, + )); + dae_model + .symbols + .functions + .insert(set_state.name.clone(), set_state); + let structural_bindings = IndexMap::new(); + let expr = rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.setState").into(), + args: vec![array(vec![real(0.25)], false)], + is_constructor: false, + span: test_span(), + }), + field: "X".to_string(), + span: test_span(), + }; + + let values = function_call_projected_scalars_with_owner( + &expr, + &dae_model, + &structural_bindings, + test_span(), + )? + .expect("projected record field should expand"); + + assert_eq!(values.len(), 2); + assert!(matches!( + &values[0], + rumoca_core::Expression::Literal { + value: Literal::Real(value), + .. + } if (*value - 0.25).abs() < 1e-12 + )); + assert!(matches!( + &values[1], + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + .. + } + )); + Ok(()) +} + +#[test] +#[allow(clippy::too_many_lines)] +fn projected_record_field_expands_stream_variable_cat_output() -> Result<(), LowerError> { + let mut dae_model = dae::Dae::default(); + dae_model.variables.algebraics.insert( + rumoca_core::VarName::new("Xi"), + dae::Variable { + name: rumoca_core::VarName::new("Xi"), + dims: vec![1], + ..rumoca_ir_dae::Variable::empty_with_span(test_span()) + }, + ); + let mut state_ctor = rumoca_core::Function::new("My.State", test_span()); + state_ctor.is_constructor = true; + state_ctor.pure = true; + state_ctor.inputs.push(function_param_with_dims("X", &[0])); + state_ctor + .outputs + .push(record_function_param("state", "My.State")); + dae_model + .symbols + .functions + .insert(state_ctor.name.clone(), state_ctor); + + let mut set_state = rumoca_core::Function::new("My.setState", test_span()); + set_state.pure = true; + set_state.inputs.push(function_param_with_dims("X", &[0])); + set_state + .outputs + .push(record_function_param("state", "My.State")); + let x_value = rumoca_core::Expression::If { + branches: vec![( + binary( + rumoca_core::OpBinary::Eq, + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + args: vec![var_ref("X"), integer(1)], + span: test_span(), + }, + integer(2), + test_span(), + ), + var_ref("X"), + )], + else_branch: Box::new(rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Cat, + args: vec![ + integer(1), + var_ref("X"), + array( + vec![binary( + rumoca_core::OpBinary::Sub, + integer(1), + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Sum, + args: vec![var_ref("X")], + span: test_span(), + }, + test_span(), + )], + false, + ), + ], + span: test_span(), + }), + span: test_span(), + }; + set_state.body.push(scalar_assignment( + "state", + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.State").into(), + args: vec![rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("__rumoca_named_arg__.X").into(), + args: vec![x_value], + is_constructor: true, + span: test_span(), + }], + is_constructor: true, + span: test_span(), + }, + )); + dae_model + .symbols + .functions + .insert(set_state.name.clone(), set_state); + let structural_bindings = IndexMap::new(); + let xi_stream = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("inStream").into(), + args: vec![var_ref("Xi")], + is_constructor: false, + span: test_span(), + }; + let expr = rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.setState").into(), + args: vec![xi_stream.clone()], + is_constructor: false, + span: test_span(), + }), + field: "X".to_string(), + span: test_span(), + }; + + let values = function_call_projected_scalars_with_owner( + &expr, + &dae_model, + &structural_bindings, + test_span(), + )? + .expect("projected stream record field should expand"); + + assert_eq!(values.len(), 2); + assert_var_ref_name(&values[0], "Xi[1]"); + assert!(matches!( + &values[1], + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + .. + } + )); + Ok(()) +} + +#[test] +fn record_like_function_input_accepts_structured_actual_dimensions() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let mut function = rumoca_core::Function::new("My.needsRecord", test_span()); + function + .inputs + .push(function_param_with_type("q", "Pkg.Quaternion")); + + let scope = analysis + .bind_inputs( + &function, + &[array( + vec![real(1.0), real(2.0), real(3.0), real(4.0)], + false, + )], + 0, + test_span(), + ) + .expect("record-like input binding should not fail") + .expect("record-like input should bind"); + + assert_eq!(scope.dims.get("q"), Some(&vec![4])); + assert_eq!(scope.scalars.get("q").map(Vec::len), Some(4)); +} + +#[test] +fn function_projection_initializes_array_local_from_declaration_binding() -> Result<(), LowerError> +{ + let mut dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let mut function = rumoca_core::Function::new("My.localDefault", test_span()); + function.outputs.push(function_param_with_dims("y", &[3])); + let mut local = function_param_with_dims("x", &[3]); + local.default = Some(array(vec![real(1.0), real(2.0), real(3.0)], false)); + function.locals.push(local); + function.body.push(scalar_assignment("y", local_var("x"))); + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.localDefault").into(), + args: Vec::new(), + is_constructor: false, + span: test_span(), + }; + + let projected = function_call_projected_scalars_with_owner( + &call, + &dae_model, + &structural_bindings, + test_span(), + )? + .expect("function call should project declaration-bound array local"); + let values = projected + .iter() + .map(|expr| match expr { + rumoca_core::Expression::Literal { + value: Literal::Real(value), + .. + } => *value, + other => panic!("expected real literal projection, got {other:?}"), + }) + .collect::>(); + + assert_eq!(values, vec![1.0, 2.0, 3.0]); + Ok(()) +} + +#[test] +fn vector_constructor_input_rejects_scalar_actual_with_span() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let scope = FunctionProjectionScope::default(); + let input = function_param_with_dims("u", &[2]); + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name( + "phase_solve_lower_derivative_rhs_function_projection_tests_source_53.mo", + ), + 5, + 11, + ); + let actual = real_with_span(1.0, span); + + let err = analysis + .optional_constructor_input_scalars(&actual, &input, &scope, 0, span) + .expect_err("scalar actual must not be projected as vector constructor input"); + + assert_eq!(err.source_span(), Some(span)); + assert_eq!( + err.reason(), + "record constructor input `u` expects dimensions [2], got []" + ); +} + +#[test] +fn scalar_constructor_input_uses_formal_scalar_dimensions() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let scope = FunctionProjectionScope::default(); + let input = scalar_function_param("u"); + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name( + "phase_solve_lower_derivative_rhs_function_projection_tests_source_57.mo", + ), + 5, + 11, + ); + let actual = rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("runtime_scalar"), + subscripts: Vec::new(), + span, + }; + + let (dims, scalars) = analysis + .optional_constructor_input_scalars(&actual, &input, &scope, 0, span) + .expect("scalar constructor input projection should not fail") + .expect("scalar primitive input should project as a scalar"); + + assert!(dims.is_empty()); + assert_eq!(scalars.len(), 1); +} + +#[test] +fn record_like_constructor_input_declines_unknown_actual_dimensions() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let scope = FunctionProjectionScope::default(); + let input = function_param_with_type("q", "Pkg.Quaternion"); + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name( + "phase_solve_lower_derivative_rhs_function_projection_tests_source_58.mo", + ), + 6, + 14, + ); + let actual = rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("runtime_record"), + subscripts: Vec::new(), + span, + }; + + let scalars = analysis + .optional_constructor_input_scalars(&actual, &input, &scope, 0, span) + .expect("unknown record-like constructor input dimensions should decline"); + + assert!(scalars.is_none()); +} + +#[test] +fn if_projection_rejects_scalar_values_without_dimensions() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let entry_scope = FunctionProjectionScope::default(); + let mut else_scope = FunctionProjectionScope::default(); + else_scope.scalars.insert("x".to_string(), vec![real(0.0)]); + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name( + "phase_solve_lower_derivative_rhs_function_projection_tests_source_59.mo", + ), + 7, + 18, + ); + + let err = match analysis.merged_if_scope(&entry_scope, &[], &[], &else_scope, span) { + Ok(_) => panic!("merged scalar projection without dimensions must fail"), + Err(err) => err, + }; + + assert_eq!(err.source_span(), Some(span)); + assert_eq!( + err.reason(), + "invalid IR contract: if-statement projection for `x` has scalar values but no dimensions" + ); +} + +#[test] +fn if_projection_rejects_conflicting_branch_dimensions() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let mut entry_scope = FunctionProjectionScope::default(); + entry_scope + .scalars + .insert("x".to_string(), vec![real(0.0), real(0.0)]); + entry_scope.dims.insert("x".to_string(), vec![2]); + let mut branch_scope = entry_scope.clone(); + branch_scope.dims.insert("x".to_string(), vec![1, 2]); + let condition = real(1.0); + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name( + "phase_solve_lower_derivative_rhs_function_projection_tests_source_60.mo", + ), + 8, + 21, + ); + + let err = match analysis.merged_if_scope( + &entry_scope, + &[condition], + &[branch_scope], + &entry_scope, + span, + ) { + Ok(_) => panic!("merged scalar projection with conflicting dimensions must fail"), + Err(err) => err, + }; + + assert_eq!(err.source_span(), Some(span)); + assert_eq!( + err.reason(), + "invalid IR contract: if-statement projection for `x` has mismatched dimensions: [2] and [1, 2]" + ); +} + +fn scalar_assignment(target: &str, value: rumoca_core::Expression) -> rumoca_core::Statement { + assignment_with_span(target, value, test_span()) +} + +fn assignment_with_span( + target: &str, + value: rumoca_core::Expression, + span: rumoca_core::Span, +) -> rumoca_core::Statement { + rumoca_core::Statement::Assignment { + comp: rumoca_core::ComponentReference { + local: false, + span, + parts: vec![rumoca_core::ComponentRefPart { + ident: target.to_string(), + span, + subs: Vec::new(), + }], + def_id: None, + }, + value, + span, + } +} + +fn indexed_assignment_with_span( + target: &str, + indices: &[i64], + value: rumoca_core::Expression, + span: rumoca_core::Span, +) -> rumoca_core::Statement { + rumoca_core::Statement::Assignment { + comp: rumoca_core::ComponentReference { + local: false, + span, + parts: vec![rumoca_core::ComponentRefPart { + ident: target.to_string(), + span, + subs: indices + .iter() + .map(|value| rumoca_core::Subscript::Index { + value: *value, + span, + }) + .collect(), + }], + def_id: None, + }, + value, + span, + } +} + +fn component_ref_target(target: &str) -> rumoca_core::ComponentReference { + rumoca_core::ComponentReference { + local: false, + span: test_span(), + parts: vec![rumoca_core::ComponentRefPart { + ident: target.to_string(), + span: test_span(), + subs: Vec::new(), + }], + def_id: None, + } +} + +#[test] +fn vector_assignment_rejects_scalar_value_with_span() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let mut function = rumoca_core::Function::new("My.badAssign", test_span()); + function.locals.push(function_param_with_dims("x", &[2])); + let mut scope = FunctionProjectionScope::default(); + let mut projected = Vec::new(); + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name( + "phase_solve_lower_derivative_rhs_function_projection_tests_source_55.mo", + ), + 2, + 9, + ); + let statement = assignment_with_span("x", real_with_span(1.0, span), span); + + let err = analysis + .apply_assignment(&function, &statement, &mut scope, &mut projected, 0, span) + .expect_err("scalar assignment to vector local must fail"); + + assert_eq!(err.source_span(), Some(span)); + assert_eq!( + err.reason(), + "function `My.badAssign` assignment to `x` expects dimensions [2], got []" + ); +} + +#[test] +fn function_call_statement_assigns_scalarized_array_output_to_target() -> Result<(), LowerError> { + let mut callee = rumoca_core::Function::new("My.makeVector", test_span()); + callee.inputs.push(scalar_function_param("n")); + callee.outputs.push(function_param_with_shape_expr( + "y", + &[0], + vec![rumoca_core::Subscript::Expr { + expr: Box::new(var_ref("n")), + span: test_span(), + }], + )); + callee.outputs.push(function_param_with_dims("empty", &[0])); + callee.body.push(indexed_assignment_with_span( + "y", + &[1], + real(1.0), + test_span(), + )); + callee.body.push(indexed_assignment_with_span( + "y", + &[2], + real(2.0), + test_span(), + )); + callee.body.push(indexed_assignment_with_span( + "y", + &[3], + real(3.0), + test_span(), + )); + + let mut caller = rumoca_core::Function::new("My.caller", test_span()); + caller.inputs.push(scalar_function_param("n")); + caller.locals.push(function_param_with_shape_expr( + "x", + &[0], + vec![rumoca_core::Subscript::Expr { + expr: Box::new(var_ref("n")), + span: test_span(), + }], + )); + caller.locals.push(function_param_with_dims("empty", &[0])); + caller.outputs.push(function_param_with_shape_expr( + "out", + &[0], + vec![rumoca_core::Subscript::Expr { + expr: Box::new(var_ref("n")), + span: test_span(), + }], + )); + caller.body.push(rumoca_core::Statement::FunctionCall { + comp: component_ref_target("My.makeVector"), + args: vec![var_ref("n")], + outputs: vec![component_ref_target("x"), component_ref_target("empty")], + span: test_span(), + }); + caller.body.push(scalar_assignment("out", var_ref("x"))); + + let mut dae_model = dae::Dae::default(); + dae_model + .symbols + .functions + .insert(callee.name.clone(), callee); + dae_model + .symbols + .functions + .insert(caller.name.clone(), caller); + let structural_bindings = IndexMap::new(); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.caller").into(), + args: vec![integer(3)], + is_constructor: false, + span: test_span(), + }; + + let values = function_call_projected_scalars_with_owner( + &call, + &dae_model, + &structural_bindings, + test_span(), + )? + .expect("caller output should project"); + + assert_eq!(values, vec![real(1.0), real(2.0), real(3.0)]); + Ok(()) +} + +#[test] +fn scalar_assignment_rejects_vector_value_with_span() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let mut function = rumoca_core::Function::new("My.badAssign", test_span()); + function.locals.push(scalar_function_param("x")); + let mut scope = FunctionProjectionScope::default(); + let mut projected = Vec::new(); + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name( + "phase_solve_lower_derivative_rhs_function_projection_tests_source_56.mo", + ), + 4, + 12, + ); + let statement = assignment_with_span( + "x", + rumoca_core::Expression::Array { + elements: vec![real(1.0), real(2.0)], + is_matrix: false, + span, + }, + span, + ); + + let err = analysis + .apply_assignment(&function, &statement, &mut scope, &mut projected, 0, span) + .expect_err("vector assignment to scalar local must fail"); + + assert_eq!(err.source_span(), Some(span)); + assert_eq!( + err.reason(), + "function `My.badAssign` assignment to `x` expects dimensions [], got [2]" + ); +} + +#[test] +fn unassigned_projected_scalar_reports_projection_span() -> Result<(), String> { + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name( + "phase_solve_lower_derivative_rhs_function_projection_tests_source_57.mo", + ), + 6, + 15, + ); + let values = vec![rumoca_core::Expression::Empty { span }]; + + let Err(err) = assigned_projected_scalar_value("x", &[1], &values, 0, span) else { + return Err("unassigned projected scalar slot succeeded".to_string()); + }; + + assert_eq!(err.source_span(), Some(span)); + assert_eq!( + err.reason(), + "projected local component `x[1]` is unassigned" + ); + Ok(()) +} + +#[test] +fn scalar_selector_rejects_colon_with_subscript_span() -> Result<(), String> { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let scope = FunctionProjectionScope::default(); + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name( + "phase_solve_lower_derivative_rhs_function_projection_tests_source_58.mo", + ), + 2, + 3, + ); + let subscript = rumoca_core::Subscript::colon(span); + + let Err(err) = subscript_selector_expr(&subscript, &analysis, &scope, 0) else { + return Err("colon scalar selector succeeded".to_string()); + }; + + assert_eq!(err.source_span(), Some(span)); + assert_eq!( + err.reason(), + "colon subscript cannot select a scalar projected value" + ); + Ok(()) +} + +#[test] +fn guarded_assignment_without_base_reports_assignment_span() { + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name( + "phase_solve_lower_derivative_rhs_function_projection_tests_source_59.mo", + ), + 9, + 21, + ); + + let err = guarded_assignment_without_base("y", span); + + assert_eq!(err.source_span(), Some(span)); + assert_eq!( + err.reason(), + "if-statement assignment to `y` requires an existing binding or an else assignment" + ); +} + +fn local_var(name: &str) -> rumoca_core::Expression { + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new(name), + subscripts: Vec::new(), + span: test_span(), + } +} + +/// A pure function whose projected output doubles in size per statement, +/// crossing `MAX_FUNCTION_PROJECTION_NODES` long before it finishes. +fn over_budget_function() -> rumoca_core::Function { + let mut body = vec![scalar_assignment( + "y", + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs: Box::new(local_var("x")), + rhs: Box::new(local_var("x")), + span: test_span(), + }, + )]; + for _ in 0..16 { + body.push(scalar_assignment( + "y", + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs: Box::new(local_var("y")), + rhs: Box::new(local_var("y")), + span: test_span(), + }, + )); + } + rumoca_core::Function { + name: rumoca_core::VarName::new("My.explode"), + def_id: None, + inputs: vec![scalar_function_param("x")], + outputs: vec![scalar_function_param("y")], + locals: vec![], + body, + is_constructor: false, + pure: true, + external: None, + derivatives: vec![], + span: test_span(), + } +} + +fn expression_over_projection_budget(mut expr: rumoca_core::Expression) -> rumoca_core::Expression { + for _ in 0..12 { + expr = rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs: Box::new(expr.clone()), + rhs: Box::new(expr), + span: test_span(), + }; + } + expr +} + +fn over_budget_wrapper_call( + name: &str, + body: Vec, +) -> (dae::Dae, rumoca_core::Expression) { + let explode = over_budget_function(); + let mut wrapper = rumoca_core::Function::new(name, test_span()); + wrapper.inputs.push(scalar_function_param("x")); + wrapper.outputs.push(scalar_function_param("y")); + wrapper.body = body; + + let mut dae_model = dae::Dae::default(); + dae_model + .symbols + .functions + .insert(explode.name.clone(), explode); + dae_model + .symbols + .functions + .insert(wrapper.name.clone(), wrapper); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new(name).into(), + args: vec![real(2.0)], + is_constructor: false, + span: test_span(), + }; + (dae_model, call) +} + +#[test] +fn over_budget_projection_is_a_typed_error_and_declines_at_the_boundary() { + let mut dae_model = dae::Dae::default(); + dae_model.symbols.functions.insert( + rumoca_core::VarName::new("My.explode"), + over_budget_function(), + ); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.explode").into(), + args: vec![real(2.0)], + is_constructor: false, + span: test_span(), + }; + + let err = analysis + .function_call_outputs_with_owner(&call, 0, test_span()) + .expect_err("over-budget projection must surface as a typed error"); + assert!(err.is_projection_budget_exceeded(), "got: {err:?}"); + assert!(err.reason().contains("My.explode"), "got: {}", err.reason()); + + // The outermost boundary resolves the decline by keeping the runtime + // call; the memoized decline must answer follow-up probes identically. + for _ in 0..2 { + let outputs = analysis + .top_level_function_call_outputs(&call, test_span()) + .expect("budget decline must not fail the outer lowering"); + assert!(outputs.is_none()); + } +} + +fn assert_internal_budget_propagates_to_top_level( + dae_model: &dae::Dae, + call: &rumoca_core::Expression, +) { + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(dae_model, &structural_bindings); + let err = analysis + .function_call_outputs_with_owner(call, 0, test_span()) + .expect_err("nested budget exhaustion must remain a typed internal error"); + assert!(err.is_projection_budget_exceeded(), "got: {err:?}"); + + let outputs = analysis + .top_level_function_call_outputs(call, test_span()) + .expect("the outer boundary must preserve the complete runtime call"); + assert!(outputs.is_none()); +} + +#[test] +fn over_budget_function_call_statement_propagates_to_top_level() { + let body = vec![rumoca_core::Statement::FunctionCall { + comp: component_ref_target("My.explode"), + args: vec![local_var("x")], + outputs: vec![component_ref_target("y")], + span: test_span(), + }]; + let (dae_model, call) = over_budget_wrapper_call("My.statementWrapper", body); + + assert_internal_budget_propagates_to_top_level(&dae_model, &call); +} + +#[test] +fn over_budget_scalar_assignment_call_propagates_to_top_level() { + let body = vec![ + scalar_assignment( + "y", + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.explode").into(), + args: vec![local_var("x")], + is_constructor: false, + span: test_span(), + }, + ), + scalar_assignment("y", real(1.0)), + ]; + let (dae_model, call) = over_budget_wrapper_call("My.scalarWrapper", body); + + assert_internal_budget_propagates_to_top_level(&dae_model, &call); +} + +#[test] +fn lane_rewritten_over_budget_call_propagates_typed_error() { + let expanded = expression_over_projection_budget(real(2.0)); + let mut function = rumoca_core::Function::new("My.laneBudget", test_span()); + function.inputs.push(function_param_with_dims("x", &[0])); + function.outputs.push(scalar_function_param("y")); + function.locals.push(scalar_function_param("scratch")); + function.body.push(rumoca_core::Statement::If { + cond_blocks: vec![rumoca_core::StatementBlock { + cond: binary( + rumoca_core::OpBinary::Gt, + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + args: vec![local_var("x"), integer(1)], + span: test_span(), + }, + integer(1), + test_span(), + ), + stmts: vec![scalar_assignment("y", local_var("scratch"))], + }], + else_block: Some(vec![scalar_assignment("y", expanded)]), + span: test_span(), + }); + + let mut dae_model = dae::Dae::default(); + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let mut scope = FunctionProjectionScope::default(); + let actual_values = vec![real(1.0), real(2.0), real(3.0), real(4.0)]; + scope + .full + .insert("u".to_string(), array(actual_values.clone(), false)); + scope.scalars.insert("u".to_string(), actual_values); + scope.dims.insert("u".to_string(), vec![4]); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.laneBudget").into(), + args: vec![local_var("u")], + is_constructor: false, + span: test_span(), + }; + + let full_outputs = analysis + .function_call_outputs_with_projection_scope(&call, 1, test_span(), Some(&scope)) + .expect("full caller-scope probe should not fail"); + assert!( + full_outputs.is_none(), + "full vector call must decline because its selected output leaks a local" + ); + + let err = analysis + .project_function_call_value(&call, &[4], 0, &scope, 0, test_span()) + .expect_err("lane-rewritten budget exhaustion must remain a typed error"); + assert!(err.is_projection_budget_exceeded(), "got: {err:?}"); +} + +fn nested_over_budget_vector_call() -> (dae::Dae, rumoca_core::Expression) { + let expanded = expression_over_projection_budget(local_var("x")); + let mut explode_vector = over_budget_function(); + explode_vector.name = rumoca_core::VarName::new("My.explodeVector"); + explode_vector.inputs[0].dims = vec![4]; + explode_vector.outputs[0].dims = vec![4]; + explode_vector.body = vec![scalar_assignment("y", expanded)]; + + let mut wrapper = rumoca_core::Function::new("My.wrapper", test_span()); + wrapper.inputs.push(function_param_with_dims("u", &[4])); + wrapper.outputs.push(function_param_with_dims("y", &[4])); + wrapper.body.push(scalar_assignment( + "y", + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.explodeVector").into(), + args: vec![local_var("u")], + is_constructor: false, + span: test_span(), + }, + )); + + let mut dae_model = dae::Dae::default(); + dae_model + .symbols + .functions + .insert(explode_vector.name.clone(), explode_vector); + dae_model + .symbols + .functions + .insert(wrapper.name.clone(), wrapper); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.wrapper").into(), + args: vec![array( + vec![real(1.0), real(2.0), real(3.0), real(4.0)], + false, + )], + is_constructor: false, + span: test_span(), + }; + (dae_model, call) +} + +fn under_budget_vector_call() -> (dae::Dae, rumoca_core::Expression) { + let mut function = rumoca_core::Function::new("My.vectorIdentity", test_span()); + function.inputs.push(function_param_with_dims("x", &[2])); + function.outputs.push(function_param_with_dims("y", &[2])); + function.body.push(scalar_assignment("y", local_var("x"))); + + let mut dae_model = dae::Dae::default(); + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.vectorIdentity").into(), + args: vec![array(vec![real(1.0), real(2.0)], false)], + is_constructor: false, + span: test_span(), + }; + (dae_model, call) +} + +fn array_like_boundary_results( + expr: &rumoca_core::Expression, + dae_model: &dae::Dae, + flat_index: usize, +) -> ( + Option>, + Option, +) { + let structural_bindings = IndexMap::new(); + let values = + project_array_like_scalars_with_owner(expr, dae_model, &structural_bindings, test_span()) + .expect("array-like projection should not fail"); + let value = project_array_like_scalar_with_owner( + expr, + flat_index, + dae_model, + &structural_bindings, + test_span(), + ) + .expect("array-like scalar projection should not fail"); + (values, value) +} + +#[test] +fn nested_over_budget_vector_call_preserves_outer_runtime_call() { + let (dae_model, call) = nested_over_budget_vector_call(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + + let err = analysis + .function_call_outputs_with_owner(&call, 0, test_span()) + .expect_err("nested budget exhaustion must remain a typed internal error"); + assert!(err.is_projection_budget_exceeded(), "got: {err:?}"); + + let outputs = analysis + .top_level_function_call_outputs(&call, test_span()) + .expect("the outer boundary must preserve the complete runtime call"); + assert!( + outputs.is_none(), + "the outer boundary must decline projection instead of scalarizing call arguments" + ); +} + +#[test] +fn over_budget_root_function_call_declines_at_array_like_scalars_boundary() { + let (dae_model, call) = nested_over_budget_vector_call(); + let structural_bindings = IndexMap::new(); + + let values = + project_array_like_scalars_with_owner(&call, &dae_model, &structural_bindings, test_span()) + .expect("the outer array-like boundary must preserve the complete runtime call"); + + assert!(values.is_none()); +} + +#[test] +fn over_budget_root_function_call_declines_at_array_like_scalar_boundary() { + let (dae_model, call) = nested_over_budget_vector_call(); + let structural_bindings = IndexMap::new(); + + let value = project_array_like_scalar_with_owner( + &call, + 0, + &dae_model, + &structural_bindings, + test_span(), + ) + .expect("the outer array-like scalar boundary must preserve the complete runtime call"); + + assert!(value.is_none()); +} + +#[test] +fn under_budget_root_function_call_projects_at_array_like_boundaries() { + let (dae_model, call) = under_budget_vector_call(); + let (Some(values), Some(value)) = array_like_boundary_results(&call, &dae_model, 1) else { + panic!("under-budget root function call should project at both boundaries"); + }; + + assert_eq!(values.len(), 2); + assert_eq!(values[1], value); +} + +#[test] +fn non_function_constructor_and_unknown_array_like_behavior_is_unchanged() { + let dae_model = dae::Dae::default(); + let literal = array(vec![real(1.0), real(2.0)], false); + let (Some(values), Some(value)) = array_like_boundary_results(&literal, &dae_model, 1) else { + panic!("literal array should keep projecting at both boundaries"); + }; + + assert_eq!(values.len(), 2); + assert_eq!(values[1], value); + for declined in [ + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.Record").into(), + args: vec![real(1.0), real(2.0)], + is_constructor: true, + span: test_span(), + }, + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.unknownVector").into(), + args: Vec::new(), + is_constructor: false, + span: test_span(), + }, + ] { + let (values, value) = array_like_boundary_results(&declined, &dae_model, 0); + assert!(values.is_none()); + assert!(value.is_none()); + } +} + +#[test] +fn external_and_impure_array_calls_decline_at_both_projection_boundaries() { + let mut external = rumoca_core::Function::new("My.externalVector", test_span()); + external.outputs.push(function_param_with_dims("y", &[2])); + external.external = Some(rumoca_core::ExternalFunction::default()); + let mut impure = rumoca_core::Function::new("My.impureVector", test_span()); + impure.outputs.push(function_param_with_dims("y", &[2])); + impure.pure = false; + + let mut dae_model = dae::Dae::default(); + for function in [external, impure] { + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + } + for name in ["My.externalVector", "My.impureVector"] { + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new(name).into(), + args: Vec::new(), + is_constructor: false, + span: test_span(), + }; + let (values, value) = array_like_boundary_results(&call, &dae_model, 1); + assert!( + values.is_none(), + "{name} must stay whole at the array-like scalars boundary" + ); + assert!( + value.is_none(), + "{name} must stay whole at the array-like scalar boundary" + ); + } +} + +#[test] +fn projection_declines_when_output_leaks_function_local_reference() { + let mut function = rumoca_core::Function::new("My.leaksLocal", test_span()); + function.outputs.push(scalar_function_param("y")); + function + .locals + .push(function_param_with_type("scratch", "Pkg.Record")); + function.body.push(scalar_assignment( + "y", + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("scratch.value"), + subscripts: Vec::new(), + span: test_span(), + }, + )); + + let mut dae_model = dae::Dae::default(); + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.leaksLocal").into(), + args: Vec::new(), + is_constructor: false, + span: test_span(), + }; + + let outputs = analysis + .top_level_function_call_outputs(&call, test_span()) + .expect("local leakage should decline optional projection"); + + assert!(outputs.is_none()); +} + +#[test] +fn projection_allows_input_actual_with_same_name_as_formal() { + let mut function = rumoca_core::Function::new("My.sameName", test_span()); + function.inputs.push(scalar_function_param("T")); + function.outputs.push(scalar_function_param("y")); + function.body.push(scalar_assignment("y", local_var("T"))); + + let mut dae_model = dae::Dae::default(); + dae_model.variables.parameters.insert( + rumoca_core::VarName::new("T"), + dae::Variable { + name: rumoca_core::VarName::new("T"), + ..rumoca_ir_dae::Variable::empty_with_span(test_span()) + }, + ); + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.sameName").into(), + args: vec![local_var("T")], + is_constructor: false, + span: test_span(), + }; + + let outputs = analysis + .top_level_function_call_outputs(&call, test_span()) + .expect("same-name formal and actual should not fail projection") + .expect("same-name actual should remain projectable"); + + assert_eq!(outputs.len(), 1); +} + +#[test] +fn static_while_projection_executes_until_condition_is_false() { + let mut function = rumoca_core::Function::new("My.staticWhile", test_span()); + function.outputs.push(scalar_function_param("y")); + let mut alpha = scalar_function_param("alpha"); + alpha.default = Some(real(1.0)); + function.locals.push(alpha); + function.body.push(rumoca_core::Statement::While { + block: rumoca_core::StatementBlock { + cond: binary( + rumoca_core::OpBinary::Lt, + local_var("alpha"), + real(4.0), + test_span(), + ), + stmts: vec![scalar_assignment( + "alpha", + binary( + rumoca_core::OpBinary::Mul, + real(2.0), + local_var("alpha"), + test_span(), + ), + )], + }, + span: test_span(), + }); + function + .body + .push(scalar_assignment("y", local_var("alpha"))); + + let mut dae_model = dae::Dae::default(); + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.staticWhile").into(), + args: Vec::new(), + is_constructor: false, + span: test_span(), + }; + + let outputs = analysis + .top_level_function_call_outputs(&call, test_span()) + .expect("static while function should project") + .expect("static while output should be available"); + let value = analysis + .compile_time_scalar_in_scope(&outputs[0].expr, &FunctionProjectionScope::default()) + .expect("projected while result should be compile-time evaluable") + .expect("projected while result should be scalar"); + + assert_eq!(value, 4.0); +} + +#[test] +fn static_while_projection_evaluates_scalar_function_condition() { + let mut residue = rumoca_core::Function::new("My.residue", test_span()); + residue.inputs.push(scalar_function_param("alpha")); + residue.outputs.push(scalar_function_param("y")); + residue.body.push(scalar_assignment( + "y", + binary( + rumoca_core::OpBinary::Sub, + local_var("alpha"), + real(4.0), + test_span(), + ), + )); + + let mut function = rumoca_core::Function::new("My.staticWhileWithCall", test_span()); + function.outputs.push(scalar_function_param("y")); + let mut alpha = scalar_function_param("alpha"); + alpha.default = Some(real(1.0)); + function.locals.push(alpha); + function.locals.push(scalar_function_param("residue")); + function.body.push(scalar_assignment( + "residue", + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.residue").into(), + args: vec![local_var("alpha")], + is_constructor: false, + span: test_span(), + }, + )); + function.body.push(rumoca_core::Statement::While { + block: rumoca_core::StatementBlock { + cond: binary( + rumoca_core::OpBinary::Lt, + local_var("residue"), + real(0.0), + test_span(), + ), + stmts: vec![ + scalar_assignment( + "alpha", + binary( + rumoca_core::OpBinary::Mul, + real(2.0), + local_var("alpha"), + test_span(), + ), + ), + scalar_assignment( + "residue", + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.residue").into(), + args: vec![local_var("alpha")], + is_constructor: false, + span: test_span(), + }, + ), + ], + }, + span: test_span(), + }); + function + .body + .push(scalar_assignment("y", local_var("alpha"))); + + let mut dae_model = dae::Dae::default(); + dae_model + .symbols + .functions + .insert(residue.name.clone(), residue); + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.staticWhileWithCall").into(), + args: Vec::new(), + is_constructor: false, + span: test_span(), + }; + + let outputs = analysis + .top_level_function_call_outputs(&call, test_span()) + .expect("static while function should project through scalar function condition") + .expect("static while output should be available"); + let value = analysis + .compile_time_scalar_in_scope(&outputs[0].expr, &FunctionProjectionScope::default()) + .expect("projected while result should be compile-time evaluable") + .expect("projected while result should be scalar"); + + assert_eq!(value, 4.0); +} + +#[test] +fn scalar_function_output_substitutes_local_bindings_before_scope_check() { + let mut function = rumoca_core::Function::new("My.localScalar", test_span()); + function.inputs.push(scalar_function_param("alpha")); + function.outputs.push(scalar_function_param("y")); + let mut beta = scalar_function_param("beta"); + beta.default = Some(real(2.0)); + function.locals.push(beta); + function.body.push(scalar_assignment( + "y", + binary( + rumoca_core::OpBinary::Sub, + local_var("alpha"), + local_var("beta"), + test_span(), + ), + )); + + let mut dae_model = dae::Dae::default(); + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.localScalar").into(), + args: vec![real(5.0)], + is_constructor: false, + span: test_span(), + }; + + let outputs = analysis + .top_level_function_call_outputs(&call, test_span()) + .expect("local scalar output should project") + .expect("local scalar output should not be rejected as unresolved"); + let value = analysis + .compile_time_scalar_in_scope(&outputs[0].expr, &FunctionProjectionScope::default()) + .expect("projected output should be compile-time evaluable") + .expect("projected output should be scalar"); + + assert_eq!(value, 3.0); +} + +#[test] +fn scalar_function_assignment_freezes_projected_array_inputs() { + let mut first = rumoca_core::Function::new("My.first", test_span()); + first.inputs.push(function_param_with_dims("u", &[1])); + first.outputs.push(scalar_function_param("y")); + first.body.push(scalar_assignment( + "y", + rumoca_core::Expression::Index { + base: Box::new(local_var("u")), + subscripts: vec![rumoca_core::Subscript::Index { + value: 1, + span: test_span(), + }], + span: test_span(), + }, + )); + + let mut function = rumoca_core::Function::new("My.freezeScalarCall", test_span()); + function.outputs.push(scalar_function_param("y")); + function.locals.push(function_param_with_dims("u", &[1])); + function.locals.push(scalar_function_param("alpha")); + function.body.push(assignment_with_span( + "u", + array(vec![real(2.0)], false), + test_span(), + )); + function.body.push(scalar_assignment( + "alpha", + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.first").into(), + args: vec![local_var("u")], + is_constructor: false, + span: test_span(), + }, + )); + function.body.push(assignment_with_span( + "u", + binary( + rumoca_core::OpBinary::Mul, + local_var("u"), + real(3.0), + test_span(), + ), + test_span(), + )); + function + .body + .push(scalar_assignment("y", local_var("alpha"))); + + let mut dae_model = dae::Dae::default(); + dae_model + .symbols + .functions + .insert(first.name.clone(), first); + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.freezeScalarCall").into(), + args: Vec::new(), + is_constructor: false, + span: test_span(), + }; + + let outputs = analysis + .top_level_function_call_outputs(&call, test_span()) + .expect("scalar function assignment should project") + .expect("scalar output should be available"); + let value = analysis + .compile_time_scalar_in_scope(&outputs[0].expr, &FunctionProjectionScope::default()) + .expect("projected scalar assignment should be compile-time evaluable") + .expect("projected scalar assignment should be scalar"); + + assert_eq!(value, 2.0); +} + +#[test] +// SPEC_0021: Exception - regression fixture builds two nested Modelica-style +// functions inline so the size/while projection scenario stays auditable. +#[allow(clippy::too_many_lines)] +fn static_while_projection_uses_nested_matrix_input_size() { + fn indexed_var(name: &str, indices: Vec) -> rumoca_core::Expression { + rumoca_core::Expression::Index { + base: Box::new(local_var(name)), + subscripts: indices, + span: test_span(), + } + } + + let mut residue = rumoca_core::Function::new("My.matrixResidue", test_span()); + residue.inputs.push(function_param_with_dims("c1", &[0])); + residue.inputs.push(function_param_with_dims("c2", &[0, 2])); + residue.inputs.push(scalar_function_param("alpha")); + residue.outputs.push(scalar_function_param("residue")); + let mut alpha2 = scalar_function_param("alpha2"); + alpha2.default = Some(binary( + rumoca_core::OpBinary::Mul, + local_var("alpha"), + local_var("alpha"), + test_span(), + )); + residue.locals.push(alpha2); + let mut a2 = scalar_function_param("A2"); + a2.default = Some(real(1.0)); + residue.locals.push(a2); + residue.body.push(rumoca_core::Statement::If { + cond_blocks: vec![rumoca_core::StatementBlock { + cond: binary( + rumoca_core::OpBinary::Eq, + builtin( + rumoca_core::BuiltinFunction::Size, + vec![local_var("c1"), integer(1)], + ), + real(1.0), + test_span(), + ), + stmts: vec![scalar_assignment( + "A2", + binary( + rumoca_core::OpBinary::Mul, + local_var("A2"), + binary( + rumoca_core::OpBinary::Add, + real(1.0), + binary( + rumoca_core::OpBinary::Mul, + indexed_var( + "c1", + vec![rumoca_core::Subscript::Index { + value: 1, + span: test_span(), + }], + ), + local_var("alpha2"), + test_span(), + ), + test_span(), + ), + test_span(), + ), + )], + }], + else_block: None, + span: test_span(), + }); + residue.body.push(rumoca_core::Statement::For { + indices: vec![rumoca_core::ForIndex { + ident: "i".to_string(), + range: rumoca_core::Expression::Range { + start: Box::new(integer(1)), + step: None, + end: Box::new(builtin( + rumoca_core::BuiltinFunction::Size, + vec![local_var("c2"), integer(1)], + )), + span: test_span(), + }, + }], + equations: vec![scalar_assignment( + "A2", + binary( + rumoca_core::OpBinary::Mul, + local_var("A2"), + binary( + rumoca_core::OpBinary::Add, + real(1.0), + binary( + rumoca_core::OpBinary::Mul, + indexed_var( + "c2", + vec![ + rumoca_core::Subscript::Expr { + expr: Box::new(local_var("i")), + span: test_span(), + }, + rumoca_core::Subscript::Index { + value: 1, + span: test_span(), + }, + ], + ), + local_var("alpha2"), + test_span(), + ), + test_span(), + ), + test_span(), + ), + )], + span: test_span(), + }); + residue.body.push(scalar_assignment( + "residue", + binary( + rumoca_core::OpBinary::Sub, + binary( + rumoca_core::OpBinary::Div, + real(1.0), + builtin(rumoca_core::BuiltinFunction::Sqrt, vec![local_var("A2")]), + test_span(), + ), + real(0.7), + test_span(), + ), + )); + + let mut function = rumoca_core::Function::new("My.matrixFind", test_span()); + function.inputs.push(function_param_with_dims("c1", &[0])); + function + .inputs + .push(function_param_with_dims("c2", &[0, 2])); + function.outputs.push(scalar_function_param("alpha")); + let mut residue_local = scalar_function_param("residue"); + residue_local.default = Some(real(1.0)); + function.locals.push(residue_local); + function.body.push(scalar_assignment("alpha", real(1.0))); + function.body.push(scalar_assignment( + "residue", + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.matrixResidue").into(), + args: vec![local_var("c1"), local_var("c2"), local_var("alpha")], + is_constructor: false, + span: test_span(), + }, + )); + function.body.push(rumoca_core::Statement::If { + cond_blocks: vec![rumoca_core::StatementBlock { + cond: binary( + rumoca_core::OpBinary::Lt, + local_var("residue"), + real(0.0), + test_span(), + ), + stmts: Vec::new(), + }], + else_block: Some(vec![rumoca_core::Statement::While { + block: rumoca_core::StatementBlock { + cond: binary( + rumoca_core::OpBinary::Ge, + local_var("residue"), + real(0.0), + test_span(), + ), + stmts: vec![ + scalar_assignment( + "alpha", + binary( + rumoca_core::OpBinary::Mul, + real(2.0), + local_var("alpha"), + test_span(), + ), + ), + scalar_assignment( + "residue", + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.matrixResidue").into(), + args: vec![local_var("c1"), local_var("c2"), local_var("alpha")], + is_constructor: false, + span: test_span(), + }, + ), + ], + }, + span: test_span(), + }]), + span: test_span(), + }); + + let mut dae_model = dae::Dae::default(); + dae_model + .symbols + .functions + .insert(residue.name.clone(), residue); + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.matrixFind").into(), + args: vec![ + array(vec![real(0.0)], false), + array(vec![array(vec![real(1.0), real(0.0)], false)], true), + ], + is_constructor: false, + span: test_span(), + }; + + let outputs = analysis + .top_level_function_call_outputs(&call, test_span()) + .expect("matrix-size while should project") + .expect("matrix-size while output should be available"); + let value = analysis + .compile_time_scalar_in_scope(&outputs[0].expr, &FunctionProjectionScope::default()) + .expect("projected matrix while result should be compile-time evaluable") + .expect("projected matrix while result should be scalar"); + + assert_eq!(value, 2.0); +} + +#[test] +fn declared_scalar_shape_initialization_seeds_plain_scalars_without_overwrite() +-> Result<(), LowerError> { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let mut function = rumoca_core::Function::new("My.runtimeScalar", test_span()); + function.locals.push(scalar_function_param("gain")); + function + .locals + .push(scalar_function_param("vectorized_gain")); + function.locals.push(function_param_with_type( + "structured_value", + "Pkg.Quaternion", + )); + function.outputs.push(scalar_function_param("result")); + function + .outputs + .push(scalar_function_param("vectorized_result")); + let mut scope = FunctionProjectionScope::default(); + scope.dims.insert("vectorized_gain".to_string(), vec![3]); + + analysis.initialize_projected_declared_arrays(&function, &mut scope, 0, test_span())?; + + assert_eq!(scope.dims.get("gain"), Some(&Vec::new())); + assert_eq!(scope.dims.get("result"), Some(&Vec::new())); + assert_eq!(scope.dims.get("vectorized_result"), Some(&Vec::new())); + assert_eq!(scope.dims.get("vectorized_gain"), Some(&vec![3])); + assert!(!scope.dims.contains_key("structured_value")); + + scope.dims.insert("vectorized_result".to_string(), vec![3]); + assert_eq!(scope.dims.get("vectorized_result"), Some(&vec![3])); + Ok(()) +} + +#[test] +fn declared_scalar_shape_input_binding_seeds_plain_scalar_without_overwrite() +-> Result<(), LowerError> { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let mut nested = rumoca_core::Function::new("My.nested", test_span()); + nested.inputs.push(scalar_function_param("nested_gain")); + nested.inputs.push(scalar_function_param("vectorized_gain")); + let nested_scope = analysis + .bind_inputs_with_projection_scope( + &nested, + &[ + local_var("runtime_scalar"), + array(vec![real(1.0), real(2.0), real(3.0)], false), + ], + 0, + test_span(), + None, + )? + .expect("declared nested scalar input should bind"); + + assert_eq!(nested_scope.dims.get("nested_gain"), Some(&Vec::new())); + assert_eq!(nested_scope.dims.get("vectorized_gain"), Some(&vec![3])); + Ok(()) +} + +fn apply_runtime_scalar_if( + analysis: &FunctionProjectionAnalysis<'_>, + function: &rumoca_core::Function, + target: &str, + scope: &mut FunctionProjectionScope, +) -> Result<(), LowerError> { + let statement = rumoca_core::Statement::If { + cond_blocks: vec![rumoca_core::StatementBlock { + cond: local_var("runtime_condition"), + stmts: vec![scalar_assignment(target, local_var("runtime_then_value"))], + }], + else_block: Some(vec![scalar_assignment( + target, + local_var("runtime_else_value"), + )]), + span: test_span(), + }; + analysis.apply_statement(function, &statement, scope, &mut Vec::new(), 0, test_span()) +} - assert_eq!(err.source_span(), Some(span)); - assert_eq!( - err.reason(), - "function `My.badAssign` assignment to `x` expects dimensions [2], got []" - ); +fn declared_scalar_three_vector_product(scalar_name: &str) -> rumoca_core::Expression { + binary( + rumoca_core::OpBinary::Mul, + local_var(scalar_name), + array(vec![real(1.0), real(2.0), real(3.0)], false), + test_span(), + ) +} + +fn assert_declared_scalar_three_vector_projection( + analysis: &FunctionProjectionAnalysis<'_>, + scalar_name: &str, + scope: &FunctionProjectionScope, +) -> Result<(), LowerError> { + let product = declared_scalar_three_vector_product(scalar_name); + let projected = analysis + .project_value_scalars(&product, &[3], scope, 0, test_span())? + .expect("declared scalar times known vector should project"); + assert_eq!(projected.len(), 3); + let dims = + analysis.known_expr_dims(&product, scope, 0, "declared scalar product", test_span())?; + assert_eq!(dims, vec![3]); + Ok(()) } #[test] -fn scalar_assignment_rejects_vector_value_with_span() { +fn declared_scalar_shape_survives_runtime_if_merge_into_vector_product() -> Result<(), LowerError> { let dae_model = dae::Dae::default(); let structural_bindings = IndexMap::new(); let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); - let mut function = rumoca_core::Function::new("My.badAssign", test_span()); - function.locals.push(scalar_function_param("x")); + let mut function = rumoca_core::Function::new("My.runtimeScalar", test_span()); + function.locals.push(scalar_function_param("gain")); let mut scope = FunctionProjectionScope::default(); - let mut projected = Vec::new(); - let span = rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name( - "phase_solve_lower_derivative_rhs_function_projection_tests_source_56.mo", - ), - 4, - 12, - ); - let statement = assignment_with_span( - "x", - rumoca_core::Expression::Array { - elements: vec![real(1.0), real(2.0)], - is_matrix: false, - span, - }, - span, - ); + analysis.initialize_projected_declared_arrays(&function, &mut scope, 0, test_span())?; - let err = analysis - .apply_assignment(&function, &statement, &mut scope, &mut projected, 0, span) - .expect_err("vector assignment to scalar local must fail"); + apply_runtime_scalar_if(&analysis, &function, "gain", &mut scope)?; - assert_eq!(err.source_span(), Some(span)); - assert_eq!( - err.reason(), - "function `My.badAssign` assignment to `x` expects dimensions [], got [2]" - ); + assert_declared_scalar_three_vector_projection(&analysis, "gain", &scope) } #[test] -fn unassigned_projected_scalar_reports_projection_span() -> Result<(), String> { - let span = rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name( - "phase_solve_lower_derivative_rhs_function_projection_tests_source_57.mo", - ), - 6, - 15, - ); - let values = vec![rumoca_core::Expression::Empty { span }]; +fn declared_scalar_shape_survives_caller_scope_substitution_into_nested_input() +-> Result<(), LowerError> { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let mut caller = rumoca_core::Function::new("My.caller", test_span()); + caller.locals.push(scalar_function_param("gain")); + let mut caller_scope = FunctionProjectionScope::default(); + analysis.initialize_projected_declared_arrays(&caller, &mut caller_scope, 0, test_span())?; + apply_runtime_scalar_if(&analysis, &caller, "gain", &mut caller_scope)?; + + let mut nested = rumoca_core::Function::new("My.nested", test_span()); + nested.inputs.push(scalar_function_param("nested_gain")); + let nested_scope = analysis + .bind_inputs_with_projection_scope( + &nested, + &[local_var("gain")], + 0, + test_span(), + Some(&caller_scope), + )? + .expect("declared nested scalar input should bind through caller scope"); - let Err(err) = assigned_projected_scalar_value("x", &[1], &values, 0, span) else { - return Err("unassigned projected scalar slot succeeded".to_string()); - }; + assert_declared_scalar_three_vector_projection(&analysis, "nested_gain", &nested_scope) +} - assert_eq!(err.source_span(), Some(span)); - assert_eq!( - err.reason(), - "projected local component `x[1]` is unassigned" - ); - Ok(()) +fn declared_scalar_normalize_function() -> rumoca_core::Function { + let mut function = rumoca_core::Function::new("My.normalize", test_span()); + function.inputs.push(function_param_with_dims("q", &[4])); + function.inputs.push(scalar_function_param("runtime_norm")); + function.outputs.push(function_param_with_dims("q_n", &[4])); + function.locals.push(scalar_function_param("n")); + function.body.push(scalar_assignment( + "n", + binary( + rumoca_core::OpBinary::Add, + local_var("runtime_norm"), + real(1.0e-10), + test_span(), + ), + )); + function.body.push(scalar_assignment( + "q_n", + binary( + rumoca_core::OpBinary::Div, + local_var("q"), + local_var("n"), + test_span(), + ), + )); + function } #[test] -fn scalar_selector_rejects_colon_with_subscript_span() -> Result<(), String> { +fn original_projection_fallback_requires_erased_declared_scalar_shape() { let dae_model = dae::Dae::default(); let structural_bindings = IndexMap::new(); let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); - let scope = FunctionProjectionScope::default(); - let span = rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name( - "phase_solve_lower_derivative_rhs_function_projection_tests_source_58.mo", - ), - 2, - 3, - ); - let subscript = rumoca_core::Subscript::colon(span); - - let Err(err) = subscript_selector_expr(&subscript, &analysis, &scope, 0) else { - return Err("colon scalar selector succeeded".to_string()); - }; - - assert_eq!(err.source_span(), Some(span)); - assert_eq!( - err.reason(), - "colon subscript cannot select a scalar projected value" + let original = binary( + rumoca_core::OpBinary::Div, + local_var("q"), + local_var("n"), + test_span(), ); - Ok(()) -} - -#[test] -fn guarded_assignment_without_base_reports_assignment_span() { - let span = rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name( - "phase_solve_lower_derivative_rhs_function_projection_tests_source_59.mo", + let mut scope = FunctionProjectionScope::default(); + scope.dims.insert("n".to_string(), Vec::new()); + scope.full.insert( + "n".to_string(), + binary( + rumoca_core::OpBinary::Add, + local_var("runtime_norm"), + real(1.0e-10), + test_span(), ), - 9, - 21, ); - let err = guarded_assignment_without_base("y", span); + assert!(analysis.original_projection_preserves_declared_scalar_shape(&original, &scope)); - assert_eq!(err.source_span(), Some(span)); - assert_eq!( - err.reason(), - "if-statement assignment to `y` requires an existing binding or an else assignment" + let undeclared_scope = FunctionProjectionScope::default(); + assert!( + !analysis.original_projection_preserves_declared_scalar_shape(&original, &undeclared_scope) ); -} -fn local_var(name: &str) -> rumoca_core::Expression { - rumoca_core::Expression::VarRef { - name: rumoca_core::Reference::new(name), - subscripts: Vec::new(), - span: test_span(), - } -} + scope.dims.insert("n".to_string(), vec![4]); + assert!(!analysis.original_projection_preserves_declared_scalar_shape(&original, &scope)); -/// A pure function whose projected output doubles in size per statement, -/// crossing `MAX_FUNCTION_PROJECTION_NODES` long before it finishes. -fn over_budget_function() -> rumoca_core::Function { - let mut body = vec![scalar_assignment( - "y", - rumoca_core::Expression::Binary { - op: rumoca_core::OpBinary::Mul, - lhs: Box::new(local_var("x")), - rhs: Box::new(local_var("x")), - span: test_span(), - }, - )]; - for _ in 0..16 { - body.push(scalar_assignment( - "y", - rumoca_core::Expression::Binary { - op: rumoca_core::OpBinary::Mul, - lhs: Box::new(local_var("y")), - rhs: Box::new(local_var("y")), - span: test_span(), - }, - )); - } - rumoca_core::Function { - name: rumoca_core::VarName::new("My.explode"), - def_id: None, - inputs: vec![scalar_function_param("x")], - outputs: vec![scalar_function_param("y")], - locals: vec![], - body, + scope.dims.insert("n".to_string(), Vec::new()); + scope.full.insert("n".to_string(), local_var("n")); + assert!(!analysis.original_projection_preserves_declared_scalar_shape(&original, &scope)); + + scope.full.insert("n".to_string(), real(2.0)); + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.boundary").into(), + args: vec![original.clone()], is_constructor: false, - pure: true, - external: None, - derivatives: vec![], span: test_span(), - } + }; + assert!(!analysis.original_projection_preserves_declared_scalar_shape(&call, &scope)); + let constructor = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.constructorBoundary").into(), + args: vec![original], + is_constructor: true, + span: test_span(), + }; + assert!(!analysis.original_projection_preserves_declared_scalar_shape(&constructor, &scope)); } #[test] -fn over_budget_projection_is_a_typed_error_and_declines_at_the_boundary() { - let mut dae_model = dae::Dae::default(); - dae_model.symbols.functions.insert( - rumoca_core::VarName::new("My.explode"), - over_budget_function(), - ); +fn original_projection_fallback_rejects_boundary_siblings() { + let dae_model = dae::Dae::default(); let structural_bindings = IndexMap::new(); let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); - let call = rumoca_core::Expression::FunctionCall { - name: rumoca_core::VarName::new("My.explode").into(), - args: vec![real(2.0)], + let mut scope = FunctionProjectionScope::default(); + scope.dims.insert("n".to_string(), Vec::new()); + scope.full.insert("n".to_string(), real(2.0)); + let boundary = || rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.boundary").into(), + args: vec![array(vec![real(1.0), real(2.0)], false)], is_constructor: false, span: test_span(), }; - let err = analysis - .function_call_outputs_with_owner(&call, 0, test_span()) - .expect_err("over-budget projection must surface as a typed error"); - assert!(err.is_projection_budget_exceeded(), "got: {err:?}"); - assert!(err.reason().contains("My.explode"), "got: {}", err.reason()); + let candidates = vec![ + binary( + rumoca_core::OpBinary::Add, + boundary(), + local_var("n"), + test_span(), + ), + array(vec![boundary(), local_var("n")], false), + rumoca_core::Expression::Tuple { + elements: vec![local_var("n"), boundary()], + span: test_span(), + }, + rumoca_core::Expression::If { + branches: vec![(local_var("condition"), boundary())], + else_branch: Box::new(local_var("n")), + span: test_span(), + }, + rumoca_core::Expression::Range { + start: Box::new(local_var("n")), + step: None, + end: Box::new(boundary()), + span: test_span(), + }, + builtin( + rumoca_core::BuiltinFunction::Max, + vec![boundary(), local_var("n")], + ), + array( + vec![binary( + rumoca_core::OpBinary::Mul, + local_var("n"), + rumoca_core::Expression::Index { + base: Box::new(local_var("vector")), + subscripts: vec![rumoca_core::Subscript::Index { + value: 1, + span: test_span(), + }], + span: test_span(), + }, + test_span(), + )], + false, + ), + ]; - // The outermost boundary resolves the decline by keeping the runtime - // call; the memoized decline must answer follow-up probes identically. - for _ in 0..2 { - let outputs = analysis - .top_level_function_call_outputs(&call, test_span()) - .expect("budget decline must not fail the outer lowering"); - assert!(outputs.is_none()); + for candidate in candidates { + assert!( + !analysis.original_projection_preserves_declared_scalar_shape(&candidate, &scope), + "projection boundary sibling must reject original projection: {candidate:?}" + ); } } #[test] -fn projection_declines_when_output_leaks_function_local_reference() { - let mut function = rumoca_core::Function::new("My.leaksLocal", test_span()); - function.outputs.push(scalar_function_param("y")); - function - .locals - .push(function_param_with_type("scratch", "Pkg.Record")); - function.body.push(scalar_assignment( - "y", - rumoca_core::Expression::VarRef { - name: rumoca_core::Reference::new("scratch.value"), - subscripts: Vec::new(), - span: test_span(), - }, - )); +fn declared_scalar_shape_survives_assignment_substitution_in_vector_division() +-> Result<(), LowerError> { + let function = declared_scalar_normalize_function(); let mut dae_model = dae::Dae::default(); dae_model @@ -1277,51 +4774,113 @@ fn projection_declines_when_output_leaks_function_local_reference() { let structural_bindings = IndexMap::new(); let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); let call = rumoca_core::Expression::FunctionCall { - name: rumoca_core::VarName::new("My.leaksLocal").into(), - args: Vec::new(), + name: rumoca_core::VarName::new("My.normalize").into(), + args: vec![ + array(vec![real(1.0), real(2.0), real(3.0), real(4.0)], false), + local_var("runtime_norm_actual"), + ], is_constructor: false, span: test_span(), }; let outputs = analysis - .top_level_function_call_outputs(&call, test_span()) - .expect("local leakage should decline optional projection"); + .top_level_function_call_outputs(&call, test_span())? + .expect("declared scalar denominator should preserve vector projection"); - assert!(outputs.is_none()); + assert_eq!(outputs.len(), 4); + assert!(outputs.iter().all(|output| matches!( + output.expr, + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Div, + .. + } + ))); + Ok(()) } #[test] -fn projection_allows_input_actual_with_same_name_as_formal() { - let mut function = rumoca_core::Function::new("My.sameName", test_span()); - function.inputs.push(scalar_function_param("T")); - function.outputs.push(scalar_function_param("y")); - function.body.push(scalar_assignment("y", local_var("T"))); +fn vector_output_function_call_assignment_keeps_vector_actuals() -> Result<(), LowerError> { + let normalize = declared_scalar_normalize_function(); + + let mut wrapper = rumoca_core::Function::new("My.wrapper", test_span()); + wrapper.inputs.push(function_param_with_dims("state", &[5])); + wrapper.inputs.push(scalar_function_param("runtime_norm")); + wrapper + .outputs + .push(function_param_with_dims("limited", &[4])); + wrapper.body.push(scalar_assignment( + "limited", + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.normalize").into(), + args: vec![ + rumoca_core::Expression::Index { + base: Box::new(local_var("state")), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(rumoca_core::Expression::Range { + start: Box::new(integer(1)), + step: None, + end: Box::new(integer(4)), + span: test_span(), + }), + span: test_span(), + }], + span: test_span(), + }, + local_var("runtime_norm"), + ], + is_constructor: false, + span: test_span(), + }, + )); let mut dae_model = dae::Dae::default(); - dae_model.variables.parameters.insert( - rumoca_core::VarName::new("T"), - dae::Variable { - name: rumoca_core::VarName::new("T"), - ..rumoca_ir_dae::Variable::empty_with_span(test_span()) - }, - ); dae_model .symbols .functions - .insert(function.name.clone(), function); + .insert(normalize.name.clone(), normalize); + dae_model + .symbols + .functions + .insert(wrapper.name.clone(), wrapper); let structural_bindings = IndexMap::new(); let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); let call = rumoca_core::Expression::FunctionCall { - name: rumoca_core::VarName::new("My.sameName").into(), - args: vec![local_var("T")], + name: rumoca_core::VarName::new("My.wrapper").into(), + args: vec![ + array( + vec![real(1.0), real(2.0), real(3.0), real(4.0), real(5.0)], + false, + ), + local_var("runtime_norm_actual"), + ], is_constructor: false, span: test_span(), }; let outputs = analysis - .top_level_function_call_outputs(&call, test_span()) - .expect("same-name formal and actual should not fail projection") - .expect("same-name actual should remain projectable"); + .top_level_function_call_outputs(&call, test_span())? + .expect("vector output call assignment should preserve its vector actual"); - assert_eq!(outputs.len(), 1); + assert_eq!(outputs.len(), 4); + Ok(()) +} + +#[test] +fn declared_scalar_shape_does_not_cover_undeclared_runtime_reference() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let scope = FunctionProjectionScope::default(); + + let product = binary( + rumoca_core::OpBinary::Mul, + local_var("undeclared_runtime_scalar"), + array(vec![real(1.0), real(2.0), real(3.0)], false), + test_span(), + ); + let err = analysis + .project_value_scalars(&product, &[3], &scope, 0, test_span()) + .expect_err("undeclared runtime scalar dimensions must remain unknown"); + + assert!(err.reason().contains("unknown dimensions"), "{err:?}"); } diff --git a/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/tests/vector_dot_projection.rs b/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/tests/vector_dot_projection.rs new file mode 100644 index 000000000..9672a2b57 --- /dev/null +++ b/crates/rumoca-phase-solve/src/lower/derivative_rhs/function_projection/tests/vector_dot_projection.rs @@ -0,0 +1,220 @@ +use super::*; + +fn evaluate_constant(expr: &rumoca_core::Expression) -> Option { + match expr { + rumoca_core::Expression::Literal { + value: Literal::Real(value), + .. + } => Some(*value), + rumoca_core::Expression::Literal { + value: Literal::Integer(value), + .. + } => Some(*value as f64), + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs, + rhs, + .. + } => Some(evaluate_constant(lhs)? + evaluate_constant(rhs)?), + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs, + rhs, + .. + } => Some(evaluate_constant(lhs)? * evaluate_constant(rhs)?), + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Sqrt, + args, + .. + } if args.len() == 1 => Some(evaluate_constant(&args[0])?.sqrt()), + _ => None, + } +} + +fn vector_norm_expr(name: &str) -> rumoca_core::Expression { + builtin( + rumoca_core::BuiltinFunction::Sqrt, + vec![binary( + rumoca_core::OpBinary::Mul, + local_var(name), + local_var(name), + test_span(), + )], + ) +} + +fn projected_scalar( + functions: Vec, + call_name: &str, + args: Vec, +) -> Result { + let mut dae_model = dae::Dae::default(); + for function in functions { + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + } + let call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new(call_name), + args, + is_constructor: false, + span: test_span(), + }; + let values = function_call_projected_scalars_with_owner( + &call, + &dae_model, + &IndexMap::new(), + test_span(), + )? + .expect("scalar function output should project"); + let [value] = values.as_slice() else { + panic!("expected one projected scalar, got {values:?}"); + }; + Ok(value.clone()) +} + +#[test] +fn array_output_bound_to_function_local_uses_complete_vector_dot_product() -> Result<(), LowerError> +{ + let mut source = rumoca_core::Function::new("My.vectorSource", test_span()); + source.outputs.push(function_param_with_dims("v", &[3])); + source.body.push(scalar_assignment( + "v", + array(vec![real(3.0), real(4.0), real(0.0)], false), + )); + + let mut norm = rumoca_core::Function::new("My.localVectorNorm", test_span()); + norm.locals.push(function_param_with_dims("v", &[3])); + norm.outputs.push(scalar_function_param("y")); + norm.body.push(scalar_assignment( + "v", + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("My.vectorSource"), + args: Vec::new(), + is_constructor: false, + span: test_span(), + }, + )); + norm.body + .push(scalar_assignment("y", vector_norm_expr("v"))); + + let value = projected_scalar(vec![source, norm], "My.localVectorNorm", Vec::new())?; + assert_eq!(evaluate_constant(&value), Some(5.0)); + Ok(()) +} + +#[test] +fn function_local_scalar_dot_product_controls_branch() -> Result<(), LowerError> { + let mut function = rumoca_core::Function::new("My.localScalarNorm", test_span()); + function.inputs.push(function_param_with_dims("v", &[3])); + function.locals.push(scalar_function_param("n")); + function.outputs.push(scalar_function_param("y")); + function + .body + .push(scalar_assignment("n", vector_norm_expr("v"))); + function.body.push(rumoca_core::Statement::If { + cond_blocks: vec![rumoca_core::StatementBlock { + cond: binary( + rumoca_core::OpBinary::Gt, + local_var("n"), + real(4.5), + test_span(), + ), + stmts: vec![scalar_assignment("y", real(5.0))], + }], + else_block: Some(vec![scalar_assignment("y", real(0.0))]), + span: test_span(), + }); + + let value = projected_scalar( + vec![function], + "My.localScalarNorm", + vec![array(vec![real(3.0), real(4.0), real(0.0)], false)], + )?; + assert_eq!(evaluate_constant(&value), Some(5.0)); + Ok(()) +} + +#[test] +fn previous_quat_shape_uses_all_four_dot_product_terms() -> Result<(), LowerError> { + let mut function = rumoca_core::Function::new("My.previousQuatNorm", test_span()); + function + .inputs + .push(function_param_with_dims("previous_quat", &[4])); + function.outputs.push(scalar_function_param("y")); + function + .body + .push(scalar_assignment("y", vector_norm_expr("previous_quat"))); + + let value = projected_scalar( + vec![function], + "My.previousQuatNorm", + vec![array( + vec![real(0.5), real(0.5), real(0.5), real(0.5)], + false, + )], + )?; + assert_eq!(evaluate_constant(&value), Some(1.0)); + Ok(()) +} + +#[test] +fn tangential_vector_norm_sums_three_squared_components() -> Result<(), LowerError> { + let mut function = rumoca_core::Function::new("My.tangentialNorm", test_span()); + function + .inputs + .push(function_param_with_dims("tangential", &[3])); + function.outputs.push(scalar_function_param("y")); + function + .body + .push(scalar_assignment("y", vector_norm_expr("tangential"))); + + let value = projected_scalar( + vec![function], + "My.tangentialNorm", + vec![array(vec![real(3.0), real(4.0), real(0.0)], false)], + )?; + assert_eq!(evaluate_constant(&value), Some(5.0)); + Ok(()) +} + +#[test] +fn vector_dot_projection_rejects_unequal_unknown_and_invalid_dimensions() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let analysis = FunctionProjectionAnalysis::new(&dae_model, &structural_bindings); + let expr = binary( + rumoca_core::OpBinary::Mul, + local_var("lhs"), + local_var("rhs"), + test_span(), + ); + + for (lhs_dims, rhs_dims, expected) in [ + (Some(vec![2]), Some(vec![3]), "incompatible"), + (None, Some(vec![3]), "unknown dimensions"), + (Some(vec![0]), Some(vec![0]), "positive"), + (Some(vec![-1]), Some(vec![-1]), "invalid dimension"), + ] { + let mut scope = FunctionProjectionScope::default(); + if let Some(dims) = lhs_dims { + scope.dims.insert("lhs".to_string(), dims); + } + if let Some(dims) = rhs_dims { + scope.dims.insert("rhs".to_string(), dims); + } + let ctx = projection_value_ctx(&[], 0, &scope, 0, test_span()); + let rumoca_core::Expression::Binary { lhs, rhs, .. } = &expr else { + unreachable!(); + }; + let err = analysis + .project_binary_value(&rumoca_core::OpBinary::Mul, lhs, rhs, &ctx) + .expect_err("invalid vector dot dimensions must fail closed"); + assert!( + err.reason().contains(expected), + "expected `{expected}` in `{}`", + err.reason() + ); + } +} diff --git a/crates/rumoca-phase-solve/src/lower/derivative_rhs/linear_parts.rs b/crates/rumoca-phase-solve/src/lower/derivative_rhs/linear_parts.rs index 33e94a727..2579ed901 100644 --- a/crates/rumoca-phase-solve/src/lower/derivative_rhs/linear_parts.rs +++ b/crates/rumoca-phase-solve/src/lower/derivative_rhs/linear_parts.rs @@ -180,16 +180,19 @@ pub(in crate::lower) fn derivative_linear_parts_any( divide_linear_parts(parts, rhs.as_ref().clone(), *span).map(Some) } _ if !expr.contains_der() => Ok(Some((IndexMap::new(), Some(expr.clone())))), - _ => derivative_terminal_linear_parts(expr, ctx.state_names, owner_span), + _ => derivative_terminal_linear_parts(expr, ctx, owner_span), } } fn derivative_terminal_linear_parts( expr: &rumoca_core::Expression, - state_names: &HashSet, + ctx: &DerivativeLinearCtx<'_>, owner_span: rumoca_core::Span, ) -> Result, LowerError> { - let Some((name, coeff)) = derivative_term_coefficient(expr, state_names, owner_span)? else { + if let Some(parts) = derivative_projected_function_call_linear_parts(expr, ctx, owner_span)? { + return Ok(Some(parts)); + } + let Some((name, coeff)) = derivative_term_coefficient(expr, ctx, owner_span)? else { return Ok(None); }; let span = linear_expr_or_owner_span(expr, owner_span)?; @@ -199,6 +202,45 @@ fn derivative_terminal_linear_parts( Ok(Some((coefficients, None))) } +fn derivative_projected_function_call_linear_parts( + expr: &rumoca_core::Expression, + ctx: &DerivativeLinearCtx<'_>, + owner_span: rumoca_core::Span, +) -> Result, LowerError> { + if !matches!( + expr, + rumoca_core::Expression::FunctionCall { + is_constructor: false, + .. + } + ) { + return Ok(None); + } + let span = linear_expr_or_owner_span(expr, owner_span)?; + let values = match function_call_projected_scalars_with_owner( + expr, + ctx.dae_model, + ctx.structural_bindings, + span, + ) { + Ok(values) => values, + Err(LowerError::MissingFunction { .. } | LowerError::MissingBinding { .. }) => { + return Ok(None); + } + Err(err) => return Err(err), + }; + let Some(values) = values else { + return Ok(None); + }; + let [projected] = values.as_slice() else { + return Ok(None); + }; + if projected == expr { + return Ok(None); + } + derivative_linear_parts_any(projected, ctx, span) +} + pub(in crate::lower) fn derivative_dot_product_linear_parts( lhs: &rumoca_core::Expression, rhs: &rumoca_core::Expression, @@ -212,14 +254,22 @@ pub(in crate::lower) fn derivative_dot_product_linear_parts( } let lhs_dims = match expression_result_dims(lhs, ctx.dae_model, ctx.structural_bindings, span) { Ok(dims) => dims, - Err(LowerError::MissingBinding { .. } | LowerError::Unsupported { .. }) => { + Err( + LowerError::MissingBinding { .. } + | LowerError::Unsupported { .. } + | LowerError::UnsupportedAt { .. }, + ) => { return Ok(None); } Err(err) => return Err(err), }; let rhs_dims = match expression_result_dims(rhs, ctx.dae_model, ctx.structural_bindings, span) { Ok(dims) => dims, - Err(LowerError::MissingBinding { .. } | LowerError::Unsupported { .. }) => { + Err( + LowerError::MissingBinding { .. } + | LowerError::Unsupported { .. } + | LowerError::UnsupportedAt { .. }, + ) => { return Ok(None); } Err(err) => return Err(err), @@ -552,10 +602,10 @@ pub(in crate::lower) fn rhs_without_remainder( pub(in crate::lower) fn derivative_term_coefficient( term: &rumoca_core::Expression, - state_names: &HashSet, + ctx: &DerivativeLinearCtx<'_>, owner_span: rumoca_core::Span, ) -> Result, LowerError> { - if let Some(name) = der_state_name(term, state_names)? { + if let Some(name) = der_state_name(term, ctx, owner_span)? { let span = linear_expr_or_owner_span(term, owner_span)?; return Ok(Some((name, one_expr_with_span(span)))); } @@ -566,12 +616,12 @@ pub(in crate::lower) fn derivative_term_coefficient( return Ok(None); } if !rhs.contains_der() - && let Some(name) = der_state_name(lhs, state_names)? + && let Some(name) = der_state_name(lhs, ctx, owner_span)? { return Ok(Some((name, rhs.as_ref().clone()))); } if !lhs.contains_der() - && let Some(name) = der_state_name(rhs, state_names)? + && let Some(name) = der_state_name(rhs, ctx, owner_span)? { return Ok(Some((name, lhs.as_ref().clone()))); } @@ -580,7 +630,8 @@ pub(in crate::lower) fn derivative_term_coefficient( pub(in crate::lower) fn der_state_name( expr: &rumoca_core::Expression, - state_names: &HashSet, + ctx: &DerivativeLinearCtx<'_>, + owner_span: rumoca_core::Span, ) -> Result, LowerError> { let rumoca_core::Expression::BuiltinCall { function, @@ -599,65 +650,24 @@ pub(in crate::lower) fn der_state_name( *span, )); }; - let key = binding_key_for_der_arg(arg) - .map_err(|err| err.with_fallback_span(arg.span().unwrap_or(*span)))?; - let Some(key) = key else { + let arg_span = arg + .span() + .unwrap_or_else(|| inherited_linear_span(*span, owner_span)); + let keys = + match derivative_arg_binding_keys(arg, ctx.dae_model, ctx.structural_bindings, arg_span) { + Ok(keys) => keys, + Err(LowerError::MissingBinding { .. }) => return Ok(None), + Err(err @ LowerError::UnsupportedAt { .. }) + if err.reason().contains("not compile-time bound") => + { + return Ok(None); + } + Err(err) => return Err(err.with_fallback_span(arg_span)), + }; + let [key] = keys.as_slice() else { return Ok(None); }; - Ok(state_names.contains(&key).then_some(key)) -} - -pub(in crate::lower) fn binding_key_for_der_arg( - expr: &rumoca_core::Expression, -) -> Result, LowerError> { - match expr { - rumoca_core::Expression::VarRef { - name, - subscripts, - span, - } => { - if subscripts.is_empty() { - return Ok(Some(name.as_str().to_string())); - } - let Some(indices) = der_arg_static_subscript_indices(subscripts, *span)? else { - return Ok(None); - }; - Ok(Some(dae::format_subscript_key(name.as_str(), &indices))) - } - rumoca_core::Expression::Index { - base, - subscripts, - span, - } => { - let Some(base_key) = binding_key_for_der_arg(base)? else { - return Ok(None); - }; - let Some(indices) = der_arg_static_subscript_indices(subscripts, *span)? else { - return Ok(None); - }; - Ok(Some(dae::format_subscript_key(&base_key, &indices))) - } - _ => Ok(None), - } -} - -fn der_arg_static_subscript_indices( - subscripts: &[rumoca_core::Subscript], - fallback_span: rumoca_core::Span, -) -> Result>, LowerError> { - super::super::static_subscript_indices_with_owner(subscripts, fallback_span) - .map_err(|err| err.with_fallback_span(der_arg_subscript_span(subscripts, fallback_span))) -} - -fn der_arg_subscript_span( - subscripts: &[rumoca_core::Subscript], - fallback_span: rumoca_core::Span, -) -> rumoca_core::Span { - subscripts - .iter() - .map(rumoca_core::Subscript::span) - .find(|span| !span.is_dummy()) - .unwrap_or(fallback_span) + Ok(ctx.state_names.contains(key).then_some(key.clone())) } pub(in crate::lower) fn split_subtraction( @@ -867,6 +877,31 @@ mod tests { } } + fn dae_with_state(name: &str, dims: &[i64]) -> dae::Dae { + let mut dae_model = dae::Dae::default(); + dae_model.variables.states.insert( + rumoca_core::VarName::new(name), + dae::Variable { + name: rumoca_core::VarName::new(name), + dims: dims.to_vec(), + ..dae::Variable::empty_with_span(span(0, 1)) + }, + ); + dae_model + } + + fn derivative_ctx<'a>( + state_names: &'a HashSet, + dae_model: &'a dae::Dae, + structural_bindings: &'a IndexMap, + ) -> DerivativeLinearCtx<'a> { + DerivativeLinearCtx { + state_names, + dae_model, + structural_bindings, + } + } + #[test] fn rhs_without_remainder_preserves_rhs_span() { let rhs_span = span(10, 20); @@ -941,14 +976,22 @@ mod tests { #[test] fn der_state_name_declines_non_state_derivative() -> Result<(), LowerError> { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let state_names = HashSet::from(["x".to_string()]); + let ctx = derivative_ctx(&state_names, &dae_model, &structural_bindings); let expr = der(var_ref("p", Vec::new(), span(1, 2)), span(0, 3)); - assert!(der_state_name(&expr, &HashSet::from(["x".to_string()]))?.is_none()); + assert!(der_state_name(&expr, &ctx, span(0, 3))?.is_none()); Ok(()) } #[test] fn der_state_name_reports_missing_argument_with_call_span() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let state_names = HashSet::from(["x".to_string()]); + let ctx = derivative_ctx(&state_names, &dae_model, &structural_bindings); let call_span = span(0, 5); let expr = rumoca_core::Expression::BuiltinCall { function: rumoca_core::BuiltinFunction::Der, @@ -956,21 +999,25 @@ mod tests { span: call_span, }; - let err = der_state_name(&expr, &HashSet::from(["x".to_string()])) - .expect_err("malformed der() call should fail"); + let err = + der_state_name(&expr, &ctx, call_span).expect_err("malformed der() call should fail"); assert_eq!(err.source_span(), Some(call_span)); assert!(err.reason().contains("der() call has no argument")); } #[test] fn der_state_name_missing_argument_dummy_span_stays_unspanned() { + let dae_model = dae::Dae::default(); + let structural_bindings = IndexMap::new(); + let state_names = HashSet::from(["x".to_string()]); + let ctx = derivative_ctx(&state_names, &dae_model, &structural_bindings); let expr = rumoca_core::Expression::BuiltinCall { function: rumoca_core::BuiltinFunction::Der, args: Vec::new(), span: unspanned_linear_parts_test_span(), }; - let err = der_state_name(&expr, &HashSet::from(["x".to_string()])) + let err = der_state_name(&expr, &ctx, unspanned_linear_parts_test_span()) .expect_err("malformed synthetic der() call should fail"); assert!( matches!(err, LowerError::UnspannedContractViolation { .. }), @@ -981,6 +1028,10 @@ mod tests { #[test] fn der_state_name_bubbles_invalid_static_subscript_span() { + let dae_model = dae_with_state("x", &[1]); + let structural_bindings = IndexMap::new(); + let state_names = HashSet::from(["x[1]".to_string()]); + let ctx = derivative_ctx(&state_names, &dae_model, &structural_bindings); let subscript_span = span(3, 4); let expr = der( var_ref( @@ -994,14 +1045,46 @@ mod tests { span(0, 5), ); - let err = der_state_name(&expr, &HashSet::from(["x[1]".to_string()])) + let err = der_state_name(&expr, &ctx, span(0, 5)) .expect_err("invalid der() subscript should fail"); assert_eq!(err.source_span(), Some(subscript_span)); - assert!(err.reason().contains("non-positive subscript")); + assert!(err.reason().contains("subscript"), "{err:?}"); + } + + #[test] + fn der_state_name_declines_slice_subscript() -> Result<(), LowerError> { + let dae_model = dae_with_state("x", &[1, 1]); + let structural_bindings = IndexMap::new(); + let state_names = HashSet::from(["x[1,1]".to_string()]); + let ctx = derivative_ctx(&state_names, &dae_model, &structural_bindings); + let expr = der( + var_ref( + "x", + vec![ + rumoca_core::Subscript::Index { + value: 1, + span: span(3, 4), + }, + rumoca_core::Subscript::Colon { span: span(5, 6) }, + ], + span(1, 6), + ), + span(0, 7), + ); + + assert_eq!( + der_state_name(&expr, &ctx, span(0, 7))?, + Some("x[1,1]".to_string()) + ); + Ok(()) } #[test] fn der_state_name_declines_dynamic_subscript() -> Result<(), LowerError> { + let dae_model = dae_with_state("x", &[1]); + let structural_bindings = IndexMap::new(); + let state_names = HashSet::from(["x[1]".to_string()]); + let ctx = derivative_ctx(&state_names, &dae_model, &structural_bindings); let subscript_span = span(3, 4); let expr = der( var_ref( @@ -1015,7 +1098,56 @@ mod tests { span(0, 5), ); - assert!(der_state_name(&expr, &HashSet::from(["x[1]".to_string()]))?.is_none()); + assert!(der_state_name(&expr, &ctx, span(0, 5))?.is_none()); + Ok(()) + } + + #[test] + fn der_state_name_resolves_structural_subscript_expression() -> Result<(), LowerError> { + let dae_model = dae_with_state("x", &[3]); + let mut structural_bindings = IndexMap::new(); + structural_bindings.insert("nr".to_string(), 1.0); + structural_bindings.insert("i".to_string(), 1.0); + let state_names = HashSet::from(["x[2]".to_string()]); + let ctx = derivative_ctx(&state_names, &dae_model, &structural_bindings); + let subscript_expr = rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: Box::new(var_ref("nr", Vec::new(), span(3, 5))), + rhs: Box::new(rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs: Box::new(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(2), + span: span(6, 7), + }), + rhs: Box::new(var_ref("i", Vec::new(), span(8, 9))), + span: span(6, 9), + }), + rhs: Box::new(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(1), + span: span(10, 11), + }), + span: span(6, 11), + }), + span: span(3, 11), + }; + let expr = der( + var_ref( + "x", + vec![rumoca_core::Subscript::Expr { + expr: Box::new(subscript_expr), + span: span(3, 11), + }], + span(1, 11), + ), + span(0, 12), + ); + + assert_eq!( + der_state_name(&expr, &ctx, span(0, 12))?, + Some("x[2]".to_string()) + ); Ok(()) } } diff --git a/crates/rumoca-phase-solve/src/lower/derivative_rhs/projection.rs b/crates/rumoca-phase-solve/src/lower/derivative_rhs/projection.rs index 7c1b7d5f8..59605f028 100644 --- a/crates/rumoca-phase-solve/src/lower/derivative_rhs/projection.rs +++ b/crates/rumoca-phase-solve/src/lower/derivative_rhs/projection.rs @@ -3,8 +3,10 @@ // proof and row assembly into focused projection submodules. use crate::lower::{ - function_calls::external_table_intrinsic_kind, - helpers::{format_i64_dims, format_usize_dims, positive_i64_index}, + function_calls::{external_table_intrinsic_kind, split_named_and_positional_call_args}, + helpers::{ + format_i64_dims, format_usize_dims, is_stream_passthrough_intrinsic, positive_i64_index, + }, unsupported_at, }; use crate::projection_suffix::parse_output_projection_suffix; @@ -824,13 +826,157 @@ fn project_function_call_scalar( ctx: &ProjectionContext<'_>, ) -> Result, LowerError> { let span = projection_expr_or_owner_span(expr, ctx.owner_span)?; + if let rumoca_core::Expression::FunctionCall { name, args, .. } = expr + && is_stream_passthrough_intrinsic(name.as_str()) + { + let Some(arg) = args.first() else { + return Ok(None); + }; + return project_operand_scalar_ctx(arg, ctx, span); + } let values = function_call_projected_scalars_with_owner( expr, ctx.dae_model, ctx.structural_bindings, span, )?; - Ok(values.and_then(|values| values.get(ctx.flat_index).cloned())) + if let Some(value) = values.and_then(|values| values.get(ctx.flat_index).cloned()) { + return Ok(Some(value)); + } + project_vectorized_scalar_function_call(expr, ctx, span) +} + +fn project_vectorized_scalar_function_call( + expr: &rumoca_core::Expression, + ctx: &ProjectionContext<'_>, + span: rumoca_core::Span, +) -> Result, LowerError> { + let rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor: false, + .. + } = expr + else { + return Ok(None); + }; + let Some(function) = ctx.dae_model.symbols.functions.get(name.var_name()) else { + return Ok(None); + }; + let [output] = function.outputs.as_slice() else { + return Ok(None); + }; + if !output.dims.is_empty() { + return Ok(None); + } + let (named_args, positional_args) = split_named_and_positional_call_args(name.as_str(), args)?; + let lane_dims = + checked_usize_dims_to_i64(ctx.dims, "vectorized function lane dimension", span)?; + let lane_indices = + dae::flat_index_to_subscripts(&lane_dims, ctx.flat_index).ok_or_else(|| { + LowerError::contract_violation( + format!( + "vectorized scalar function flat index {} is outside lane dimensions {}", + ctx.flat_index, + format_usize_dims(ctx.dims) + ), + span, + ) + })?; + let mut projected_args = derivative_vec_with_capacity( + function.inputs.len(), + "vectorized scalar function projected argument count", + span, + )?; + let mut positional_idx = 0usize; + let mut projected_any = false; + for input in &function.inputs { + let actual = if let Some(actual) = named_args.get(input.name.as_str()) { + *actual + } else if let Some(actual) = positional_args.get(positional_idx).copied() { + positional_idx += 1; + actual + } else if let Some(default) = input.default.as_ref() { + default + } else { + return Ok(None); + }; + let Some((projected, did_project)) = + project_vectorized_function_actual(actual, input.dims.len(), &lane_indices, ctx, span)? + else { + return Ok(None); + }; + projected_any |= did_project; + projected_args.push(projected); + } + if !projected_any { + return Ok(None); + } + Ok(Some(rumoca_core::Expression::FunctionCall { + name: name.clone(), + args: projected_args, + is_constructor: false, + span, + })) +} + +fn project_vectorized_function_actual( + actual: &rumoca_core::Expression, + formal_rank: usize, + lane_indices: &[usize], + ctx: &ProjectionContext<'_>, + span: rumoca_core::Span, +) -> Result, LowerError> { + let actual_dims = expression_result_dims(actual, ctx.dae_model, ctx.structural_bindings, span)?; + if actual_dims.is_empty() || actual_dims.len() == formal_rank { + return Ok(Some((actual.clone(), false))); + } + if formal_rank == 0 && actual_dims == ctx.dims { + let Some(projected) = project_expression_scalar_with_owner( + actual, + ctx.dims, + ctx.flat_index, + ctx.dae_model, + ctx.structural_bindings, + span, + )? + else { + return Ok(None); + }; + return Ok(Some((projected, true))); + } + if actual_dims.len() != ctx.dims.len() + formal_rank + || actual_dims[..ctx.dims.len()] != *ctx.dims + { + return Ok(None); + } + let mut subscripts = derivative_vec_with_capacity( + actual_dims.len(), + "vectorized function argument subscript count", + span, + )?; + for index in lane_indices { + let index = checked_usize_to_i64(*index, "vectorized function lane index", span)?; + subscripts.push(rumoca_core::Subscript::try_generated_index( + index, + span, + "vectorized function lane index", + )?); + } + for _ in 0..formal_rank { + subscripts.push(rumoca_core::Subscript::try_generated_colon( + span, + "vectorized function formal slice", + )?); + } + Ok(Some(( + rumoca_core::Expression::Index { + base: Box::new(actual.clone()), + subscripts, + span, + }, + true, + ))) } fn project_cross_scalar( @@ -1430,11 +1576,14 @@ pub(in crate::lower) fn builtin_size_args_dims( let span = projection_first_expr_or_owner_span(args, owner_span)?; let mut dims = derivative_vec_with_capacity(args.len(), "size() dimension count", span)?; for arg in args { - dims.push(super::super::compile_time_index_expr_with_owner( - arg, - structural_bindings, - span, - )?); + dims.push( + super::super::compile_time_non_negative_dimension_expr_with_owner( + arg, + structural_bindings, + span, + "builtin array dimension", + )?, + ); } Ok(dims) } @@ -1466,6 +1615,9 @@ pub(in crate::lower) fn expression_dims_for_subscripted_binding( } if subscripts.is_empty() { + if scalarized_variable_name_is_declared(dae_model, base)? { + return Ok(Vec::new()); + } if let Some(dims) = scalarized_child_dims(dae_model, base, fallback_span)? { return Ok(dims); } @@ -1518,6 +1670,9 @@ fn required_declared_function_output_dims( if external_table_intrinsic_kind(requested).is_some() { return Ok(Vec::new()); } + if is_stream_passthrough_intrinsic(requested) { + return Ok(Vec::new()); + } Err(LowerError::MissingFunction { name: name.to_string(), } diff --git a/crates/rumoca-phase-solve/src/lower/derivative_rhs/projection/binding_expressions.rs b/crates/rumoca-phase-solve/src/lower/derivative_rhs/projection/binding_expressions.rs index b72405c12..f07af8001 100644 --- a/crates/rumoca-phase-solve/src/lower/derivative_rhs/projection/binding_expressions.rs +++ b/crates/rumoca-phase-solve/src/lower/derivative_rhs/projection/binding_expressions.rs @@ -52,77 +52,24 @@ fn binding_expressions_for_subscripted_reference( structural_bindings: &IndexMap, ) -> Result, LowerError> { if let Some(dims) = variable_dims(dae_model, name.as_str())? { - let selections = if subscripts.is_empty() { - let mut selections = derivative_vec_with_capacity( - dims.len(), - "derivative slice selection dimension count", - span, - )?; - for dim in dims { - selections.push(one_based_index_range( - dim, - "derivative full-slice index count", - span, - )?); - } - selections - } else { - slice_selections(subscripts, &dims, structural_bindings, span)? - }; - let expression_count = slice_selection_count(&selections, span)?; - let mut expressions = derivative_vec_with_capacity( - expression_count, - "derivative slice expression count", - span, - )?; - let mut current = - derivative_vec_with_capacity(selections.len(), "derivative slice index depth", span)?; - collect_slice_reference_expressions( + return collect_regular_slice_binding_expressions( name, + &dims, + subscripts, + structural_bindings, span, - &selections, - 0, - &mut current, - &mut expressions, - )?; - return Ok(expressions); + ); } if let Some(dims) = scalarized_child_dims(dae_model, name.as_str(), span)? { - let selections = if subscripts.is_empty() { - let mut selections = derivative_vec_with_capacity( - dims.len(), - "scalarized derivative slice selection dimension count", - span, - )?; - for dim in dims { - selections.push(one_based_index_range( - dim, - "scalarized derivative full-slice index count", - span, - )?); - } - selections - } else { - slice_selections(subscripts, &dims, structural_bindings, span)? - }; - let key_count = slice_selection_count(&selections, span)?; - let mut keys = - derivative_vec_with_capacity(key_count, "scalarized derivative slice key count", span)?; - let mut current = - derivative_vec_with_capacity(selections.len(), "derivative slice key depth", span)?; - collect_slice_keys(name.as_str(), &selections, span, 0, &mut current, &mut keys)?; - let mut expressions = derivative_vec_with_capacity( - keys.len(), - "scalarized derivative slice expression count", + return collect_scalarized_slice_binding_expressions( + dae_model, + name, + &dims, + subscripts, + structural_bindings, span, - )?; - for key in keys { - let variable = variable_by_name(dae_model, &key) - .ok_or_else(|| LowerError::MissingBinding { name: key.clone() })?; - expressions.push(dae_variable_ref_expr(&key, variable, span, Vec::new())?); - } - return Ok(expressions); + ); } if subscripts.is_empty() { @@ -138,6 +85,17 @@ fn binding_expressions_for_subscripted_reference( ); } + if let Some(variable) = variable_by_name(dae_model, name.as_str()) + && variable.dims.is_empty() + && subscript_indices_are_all_singleton(subscripts, structural_bindings, span)? + { + return single_expression_vec( + dae_variable_ref_expr(name.as_str(), variable, span, Vec::new())?, + "derivative scalar singleton binding expression count", + span, + ); + } + let indices = binding_subscript_indices(name, subscripts, structural_bindings, span)?; let scalarized_key = dae::format_subscript_key(name.as_str(), &indices); let variable = @@ -151,6 +109,99 @@ fn binding_expressions_for_subscripted_reference( ) } +fn collect_regular_slice_binding_expressions( + name: &rumoca_core::Reference, + dims: &[usize], + subscripts: &[rumoca_core::Subscript], + structural_bindings: &IndexMap, + span: rumoca_core::Span, +) -> Result, LowerError> { + let selections = binding_slice_selections( + dims, + subscripts, + structural_bindings, + "derivative slice selection dimension count", + "derivative full-slice index count", + span, + )?; + let expression_count = slice_selection_count(&selections, span)?; + let mut expressions = + derivative_vec_with_capacity(expression_count, "derivative slice expression count", span)?; + let mut current = + derivative_vec_with_capacity(selections.len(), "derivative slice index depth", span)?; + collect_slice_reference_expressions( + name, + span, + &selections, + 0, + &mut current, + &mut expressions, + )?; + Ok(expressions) +} + +fn collect_scalarized_slice_binding_expressions( + dae_model: &dae::Dae, + name: &rumoca_core::Reference, + dims: &[usize], + subscripts: &[rumoca_core::Subscript], + structural_bindings: &IndexMap, + span: rumoca_core::Span, +) -> Result, LowerError> { + let selections = binding_slice_selections( + dims, + subscripts, + structural_bindings, + "scalarized derivative slice selection dimension count", + "scalarized derivative full-slice index count", + span, + )?; + let key_count = slice_selection_count(&selections, span)?; + let mut keys = + derivative_vec_with_capacity(key_count, "scalarized derivative slice key count", span)?; + let mut current = + derivative_vec_with_capacity(selections.len(), "derivative slice key depth", span)?; + collect_slice_keys(name.as_str(), &selections, span, 0, &mut current, &mut keys)?; + let mut expressions = derivative_vec_with_capacity( + keys.len(), + "scalarized derivative slice expression count", + span, + )?; + for key in keys { + let variable = variable_by_name(dae_model, &key) + .ok_or_else(|| LowerError::MissingBinding { name: key.clone() })?; + expressions.push(dae_variable_ref_expr(&key, variable, span, Vec::new())?); + } + Ok(expressions) +} + +fn binding_slice_selections( + dims: &[usize], + subscripts: &[rumoca_core::Subscript], + structural_bindings: &IndexMap, + capacity_context: &'static str, + range_context: &'static str, + span: rumoca_core::Span, +) -> Result>, LowerError> { + if !subscripts.is_empty() { + return slice_selections(subscripts, dims, structural_bindings, span); + } + let mut selections = derivative_vec_with_capacity(dims.len(), capacity_context, span)?; + for dim in dims { + selections.push(one_based_index_range(*dim, range_context, span)?); + } + Ok(selections) +} + +fn subscript_indices_are_all_singleton( + subscripts: &[rumoca_core::Subscript], + structural_bindings: &IndexMap, + span: rumoca_core::Span, +) -> Result { + let indices = compile_time_subscript_indices_with_owner(subscripts, structural_bindings, span)?; + Ok(indices.iter().all(|index| *index == 1)) +} + fn binding_subscript_indices( name: &rumoca_core::Reference, subscripts: &[rumoca_core::Subscript], @@ -265,6 +316,9 @@ pub(in crate::lower) fn binding_keys_for_subscripted_name( if variable_by_name(dae_model, base).is_some() { return Ok(vec![base.to_string()]); } + if scalarized_variable_name_is_declared(dae_model, base)? { + return Ok(vec![base.to_string()]); + } if let Some(dims) = scalarized_child_dims(dae_model, base, fallback_span)? { return scalar_keys_for_dims(base, &dims, fallback_span); } @@ -291,6 +345,12 @@ pub(in crate::lower) fn binding_keys_for_subscripted_name( collect_slice_keys(base, &selections, span, 0, &mut current, &mut keys)?; return Ok(keys); } + if (variable_by_name(dae_model, base).is_some() + || scalarized_variable_name_is_declared(dae_model, base)?) + && subscript_indices_are_all_singleton(subscripts, structural_bindings, span)? + { + return Ok(vec![base.to_string()]); + } let scalarized_key = scalarized_binding_key(base, subscripts, structural_bindings, fallback_span)?; if variable_by_name(dae_model, &scalarized_key).is_some() { @@ -313,6 +373,26 @@ pub(in crate::lower) fn binding_keys_for_subscripted_name( Ok(keys) } +pub(in crate::lower) fn scalarized_variable_name_is_declared( + dae_model: &dae::Dae, + key: &str, +) -> Result { + let Some(scalar) = rumoca_core::parse_scalar_name(key) else { + return Ok(false); + }; + let Some(dims) = variable_dims(dae_model, scalar.base)? else { + return Ok(false); + }; + if scalar.indices.len() != dims.len() { + return Ok(false); + } + Ok(scalar + .indices + .iter() + .zip(dims.iter()) + .all(|(index, dim)| usize::try_from(*index).is_ok_and(|index| index > 0 && index <= *dim))) +} + pub(in crate::lower) fn scalarized_binding_key( base: &str, subscripts: &[rumoca_core::Subscript], @@ -587,6 +667,9 @@ pub(in crate::lower) fn result_dims_for_subscripts( structural_bindings: &IndexMap, fallback_span: rumoca_core::Span, ) -> Result, LowerError> { + if dims.is_empty() && singleton_scalar_projection_subscripts(subscripts) { + return Ok(Vec::new()); + } if subscripts.is_empty() { let mut copied = derivative_vec_with_capacity( dims.len(), @@ -605,3 +688,18 @@ pub(in crate::lower) fn result_dims_for_subscripts( } Ok(result_dims) } + +fn singleton_scalar_projection_subscripts(subscripts: &[rumoca_core::Subscript]) -> bool { + !subscripts.is_empty() + && subscripts.iter().all(|subscript| match subscript { + rumoca_core::Subscript::Index { value, .. } => *value == 1, + rumoca_core::Subscript::Expr { expr, .. } => matches!( + expr.as_ref(), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(1), + .. + } + ), + rumoca_core::Subscript::Colon { .. } => false, + }) +} diff --git a/crates/rumoca-phase-solve/src/lower/derivative_rhs/projection/slices.rs b/crates/rumoca-phase-solve/src/lower/derivative_rhs/projection/slices.rs index d195eae08..5f88bb66b 100644 --- a/crates/rumoca-phase-solve/src/lower/derivative_rhs/projection/slices.rs +++ b/crates/rumoca-phase-solve/src/lower/derivative_rhs/projection/slices.rs @@ -34,6 +34,12 @@ pub(in crate::lower) fn slice_selections( owner_span: rumoca_core::Span, ) -> Result>, LowerError> { let span = subscript_list_span_or_owner(subscripts, owner_span)?; + let subscripts = normalize_overspecified_scalar_slice_subscripts( + subscripts, + dims, + structural_bindings, + span, + )?; if subscripts.len() > dims.len() { let span = subscripts .get(dims.len()) @@ -64,6 +70,29 @@ pub(in crate::lower) fn slice_selections( Ok(selections) } +fn normalize_overspecified_scalar_slice_subscripts<'a>( + subscripts: &'a [rumoca_core::Subscript], + dims: &[usize], + structural_bindings: &IndexMap, + span: rumoca_core::Span, +) -> Result<&'a [rumoca_core::Subscript], LowerError> { + if subscripts.len() <= dims.len() { + return Ok(subscripts); + } + let (declared_subscripts, extra_subscripts) = subscripts.split_at(dims.len()); + for (subscript, dim) in declared_subscripts.iter().zip(dims.iter().copied()) { + if slice_subscript_indices(subscript, dim, structural_bindings, span)?.len() != 1 { + return Ok(subscripts); + } + } + for subscript in extra_subscripts { + if compile_time_subscript_index_with_owner(subscript, structural_bindings, span)? != 1 { + return Ok(subscripts); + } + } + Ok(declared_subscripts) +} + pub(in crate::lower) fn slice_subscript_indices( subscript: &rumoca_core::Subscript, dim: usize, diff --git a/crates/rumoca-phase-solve/src/lower/derivative_rhs/projection/tests.rs b/crates/rumoca-phase-solve/src/lower/derivative_rhs/projection/tests.rs index 997b034c0..41504b9c3 100644 --- a/crates/rumoca-phase-solve/src/lower/derivative_rhs/projection/tests.rs +++ b/crates/rumoca-phase-solve/src/lower/derivative_rhs/projection/tests.rs @@ -112,6 +112,156 @@ fn derivative_slice_subscript_bounds_error_reports_subscript_span() { ); } +#[test] +fn derivative_slice_accepts_trailing_singleton_after_scalar_selection() { + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("derivative_slice_singleton.mo"), + 8, + 16, + ); + let selections = slice_selections( + &[ + rumoca_core::Subscript::index(3, span), + rumoca_core::Subscript::index(1, span), + ], + &[3], + &IndexMap::new(), + span, + ) + .expect("trailing singleton after scalar vector selection should be consumed"); + + assert_eq!(selections, vec![vec![3]]); +} + +#[test] +fn derivative_slice_rejects_non_singleton_extra_subscript() { + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("derivative_slice_bad_extra.mo"), + 8, + 16, + ); + let err = slice_selections( + &[ + rumoca_core::Subscript::index(3, span), + rumoca_core::Subscript::index(2, span), + ], + &[3], + &IndexMap::new(), + span, + ) + .expect_err("non-singleton extra subscript should remain invalid"); + + assert_eq!( + err.reason(), + "array derivative slice has more subscripts than dimensions" + ); +} + +#[test] +fn binding_expression_consumes_singleton_subscript_on_scalarized_variable() -> Result<(), LowerError> +{ + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("derivative_scalarized_singleton_binding.mo"), + 8, + 16, + ); + let mut dae_model = dae::Dae::new(); + dae_model + .variables + .outputs + .insert(rumoca_core::VarName::new("y[2]"), { + let mut variable = dae::Variable { + name: rumoca_core::VarName::new("y[2]"), + ..rumoca_ir_dae::Variable::empty_with_span(span) + }; + variable.origin = dae::VariableOrigin::Generated; + variable + }); + let expr = rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("y[2]"), + subscripts: vec![rumoca_core::Subscript::index(1, span)], + span, + }; + + let bindings = expression_binding_expressions(&expr, &dae_model, &IndexMap::new(), span)? + .expect("scalarized variable should resolve"); + + assert_eq!(bindings.len(), 1); + assert!(matches!( + &bindings[0], + rumoca_core::Expression::VarRef { name, subscripts, .. } + if name.as_str() == "y[2]" && subscripts.is_empty() + )); + Ok(()) +} + +#[test] +fn binding_keys_consume_singleton_subscript_on_scalarized_variable() -> Result<(), LowerError> { + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("derivative_scalarized_singleton_key.mo"), + 8, + 16, + ); + let mut dae_model = dae::Dae::new(); + dae_model + .variables + .outputs + .insert(rumoca_core::VarName::new("y[2]"), { + let mut variable = dae::Variable { + name: rumoca_core::VarName::new("y[2]"), + ..rumoca_ir_dae::Variable::empty_with_span(span) + }; + variable.origin = dae::VariableOrigin::Generated; + variable + }); + let keys = binding_keys_for_subscripted_name( + "y[2]", + &[rumoca_core::Subscript::index(1, span)], + &dae_model, + &IndexMap::new(), + span, + )?; + + assert_eq!(keys, vec!["y[2]"]); + Ok(()) +} + +#[test] +fn binding_keys_consume_singleton_subscript_after_vector_scalar_selection() -> Result<(), LowerError> +{ + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("derivative_vector_singleton_key.mo"), + 8, + 16, + ); + let mut dae_model = dae::Dae::new(); + dae_model + .variables + .outputs + .insert(rumoca_core::VarName::new("y"), { + let mut variable = dae::Variable { + name: rumoca_core::VarName::new("y"), + dims: vec![3], + ..rumoca_ir_dae::Variable::empty_with_span(span) + }; + variable.origin = dae::VariableOrigin::Generated; + variable + }); + let keys = binding_keys_for_subscripted_name( + "y", + &[ + rumoca_core::Subscript::index(2, span), + rumoca_core::Subscript::index(1, span), + ], + &dae_model, + &IndexMap::new(), + span, + )?; + + assert_eq!(keys, vec!["y[2]"]); + Ok(()) +} + #[test] fn scalar_binding_indexed_dimension_error_reports_subscript_span() -> Result<(), String> { let subscript_span = rumoca_core::Span::from_offsets( @@ -308,6 +458,52 @@ fn project_expression_scalars_rejects_der_without_argument() { ); } +#[test] +fn project_expression_scalars_passes_through_stream_operator() { + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name( + "phase_solve_lower_derivative_rhs_projection_tests_source_29.mo", + ), + 3, + 18, + ); + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("inStream").into(), + args: vec![rumoca_core::Expression::Array { + elements: vec![ + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(1.0), + span, + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(2.0), + span, + }, + ], + is_matrix: false, + span, + }], + is_constructor: false, + span, + }; + + let scalars = project_expression_scalars(&expr, &[2], &dae::Dae::new(), &IndexMap::new(), span) + .expect("stream passthrough should project its argument") + .expect("stream passthrough argument should have projected scalars"); + let values = scalars + .iter() + .map(|expr| match expr { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(value), + .. + } => *value, + other => panic!("expected scalar literal, got {other:?}"), + }) + .collect::>(); + + assert_eq!(values, vec![1.0, 2.0]); +} + #[test] fn project_expression_scalars_rejects_fill_without_value_argument() { let span = rumoca_core::Span::from_offsets( @@ -389,6 +585,32 @@ fn project_expression_scalars_rejects_zeros_without_dimension_argument() { ); } +#[test] +fn project_expression_scalars_accepts_zero_length_zeros_dimension() -> Result<(), LowerError> { + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name( + "phase_solve_lower_derivative_rhs_projection_tests_source_19.mo", + ), + 9, + 17, + ); + let expr = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Zeros, + args: vec![rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(0), + span, + }], + span, + }; + + let projected = + project_expression_scalars(&expr, &[0], &dae::Dae::new(), &IndexMap::new(), span)? + .expect("zero-length zeros() should project to an empty scalar list"); + + assert!(projected.is_empty()); + Ok(()) +} + #[test] fn project_expression_scalars_rejects_ones_without_dimension_argument() { let span = rumoca_core::Span::from_offsets( diff --git a/crates/rumoca-phase-solve/src/lower/derivative_rhs/projection_selection.rs b/crates/rumoca-phase-solve/src/lower/derivative_rhs/projection_selection.rs index 91667eb16..45f0bb461 100644 --- a/crates/rumoca-phase-solve/src/lower/derivative_rhs/projection_selection.rs +++ b/crates/rumoca-phase-solve/src/lower/derivative_rhs/projection_selection.rs @@ -233,6 +233,9 @@ pub(super) fn static_subscript_indices( let Some(index) = (match subscript { rumoca_core::Subscript::Index { value, .. } if *value > 0 => Some(*value), rumoca_core::Subscript::Expr { expr, .. } => { + if plain_runtime_selector(expr, analysis, scope) { + return Ok(None); + } let expr = analysis.substitute(expr, scope)?; let value = match analysis.compile_time_scalar(&expr) { Some(value) => value, @@ -249,6 +252,32 @@ pub(super) fn static_subscript_indices( Ok(Some(indices)) } +fn plain_runtime_selector( + expr: &rumoca_core::Expression, + analysis: &FunctionProjectionAnalysis<'_>, + scope: &FunctionProjectionScope, +) -> bool { + let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = expr + else { + return false; + }; + if !subscripts.is_empty() { + return false; + } + if analysis.structural_bindings.contains_key(name.as_str()) { + return false; + } + let Some(replacement) = scope.full.get(name.as_str()) else { + return true; + }; + if analysis.compile_time_scalar(replacement).is_some() { + return false; + } + true +} + fn positive_i64_from_compile_time_scalar(value: f64) -> Option { let rounded = value.round(); if !rounded.is_finite() @@ -314,6 +343,9 @@ pub(super) fn subscript_selector_expr( span: *span, }), rumoca_core::Subscript::Expr { expr, .. } => { + if plain_runtime_selector(expr, analysis, scope) { + return Ok(*expr.clone()); + } Ok(analysis.scalar_assignment_value(expr, scope, depth + 1)?) } rumoca_core::Subscript::Colon { span } => Err(unsupported_at( @@ -323,12 +355,6 @@ pub(super) fn subscript_selector_expr( } } -pub(super) fn is_function_output_target(function: &rumoca_core::Function, target: &str) -> bool { - function.outputs.iter().any(|output| { - target == output.name || target.starts_with(format!("{}.", output.name).as_str()) - }) -} - pub(super) fn projection_scope_names( entry_scope: &FunctionProjectionScope, branch_scopes: &[FunctionProjectionScope], diff --git a/crates/rumoca-phase-solve/src/lower/derivative_rhs/row_projection.rs b/crates/rumoca-phase-solve/src/lower/derivative_rhs/row_projection.rs new file mode 100644 index 000000000..20e822b50 --- /dev/null +++ b/crates/rumoca-phase-solve/src/lower/derivative_rhs/row_projection.rs @@ -0,0 +1,109 @@ +use super::*; + +pub(super) fn project_derivative_row_expr( + builder: &LowerBuilder, + expr: &rumoca_core::Expression, + row_index: usize, + row_count: usize, + source_context_span: rumoca_core::Span, + scope: &Scope, +) -> Result { + let dims = builder.infer_expr_dims(expr, scope)?; + if dims.is_empty() { + return Ok(expr.clone()); + } + let span = expr.span().unwrap_or(source_context_span); + match expr { + rumoca_core::Expression::VarRef { + name, + subscripts, + span, + } if subscripts.is_empty() && dims.as_slice() == [row_count] => { + let index = checked_usize_to_i64( + row_index + 1, + "derivative row projection subscript", + source_context_span, + )?; + let subscript = rumoca_core::Subscript::try_generated_index( + index, + *span, + "derivative row projection subscript", + )?; + Ok(rumoca_core::Expression::VarRef { + name: name.clone(), + subscripts: vec![subscript], + span: *span, + }) + } + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Der, + args, + span, + } => { + let [arg] = args.as_slice() else { + return Ok(expr.clone()); + }; + Ok(rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Der, + args: vec![project_derivative_row_expr( + builder, + arg, + row_index, + row_count, + source_context_span, + scope, + )?], + span: *span, + }) + } + rumoca_core::Expression::Binary { op, lhs, rhs, span } => { + Ok(rumoca_core::Expression::Binary { + op: op.clone(), + lhs: Box::new(project_derivative_row_expr( + builder, + lhs, + row_index, + row_count, + source_context_span, + scope, + )?), + rhs: Box::new(project_derivative_row_expr( + builder, + rhs, + row_index, + row_count, + source_context_span, + scope, + )?), + span: *span, + }) + } + rumoca_core::Expression::Unary { op, rhs, span } => Ok(rumoca_core::Expression::Unary { + op: op.clone(), + rhs: Box::new(project_derivative_row_expr( + builder, + rhs, + row_index, + row_count, + source_context_span, + scope, + )?), + span: *span, + }), + _ if dims.as_slice() == [row_count] => { + let index = + checked_usize_to_i64(row_index + 1, "derivative row projection subscript", span)?; + let subscript = rumoca_core::Subscript::try_generated_index( + index, + span, + "derivative row projection subscript", + )?; + Ok(rumoca_core::Expression::Index { + base: Box::new(expr.clone()), + subscripts: vec![subscript], + span, + }) + } + _ => Ok(expr.clone()), + } +} diff --git a/crates/rumoca-phase-solve/src/lower/derivative_rhs/tests.rs b/crates/rumoca-phase-solve/src/lower/derivative_rhs/tests.rs index 19c44ebb5..6cc0dce2d 100644 --- a/crates/rumoca-phase-solve/src/lower/derivative_rhs/tests.rs +++ b/crates/rumoca-phase-solve/src/lower/derivative_rhs/tests.rs @@ -79,6 +79,16 @@ fn scalar_var(name: &str) -> dae::Variable { } } +fn source_scalar_var(name: &str) -> dae::Variable { + let span = derivative_rhs_test_span(); + let var_name = rumoca_core::VarName::new(name); + dae::Variable { + name: var_name.clone(), + component_ref: rumoca_core::component_reference_from_flat_name(&var_name, span), + ..dae::Variable::empty_with_span(span) + } +} + fn var_ref(name: &str) -> rumoca_core::Expression { rumoca_core::Expression::VarRef { name: rumoca_core::Reference::new(name), @@ -87,6 +97,105 @@ fn var_ref(name: &str) -> rumoca_core::Expression { } } +fn derivative_test_var(name: &str) -> rumoca_core::Expression { + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new(name), + subscripts: Vec::new(), + span: derivative_rhs_test_span(), + } +} + +fn derivative_test_derivative(expr: rumoca_core::Expression) -> rumoca_core::Expression { + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Der, + args: vec![expr], + span: derivative_rhs_test_span(), + } +} + +fn derivative_test_binary( + op: rumoca_core::OpBinary, + lhs: rumoca_core::Expression, + rhs: rumoca_core::Expression, +) -> rumoca_core::Expression { + rumoca_core::Expression::Binary { + op, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span: derivative_rhs_test_span(), + } +} + +#[test] +fn noncontiguous_coupled_derivative_states_keep_their_output_ownership() { + let mut dae_model = dae::Dae::default(); + for name in ["x", "y", "z"] { + dae_model + .variables + .states + .insert(rumoca_core::VarName::new(name), source_scalar_var(name)); + } + for name in ["u", "v", "w"] { + dae_model + .variables + .parameters + .insert(rumoca_core::VarName::new(name), source_scalar_var(name)); + } + let sub = |lhs, rhs| derivative_test_binary(rumoca_core::OpBinary::Sub, lhs, rhs); + let add = |lhs, rhs| derivative_test_binary(rumoca_core::OpBinary::Add, lhs, rhs); + dae_model.continuous.equations.extend([ + dae::Equation::residual( + sub( + add( + derivative_test_derivative(derivative_test_var("x")), + derivative_test_derivative(derivative_test_var("z")), + ), + derivative_test_var("u"), + ), + derivative_rhs_test_span(), + "x_z_coupling", + ), + dae::Equation::residual( + sub( + derivative_test_derivative(derivative_test_var("z")), + derivative_test_var("v"), + ), + derivative_rhs_test_span(), + "z_equation", + ), + dae::Equation::residual( + sub( + derivative_test_derivative(derivative_test_var("y")), + derivative_test_var("w"), + ), + derivative_rhs_test_span(), + "y_equation", + ), + ]); + + let layout = crate::build_var_layout(&dae_model).expect("test DAE layout should build"); + let block = lower_derivative_rhs(&dae_model, &layout) + .expect("noncontiguous coupled derivatives should lower"); + assert!(matches!( + block.nodes.first(), + Some(rumoca_ir_solve::ComputeNode::LinSolve { + n: 2, + output_indices, + .. + }) if output_indices == &[0, 2] + )); + assert_eq!( + block + .len() + .expect("derivative block output count should be valid"), + 3 + ); + let scalar = rumoca_eval_solve::to_scalar_program_block(&block) + .expect("derivative block should scalarize"); + + assert_eq!(scalar.output_indices, vec![0, 2, 1]); +} + #[test] fn expression_result_dims_rejects_missing_scalar_binding() { let dae_model = dae::Dae::default(); diff --git a/crates/rumoca-phase-solve/src/lower/emit.rs b/crates/rumoca-phase-solve/src/lower/emit.rs index b3435de39..8e1b7430f 100644 --- a/crates/rumoca-phase-solve/src/lower/emit.rs +++ b/crates/rumoca-phase-solve/src/lower/emit.rs @@ -1,4 +1,6 @@ -use rumoca_ir_solve::{BinaryOp, CompareOp, LinearOp, RandomGenerator, Reg, ScalarSlot, UnaryOp}; +use rumoca_ir_solve::{ + BinaryOp, CompareOp, ExternalFunctionKind, LinearOp, RandomGenerator, Reg, ScalarSlot, UnaryOp, +}; use super::cse::{SlotLoadKey, canonical_binary_key}; use super::{LowerBuilder, LowerError}; @@ -355,6 +357,35 @@ impl LowerBuilder<'_> { Ok(dst) } + pub(super) fn emit_external_call( + &mut self, + function: ExternalFunctionKind, + args: &[Reg], + output_index: usize, + span: rumoca_core::Span, + ) -> Result { + if args.len() > 8 { + return Err(LowerError::contract_violation( + format!( + "external call `{function:?}` has {} scalar argument registers; max supported is 8", + args.len() + ), + span, + )); + } + let dst = self.try_alloc_reg(span)?; + let mut op_args = [0; 8]; + op_args[..args.len()].copy_from_slice(args); + self.ops.push(LinearOp::ExternalCall { + dst, + function, + args: op_args, + arg_count: args.len(), + output_index, + }); + Ok(dst) + } + pub(super) fn try_pack_registers( &mut self, regs: &[Reg], diff --git a/crates/rumoca-phase-solve/src/lower/error.rs b/crates/rumoca-phase-solve/src/lower/error.rs index ccc1f3e58..e23ec13fc 100644 --- a/crates/rumoca-phase-solve/src/lower/error.rs +++ b/crates/rumoca-phase-solve/src/lower/error.rs @@ -37,6 +37,13 @@ pub enum LowerError { function: String, span: rumoca_core::Span, }, + /// Function-output projection encountered a runtime-dependent while loop. + /// The outer projection boundary declines so ordinary Solve-IR function + /// lowering can preserve the loop with conditional register updates. + DynamicWhileProjection { + function: String, + span: rumoca_core::Span, + }, /// A subscript whose value is only known at runtime, in a position that /// requires a compile-time index. Observation lowering declines on this. DynamicSubscript, @@ -94,6 +101,10 @@ impl std::fmt::Display for LowerError { f, "function `{function}` projection exceeded the inline node budget" ), + Self::DynamicWhileProjection { function, .. } => write!( + f, + "function `{function}` has a runtime-dependent while loop" + ), Self::DynamicSubscript | Self::ForRangeUnknownDimension { .. } | Self::DynamicBindingBase { .. } => { @@ -166,6 +177,9 @@ impl LowerError { Self::ProjectionBudgetExceeded { function, .. } => { format!("function `{function}` projection exceeded the inline node budget") } + Self::DynamicWhileProjection { function, .. } => { + format!("function `{function}` has a runtime-dependent while loop") + } Self::DynamicSubscript => "dynamic subscript expressions are unsupported".to_string(), Self::ForRangeUnknownDimension { name } => { format!("size() in for-loop range requires known dimension `{name}`") @@ -204,6 +218,9 @@ impl LowerError { Self::ProjectionBudgetExceeded { function, .. } => { format!("function `{function}` projection exceeded the inline node budget") } + Self::DynamicWhileProjection { function, .. } => { + format!("function `{function}` has a runtime-dependent while loop") + } Self::DynamicSubscript | Self::ForRangeUnknownDimension { .. } | Self::MissingActualArgument { .. } @@ -219,6 +236,7 @@ impl LowerError { Self::ContractViolation { span, .. } if !span.is_dummy() => Some(*span), Self::Scalarization { span, .. } => *span, Self::ProjectionBudgetExceeded { span, .. } if !span.is_dummy() => Some(*span), + Self::DynamicWhileProjection { span, .. } if !span.is_dummy() => Some(*span), Self::MissingActualArgument { span, .. } if !span.is_dummy() => Some(*span), Self::Spanned { source, span } => source .source_span() @@ -245,6 +263,18 @@ impl LowerError { self.projection_budget_exceeded_parts().is_some() } + /// True when projection must defer a dynamic while loop to ordinary + /// Solve-IR function lowering, looking through context/span wrappers. + pub fn is_dynamic_while_projection(&self) -> bool { + match self { + Self::DynamicWhileProjection { .. } => true, + Self::Spanned { source, .. } | Self::WithContext { source, .. } => { + source.is_dynamic_while_projection() + } + _ => false, + } + } + pub fn is_missing_binding(&self) -> bool { match self { Self::MissingBinding { .. } => true, @@ -255,6 +285,16 @@ impl LowerError { } } + pub fn is_missing_binding_or_function(&self) -> bool { + match self { + Self::MissingBinding { .. } | Self::MissingFunction { .. } => true, + Self::Spanned { source, .. } | Self::WithContext { source, .. } => { + source.is_missing_binding_or_function() + } + _ => false, + } + } + pub fn with_fallback_span(self, span: rumoca_core::Span) -> Self { if span.is_dummy() || self.source_span().is_some() { return self; diff --git a/crates/rumoca-phase-solve/src/lower/expression_rows.rs b/crates/rumoca-phase-solve/src/lower/expression_rows.rs index eeb0d8dee..9c8f7b3c9 100644 --- a/crates/rumoca-phase-solve/src/lower/expression_rows.rs +++ b/crates/rumoca-phase-solve/src/lower/expression_rows.rs @@ -10,12 +10,15 @@ use rumoca_ir_solve::{ use super::{ DirectAssignmentValue, IndexedBindingMap, LowerBuilder, LowerBuilderMetadata, LowerError, Scope, compile_time, derivative_rhs, + function_calls::external_table_intrinsic_kind, helpers::{ build_indexed_binding_map, format_usize_dims, parse_indexed_binding_key, variable_size, }, unsupported_at, }; +mod residual_projection; + pub(super) type LoweredRowsAndTargets = (Vec>, Vec>); struct RowLoweringContext<'a> { @@ -1127,6 +1130,10 @@ fn lower_residual_rows_from_equations_core<'a>( let mut rows = expression_vec_with_capacity(equations.len(), "residual row count", equation_span)?; for (row_idx, eq) in equations { + if eq.scalar_count == 0 { + after_equation(eq, 0)?; + continue; + } let start = rows.len(); let ctx = RowLoweringContext { layout, @@ -1154,7 +1161,7 @@ fn lower_residual_rows_from_equations_core<'a>( )?; rows.extend(record_rows); let row_count = rows.len() - start; - validate_equation_row_count(eq, row_count)?; + validate_equation_row_count(eq, row_count, &ctx)?; after_equation(eq, row_count)?; continue; } @@ -1168,14 +1175,18 @@ fn lower_residual_rows_from_equations_core<'a>( )?; rows.extend(lowered_rows); let row_count = rows.len() - start; - validate_equation_row_count(eq, row_count)?; + validate_equation_row_count(eq, row_count, &ctx)?; after_equation(eq, row_count)?; } Ok(rows) } -fn validate_equation_row_count(eq: &dae::Equation, actual: usize) -> Result<(), LowerError> { - let expected = eq.scalar_count.max(1); +fn validate_equation_row_count( + eq: &dae::Equation, + actual: usize, + ctx: &RowLoweringContext<'_>, +) -> Result<(), LowerError> { + let expected = residual_equation_row_count(eq, ctx)?.unwrap_or(eq.scalar_count); if actual == expected { return Ok(()); } @@ -1188,6 +1199,63 @@ fn validate_equation_row_count(eq: &dae::Equation, actual: usize) -> Result<(), )) } +pub(crate) fn residual_equation_effective_row_count( + dae_model: &dae::Dae, + eq: &dae::Equation, +) -> Result { + let rumoca_core::Expression::Binary { + op: OpBinary::Sub, + lhs, + rhs, + span, + } = &eq.rhs + else { + return Ok(eq.scalar_count); + }; + let structural_bindings = compile_time::structural_bindings(dae_model)?; + residual_projection::scalarized_tuple_residual_binding_count( + lhs, + rhs, + dae_model, + &structural_bindings, + *span, + ) + .map(|count| count.unwrap_or(eq.scalar_count)) +} + +fn residual_equation_row_count( + eq: &dae::Equation, + ctx: &RowLoweringContext<'_>, +) -> Result, LowerError> { + let rumoca_core::Expression::Binary { + op: OpBinary::Sub, + lhs, + rhs, + span, + } = &eq.rhs + else { + return Ok(None); + }; + let Some(structural_bindings) = ctx.structural_bindings.as_ref() else { + return Ok(None); + }; + let Some(dae_variables) = ctx.dae_variables else { + return Ok(None); + }; + let mut dae_model = dae::Dae { + variables: dae_variables.clone(), + ..Default::default() + }; + dae_model.symbols.functions = ctx.functions.clone(); + residual_projection::scalarized_tuple_residual_binding_count( + lhs, + rhs, + &dae_model, + structural_bindings, + *span, + ) +} + fn lower_equation_residual_rows( eq: &dae::Equation, row_idx: usize, @@ -1218,15 +1286,42 @@ fn lower_equation_residual_rows( } else { &eq.rhs }; - - let scalar_count = eq.scalar_count.max(1); + let scalar_count = residual_equation_row_count(eq, ctx)? + .unwrap_or(eq.scalar_count) + .max(1); let values = if scalar_count == 1 { let mut values = expression_vec_with_capacity(1, "scalar residual value count", eq.span)?; - values.push( - builder - .lower_expr_with_source_context(expr, eq.span, &scope, 0) - .map_err(|err| residual_row_context(err, row_idx, eq))?, - ); + if row_idx >= state_scalar_count + && let Some(projected_values) = + scalarized_explicit_residual_values(eq, scalar_count, ctx, &mut builder)? + && let [value] = projected_values.as_slice() + { + values.push(*value); + } else if let Some(projected_values) = + scalarized_binary_residual_values(expr, scalar_count, ctx, &mut builder)? + && let [value] = projected_values.as_slice() + { + values.push(*value); + } else if let Some(value) = + lower_singleton_array_residual_value(&mut builder, expr, eq.span, &scope) + .map_err(|err| residual_row_context(err, row_idx, eq))? + { + values.push(value); + } else { + values.push( + builder + .lower_expr_with_source_context(expr, eq.span, &scope, 0) + .map_err(|err| residual_row_context(err, row_idx, eq))?, + ); + } + values + } else if let Some(values) = + scalarized_explicit_residual_values(eq, scalar_count, ctx, &mut builder)? + { + values + } else if let Some(values) = + scalarized_binary_residual_values(expr, scalar_count, ctx, &mut builder)? + { values } else { let values = builder @@ -1254,6 +1349,421 @@ fn lower_equation_residual_rows( Ok(rows) } +fn lower_singleton_array_residual_value( + builder: &mut LowerBuilder<'_>, + expr: &rumoca_core::Expression, + span: rumoca_core::Span, + scope: &Scope, +) -> Result, LowerError> { + let dims = match builder.infer_expr_dims(expr, scope) { + Ok(dims) => dims, + Err(_) => return Ok(None), + }; + if dims.is_empty() { + return Ok(None); + } + let mut count = 1usize; + for dim in dims { + count = count.checked_mul(dim).ok_or_else(|| { + LowerError::contract_violation( + "singleton residual shape overflows host index range", + span, + ) + })?; + } + if count != 1 { + return Ok(None); + } + let values = builder.lower_array_like_values_with_source_context(expr, span, scope, 0)?; + match values.as_slice() { + [value] => Ok(Some(*value)), + _ => Ok(None), + } +} + +fn scalarized_explicit_residual_values( + eq: &dae::Equation, + scalar_count: usize, + ctx: &RowLoweringContext<'_>, + builder: &mut LowerBuilder<'_>, +) -> Result>, LowerError> { + let Some(lhs) = eq.lhs.as_ref() else { + return Ok(None); + }; + let target = scalarized_explicit_residual_target(lhs, ctx.layout, eq.span)?; + if scalar_count == 1 && is_plain_scalar_residual_target(&target) { + return Ok(None); + } + scalarized_binary_residual_operands(&target, &eq.rhs, scalar_count, ctx, builder, eq.span) +} + +fn is_plain_scalar_residual_target(target: &rumoca_core::Expression) -> bool { + matches!( + target, + rumoca_core::Expression::VarRef { subscripts, .. } if subscripts.is_empty() + ) +} + +fn scalarized_explicit_residual_target( + lhs: &rumoca_core::Reference, + layout: &VarLayout, + span: rumoca_core::Span, +) -> Result { + if let Some((base, indices)) = parse_indexed_binding_key(lhs.as_str()) + && layout.shape(base.as_str()).is_some() + { + return Ok(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new(base), + subscripts: generated_subscripts_from_usize(&indices, span)?, + span, + }); + } + + Ok(rumoca_core::Expression::VarRef { + name: lhs.clone(), + subscripts: Vec::new(), + span, + }) +} + +fn scalarized_binary_residual_values( + expr: &rumoca_core::Expression, + scalar_count: usize, + ctx: &RowLoweringContext<'_>, + builder: &mut LowerBuilder<'_>, +) -> Result>, LowerError> { + let rumoca_core::Expression::Binary { + op: OpBinary::Sub, + lhs, + rhs, + span, + } = expr + else { + return Ok(None); + }; + if !requires_projected_function_scalars(rhs) { + return Ok(None); + } + if scalar_count == 1 + && matches!(lhs.as_ref(), rumoca_core::Expression::VarRef { name, subscripts, .. } + if subscripts.is_empty() && ctx.layout.binding(name.as_str()).is_some()) + { + return Ok(None); + } + scalarized_binary_residual_operands(lhs, rhs, scalar_count, ctx, builder, *span) +} + +fn requires_projected_function_scalars(expr: &rumoca_core::Expression) -> bool { + match expr { + rumoca_core::Expression::FunctionCall { + name, + is_constructor: false, + .. + } => external_table_intrinsic_kind(name.as_str()).is_none(), + rumoca_core::Expression::FieldAccess { base, .. } => { + matches!(base.as_ref(), rumoca_core::Expression::FunctionCall { .. }) + || requires_projected_function_scalars(base) + } + rumoca_core::Expression::Index { base, .. } => requires_projected_function_scalars(base), + rumoca_core::Expression::Binary { lhs, rhs, .. } => { + requires_projected_function_scalars(lhs) || requires_projected_function_scalars(rhs) + } + rumoca_core::Expression::Unary { rhs, .. } => requires_projected_function_scalars(rhs), + rumoca_core::Expression::BuiltinCall { args, .. } => { + args.iter().any(requires_projected_function_scalars) + } + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + branches + .iter() + .any(|(_, value)| requires_projected_function_scalars(value)) + || requires_projected_function_scalars(else_branch) + } + _ => false, + } +} + +#[allow(clippy::too_many_lines)] +fn scalarized_binary_residual_operands( + target: &rumoca_core::Expression, + rhs: &rumoca_core::Expression, + scalar_count: usize, + ctx: &RowLoweringContext<'_>, + builder: &mut LowerBuilder<'_>, + span: rumoca_core::Span, +) -> Result>, LowerError> { + let Some(structural_bindings) = ctx.structural_bindings.as_ref() else { + return Ok(None); + }; + let Some(dae_variables) = ctx.dae_variables else { + return Ok(None); + }; + let mut dae_model = dae::Dae { + variables: dae_variables.clone(), + ..Default::default() + }; + dae_model.symbols.functions = ctx.functions.clone(); + residual_projection::reject_array_denominator_division(rhs, &dae_model, structural_bindings)?; + if matches!(target, rumoca_core::Expression::Tuple { .. }) { + return residual_projection::scalarized_tuple_residual_operands( + target, + rhs, + scalar_count, + &dae_model, + structural_bindings, + builder, + span, + ); + } + let lhs_values = derivative_rhs::scalarized_rhs_expressions_with_owner( + target, + target, + scalar_count, + &dae_model, + structural_bindings, + span, + )?; + if scalar_count == 1 + && let Some(lane_index) = target_scalarized_lane_index(target, ctx.layout, span)? + && let Some(rhs_value) = derivative_rhs::project_array_like_scalar_with_owner( + rhs, + lane_index, + &dae_model, + structural_bindings, + span, + )? + { + let scope = Scope::new(); + let mut values = + expression_vec_with_capacity(1, "scalarized explicit residual value count", span)?; + let lhs = lhs_values.into_iter().next().ok_or_else(|| { + LowerError::contract_violation("scalarized residual target produced no lhs value", span) + })?; + let residual = rumoca_core::Expression::Binary { + op: OpBinary::Sub, + lhs: Box::new(lhs), + rhs: Box::new(rhs_value), + span, + }; + values.push(builder.lower_expr_with_source_context(&residual, span, &scope, 0)?); + return Ok(Some(values)); + } + if let Some(rhs_values) = matching_function_output_residual_values( + target, + rhs, + scalar_count, + &dae_model, + structural_bindings, + span, + )? { + return scalarized_residual_registers(lhs_values, rhs_values, scalar_count, builder, span); + } + let rhs_values = if requires_projected_function_scalars(rhs) { + derivative_rhs::function_call_projected_scalars_with_owner( + rhs, + &dae_model, + structural_bindings, + span, + )? + .or(derivative_rhs::project_array_like_scalars_with_owner( + rhs, + &dae_model, + structural_bindings, + span, + )?) + } else { + derivative_rhs::project_array_like_scalars_with_owner( + rhs, + &dae_model, + structural_bindings, + span, + )? + }; + let Some(mut rhs_values) = rhs_values else { + return Ok(None); + }; + if rhs_values.len() == 1 && scalar_count > 1 { + let Some(projected) = derivative_rhs::project_array_like_scalars_with_owner( + &rhs_values[0], + &dae_model, + structural_bindings, + span, + )? + else { + return Ok(None); + }; + rhs_values = projected; + } + if scalar_count == 1 + && rhs_values.len() > 1 + && let Some(lane_index) = target_scalarized_lane_index(target, ctx.layout, span)? + { + let Some(rhs_value) = rhs_values.get(lane_index).cloned() else { + return Ok(None); + }; + rhs_values = vec![rhs_value]; + } + if lhs_values.len() != scalar_count || rhs_values.len() != scalar_count { + return Ok(None); + } + scalarized_residual_registers(lhs_values, rhs_values, scalar_count, builder, span) +} + +fn matching_function_output_residual_values( + target: &rumoca_core::Expression, + rhs: &rumoca_core::Expression, + scalar_count: usize, + dae_model: &dae::Dae, + structural_bindings: &IndexMap, + span: rumoca_core::Span, +) -> Result>, LowerError> { + let Some(target_leaf) = residual_target_leaf_name(target) else { + return Ok(None); + }; + let rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor: false, + .. + } = rhs + else { + return Ok(None); + }; + let Some(function) = dae_model.symbols.functions.get(name.var_name()) else { + return Ok(None); + }; + if function.inputs.iter().any(|input| { + matches!( + input.type_class, + Some(rumoca_core::ClassType::Record | rumoca_core::ClassType::Connector) + ) + }) { + return Ok(None); + } + if !function + .outputs + .iter() + .any(|output| output.name == target_leaf) + { + return Ok(None); + } + let selected_call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new(format!("{}.{}", name.as_str(), target_leaf)).into(), + args: args.clone(), + is_constructor: false, + span, + }; + let Some(values) = (match derivative_rhs::function_call_projected_scalars_with_owner( + &selected_call, + dae_model, + structural_bindings, + span, + )? { + Some(values) => Some(values), + None => derivative_rhs::project_array_like_scalars_with_owner( + &selected_call, + dae_model, + structural_bindings, + span, + )?, + }) else { + return Ok(None); + }; + if values.len() == scalar_count { + Ok(Some(values)) + } else { + Ok(None) + } +} + +fn residual_target_leaf_name(target: &rumoca_core::Expression) -> Option<&str> { + match target { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } if subscripts.is_empty() => Some(name.last_segment()), + _ => None, + } +} + +fn scalarized_residual_registers( + lhs_values: Vec, + rhs_values: Vec, + scalar_count: usize, + builder: &mut LowerBuilder<'_>, + span: rumoca_core::Span, +) -> Result>, LowerError> { + if lhs_values.len() != scalar_count || rhs_values.len() != scalar_count { + return Ok(None); + } + let scope = Scope::new(); + let mut values = expression_vec_with_capacity( + scalar_count, + "scalarized explicit residual value count", + span, + )?; + for (lhs, rhs) in lhs_values.into_iter().zip(rhs_values) { + let residual = rumoca_core::Expression::Binary { + op: OpBinary::Sub, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span, + }; + values.push(builder.lower_expr_with_source_context(&residual, span, &scope, 0)?); + } + Ok(Some(values)) +} + +fn target_scalarized_lane_index( + target: &rumoca_core::Expression, + layout: &VarLayout, + span: rumoca_core::Span, +) -> Result, LowerError> { + let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = target + else { + return Ok(None); + }; + if subscripts.is_empty() { + return Ok(None); + } + let Some(dims) = layout.shape(name.as_str()) else { + return Ok(None); + }; + let Some(indices) = super::helpers::static_subscript_indices_with_owner(subscripts, span)? + else { + return Ok(None); + }; + if indices.len() != dims.len() { + return Ok(None); + } + let mut flat_index = 0usize; + for (dim_index, index) in indices.iter().copied().enumerate() { + let dim = dims[dim_index]; + if index == 0 || index > dim { + return Ok(None); + } + let stride = dims[dim_index + 1..].iter().product::(); + flat_index = flat_index + .checked_add((index - 1).checked_mul(stride).ok_or_else(|| { + LowerError::contract_violation( + "scalarized target lane index overflows host index range", + span, + ) + })?) + .ok_or_else(|| { + LowerError::contract_violation( + "scalarized target lane index overflows host index range", + span, + ) + })?; + } + Ok(Some(flat_index)) +} + fn lower_scalarized_record_residual_rows( eq: &dae::Equation, row_idx: usize, @@ -1418,195 +1928,4 @@ fn lower_builder_for_context<'a>( } #[cfg(test)] -mod tests { - use super::*; - - fn layout_with_bindings(names: &[&str]) -> VarLayout { - let mut bindings = IndexMap::new(); - for (index, name) in names.iter().enumerate() { - bindings.insert( - (*name).to_string(), - ScalarSlot::Y { - index, - byte_offset: index * std::mem::size_of::(), - }, - ); - } - VarLayout::from_parts(bindings, names.len(), 0) - } - - #[test] - fn scalarized_record_fields_ignore_only_top_level_nested_suffixes() { - let layout = layout_with_bindings(&[ - "state.p", - "state.v[index.with.dot]", - "state.nested.q", - "other.p", - ]); - - let fields = scalarized_record_fields("state", &layout).expect("state fields"); - let suffixes = fields - .iter() - .map(|field| field.suffix.as_str()) - .collect::>(); - - assert!(suffixes.contains(&"p")); - assert!(suffixes.contains(&"v[index.with.dot]")); - assert!(!suffixes.contains(&"nested.q")); - } - - #[test] - fn scalarized_record_fields_ignore_scalar_binding_with_enum_alias_prefix() { - let layout = layout_with_bindings(&["Th", "Th.default"]); - - assert!(scalarized_record_fields("Th", &layout).is_none()); - } - - #[test] - fn scalar_row_namespace_rejects_overflow() { - let span = rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name( - "phase_solve_lower_expression_rows_source_19.mo", - ), - 1, - 4, - ); - let err = scalar_row_namespace(u64::MAX, 1, span) - .expect_err("row namespace multiplication should reject overflow"); - - assert_eq!(err.source_span(), Some(span)); - assert!(matches!(err, LowerError::ContractViolation { .. })); - } - - fn unspanned_expression_rows_test_span() -> rumoca_core::Span { - rumoca_core::Span::DUMMY - } - - #[test] - fn scalar_row_namespace_rejects_overflow_without_dummy_span() { - let err = scalar_row_namespace(u64::MAX, 1, unspanned_expression_rows_test_span()) - .expect_err("row namespace multiplication should reject overflow"); - - assert_eq!(err.source_span(), None); - assert!( - err.reason().contains("scalar row namespace overflows"), - "{err:?}" - ); - } - - #[test] - fn row_namespace_rejects_u64_overflow_with_span() { - let Some(row_idx) = usize::try_from(u64::MAX) - .ok() - .and_then(|value| value.checked_add(1)) - else { - return; - }; - let span = rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name( - "phase_solve_lower_expression_rows_source_21.mo", - ), - 2, - 7, - ); - - let err = row_namespace_from_usize(row_idx, Some(span)) - .expect_err("row namespace must fit in u64"); - - assert_eq!(err.source_span(), Some(span)); - assert!(matches!(err, LowerError::ContractViolation { .. })); - assert!(err.reason().contains("exceeds u64 namespace")); - } - - #[test] - fn expression_contract_violation_without_span_stays_unspanned() { - let err = expression_contract_violation("expression row metadata mismatch", None); - - assert_eq!(err.source_span(), None); - assert!(matches!(err, LowerError::UnspannedContractViolation { .. })); - assert!( - err.reason().contains("expression row metadata mismatch"), - "{err:?}" - ); - } - - #[test] - fn generated_subscript_rejects_i64_overflow_with_span() { - let Some(index) = usize::try_from(i64::MAX) - .ok() - .and_then(|value| value.checked_add(1)) - else { - return; - }; - let span = rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name( - "phase_solve_lower_expression_rows_source_20.mo", - ), - 3, - 8, - ); - - let err = generated_subscript_from_usize(index, span) - .expect_err("generated subscript must fit in Modelica integer range"); - - assert_eq!(err.source_span(), Some(span)); - assert!(matches!( - err, - LowerError::ContractViolation { reason, span: actual } - if actual == span && reason.contains("exceeds i64 range") - )); - } - - #[test] - fn indexed_sample_value_rejects_i64_dimension_overflow_with_span() { - let Some(dim) = usize::try_from(i64::MAX) - .ok() - .and_then(|value| value.checked_add(1)) - else { - return; - }; - let span = rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name( - "phase_solve_lower_expression_rows_source_21.mo", - ), - 5, - 11, - ); - let value = rumoca_core::Expression::Literal { - value: rumoca_core::Literal::Real(1.0), - span, - }; - - let err = indexed_sample_value(&value, &[dim], 0, span) - .expect_err("sample array dimension must fit in Modelica integer range"); - - assert_eq!(err.source_span(), Some(span)); - assert!(matches!( - err, - LowerError::ContractViolation { reason, span: actual } - if actual == span && reason.contains("sample array dimension") - )); - } - - #[test] - fn expand_row_values_reports_capacity_overflow_with_owner_span() { - let span = rumoca_core::Span::from_offsets( - rumoca_core::SourceId::from_source_name( - "phase_solve_lower_expression_rows_source_22.mo", - ), - 1, - 4, - ); - let err = match expand_row_values(vec![0], usize::MAX, span) { - Ok(_) => panic!("oversized expression value expansion should fail before allocating"), - Err(err) => err, - }; - - assert_eq!(err.source_span(), Some(span)); - assert!( - err.reason() - .contains("expanded expression value count capacity"), - "unexpected error: {err}" - ); - } -} +mod tests; diff --git a/crates/rumoca-phase-solve/src/lower/expression_rows/residual_projection.rs b/crates/rumoca-phase-solve/src/lower/expression_rows/residual_projection.rs new file mode 100644 index 000000000..0d33bd092 --- /dev/null +++ b/crates/rumoca-phase-solve/src/lower/expression_rows/residual_projection.rs @@ -0,0 +1,713 @@ +use indexmap::IndexMap; +use rumoca_core::OpBinary; +use rumoca_ir_dae as dae; +use rumoca_ir_solve::{BinaryOp, Reg}; + +use crate::lower::{ + LowerBuilder, LowerError, Scope, compile_time, derivative_rhs, helpers::format_usize_dims, + unsupported_at, +}; + +pub(super) fn scalarized_tuple_residual_operands( + target: &rumoca_core::Expression, + rhs: &rumoca_core::Expression, + scalar_count: usize, + dae_model: &dae::Dae, + structural_bindings: &IndexMap, + builder: &mut LowerBuilder<'_>, + span: rumoca_core::Span, +) -> Result>, LowerError> { + let scope = Scope::new(); + if let rumoca_core::Expression::Tuple { elements, .. } = target + && let Some(output_groups) = match selected_output_groups_from_declared_shapes( + rhs, + dae_model, + structural_bindings, + span, + )? { + Some(groups) => Some(groups), + None => derivative_rhs::function_call_projected_output_groups_with_owner( + rhs, + dae_model, + structural_bindings, + span, + )?, + } + { + if elements.len() != output_groups.len() { + return Ok(None); + } + let mut values = super::expression_vec_with_capacity( + scalar_count, + "scalarized tuple residual value count", + span, + )?; + for (target_element, rhs_values) in elements.iter().zip(output_groups) { + if rhs_values.is_empty() || !tuple_target_element_has_binding(target_element) { + continue; + } + let lhs_values = lower_tuple_target_element_values( + builder, + target_element, + rhs_values.len(), + span, + &scope, + )?; + if lhs_values.len() != rhs_values.len() { + return Ok(None); + } + for (lhs, rhs) in lhs_values.into_iter().zip(rhs_values) { + let rhs = builder.lower_expr_with_source_context(&rhs, span, &scope, 0)?; + values.push(builder.emit_binary_at(BinaryOp::Sub, lhs, rhs, span)?); + } + } + if values.is_empty() { + return Ok(None); + } + return Ok(Some(values)); + } + let lhs_values = + builder.lower_array_like_values_with_source_context(target, span, &scope, 0)?; + let Some(rhs_values) = derivative_rhs::function_call_projected_scalars_with_owner( + rhs, + dae_model, + structural_bindings, + span, + )? + .or(derivative_rhs::project_array_like_scalars_with_owner( + rhs, + dae_model, + structural_bindings, + span, + )?) else { + return Ok(None); + }; + if lhs_values.len() != scalar_count || rhs_values.len() != scalar_count { + return Ok(None); + } + let mut values = super::expression_vec_with_capacity( + scalar_count, + "scalarized tuple residual value count", + span, + )?; + for (lhs, rhs) in lhs_values.into_iter().zip(rhs_values) { + let rhs = builder.lower_expr_with_source_context(&rhs, span, &scope, 0)?; + values.push(builder.emit_binary_at(BinaryOp::Sub, lhs, rhs, span)?); + } + Ok(Some(values)) +} + +fn selected_output_groups_from_declared_shapes( + rhs: &rumoca_core::Expression, + dae_model: &dae::Dae, + structural_bindings: &IndexMap, + span: rumoca_core::Span, +) -> Result>>, LowerError> { + let rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor: false, + .. + } = rhs + else { + return Ok(None); + }; + let Some(function) = dae_model.symbols.functions.get(name.var_name()) else { + return Ok(None); + }; + if !args_are_compile_time_scalars(args, structural_bindings)? { + return Ok(None); + } + let Some(shape_env) = function_call_shape_env( + name.as_str(), + &function.inputs, + args, + dae_model, + structural_bindings, + span, + )? + else { + return Ok(None); + }; + let mut groups = super::expression_vec_with_capacity( + function.outputs.len(), + "declared selected-output group count", + span, + )?; + for output in &function.outputs { + let Some(count) = function_param_scalar_count(output, &shape_env, span)? else { + return Ok(None); + }; + let mut group = super::expression_vec_with_capacity( + count, + "declared selected-output scalar count", + span, + )?; + if count == 1 { + group.push(selected_output_expression( + dae_model, name, args, output, None, span, + )); + } else { + for index in 1..=count { + group.push(selected_output_expression( + dae_model, + name, + args, + output, + Some(index), + span, + )); + } + } + groups.push(group); + } + Ok(Some(groups)) +} + +fn args_are_compile_time_scalars( + args: &[rumoca_core::Expression], + structural_bindings: &IndexMap, +) -> Result { + for arg in args { + if static_shape_scalar(arg, &IndexMap::new(), structural_bindings)?.is_none() { + return Ok(false); + } + } + Ok(true) +} + +fn selected_output_expression( + dae_model: &dae::Dae, + function_name: &rumoca_core::Reference, + args: &[rumoca_core::Expression], + output: &rumoca_core::FunctionParam, + one_based_index: Option, + span: rumoca_core::Span, +) -> rumoca_core::Expression { + let selection_indices = selected_output_eval_indices(output, one_based_index); + if let Some(value) = compile_time::eval_selected_function_output( + dae_model, + function_name.var_name(), + output.name.as_str(), + &selection_indices, + args, + ) { + return rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(value), + span, + }; + } + let selector = match one_based_index { + Some(index) => format!("{}[{index}]", output.name), + None => output.name.to_string(), + }; + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new(format!("{}.{}", function_name.as_str(), selector)).into(), + args: args.to_vec(), + is_constructor: false, + span, + } +} + +fn selected_output_eval_indices( + output: &rumoca_core::FunctionParam, + one_based_index: Option, +) -> Vec { + match one_based_index { + Some(index) => vec![index as i64], + None if output.dims.is_empty() && output.shape_expr.is_empty() => Vec::new(), + None => vec![1], + } +} + +pub(super) fn scalarized_tuple_residual_binding_count( + target: &rumoca_core::Expression, + rhs: &rumoca_core::Expression, + dae_model: &dae::Dae, + structural_bindings: &IndexMap, + span: rumoca_core::Span, +) -> Result, LowerError> { + let rumoca_core::Expression::Tuple { elements, .. } = target else { + return Ok(None); + }; + if let Some(output_lengths) = + function_call_declared_output_group_lengths(rhs, dae_model, structural_bindings, span)? + && let Some(count) = tuple_bound_output_count(elements, &output_lengths, span)? + { + return Ok(Some(count)); + } + let Some(output_groups) = derivative_rhs::function_call_projected_output_groups_with_owner( + rhs, + dae_model, + structural_bindings, + span, + )? + else { + return Ok(None); + }; + if elements.len() != output_groups.len() { + return Ok(None); + } + let output_lengths: Vec<_> = output_groups.into_iter().map(|group| group.len()).collect(); + tuple_bound_output_count(elements, &output_lengths, span) +} + +fn function_call_declared_output_group_lengths( + rhs: &rumoca_core::Expression, + dae_model: &dae::Dae, + structural_bindings: &IndexMap, + span: rumoca_core::Span, +) -> Result>, LowerError> { + let rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor: false, + .. + } = rhs + else { + return Ok(None); + }; + let Some(function) = dae_model.symbols.functions.get(name.var_name()) else { + return Ok(None); + }; + let Some(shape_env) = function_call_shape_env( + name.as_str(), + &function.inputs, + args, + dae_model, + structural_bindings, + span, + )? + else { + return Ok(None); + }; + let mut lengths = super::expression_vec_with_capacity( + function.outputs.len(), + "declared function output group count", + span, + )?; + for output in &function.outputs { + let Some(count) = function_param_scalar_count(output, &shape_env, span)? else { + return Ok(None); + }; + lengths.push(count); + } + Ok(Some(lengths)) +} + +struct FunctionCallShapeEnv { + input_dims: IndexMap>, + input_scalars: IndexMap, +} + +fn function_call_shape_env( + function_name: &str, + inputs: &[rumoca_core::FunctionParam], + args: &[rumoca_core::Expression], + dae_model: &dae::Dae, + structural_bindings: &IndexMap, + span: rumoca_core::Span, +) -> Result, LowerError> { + let (named_args, positional_args) = + crate::lower::function_calls::split_named_and_positional_call_args(function_name, args)?; + let mut positional_idx = 0usize; + let mut dims_by_input = IndexMap::new(); + let mut scalars_by_input = IndexMap::new(); + for input in inputs { + let actual = if let Some(actual) = named_args.get(input.name.as_str()).copied() { + actual + } else if let Some(actual) = positional_args.get(positional_idx).copied() { + positional_idx += 1; + actual + } else if let Some(default) = input.default.as_ref() { + default + } else if input.dims.iter().all(|dim| *dim >= 0) { + let dims = input + .dims + .iter() + .map(|dim| usize::try_from(*dim)) + .collect::, _>>() + .map_err(|_| { + LowerError::contract_violation( + format!( + "function `{function_name}` input `{}` has invalid declared dimension", + input.name + ), + span, + ) + })?; + dims_by_input.insert(input.name.clone(), dims); + continue; + } else { + return Ok(None); + }; + let mut dims = + derivative_rhs::expression_result_dims(actual, dae_model, structural_bindings, span)?; + if let Some(value) = static_shape_scalar(actual, &scalars_by_input, structural_bindings)? { + scalars_by_input.insert(input.name.clone(), value); + } + if input.dims.as_slice() == [0] + && dims.is_empty() + && is_zero_length_component_placeholder(actual) + { + dims.push(0); + } + dims_by_input.insert(input.name.clone(), dims); + } + Ok(Some(FunctionCallShapeEnv { + input_dims: dims_by_input, + input_scalars: scalars_by_input, + })) +} + +fn function_param_scalar_count( + param: &rumoca_core::FunctionParam, + shape_env: &FunctionCallShapeEnv, + span: rumoca_core::Span, +) -> Result, LowerError> { + if !param.shape_expr.is_empty() { + let mut count = 1usize; + for shape in ¶m.shape_expr { + let Some(dim) = shape_expr_scalar_count_dim(shape, shape_env, span)? else { + return Ok(None); + }; + count = count.checked_mul(dim).ok_or_else(|| { + LowerError::contract_violation( + format!( + "function output `{}` scalar count overflows host index range", + param.name + ), + span, + ) + })?; + } + return Ok(Some(count)); + } + if param.dims.is_empty() { + return Ok(Some(1)); + } + let mut count = 1usize; + for dim in ¶m.dims { + let dim = usize::try_from(*dim).map_err(|_| { + LowerError::contract_violation( + format!( + "function output `{}` has invalid declared dimension {dim}", + param.name + ), + span, + ) + })?; + count = count.checked_mul(dim).ok_or_else(|| { + LowerError::contract_violation( + format!( + "function output `{}` scalar count overflows host index range", + param.name + ), + span, + ) + })?; + } + Ok(Some(count)) +} + +fn shape_expr_scalar_count_dim( + shape: &rumoca_core::Subscript, + shape_env: &FunctionCallShapeEnv, + span: rumoca_core::Span, +) -> Result, LowerError> { + match shape { + rumoca_core::Subscript::Index { value, .. } => usize::try_from(*value) + .map(Some) + .map_err(|_| LowerError::contract_violation("negative function output shape", span)), + rumoca_core::Subscript::Expr { expr, .. } => { + match shape_expr_builtin_size_dim(expr, shape_env, span)? { + Some(dim) => Ok(Some(dim)), + None => static_shape_scalar(expr, &shape_env.input_scalars, &IndexMap::new())? + .map(|value| usize_shape_dim(value, expr.span().unwrap_or(span))) + .transpose(), + } + } + rumoca_core::Subscript::Colon { .. } => Ok(None), + } +} + +fn shape_expr_builtin_size_dim( + expr: &rumoca_core::Expression, + shape_env: &FunctionCallShapeEnv, + span: rumoca_core::Span, +) -> Result, LowerError> { + let rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + args, + .. + } = expr + else { + return Ok(None); + }; + let [base, dim] = args.as_slice() else { + return Ok(None); + }; + let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = base + else { + return Ok(None); + }; + if !subscripts.is_empty() { + return Ok(None); + } + let rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(dim), + .. + } = dim + else { + return Ok(None); + }; + let dim = usize::try_from(*dim) + .ok() + .and_then(|dim| dim.checked_sub(1)) + .ok_or_else(|| LowerError::contract_violation("size dimension must be positive", span))?; + Ok(shape_env + .input_dims + .get(name.as_str()) + .and_then(|dims| dims.get(dim)) + .copied()) +} + +fn static_shape_scalar( + expr: &rumoca_core::Expression, + local_scalars: &IndexMap, + structural_bindings: &IndexMap, +) -> Result, LowerError> { + match expr { + rumoca_core::Expression::Literal { value, .. } => Ok(literal_scalar(value)), + rumoca_core::Expression::VarRef { + name, subscripts, .. + } if subscripts.is_empty() => Ok(local_scalars + .get(name.as_str()) + .copied() + .or_else(|| structural_bindings.get(name.as_str()).copied())), + rumoca_core::Expression::Unary { op, rhs, .. } => { + let Some(value) = static_shape_scalar(rhs, local_scalars, structural_bindings)? else { + return Ok(None); + }; + match op { + rumoca_core::OpUnary::Minus | rumoca_core::OpUnary::DotMinus => Ok(Some(-value)), + rumoca_core::OpUnary::Plus + | rumoca_core::OpUnary::DotPlus + | rumoca_core::OpUnary::Empty => Ok(Some(value)), + rumoca_core::OpUnary::Not => Ok(None), + } + } + rumoca_core::Expression::Binary { op, lhs, rhs, .. } => { + let Some(lhs) = static_shape_scalar(lhs, local_scalars, structural_bindings)? else { + return Ok(None); + }; + let Some(rhs) = static_shape_scalar(rhs, local_scalars, structural_bindings)? else { + return Ok(None); + }; + match op { + rumoca_core::OpBinary::Add => Ok(Some(lhs + rhs)), + rumoca_core::OpBinary::Sub => Ok(Some(lhs - rhs)), + rumoca_core::OpBinary::Mul => Ok(Some(lhs * rhs)), + rumoca_core::OpBinary::Div => Ok(Some(lhs / rhs)), + _ => Ok(None), + } + } + rumoca_core::Expression::BuiltinCall { function, args, .. } => { + static_shape_builtin_scalar(function, args, local_scalars, structural_bindings) + } + _ => Ok(None), + } +} + +fn static_shape_builtin_scalar( + function: &rumoca_core::BuiltinFunction, + args: &[rumoca_core::Expression], + local_scalars: &IndexMap, + structural_bindings: &IndexMap, +) -> Result, LowerError> { + match (function, args) { + (rumoca_core::BuiltinFunction::Integer, [arg]) => { + Ok(static_shape_scalar(arg, local_scalars, structural_bindings)?.map(f64::trunc)) + } + (rumoca_core::BuiltinFunction::Mod, [lhs, rhs]) => { + let Some(lhs) = static_shape_scalar(lhs, local_scalars, structural_bindings)? else { + return Ok(None); + }; + let Some(rhs) = static_shape_scalar(rhs, local_scalars, structural_bindings)? else { + return Ok(None); + }; + Ok(Some(lhs - rhs * (lhs / rhs).floor())) + } + _ => Ok(None), + } +} + +fn literal_scalar(value: &rumoca_core::Literal) -> Option { + match value { + rumoca_core::Literal::Integer(value) => Some(*value as f64), + rumoca_core::Literal::Real(value) => Some(*value), + rumoca_core::Literal::Boolean(value) => Some(if *value { 1.0 } else { 0.0 }), + rumoca_core::Literal::String(_) => None, + } +} + +fn usize_shape_dim(value: f64, span: rumoca_core::Span) -> Result { + if value.is_finite() && value >= 0.0 && value.fract() == 0.0 && value <= usize::MAX as f64 { + Ok(value as usize) + } else { + Err(LowerError::contract_violation( + format!("function output shape evaluated to invalid dimension {value}"), + span, + )) + } +} + +fn is_zero_length_component_placeholder(expr: &rumoca_core::Expression) -> bool { + matches!( + expr, + rumoca_core::Expression::FunctionCall { + name, + is_constructor: true, + .. + } if name.as_str() == "Real" + ) +} + +fn tuple_bound_output_count( + elements: &[rumoca_core::Expression], + output_lengths: &[usize], + span: rumoca_core::Span, +) -> Result, LowerError> { + if elements.len() != output_lengths.len() { + return Ok(None); + } + let mut count = 0usize; + for (target_element, rhs_len) in elements.iter().zip(output_lengths.iter().copied()) { + let target_has_binding = tuple_target_element_has_binding(target_element); + if rhs_len == 0 && target_has_binding { + return Ok(None); + } + if rhs_len == 0 || !target_has_binding { + continue; + } + count = count.checked_add(rhs_len).ok_or_else(|| { + LowerError::contract_violation( + "scalarized tuple residual binding count overflows host index range", + span, + ) + })?; + } + Ok(Some(count)) +} + +fn tuple_target_element_has_binding(target: &rumoca_core::Expression) -> bool { + match target { + rumoca_core::Expression::VarRef { .. } => true, + rumoca_core::Expression::Array { elements, .. } + | rumoca_core::Expression::Tuple { elements, .. } => { + elements.iter().any(tuple_target_element_has_binding) + } + rumoca_core::Expression::FieldAccess { base, .. } + | rumoca_core::Expression::Index { base, .. } => tuple_target_element_has_binding(base), + _ => false, + } +} + +fn lower_tuple_target_element_values( + builder: &mut LowerBuilder<'_>, + target: &rumoca_core::Expression, + expected_count: usize, + span: rumoca_core::Span, + scope: &Scope, +) -> Result, LowerError> { + if expected_count == 1 + && builder + .infer_expr_dims(target, scope) + .map(|dims| dims.is_empty()) + .unwrap_or(false) + { + let mut values = + super::expression_vec_with_capacity(1, "scalarized tuple target scalar count", span)?; + values.push(builder.lower_expr_with_source_context(target, span, scope, 0)?); + return Ok(values); + } + builder.lower_array_like_values_with_source_context(target, span, scope, 0) +} + +pub(super) fn reject_array_denominator_division( + expr: &rumoca_core::Expression, + dae_model: &dae::Dae, + structural_bindings: &IndexMap, +) -> Result<(), LowerError> { + match expr { + rumoca_core::Expression::Binary { + op: OpBinary::Div, + lhs, + rhs, + span, + } => { + let lhs_dims = + derivative_rhs::expression_result_dims(lhs, dae_model, structural_bindings, *span)?; + let rhs_dims = + derivative_rhs::expression_result_dims(rhs, dae_model, structural_bindings, *span)?; + if !rhs_dims.is_empty() { + return Err(unsupported_at( + format!( + "array division requires a scalar denominator \ + (lhs_shape={}, lhs_values={}, rhs_shape={}, rhs_values={})", + format_usize_dims(&lhs_dims), + dims_value_count(&lhs_dims, *span)?, + format_usize_dims(&rhs_dims), + dims_value_count(&rhs_dims, *span)?, + ), + *span, + )); + } + } + rumoca_core::Expression::Binary { lhs, rhs, .. } => { + reject_array_denominator_division(lhs, dae_model, structural_bindings)?; + reject_array_denominator_division(rhs, dae_model, structural_bindings)?; + } + rumoca_core::Expression::Unary { rhs, .. } => { + reject_array_denominator_division(rhs, dae_model, structural_bindings)?; + } + rumoca_core::Expression::FieldAccess { base, .. } + | rumoca_core::Expression::Index { base, .. } => { + reject_array_denominator_division(base, dae_model, structural_bindings)?; + } + rumoca_core::Expression::Array { elements, .. } + | rumoca_core::Expression::Tuple { elements, .. } => { + for element in elements { + reject_array_denominator_division(element, dae_model, structural_bindings)?; + } + } + rumoca_core::Expression::BuiltinCall { args, .. } + | rumoca_core::Expression::FunctionCall { args, .. } => { + for arg in args { + reject_array_denominator_division(arg, dae_model, structural_bindings)?; + } + } + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + for (_, branch_expr) in branches { + reject_array_denominator_division(branch_expr, dae_model, structural_bindings)?; + } + reject_array_denominator_division(else_branch, dae_model, structural_bindings)?; + } + _ => {} + } + Ok(()) +} + +fn dims_value_count(dims: &[usize], span: rumoca_core::Span) -> Result { + dims.iter().try_fold(1usize, |count, dim| { + count.checked_mul(*dim).ok_or_else(|| { + LowerError::contract_violation("array division shape overflows host index range", span) + }) + }) +} diff --git a/crates/rumoca-phase-solve/src/lower/expression_rows/tests.rs b/crates/rumoca-phase-solve/src/lower/expression_rows/tests.rs new file mode 100644 index 000000000..d9f99dee1 --- /dev/null +++ b/crates/rumoca-phase-solve/src/lower/expression_rows/tests.rs @@ -0,0 +1,180 @@ +use super::*; + +fn layout_with_bindings(names: &[&str]) -> VarLayout { + let mut bindings = IndexMap::new(); + for (index, name) in names.iter().enumerate() { + bindings.insert( + (*name).to_string(), + ScalarSlot::Y { + index, + byte_offset: index * std::mem::size_of::(), + }, + ); + } + VarLayout::from_parts(bindings, names.len(), 0) +} + +#[test] +fn scalarized_record_fields_ignore_only_top_level_nested_suffixes() { + let layout = layout_with_bindings(&[ + "state.p", + "state.v[index.with.dot]", + "state.nested.q", + "other.p", + ]); + + let fields = scalarized_record_fields("state", &layout).expect("state fields"); + let suffixes = fields + .iter() + .map(|field| field.suffix.as_str()) + .collect::>(); + + assert!(suffixes.contains(&"p")); + assert!(suffixes.contains(&"v[index.with.dot]")); + assert!(!suffixes.contains(&"nested.q")); +} + +#[test] +fn scalarized_record_fields_ignore_scalar_binding_with_enum_alias_prefix() { + let layout = layout_with_bindings(&["Th", "Th.default"]); + + assert!(scalarized_record_fields("Th", &layout).is_none()); +} + +#[test] +fn scalar_row_namespace_rejects_overflow() { + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("phase_solve_lower_expression_rows_source_19.mo"), + 1, + 4, + ); + let err = scalar_row_namespace(u64::MAX, 1, span) + .expect_err("row namespace multiplication should reject overflow"); + + assert_eq!(err.source_span(), Some(span)); + assert!(matches!(err, LowerError::ContractViolation { .. })); +} + +fn unspanned_expression_rows_test_span() -> rumoca_core::Span { + rumoca_core::Span::DUMMY +} + +#[test] +fn scalar_row_namespace_rejects_overflow_without_dummy_span() { + let err = scalar_row_namespace(u64::MAX, 1, unspanned_expression_rows_test_span()) + .expect_err("row namespace multiplication should reject overflow"); + + assert_eq!(err.source_span(), None); + assert!( + err.reason().contains("scalar row namespace overflows"), + "{err:?}" + ); +} + +#[test] +fn row_namespace_rejects_u64_overflow_with_span() { + let Some(row_idx) = usize::try_from(u64::MAX) + .ok() + .and_then(|value| value.checked_add(1)) + else { + return; + }; + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("phase_solve_lower_expression_rows_source_21.mo"), + 2, + 7, + ); + + let err = + row_namespace_from_usize(row_idx, Some(span)).expect_err("row namespace must fit in u64"); + + assert_eq!(err.source_span(), Some(span)); + assert!(matches!(err, LowerError::ContractViolation { .. })); + assert!(err.reason().contains("exceeds u64 namespace")); +} + +#[test] +fn expression_contract_violation_without_span_stays_unspanned() { + let err = expression_contract_violation("expression row metadata mismatch", None); + + assert_eq!(err.source_span(), None); + assert!(matches!(err, LowerError::UnspannedContractViolation { .. })); + assert!( + err.reason().contains("expression row metadata mismatch"), + "{err:?}" + ); +} + +#[test] +fn generated_subscript_rejects_i64_overflow_with_span() { + let Some(index) = usize::try_from(i64::MAX) + .ok() + .and_then(|value| value.checked_add(1)) + else { + return; + }; + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("phase_solve_lower_expression_rows_source_20.mo"), + 3, + 8, + ); + + let err = generated_subscript_from_usize(index, span) + .expect_err("generated subscript must fit in Modelica integer range"); + + assert_eq!(err.source_span(), Some(span)); + assert!(matches!( + err, + LowerError::ContractViolation { reason, span: actual } + if actual == span && reason.contains("exceeds i64 range") + )); +} + +#[test] +fn indexed_sample_value_rejects_i64_dimension_overflow_with_span() { + let Some(dim) = usize::try_from(i64::MAX) + .ok() + .and_then(|value| value.checked_add(1)) + else { + return; + }; + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("phase_solve_lower_expression_rows_source_21.mo"), + 5, + 11, + ); + let value = rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(1.0), + span, + }; + + let err = indexed_sample_value(&value, &[dim], 0, span) + .expect_err("sample array dimension must fit in Modelica integer range"); + + assert_eq!(err.source_span(), Some(span)); + assert!(matches!( + err, + LowerError::ContractViolation { reason, span: actual } + if actual == span && reason.contains("sample array dimension") + )); +} + +#[test] +fn expand_row_values_reports_capacity_overflow_with_owner_span() { + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("phase_solve_lower_expression_rows_source_22.mo"), + 1, + 4, + ); + let err = match expand_row_values(vec![0], usize::MAX, span) { + Ok(_) => panic!("oversized expression value expansion should fail before allocating"), + Err(err) => err, + }; + + assert_eq!(err.source_span(), Some(span)); + assert!( + err.reason() + .contains("expanded expression value count capacity"), + "unexpected error: {err}" + ); +} diff --git a/crates/rumoca-phase-solve/src/lower/fft.rs b/crates/rumoca-phase-solve/src/lower/fft.rs index 416d6508b..7f2568992 100644 --- a/crates/rumoca-phase-solve/src/lower/fft.rs +++ b/crates/rumoca-phase-solve/src/lower/fft.rs @@ -189,6 +189,27 @@ impl<'a> LowerBuilder<'a> { }; return self.eval_compile_time_expr(value, const_scope); } + if matches!( + name.as_str(), + "Modelica.Utilities.Strings.isEqual" | "Strings.isEqual" | "isEqual" + ) { + let left = args.first().ok_or_else(|| { + unsupported_at( + "Modelica.Utilities.Strings.isEqual requires a left string argument", + span, + ) + })?; + let right = args.get(1).ok_or_else(|| { + unsupported_at( + "Modelica.Utilities.Strings.isEqual requires a right string argument", + span, + ) + })?; + return Ok(bool_to_f64( + self.eval_compile_time_string(left, const_scope)? + == self.eval_compile_time_string(right, const_scope)?, + )); + } if name.last_segment() != "realFFTsamplePoints" { return Err(unsupported_at( "unsupported expression in for-loop range", diff --git a/crates/rumoca-phase-solve/src/lower/function_calls.rs b/crates/rumoca-phase-solve/src/lower/function_calls.rs index b859b663e..815045a91 100644 --- a/crates/rumoca-phase-solve/src/lower/function_calls.rs +++ b/crates/rumoca-phase-solve/src/lower/function_calls.rs @@ -11,13 +11,210 @@ mod random; mod runtime_intrinsics; use helpers::{ ComplexProjectionComprehensionCtx, FlattenedRecordInputRequest, - FlattenedRecordPositionalInputRequest, NamedOrPositionalArg, append_complex_projection_values, - checked_usize_dims_to_i64, complex_projection_vec_with_capacity, flattened_input_has_prefix, - function_input_actual_dim, missing_intrinsic_argument, missing_required_function_input, - record_constructor_field, split_flattened_record_input_name, validate_complex_component_width, + FlattenedRecordPositionalInputRequest, FunctionInputRequest, NamedOrPositionalArg, + append_complex_projection_values, checked_usize_dims_to_i64, + complex_projection_vec_with_capacity, flattened_input_has_prefix, function_input_actual_dim, + missing_intrinsic_argument, missing_required_function_input, record_constructor_field, + split_flattened_record_input_name, synthesize_missing_flattened_record_field_arg, + validate_complex_component_width, }; pub(super) use runtime_intrinsics::*; +pub(super) fn is_flattened_record_field_actual( + expr: &rumoca_core::Expression, + field: &str, +) -> bool { + match expr { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } if subscripts.is_empty() => { + name.as_str().ends_with(&format!(".{field}")) + || name.as_str().ends_with(&format!("_{field}")) + } + rumoca_core::Expression::FieldAccess { + field: actual_field, + .. + } => actual_field == field, + _ => false, + } +} + +pub(super) fn referenced_function_input_names( + function: &rumoca_core::Function, +) -> IndexSet { + let input_names: IndexSet<&str> = function + .inputs + .iter() + .map(|input| input.name.as_str()) + .collect(); + let mut used = IndexSet::new(); + for statement in &function.body { + collect_statement_var_refs(statement, &input_names, &mut used); + } + used +} + +fn collect_statement_var_refs( + statement: &rumoca_core::Statement, + input_names: &IndexSet<&str>, + used: &mut IndexSet, +) { + match statement { + rumoca_core::Statement::Assignment { value, .. } => { + collect_expression_var_refs(value, input_names, used); + } + rumoca_core::Statement::If { + cond_blocks, + else_block, + .. + } => { + for block in cond_blocks { + collect_expression_var_refs(&block.cond, input_names, used); + for statement in &block.stmts { + collect_statement_var_refs(statement, input_names, used); + } + } + if let Some(else_block) = else_block { + for statement in else_block { + collect_statement_var_refs(statement, input_names, used); + } + } + } + rumoca_core::Statement::For { + indices, equations, .. + } => { + for index in indices { + collect_expression_var_refs(&index.range, input_names, used); + } + for statement in equations { + collect_statement_var_refs(statement, input_names, used); + } + } + rumoca_core::Statement::While { block, .. } => { + collect_expression_var_refs(&block.cond, input_names, used); + for statement in &block.stmts { + collect_statement_var_refs(statement, input_names, used); + } + } + rumoca_core::Statement::When { blocks, .. } => { + for block in blocks { + collect_expression_var_refs(&block.cond, input_names, used); + for statement in &block.stmts { + collect_statement_var_refs(statement, input_names, used); + } + } + } + rumoca_core::Statement::FunctionCall { args, .. } => { + for arg in args { + collect_expression_var_refs(arg, input_names, used); + } + } + rumoca_core::Statement::Reinit { value, .. } => { + collect_expression_var_refs(value, input_names, used); + } + rumoca_core::Statement::Assert { + condition, + message, + level, + .. + } => { + collect_expression_var_refs(condition, input_names, used); + collect_expression_var_refs(message, input_names, used); + if let Some(level) = level { + collect_expression_var_refs(level, input_names, used); + } + } + rumoca_core::Statement::Empty { .. } + | rumoca_core::Statement::Return { .. } + | rumoca_core::Statement::Break { .. } => {} + } +} + +fn collect_expression_var_refs( + expr: &rumoca_core::Expression, + input_names: &IndexSet<&str>, + used: &mut IndexSet, +) { + match expr { + rumoca_core::Expression::VarRef { name, .. } => { + if input_names.contains(name.as_str()) { + used.insert(name.as_str().to_string()); + } else { + let flattened_name = name.as_str().replace('.', "_"); + if input_names.contains(flattened_name.as_str()) { + used.insert(flattened_name); + } + } + } + rumoca_core::Expression::Unary { rhs, .. } => { + collect_expression_var_refs(rhs, input_names, used); + } + rumoca_core::Expression::Binary { lhs, rhs, .. } => { + collect_expression_var_refs(lhs, input_names, used); + collect_expression_var_refs(rhs, input_names, used); + } + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + for (condition, value) in branches { + collect_expression_var_refs(condition, input_names, used); + collect_expression_var_refs(value, input_names, used); + } + collect_expression_var_refs(else_branch, input_names, used); + } + rumoca_core::Expression::FunctionCall { args, .. } + | rumoca_core::Expression::BuiltinCall { args, .. } + | rumoca_core::Expression::Tuple { elements: args, .. } + | rumoca_core::Expression::Array { elements: args, .. } => { + for arg in args { + collect_expression_var_refs(arg, input_names, used); + } + } + rumoca_core::Expression::FieldAccess { base, field, .. } => { + if let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = base.as_ref() + && subscripts.is_empty() + { + let flattened_name = format!("{}_{}", name.as_str(), field); + if input_names.contains(flattened_name.as_str()) { + used.insert(flattened_name); + } + } + collect_expression_var_refs(base, input_names, used); + } + rumoca_core::Expression::Index { base, .. } => { + collect_expression_var_refs(base, input_names, used); + } + rumoca_core::Expression::Range { + start, step, end, .. + } => { + collect_expression_var_refs(start, input_names, used); + if let Some(step) = step { + collect_expression_var_refs(step, input_names, used); + } + collect_expression_var_refs(end, input_names, used); + } + rumoca_core::Expression::ArrayComprehension { + expr, + indices, + filter, + .. + } => { + collect_expression_var_refs(expr, input_names, used); + for index in indices { + collect_expression_var_refs(&index.range, input_names, used); + } + if let Some(filter) = filter { + collect_expression_var_refs(filter, input_names, used); + } + } + rumoca_core::Expression::Literal { .. } | rumoca_core::Expression::Empty { .. } => {} + } +} + impl<'a> LowerBuilder<'a> { pub(super) fn lower_qualified_standard_numeric_intrinsic( &mut self, @@ -33,6 +230,19 @@ impl<'a> LowerBuilder<'a> { // solve-IR expressions instead of inlining their function bodies, // so local helper arrays in the library function remain local and // cannot leak into the runtime variable layout. + "Modelica.ComplexMath.real" => { + let arg = args + .first() + .ok_or_else(|| missing_intrinsic_argument(name.as_str(), "argument 1", span))?; + self.lower_expr(arg, scope, call_depth).map(Some) + } + "Modelica.ComplexMath.imag" => { + let arg = args + .get(1) + .or_else(|| args.first()) + .ok_or_else(|| missing_intrinsic_argument(name.as_str(), "argument 1", span))?; + self.lower_expr(arg, scope, call_depth).map(Some) + } "Modelica.Math.Distributions.Uniform.quantile" => self .lower_uniform_quantile(args, span, scope, call_depth) .map(Some), @@ -146,6 +356,11 @@ impl<'a> LowerBuilder<'a> { if let Some(reg) = self.lower_runtime_string_special_intrinsic(call_name, args, span)? { return Ok(Some(reg)); } + if let Some(reg) = self.lower_buildings_energyplus_external_intrinsic( + call_name, args, span, scope, call_depth, + )? { + return Ok(Some(reg)); + } if let Some(reg) = self.lower_external_table_intrinsic(call_name, args, span, scope, call_depth)? { @@ -166,6 +381,31 @@ impl<'a> LowerBuilder<'a> { Ok(None) } + fn lower_buildings_energyplus_external_intrinsic( + &mut self, + call_name: &str, + args: &[rumoca_core::Expression], + span: rumoca_core::Span, + scope: &Scope, + call_depth: usize, + ) -> Result, LowerError> { + let Some(kind) = buildings_energyplus_external_kind(call_name) else { + return Ok(None); + }; + let (named_args, positional_args) = split_named_and_positional_call_args(call_name, args)?; + let mut scalar_args = Vec::new(); + if kind == rumoca_ir_solve::ExternalFunctionKind::BuildingsEnergyPlusInitialize + && let Some(expr) = named_args + .get("isSynchronized") + .copied() + .or_else(|| positional_args.last().copied()) + { + scalar_args.push(self.lower_expr(expr, scope, call_depth)?); + } + self.emit_external_call(kind, &scalar_args, 0, span) + .map(Some) + } + fn lower_synchronous_value_intrinsic( &mut self, call_name: &str, @@ -354,26 +594,170 @@ impl<'a> LowerBuilder<'a> { // compiled kernels stay on the compiled path instead of falling back to // runtime expression evaluation. match external_table_intrinsic_kind(call_name) { - Some(ExternalTableIntrinsicKind::Bounds { upper }) => { - let table_id = self.lower_optional_arg(args, 0, span, scope, call_depth)?; + Some(ExternalTableIntrinsicKind::Bounds { table, upper }) => { + let (table_id, _) = + self.lower_external_table_id_arg(table, args, span, scope, call_depth)?; self.emit_table_bounds(table_id, upper, span).map(Some) } - Some(ExternalTableIntrinsicKind::Lookup) => { - let table_id = self.lower_optional_arg(args, 0, span, scope, call_depth)?; - let column = self.lower_optional_arg(args, 1, span, scope, call_depth)?; - let input = self.lower_optional_arg(args, 2, span, scope, call_depth)?; + Some(ExternalTableIntrinsicKind::Lookup { table }) => { + let (table_id, next_arg_idx) = + self.lower_external_table_id_arg(table, args, span, scope, call_depth)?; + let column = + self.lower_optional_arg(args, next_arg_idx, span, scope, call_depth)?; + let input = + self.lower_optional_arg(args, next_arg_idx + 1, span, scope, call_depth)?; self.emit_table_lookup(table_id, column, input, span) .map(Some) } - Some(ExternalTableIntrinsicKind::NextEvent) => { - let table_id = self.lower_optional_arg(args, 0, span, scope, call_depth)?; - let time = self.lower_optional_arg(args, 1, span, scope, call_depth)?; + Some(ExternalTableIntrinsicKind::NextEvent { table }) => { + let (table_id, next_arg_idx) = + self.lower_external_table_id_arg(table, args, span, scope, call_depth)?; + let time = self.lower_optional_arg(args, next_arg_idx, span, scope, call_depth)?; self.emit_table_next_event(table_id, time, span).map(Some) } _ => Ok(None), } } + fn lower_external_table_id_arg( + &mut self, + table: ExternalTableRecordKind, + args: &[rumoca_core::Expression], + span: rumoca_core::Span, + scope: &Scope, + call_depth: usize, + ) -> Result<(Reg, usize), LowerError> { + let Some(arg) = args.first() else { + return self + .lower_optional_arg(args, 0, span, scope, call_depth) + .map(|reg| (reg, 1)); + }; + if let Some(table_id) = self.structural_table_id_for_constructor_arg(table, arg) { + let consumed = if repeated_external_table_constructor_record_args(table, args) { + table.flattened_field_count() + } else { + 1 + }; + return self + .emit_const_at(table_id, arg.span().unwrap_or(span)) + .map(|reg| (reg, consumed)); + } + if external_table_constructor_call(arg) + && let Some(table_id) = self.eval_external_table_constructor_arg(arg) + { + let consumed = if repeated_external_table_constructor_record_args(table, args) { + table.flattened_field_count() + } else { + 1 + }; + return self + .emit_const_at(table_id, arg.span().unwrap_or(span)) + .map(|reg| (reg, consumed)); + } + if let Some(table_id) = self.structural_table_id_for_flattened_args(table, args) { + return self + .emit_const_at(table_id, arg.span().unwrap_or(span)) + .map(|reg| (reg, table.flattened_field_count())); + } + if let Some(flattened_constructor) = flattened_external_table_constructor(table, args, span) + && let Some(table_id) = self.eval_external_table_constructor_arg(&flattened_constructor) + { + return self + .emit_const_at(table_id, flattened_constructor.span().unwrap_or(span)) + .map(|reg| (reg, table.flattened_field_count())); + } + if let Some(table_id) = self.eval_external_table_constructor_arg(arg) { + return self + .emit_const_at(table_id, arg.span().unwrap_or(span)) + .map(|reg| (reg, 1)); + } + self.lower_optional_arg(args, 0, span, scope, call_depth) + .map(|reg| (reg, 1)) + } + + fn structural_table_id_for_flattened_args( + &self, + table: ExternalTableRecordKind, + args: &[rumoca_core::Expression], + ) -> Option { + if args.len() < table.flattened_field_count() + || external_table_constructor_call(args.first()?) + { + return None; + } + let fields = table.flattened_field_names(); + let first_prefix = external_table_field_prefix(args.first()?, fields.first()?)?; + for (arg, field) in args.iter().zip(fields.iter()).take(3) { + if external_table_field_prefix(arg, field)? != first_prefix { + return None; + } + } + self.structural_bindings + .get(format!("{first_prefix}.tableID").as_str()) + .copied() + } + + fn structural_table_id_for_constructor_arg( + &self, + table: ExternalTableRecordKind, + expr: &rumoca_core::Expression, + ) -> Option { + if !external_table_constructor_call(expr) { + return None; + } + let rumoca_core::Expression::FunctionCall { args, .. } = expr else { + return None; + }; + let mut prefix = None; + let fields = table.flattened_field_names(); + for arg in args { + let Some(field_prefix) = external_table_field_prefix_anywhere(arg, fields) else { + continue; + }; + if prefix + .as_ref() + .is_some_and(|existing_prefix| existing_prefix != &field_prefix) + { + return None; + } + prefix = Some(field_prefix); + } + self.structural_bindings + .get(format!("{}.tableID", prefix?).as_str()) + .copied() + } + + fn eval_external_table_constructor_arg(&self, expr: &rumoca_core::Expression) -> Option { + let env = self.external_table_eval_env(); + rumoca_eval_dae::eval_expr::(expr, &env).ok() + } + + fn external_table_eval_env(&self) -> rumoca_eval_dae::VarEnv { + let mut env = rumoca_eval_dae::VarEnv::new(); + env.functions = Arc::new( + self.functions + .iter() + .map(|(name, func)| (name.as_str().to_string(), func.clone())) + .collect(), + ); + if let Some(starts) = self.variable_starts { + env.start_exprs = Arc::new(starts.clone()); + } + if let Some(clock_intervals) = self.clock_intervals { + env.clock_intervals = Arc::new(clock_intervals.clone()); + } + if let Some(dae_variables) = self.dae_variables { + env.dims = Arc::new(dae_variable_dims(dae_variables)); + } + for (name, value) in self.structural_bindings.iter() { + env.set(name, *value); + } + for (name, value) in &self.local_const_bindings { + env.set(name, *value); + } + env + } + pub(super) fn lower_complex_math_sum_projection( &mut self, call_name: &str, @@ -598,6 +982,24 @@ impl<'a> LowerBuilder<'a> { ) } + pub(super) fn bind_used_function_inputs( + &mut self, + function: &rumoca_core::Function, + args: &[rumoca_core::Expression], + caller_scope: &Scope, + call_depth: usize, + ) -> Result { + let used_inputs = referenced_function_input_names(function); + self.bind_function_inputs_for_name_with_used( + function.name.as_str(), + &function.inputs, + args, + caller_scope, + call_depth, + Some(&used_inputs), + ) + } + pub(super) fn bind_function_inputs_for_name( &mut self, function_name: &str, @@ -605,6 +1007,26 @@ impl<'a> LowerBuilder<'a> { args: &[rumoca_core::Expression], caller_scope: &Scope, call_depth: usize, + ) -> Result { + self.bind_function_inputs_for_name_with_used( + function_name, + inputs, + args, + caller_scope, + call_depth, + None, + ) + } + + #[allow(clippy::excessive_nesting)] + fn bind_function_inputs_for_name_with_used( + &mut self, + function_name: &str, + inputs: &[rumoca_core::FunctionParam], + args: &[rumoca_core::Expression], + caller_scope: &Scope, + call_depth: usize, + used_inputs: Option<&IndexSet>, ) -> Result { let (named_args, positional_args) = split_named_and_positional_call_args(function_name, args)?; @@ -614,32 +1036,63 @@ impl<'a> LowerBuilder<'a> { let mut positional_idx = 0usize; for (input_idx, input) in inputs.iter().enumerate() { - if let Some(arg_expr) = named_args.get(input.name.as_str()) { - self.bind_function_input_arg_or_closure( - FunctionInputBindState { - scope: &mut scope, - const_scope: &mut const_scope, - const_bindings: &mut const_bindings, - }, - function_name, - input, - arg_expr, - caller_scope, - call_depth + 1, - )?; - continue; + if used_inputs.is_some_and(|used_inputs| !used_inputs.contains(&input.name)) { + let Some((prefix, field)) = split_flattened_record_input_name(&input.name) else { + if self.try_bind_function_input( + FunctionInputBindState { + scope: &mut scope, + const_scope: &mut const_scope, + const_bindings: &mut const_bindings, + }, + FunctionInputRequest { + function_name, + input, + inputs, + input_idx, + named_args: &named_args, + positional_args: &positional_args, + positional_idx: &mut positional_idx, + caller_scope, + call_depth: call_depth + 1, + }, + )? { + continue; + } + return missing_required_function_input(function_name, input); + }; + let flattened_group_has_used_sibling = inputs.iter().any(|candidate| { + flattened_input_has_prefix(&candidate.name, prefix) + && used_inputs.is_some_and(|used| used.contains(&candidate.name)) + }); + if flattened_group_has_used_sibling && !named_args.contains_key(input.name.as_str()) + { + let later_flattened_sibling_is_used = + inputs.iter().skip(input_idx + 1).any(|next| { + flattened_input_has_prefix(&next.name, prefix) + && used_inputs.is_some_and(|used| used.contains(&next.name)) + }); + if !later_flattened_sibling_is_used + || positional_args + .get(positional_idx) + .is_some_and(|arg| is_flattened_record_field_actual(arg, field)) + { + positional_idx += usize::from(positional_idx < positional_args.len()); + } + continue; + } } - if self.bind_flattened_record_positional_input( + if self.try_bind_function_input( FunctionInputBindState { scope: &mut scope, const_scope: &mut const_scope, const_bindings: &mut const_bindings, }, - FlattenedRecordPositionalInputRequest { + FunctionInputRequest { function_name, input, inputs, input_idx, + named_args: &named_args, positional_args: &positional_args, positional_idx: &mut positional_idx, caller_scope, @@ -648,59 +1101,6 @@ impl<'a> LowerBuilder<'a> { )? { continue; } - if positional_args.len() > inputs.len() - && let Some(fields) = self.record_constructor_fields(&input.type_name) - && self.bind_flattened_record_function_input( - FunctionInputBindState { - scope: &mut scope, - const_scope: &mut const_scope, - const_bindings: &mut const_bindings, - }, - FlattenedRecordInputRequest { - input, - fields: &fields, - positional_args: &positional_args, - positional_idx: &mut positional_idx, - caller_scope, - call_depth: call_depth + 1, - }, - )? - { - continue; - } - if let Some(arg_expr) = - next_positional_function_input_arg(input, &positional_args, &mut positional_idx) - { - self.bind_function_input_arg_or_closure( - FunctionInputBindState { - scope: &mut scope, - const_scope: &mut const_scope, - const_bindings: &mut const_bindings, - }, - function_name, - input, - arg_expr, - caller_scope, - call_depth + 1, - )?; - continue; - } - if let Some(default) = input.default.as_ref() { - let local_scope = scope.clone(); - self.bind_function_input_arg_or_closure( - FunctionInputBindState { - scope: &mut scope, - const_scope: &mut const_scope, - const_bindings: &mut const_bindings, - }, - function_name, - input, - default, - &local_scope, - call_depth + 1, - )?; - continue; - } return missing_required_function_input(function_name, input); } @@ -711,6 +1111,129 @@ impl<'a> LowerBuilder<'a> { }) } + fn try_bind_function_input( + &mut self, + state: FunctionInputBindState<'_>, + request: FunctionInputRequest<'_, '_>, + ) -> Result { + if let Some(arg_expr) = request.named_args.get(request.input.name.as_str()) { + return self + .bind_function_input_arg_or_closure( + state, + request.function_name, + request.input, + arg_expr, + request.caller_scope, + request.call_depth, + ) + .map(|()| true); + } + self.try_bind_non_named_function_input(state, request) + } + + fn try_bind_non_named_function_input( + &mut self, + state: FunctionInputBindState<'_>, + request: FunctionInputRequest<'_, '_>, + ) -> Result { + if self.bind_flattened_record_positional_input( + FunctionInputBindState { + scope: state.scope, + const_scope: state.const_scope, + const_bindings: state.const_bindings, + }, + FlattenedRecordPositionalInputRequest { + function_name: request.function_name, + input: request.input, + inputs: request.inputs, + input_idx: request.input_idx, + positional_args: request.positional_args, + positional_idx: request.positional_idx, + caller_scope: request.caller_scope, + call_depth: request.call_depth, + }, + )? { + return Ok(true); + } + if request.positional_args.len() > request.inputs.len() + && let Some(fields) = self.record_constructor_fields(&request.input.type_name) + && self.bind_flattened_record_function_input( + FunctionInputBindState { + scope: state.scope, + const_scope: state.const_scope, + const_bindings: state.const_bindings, + }, + FlattenedRecordInputRequest { + input: request.input, + fields: &fields, + positional_args: request.positional_args, + positional_idx: request.positional_idx, + caller_scope: request.caller_scope, + call_depth: request.call_depth, + }, + )? + { + return Ok(true); + } + self.try_bind_simple_function_input(state, request) + } + + fn try_bind_simple_function_input( + &mut self, + state: FunctionInputBindState<'_>, + request: FunctionInputRequest<'_, '_>, + ) -> Result { + if let Some(arg_expr) = next_positional_function_input_arg( + request.input, + request.positional_args, + request.positional_idx, + ) { + return self + .bind_function_input_arg_or_closure( + state, + request.function_name, + request.input, + arg_expr, + request.caller_scope, + request.call_depth, + ) + .map(|()| true); + } + if let Some(default) = request.input.default.as_ref() { + let local_scope = state.scope.clone(); + return self + .bind_function_input_arg_or_closure( + state, + request.function_name, + request.input, + default, + &local_scope, + request.call_depth, + ) + .map(|()| true); + } + if let Some(synthesized) = synthesize_missing_flattened_record_field_arg( + request.input, + request.inputs, + request.input_idx, + request.positional_args, + *request.positional_idx, + ) { + return self + .bind_function_input_arg_or_closure( + state, + request.function_name, + request.input, + &synthesized, + request.caller_scope, + request.call_depth, + ) + .map(|()| true); + } + Ok(false) + } + + #[allow(clippy::excessive_nesting)] fn bind_flattened_record_positional_input( &mut self, state: FunctionInputBindState<'_>, @@ -726,6 +1249,48 @@ impl<'a> LowerBuilder<'a> { return Ok(false); } + if let rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor: false, + .. + } = arg_expr + && let Some(materialized) = self.materialize_single_record_function_call_components( + name, + args, + arg_expr.span().unwrap_or(request.input.span), + request.caller_scope, + request.call_depth, + )? + { + let local_key = request.input.name.clone(); + if let Some(component) = materialized + .components + .into_iter() + .find(|component| component.suffix == field) + { + state + .scope + .insert(generated_scope_key(&local_key), component.reg); + if let Some(dims) = component.dims { + self.local_binding_dims.insert(local_key.clone(), dims); + self.update_known_empty_local_array(&local_key, component.known_empty); + } + self.advance_flattened_record_positional(request, prefix); + return Ok(true); + } + if let Some((_suffix, bindings)) = materialized + .indexed_components + .into_iter() + .find(|(suffix, _bindings)| suffix == field) + { + self.local_indexed_bindings + .insert(local_key.clone(), bindings); + self.advance_flattened_record_positional(request, prefix); + return Ok(true); + } + } + let projected = rumoca_core::Expression::FieldAccess { base: Box::new((*arg_expr).clone()), field: field.to_string(), @@ -739,6 +1304,15 @@ impl<'a> LowerBuilder<'a> { request.caller_scope, request.call_depth, )?; + self.advance_flattened_record_positional(request, prefix); + Ok(true) + } + + fn advance_flattened_record_positional( + &self, + request: FlattenedRecordPositionalInputRequest<'_, '_>, + prefix: &str, + ) { if !request .inputs .iter() @@ -747,7 +1321,6 @@ impl<'a> LowerBuilder<'a> { { *request.positional_idx += 1; } - Ok(true) } fn bind_function_input_arg_or_closure( @@ -840,16 +1413,15 @@ impl<'a> LowerBuilder<'a> { caller_scope: &Scope, ) -> Result, LowerError> { match arg_expr { - rumoca_core::Expression::FunctionCall { - name, - args, - is_constructor: false, - .. - } if self.lookup_function(name).is_some() => Ok(Some(FunctionClosure { - target_name: name.clone(), - bound_args: args.clone(), - captured_scope: caller_scope.clone(), - })), + rumoca_core::Expression::FunctionCall { name, args, .. } + if self.lookup_function(name).is_some() => + { + Ok(Some(FunctionClosure { + target_name: name.clone(), + bound_args: args.clone(), + captured_scope: caller_scope.clone(), + })) + } rumoca_core::Expression::VarRef { name, subscripts, @@ -938,8 +1510,10 @@ impl<'a> LowerBuilder<'a> { self.functions .iter() .find(|(name, function)| { - function.is_constructor - && rumoca_core::qualified_type_name_matches(name.as_str(), type_name) + (function.is_constructor + || is_record_constructor_signature(name.as_str(), function)) + && (rumoca_core::qualified_type_name_matches(name.as_str(), type_name) + || rumoca_core::qualified_type_name_matches(type_name, name.as_str())) }) .map(|(_, function)| function.inputs.clone()) .filter(|fields| !fields.is_empty()) @@ -1081,24 +1655,28 @@ impl<'a> LowerBuilder<'a> { return Ok(false); } - // SPEC_0008: a call recognized as a record constructor must have a - // registered field list; fabricating scalar Real params for missing - // registrations silently mis-binds record fields. - let Some(fields) = self.record_constructor_fields(name.as_str()) else { + let (named_args, positional_args) = + split_named_and_positional_call_args(name.as_str(), args)?; + let fields = self + .record_constructor_fields(name.as_str()) + .or_else(|| self.record_constructor_fields(&input.type_name)); + if fields.is_none() && named_args.is_empty() { return Err(LowerError::InvalidFunction { name: name.as_str().to_string(), reason: "record constructor has no registered field list".to_string(), }); - }; - let (named_args, positional_args) = - split_named_and_positional_call_args(name.as_str(), args)?; + } let mut bound_any = false; for (field_name, arg_expr) in named_args .into_iter() .filter(|(_, arg_expr)| !non_numeric_record_metadata(arg_expr)) { let local_name = format!("{}.{}", input.name, field_name); - let field = record_constructor_field(name.as_str(), &fields, &field_name, *span)?; + let field = if let Some(fields) = fields.as_ref() { + record_constructor_field(name.as_str(), fields, &field_name, *span)? + } else { + self.record_constructor_named_arg_field(&field_name, arg_expr, *span, expr_scope)? + }; self.bind_flattened_record_field_value( FunctionInputBindState { scope: &mut *state.scope, @@ -1114,7 +1692,14 @@ impl<'a> LowerBuilder<'a> { bound_any = true; } - if !positional_args.is_empty() && !fields.is_empty() { + if !positional_args.is_empty() { + let Some(fields) = fields.as_ref() else { + return Err(LowerError::InvalidFunction { + name: name.as_str().to_string(), + reason: "positional record constructor requires a registered field list" + .to_string(), + }); + }; for (field, arg_expr) in fields .iter() .zip(positional_args) @@ -1140,6 +1725,18 @@ impl<'a> LowerBuilder<'a> { Ok(bound_any) } + fn record_constructor_named_arg_field( + &self, + field_name: &str, + arg_expr: &rumoca_core::Expression, + span: rumoca_core::Span, + scope: &Scope, + ) -> Result { + let dims = self.infer_expr_dims(arg_expr, scope)?; + let dims = checked_usize_dims_to_i64(&dims, "record constructor named field shape", span)?; + Ok(rumoca_core::FunctionParam::new(field_name, "Real", span).with_dims(dims)) + } + fn bind_function_input_value( &mut self, mut state: FunctionInputBindState<'_>, @@ -1149,6 +1746,9 @@ impl<'a> LowerBuilder<'a> { expr_scope: &Scope, call_depth: usize, ) -> Result<(), LowerError> { + if self.bind_function_input_closure(function_name, input, expr, expr_scope)? { + return Ok(()); + } if self.bind_record_constructor_function_input( input, expr, expr_scope, &mut state, call_depth, )? { @@ -1190,6 +1790,21 @@ impl<'a> LowerBuilder<'a> { )? { return Ok(()); } + if input.dims.is_empty() { + let inferred_dims = self.infer_expr_dims(expr, expr_scope)?; + if !inferred_dims.is_empty() + && inferred_dims + .iter() + .try_fold(1usize, |total, dim| total.checked_mul(*dim)) + == Some(1) + && let Ok(reg) = self.lower_expr(expr, expr_scope, call_depth) + { + self.insert_optional_compile_time_input_binding(input, expr, &mut state); + self.clear_local_array_metadata(&input.name, expr.span().unwrap_or(input.span))?; + state.scope.insert(generated_scope_key(&input.name), reg); + return Ok(()); + } + } if self .bind_inferred_array_function_input(&mut state, input, expr, expr_scope, call_depth)? { @@ -1210,10 +1825,38 @@ impl<'a> LowerBuilder<'a> { if input.dims.is_empty() { return Ok(false); } - let values = self.lower_array_like_values(expr, expr_scope, call_depth)?; - let binding_dims = - self.resolve_function_input_binding_dims(function_name, input, expr, expr_scope)?; let span = expr.span().unwrap_or(input.span); + let binding_dims = match self.resolve_function_input_binding_dims( + function_name, + input, + expr, + expr_scope, + ) { + Ok(dims) => dims, + Err(err) + if err.is_missing_binding_or_function() + && input.dims.contains(&0) + && input.dims.iter().all(|dim| *dim >= 0) => + { + input.dims.clone() + } + Err(err) => return Err(err), + }; + let values = if binding_dims.iter().any(|dim| *dim <= 0) { + Some(self.lower_array_like_values(expr, expr_scope, call_depth)?) + } else { + None + }; + let binding_dims = if let Some(values) = values.as_ref() { + resolve_array_dims_for_value_count( + &binding_dims, + values.len(), + "function input dynamic shape dimension resolution", + span, + )? + } else { + binding_dims + }; let expected = dims_scalar_count( &binding_dims, format!( @@ -1222,6 +1865,11 @@ impl<'a> LowerBuilder<'a> { ), span, )?; + let values = match values { + Some(values) => values, + None if expected == 0 => Vec::new(), + None => self.lower_array_like_values(expr, expr_scope, call_depth)?, + }; if values.len() != expected { let shape = format_i64_dims(&binding_dims); return Err(LowerError::contract_violation( @@ -1240,6 +1888,14 @@ impl<'a> LowerBuilder<'a> { &binding_dims, span, )?; + self.bind_compile_time_array_input_scalars( + &input.name, + &binding_dims, + expr, + &mut *state.const_scope, + &mut *state.const_bindings, + span, + )?; self.bind_local_array_shape(&input.name, &binding_dims, values.len(), span)?; Ok(true) } @@ -1274,10 +1930,67 @@ impl<'a> LowerBuilder<'a> { &dims, span, )?; + self.bind_compile_time_array_input_scalars( + &input.name, + &dims, + expr, + &mut *state.const_scope, + &mut *state.const_bindings, + span, + )?; self.bind_local_array_shape(&input.name, &dims, values.len(), span)?; Ok(true) } + fn bind_compile_time_array_input_scalars( + &self, + input_name: &str, + dims: &[i64], + expr: &rumoca_core::Expression, + const_scope: &mut IndexMap, + const_bindings: &mut IndexMap, + span: rumoca_core::Span, + ) -> Result<(), LowerError> { + let mut values = Vec::new(); + if self + .collect_compile_time_array_input_scalars(expr, const_scope, &mut values) + .is_err() + { + return Ok(()); + } + let expected = dims_scalar_count(dims, "function input compile-time array shape", span)?; + if values.len() != expected { + return Ok(()); + } + for (flat_index, value) in values.into_iter().enumerate() { + let key = dae::scalar_name_text_for_flat_index(input_name, dims, flat_index); + const_scope.insert(key.clone(), value); + const_bindings.insert(key, value); + } + Ok(()) + } + + fn collect_compile_time_array_input_scalars( + &self, + expr: &rumoca_core::Expression, + const_scope: &IndexMap, + values: &mut Vec, + ) -> Result<(), LowerError> { + match expr { + rumoca_core::Expression::Array { elements, .. } + | rumoca_core::Expression::Tuple { elements, .. } => { + for element in elements { + self.collect_compile_time_array_input_scalars(element, const_scope, values)?; + } + Ok(()) + } + _ => { + values.push(self.eval_compile_time_expr(expr, const_scope)?); + Ok(()) + } + } + } + fn bind_scalar_function_input_value( &mut self, state: FunctionInputBindState<'_>, @@ -1311,6 +2024,19 @@ impl<'a> LowerBuilder<'a> { Ok(()) } + fn insert_optional_compile_time_input_binding( + &self, + input: &rumoca_core::FunctionParam, + expr: &rumoca_core::Expression, + state: &mut FunctionInputBindState<'_>, + ) { + let Ok(value) = self.eval_compile_time_expr(expr, state.const_scope) else { + return; + }; + state.const_scope.insert(input.name.clone(), value); + state.const_bindings.insert(input.name.clone(), value); + } + fn resolve_function_input_binding_dims( &self, function_name: &str, @@ -1333,6 +2059,18 @@ impl<'a> LowerBuilder<'a> { )); } let actual_dims = self.infer_expr_dims(expr, expr_scope)?; + let actual_dims = if actual_dims.len() > input.dims.len() + && actual_dims[..actual_dims.len() - input.dims.len()] + .iter() + .all(|dim| *dim == 1) + { + actual_dims[actual_dims.len() - input.dims.len()..].to_vec() + } else { + actual_dims + }; + if input.dims.as_slice() == [0] && actual_dims.is_empty() { + return Ok(vec![1]); + } if actual_dims.len() != input.dims.len() { let shape = format_i64_dims(&input.dims); return Err(LowerError::contract_violation( @@ -1719,6 +2457,18 @@ impl<'a> LowerBuilder<'a> { result } + pub(super) fn with_local_const_bindings( + &mut self, + const_bindings: &IndexMap, + f: impl FnOnce(&mut Self) -> Result, + ) -> Result { + let saved = self.local_const_bindings.clone(); + self.local_const_bindings.extend(const_bindings.clone()); + let result = f(self); + self.local_const_bindings = saved; + result + } + pub(super) fn initialize_function_output_scope( &mut self, function: &rumoca_core::Function, @@ -1744,6 +2494,13 @@ impl<'a> LowerBuilder<'a> { self.guarded_uninitialized_locals .shift_remove(param.name.as_str()); if param.default.is_some() { + if param.dims.iter().any(|dim| *dim < 0) { + self.local_binding_dims + .insert(param.name.clone(), param.dims.clone()); + let marker = self.emit_const_at(0.0, param.span)?; + scope.insert(generated_scope_key(¶m.name), marker); + return Ok(()); + } let values = self.initial_function_param_values(param, scope, call_depth)?; self.bind_assignment_values_with_dims( scope, @@ -1948,5 +2705,142 @@ impl<'a> LowerBuilder<'a> { } } +fn external_table_constructor_call(expr: &rumoca_core::Expression) -> bool { + let rumoca_core::Expression::FunctionCall { name, .. } = expr else { + return false; + }; + matches!( + name.last_segment(), + "ExternalCombiTimeTable" | "ExternalCombiTable1D" + ) +} + +fn repeated_external_table_constructor_record_args( + table: ExternalTableRecordKind, + args: &[rumoca_core::Expression], +) -> bool { + let field_count = table.flattened_field_count(); + args.len() > field_count + && args[..field_count] + .iter() + .all(external_table_constructor_call) +} + +fn external_table_field_prefix(expr: &rumoca_core::Expression, field: &str) -> Option { + let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = expr + else { + return None; + }; + if !subscripts.is_empty() { + return None; + } + name.as_str() + .strip_suffix(format!(".{field}").as_str()) + .map(str::to_string) +} + +fn external_table_field_prefix_anywhere( + expr: &rumoca_core::Expression, + fields: &[&str], +) -> Option { + let direct = fields + .iter() + .find_map(|field| external_table_field_prefix(expr, field)); + if direct.is_some() { + return direct; + } + match expr { + rumoca_core::Expression::Binary { lhs, rhs, .. } => { + first_matching_external_table_field_prefix([lhs.as_ref(), rhs.as_ref()], fields) + } + rumoca_core::Expression::Unary { rhs, .. } => { + external_table_field_prefix_anywhere(rhs, fields) + } + rumoca_core::Expression::BuiltinCall { args, .. } + | rumoca_core::Expression::FunctionCall { args, .. } + | rumoca_core::Expression::Array { elements: args, .. } + | rumoca_core::Expression::Tuple { elements: args, .. } => { + first_matching_external_table_field_prefix(args.iter(), fields) + } + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => branches + .iter() + .flat_map(|(condition, value)| [condition, value]) + .chain([else_branch.as_ref()]) + .find_map(|expr| external_table_field_prefix_anywhere(expr, fields)), + rumoca_core::Expression::Range { + start, step, end, .. + } => [Some(start.as_ref()), step.as_deref(), Some(end.as_ref())] + .into_iter() + .flatten() + .find_map(|expr| external_table_field_prefix_anywhere(expr, fields)), + rumoca_core::Expression::ArrayComprehension { expr, filter, .. } => { + [Some(expr.as_ref()), filter.as_deref()] + .into_iter() + .flatten() + .find_map(|expr| external_table_field_prefix_anywhere(expr, fields)) + } + rumoca_core::Expression::Index { base, .. } + | rumoca_core::Expression::FieldAccess { base, .. } => { + external_table_field_prefix_anywhere(base, fields) + } + rumoca_core::Expression::VarRef { .. } + | rumoca_core::Expression::Literal { .. } + | rumoca_core::Expression::Empty { .. } => None, + } +} + +fn first_matching_external_table_field_prefix<'a>( + exprs: impl IntoIterator, + fields: &[&str], +) -> Option { + exprs + .into_iter() + .find_map(|expr| external_table_field_prefix_anywhere(expr, fields)) +} + +fn flattened_external_table_constructor( + table: ExternalTableRecordKind, + args: &[rumoca_core::Expression], + span: rumoca_core::Span, +) -> Option { + let field_count = table.flattened_field_count(); + if args.len() < field_count { + return None; + } + if external_table_constructor_call(args.first()?) { + return None; + } + Some(rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from(table.constructor_name()), + args: args[..field_count].to_vec(), + is_constructor: true, + span, + }) +} + +fn dae_variable_dims(variables: &dae::DaeVariables) -> IndexMap> { + let mut dims = IndexMap::new(); + for (name, var) in variables + .states + .iter() + .chain(variables.algebraics.iter()) + .chain(variables.outputs.iter()) + .chain(variables.parameters.iter()) + .chain(variables.constants.iter()) + .chain(variables.inputs.iter()) + .chain(variables.discrete_reals.iter()) + .chain(variables.discrete_valued.iter()) + { + dims.insert(name.as_str().to_string(), var.dims.clone()); + } + dims +} + #[cfg(test)] mod tests; diff --git a/crates/rumoca-phase-solve/src/lower/function_calls/complex_inputs.rs b/crates/rumoca-phase-solve/src/lower/function_calls/complex_inputs.rs index 68966e74c..4fea1dcf2 100644 --- a/crates/rumoca-phase-solve/src/lower/function_calls/complex_inputs.rs +++ b/crates/rumoca-phase-solve/src/lower/function_calls/complex_inputs.rs @@ -1,6 +1,7 @@ use super::*; impl<'a> LowerBuilder<'a> { + #[allow(clippy::too_many_lines)] pub(super) fn bind_complex_input( &mut self, scope: &mut Scope, @@ -61,7 +62,19 @@ impl<'a> LowerBuilder<'a> { span, )?; } - self.bind_local_array_shape(&input.name, &binding_dims, expected, span)?; + self.set_known_local_array_dims(&input.name, binding_dims.clone(), expected, span)?; + self.set_known_local_array_dims( + &format!("{}.re", input.name), + binding_dims.clone(), + expected, + span, + )?; + self.set_known_local_array_dims( + &format!("{}.im", input.name), + binding_dims.clone(), + expected, + span, + )?; return Ok(()); } @@ -102,7 +115,19 @@ impl<'a> LowerBuilder<'a> { expected, span, )?; - self.bind_local_array_shape(&input.name, &binding_dims, expected, span)?; + self.set_known_local_array_dims(&input.name, binding_dims.clone(), expected, span)?; + self.set_known_local_array_dims( + &format!("{}.re", input.name), + binding_dims.clone(), + expected, + span, + )?; + self.set_known_local_array_dims( + &format!("{}.im", input.name), + binding_dims.clone(), + expected, + span, + )?; Ok(()) } @@ -150,6 +175,8 @@ impl<'a> LowerBuilder<'a> { if let Some(im) = im_values.first().copied() { scope.insert(generated_scope_key(format!("{base_name}.im")), im); } + self.bind_assignment_values_at(scope, &format!("{base_name}.re"), re_values, span)?; + self.bind_assignment_values_at(scope, &format!("{base_name}.im"), im_values, span)?; self.bind_assignment_values_at(scope, &format!("{base_name}[:].re"), re_values, span)?; self.bind_assignment_values_at(scope, &format!("{base_name}[:].im"), im_values, span)?; self.bind_assignment_values_at(scope, &format!("{base_name}[:].re.re"), re_values, span)?; diff --git a/crates/rumoca-phase-solve/src/lower/function_calls/helpers.rs b/crates/rumoca-phase-solve/src/lower/function_calls/helpers.rs index b491687fa..e42cf95ed 100644 --- a/crates/rumoca-phase-solve/src/lower/function_calls/helpers.rs +++ b/crates/rumoca-phase-solve/src/lower/function_calls/helpers.rs @@ -30,6 +30,18 @@ pub(super) struct FlattenedRecordPositionalInputRequest<'a, 'b> { pub(super) call_depth: usize, } +pub(super) struct FunctionInputRequest<'a, 'b> { + pub(super) function_name: &'a str, + pub(super) input: &'a rumoca_core::FunctionParam, + pub(super) inputs: &'a [rumoca_core::FunctionParam], + pub(super) input_idx: usize, + pub(super) named_args: &'a IndexMap, + pub(super) positional_args: &'a [&'a rumoca_core::Expression], + pub(super) positional_idx: &'b mut usize, + pub(super) caller_scope: &'a Scope, + pub(super) call_depth: usize, +} + pub(super) struct NamedOrPositionalArg<'a> { pub(super) name: &'a str, pub(super) idx: usize, @@ -57,6 +69,66 @@ pub(super) fn missing_required_function_input( }) } +pub(super) fn synthesize_missing_flattened_record_field_arg( + input: &rumoca_core::FunctionParam, + inputs: &[rumoca_core::FunctionParam], + input_idx: usize, + positional_args: &[&rumoca_core::Expression], + positional_idx: usize, +) -> Option { + let (prefix, field) = split_flattened_record_input_name(&input.name)?; + let search_len = positional_idx.min(positional_args.len()).min(input_idx); + for previous_idx in (0..search_len).rev() { + let (previous_prefix, previous_field) = + split_flattened_record_input_name(&inputs.get(previous_idx)?.name)?; + if previous_prefix != prefix { + continue; + } + let base = + flattened_record_field_actual_base(positional_args[previous_idx], previous_field)?; + return record_field_access_expr(base, field); + } + None +} + +fn flattened_record_field_actual_base( + expr: &rumoca_core::Expression, + field: &str, +) -> Option { + match expr { + rumoca_core::Expression::VarRef { + name, + subscripts, + span, + } if subscripts.is_empty() => { + let base = name.as_str().strip_suffix(&format!(".{field}"))?; + Some(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new(base), + subscripts: Vec::new(), + span: *span, + }) + } + rumoca_core::Expression::FieldAccess { + base, + field: actual_field, + .. + } if actual_field == field => Some((**base).clone()), + _ => None, + } +} + +fn record_field_access_expr( + base: rumoca_core::Expression, + field: &str, +) -> Option { + let span = base.span()?; + Some(rumoca_core::Expression::FieldAccess { + span, + base: Box::new(base), + field: field.to_string(), + }) +} + pub(super) fn missing_intrinsic_argument( function_name: &str, argument: &'static str, diff --git a/crates/rumoca-phase-solve/src/lower/function_calls/runtime_intrinsics.rs b/crates/rumoca-phase-solve/src/lower/function_calls/runtime_intrinsics.rs index e798ee2d7..f3572a6a9 100644 --- a/crates/rumoca-phase-solve/src/lower/function_calls/runtime_intrinsics.rs +++ b/crates/rumoca-phase-solve/src/lower/function_calls/runtime_intrinsics.rs @@ -1,30 +1,126 @@ use super::*; -use rumoca_ir_solve::RandomGenerator; +use rumoca_ir_solve::{ExternalFunctionKind, RandomGenerator}; #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub(in crate::lower) enum ExternalTableIntrinsicKind { - Bounds { upper: bool }, - Lookup, - NextEvent, + Bounds { + table: ExternalTableRecordKind, + upper: bool, + }, + Lookup { + table: ExternalTableRecordKind, + }, + NextEvent { + table: ExternalTableRecordKind, + }, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(in crate::lower) enum ExternalTableRecordKind { + CombiTimeTable, + CombiTable1D, +} + +impl ExternalTableRecordKind { + pub(in crate::lower) fn constructor_name(self) -> &'static str { + match self { + Self::CombiTimeTable => "Modelica.Blocks.Types.ExternalCombiTimeTable", + Self::CombiTable1D => "Modelica.Blocks.Types.ExternalCombiTable1D", + } + } + + pub(in crate::lower) fn flattened_field_count(self) -> usize { + self.flattened_field_names().len() + } + + pub(in crate::lower) fn flattened_field_names(self) -> &'static [&'static str] { + match self { + Self::CombiTimeTable => &[ + "tableName", + "fileName", + "table", + "startTime", + "columns", + "smoothness", + "extrapolation", + "shiftTime", + "timeEvents", + "verboseRead", + "delimiter", + "nHeaderLines", + ], + Self::CombiTable1D => &[ + "tableName", + "fileName", + "table", + "columns", + "smoothness", + "extrapolation", + "verboseRead", + "delimiter", + "nHeaderLines", + ], + } + } } pub(in crate::lower) fn external_table_intrinsic_kind( call_name: &str, ) -> Option { - match intrinsic_short_name(call_name) { - "getTimeTableTmin" | "getTable1DAbscissaUmin" => { - Some(ExternalTableIntrinsicKind::Bounds { upper: false }) + let short_name = match crate::path_utils::scope_split(call_name) { + Some((function_name, "y")) => intrinsic_short_name(function_name), + _ => intrinsic_short_name(call_name), + }; + match short_name { + "getTimeTableTmin" => Some(ExternalTableIntrinsicKind::Bounds { + table: ExternalTableRecordKind::CombiTimeTable, + upper: false, + }), + "getTable1DAbscissaUmin" => Some(ExternalTableIntrinsicKind::Bounds { + table: ExternalTableRecordKind::CombiTable1D, + upper: false, + }), + "getTimeTableTmax" => Some(ExternalTableIntrinsicKind::Bounds { + table: ExternalTableRecordKind::CombiTimeTable, + upper: true, + }), + "getTable1DAbscissaUmax" => Some(ExternalTableIntrinsicKind::Bounds { + table: ExternalTableRecordKind::CombiTable1D, + upper: true, + }), + "getTimeTableValueNoDer" | "getTimeTableValueNoDer2" | "getTimeTableValue" => { + Some(ExternalTableIntrinsicKind::Lookup { + table: ExternalTableRecordKind::CombiTimeTable, + }) + } + "getTable1DValueNoDer" | "getTable1DValueNoDer2" | "getTable1DValue" => { + Some(ExternalTableIntrinsicKind::Lookup { + table: ExternalTableRecordKind::CombiTable1D, + }) + } + "getNextTimeEvent" => Some(ExternalTableIntrinsicKind::NextEvent { + table: ExternalTableRecordKind::CombiTimeTable, + }), + _ => None, + } +} + +pub(in crate::lower) fn buildings_energyplus_external_kind( + call_name: &str, +) -> Option { + match call_name { + "Buildings.ThermalZones.EnergyPlus_9_6_0.BaseClasses.initialize" => { + Some(ExternalFunctionKind::BuildingsEnergyPlusInitialize) + } + "Buildings.ThermalZones.EnergyPlus_9_6_0.BaseClasses.getParameters" => { + Some(ExternalFunctionKind::BuildingsEnergyPlusGetParameters) + } + "Buildings.ThermalZones.EnergyPlus_9_6_0.BaseClasses.exchange" => { + Some(ExternalFunctionKind::BuildingsEnergyPlusExchange) } - "getTimeTableTmax" | "getTable1DAbscissaUmax" => { - Some(ExternalTableIntrinsicKind::Bounds { upper: true }) + "Buildings.ThermalZones.EnergyPlus_9_6_0.BaseClasses.SpawnExternalObject" => { + Some(ExternalFunctionKind::BuildingsEnergyPlusSpawnExternalObject) } - "getTimeTableValueNoDer" - | "getTimeTableValueNoDer2" - | "getTimeTableValue" - | "getTable1DValueNoDer" - | "getTable1DValueNoDer2" - | "getTable1DValue" => Some(ExternalTableIntrinsicKind::Lookup), - "getNextTimeEvent" => Some(ExternalTableIntrinsicKind::NextEvent), _ => None, } } diff --git a/crates/rumoca-phase-solve/src/lower/function_calls/tests.rs b/crates/rumoca-phase-solve/src/lower/function_calls/tests.rs index a94636c05..4cccdfe3c 100644 --- a/crates/rumoca-phase-solve/src/lower/function_calls/tests.rs +++ b/crates/rumoca-phase-solve/src/lower/function_calls/tests.rs @@ -1,5 +1,103 @@ use super::*; +fn test_span() -> rumoca_core::Span { + rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("phase_solve_lower_function_calls_tests.mo"), + 1, + 2, + ) +} + +fn real(value: f64) -> rumoca_core::Expression { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(value), + span: test_span(), + } +} + +fn int(value: i64) -> rumoca_core::Expression { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + span: test_span(), + } +} + +fn bool_lit(value: bool) -> rumoca_core::Expression { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Boolean(value), + span: test_span(), + } +} + +fn string(value: &str) -> rumoca_core::Expression { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String(value.to_string()), + span: test_span(), + } +} + +fn var_ref(name: &str) -> rumoca_core::Expression { + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new(name), + subscripts: Vec::new(), + span: test_span(), + } +} + +fn array(elements: Vec) -> rumoca_core::Expression { + array_with_matrix_flag(elements, false) +} + +fn matrix(elements: Vec) -> rumoca_core::Expression { + array_with_matrix_flag(elements, true) +} + +fn array_with_matrix_flag( + elements: Vec, + is_matrix: bool, +) -> rumoca_core::Expression { + rumoca_core::Expression::Array { + elements, + is_matrix, + span: test_span(), + } +} + +fn external_time_table_constructor() -> rumoca_core::Expression { + external_time_table_constructor_with( + matrix(vec![ + array(vec![real(0.0), real(2.0)]), + array(vec![real(1.0), real(4.0)]), + ]), + real(0.0), + ) +} + +fn external_time_table_constructor_with( + table: rumoca_core::Expression, + start_time: rumoca_core::Expression, +) -> rumoca_core::Expression { + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from("Modelica.Blocks.Types.ExternalCombiTimeTable"), + args: vec![ + string("NoName"), + string("NoName"), + table, + start_time, + array(vec![int(2)]), + int(1), + int(1), + real(0.0), + int(1), + bool_lit(false), + string(","), + int(0), + ], + is_constructor: true, + span: test_span(), + } +} + #[test] fn checked_usize_dims_to_i64_rejects_overflow_with_span() { let Some(dim) = usize::try_from(i64::MAX) @@ -27,3 +125,215 @@ fn checked_usize_dims_to_i64_rejects_overflow_with_span() { ) ); } + +#[test] +fn lower_external_table_lookup_accepts_flattened_time_table_record_args() { + let table = matrix(vec![ + array(vec![real(0.0), real(2.0)]), + array(vec![real(1.0), real(4.0)]), + ]); + let columns = array(vec![int(2)]); + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from( + "Modelica.Blocks.Tables.Internal.getTimeTableValueNoDer", + ), + args: vec![ + string("NoName"), + string("NoName"), + table, + real(0.0), + columns, + int(1), + int(1), + real(0.0), + int(1), + bool_lit(false), + string(","), + int(0), + int(1), + real(0.5), + real(1.0), + real(0.0), + ], + is_constructor: false, + span: test_span(), + }; + + let lowered = lower_expression( + &expr, + &rumoca_ir_solve::VarLayout::default(), + &IndexMap::new(), + ) + .expect("flattened table record args should lower as external table lookup"); + + assert!( + lowered + .ops + .iter() + .any(|op| matches!(op, rumoca_ir_solve::LinearOp::TableLookup { .. })), + "lowered ops should include a host-backed table lookup: {:?}", + lowered.ops + ); +} + +#[test] +fn lower_external_table_lookup_consumes_repeated_flattened_constructor_record_args() { + let constructor = external_time_table_constructor(); + let mut args = + vec![constructor; ExternalTableRecordKind::CombiTimeTable.flattened_field_count()]; + args.extend([int(1), real(0.5), real(1.0), real(0.0)]); + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from( + "Modelica.Blocks.Tables.Internal.getTimeTableValueNoDer", + ), + args, + is_constructor: false, + span: test_span(), + }; + + let lowered = lower_expression( + &expr, + &rumoca_ir_solve::VarLayout::default(), + &IndexMap::new(), + ) + .expect("repeated flattened constructor record args should lower"); + + assert!( + lowered + .ops + .iter() + .any(|op| matches!(op, rumoca_ir_solve::LinearOp::TableLookup { .. })), + "lowered ops should include a host-backed table lookup: {:?}", + lowered.ops + ); +} + +#[test] +fn lower_external_table_lookup_accepts_projected_output_with_constructor_arg() { + let mut args = vec![external_time_table_constructor()]; + args.extend([int(1), real(0.5), real(1.0), real(0.0)]); + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from( + "Modelica.Blocks.Tables.Internal.getTimeTableValueNoDer.y", + ), + args, + is_constructor: false, + span: test_span(), + }; + + let lowered = lower_expression( + &expr, + &rumoca_ir_solve::VarLayout::default(), + &IndexMap::new(), + ) + .expect("projected table lookup with constructor arg should lower"); + + assert!( + lowered + .ops + .iter() + .any(|op| matches!(op, rumoca_ir_solve::LinearOp::TableLookup { .. })), + "lowered ops should include a host-backed table lookup: {:?}", + lowered.ops + ); +} + +#[test] +fn lower_external_table_lookup_reuses_structural_table_id_for_flattened_record_fields() { + let mut args = vec![ + var_ref("block.table.tableName"), + var_ref("block.table.fileName"), + var_ref("block.table.table"), + var_ref("block.table.startTime"), + var_ref("block.table.columns"), + var_ref("block.table.smoothness"), + var_ref("block.table.extrapolation"), + var_ref("block.table.shiftTime"), + var_ref("block.table.timeEvents"), + var_ref("block.table.verboseRead"), + var_ref("block.table.delimiter"), + var_ref("block.table.nHeaderLines"), + ]; + args.extend([int(1), real(0.5), real(1.0), real(0.0)]); + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from( + "Modelica.Blocks.Tables.Internal.getTimeTableValueNoDer", + ), + args, + is_constructor: false, + span: test_span(), + }; + let structural_bindings = Arc::new(IndexMap::from([( + "block.table.tableID".to_string(), + 12345.0, + )])); + let layout = rumoca_ir_solve::VarLayout::default(); + let functions = IndexMap::new(); + let mut builder = + LowerBuilder::new(&layout, &functions).with_structural_bindings(structural_bindings); + + builder + .lower_expr(&expr, &Scope::new(), 0) + .expect("flattened table record should lower via structural tableID"); + + assert!( + builder.ops.iter().any( + |op| matches!(op, rumoca_ir_solve::LinearOp::Const { value, .. } if *value == 12345.0) + ), + "lowered ops should load the structural table id: {:?}", + builder.ops + ); + assert!( + builder + .ops + .iter() + .any(|op| matches!(op, rumoca_ir_solve::LinearOp::TableLookup { .. })), + "lowered ops should include a host-backed table lookup: {:?}", + builder.ops + ); +} + +#[test] +fn lower_external_table_lookup_reuses_structural_table_id_for_repeated_constructor_record_args() { + let constructor = + external_time_table_constructor_with(matrix(Vec::new()), var_ref("block.table.startTime")); + let mut args = + vec![constructor; ExternalTableRecordKind::CombiTimeTable.flattened_field_count()]; + args.extend([int(1), real(0.5), real(1.0), real(0.0)]); + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from( + "Modelica.Blocks.Tables.Internal.getTimeTableValueNoDer", + ), + args, + is_constructor: false, + span: test_span(), + }; + let structural_bindings = Arc::new(IndexMap::from([( + "block.table.tableID".to_string(), + 12345.0, + )])); + let layout = rumoca_ir_solve::VarLayout::default(); + let functions = IndexMap::new(); + let mut builder = + LowerBuilder::new(&layout, &functions).with_structural_bindings(structural_bindings); + + builder + .lower_expr(&expr, &Scope::new(), 0) + .expect("repeated constructor record args should reuse structural tableID"); + + assert!( + builder.ops.iter().any( + |op| matches!(op, rumoca_ir_solve::LinearOp::Const { value, .. } if *value == 12345.0) + ), + "lowered ops should load the structural table id: {:?}", + builder.ops + ); + assert!( + builder + .ops + .iter() + .any(|op| matches!(op, rumoca_ir_solve::LinearOp::TableLookup { .. })), + "lowered ops should include a host-backed table lookup: {:?}", + builder.ops + ); +} diff --git a/crates/rumoca-phase-solve/src/lower/function_dispatch.rs b/crates/rumoca-phase-solve/src/lower/function_dispatch.rs index a859bd457..d6398ec05 100644 --- a/crates/rumoca-phase-solve/src/lower/function_dispatch.rs +++ b/crates/rumoca-phase-solve/src/lower/function_dispatch.rs @@ -1,6 +1,7 @@ use super::*; impl<'a> LowerBuilder<'a> { + #[allow(clippy::too_many_lines)] pub(super) fn lower_function_call( &mut self, name: &rumoca_core::Reference, @@ -13,14 +14,26 @@ impl<'a> LowerBuilder<'a> { if self.is_record_constructor_call(name, is_constructor) { let (named_args, positional_args) = function_calls::split_named_and_positional_call_args(name.as_str(), args)?; + if let Some(first_input) = self + .lookup_function(name) + .and_then(|constructor| constructor.inputs.first()) + && let Some(expr) = named_args + .get(first_input.name.as_str()) + .copied() + .or(first_input.default.as_ref()) + { + return self.lower_expr(expr, caller_scope, call_depth + 1); + } if let Some(expr) = named_args .get("re") .copied() + .or_else(|| named_args.get("start").copied()) .or_else(|| positional_args.first().copied()) { // Modelica.Complex and other scalar record constructors use // declared field order; numeric scalar contexts read the first - // field unless a projection selects another component. + // field unless a projection selects another component. Scalar + // type constructors may carry only attributes such as start. return self.lower_expr(expr, caller_scope, call_depth + 1); } if let Some(default_expr) = self @@ -106,19 +119,140 @@ impl<'a> LowerBuilder<'a> { )); } + let projection_candidate = + function.pure && is_projected_scalar_function_candidate(&function); + if projection_candidate + && let Some(reg) = self.lower_projected_scalar_function_call( + name, + args, + span, + caller_scope, + call_depth, + )? + { + return Ok(reg); + } + self.with_local_lower_frame(|this| { let bindings = - this.bind_function_inputs(name, &function.inputs, args, caller_scope, call_depth)?; + this.bind_used_function_inputs(&function, args, caller_scope, call_depth)?; let mut scope = bindings.scope; this.local_const_bindings.extend(bindings.const_bindings); this.initialize_function_output_scope(&function, &mut scope, call_depth)?; let _returned = this.lower_statements(&function.body, &mut scope, call_depth + 1)?; - this.lower_scalar_function_output_value(name.as_str(), &function, &scope, span) }) } + fn lower_projected_scalar_function_call( + &mut self, + name: &rumoca_core::Reference, + args: &[rumoca_core::Expression], + span: rumoca_core::Span, + caller_scope: &Scope, + call_depth: usize, + ) -> Result, LowerError> { + let mut dae_model = dae::Dae::default(); + dae_model.symbols.functions = self.functions.clone(); + if let Some(variables) = self.dae_variables { + dae_model.variables = variables.clone(); + } + let expr = rumoca_core::Expression::FunctionCall { + name: name.clone(), + args: args.to_vec(), + is_constructor: false, + span, + }; + let Some(mut values) = (match derivative_rhs::function_call_projected_scalars_with_owner( + &expr, + &dae_model, + self.structural_bindings.as_ref(), + span, + ) { + Ok(values) => values, + Err(_) => return Ok(None), + }) else { + return Ok(None); + }; + if values.len() != 1 { + return Ok(None); + } + let value = values.remove(0); + if matches!( + &value, + rumoca_core::Expression::FunctionCall { + name: projected_name, + args: projected_args, + is_constructor: false, + .. + } if projected_name == name && projected_args == args + ) { + return Ok(None); + } + match self.lower_expr(&value, caller_scope, call_depth + 1) { + Ok(reg) => Ok(Some(reg)), + Err(_) => Ok(None), + } + } + + pub(super) fn lower_projected_scalar_function_call_values( + &mut self, + name: &rumoca_core::Reference, + args: &[rumoca_core::Expression], + span: rumoca_core::Span, + caller_scope: &Scope, + call_depth: usize, + ) -> Result>, LowerError> { + let mut dae_model = dae::Dae::default(); + dae_model.symbols.functions = self.functions.clone(); + if let Some(variables) = self.dae_variables { + dae_model.variables = variables.clone(); + } + let expr = rumoca_core::Expression::FunctionCall { + name: name.clone(), + args: args.to_vec(), + is_constructor: false, + span, + }; + let Some(values) = (match derivative_rhs::function_call_projected_scalars_with_owner( + &expr, + &dae_model, + self.structural_bindings.as_ref(), + span, + ) { + Ok(values) => values, + Err(_) => return Ok(None), + }) else { + return Ok(None); + }; + if values.len() == 1 + && matches!( + &values[0], + rumoca_core::Expression::FunctionCall { + name: projected_name, + args: projected_args, + is_constructor: false, + .. + } if projected_name == name && projected_args == args + ) + { + return Ok(None); + } + let mut regs = + crate::lower_vec_with_capacity(values.len(), "projected function value count", span)?; + for value in values { + let array_values = + self.lower_array_like_values(&value, caller_scope, call_depth + 1)?; + if array_values.len() == 1 { + regs.push(array_values[0]); + } else { + regs.extend(array_values); + } + } + Ok(Some(regs)) + } + pub(super) fn lookup_function_closure( &self, name: &rumoca_core::Reference, @@ -134,7 +268,10 @@ impl<'a> LowerBuilder<'a> { return Ok(None); } let key = self.scope_key_from_reference(name, span)?; - Ok(self.function_closures.get(&key)) + Ok(self.function_closures.get(&key).or_else(|| { + self.function_closures + .get(&generated_scope_key(name.as_str())) + })) } pub(super) fn lower_function_closure_call( @@ -255,3 +392,59 @@ impl<'a> LowerBuilder<'a> { .is_some_and(|function| is_record_constructor_signature(name, function)) } } + +fn is_projected_scalar_function_candidate(function: &rumoca_core::Function) -> bool { + function.body.iter().all(|statement| match statement { + rumoca_core::Statement::Empty { .. } | rumoca_core::Statement::Return { .. } => true, + rumoca_core::Statement::Assignment { value, .. } => !expr_contains_size_builtin(value), + _ => false, + }) +} + +fn expr_contains_size_builtin(expr: &rumoca_core::Expression) -> bool { + match expr { + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Size, + .. + } => true, + rumoca_core::Expression::BuiltinCall { args, .. } + | rumoca_core::Expression::FunctionCall { args, .. } => { + args.iter().any(expr_contains_size_builtin) + } + rumoca_core::Expression::Unary { rhs, .. } => expr_contains_size_builtin(rhs), + rumoca_core::Expression::Binary { lhs, rhs, .. } => { + expr_contains_size_builtin(lhs) || expr_contains_size_builtin(rhs) + } + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + branches.iter().any(|(condition, value)| { + expr_contains_size_builtin(condition) || expr_contains_size_builtin(value) + }) || expr_contains_size_builtin(else_branch) + } + rumoca_core::Expression::Array { elements, .. } + | rumoca_core::Expression::Tuple { elements, .. } => { + elements.iter().any(expr_contains_size_builtin) + } + rumoca_core::Expression::Index { + base, subscripts, .. + } => { + expr_contains_size_builtin(base) + || subscripts.iter().any(|subscript| match subscript { + rumoca_core::Subscript::Expr { expr, .. } => expr_contains_size_builtin(expr), + _ => false, + }) + } + rumoca_core::Expression::FieldAccess { base, .. } => expr_contains_size_builtin(base), + rumoca_core::Expression::Range { + start, step, end, .. + } => { + expr_contains_size_builtin(start) + || step.as_deref().is_some_and(expr_contains_size_builtin) + || expr_contains_size_builtin(end) + } + _ => false, + } +} diff --git a/crates/rumoca-phase-solve/src/lower/function_projection.rs b/crates/rumoca-phase-solve/src/lower/function_projection.rs index a1dcbf143..cc4b1f471 100644 --- a/crates/rumoca-phase-solve/src/lower/function_projection.rs +++ b/crates/rumoca-phase-solve/src/lower/function_projection.rs @@ -60,10 +60,16 @@ impl<'a> LowerBuilder<'a> { } else { return Ok(None); }; - if let Some(field) = output_field.as_deref() - && (!output_is_complex_record(output) || !matches!(field, "re" | "im")) - { - return Ok(None); + if let Some(field) = output_field.as_deref() { + let constructor_field = function.is_constructor + && output.type_class == Some(rumoca_core::ClassType::Record) + && function.inputs.iter().any(|input| input.name == field); + let ordinary_complex_field = !function.is_constructor + && output_is_complex_record(output) + && matches!(field, "re" | "im"); + if !constructor_field && !ordinary_complex_field { + return Ok(None); + } } let indices = if output_has_dynamic_dims(output) { @@ -146,6 +152,16 @@ impl<'a> LowerBuilder<'a> { )); } + if function.is_constructor { + return self.lower_record_constructor_output_projection( + &function, + projection, + args, + caller_scope, + call_depth, + ); + } + if projection.indices.is_empty() && projection.output_field.is_none() { let values = self.lower_user_function_call_named_output_values( &projection.base_function_name, @@ -194,6 +210,68 @@ impl<'a> LowerBuilder<'a> { }) } + fn lower_record_constructor_output_projection( + &mut self, + constructor: &rumoca_core::Function, + projection: &FunctionOutputProjection, + args: &[rumoca_core::Expression], + caller_scope: &Scope, + call_depth: usize, + ) -> Result { + let Some(field) = projection.output_field.as_deref() else { + return Err(LowerError::InvalidFunction { + name: constructor.name.as_str().to_string(), + reason: format!( + "constructor output `{}` requires a selected record field", + projection.output_name + ), + } + .with_fallback_span(projection.span)); + }; + let output_is_record = constructor.outputs.iter().any(|output| { + output.name == projection.output_name + && output.type_class == Some(rumoca_core::ClassType::Record) + }); + if !output_is_record || !projection.indices.is_empty() { + return Err(LowerError::InvalidFunction { + name: constructor.name.as_str().to_string(), + reason: format!( + "constructor output field `{}.{field}` cannot be resolved", + projection.output_name + ), + } + .with_fallback_span(projection.span)); + } + let Some(input) = constructor.inputs.iter().find(|input| input.name == field) else { + return Err(LowerError::InvalidFunction { + name: constructor.name.as_str().to_string(), + reason: format!("constructor does not define field `{field}`"), + } + .with_fallback_span(projection.span)); + }; + + self.with_local_lower_frame(|this| { + let bindings = this.bind_function_inputs_for_name( + constructor.name.as_str(), + &constructor.inputs, + args, + caller_scope, + call_depth, + )?; + bindings + .scope + .get(&generated_scope_key(&input.name)) + .copied() + .ok_or_else(|| { + LowerError::InvalidFunction { + name: constructor.name.as_str().to_string(), + reason: format!("constructor field `{field}` has no bound scalar value"), + } + .with_fallback_span(projection.span) + }) + }) + } + pub(super) fn bind_assignment_values_at( &mut self, scope: &mut Scope, diff --git a/crates/rumoca-phase-solve/src/lower/helpers.rs b/crates/rumoca-phase-solve/src/lower/helpers.rs index 386b1f451..6a26cd2dc 100644 --- a/crates/rumoca-phase-solve/src/lower/helpers.rs +++ b/crates/rumoca-phase-solve/src/lower/helpers.rs @@ -254,16 +254,19 @@ pub(super) fn indexed_key_for_reference( { return Ok(contained(key)); } - let key = ComponentReferenceKey::from_component_reference(component_ref).map_err(|err| { - LowerError::contract_violation( + match ComponentReferenceKey::from_component_reference(component_ref) { + Ok(key) => Ok(contained(key)), + Err(err) if err.kind == rumoca_ir_solve::ComponentReferenceKeyErrorKind::MissingDefId => { + Ok(None) + } + Err(err) => Err(LowerError::contract_violation( format!( "indexed solve-layout lookup for `{}` has non-static component reference: {err}", reference.as_str(), ), err.span, - ) - })?; - Ok(contained(key)) + )), + } } pub(super) fn parse_indexed_binding_key(key: &str) -> Option<(String, Vec)> { @@ -447,7 +450,13 @@ pub(super) fn resolve_array_dims_for_value_count( context: &'static str, span: rumoca_core::Span, ) -> Result, LowerError> { - let unknown_count = dims.iter().filter(|dim| **dim <= 0).count(); + if let Some(dim) = dims.iter().find(|dim| **dim < 0) { + return Err(LowerError::contract_violation( + format!("{context} has negative dimension in declared shape {dims:?}: {dim}"), + span, + )); + } + let unknown_count = dims.iter().filter(|dim| **dim == 0).count(); if dims.is_empty() || unknown_count != 1 || value_count == 0 { return copy_i64_dims(dims, context, span); } @@ -541,8 +550,20 @@ pub(super) fn static_subscript_indices_with_owner( pub(super) fn is_static_singleton_scalar_projection( base: &rumoca_core::Expression, subscripts: &[rumoca_core::Subscript], + owner_span: Option, ) -> Result { - let owner_span = base.require_span("singleton scalar projection")?.span(); + if subscripts + .iter() + .any(|subscript| matches!(subscript, rumoca_core::Subscript::Colon { .. })) + { + return Ok(false); + } + let owner_span = + base.span() + .or(owner_span) + .ok_or_else(|| LowerError::UnspannedContractViolation { + reason: "missing source provenance for singleton scalar projection".to_string(), + })?; let Some(indices) = static_subscript_indices_with_owner(subscripts, owner_span)? else { return Ok(false); }; @@ -734,6 +755,9 @@ pub(super) fn lower_static_index_expr_with_owner( index_expr_owner_span(expr, owner_span, "static subscript expression")?, )?)); } + if matches!(expr, rumoca_core::Expression::VarRef { .. }) { + return Ok(None); + } Err(unsupported_at( "subscript expression did not evaluate to a positive integer", @@ -958,6 +982,31 @@ pub(super) fn compile_time_index_expr_with_owner( )) } +pub(super) fn compile_time_non_negative_dimension_expr_with_owner( + expr: &rumoca_core::Expression, + const_scope: &IndexMap, + owner_span: rumoca_core::Span, + context: &str, +) -> Result { + let raw = compile_time_index_raw_with_owner(expr, const_scope, owner_span)?; + let rounded = raw.round(); + let span = index_expr_owner_span(expr, owner_span, "compile-time dimension expression")?; + if rounded.is_finite() && rounded >= 0.0 && (rounded - raw).abs() < f64::EPSILON { + if rounded < usize::MAX as f64 { + // Bounds and integrality are checked above; Rust has no TryFrom. + return Ok(rounded as usize); + } + return Err(LowerError::contract_violation( + format!("{context} {rounded} exceeds host index range"), + span, + )); + } + Err(unsupported_at( + format!("{context} must be non-negative"), + span, + )) +} + pub(in crate::lower) fn positive_i64_index( value: i64, span: rumoca_core::Span, @@ -1081,18 +1130,33 @@ fn compile_time_index_builtin( const_scope: &IndexMap, span: rumoca_core::Span, ) -> Result { - let Some(arg) = args.first() else { - return Err(unsupported_at( - "compile-time subscript builtin requires an argument", - span, - )); - }; - let value = compile_time_index_raw_with_owner(arg, const_scope, span)?; match function { rumoca_core::BuiltinFunction::Floor | rumoca_core::BuiltinFunction::Integer => { + let value = compile_time_unary_numeric_builtin_arg(args, const_scope, span)?; Ok(value.floor()) } - rumoca_core::BuiltinFunction::Ceil => Ok(value.ceil()), + rumoca_core::BuiltinFunction::Ceil => { + let value = compile_time_unary_numeric_builtin_arg(args, const_scope, span)?; + Ok(value.ceil()) + } + rumoca_core::BuiltinFunction::Atan2 + | rumoca_core::BuiltinFunction::Min + | rumoca_core::BuiltinFunction::Max + | rumoca_core::BuiltinFunction::Div + | rumoca_core::BuiltinFunction::Mod + | rumoca_core::BuiltinFunction::Rem => { + let (lhs, rhs) = compile_time_binary_numeric_builtin_args(args, const_scope, span)?; + rumoca_core::apply_scalar_binary_math(function, lhs, rhs).ok_or_else(|| { + unsupported_at( + format!( + "builtin `{}` is undefined for compile-time subscript arguments", + function.name() + ), + span, + ) + }) + } + rumoca_core::BuiltinFunction::Size => compile_time_size_builtin(args, const_scope, span), _ => Err(unsupported_at( format!( "builtin `{}` is unsupported in compile-time subscript expression", @@ -1103,6 +1167,110 @@ fn compile_time_index_builtin( } } +fn compile_time_unary_numeric_builtin_arg( + args: &[rumoca_core::Expression], + const_scope: &IndexMap, + span: rumoca_core::Span, +) -> Result { + let Some(arg) = args.first() else { + return Err(unsupported_at( + "compile-time subscript builtin requires an argument", + span, + )); + }; + compile_time_index_raw_with_owner(arg, const_scope, span) +} + +fn compile_time_binary_numeric_builtin_args( + args: &[rumoca_core::Expression], + const_scope: &IndexMap, + span: rumoca_core::Span, +) -> Result<(f64, f64), LowerError> { + let Some(lhs) = args.first() else { + return Err(unsupported_at( + "compile-time subscript builtin requires two arguments", + span, + )); + }; + let Some(rhs) = args.get(1) else { + return Err(unsupported_at( + "compile-time subscript builtin requires two arguments", + span, + )); + }; + Ok(( + compile_time_index_raw_with_owner(lhs, const_scope, span)?, + compile_time_index_raw_with_owner(rhs, const_scope, span)?, + )) +} + +fn compile_time_size_builtin( + args: &[rumoca_core::Expression], + const_scope: &IndexMap, + span: rumoca_core::Span, +) -> Result { + let [array, dim] = args else { + return Err(unsupported_at( + "size() in compile-time subscript expression requires array and dimension arguments", + span, + )); + }; + let dim = compile_time_index_expr_with_owner(dim, const_scope, span)?; + let Some(zero_based_dim) = dim.checked_sub(1) else { + return Err(unsupported_at("size dimension must be positive", span)); + }; + if let Some(shape) = literal_array_shape(array) { + return shape + .get(zero_based_dim) + .copied() + .map(|value| value as f64) + .ok_or_else(|| { + unsupported_at( + format!("size dimension {dim} exceeds array rank {}", shape.len()), + array.span().unwrap_or(span), + ) + }); + } + if let rumoca_core::Expression::VarRef { + name, subscripts, .. + } = array + && subscripts.is_empty() + && let Some(value) = const_scope.get(&super::size_binding_key(name.as_str(), dim)) + { + return Ok(*value); + } + Err(unsupported_at( + "size() requires a compile-time array shape", + array.span().unwrap_or(span), + )) +} + +fn literal_array_shape(expr: &rumoca_core::Expression) -> Option> { + match expr { + rumoca_core::Expression::Array { + elements, + is_matrix: false, + .. + } => Some(vec![elements.len()]), + rumoca_core::Expression::Array { + elements, + is_matrix: true, + .. + } => { + let rows = elements.len(); + let cols = elements + .first() + .and_then(|row| match row { + rumoca_core::Expression::Array { elements, .. } => Some(elements.len()), + _ => None, + }) + .unwrap_or(0); + Some(vec![rows, cols]) + } + _ => None, + } +} + pub(super) struct AssignmentTarget { pub base: String, pub indices: Option>, @@ -1611,6 +1779,34 @@ mod tests { ); } + #[test] + fn compile_time_index_expr_evaluates_modelica_mod_builtin() { + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("compile_time_mod.mo"), + 2, + 11, + ); + let expr = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Mod, + args: vec![ + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(5), + span, + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(2), + span, + }, + ], + span, + }; + + let index = compile_time_index_expr_with_owner(&expr, &IndexMap::new(), span) + .expect("mod(5, 2) should evaluate to a positive compile-time index"); + + assert_eq!(index, 1); + } + #[test] fn positive_f64_index_rejects_non_positive_value_with_span() { let span = rumoca_core::Span::from_offsets( diff --git a/crates/rumoca-phase-solve/src/lower/initial_residual.rs b/crates/rumoca-phase-solve/src/lower/initial_residual.rs index fc87149ca..88cf27252 100644 --- a/crates/rumoca-phase-solve/src/lower/initial_residual.rs +++ b/crates/rumoca-phase-solve/src/lower/initial_residual.rs @@ -19,6 +19,42 @@ pub fn lower_initial_residual( ) } +/// Lower exactly one already-materialized initial residual cell. +pub(crate) fn lower_initial_residual_cell( + dae_model: &dae::Dae, + layout: &VarLayout, + equation_index: usize, + equation: &dae::Equation, +) -> Result, super::LowerError> { + let mut rows = expression_rows::lower_residual_rows_from_equations_with_mode( + dae_model, + layout, + [(equation_index, equation)], + 0, + true, + )?; + if rows.len() != 1 { + return Err(super::LowerError::contract_violation( + "structured initial base cell must lower to exactly one residual row", + equation.span, + )); + } + Ok(rows.remove(0)) +} + +/// Lower a complete materialized initial family in one shared lowering context. +/// GPU affine proof uses the batch form so full-domain verification remains +/// linear instead of rebuilding expression-lowering metadata for every cell. +pub(crate) fn lower_initial_residual_cells<'a>( + dae_model: &dae::Dae, + layout: &VarLayout, + equations: impl IntoIterator, +) -> Result>, super::LowerError> { + expression_rows::lower_residual_rows_from_equations_with_mode( + dae_model, layout, equations, 0, true, + ) +} + pub fn initial_residual_equations<'a>( dae_model: &'a dae::Dae, layout: &VarLayout, diff --git a/crates/rumoca-phase-solve/src/lower/misc_helpers.rs b/crates/rumoca-phase-solve/src/lower/misc_helpers.rs index 6141cb4a7..439f00004 100644 --- a/crates/rumoca-phase-solve/src/lower/misc_helpers.rs +++ b/crates/rumoca-phase-solve/src/lower/misc_helpers.rs @@ -34,7 +34,7 @@ pub(super) fn scalar_literal_projection( reason: "literal scalar projection requires source span".to_string(), })?; let Some(indices) = static_subscript_indices_with_owner(subscripts, span)? else { - return Ok(false); + return Ok(true); }; Ok(indices.first().is_some_and(|index| *index > 0)) } diff --git a/crates/rumoca-phase-solve/src/lower/root_conditions.rs b/crates/rumoca-phase-solve/src/lower/root_conditions.rs index c34940155..ce2ffaab0 100644 --- a/crates/rumoca-phase-solve/src/lower/root_conditions.rs +++ b/crates/rumoca-phase-solve/src/lower/root_conditions.rs @@ -9,6 +9,7 @@ struct RootRuntime<'a> { clock_timings: &'a IndexMap, triggered_clock_conditions: &'a [rumoca_core::Expression], variable_starts: &'a IndexMap, + dae_variables: &'a dae::DaeVariables, } pub(super) fn lower_root_conditions( @@ -21,6 +22,7 @@ pub(super) fn lower_root_conditions( clock_timings: &dae_model.clocks.timings, triggered_clock_conditions: &dae_model.clocks.triggered_conditions, variable_starts: &dae_model.metadata.variable_starts, + dae_variables: &dae_model.variables, }; let span = root_condition_context_span(dae_model); let row_count = root_condition_count(dae_model, span)?; @@ -201,11 +203,14 @@ fn lower_root_condition_row( let mut builder = LowerBuilder::new_with_runtime_metadata( layout, runtime.functions, - runtime.clock_intervals, - runtime.clock_timings, - runtime.triggered_clock_conditions, - runtime.variable_starts, - false, + super::RuntimeLowerBuilderMetadata { + clock_intervals: runtime.clock_intervals, + clock_timings: runtime.clock_timings, + triggered_clock_conditions: runtime.triggered_clock_conditions, + variable_starts: runtime.variable_starts, + dae_variables: runtime.dae_variables, + is_initial_mode: false, + }, ); let scope = Scope::new(); let span = root_condition_span(condition)?; @@ -448,11 +453,14 @@ fn lower_synthetic_root_condition_row( let mut builder = LowerBuilder::new_with_runtime_metadata( layout, runtime.functions, - runtime.clock_intervals, - runtime.clock_timings, - runtime.triggered_clock_conditions, - runtime.variable_starts, - false, + super::RuntimeLowerBuilderMetadata { + clock_intervals: runtime.clock_intervals, + clock_timings: runtime.clock_timings, + triggered_clock_conditions: runtime.triggered_clock_conditions, + variable_starts: runtime.variable_starts, + dae_variables: runtime.dae_variables, + is_initial_mode: false, + }, ); let root_value = builder.lower_expr(condition, &Scope::new(), 0)?; builder.ops.push(LinearOp::StoreOutput { src: root_value }); @@ -467,11 +475,14 @@ fn lower_triggered_clock_condition_row( let mut builder = LowerBuilder::new_with_runtime_metadata( layout, runtime.functions, - runtime.clock_intervals, - runtime.clock_timings, - runtime.triggered_clock_conditions, - runtime.variable_starts, - false, + super::RuntimeLowerBuilderMetadata { + clock_intervals: runtime.clock_intervals, + clock_timings: runtime.clock_timings, + triggered_clock_conditions: runtime.triggered_clock_conditions, + variable_starts: runtime.variable_starts, + dae_variables: runtime.dae_variables, + is_initial_mode: false, + }, ); let root_value = lower_bool_condition_as_root(condition, &mut builder, &Scope::new())?; builder.ops.push(LinearOp::StoreOutput { src: root_value }); diff --git a/crates/rumoca-phase-solve/src/lower/source_refs.rs b/crates/rumoca-phase-solve/src/lower/source_refs.rs index 5ef9269c2..c794b0617 100644 --- a/crates/rumoca-phase-solve/src/lower/source_refs.rs +++ b/crates/rumoca-phase-solve/src/lower/source_refs.rs @@ -182,7 +182,7 @@ pub(super) fn component_reference_key_for_expr( let Some(component_ref) = component_reference_for_expr(expr)? else { return Ok(None); }; - component_reference_key(component_ref).map(Some) + optional_component_reference_key(component_ref) } pub(super) fn component_reference_key_for_field_base( @@ -192,7 +192,7 @@ pub(super) fn component_reference_key_for_field_base( let Some(component_ref) = component_reference_for_field_base(base, field)? else { return Ok(None); }; - component_reference_key(component_ref).map(Some) + optional_component_reference_key(component_ref) } pub(super) fn component_reference_for_field_base( @@ -270,6 +270,12 @@ pub(super) fn component_reference_key( return Ok(key); } } + component_reference_key_strict(component_ref) +} + +pub(super) fn component_reference_key_strict( + component_ref: rumoca_core::ComponentReference, +) -> Result { ComponentReferenceKey::from_component_reference(&component_ref).map_err(|err| { LowerError::contract_violation( format!("indexed solve-layout lookup has non-static component reference: {err}"), @@ -277,3 +283,18 @@ pub(super) fn component_reference_key( ) }) } + +pub(super) fn optional_component_reference_key( + component_ref: rumoca_core::ComponentReference, +) -> Result, LowerError> { + match ComponentReferenceKey::from_component_reference(&component_ref) { + Ok(key) => Ok(Some(key)), + Err(err) if err.kind == rumoca_ir_solve::ComponentReferenceKeyErrorKind::MissingDefId => { + Ok(None) + } + Err(err) => Err(LowerError::contract_violation( + format!("indexed solve-layout lookup has non-static component reference: {err}"), + err.span, + )), + } +} diff --git a/crates/rumoca-phase-solve/src/lower/statements.rs b/crates/rumoca-phase-solve/src/lower/statements.rs index 7760b5ef9..70f9cf793 100644 --- a/crates/rumoca-phase-solve/src/lower/statements.rs +++ b/crates/rumoca-phase-solve/src/lower/statements.rs @@ -1,3 +1,6 @@ +// SPEC_0021 file-size exception: statement lowering still owns function body, +// loop, branch, and assignment lowering together. split plan: split loop/body +// lowering and assignment projection into dedicated modules. use indexmap::{IndexMap, IndexSet}; use rumoca_ir_solve::{BinaryOp, Reg, ScalarSlot}; @@ -35,6 +38,33 @@ fn copy_statement_call_args( Ok(copied) } +fn bind_merged_indexed_scope_value( + scope: &mut Scope, + base: &str, + indices: &[usize], + reg: Reg, + span: rumoca_core::Span, +) -> Result<(), LowerError> { + let key = format_subscript_binding_key(base, indices); + scope.insert(generated_scope_key(&key), reg); + let target_path = generated_scope_key(base); + scope.insert_indexed(&target_path, indices, reg, span)?; + if indices.iter().all(|index| *index == 1) { + scope.insert(target_path, reg); + } + Ok(()) +} + +struct BranchIndexedMerge<'a> { + scope: &'a mut Scope, + entry_indexed: &'a IndexMap>, + cond_regs: &'a [Reg], + cond_spans: &'a [rumoca_core::Span], + branch_indexed: &'a [IndexMap>], + else_indexed: &'a IndexMap>, + span: rumoca_core::Span, +} + impl<'a> LowerBuilder<'a> { /// Returns `true` when lowering should stop due to `return`. pub(super) fn lower_statements( @@ -71,7 +101,6 @@ impl<'a> LowerBuilder<'a> { } return Ok(false); } - let entry_scope = scope.clone(); // Each branch lowers inside `with_local_lower_frame`, which rolls back // the builder-level `local_indexed_bindings` cache afterwards. Snapshot @@ -156,14 +185,15 @@ impl<'a> LowerBuilder<'a> { // branch frame. Reconstruct array elements assigned inside branches so // a partially-filled array (e.g. a small-angle guard that writes only // some Lie-algebra components per branch) survives intact. - self.merge_branch_indexed_bindings( - &entry_indexed, - &cond_regs, - &cond_spans, - &branch_indexed, - &else_indexed, - branch_span, - )?; + self.merge_branch_indexed_bindings(BranchIndexedMerge { + scope, + entry_indexed: &entry_indexed, + cond_regs: &cond_regs, + cond_spans: &cond_spans, + branch_indexed: &branch_indexed, + else_indexed: &else_indexed, + span: branch_span, + })?; Ok(false) } @@ -181,12 +211,7 @@ impl<'a> LowerBuilder<'a> { /// (never an unconditional leak). fn merge_branch_indexed_bindings( &mut self, - entry_indexed: &IndexMap>, - cond_regs: &[Reg], - cond_spans: &[rumoca_core::Span], - branch_indexed: &[IndexMap>], - else_indexed: &IndexMap>, - merge_span: rumoca_core::Span, + ctx: BranchIndexedMerge<'_>, ) -> Result<(), LowerError> { fn lookup( store: &IndexMap>, @@ -202,29 +227,31 @@ impl<'a> LowerBuilder<'a> { // Every (base, indices) a branch or the else-branch assigns differently // from the entry state needs a merged value. let mut candidates: IndexSet<(String, Vec)> = IndexSet::new(); - let assigned_entries = branch_indexed + let assigned_entries = ctx + .branch_indexed .iter() - .chain(std::iter::once(else_indexed)) + .chain(std::iter::once(ctx.else_indexed)) .flat_map(|store| store.iter()) .flat_map(|(base, entries)| entries.iter().map(move |entry| (base, entry))); for (base, entry) in assigned_entries { - if lookup(entry_indexed, base, &entry.indices) != Some(entry.reg) { + if lookup(ctx.entry_indexed, base, &entry.indices) != Some(entry.reg) { candidates.insert((base.clone(), entry.indices.clone())); } } for (base, indices) in candidates { - let mut merged = if let Some(merged) = lookup(else_indexed, &base, &indices) - .or_else(|| lookup(entry_indexed, &base, &indices)) + let mut merged = if let Some(merged) = lookup(ctx.else_indexed, &base, &indices) + .or_else(|| lookup(ctx.entry_indexed, &base, &indices)) { merged } else { - self.emit_const_at(0.0, merge_span)? + self.emit_const_at(0.0, ctx.span)? }; - let branch_values = cond_regs + let branch_values = ctx + .cond_regs .iter() - .zip(cond_spans.iter()) - .zip(branch_indexed.iter()) + .zip(ctx.cond_spans.iter()) + .zip(ctx.branch_indexed.iter()) .rev() .filter_map(|((cond, span), store)| { lookup(store, &base, &indices).map(|reg| (cond, span, reg)) @@ -233,11 +260,12 @@ impl<'a> LowerBuilder<'a> { merged = self.emit_select_at(*cond, reg, merged, *span)?; } upsert_local_indexed_binding( - self.local_indexed_bindings.entry(base).or_default(), + self.local_indexed_bindings.entry(base.clone()).or_default(), &indices, merged, - merge_span, + ctx.span, )?; + bind_merged_indexed_scope_value(ctx.scope, &base, &indices, merged, ctx.span)?; } Ok(()) } @@ -489,6 +517,84 @@ impl<'a> LowerBuilder<'a> { } } + pub(super) fn eval_compile_time_string( + &self, + expr: &rumoca_core::Expression, + const_scope: &IndexMap, + ) -> Result { + match expr { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String(value), + .. + } => Ok(value.clone()), + rumoca_core::Expression::VarRef { + name, + subscripts, + span, + } => { + let key = compile_time_var_key(name, subscripts, const_scope, *span)?; + let Some(start) = self + .variable_starts + .and_then(|starts| starts.get(key.as_str())) + else { + return Err(unsupported_at( + format!("compile-time string expression requires constant `{key}`"), + *span, + )); + }; + self.eval_compile_time_string(start, const_scope) + } + rumoca_core::Expression::Index { + base, + subscripts, + span, + } => { + let [subscript] = subscripts.as_slice() else { + return Err(unsupported_at( + "compile-time string array index requires one subscript", + *span, + )); + }; + let index = match subscript { + rumoca_core::Subscript::Index { value, span } if *value > 0 => { + positive_i64_index(*value, *span)? + } + rumoca_core::Subscript::Expr { expr, .. } => { + let value = self.eval_compile_time_int( + expr, + const_scope, + "compile-time string array index", + )?; + positive_compile_time_string_index(value, expr.span().unwrap_or(*span))? + } + _ => { + return Err(unsupported_at( + "compile-time string array index requires a positive scalar index", + subscript.span(), + )); + } + }; + let rumoca_core::Expression::Array { elements, .. } = base.as_ref() else { + return Err(unsupported_at( + "compile-time string index requires a literal string array", + *span, + )); + }; + let Some(element) = elements.get(index - 1) else { + return Err(unsupported_at( + "compile-time string array index is out of bounds", + *span, + )); + }; + self.eval_compile_time_string(element, const_scope) + } + _ => Err(unsupported_at( + "unsupported compile-time string expression", + self.statement_expr_span(expr)?, + )), + } + } + pub(super) fn eval_compile_time_var_ref( &self, name: &rumoca_core::Reference, @@ -513,6 +619,15 @@ impl<'a> LowerBuilder<'a> { { return Ok(value); } + if let Some(value) = self.eval_compile_time_var_ref_by_def_id(name, const_scope) { + return Ok(value); + } + if !name.as_str().contains('.') + && subscripts.is_empty() + && let Some(value) = self.eval_unique_compile_time_suffix(name.as_str(), const_scope) + { + return Ok(value); + } match self.layout.binding(key.as_str()) { Some(ScalarSlot::Constant(value)) => Ok(value), Some(_) | None => Err(unsupported_at( @@ -522,6 +637,77 @@ impl<'a> LowerBuilder<'a> { } } + fn eval_compile_time_var_ref_by_def_id( + &self, + name: &rumoca_core::Reference, + const_scope: &IndexMap, + ) -> Option { + let def_id = name.target_def_id()?; + let variables = self.dae_variables?; + let variable = variables + .parameters + .values() + .chain(variables.constants.values()) + .chain(variables.discrete_valued.values()) + .chain(variables.discrete_reals.values()) + .find(|variable| { + variable + .component_ref + .as_ref() + .and_then(|component_ref| component_ref.def_id) + == Some(def_id) + })?; + let start = variable.start.as_ref()?; + if start_metadata_refers_to_key(start, variable.name.as_str()) { + return None; + } + self.eval_compile_time_expr(start, const_scope).ok() + } + + #[allow(clippy::excessive_nesting)] + fn eval_unique_compile_time_suffix( + &self, + name: &str, + const_scope: &IndexMap, + ) -> Option { + let suffix = format!(".{name}"); + let mut matched = None; + if let Some(value) = self.structural_bindings.get(name).copied() { + matched = Some(value); + } + for (key, value) in self + .structural_bindings + .iter() + .filter(|(key, _)| key.ends_with(&suffix)) + { + let _ = key; + let value = *value; + if let Some(previous) = matched + && (previous - value).abs() > 1.0e-9 + { + return None; + } + matched = Some(value); + } + if let Some(starts) = self.variable_starts { + for (key, start) in starts.iter().filter(|(key, _)| key.ends_with(&suffix)) { + if start_metadata_refers_to_key(start, key.as_str()) { + continue; + } + let Ok(value) = self.eval_compile_time_expr(start, const_scope) else { + continue; + }; + if let Some(previous) = matched + && (previous - value).abs() > 1.0e-9 + { + return None; + } + matched = Some(value); + } + } + matched + } + pub(super) fn eval_compile_time_unary( &self, op: &rumoca_core::OpUnary, @@ -594,6 +780,25 @@ impl<'a> LowerBuilder<'a> { if matches!(function, rumoca_core::BuiltinFunction::Size) { return self.eval_compile_time_size(args, span, const_scope); } + match function { + rumoca_core::BuiltinFunction::Min if args.len() == 1 => { + return self.eval_compile_time_array_reduction( + &args[0], + const_scope, + f64::min, + "min", + ); + } + rumoca_core::BuiltinFunction::Max if args.len() == 1 => { + return self.eval_compile_time_array_reduction( + &args[0], + const_scope, + f64::max, + "max", + ); + } + _ => {} + } let arg0 = eval_builtin_arg(self, args, 0, const_scope)?; match function { rumoca_core::BuiltinFunction::Abs => Ok(arg0.abs()), @@ -604,10 +809,26 @@ impl<'a> LowerBuilder<'a> { } rumoca_core::BuiltinFunction::Ceil => Ok(arg0.ceil()), rumoca_core::BuiltinFunction::Min => { + if args.len() == 1 { + return self.eval_compile_time_array_reduction( + &args[0], + const_scope, + f64::min, + "min", + ); + } let arg1 = eval_builtin_arg(self, args, 1, const_scope)?; Ok(arg0.min(arg1)) } rumoca_core::BuiltinFunction::Max => { + if args.len() == 1 { + return self.eval_compile_time_array_reduction( + &args[0], + const_scope, + f64::max, + "max", + ); + } let arg1 = eval_builtin_arg(self, args, 1, const_scope)?; Ok(arg0.max(arg1)) } @@ -657,6 +878,46 @@ impl<'a> LowerBuilder<'a> { } } + fn eval_compile_time_array_reduction( + &self, + expr: &rumoca_core::Expression, + const_scope: &IndexMap, + op: fn(f64, f64) -> f64, + name: &'static str, + ) -> Result { + let mut values = Vec::new(); + self.collect_compile_time_array_scalars(expr, const_scope, &mut values)?; + let mut iter = values.into_iter(); + let Some(first) = iter.next() else { + return Err(unsupported_at( + format!("{name}() requires a non-empty array argument"), + self.statement_expr_span(expr)?, + )); + }; + Ok(iter.fold(first, op)) + } + + fn collect_compile_time_array_scalars( + &self, + expr: &rumoca_core::Expression, + const_scope: &IndexMap, + values: &mut Vec, + ) -> Result<(), LowerError> { + match expr { + rumoca_core::Expression::Array { elements, .. } + | rumoca_core::Expression::Tuple { elements, .. } => { + for element in elements { + self.collect_compile_time_array_scalars(element, const_scope, values)?; + } + Ok(()) + } + _ => { + values.push(self.eval_compile_time_expr(expr, const_scope)?); + Ok(()) + } + } + } + pub(super) fn eval_compile_time_size( &self, args: &[rumoca_core::Expression], @@ -674,6 +935,24 @@ impl<'a> LowerBuilder<'a> { let Some(expr) = args.first() else { return Err(unsupported_at("size() requires an array expression", span)); }; + if let rumoca_core::Expression::Range { + start, + step, + end, + span: range_span, + } = expr + && dim == 1 + { + let values = self.eval_compile_time_range_values( + start, + step.as_deref(), + end, + *range_span, + const_scope, + "size range dimension", + )?; + return Ok(values.len() as f64); + } let rumoca_core::Expression::VarRef { name, subscripts, .. } = expr @@ -695,6 +974,12 @@ impl<'a> LowerBuilder<'a> { if !subscripts.is_empty() { return Ok(1.0); } + if let Some(dims) = self.local_binding_dims.get(name.as_str()) + && let Some(value) = dims.get(dim - 1) + && *value > 0 + { + return Ok(*value as f64); + } let key = size_binding_key(name.as_str(), dim); self.structural_bindings .get(key.as_str()) @@ -731,7 +1016,7 @@ impl<'a> LowerBuilder<'a> { { return Ok(false); } - let values = self.lower_array_like_values(value, scope, call_depth)?; + let values = self.lower_assignment_values(&target, value, scope, call_depth)?; if let Some(indices) = target .indices .as_deref() @@ -798,6 +1083,33 @@ impl<'a> LowerBuilder<'a> { } } + fn lower_assignment_values( + &mut self, + target: &AssignmentTarget, + value: &rumoca_core::Expression, + scope: &Scope, + call_depth: usize, + ) -> Result, LowerError> { + if target.indices.is_none() && !self.assignment_target_has_array_shape(&target.base) { + return Ok(vec![self.lower_expr(value, scope, call_depth)?]); + } + self.lower_array_like_values(value, scope, call_depth) + } + + fn assignment_target_has_array_shape(&self, target: &str) -> bool { + self.local_binding_dims + .get(target) + .is_some_and(|dims| !dims.is_empty()) + || self + .layout + .shape(target) + .is_some_and(|dims| !dims.is_empty()) + || self + .local_indexed_bindings + .get(target) + .is_some_and(|bindings| !bindings.is_empty()) + } + fn bind_record_component_assignment( &mut self, scope: &mut Scope, @@ -940,7 +1252,7 @@ impl<'a> LowerBuilder<'a> { assignment_span: rumoca_core::Span, call_depth: usize, ) -> Result { - let Some(fields) = record_if_assignment_fields(value) else { + let Some(fields) = self.record_if_assignment_fields(value) else { return Ok(false); }; let span = self.statement_expr_or_context_span(value, assignment_span)?; @@ -959,6 +1271,58 @@ impl<'a> LowerBuilder<'a> { Ok(true) } + fn record_if_assignment_fields(&self, value: &rumoca_core::Expression) -> Option> { + let rumoca_core::Expression::If { + branches, + else_branch, + .. + } = value + else { + return None; + }; + + let mut fields = IndexSet::new(); + for (_, branch_expr) in branches { + self.collect_record_constructor_fields(branch_expr, &mut fields); + } + self.collect_record_constructor_fields(else_branch, &mut fields); + (!fields.is_empty()).then(|| fields.into_iter().collect()) + } + + fn collect_record_constructor_fields( + &self, + expr: &rumoca_core::Expression, + fields: &mut IndexSet, + ) { + let rumoca_core::Expression::FunctionCall { + name, + args, + is_constructor, + .. + } = expr + else { + return; + }; + if !self.is_record_constructor_call(name, *is_constructor) { + return; + } + if let Some(constructor) = self.lookup_function(name) { + for input in &constructor.inputs { + fields.insert(input.name.clone()); + } + return; + } + for arg in args { + if let Some((field, _)) = super::function_calls::decode_named_function_arg(arg) { + fields.insert(field.to_string()); + } + } + if fields.is_empty() && name.last_segment() == "Complex" { + fields.insert("re".to_string()); + fields.insert("im".to_string()); + } + } + fn bind_record_function_call_assignment( &mut self, caller_scope: &mut Scope, @@ -1554,14 +1918,29 @@ impl<'a> LowerBuilder<'a> { if !self.is_record_constructor_call(name, *is_constructor) { return Ok(()); } - let Some(constructor) = self.lookup_function(name).cloned() else { + let constructor_inputs = if let Some(constructor) = self.lookup_function(name).cloned() { + constructor.inputs + } else if name.last_segment() == "Complex" { + let span = value + .span() + .ok_or_else(|| LowerError::UnspannedContractViolation { + reason: format!( + "built-in Complex constructor `{}` requires source provenance for field binding", + name.as_str() + ), + })?; + vec![ + rumoca_core::FunctionParam::new("re", "Real", span), + rumoca_core::FunctionParam::new("im", "Real", span), + ] + } else { return Ok(()); }; let (named_args, positional_args) = super::function_calls::split_named_and_positional_call_args(name.as_str(), args)?; let mut positional_idx = 0usize; - for input in &constructor.inputs { + for input in &constructor_inputs { let arg_expr = named_args.get(input.name.as_str()).copied().or_else(|| { let positional = positional_args.get(positional_idx).copied(); positional_idx += usize::from(positional.is_some()); @@ -1598,6 +1977,24 @@ impl<'a> LowerBuilder<'a> { } } +fn positive_compile_time_string_index( + value: i64, + span: rumoca_core::Span, +) -> Result { + if value <= 0 { + return Err(unsupported_at( + "compile-time string array index requires a positive scalar index", + span, + )); + } + usize::try_from(value).map_err(|_| { + unsupported_at( + "compile-time string array index is outside supported range", + span, + ) + }) +} + fn start_metadata_refers_to_key(expr: &rumoca_core::Expression, key: &str) -> bool { binding_base_key(expr).is_ok_and(|start_key| start_key == key) } @@ -1611,43 +2008,6 @@ fn is_control_flag(name: &rumoca_ir_solve::ComponentReferenceKey) -> bool { } } -fn record_if_assignment_fields(value: &rumoca_core::Expression) -> Option> { - let rumoca_core::Expression::If { - branches, - else_branch, - .. - } = value - else { - return None; - }; - - let mut fields = IndexSet::new(); - for (_, branch_expr) in branches { - collect_record_constructor_fields(branch_expr, &mut fields); - } - collect_record_constructor_fields(else_branch, &mut fields); - (!fields.is_empty()).then(|| fields.into_iter().collect()) -} - -fn collect_record_constructor_fields( - expr: &rumoca_core::Expression, - fields: &mut IndexSet, -) { - let rumoca_core::Expression::FunctionCall { - args, - is_constructor: true, - .. - } = expr - else { - return; - }; - for arg in args { - if let Some((field, _)) = super::function_calls::decode_named_function_arg(arg) { - fields.insert(field.to_string()); - } - } -} - fn record_if_field_expression( value: &rumoca_core::Expression, field: &str, diff --git a/crates/rumoca-phase-solve/src/lower/statements/tests.rs b/crates/rumoca-phase-solve/src/lower/statements/tests.rs index 2227dffac..02ca20906 100644 --- a/crates/rumoca-phase-solve/src/lower/statements/tests.rs +++ b/crates/rumoca-phase-solve/src/lower/statements/tests.rs @@ -11,6 +11,13 @@ fn literal_i64(value: i64, span: rumoca_core::Span) -> rumoca_core::Expression { } } +fn literal_string(value: &str, span: rumoca_core::Span) -> rumoca_core::Expression { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String(value.to_string()), + span, + } +} + fn lower_builder<'a>( layout: &'a rumoca_ir_solve::VarLayout, functions: &'a IndexMap, @@ -72,6 +79,44 @@ fn eval_compile_time_expr_rejects_unsupported_form_with_span() { assert_eq!(err.reason(), "unsupported expression in for-loop range"); } +#[test] +fn eval_compile_time_expr_evaluates_modelica_strings_is_equal() { + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("compile_time_string_equal.mo"), + 3, + 32, + ); + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Modelica.Utilities.Strings.isEqual"), + args: vec![ + rumoca_core::Expression::Index { + base: Box::new(rumoca_core::Expression::Array { + elements: vec![literal_string("CO2", span)], + is_matrix: false, + span, + }), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(literal_i64(1, span)), + span, + }], + span, + }, + literal_string("CO2", span), + ], + is_constructor: false, + span, + }; + let layout = rumoca_ir_solve::VarLayout::default(); + let functions = IndexMap::new(); + let builder = lower_builder(&layout, &functions); + + let value = builder + .eval_compile_time_expr(&expr, &IndexMap::new()) + .expect("Strings.isEqual over literal strings should fold"); + + assert_eq!(value, 1.0); +} + #[test] fn eval_compile_time_builtin_reports_zero_denominator_span() { let call_span = rumoca_core::Span::from_offsets( diff --git a/crates/rumoca-phase-solve/src/lower/tests.rs b/crates/rumoca-phase-solve/src/lower/tests.rs index d1abce8ee..82926a84c 100644 --- a/crates/rumoca-phase-solve/src/lower/tests.rs +++ b/crates/rumoca-phase-solve/src/lower/tests.rs @@ -3,7 +3,7 @@ // behavior families into focused test modules with shared fixtures. use super::{ - Scope, expression_rows::lower_expression_rows_with_mode, lower_derivative_rhs, + LowerBuilder, Scope, expression_rows::lower_expression_rows_with_mode, lower_derivative_rhs, lower_derivative_rhs_scalar_programs, lower_discrete_rhs, lower_expression, lower_expression_rows_from_expressions_with_runtime_metadata, lower_initial_expression_rows_from_expressions, lower_initial_residual, lower_residual, @@ -14,8 +14,10 @@ use crate::lower_solve_problem; use indexmap::IndexMap; use rumoca_ir_dae as dae; use rumoca_ir_solve::{ - BinaryOp, CompareOp, ComputeNode, LinearOp, Reg, ScalarSlot, UnaryOp, VarLayout, + BinaryOp, CompareOp, ComputeNode, ExternalFunctionKind, LinearOp, Reg, ScalarSlot, UnaryOp, + VarLayout, }; +use std::sync::Arc; mod array_operator_tests; mod discrete_expression_tests; @@ -66,6 +68,61 @@ fn lower_builder_try_pack_registers_rejects_overflow() { assert!(err.reason().contains("Solve register allocation overflow")); } +#[test] +fn lower_expression_emits_energyplus_initialize_external_call() { + let span = lower_test_span(); + let mut functions = IndexMap::new(); + let mut initialize = rumoca_core::Function::new( + "Buildings.ThermalZones.EnergyPlus_9_6_0.BaseClasses.initialize", + span, + ); + initialize.external = Some(rumoca_core::ExternalFunction::default()); + initialize.add_input(rumoca_core::FunctionParam::new( + "isSynchronized", + "Real", + span, + )); + initialize + .outputs + .push(rumoca_core::FunctionParam::new("nObj", "Integer", span)); + functions.insert( + rumoca_core::VarName::new("Buildings.ThermalZones.EnergyPlus_9_6_0.BaseClasses.initialize"), + initialize, + ); + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new( + "Buildings.ThermalZones.EnergyPlus_9_6_0.BaseClasses.initialize", + ) + .into(), + args: vec![rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("__rumoca_named_arg__.isSynchronized").into(), + args: vec![rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(1.0), + span, + }], + is_constructor: true, + span, + }], + is_constructor: false, + span, + }; + + let lowered = lower_expression(&expr, &VarLayout::default(), &functions) + .expect("supported EnergyPlus external initialize should lower"); + assert!(matches!( + lowered.ops.as_slice(), + [ + LinearOp::Const { .. }, + LinearOp::ExternalCall { + function: ExternalFunctionKind::BuildingsEnergyPlusInitialize, + arg_count: 1, + output_index: 0, + .. + } + ] + )); +} + #[test] fn dynamic_layout_entries_rejects_missing_source_binding_group_with_span() { let layout = VarLayout::default(); @@ -443,6 +500,92 @@ fn size_builtin_rejects_unspanned_base_without_fabricating_span() { assert!(builder.ops.is_empty()); } +#[test] +fn size_builtin_lowers_compile_time_builtin_array_dimension() { + let layout = VarLayout::default(); + let functions = IndexMap::new(); + let mut builder = super::LowerBuilder::new(&layout, &functions); + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("phase_solve_lower_tests_source_32.mo"), + 4, + 24, + ); + let fill = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Fill, + args: vec![ + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(0.0), + span, + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(0), + span, + }, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(2), + span, + }, + ], + span, + }; + + let reg = builder + .lower_size_builtin(&[fill, int_lit(2)], span, &Scope::new(), 0) + .expect("size(fill(...), 2) should lower from compile-time shape"); + + assert_eq!(reg, 0); + assert!(matches!( + builder.ops.as_slice(), + [LinearOp::Const { dst: 0, value }] if (*value - 2.0).abs() < f64::EPSILON + )); +} + +#[test] +fn scalar_type_constructor_lowers_start_attribute_value() { + let span = lower_test_span(); + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("SI.Time").into(), + args: vec![ + named_arg("min", real_lit(f64::MIN_POSITIVE)), + named_arg("start", real_lit(1.0)), + ], + is_constructor: true, + span, + }; + + let lowered = lower_expression(&expr, &VarLayout::default(), &IndexMap::new()) + .expect("scalar type constructor with start attribute should lower"); + let (regs, _) = eval_linear_ops(&lowered.ops, &[], &[], 0.0); + + assert!((read_reg(®s, lowered.result) - 1.0).abs() < 1e-12); +} + +#[test] +fn scalar_type_constructor_in_multiplication_infers_scalar_shape() { + let span = lower_test_span(); + let constructor = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("SI.Time").into(), + args: vec![ + named_arg("min", real_lit(f64::MIN_POSITIVE)), + named_arg("start", real_lit(1.0)), + ], + is_constructor: true, + span, + }; + let expr = rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs: Box::new(constructor), + rhs: Box::new(real_lit(2.0)), + span, + }; + + let lowered = lower_expression(&expr, &VarLayout::default(), &IndexMap::new()) + .expect("scalar type constructor should infer scalar dimensions before multiplication"); + let (regs, _) = eval_linear_ops(&lowered.ops, &[], &[], 0.0); + + assert!((read_reg(®s, lowered.result) - 2.0).abs() < 1e-12); +} + #[test] fn size_from_dims_rejects_register_allocation_overflow_with_span() { let layout = VarLayout::default(); @@ -764,6 +907,547 @@ fn source_ref(name: &str) -> rumoca_core::Reference { rumoca_core::Reference::from_component_reference(source_component_ref_from_name(name)) } +#[test] +fn compile_time_subscript_indices_use_structural_bindings() { + let span = lower_test_span(); + let mut structural_bindings = IndexMap::new(); + structural_bindings.insert("nEle".to_string(), 4.0); + let layout = VarLayout::default(); + let functions = IndexMap::new(); + let builder = LowerBuilder::new(&layout, &functions) + .with_structural_bindings(Arc::new(structural_bindings)); + let subscripts = vec![rumoca_core::Subscript::generated_expr( + Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("nEle"), + subscripts: Vec::new(), + span, + }), + span, + )]; + + let indices = builder + .compile_time_subscript_indices(&subscripts, span) + .expect("structural subscript folding should not fail") + .expect("structural parameter subscript should fold"); + + assert_eq!(indices, vec![4]); +} + +#[test] +fn indexed_field_access_uses_structural_subscript_bindings() { + let span = lower_test_span(); + let key = "ele[4].T"; + let layout = VarLayout::from_parts_with_shapes_spans_and_indexed_bindings( + IndexMap::from([(key.to_string(), rumoca_ir_solve::scalar_slot_y(0))]), + IndexMap::new(), + IndexMap::new(), + IndexMap::new(), + 1, + 0, + ) + .expect("indexed field fixture layout should satisfy shape contract"); + let functions = IndexMap::new(); + let mut structural_bindings = IndexMap::new(); + structural_bindings.insert("nEle".to_string(), 4.0); + let mut builder = LowerBuilder::new(&layout, &functions) + .with_structural_bindings(Arc::new(structural_bindings)); + let expr = rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::Index { + base: Box::new(source_var("ele")), + subscripts: vec![rumoca_core::Subscript::generated_expr( + Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("nEle"), + subscripts: Vec::new(), + span, + }), + span, + )], + span, + }), + field: "T".to_string(), + span, + }; + + let reg = builder + .lower_expr(&expr, &Scope::new(), 0) + .expect("structural indexed field access should lower"); + let (regs, _) = eval_linear_ops(&builder.ops, &[289.0], &[], 0.0); + + assert!((read_reg(®s, reg) - 289.0).abs() < 1e-12); +} + +#[test] +fn nested_indexed_field_access_uses_structural_subscript_bindings() { + let span = lower_test_span(); + let key = "ele[4].vol2.T"; + let layout = VarLayout::from_parts_with_shapes_spans_and_indexed_bindings( + IndexMap::from([(key.to_string(), rumoca_ir_solve::scalar_slot_y(0))]), + IndexMap::new(), + IndexMap::new(), + IndexMap::new(), + 1, + 0, + ) + .expect("nested indexed field fixture layout should satisfy shape contract"); + let functions = IndexMap::new(); + let mut structural_bindings = IndexMap::new(); + structural_bindings.insert("nEle".to_string(), 4.0); + let mut builder = LowerBuilder::new(&layout, &functions) + .with_structural_bindings(Arc::new(structural_bindings)); + let indexed_ele = rumoca_core::Expression::Index { + base: Box::new(source_var("ele")), + subscripts: vec![rumoca_core::Subscript::generated_expr( + Box::new(rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("nEle"), + subscripts: Vec::new(), + span, + }), + span, + )], + span, + }; + let expr = rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FieldAccess { + base: Box::new(indexed_ele), + field: "vol2".to_string(), + span, + }), + field: "T".to_string(), + span, + }; + + let reg = builder + .lower_expr(&expr, &Scope::new(), 0) + .expect("nested structural indexed field access should lower"); + let (regs, _) = eval_linear_ops(&builder.ops, &[291.0], &[], 0.0); + + assert!((read_reg(®s, reg) - 291.0).abs() < 1e-12); +} + +#[test] +fn array_comprehension_index_lowers_as_compile_time_subscript() { + let span = lower_test_span(); + let layout = VarLayout::from_parts_with_shapes_spans_and_indexed_bindings( + IndexMap::from([ + ( + "ele[1].vol2.mWat_flow".to_string(), + rumoca_ir_solve::scalar_slot_y(0), + ), + ( + "ele[2].vol2.mWat_flow".to_string(), + rumoca_ir_solve::scalar_slot_y(1), + ), + ( + "ele[3].vol2.mWat_flow".to_string(), + rumoca_ir_solve::scalar_slot_y(2), + ), + ( + "ele[4].vol2.mWat_flow".to_string(), + rumoca_ir_solve::scalar_slot_y(3), + ), + ]), + IndexMap::new(), + IndexMap::new(), + IndexMap::new(), + 4, + 0, + ) + .expect("comprehension fixture layout should satisfy shape contract"); + let functions = IndexMap::new(); + let mut builder = LowerBuilder::new(&layout, &functions); + let indexed_ele = rumoca_core::Expression::Index { + base: Box::new(source_var("ele")), + subscripts: vec![rumoca_core::Subscript::generated_expr( + Box::new(var("i")), + span, + )], + span, + }; + let member_expr = rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FieldAccess { + base: Box::new(indexed_ele), + field: "vol2".to_string(), + span, + }), + field: "mWat_flow".to_string(), + span, + }; + let expr = rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Sum, + args: vec![rumoca_core::Expression::ArrayComprehension { + expr: Box::new(member_expr), + indices: vec![rumoca_core::ComprehensionIndex { + name: "i".to_string(), + range: rumoca_core::Expression::Range { + start: Box::new(int_lit(1)), + step: None, + end: Box::new(int_lit(4)), + span, + }, + }], + filter: None, + span, + }], + span, + }; + + let reg = builder + .lower_expr(&expr, &Scope::new(), 0) + .expect("array comprehension index should lower as compile-time subscript"); + let (regs, _) = eval_linear_ops(&builder.ops, &[1.0, 2.0, 3.0, 4.0], &[], 0.0); + + assert!((read_reg(®s, reg) - 10.0).abs() < 1e-12); +} + +#[test] +fn scalar_field_access_reads_first_record_array_member_field() { + let span = lower_test_span(); + let layout = VarLayout::from_parts_with_shapes_spans_and_indexed_bindings( + IndexMap::from([( + "zon[1].ports[1].p".to_string(), + rumoca_ir_solve::scalar_slot_y(0), + )]), + IndexMap::new(), + IndexMap::new(), + IndexMap::new(), + 1, + 0, + ) + .expect("record-array member fixture layout should satisfy shape contract"); + let functions = IndexMap::new(); + let mut builder = LowerBuilder::new(&layout, &functions); + let expr = rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FieldAccess { + base: Box::new(var("zon[1]")), + field: "ports".to_string(), + span, + }), + field: "p".to_string(), + span, + }; + + let reg = builder + .lower_expr(&expr, &Scope::new(), 0) + .expect("scalar field access should read first record-array member"); + let (regs, _) = eval_linear_ops(&builder.ops, &[101325.0], &[], 0.0); + + assert!((read_reg(®s, reg) - 101325.0).abs() < 1e-12); +} + +#[test] +fn scalar_field_access_reads_indexed_source_record_array_member_without_def_id() { + let span = lower_test_span(); + let layout = VarLayout::from_parts_with_shapes_spans_and_indexed_bindings( + IndexMap::from([( + "zon[1].ports[1].p".to_string(), + rumoca_ir_solve::scalar_slot_y(0), + )]), + IndexMap::new(), + IndexMap::new(), + IndexMap::new(), + 1, + 0, + ) + .expect("indexed source record-array fixture layout should satisfy shape contract"); + let functions = IndexMap::new(); + let mut builder = LowerBuilder::new(&layout, &functions); + let expr = rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FieldAccess { + base: Box::new(unresolved_source_var("zon[1]")), + field: "ports".to_string(), + span, + }), + field: "p".to_string(), + span, + }; + + let reg = builder + .lower_expr(&expr, &Scope::new(), 0) + .expect("scalar field access should use scalarized layout before requiring def identity"); + let (regs, _) = eval_linear_ops(&builder.ops, &[101325.0], &[], 0.0); + + assert!((read_reg(®s, reg) - 101325.0).abs() < 1e-12); +} + +#[test] +fn size_infers_indexed_source_field_shape_without_def_id() { + let span = lower_test_span(); + let layout = VarLayout::from_parts_with_shapes( + IndexMap::from([( + "zon[1].ports".to_string(), + rumoca_ir_solve::scalar_slot_y(0), + )]), + IndexMap::from([("zon[1].ports".to_string(), vec![5])]), + 5, + 0, + ) + .expect("indexed source field shape fixture layout should satisfy shape contract"); + let functions = IndexMap::new(); + let mut builder = LowerBuilder::new(&layout, &functions); + let expr = size_expr( + rumoca_core::Expression::FieldAccess { + base: Box::new(unresolved_source_var("zon[1]")), + field: "ports".to_string(), + span, + }, + 1, + ); + + let reg = builder + .lower_expr(&expr, &Scope::new(), 0) + .expect("size() should infer display-keyed field shape before requiring def identity"); + let (regs, _) = eval_linear_ops(&builder.ops, &[], &[], 0.0); + + assert!((read_reg(®s, reg) - 5.0).abs() < 1e-12); +} + +#[test] +fn size_of_unresolved_indexed_source_ref_does_not_require_def_id() { + let layout = VarLayout::default(); + let functions = IndexMap::new(); + let mut builder = LowerBuilder::new(&layout, &functions); + let expr = size_expr(unresolved_source_var("zon[1]"), 1); + + let reg = builder + .lower_expr(&expr, &Scope::new(), 0) + .expect("size() fallback should not require missing source def identity"); + let (regs, _) = eval_linear_ops(&builder.ops, &[], &[], 0.0); + + assert!((read_reg(®s, reg) - 1.0).abs() < 1e-12); +} + +#[test] +fn size_of_stream_passthrough_uses_argument_shape() { + let layout = VarLayout::from_parts_with_shapes( + IndexMap::from([("x".to_string(), rumoca_ir_solve::scalar_slot_y(0))]), + IndexMap::from([("x".to_string(), vec![3])]), + 3, + 0, + ) + .expect("stream passthrough size fixture layout should satisfy shape contract"); + let functions = IndexMap::new(); + let mut builder = LowerBuilder::new(&layout, &functions); + let expr = size_expr(stream_passthrough_expr("inStream", var("x")), 1); + + let reg = builder + .lower_expr(&expr, &Scope::new(), 0) + .expect("size() should use stream passthrough argument shape"); + let (regs, _) = eval_linear_ops(&builder.ops, &[], &[], 0.0); + + assert!((read_reg(®s, reg) - 3.0).abs() < 1e-12); +} + +#[test] +fn index_over_stream_passthrough_uses_argument_bindings() { + let span = lower_test_span(); + let layout = VarLayout::from_parts_with_shapes( + IndexMap::from([("x".to_string(), rumoca_ir_solve::scalar_slot_y(0))]), + IndexMap::from([("x".to_string(), vec![3])]), + 3, + 0, + ) + .expect("stream passthrough index fixture layout should satisfy shape contract"); + let functions = IndexMap::new(); + let mut builder = LowerBuilder::new(&layout, &functions); + let expr = rumoca_core::Expression::Index { + base: Box::new(stream_passthrough_expr("inStream", var("x"))), + subscripts: vec![rumoca_core::Subscript::generated_index(2, span)], + span, + }; + + let reg = builder + .lower_expr(&expr, &Scope::new(), 0) + .expect("index should use stream passthrough argument binding"); + let (regs, _) = eval_linear_ops(&builder.ops, &[10.0, 20.0, 30.0], &[], 0.0); + + assert!((read_reg(®s, reg) - 20.0).abs() < 1e-12); +} + +#[test] +fn scalar_var_ref_consumes_singleton_projection_after_array_element_selection() { + let span = lower_test_span(); + let mut dae_model = dae::Dae::new(); + dae_model + .variables + .outputs + .insert(rumoca_core::VarName::new("y"), { + let mut variable = dae::Variable { + name: rumoca_core::VarName::new("y"), + dims: vec![3], + ..dae::Variable::empty_with_span(span) + }; + variable.origin = dae::VariableOrigin::Generated; + variable + }); + let layout = build_var_layout(&dae_model).expect("vector layout should build"); + let functions = IndexMap::new(); + let mut builder = LowerBuilder::new(&layout, &functions); + let expr = rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("y"), + subscripts: vec![ + rumoca_core::Subscript::index(2, span), + rumoca_core::Subscript::index(1, span), + ], + span, + }; + + let reg = builder + .lower_expr(&expr, &Scope::new(), 0) + .expect("scalarized vector element should accept trailing singleton projection"); + let (regs, _) = eval_linear_ops(&builder.ops, &[10.0, 20.0, 30.0], &[], 0.0); + + assert!((read_reg(®s, reg) - 20.0).abs() < 1e-12); +} + +#[test] +fn zero_length_function_input_accepts_unknown_rank_empty_actual() { + let span = lower_test_span(); + let mut function = rumoca_core::Function::new("Pkg.zeroLength", span); + function.inputs.push(function_param_with_dims("X", &[0])); + function.outputs.push(function_param("y")); + function.body.push(rumoca_core::Statement::Assignment { + comp: component_ref("y"), + value: real_lit(1.0), + span, + }); + let mut functions = IndexMap::new(); + functions.insert(function.name.clone(), function); + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("Pkg.zeroLength").into(), + args: vec![stream_passthrough_expr("inStream", var("missingXi"))], + is_constructor: false, + span, + }; + + let lowered = lower_expression(&expr, &VarLayout::default(), &functions) + .expect("zero-length formal should bind unknown-rank empty actual as an empty array"); + let (regs, _) = eval_linear_ops(&lowered.ops, &[], &[], 0.0); + + assert!((read_reg(®s, lowered.result) - 1.0).abs() < 1e-12); +} + +#[test] +fn stream_passthrough_zero_shape_lowers_to_empty_array_values() { + let layout = VarLayout::from_parts_with_shapes( + IndexMap::new(), + IndexMap::from([("x".to_string(), vec![0])]), + 0, + 0, + ) + .expect("zero-size stream argument fixture layout should satisfy shape contract"); + let functions = IndexMap::new(); + let mut builder = LowerBuilder::new(&layout, &functions); + let values = builder + .lower_array_like_values( + &stream_passthrough_expr("inStream", var("x")), + &Scope::new(), + 0, + ) + .expect("zero-size stream passthrough should lower to empty values"); + + assert!(values.is_empty()); +} + +#[test] +fn multiply_zero_shape_stream_passthrough_lowers_to_empty_array_values() { + let span = lower_test_span(); + let layout = VarLayout::from_parts_with_shapes( + IndexMap::new(), + IndexMap::from([("x".to_string(), vec![0])]), + 0, + 0, + ) + .expect("zero-size stream multiplication fixture layout should satisfy shape contract"); + let functions = IndexMap::new(); + let mut builder = LowerBuilder::new(&layout, &functions); + let expr = rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs: Box::new(stream_passthrough_expr("inStream", var("x"))), + rhs: Box::new(real_lit(2.0)), + span, + }; + let values = builder + .lower_array_like_values(&expr, &Scope::new(), 0) + .expect("zero-size stream multiplication should lower to empty values"); + + assert!(values.is_empty()); +} + +#[test] +fn elementwise_zero_shape_stream_passthrough_lowers_to_empty_array_values() { + let span = lower_test_span(); + let layout = VarLayout::from_parts_with_shapes( + IndexMap::new(), + IndexMap::from([("x".to_string(), vec![0])]), + 0, + 0, + ) + .expect("zero-size stream elementwise fixture layout should satisfy shape contract"); + let functions = IndexMap::new(); + let mut builder = LowerBuilder::new(&layout, &functions); + let expr = rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: Box::new(stream_passthrough_expr("inStream", var("x"))), + rhs: Box::new(real_lit(2.0)), + span, + }; + let values = builder + .lower_array_like_values(&expr, &Scope::new(), 0) + .expect("zero-size stream elementwise operation should lower to empty values"); + + assert!(values.is_empty()); +} + +#[test] +fn scalar_left_elementwise_zero_shape_stream_passthrough_lowers_to_empty_array_values() { + let span = lower_test_span(); + let layout = VarLayout::from_parts_with_shapes( + IndexMap::new(), + IndexMap::from([("x".to_string(), vec![0])]), + 0, + 0, + ) + .expect("zero-size stream elementwise fixture layout should satisfy shape contract"); + let functions = IndexMap::new(); + let mut builder = LowerBuilder::new(&layout, &functions); + let expr = rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: Box::new(real_lit(2.0)), + rhs: Box::new(stream_passthrough_expr("inStream", var("x"))), + span, + }; + let values = builder + .lower_array_like_values(&expr, &Scope::new(), 0) + .expect("scalar plus zero-size stream elementwise operation should lower to empty values"); + + assert!(values.is_empty()); +} + +#[test] +fn divide_zero_shape_stream_passthrough_lowers_to_empty_array_values() { + let span = lower_test_span(); + let layout = VarLayout::from_parts_with_shapes( + IndexMap::new(), + IndexMap::from([("x".to_string(), vec![0]), ("y".to_string(), vec![0])]), + 0, + 0, + ) + .expect("zero-size stream division fixture layout should satisfy shape contract"); + let functions = IndexMap::new(); + let mut builder = LowerBuilder::new(&layout, &functions); + let expr = rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Div, + lhs: Box::new(stream_passthrough_expr("inStream", var("x"))), + rhs: Box::new(stream_passthrough_expr("inStream", var("y"))), + span, + }; + let values = builder + .lower_array_like_values(&expr, &Scope::new(), 0) + .expect("zero-size stream division should lower to empty values"); + + assert!(values.is_empty()); +} + fn source_fixture_def_id(name: &str) -> rumoca_core::DefId { let hash = name.bytes().fold(2_166_136_261_u32, |hash, byte| { hash.wrapping_mul(16_777_619) ^ u32::from(byte) @@ -978,6 +1662,7 @@ fn eval_linear_ops_collect(ops: &[LinearOp], y: &[f64], p: &[f64], t: f64) -> (V | LinearOp::ImpureRandomInit { dst, .. } | LinearOp::ImpureRandom { dst, .. } | LinearOp::ImpureRandomInteger { dst, .. } => write_reg(&mut regs, dst, 1.0), + LinearOp::ExternalCall { dst, .. } => write_reg(&mut regs, dst, 0.0), LinearOp::Unary { dst, op, arg } => { let value = read_reg(®s, arg); write_reg(&mut regs, dst, apply_unary(op, value)); @@ -1151,6 +1836,19 @@ fn source_var(name: &str) -> rumoca_core::Expression { } } +fn unresolved_source_var(name: &str) -> rumoca_core::Expression { + let mut component_ref = test_component_ref_from_name(name); + component_ref.span = lower_test_span(); + for part in &mut component_ref.parts { + part.span = lower_test_span(); + } + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::from_component_reference(component_ref), + subscripts: vec![], + span: lower_test_span(), + } +} + fn var_index(name: &str, index: i64) -> rumoca_core::Expression { let span = lower_test_span(); rumoca_core::Expression::VarRef { @@ -1213,6 +1911,15 @@ fn der(expr: rumoca_core::Expression) -> rumoca_core::Expression { } } +fn stream_passthrough_expr(name: &str, arg: rumoca_core::Expression) -> rumoca_core::Expression { + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new(name).into(), + args: vec![arg], + is_constructor: false, + span: lower_test_span(), + } +} + fn size_expr(expr: rumoca_core::Expression, dim: i64) -> rumoca_core::Expression { let span = lower_test_span(); rumoca_core::Expression::BuiltinCall { diff --git a/crates/rumoca-phase-solve/src/lower/tests/array_operator_tests.rs b/crates/rumoca-phase-solve/src/lower/tests/array_operator_tests.rs index 7d18de5eb..8c51e5c7c 100644 --- a/crates/rumoca-phase-solve/src/lower/tests/array_operator_tests.rs +++ b/crates/rumoca-phase-solve/src/lower/tests/array_operator_tests.rs @@ -49,6 +49,14 @@ fn field_access(base: rumoca_core::Expression, field: &str) -> rumoca_core::Expr } } +fn generated_var_with_span(name: &str) -> rumoca_core::Expression { + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::generated(name), + subscripts: vec![], + span: lower_test_span(), + } +} + #[test] fn lower_expression_lowers_vector_vector_multiply_as_scalar_product() { let mut dae_model = dae::Dae::default(); @@ -82,6 +90,167 @@ fn lower_expression_lowers_vector_vector_multiply_as_scalar_product() { assert_eq!(read_reg(®s, lowered.result), 65.0); } +#[test] +fn lower_expression_broadcasts_dynamic_scalar_literal_projection_as_array_operand() { + let span = lower_test_span(); + let mut dae_model = dae::Dae::default(); + dae_model + .variables + .parameters + .insert(rumoca_core::VarName::new("i"), scalar_var("i")); + dae_model.variables.parameters.insert( + rumoca_core::VarName::new("v"), + dae::Variable { + dims: vec![2], + ..scalar_var("v") + }, + ); + + let literal_projection = rumoca_core::Expression::Index { + base: Box::new(real_lit(4.0)), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(var("i")), + span, + }], + span, + }; + let expr = source_builtin( + rumoca_core::BuiltinFunction::Sum, + vec![mul(literal_projection, var("v"))], + ); + + let layout = build_var_layout(&dae_model).expect("test DAE layout should build"); + let lowered = lower_expression(&expr, &layout, &IndexMap::new()) + .expect("dynamic scalar literal projection should remain a scalar array operand"); + let mut p = vec![0.0; layout.p_scalars()]; + set_p_value(&layout, &mut p, "i", 2.0); + set_p_value(&layout, &mut p, "v[1]", 3.0); + set_p_value(&layout, &mut p, "v[2]", 5.0); + + let (regs, _) = eval_linear_ops(&lowered.ops, &[], &p, 0.0); + + assert_eq!(read_reg(®s, lowered.result), 32.0); +} + +#[test] +fn lower_expression_lowers_singleton_array_power_as_elementwise_pow() { + let mut dae_model = dae::Dae::default(); + dae_model.variables.parameters.insert( + rumoca_core::VarName::new("x"), + dae::Variable { + dims: vec![1], + ..scalar_var("x") + }, + ); + + let layout = build_var_layout(&dae_model).expect("test DAE layout should build"); + let expr = rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Exp, + lhs: Box::new(var("x")), + rhs: Box::new(real_lit(2.0)), + span: lower_test_span(), + }; + let lowered = lower_expression(&expr, &layout, &IndexMap::new()) + .expect("singleton array power should lower as elementwise pow"); + let mut p = vec![0.0; layout.p_scalars()]; + set_p_value(&layout, &mut p, "x[1]", 3.0); + + let (regs, _) = eval_linear_ops(&lowered.ops, &[], &p, 0.0); + + assert_eq!(read_reg(®s, lowered.result), 9.0); +} + +#[test] +fn lower_expression_projects_complex_vector_dot_from_scalarized_record_fields() { + let mut dae_model = dae::Dae::default(); + for name in ["powerSensor.sum.k[1].re", "powerSensor.sum.k[1].im"] { + dae_model + .variables + .parameters + .insert(rumoca_core::VarName::new(name), scalar_var(name)); + } + for name in [ + "powerSensor.sum.uInternal[1].re", + "powerSensor.sum.uInternal[1].im", + "powerSensor.sum.uInternal[2].re", + "powerSensor.sum.uInternal[2].im", + ] { + dae_model + .variables + .algebraics + .insert(rumoca_core::VarName::new(name), scalar_var(name)); + } + + let layout = build_var_layout(&dae_model).expect("test DAE layout should build"); + let expr = field_access( + mul(var("powerSensor.sum.k"), var("powerSensor.sum.uInternal")), + "re", + ); + let lowered = lower_expression(&expr, &layout, &IndexMap::new()) + .expect("Complex vector product field projection should lower"); + let mut y = vec![0.0; layout.y_scalars()]; + let mut p = vec![0.0; layout.p_scalars()]; + set_p_value(&layout, &mut p, "powerSensor.sum.k[1].re", 2.0); + set_p_value(&layout, &mut p, "powerSensor.sum.k[1].im", 3.0); + set_y_value(&layout, &mut y, "powerSensor.sum.uInternal[1].re", 7.0); + set_y_value(&layout, &mut y, "powerSensor.sum.uInternal[1].im", 11.0); + set_y_value(&layout, &mut y, "powerSensor.sum.uInternal[2].re", 13.0); + set_y_value(&layout, &mut y, "powerSensor.sum.uInternal[2].im", 17.0); + + let (regs, _) = eval_linear_ops(&lowered.ops, &y, &p, 0.0); + + assert_eq!(read_reg(®s, lowered.result), -44.0); +} + +#[test] +fn lower_expression_projects_complex_vector_dot_with_fill_constructor_operand() { + let mut dae_model = dae::Dae::default(); + for name in [ + "powerSensor.sum.uInternal[1].re", + "powerSensor.sum.uInternal[1].im", + "powerSensor.sum.uInternal[2].re", + "powerSensor.sum.uInternal[2].im", + ] { + dae_model + .variables + .algebraics + .insert(rumoca_core::VarName::new(name), scalar_var(name)); + } + + let layout = build_var_layout(&dae_model).expect("test DAE layout should build"); + let fill_complex = builtin( + rumoca_core::BuiltinFunction::Fill, + vec![ + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("Complex").into(), + args: vec![real_lit(1.0), real_lit(0.0)], + is_constructor: true, + span: lower_test_span(), + }, + int_lit(1), + ], + ); + let expr = field_access( + mul( + fill_complex, + generated_var_with_span("powerSensor.sum.uInternal"), + ), + "re", + ); + let lowered = lower_expression(&expr, &layout, &IndexMap::new()) + .expect("Complex fill-vector product field projection should lower"); + let mut y = vec![0.0; layout.y_scalars()]; + let p = vec![0.0; layout.p_scalars()]; + set_y_value(&layout, &mut y, "powerSensor.sum.uInternal[1].re", 7.0); + set_y_value(&layout, &mut y, "powerSensor.sum.uInternal[1].im", 11.0); + set_y_value(&layout, &mut y, "powerSensor.sum.uInternal[2].re", 13.0); + set_y_value(&layout, &mut y, "powerSensor.sum.uInternal[2].im", 17.0); + + let (regs, _) = eval_linear_ops(&lowered.ops, &y, &p, 0.0); + + assert_eq!(read_reg(®s, lowered.result), 20.0); +} + #[test] fn lower_expression_reports_linspace_arity_error_with_argument_span() { let span = rumoca_core::Span::from_offsets( @@ -145,7 +314,7 @@ fn lower_expression_reports_non_scalar_multiply_shape_error_with_operation_span( assert!( err.reason().contains( "non-scalar multiplication result with width 2 is unsupported in scalar context \ - (lhs_shape=[2], rhs_shape=[2, 2], result_shape=[2])" + (lhs_shape=[2], lhs_values=2, rhs_shape=[2, 2], rhs_values=4, result_shape=[2])" ), "{err:?}" ); @@ -203,6 +372,42 @@ fn lower_residual_reports_array_division_denominator_shape_with_operation_span() assert!(!err.reason().contains("rhs_span"), "{err:?}"); } +#[test] +fn lower_expression_reduces_singleton_array_division_inside_max_abs() { + let mut dae_model = dae::Dae::default(); + dae_model.variables.parameters.insert( + rumoca_core::VarName::new("p_deltaq"), + array_var("p_deltaq", &[1]), + ); + dae_model.variables.parameters.insert( + rumoca_core::VarName::new("p_qd_max"), + array_var("p_qd_max", &[1]), + ); + + let layout = build_var_layout(&dae_model).expect("test DAE layout should build"); + let expr = builtin( + rumoca_core::BuiltinFunction::Max, + vec![builtin( + rumoca_core::BuiltinFunction::Abs, + vec![rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Div, + lhs: Box::new(var("p_deltaq")), + rhs: Box::new(var("p_qd_max")), + span: lower_test_span(), + }], + )], + ); + let lowered = lower_expression(&expr, &layout, &IndexMap::new()) + .expect("singleton array division should lower as a scalar quotient"); + let mut p = vec![0.0; layout.p_scalars()]; + set_p_value(&layout, &mut p, "p_deltaq[1]", -10.0); + set_p_value(&layout, &mut p, "p_qd_max[1]", 4.0); + + let (regs, _) = eval_linear_ops(&lowered.ops, &[], &p, 0.0); + + assert_eq!(read_reg(®s, lowered.result), 2.5); +} + #[test] fn lower_residual_projects_scalarized_record_refs_in_array_row() { let mut dae_model = dae::Dae::default(); @@ -553,6 +758,78 @@ fn lower_residual_lowers_slice_of_scalarized_record_field_array() { assert_eq!(outputs, vec![-10.0, -20.0, -30.0]); } +#[test] +fn lower_residual_scalarizes_matrix_column_slice_with_vector_cross_rhs() { + let mut dae_model = dae::Dae::default(); + for (name, dims) in [ + ("leg_v_b", vec![3, 4]), + ("leg_r_b", vec![3, 4]), + ("v_b", vec![3]), + ("omega", vec![3]), + ] { + dae_model.variables.algebraics.insert( + rumoca_core::VarName::new(name), + dae::Variable { + dims, + ..scalar_var(name) + }, + ); + } + + let column = |name: &str| rumoca_core::Expression::Index { + base: Box::new(var(name)), + subscripts: vec![ + rumoca_core::Subscript::generated_colon(lower_test_span()), + rumoca_core::Subscript::generated_index(1, lower_test_span()), + ], + span: lower_test_span(), + }; + dae_model.continuous.equations.push(dae::Equation { + lhs: None, + rhs: sub( + column("leg_v_b"), + add( + var("v_b"), + builtin( + rumoca_core::BuiltinFunction::Cross, + vec![var("omega"), column("leg_r_b")], + ), + ), + ), + span: lower_test_span(), + origin: "matrix column vector kinematics residual".to_string(), + scalar_count: 3, + }); + + let layout = build_var_layout(&dae_model).expect("test DAE layout should build"); + let rows = lower_residual(&dae_model, &layout) + .expect("matrix column slices should scalarize through vector RHS projection"); + let mut y = vec![0.0; layout.y_scalars()]; + for (name, value) in [ + ("leg_v_b[1,1]", 20.0), + ("leg_v_b[2,1]", 30.0), + ("leg_v_b[3,1]", 40.0), + ("leg_r_b[1,1]", 1.0), + ("leg_r_b[2,1]", 2.0), + ("leg_r_b[3,1]", 3.0), + ("v_b[1]", 10.0), + ("v_b[2]", 11.0), + ("v_b[3]", 12.0), + ("omega[1]", 4.0), + ("omega[2]", 5.0), + ("omega[3]", 6.0), + ] { + set_y_value(&layout, &mut y, name, value); + } + + let outputs = rows + .iter() + .map(|row| eval_linear_ops(row, &y, &[], 0.0).1.expect("row output")) + .collect::>(); + + assert_eq!(outputs, vec![7.0, 25.0, 25.0]); +} + #[test] fn lower_residual_lowers_scalarized_record_matrix_field() { let mut dae_model = dae::Dae::default(); @@ -671,6 +948,54 @@ fn lower_residual_lowers_sum_of_nested_scalarized_record_field_array() { assert_eq!(output, Some(0.0)); } +#[test] +fn lower_residual_lowers_sum_of_indexed_record_nested_field_array() { + let mut dae_model = dae::Dae::default(); + for idx in 1..=3 { + let name = format!("source[{idx}].medium.X"); + dae_model + .variables + .algebraics + .insert(rumoca_core::VarName::new(&name), scalar_var(&name)); + } + dae_model.continuous.equations.push(residual(sub( + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(6.0), + span: lower_test_span(), + }, + source_builtin( + rumoca_core::BuiltinFunction::Sum, + vec![rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new("source.medium.X").into(), + subscripts: vec![], + span: lower_test_span(), + }], + ), + ))); + + let layout = build_var_layout(&dae_model).expect("test DAE layout should build"); + let rows = lower_residual(&dae_model, &layout) + .expect("sum over indexed record nested field array should lower"); + let mut y = vec![0.0; layout.y_scalars()]; + for idx in 1..=3 { + set_y_value( + &layout, + &mut y, + &format!("source[{idx}].medium.X"), + idx as f64, + ); + } + + let (_, output) = eval_linear_ops( + rows.last().expect("indexed nested field sum row"), + &y, + &[], + 0.0, + ); + + assert_eq!(output, Some(0.0)); +} + #[test] fn lower_residual_lowers_sum_of_assignment_only_scalarized_record_field_array() { let mut dae_model = dae::Dae::default(); diff --git a/crates/rumoca-phase-solve/src/lower/tests/array_operator_tests/discrete_array_tests.rs b/crates/rumoca-phase-solve/src/lower/tests/array_operator_tests/discrete_array_tests.rs index 6183a93d0..7354b6d1d 100644 --- a/crates/rumoca-phase-solve/src/lower/tests/array_operator_tests/discrete_array_tests.rs +++ b/crates/rumoca-phase-solve/src/lower/tests/array_operator_tests/discrete_array_tests.rs @@ -1707,6 +1707,117 @@ fn lower_residual_flattens_all_outputs_of_tuple_function_call() { assert_eq!(actual, vec![-1.0, -2.0, -3.0]); } +#[test] +fn lower_residual_counts_only_bound_tuple_function_outputs() { + let mut dae_model = dae::Dae::default(); + dae_model.variables.parameters.insert( + rumoca_core::VarName::new("r"), + dae::Variable { + dims: vec![3], + ..scalar_var("r") + }, + ); + dae_model.variables.parameters.insert( + rumoca_core::VarName::new("cr"), + dae::Variable { + dims: vec![3], + ..scalar_var("cr") + }, + ); + + let mut function = rumoca_core::Function::new("Pkg.roots", lower_test_span()); + function + .inputs + .push(function_param_with_dims("cr_in", &[0])); + function + .inputs + .push(function_param_with_dims("c0_in", &[0])); + function + .inputs + .push(function_param_with_dims("c1_in", &[0])); + function.inputs.push(function_param_with_dims("f_cut", &[])); + function + .outputs + .push(function_param_with_dims("r", &[0]).with_shape_expr(vec![size_shape_expr("cr_in")])); + for output in ["a", "b", "ku"] { + function.outputs.push( + function_param_with_dims(output, &[0]).with_shape_expr(vec![size_shape_expr("c0_in")]), + ); + } + function.body.push(rumoca_core::Statement::Assignment { + comp: component_ref("r"), + value: rumoca_core::Expression::Array { + elements: vec![real_lit(1.0), real_lit(2.0), real_lit(3.0)], + is_matrix: false, + span: lower_test_span(), + }, + span: lower_test_span(), + }); + dae_model + .symbols + .functions + .insert(function.name.clone(), function); + let placeholder = || rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("Real").into(), + args: vec![rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Boolean(false), + span: lower_test_span(), + }], + is_constructor: true, + span: lower_test_span(), + }; + dae_model.continuous.equations.push(dae::Equation { + lhs: None, + rhs: sub( + rumoca_core::Expression::Tuple { + elements: vec![var("r"), placeholder(), placeholder(), placeholder()], + span: lower_test_span(), + }, + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("Pkg.roots").into(), + args: vec![var("cr"), placeholder(), placeholder(), real_lit(1.0)], + is_constructor: false, + span: lower_test_span(), + }, + ), + span: lower_test_span(), + origin: "tuple function output residual with placeholders".to_string(), + scalar_count: 6, + }); + assert_eq!( + expression_rows::residual_equation_effective_row_count( + &dae_model, + dae_model + .continuous + .equations + .last() + .expect("test equation should exist"), + ) + .expect("call-site output shape should determine bound residual rows"), + 3 + ); + + let layout = build_var_layout(&dae_model).expect("test DAE layout should build"); + let rows = lower_residual(&dae_model, &layout) + .expect("placeholder tuple outputs should not declare solve rows"); + let p = vec![0.0; layout.p_scalars()]; + let actual = eval_programs_all_outputs(&rows, &[], &p, 0.0); + + assert_eq!(rows.len(), 3); + assert_eq!(actual, vec![-1.0, -2.0, -3.0]); +} + +fn size_shape_expr(name: &str) -> rumoca_core::Subscript { + rumoca_core::Subscript::Expr { + expr: Box::new(builtin_with_span( + rumoca_core::BuiltinFunction::Size, + vec![var(name), int_lit_with_span(1, lower_test_span())], + lower_test_span(), + )), + span: lower_test_span(), + } +} + #[test] fn lower_discrete_rhs_lowers_easy_array_builtins() { let mut dae_model = dae::Dae::default(); diff --git a/crates/rumoca-phase-solve/src/lower/tests/function_expression_tests.rs b/crates/rumoca-phase-solve/src/lower/tests/function_expression_tests.rs index 49cfc2aad..d65ff7368 100644 --- a/crates/rumoca-phase-solve/src/lower/tests/function_expression_tests.rs +++ b/crates/rumoca-phase-solve/src/lower/tests/function_expression_tests.rs @@ -1,3 +1,6 @@ +// SPEC_0021 file-size exception: function-expression lowering tests still +// share builders across scalar, array, and record cases. split plan: move +// builtin, projection, and record tests into sibling modules. use super::*; mod shape_diagnostic_tests; @@ -353,6 +356,540 @@ fn lower_expression_binds_named_record_constructor_input_fields() { assert_eq!(read_reg(®s, lowered.result), 7.0); } +#[test] +fn lower_expression_binds_alias_qualified_record_constructor_input_fields() { + let mut function = rumoca_core::Function::new("Pkg.recordInput", lower_test_span()); + function.inputs.push(rumoca_core::FunctionParam { + type_class: Some(rumoca_core::ClassType::Record), + type_name: "Modelica.Media.IdealGases.Common.DataRecord".to_string(), + ..rumoca_core::FunctionParam::new( + "data", + "Modelica.Media.IdealGases.Common.DataRecord", + lower_test_span(), + ) + }); + function.outputs.push(function_param("y")); + function.body.push(rumoca_core::Statement::Assignment { + comp: component_ref("y"), + value: rumoca_core::Expression::FieldAccess { + base: Box::new(var("data")), + field: "R_s".to_string(), + span: lower_test_span(), + }, + span: lower_test_span(), + }); + + let mut constructor = rumoca_core::Function::new( + "Modelica.Media.IdealGases.Common.DataRecord", + lower_test_span(), + ); + constructor.inputs.push(rumoca_core::FunctionParam::new( + "R_s", + "Real", + lower_test_span(), + )); + + let mut functions = IndexMap::new(); + functions.insert(function.name.clone(), function); + functions.insert(constructor.name.clone(), constructor); + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( + "Pkg.recordInput", + )), + args: vec![rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( + "IdealGases.Common.DataRecord", + )), + args: vec![named_arg("R_s", real_lit(287.0))], + is_constructor: true, + span: lower_test_span(), + }], + is_constructor: false, + span: lower_test_span(), + }; + + let lowered = lower_expression(&expr, &VarLayout::default(), &functions) + .expect("alias-qualified record constructor input fields should lower"); + let (regs, _) = eval_linear_ops(&lowered.ops, &[], &[], 0.0); + + assert_eq!(read_reg(®s, lowered.result), 287.0); +} + +#[test] +fn lower_expression_binds_named_record_constructor_without_registered_field_list() { + let mut function = rumoca_core::Function::new("Pkg.recordInput", lower_test_span()); + function.inputs.push(rumoca_core::FunctionParam { + type_class: Some(rumoca_core::ClassType::Record), + type_name: "Pkg.UnregisteredData".to_string(), + ..rumoca_core::FunctionParam::new("data", "Pkg.UnregisteredData", lower_test_span()) + }); + function.outputs.push(function_param("y")); + function.body.push(rumoca_core::Statement::Assignment { + comp: component_ref("y"), + value: rumoca_core::Expression::FieldAccess { + base: Box::new(var("data")), + field: "R_s".to_string(), + span: lower_test_span(), + }, + span: lower_test_span(), + }); + + let mut functions = IndexMap::new(); + functions.insert(function.name.clone(), function); + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( + "Pkg.recordInput", + )), + args: vec![rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( + "Pkg.UnregisteredData", + )), + args: vec![named_arg("R_s", real_lit(287.0))], + is_constructor: true, + span: lower_test_span(), + }], + is_constructor: false, + span: lower_test_span(), + }; + + let lowered = lower_expression(&expr, &VarLayout::default(), &functions) + .expect("named record constructor fields should bind without constructor metadata"); + let (regs, _) = eval_linear_ops(&lowered.ops, &[], &[], 0.0); + + assert_eq!(read_reg(®s, lowered.result), 287.0); +} + +#[test] +fn lower_expression_binds_partial_function_input_closure() { + let mut scale = rumoca_core::Function::new("Pkg.scale", lower_test_span()); + scale.inputs.push(function_param("u")); + scale.outputs.push(function_param("y")); + scale.body.push(rumoca_core::Statement::Assignment { + comp: component_ref("y"), + value: rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs: Box::new(var("u")), + rhs: Box::new(real_lit(2.0)), + span: lower_test_span(), + }, + span: lower_test_span(), + }); + + let mut apply = rumoca_core::Function::new("Pkg.apply", lower_test_span()); + apply.inputs.push(rumoca_core::FunctionParam { + type_class: Some(rumoca_core::ClassType::Function), + type_name: "Modelica.Math.Nonlinear.Interfaces.partialScalarFunction".to_string(), + ..rumoca_core::FunctionParam::new( + "f", + "Modelica.Math.Nonlinear.Interfaces.partialScalarFunction", + lower_test_span(), + ) + }); + apply.outputs.push(function_param("y")); + apply.body.push(rumoca_core::Statement::Assignment { + comp: component_ref("y"), + value: rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference(source_component_ref_from_name( + "f", + )), + args: vec![real_lit(3.0)], + is_constructor: false, + span: lower_test_span(), + }, + span: lower_test_span(), + }); + + let mut functions = IndexMap::new(); + functions.insert(scale.name.clone(), scale); + functions.insert(apply.name.clone(), apply); + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( + "Pkg.apply", + )), + args: vec![rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( + "Pkg.scale", + )), + args: Vec::new(), + is_constructor: true, + span: lower_test_span(), + }], + is_constructor: false, + span: lower_test_span(), + }; + + let lowered = lower_expression(&expr, &VarLayout::default(), &functions) + .expect("partial function input closure should lower"); + let (regs, _) = eval_linear_ops(&lowered.ops, &[], &[], 0.0); + + assert_eq!(read_reg(®s, lowered.result), 6.0); +} + +#[test] +#[allow(clippy::too_many_lines)] +fn lower_expression_evaluates_captured_partial_function_in_dynamic_while() { + let mut sine_residual = rumoca_core::Function::new("Pkg.sineResidual", lower_test_span()); + for input in ["u", "A", "w", "s"] { + sine_residual.inputs.push(function_param(input)); + } + sine_residual.outputs.push(function_param("y")); + sine_residual.body.push(rumoca_core::Statement::Assignment { + comp: component_ref("y"), + value: rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: Box::new(rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs: Box::new(var("A")), + rhs: Box::new(rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Sin, + args: vec![rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Mul, + lhs: Box::new(var("w")), + rhs: Box::new(var("u")), + span: lower_test_span(), + }], + span: lower_test_span(), + }), + span: lower_test_span(), + }), + rhs: Box::new(var("s")), + span: lower_test_span(), + }, + span: lower_test_span(), + }); + + let mut solve = rumoca_core::Function::new("Pkg.solve", lower_test_span()); + solve.inputs.push(rumoca_core::FunctionParam { + type_class: Some(rumoca_core::ClassType::Function), + type_name: "Modelica.Math.Nonlinear.Interfaces.partialScalarFunction".to_string(), + ..rumoca_core::FunctionParam::new( + "f", + "Modelica.Math.Nonlinear.Interfaces.partialScalarFunction", + lower_test_span(), + ) + }); + solve.inputs.push(function_param("u_min")); + solve.inputs.push(function_param("u_max")); + solve.outputs.push(function_param("root")); + for local in ["a", "b", "mid", "f_mid", "i"] { + solve.locals.push(function_param(local)); + } + solve.locals.push( + rumoca_core::FunctionParam::new("found", "Boolean", lower_test_span()).with_default( + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Boolean(false), + span: lower_test_span(), + }, + ), + ); + let assign = |name: &str, value: rumoca_core::Expression| rumoca_core::Statement::Assignment { + comp: component_ref(name), + value, + span: lower_test_span(), + }; + let binary = |op, lhs, rhs| rumoca_core::Expression::Binary { + op, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span: lower_test_span(), + }; + solve.body = vec![ + assign("a", var("u_min")), + assign("b", var("u_max")), + assign("i", real_lit(0.0)), + rumoca_core::Statement::While { + block: rumoca_core::StatementBlock { + cond: rumoca_core::Expression::Unary { + op: rumoca_core::OpUnary::Not, + rhs: Box::new(var("found")), + span: lower_test_span(), + }, + stmts: vec![ + assign( + "mid", + binary( + rumoca_core::OpBinary::Div, + binary(rumoca_core::OpBinary::Add, var("a"), var("b")), + real_lit(2.0), + ), + ), + assign( + "f_mid", + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference( + source_component_ref_from_name("f"), + ), + args: vec![var("mid")], + is_constructor: false, + span: lower_test_span(), + }, + ), + rumoca_core::Statement::If { + cond_blocks: vec![rumoca_core::StatementBlock { + cond: binary(rumoca_core::OpBinary::Ge, var("i"), real_lit(59.0)), + stmts: vec![ + assign( + "found", + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Boolean(true), + span: lower_test_span(), + }, + ), + assign("root", var("mid")), + ], + }], + else_block: Some(vec![ + rumoca_core::Statement::If { + cond_blocks: vec![rumoca_core::StatementBlock { + cond: binary( + rumoca_core::OpBinary::Gt, + var("f_mid"), + real_lit(0.0), + ), + stmts: vec![assign("b", var("mid"))], + }], + else_block: Some(vec![assign("a", var("mid"))]), + span: lower_test_span(), + }, + assign( + "i", + binary(rumoca_core::OpBinary::Add, var("i"), real_lit(1.0)), + ), + ]), + span: lower_test_span(), + }, + ], + }, + span: lower_test_span(), + }, + ]; + + let mut functions = IndexMap::new(); + functions.insert(sine_residual.name.clone(), sine_residual); + functions.insert(solve.name.clone(), solve); + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( + "Pkg.solve", + )), + args: vec![ + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference( + test_component_ref_from_name("Pkg.sineResidual"), + ), + args: vec![ + named_arg("A", real_lit(1.0)), + named_arg("w", real_lit(1.0)), + named_arg("s", real_lit(-0.5)), + ], + is_constructor: true, + span: lower_test_span(), + }, + real_lit(-1.7), + real_lit(1.7), + ], + is_constructor: false, + span: lower_test_span(), + }; + + let mut dae_model = dae::Dae::default(); + dae_model.symbols.functions = functions.clone(); + let projected = crate::lower::derivative_rhs::function_call_projected_scalars_with_owner( + &expr, + &dae_model, + &IndexMap::new(), + lower_test_span(), + ) + .expect("dynamic scalar function calls should decline projection without failing lowering"); + assert!(projected.is_none()); + + let lowered = lower_expression(&expr, &VarLayout::default(), &functions) + .expect("captured partial function in a dynamic while should lower"); + let (regs, _) = eval_linear_ops(&lowered.ops, &[], &[], 0.0); + + assert!((read_reg(®s, lowered.result) - 0.5_f64.asin()).abs() <= 1e-12); + + dae_model + .variables + .algebraics + .insert(rumoca_core::VarName::new("x_zero"), scalar_var("x_zero")); + dae_model.continuous.equations.push(dae::Equation { + lhs: None, + rhs: binary(rumoca_core::OpBinary::Sub, var("x_zero"), expr), + span: lower_test_span(), + origin: "dynamic while function residual".to_string(), + scalar_count: 1, + }); + let layout = build_var_layout(&dae_model).expect("residual layout should build"); + let rows = lower_residual(&dae_model, &layout) + .expect("dynamic while projection should decline to ordinary residual lowering"); + let residual = eval_linear_ops(&rows[0], &[0.0], &[], 0.0) + .1 + .expect("residual output"); + assert!((residual + 0.5_f64.asin()).abs() <= 1e-12); +} + +#[test] +fn unprojectable_array_output_declines_scalar_lane_fallback_and_uses_array_runtime() { + let span = lower_test_span(); + let mut projection_declined = rumoca_core::Function::new("Pkg.projectionDeclinedArray", span); + projection_declined.inputs.push(function_param("u")); + projection_declined + .outputs + .push(function_param_with_dims("y", &[2])); + projection_declined + .locals + .push(record_param("scratch", "Pkg.Record")); + projection_declined.body = vec![ + rumoca_core::Statement::Assignment { + comp: component_ref("scratch"), + value: rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference( + test_component_ref_from_name("Pkg.Record"), + ), + args: vec![ + named_arg("a", var("u")), + named_arg("b", add(var("u"), real_lit(1.0))), + ], + is_constructor: true, + span, + }, + span, + }, + rumoca_core::Statement::Assignment { + comp: component_ref("y"), + value: rumoca_core::Expression::Array { + elements: vec![var("scratch.a"), var("scratch.b")], + is_matrix: false, + span, + }, + span, + }, + ]; + + let projection_call = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_var_name(projection_declined.name.clone()), + args: vec![source_var("u")], + is_constructor: false, + span, + }; + let mut dae_model = dae::Dae::default(); + let mut record_constructor = rumoca_core::Function::new("Pkg.Record", span); + record_constructor.is_constructor = true; + record_constructor.inputs.push(function_param("a")); + record_constructor.inputs.push(function_param("b")); + record_constructor + .outputs + .push(record_param("record", "Pkg.Record")); + dae_model + .variables + .parameters + .insert(rumoca_core::VarName::new("u"), scalar_var("u")); + dae_model.variables.algebraics.insert( + rumoca_core::VarName::new("target"), + array_var("target", &[2]), + ); + dae_model + .symbols + .functions + .insert(record_constructor.name.clone(), record_constructor); + dae_model + .symbols + .functions + .insert(projection_declined.name.clone(), projection_declined); + + let projected = crate::lower::derivative_rhs::project_array_like_scalars_with_owner( + &projection_call, + &dae_model, + &IndexMap::new(), + span, + ) + .expect("unprojectable array output should decline without an error"); + assert!( + projected.is_none(), + "a whole array call must not be duplicated as scalar lanes: {projected:?}" + ); + + dae_model.continuous.equations.push(dae::Equation { + lhs: None, + rhs: sub(source_var("target"), projection_call), + span, + origin: "unprojectable array function residual".to_string(), + scalar_count: 2, + }); + let layout = build_var_layout(&dae_model).expect("array residual layout should build"); + let rows = lower_residual(&dae_model, &layout) + .expect("unprojectable array output should use array runtime lowering"); + + assert_eq!(rows.len(), 2); + let mut p = vec![0.0; layout.p_scalars()]; + set_p_value(&layout, &mut p, "u", 2.0); + let values = eval_programs_all_outputs(&rows, &[0.0, 0.0], &p, 0.0); + assert_eq!(values, vec![-2.0, -3.0]); +} + +#[test] +fn lower_expression_binds_constructor_actual_to_flattened_record_inputs_with_defaults() { + let mut function = rumoca_core::Function::new("Pkg.drop", lower_test_span()); + function + .inputs + .push(function_param("brushParameters_V").with_default(real_lit(0.0))); + function + .inputs + .push(function_param("brushParameters_ILinear")); + function.inputs.push(function_param("i")); + function.outputs.push(function_param("v")); + function.body.push(rumoca_core::Statement::Assignment { + comp: component_ref("v"), + value: rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: Box::new(var("brushParameters_V")), + rhs: Box::new(var("brushParameters_ILinear")), + span: lower_test_span(), + }, + span: lower_test_span(), + }); + + let mut constructor = rumoca_core::Function::new( + "Modelica.Electrical.Machines.Losses.BrushParameters", + lower_test_span(), + ); + constructor.is_constructor = true; + constructor + .inputs + .push(function_param("V").with_default(real_lit(0.0))); + constructor.inputs.push(function_param("ILinear")); + + let mut functions = IndexMap::new(); + functions.insert(function.name.clone(), function); + functions.insert(constructor.name.clone(), constructor); + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( + "Pkg.drop", + )), + args: vec![ + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference( + test_component_ref_from_name( + "Modelica.Electrical.Machines.Losses.BrushParameters", + ), + ), + args: vec![named_arg("ILinear", real_lit(4.0))], + is_constructor: true, + span: lower_test_span(), + }, + real_lit(100.0), + ], + is_constructor: false, + span: lower_test_span(), + }; + + let lowered = lower_expression(&expr, &VarLayout::default(), &functions) + .expect("record constructor actual should bind flattened input fields"); + let (regs, _) = eval_linear_ops(&lowered.ops, &[], &[], 0.0); + + assert_eq!(read_reg(®s, lowered.result), 4.0); +} + #[test] fn lower_expression_projects_record_output_assigned_from_if_constructor() { let span = lower_test_span(); @@ -1013,32 +1550,184 @@ fn lower_expression_binds_same_named_local_record_actual_to_record_input() { dae_model .symbols .functions - .insert(use_local.name.clone(), use_local); + .insert(use_local.name.clone(), use_local); + + let mut build_aux = rumoca_core::Function::new("My.buildAux", span); + build_aux.inputs.push(function_param("u")); + build_aux.outputs.push(record_param("aux", "My.AuxRecord")); + build_aux.locals.push(record_param("f", "My.LocalRecord")); + build_aux.body.push(rumoca_core::Statement::Assignment { + comp: component_ref("f"), + value: rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( + "My.makeLocal", + )), + args: vec![var("u")], + is_constructor: false, + span, + }, + span, + }); + build_aux.body.push(rumoca_core::Statement::Assignment { + comp: component_ref("aux"), + value: rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( + "My.useLocal", + )), + args: vec![var("f")], + is_constructor: false, + span, + }, + span, + }); + dae_model + .symbols + .functions + .insert(build_aux.name.clone(), build_aux); + + let expr = rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( + "My.buildAux", + )), + args: vec![var("u")], + is_constructor: false, + span, + }), + field: "rho".to_string(), + span, + }; + + let layout = build_var_layout(&dae_model).expect("test DAE layout should build"); + let lowered = lower_expression(&expr, &layout, &dae_model.symbols.functions) + .expect("same-named record actual should bind components into callee input scope"); + let mut y = vec![0.0; layout.y_scalars()]; + let p = vec![]; + set_y_value(&layout, &mut y, "u", 5.0); + + let (regs, _output) = eval_linear_ops(&lowered.ops, &y, &p, 0.0); + let compiled = read_reg(®s, lowered.result); + assert!((compiled - 15.0).abs() <= 1e-12); +} + +#[test] +fn lower_expression_projects_record_field_from_function_result() { + let mut dae_model = dae::Dae::default(); + let span = lower_test_span(); + dae_model + .variables + .algebraics + .insert(rumoca_core::VarName::new("p"), scalar_var("p")); + dae_model + .variables + .algebraics + .insert(rumoca_core::VarName::new("temp"), scalar_var("temp")); + + let mut state_ctor = rumoca_core::Function::new("My.State", lower_test_span()); + state_ctor.is_constructor = true; + state_ctor.inputs.push(function_param("p")); + state_ctor.inputs.push(function_param("T")); + state_ctor.outputs.push(record_param("state", "My.State")); + dae_model + .symbols + .functions + .insert(state_ctor.name.clone(), state_ctor); + + let mut make_state = rumoca_core::Function::new("My.makeState", lower_test_span()); + make_state.inputs.push(function_param("p")); + make_state.inputs.push(function_param("T")); + make_state.outputs.push(record_param("state", "My.State")); + make_state.body.push(rumoca_core::Statement::Assignment { + comp: component_ref("state"), + value: rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( + "My.State", + )), + args: vec![var("p"), var("T")], + is_constructor: true, + span, + }, + span, + }); + dae_model + .symbols + .functions + .insert(make_state.name.clone(), make_state); + + let expr = rumoca_core::Expression::FieldAccess { + base: Box::new(rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( + "My.makeState", + )), + args: vec![var("p"), var("temp")], + is_constructor: false, + span, + }), + field: "T".to_string(), + span, + }; + + let layout = build_var_layout(&dae_model).expect("test DAE layout should build"); + let lowered = lower_expression(&expr, &layout, &dae_model.symbols.functions) + .expect("record-valued function field projection should lower"); + let mut y = vec![0.0; layout.y_scalars()]; + let p = vec![]; + set_y_value(&layout, &mut y, "p", 101325.0); + set_y_value(&layout, &mut y, "temp", 360.0); + + let (regs, _output) = eval_linear_ops(&lowered.ops, &y, &p, 0.0); + let compiled = read_reg(®s, lowered.result); + assert!((compiled - 360.0).abs() <= 1e-12); +} + +#[test] +fn lower_expression_projects_named_record_field_from_function_result() { + let mut dae_model = dae::Dae::default(); + let span = lower_test_span(); + dae_model + .variables + .algebraics + .insert(rumoca_core::VarName::new("p"), scalar_var("p")); + dae_model + .variables + .algebraics + .insert(rumoca_core::VarName::new("temp"), scalar_var("temp")); + + let mut state_ctor = rumoca_core::Function::new("My.State", lower_test_span()); + state_ctor.is_constructor = true; + state_ctor.inputs.push(function_param("p")); + state_ctor.inputs.push(function_param("T")); + state_ctor.outputs.push(record_param("state", "My.State")); + dae_model + .symbols + .functions + .insert(state_ctor.name.clone(), state_ctor); - let mut build_aux = rumoca_core::Function::new("My.buildAux", span); - build_aux.inputs.push(function_param("u")); - build_aux.outputs.push(record_param("aux", "My.AuxRecord")); - build_aux.locals.push(record_param("f", "My.LocalRecord")); - build_aux.body.push(rumoca_core::Statement::Assignment { - comp: component_ref("f"), - value: rumoca_core::Expression::FunctionCall { - name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( - "My.makeLocal", - )), - args: vec![var("u")], - is_constructor: false, - span, - }, - span, - }); - build_aux.body.push(rumoca_core::Statement::Assignment { - comp: component_ref("aux"), + let mut make_state = rumoca_core::Function::new("My.makeState", lower_test_span()); + make_state.inputs.push(function_param("p")); + make_state.inputs.push(function_param("T")); + make_state.outputs.push(record_param("state", "My.State")); + make_state.body.push(rumoca_core::Statement::Assignment { + comp: component_ref("state"), value: rumoca_core::Expression::FunctionCall { name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( - "My.useLocal", + "My.State", )), - args: vec![var("f")], - is_constructor: false, + args: vec![ + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("__rumoca_named_arg__.p").into(), + args: vec![var("p")], + is_constructor: true, + span, + }, + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("__rumoca_named_arg__.T").into(), + args: vec![var("T")], + is_constructor: true, + span, + }, + ], + is_constructor: true, span, }, span, @@ -1046,35 +1735,50 @@ fn lower_expression_binds_same_named_local_record_actual_to_record_input() { dae_model .symbols .functions - .insert(build_aux.name.clone(), build_aux); + .insert(make_state.name.clone(), make_state); let expr = rumoca_core::Expression::FieldAccess { base: Box::new(rumoca_core::Expression::FunctionCall { name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( - "My.buildAux", + "My.makeState", )), - args: vec![var("u")], + args: vec![ + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("__rumoca_named_arg__.p").into(), + args: vec![var("p")], + is_constructor: true, + span, + }, + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("__rumoca_named_arg__.T").into(), + args: vec![var("temp")], + is_constructor: true, + span, + }, + ], is_constructor: false, span, }), - field: "rho".to_string(), + field: "p".to_string(), span, }; let layout = build_var_layout(&dae_model).expect("test DAE layout should build"); let lowered = lower_expression(&expr, &layout, &dae_model.symbols.functions) - .expect("same-named record actual should bind components into callee input scope"); + .expect("named record-valued function field projection should lower"); let mut y = vec![0.0; layout.y_scalars()]; let p = vec![]; - set_y_value(&layout, &mut y, "u", 5.0); + set_y_value(&layout, &mut y, "p", 101325.0); + set_y_value(&layout, &mut y, "temp", 360.0); let (regs, _output) = eval_linear_ops(&lowered.ops, &y, &p, 0.0); let compiled = read_reg(®s, lowered.result); - assert!((compiled - 15.0).abs() <= 1e-12); + assert!((compiled - 101325.0).abs() <= 1e-12); } #[test] -fn lower_expression_projects_record_field_from_function_result() { +#[allow(clippy::too_many_lines)] +fn lower_expression_binds_flattened_named_record_field_actual() { let mut dae_model = dae::Dae::default(); let span = lower_test_span(); dae_model @@ -1102,12 +1806,49 @@ fn lower_expression_projects_record_field_from_function_result() { make_state.outputs.push(record_param("state", "My.State")); make_state.body.push(rumoca_core::Statement::Assignment { comp: component_ref("state"), - value: rumoca_core::Expression::FunctionCall { - name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( - "My.State", - )), - args: vec![var("p"), var("T")], - is_constructor: true, + value: rumoca_core::Expression::If { + branches: vec![( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Eq, + lhs: Box::new(var("p")), + rhs: Box::new(var("p")), + span, + }, + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference( + test_component_ref_from_name("My.State"), + ), + args: vec![ + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("__rumoca_named_arg__.p").into(), + args: vec![var("p")], + is_constructor: true, + span, + }, + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("__rumoca_named_arg__.T").into(), + args: vec![rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: Box::new(var("p")), + rhs: Box::new(var("T")), + span, + }], + is_constructor: true, + span, + }, + ], + is_constructor: true, + span, + }, + )], + else_branch: Box::new(rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference( + test_component_ref_from_name("My.State"), + ), + args: vec![var("p"), var("T")], + is_constructor: true, + span, + }), span, }, span, @@ -1117,22 +1858,57 @@ fn lower_expression_projects_record_field_from_function_result() { .functions .insert(make_state.name.clone(), make_state); - let expr = rumoca_core::Expression::FieldAccess { + let mut density = rumoca_core::Function::new("My.density", lower_test_span()); + density.inputs.push(function_param("state_p")); + density.inputs.push(function_param("state_T")); + density.outputs.push(function_param("d")); + density.body.push(rumoca_core::Statement::Assignment { + comp: component_ref("d"), + value: var("state_p"), + span, + }); + dae_model + .symbols + .functions + .insert(density.name.clone(), density); + + let state_p = rumoca_core::Expression::FieldAccess { base: Box::new(rumoca_core::Expression::FunctionCall { name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( "My.makeState", )), - args: vec![var("p"), var("temp")], + args: vec![ + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("__rumoca_named_arg__.p").into(), + args: vec![var("p")], + is_constructor: true, + span, + }, + rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("__rumoca_named_arg__.T").into(), + args: vec![var("temp")], + is_constructor: true, + span, + }, + ], is_constructor: false, span, }), - field: "T".to_string(), + field: "p".to_string(), + span, + }; + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( + "My.density", + )), + args: vec![state_p], + is_constructor: false, span, }; let layout = build_var_layout(&dae_model).expect("test DAE layout should build"); let lowered = lower_expression(&expr, &layout, &dae_model.symbols.functions) - .expect("record-valued function field projection should lower"); + .expect("flattened named record field actual should lower"); let mut y = vec![0.0; layout.y_scalars()]; let p = vec![]; set_y_value(&layout, &mut y, "p", 101325.0); @@ -1140,7 +1916,7 @@ fn lower_expression_projects_record_field_from_function_result() { let (regs, _output) = eval_linear_ops(&lowered.ops, &y, &p, 0.0); let compiled = read_reg(®s, lowered.result); - assert!((compiled - 360.0).abs() <= 1e-12); + assert!((compiled - 101325.0).abs() <= 1e-12); } #[test] @@ -1754,6 +2530,47 @@ fn lower_expression_binds_projected_real_component_to_complex_input() { assert!((read_reg(®s, lowered.result) - 4.5).abs() < 1e-12); } +#[test] +fn lower_expression_synthesizes_sibling_flattened_record_input_field() { + let mut dae_model = dae::Dae::default(); + let mut use_complex = rumoca_core::Function::new("My.useComplex", lower_test_span()); + use_complex.inputs.push(function_param("c1_re")); + use_complex.inputs.push(function_param("c1_im")); + use_complex.outputs.push(function_param("y")); + use_complex.body.push(rumoca_core::Statement::Assignment { + comp: component_ref("y"), + value: var("c1_im"), + span: lower_test_span(), + }); + dae_model + .symbols + .functions + .insert(use_complex.name.clone(), use_complex); + dae_model + .variables + .algebraics + .insert(rumoca_core::VarName::new("u.re"), scalar_var("u.re")); + dae_model + .variables + .algebraics + .insert(rumoca_core::VarName::new("u.im"), scalar_var("u.im")); + + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( + "My.useComplex", + )), + args: vec![source_var("u.re")], + is_constructor: false, + span: lower_test_span(), + }; + let layout = build_var_layout(&dae_model).expect("test DAE layout should build"); + let lowered = lower_expression(&expr, &layout, &dae_model.symbols.functions) + .expect("flattened record sibling field should be projected from the same actual base"); + let (regs, _) = eval_linear_ops(&lowered.ops, &[4.5, -2.0], &[], 0.0); + + assert!((read_reg(®s, lowered.result) + 2.0).abs() < 1e-12); +} + #[test] fn lower_expression_rebinds_flattened_record_input_components() { let mut dae_model = dae::Dae::default(); @@ -1821,6 +2638,164 @@ fn lower_expression_rebinds_flattened_record_input_components() { assert!((read_reg(®s, lowered.result) - 4.0).abs() < 1e-12); } +#[test] +fn lower_expression_binds_zero_dim_flattened_record_array_field_from_actual_shape() { + let mut dae_model = dae::Dae::default(); + let span = lower_test_span(); + + let mut cp = rumoca_core::Function::new("My.cp", span); + cp.inputs.push(function_param("state_p")); + cp.inputs.push(function_param("state_T")); + cp.inputs.push( + function_param_with_dims("state_X", &[0]).with_shape_expr(vec![ + rumoca_core::Subscript::generated_expr(Box::new(var("nX")), span), + ]), + ); + cp.outputs.push(function_param("y")); + cp.body.push(rumoca_core::Statement::Assignment { + comp: component_ref("y"), + value: var_index("state_X", 1), + span, + }); + dae_model.symbols.functions.insert(cp.name.clone(), cp); + dae_model + .variables + .algebraics + .insert(rumoca_core::VarName::new("s.p"), source_scalar_var("s.p")); + dae_model + .variables + .algebraics + .insert(rumoca_core::VarName::new("s.T"), source_scalar_var("s.T")); + dae_model.variables.algebraics.insert( + rumoca_core::VarName::new("s.X"), + source_array_var("s.X", &[2]), + ); + + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( + "My.cp", + )), + args: vec![source_var("s.p"), source_var("s.T"), source_var("s.X")], + is_constructor: false, + span, + }; + let layout = build_var_layout(&dae_model).expect("test DAE layout should build"); + let lowered = lower_expression(&expr, &layout, &dae_model.symbols.functions) + .expect("zero-dim flattened record array field should bind from actual shape"); + let (regs, _) = eval_linear_ops(&lowered.ops, &[101325.0, 300.0, 0.42, 0.58], &[], 0.0); + + assert!((read_reg(®s, lowered.result) - 0.42).abs() < 1e-12); +} + +#[test] +fn lower_expression_binds_singleton_vectorized_record_array_field_to_vector_input() { + let mut dae_model = dae::Dae::default(); + let span = lower_test_span(); + + let mut cp = rumoca_core::Function::new("My.cp", span); + cp.inputs.push(function_param("state_p")); + cp.inputs.push(function_param("state_T")); + cp.inputs.push( + function_param_with_dims("state_X", &[0]).with_shape_expr(vec![ + rumoca_core::Subscript::generated_expr(Box::new(var("nX")), span), + ]), + ); + cp.outputs.push(function_param("y")); + cp.body.push(rumoca_core::Statement::Assignment { + comp: component_ref("y"), + value: var_index("state_X", 2), + span, + }); + dae_model.symbols.functions.insert(cp.name.clone(), cp); + dae_model.variables.algebraics.insert( + rumoca_core::VarName::new("states.p"), + source_array_var("states.p", &[1]), + ); + dae_model.variables.algebraics.insert( + rumoca_core::VarName::new("states.T"), + source_array_var("states.T", &[1]), + ); + dae_model.variables.algebraics.insert( + rumoca_core::VarName::new("states.X"), + source_array_var("states.X", &[1, 2]), + ); + + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( + "My.cp", + )), + args: vec![ + source_var("states.p"), + source_var("states.T"), + source_var("states.X"), + ], + is_constructor: false, + span, + }; + let layout = build_var_layout(&dae_model).expect("test DAE layout should build"); + let lowered = lower_expression(&expr, &layout, &dae_model.symbols.functions) + .expect("singleton vectorized record array field should bind to vector input"); + let (regs, _) = eval_linear_ops(&lowered.ops, &[101325.0, 300.0, 0.42, 0.58], &[], 0.0); + + assert!((read_reg(®s, lowered.result) - 0.58).abs() < 1e-12); +} + +#[test] +fn lower_expression_rejects_non_singleton_vectorized_record_array_field_for_vector_input() { + let mut dae_model = dae::Dae::default(); + let span = lower_test_span(); + + let mut cp = rumoca_core::Function::new("My.cp", span); + cp.inputs.push(function_param("state_p")); + cp.inputs.push(function_param("state_T")); + cp.inputs.push( + function_param_with_dims("state_X", &[0]).with_shape_expr(vec![ + rumoca_core::Subscript::generated_expr(Box::new(var("nX")), span), + ]), + ); + cp.outputs.push(function_param("y")); + cp.body.push(rumoca_core::Statement::Assignment { + comp: component_ref("y"), + value: var_index("state_X", 1), + span, + }); + dae_model.symbols.functions.insert(cp.name.clone(), cp); + dae_model.variables.algebraics.insert( + rumoca_core::VarName::new("states.p"), + source_array_var("states.p", &[2]), + ); + dae_model.variables.algebraics.insert( + rumoca_core::VarName::new("states.T"), + source_array_var("states.T", &[2]), + ); + dae_model.variables.algebraics.insert( + rumoca_core::VarName::new("states.X"), + source_array_var("states.X", &[2, 2]), + ); + + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( + "My.cp", + )), + args: vec![ + source_var("states.p"), + source_var("states.T"), + source_var("states.X"), + ], + is_constructor: false, + span, + }; + let layout = build_var_layout(&dae_model).expect("test DAE layout should build"); + let err = lower_expression(&expr, &layout, &dae_model.symbols.functions) + .expect_err("non-singleton vectorized array field must not bind to vector input"); + + assert!( + err.to_string() + .contains("input `state_X` expected rank 1 for declared shape [0], got rank 2"), + "unexpected error: {err}" + ); +} + #[test] fn lower_expression_rejects_unknown_record_constructor_input_field_with_span() { let span = rumoca_core::Span::from_offsets( @@ -1909,3 +2884,37 @@ fn lower_expression_rejects_unknown_record_constructor_input_field_with_span() { "unexpected error: {err}" ); } + +#[test] +fn lower_expression_projects_record_constructor_output_field_with_default() { + let span = lower_test_span(); + let mut constructor = rumoca_core::Function::new("Pkg.RecordCtor", span); + constructor.is_constructor = true; + constructor + .inputs + .push(rumoca_core::FunctionParam::new("re", "Real", span)); + constructor + .inputs + .push(rumoca_core::FunctionParam::new("im", "Real", span).with_default(real_lit(0.0))); + constructor.outputs.push( + rumoca_core::FunctionParam::new("result", "Pkg.RecordValue", span) + .with_type_class(rumoca_core::ClassType::Record), + ); + + let mut functions = IndexMap::new(); + functions.insert(constructor.name.clone(), constructor); + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( + "Pkg.RecordCtor.result.im", + )), + args: vec![real_lit(2.0)], + is_constructor: false, + span, + }; + + let lowered = lower_expression(&expr, &VarLayout::default(), &functions) + .expect("record constructor output fields should project from bound inputs"); + let (regs, _) = eval_linear_ops(&lowered.ops, &[], &[], 0.0); + + assert_eq!(read_reg(®s, lowered.result), 0.0); +} diff --git a/crates/rumoca-phase-solve/src/lower/tests/function_expression_tests/statement_and_projection_tests.rs b/crates/rumoca-phase-solve/src/lower/tests/function_expression_tests/statement_and_projection_tests.rs index 90f72d216..19b92f34f 100644 --- a/crates/rumoca-phase-solve/src/lower/tests/function_expression_tests/statement_and_projection_tests.rs +++ b/crates/rumoca-phase-solve/src/lower/tests/function_expression_tests/statement_and_projection_tests.rs @@ -766,8 +766,10 @@ fn lower_expression_inlines_user_function_for_statement() { assert!((read_reg(®s, lowered.result) - 6.0).abs() <= 1e-12); } -#[test] -fn lower_expression_inlines_function_row_slice_from_scoped_matrix_input() { +fn lower_scoped_matrix_slice( + selector: rumoca_core::Expression, + selector_first: bool, +) -> Result { let mut dae_model = dae::Dae::default(); let mut pick_row = rumoca_core::Function::new("My.pickRow", lower_test_span()); pick_row.inputs = vec![ @@ -776,27 +778,42 @@ fn lower_expression_inlines_function_row_slice_from_scoped_matrix_input() { ]; pick_row.outputs = vec![function_param("out")]; pick_row.locals = vec![function_param_with_dims("e3_1", &[3])]; + let colon = rumoca_core::Subscript::Colon { + span: lower_test_span(), + }; + let selector = rumoca_core::Subscript::generated_expr(Box::new(selector), lower_test_span()); + let subscripts = if selector_first { + vec![selector, colon] + } else { + vec![colon, selector] + }; pick_row.body = vec![ rumoca_core::Statement::Assignment { comp: component_ref("e3_1"), value: rumoca_core::Expression::Index { base: Box::new(var("R_T")), - subscripts: vec![ - rumoca_core::Subscript::generated_expr( - Box::new(var("sequence[3]")), - lower_test_span(), - ), - rumoca_core::Subscript::Colon { - span: lower_test_span(), - }, - ], + subscripts, span: lower_test_span(), }, span: lower_test_span(), }, rumoca_core::Statement::Assignment { comp: component_ref("out"), - value: var_index("e3_1", 2), + value: add( + var_index("e3_1", 1), + add( + binary( + rumoca_core::OpBinary::Mul, + real_lit(10.0), + var_index("e3_1", 2), + ), + binary( + rumoca_core::OpBinary::Mul, + real_lit(100.0), + var_index("e3_1", 3), + ), + ), + ), span: lower_test_span(), }, ]; @@ -811,7 +828,8 @@ fn lower_expression_inlines_function_row_slice_from_scoped_matrix_input() { )), args: vec![ rumoca_core::Expression::Array { - elements: (1..=9) + elements: [2, 3, 5, 7, 11, 13, 17, 19, 23] + .into_iter() .map(|value| rumoca_core::Expression::Literal { value: rumoca_core::Literal::Integer(value), span: lower_test_span(), @@ -836,11 +854,224 @@ fn lower_expression_inlines_function_row_slice_from_scoped_matrix_input() { span: lower_test_span(), }; - let lowered = lower_expression(&expr, &VarLayout::default(), &dae_model.symbols.functions) - .expect("function-local matrix row slice should lower"); + let lowered = lower_expression(&expr, &VarLayout::default(), &dae_model.symbols.functions)?; let (regs, _) = eval_linear_ops(&lowered.ops, &[], &[], 0.0); + Ok(read_reg(®s, lowered.result)) +} + +#[test] +fn lower_expression_inlines_function_row_slice_from_scoped_matrix_input() { + let value = lower_scoped_matrix_slice(var("sequence[3]"), true) + .expect("function-local matrix row slice should lower"); + assert_eq!(value, 1417.0); +} + +#[test] +fn lower_expression_inlines_local_matrix_slice_with_nested_index_selector() { + let selector = rumoca_core::Expression::Index { + base: Box::new(var("sequence")), + subscripts: vec![rumoca_core::Subscript::index(3, lower_test_span())], + span: lower_test_span(), + }; + let value = lower_scoped_matrix_slice(selector, true) + .expect("function-local matrix row slice with nested selector should lower"); + assert_eq!(value, 1417.0); +} + +#[test] +fn lower_expression_inlines_local_matrix_column_slice_with_nested_index_selector() { + let selector = rumoca_core::Expression::Index { + base: Box::new(var("sequence")), + subscripts: vec![rumoca_core::Subscript::index(3, lower_test_span())], + span: lower_test_span(), + }; + let value = lower_scoped_matrix_slice(selector, false) + .expect("function-local matrix column slice with nested selector should lower"); + assert_eq!(value, 2013.0); +} + +#[test] +fn lower_expression_keeps_static_local_matrix_slice() { + let selector = rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(2), + span: lower_test_span(), + }; + let value = lower_scoped_matrix_slice(selector, true) + .expect("static function-local matrix row slice should lower"); + assert_eq!(value, 1417.0); +} + +fn lower_rank_three_slice_with_companion( + companion: rumoca_core::Subscript, +) -> Result<(), crate::LowerError> { + let mut dae_model = dae::Dae::default(); + let mut pick_slice = rumoca_core::Function::new("My.pickRankThreeSlice", lower_test_span()); + pick_slice.inputs = vec![ + function_param_with_dims("a", &[2, 3, 2]), + function_param_with_dims("sequence", &[3]), + ]; + pick_slice.outputs = vec![function_param("out")]; + pick_slice.locals = vec![function_param_with_dims("selected", &[2])]; + pick_slice.body = vec![ + rumoca_core::Statement::Assignment { + comp: component_ref("selected"), + value: rumoca_core::Expression::Index { + base: Box::new(var("a")), + subscripts: vec![ + rumoca_core::Subscript::generated_expr( + Box::new(rumoca_core::Expression::Index { + base: Box::new(var("sequence")), + subscripts: vec![rumoca_core::Subscript::index(3, lower_test_span())], + span: lower_test_span(), + }), + lower_test_span(), + ), + companion, + rumoca_core::Subscript::Colon { + span: lower_test_span(), + }, + ], + span: lower_test_span(), + }, + span: lower_test_span(), + }, + rumoca_core::Statement::Assignment { + comp: component_ref("out"), + value: var_index("selected", 1), + span: lower_test_span(), + }, + ]; + dae_model.symbols.functions.insert( + rumoca_core::VarName::new("My.pickRankThreeSlice"), + pick_slice, + ); + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::from_component_reference(test_component_ref_from_name( + "My.pickRankThreeSlice", + )), + args: vec![ + rumoca_core::Expression::Array { + elements: (1..=12) + .map(|value| rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + span: lower_test_span(), + }) + .collect(), + is_matrix: false, + span: lower_test_span(), + }, + rumoca_core::Expression::Array { + elements: [1, 2, 1] + .into_iter() + .map(|value| rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + span: lower_test_span(), + }) + .collect(), + is_matrix: false, + span: lower_test_span(), + }, + ], + is_constructor: false, + span: lower_test_span(), + }; - assert_eq!(read_reg(®s, lowered.result), 5.0); + lower_expression(&expr, &VarLayout::default(), &dae_model.symbols.functions).map(|_| ()) +} + +fn assert_rank_three_slice_unsupported(err: crate::LowerError, expected_reason: &str) { + let crate::LowerError::UnsupportedAt { + reason, + contexts, + span, + } = err + else { + panic!("expected source-spanned out-of-bounds selector error"); + }; + assert_eq!(reason, expected_reason); + assert_eq!( + contexts, + ["resolving subscripted local array `a` with shape [2, 3, 2]"] + ); + assert!(!span.is_dummy()); +} + +#[test] +fn lower_expression_rejects_out_of_bounds_companion_to_nested_slice_selector() { + let companion = rumoca_core::Subscript::generated_expr( + Box::new(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(4), + span: lower_test_span(), + }), + lower_test_span(), + ); + let err = lower_rank_three_slice_with_companion(companion) + .expect_err("out-of-bounds companion selector must fail before dynamic selection"); + assert_rank_three_slice_unsupported(err, "array slice index is outside dimension bounds"); +} + +#[test] +fn lower_expression_rejects_unsupported_companion_to_nested_slice_selector() { + let companion = rumoca_core::Subscript::generated_expr( + Box::new(rumoca_core::Expression::Tuple { + elements: vec![real_lit(2.0)], + span: lower_test_span(), + }), + lower_test_span(), + ); + let err = lower_rank_three_slice_with_companion(companion) + .expect_err("unsupported companion selector must fail before dynamic selection"); + assert_rank_three_slice_unsupported(err, "unsupported expression in for-loop range"); +} + +#[test] +fn lower_expression_preserves_local_context_for_dynamic_slice_error() { + let selector = rumoca_core::Expression::Index { + base: Box::new(var("missing_sequence")), + subscripts: vec![rumoca_core::Subscript::index(1, lower_test_span())], + span: lower_test_span(), + }; + let err = lower_scoped_matrix_slice(selector, true) + .expect_err("missing dynamic selector binding should fail with local context"); + let crate::LowerError::Spanned { source, span } = &err else { + panic!("expected source-spanned dynamic selector error, got {err:?}"); + }; + assert!(!span.is_dummy()); + let crate::LowerError::WithContext { source, contexts } = source.as_ref() else { + panic!("expected local context on dynamic selector error, got {err:?}"); + }; + assert!(matches!( + source.as_ref(), + crate::LowerError::MissingBinding { name } if name == "missing_sequence[1]" + )); + assert_eq!( + contexts.as_slice(), + ["resolving subscripted local array `R_T` with shape [3, 3]"] + ); +} + +#[test] +fn lower_expression_rejects_unsupported_local_matrix_slice_selector() { + let selector = rumoca_core::Expression::Tuple { + elements: vec![real_lit(2.0)], + span: lower_test_span(), + }; + let err = lower_scoped_matrix_slice(selector, true) + .expect_err("tuple slice selector should remain unsupported"); + let crate::LowerError::UnsupportedAt { + reason, + contexts, + span, + } = err + else { + panic!("expected source-spanned unsupported selector error"); + }; + assert_eq!(reason, "unsupported expression in for-loop range"); + assert_eq!( + contexts, + ["resolving subscripted local array `R_T` with shape [3, 3]"] + ); + assert!(!span.is_dummy()); } #[test] diff --git a/crates/rumoca-phase-solve/src/lower/tests/function_loop_tests.rs b/crates/rumoca-phase-solve/src/lower/tests/function_loop_tests.rs index 183a42519..14e3002a6 100644 --- a/crates/rumoca-phase-solve/src/lower/tests/function_loop_tests.rs +++ b/crates/rumoca-phase-solve/src/lower/tests/function_loop_tests.rs @@ -105,6 +105,70 @@ fn lower_expression_unrolls_function_for_loop_over_input_size() { assert!((read_reg(®s, lowered.result) - 6.0).abs() <= 1e-12); } +#[test] +fn lower_expression_binds_scalar_actual_as_singleton_dynamic_vector_input() { + let mut dae_model = dae::Dae::default(); + let mut evaluate = rumoca_core::Function::new("My.evaluate", lower_test_span()); + evaluate.inputs.push(function_param_with_dims("p", &[0])); + evaluate.inputs.push(function_param("u")); + evaluate.outputs.push(function_param("y")); + evaluate.body.push(rumoca_core::Statement::Assignment { + comp: component_ref("y"), + value: rumoca_core::Expression::Index { + base: Box::new(var("p")), + subscripts: vec![rumoca_core::Subscript::generated_index( + 1, + lower_test_span(), + )], + span: lower_test_span(), + }, + span: lower_test_span(), + }); + evaluate.body.push(rumoca_core::Statement::For { + indices: vec![rumoca_core::ForIndex { + ident: "j".to_string(), + range: rumoca_core::Expression::Range { + start: Box::new(int_lit(2)), + step: None, + end: Box::new(size_expr(var("p"), 1)), + span: lower_test_span(), + }, + }], + equations: vec![rumoca_core::Statement::Assignment { + comp: component_ref("y"), + value: add( + rumoca_core::Expression::Index { + base: Box::new(var("p")), + subscripts: vec![rumoca_core::Subscript::generated_expr( + Box::new(var("j")), + lower_test_span(), + )], + span: lower_test_span(), + }, + mul(var("u"), var("y")), + ), + span: lower_test_span(), + }], + span: lower_test_span(), + }); + dae_model + .symbols + .functions + .insert(rumoca_core::VarName::new("My.evaluate"), evaluate); + + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::VarName::new("My.evaluate").into(), + args: vec![real_lit(5.0), real_lit(2.0)], + is_constructor: false, + span: lower_test_span(), + }; + let lowered = lower_expression(&expr, &VarLayout::default(), &dae_model.symbols.functions) + .expect("scalar actual should bind as singleton p[:] vector"); + + let (regs, _) = eval_linear_ops(&lowered.ops, &[], &[], 0.0); + assert!((read_reg(®s, lowered.result) - 5.0).abs() <= 1e-12); +} + #[test] fn lower_expression_unrolls_function_for_loop_over_local_input_size() { let mut dae_model = dae::Dae::default(); diff --git a/crates/rumoca-phase-solve/src/lower/tests/intrinsics.rs b/crates/rumoca-phase-solve/src/lower/tests/intrinsics.rs index 84ee93090..abd628554 100644 --- a/crates/rumoca-phase-solve/src/lower/tests/intrinsics.rs +++ b/crates/rumoca-phase-solve/src/lower/tests/intrinsics.rs @@ -245,6 +245,50 @@ fn lower_discrete_rhs_holds_clocked_sample_between_clock_ticks() { ); } +#[test] +fn lower_discrete_rhs_holds_implicit_clock_sample_time_between_target_ticks() { + let mut dae_model = dae::Dae::default(); + dae_model + .variables + .discrete_reals + .insert(rumoca_core::VarName::new("simTime"), scalar_var("simTime")); + dae_model.clocks.timings.insert( + "simTime".to_string(), + dae::ClockSchedule { + period_seconds: 0.02, + phase_seconds: 0.0, + source_span: test_span(), + }, + ); + dae_model.discrete.real_updates.push(dae::Equation { + lhs: Some(rumoca_core::VarName::new("simTime").into()), + rhs: rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Sample, + args: vec![rumoca_core::Expression::VarRef { + name: rumoca_core::VarName::new("time").into(), + subscripts: vec![], + span: test_span(), + }], + span: test_span(), + }, + span: test_span(), + origin: "simTime = sample(time)".to_string(), + scalar_count: 1, + }); + + let layout = build_var_layout(&dae_model).expect("test DAE layout should build"); + let rows = lower_discrete_rhs(&dae_model, &layout) + .expect("implicit-clock sample(time) should lower with target hold"); + let mut p = vec![0.0; layout.p_scalars()]; + set_p_value(&layout, &mut p, "simTime", 0.0); + + let (_, held) = eval_linear_ops(&rows[0], &[], &p, 0.01); + assert_eq!(held.expect("row output"), 0.0); + + let (_, sampled) = eval_linear_ops(&rows[0], &[], &p, 0.02); + assert_eq!(sampled.expect("row output"), 0.02); +} + #[test] fn lower_discrete_rhs_holds_vector_clocked_sample_elements_between_ticks() { let mut dae_model = dae::Dae::default(); diff --git a/crates/rumoca-phase-solve/src/observation_refresh.rs b/crates/rumoca-phase-solve/src/observation_refresh.rs index 805b85038..be49687c2 100644 --- a/crates/rumoca-phase-solve/src/observation_refresh.rs +++ b/crates/rumoca-phase-solve/src/observation_refresh.rs @@ -134,7 +134,43 @@ fn expression_safe_for_observation_refresh( ) -> bool { !(expression_contains_clocked_value_operator(dae_model, expr) || target_has_clock_metadata(dae_model, lhs.as_str()) - && expression_contains_lowered_pre_ref(expr)) + && expression_contains_lowered_pre_ref(expr) + && !expression_is_direct_lowered_change_relation(expr)) +} + +fn expression_is_direct_lowered_change_relation(expr: &rumoca_core::Expression) -> bool { + let rumoca_core::Expression::Binary { op, lhs, rhs, .. } = expr else { + return false; + }; + if !matches!(op, rumoca_core::OpBinary::Eq | rumoca_core::OpBinary::Neq) { + return false; + } + current_and_lowered_pre_ref_match(lhs, rhs) || current_and_lowered_pre_ref_match(rhs, lhs) +} + +fn current_and_lowered_pre_ref_match( + current: &rumoca_core::Expression, + previous: &rumoca_core::Expression, +) -> bool { + let rumoca_core::Expression::VarRef { + name: current_name, + subscripts: current_subscripts, + .. + } = current + else { + return false; + }; + let rumoca_core::Expression::VarRef { + name: previous_name, + subscripts: previous_subscripts, + .. + } = previous + else { + return false; + }; + current_subscripts == previous_subscripts + && rumoca_core::pre_slot_base(previous_name.as_str()) + .is_some_and(|base| base == current_name.as_str()) } fn expression_has_observation_pulse(dae_model: &dae::Dae, expr: &rumoca_core::Expression) -> bool { @@ -230,6 +266,8 @@ impl ExpressionVisitor for EventRelationChecker { | rumoca_core::OpBinary::Le | rumoca_core::OpBinary::Gt | rumoca_core::OpBinary::Ge + | rumoca_core::OpBinary::Eq + | rumoca_core::OpBinary::Neq ) { self.found = true; diff --git a/crates/rumoca-phase-solve/src/projection_plan.rs b/crates/rumoca-phase-solve/src/projection_plan.rs new file mode 100644 index 000000000..d65b22c55 --- /dev/null +++ b/crates/rumoca-phase-solve/src/projection_plan.rs @@ -0,0 +1,818 @@ +use super::*; + +pub(super) fn lower_algebraic_projection_plan( + rows: &[Vec], + row_targets: &[Option], + state_scalar_count: usize, + solver_scalar_count: usize, + context_span: rumoca_core::Span, +) -> Result { + let projection_count = solver_scalar_count + .checked_sub(state_scalar_count) + .ok_or_else(|| { + lower_contract_violation( + "algebraic projection range starts after solver scalar count".to_string(), + context_span, + ) + })?; + let mut projection_indices = lower_vec_with_capacity( + projection_count, + "algebraic projection index count", + context_span, + )?; + projection_indices.extend(state_scalar_count..solver_scalar_count); + lower_projection_plan( + rows, + row_targets, + &projection_indices, + state_scalar_count..solver_scalar_count, + ProjectionPlanPolicy { + include_explicit_row_targets: true, + require_complete_algebraic_coverage: true, + }, + None, + context_span, + ) +} + +#[derive(Clone, Copy)] +pub(super) struct ProjectionPlanPolicy { + pub(super) include_explicit_row_targets: bool, + pub(super) require_complete_algebraic_coverage: bool, +} + +pub(super) fn lower_projection_plan( + rows: &[Vec], + row_targets: &[Option], + projection_indices: &[usize], + row_indices: std::ops::Range, + policy: ProjectionPlanPolicy, + implicit_incidence_rows: Option<&BTreeSet>, + context_span: rumoca_core::Span, +) -> Result { + let mut row_to_vars = BTreeMap::>::new(); + let projection_set = projection_indices.iter().copied().collect::>(); + let row_indices = row_indices.collect::>(); + let identity_projection_rows = row_indices + .iter() + .filter_map(|row_idx| { + identity_projection_y_index(rows.get(*row_idx)?.as_slice(), &projection_set) + .map(|y_idx| (y_idx, *row_idx)) + }) + .collect::>(); + + for &row_idx in &row_indices { + let collect_implicit_incidence = + implicit_incidence_rows.is_none_or(|rows| rows.contains(&row_idx)); + let row_target = row_targets.get(row_idx).copied().flatten(); + let Some(y_indices) = projection_row_y_indices_for_plan( + rows[row_idx].as_slice(), + row_target, + row_idx, + collect_implicit_incidence, + &projection_set, + policy.include_explicit_row_targets, + &identity_projection_rows, + ) else { + continue; + }; + row_to_vars.insert(row_idx, y_indices); + } + + let projection_incidence = algebraic_projection_incidence(&row_to_vars, context_span)?; + let (blocks, dropped_equations) = projection_blt_blocks(&projection_incidence)?; + let mut blocks = + lower_blt_projection_blocks(&blocks, row_targets, &projection_incidence, context_span)?; + if policy.require_complete_algebraic_coverage { + blocks = retain_dropped_projection_rows( + blocks, + &dropped_equations, + &projection_incidence, + context_span, + )?; + validate_complete_algebraic_projection_plan( + &blocks, + rows, + &row_indices, + &projection_incidence, + context_span, + )?; + } + Ok(solve::AlgebraicProjectionPlan { blocks }) +} + +pub(super) fn projection_row_y_indices_for_plan( + row: &[solve::LinearOp], + row_target: Option, + row_idx: usize, + collect_implicit_incidence: bool, + projection_set: &BTreeSet, + include_explicit_row_targets: bool, + identity_projection_rows: &BTreeMap, +) -> Option> { + let mut y_indices = if collect_implicit_incidence { + collect_algebraic_y_indices_for_row(row, projection_set) + } else { + BTreeSet::new() + }; + let explicit_target = match row_target { + Some(solve::ScalarSlot::Y { index, .. }) if projection_set.contains(&index) => Some(index), + _ => None, + }; + if let Some(index) = explicit_target { + if !include_explicit_row_targets { + return Some(BTreeSet::from([index])); + } + y_indices.insert(index); + return Some(y_indices); + } + if let Some(index) = identity_projection_y_index(row, projection_set) { + return Some(BTreeSet::from([index])); + } + y_indices.retain(|index| { + projection_index_not_claimed_by_identity(identity_projection_rows, *index, row_idx) + }); + (!y_indices.is_empty()).then_some(y_indices) +} + +pub(super) fn projection_index_not_claimed_by_identity( + identity_projection_rows: &BTreeMap, + index: usize, + row_idx: usize, +) -> bool { + identity_projection_rows + .get(&index) + .is_none_or(|identity_row| *identity_row == row_idx) +} + +pub(super) fn identity_projection_y_index( + row: &[solve::LinearOp], + projection_set: &BTreeSet, +) -> Option { + let [ + solve::LinearOp::LoadY { + dst: load_dst, + index, + }, + solve::LinearOp::StoreOutput { src }, + ] = row + else { + return None; + }; + (*load_dst == *src && projection_set.contains(index)).then_some(*index) +} + +pub(super) fn projection_blt_blocks( + projection_incidence: &ProjectionIncidence, +) -> Result<(Vec, Vec), LowerError> { + if projection_incidence.incidence.n_eq == 0 && projection_incidence.incidence.n_var == 0 { + return Ok((Vec::new(), Vec::new())); + } + let regular = + rumoca_phase_structural::maximum_regular_subsystem(&projection_incidence.incidence) + .map_err(|err| LowerError::Unsupported { + reason: format!("lower algebraic projection BLT: {err}"), + })?; + let blocks = + rumoca_phase_structural::build_blt_from_incidence(®ular.incidence).map_err(|err| { + LowerError::Unsupported { + reason: format!("lower algebraic projection BLT: {err}"), + } + })?; + Ok((blocks, regular.dropped_equations)) +} + +pub(super) fn retain_dropped_projection_rows( + mut blocks: Vec, + dropped_equations: &[EquationRef], + projection_incidence: &ProjectionIncidence, + context_span: rumoca_core::Span, +) -> Result, LowerError> { + for equation in dropped_equations { + let row_y_indices = projection_row_y_indices(equation.0, projection_incidence); + let mut retained = lower_vec_with_capacity( + blocks.len(), + "retained algebraic projection block count", + context_span, + )?; + let mut insertion_index = None; + let mut merged = None; + for block in blocks { + if block + .y_indices + .iter() + .any(|index| row_y_indices.contains(index)) + { + insertion_index.get_or_insert(retained.len()); + merged = Some(match merged { + Some(previous) => combine_projection_blocks(previous, block, context_span)?, + None => block, + }); + } else { + retained.push(block); + } + } + if let (Some(mut block), Some(index)) = (merged, insertion_index) { + reserve_lower_capacity( + &mut block.rows, + 1, + "retained algebraic projection row count", + context_span, + )?; + block.rows.push(equation.0); + block.rows.sort_unstable(); + block.rows.dedup(); + retained.insert(index, block); + } + blocks = retained; + } + Ok(blocks) +} + +pub(super) fn projection_row_y_indices( + row: usize, + projection_incidence: &ProjectionIncidence, +) -> BTreeSet { + let Some(position) = projection_incidence + .incidence + .equation_refs + .iter() + .position(|equation| equation.0 == row) + else { + return BTreeSet::new(); + }; + projection_incidence.incidence.eq_unknowns[position] + .iter() + .filter_map(|unknown_index| { + projection_incidence + .incidence + .unknown_names + .get(*unknown_index) + .and_then(|unknown| projection_y_index(unknown, projection_incidence)) + }) + .collect() +} + +pub(super) fn validate_complete_algebraic_projection_plan( + blocks: &[solve::AlgebraicProjectionBlock], + rows: &[Vec], + expected_rows: &[usize], + projection_incidence: &ProjectionIncidence, + context_span: rumoca_core::Span, +) -> Result<(), LowerError> { + let covered_rows = blocks + .iter() + .flat_map(|block| block.rows.iter().copied()) + .collect::>(); + for row in expected_rows { + if covered_rows.contains(row) { + continue; + } + let has_projection_incidence = projection_incidence + .incidence + .equation_refs + .iter() + .any(|equation| equation.0 == *row); + if has_projection_incidence || statically_nonzero_projection_row(&rows[*row]).is_some() { + return Err(lower_contract_violation( + format!("algebraic projection plan omits implicit row {row}"), + context_span, + )); + } + } + + let covered_y_indices = blocks + .iter() + .flat_map(|block| block.y_indices.iter().copied()) + .collect::>(); + for y_index in &projection_incidence.unknown_y_indices { + if !covered_y_indices.contains(y_index) { + return Err(lower_contract_violation( + format!("algebraic projection plan omits residual target y[{y_index}]"), + context_span, + )); + } + } + Ok(()) +} + +pub(super) fn statically_nonzero_projection_row(row: &[solve::LinearOp]) -> Option { + if !row.iter().all(|op| { + matches!( + op, + solve::LinearOp::Const { .. } + | solve::LinearOp::Move { .. } + | solve::LinearOp::Unary { .. } + | solve::LinearOp::Binary { .. } + | solve::LinearOp::Compare { .. } + | solve::LinearOp::Select { .. } + | solve::LinearOp::StoreOutput { .. } + ) + }) { + return None; + } + let value = rumoca_eval_solve::eval_row_with_context( + row, + &[], + &[], + 0.0, + rumoca_eval_solve::RowEvalContext::default(), + ) + .ok()?; + (value.is_finite() && value != 0.0).then_some(value) +} + +pub(super) fn collect_algebraic_y_indices_for_row( + row: &[solve::LinearOp], + projection_set: &BTreeSet, +) -> BTreeSet { + let mut defs = BTreeMap::::new(); + let mut outputs = Vec::new(); + for op in row { + match row_def_use(op) { + RowDefUseOp::Def { dst, def_use } => { + defs.insert(dst, def_use); + } + RowDefUseOp::Store { src } => outputs.push(src), + } + } + let mut y_indices = BTreeSet::new(); + let mut visited = BTreeSet::new(); + let mut stack = outputs; + while let Some(reg) = stack.pop() { + if !visited.insert(reg) { + continue; + } + let Some(def_use) = defs.get(®) else { + continue; + }; + if let Some(index) = def_use.loaded_y + && projection_set.contains(&index) + { + y_indices.insert(index); + } + stack.extend(def_use.inputs.iter().copied()); + } + y_indices +} + +#[derive(Debug)] +pub(super) struct RowDefUse { + loaded_y: Option, + inputs: Vec, +} + +pub(super) enum RowDefUseOp { + Def { dst: solve::Reg, def_use: RowDefUse }, + Store { src: solve::Reg }, +} + +pub(super) fn row_def_use(op: &solve::LinearOp) -> RowDefUseOp { + use solve::LinearOp as Op; + match *op { + Op::Const { dst, .. } | Op::LoadTime { dst } | Op::LoadP { dst, .. } => { + def_use(dst, None, Vec::new()) + } + Op::LoadY { dst, index } => def_use(dst, Some(index), Vec::new()), + Op::LoadSeed { dst, .. } => def_use(dst, None, Vec::new()), + Op::LoadIndexedP { dst, index, .. } | Op::LoadIndexedSeed { dst, index, .. } => { + def_use(dst, None, vec![index]) + } + Op::Move { dst, src } | Op::Unary { dst, arg: src, .. } => def_use(dst, None, vec![src]), + Op::Binary { dst, lhs, rhs, .. } | Op::Compare { dst, lhs, rhs, .. } => { + def_use(dst, None, vec![lhs, rhs]) + } + Op::Select { + dst, + cond, + if_true, + if_false, + } => def_use(dst, None, vec![cond, if_true, if_false]), + Op::LinearSolveComponent { + dst, + matrix_start, + rhs_start, + n, + .. + } => def_use( + dst, + None, + reg_range(matrix_start, n * n) + .chain(reg_range(rhs_start, n)) + .collect(), + ), + Op::TableBounds { dst, table_id, .. } => def_use(dst, None, vec![table_id]), + Op::TableLookup { + dst, + table_id, + column, + input, + } + | Op::TableLookupSlope { + dst, + table_id, + column, + input, + } => def_use(dst, None, vec![table_id, column, input]), + Op::TableNextEvent { + dst, + table_id, + time, + } => def_use(dst, None, vec![table_id, time]), + Op::RandomInitialState { + dst, + local_seed, + global_seed, + .. + } => def_use(dst, None, vec![local_seed, global_seed]), + Op::RandomResult { + dst, + state_start, + state_len, + .. + } + | Op::RandomState { + dst, + state_start, + state_len, + .. + } => def_use(dst, None, reg_range(state_start, state_len).collect()), + Op::ImpureRandomInit { dst, seed } => def_use(dst, None, vec![seed]), + Op::ImpureRandom { dst, id, .. } => def_use(dst, None, vec![id]), + Op::ImpureRandomInteger { + dst, + id, + imin, + imax, + .. + } => def_use(dst, None, vec![id, imin, imax]), + Op::ExternalCall { + dst, + args, + arg_count, + .. + } => def_use(dst, None, args.into_iter().take(arg_count).collect()), + Op::StoreOutput { src } => RowDefUseOp::Store { src }, + } +} + +pub(super) fn def_use( + dst: solve::Reg, + loaded_y: Option, + inputs: Vec, +) -> RowDefUseOp { + RowDefUseOp::Def { + dst, + def_use: RowDefUse { loaded_y, inputs }, + } +} + +pub(super) fn reg_range(start: solve::Reg, len: usize) -> impl Iterator { + (0..len).filter_map(move |offset| start.checked_add(offset.try_into().ok()?)) +} + +pub(super) struct ProjectionIncidence { + pub(super) incidence: Incidence, + pub(super) unknown_y_indices: Vec, +} + +pub(super) fn algebraic_projection_incidence( + row_to_vars: &BTreeMap>, + context_span: rumoca_core::Span, +) -> Result { + let unknown_y_set = row_to_vars + .values() + .flat_map(|vars| vars.iter().copied()) + .collect::>(); + let mut unknown_y_indices = lower_vec_with_capacity( + unknown_y_set.len(), + "projection unknown index count", + context_span, + )?; + unknown_y_indices.extend(unknown_y_set); + + let mut unknown_names = lower_vec_with_capacity( + unknown_y_indices.len(), + "projection unknown name count", + context_span, + )?; + for y_idx in &unknown_y_indices { + unknown_names.push(projection_unknown_id(*y_idx)); + } + + let unknown_positions = unknown_y_indices + .iter() + .copied() + .enumerate() + .map(|(local_idx, y_idx)| (y_idx, local_idx)) + .collect::>(); + + let mut equation_refs = lower_vec_with_capacity( + row_to_vars.len(), + "projection equation ref count", + context_span, + )?; + let mut eq_unknowns = lower_vec_with_capacity( + row_to_vars.len(), + "projection equation unknown count", + context_span, + )?; + for (row_idx, vars) in row_to_vars { + equation_refs.push(EquationRef(*row_idx)); + let mut unknowns = + lower_hash_set_with_capacity(vars.len(), "projection row unknown count", context_span)?; + for y_idx in vars { + if let Some(local_idx) = unknown_positions.get(y_idx).copied() { + unknowns.insert(local_idx); + } + } + eq_unknowns.push(unknowns); + } + + Ok(ProjectionIncidence { + incidence: Incidence::new(eq_unknowns, equation_refs, unknown_names), + unknown_y_indices, + }) +} + +pub(super) fn projection_unknown_id(y_idx: usize) -> UnknownId { + UnknownId::SolverY(y_idx) +} + +pub(super) fn projection_y_index( + unknown: &UnknownId, + projection_incidence: &ProjectionIncidence, +) -> Option { + projection_incidence + .incidence + .unknown_names + .iter() + .position(|candidate| candidate == unknown) + .and_then(|idx| projection_incidence.unknown_y_indices.get(idx).copied()) +} + +pub(super) fn lower_blt_projection_blocks( + blocks: &[BltBlock], + row_targets: &[Option], + projection_incidence: &ProjectionIncidence, + context_span: rumoca_core::Span, +) -> Result, LowerError> { + let mut lowered = lower_vec_with_capacity( + blocks.len(), + "algebraic projection block count", + context_span, + )?; + for block in blocks { + let block = match block { + BltBlock::Scalar { equation, unknown } => { + projection_y_index(unknown, projection_incidence) + .map(|y_index| { + scalar_projection_block( + equation.0, + y_index, + row_targets, + projection_incidence, + context_span, + ) + }) + .transpose()? + } + BltBlock::AlgebraicLoop { + equations, + unknowns, + } => lower_algebraic_loop_projection_block( + equations, + unknowns, + row_targets, + projection_incidence, + context_span, + )?, + }; + if let Some(block) = block { + lowered.push(block); + } + } + merge_overlapping_projection_blocks(lowered, context_span) +} + +pub(super) fn merge_overlapping_projection_blocks( + blocks: Vec, + context_span: rumoca_core::Span, +) -> Result, LowerError> { + let mut merged = lower_vec_with_capacity( + blocks.len(), + "merged algebraic projection block count", + context_span, + )?; + for block in blocks { + merge_projection_block(&mut merged, block, context_span)?; + } + Ok(merged) +} + +pub(super) fn merge_projection_block( + merged: &mut Vec, + mut block: solve::AlgebraicProjectionBlock, + context_span: rumoca_core::Span, +) -> Result<(), LowerError> { + let mut idx = 0; + while idx < merged.len() { + if projection_blocks_overlap(&merged[idx], &block) { + let previous = merged.remove(idx); + block = combine_projection_blocks(previous, block, context_span)?; + idx = 0; + } else { + idx += 1; + } + } + merged.push(block); + Ok(()) +} + +pub(super) fn projection_blocks_overlap( + lhs: &solve::AlgebraicProjectionBlock, + rhs: &solve::AlgebraicProjectionBlock, +) -> bool { + lhs.y_indices + .iter() + .any(|index| rhs.y_indices.binary_search(index).is_ok()) +} + +pub(super) fn combine_projection_blocks( + lhs: solve::AlgebraicProjectionBlock, + rhs: solve::AlgebraicProjectionBlock, + context_span: rumoca_core::Span, +) -> Result { + let causal_step_count = lhs + .causal_steps + .len() + .checked_add(rhs.causal_steps.len()) + .ok_or_else(|| { + lower_contract_violation( + "merged algebraic projection causal-step count overflows host index range" + .to_string(), + context_span, + ) + })?; + let mut causal_steps = lower_vec_with_capacity( + causal_step_count, + "merged algebraic projection causal-step count", + context_span, + )?; + causal_steps.extend(lhs.causal_steps); + causal_steps.extend(rhs.causal_steps); + Ok(solve::AlgebraicProjectionBlock { + rows: merge_unique( + lhs.rows, + rhs.rows, + "merged algebraic projection row count", + context_span, + )?, + y_indices: merge_unique( + lhs.y_indices, + rhs.y_indices, + "merged algebraic projection target count", + context_span, + )?, + causal_steps, + }) +} + +pub(super) fn merge_unique( + lhs: Vec, + rhs: Vec, + context: &'static str, + context_span: rumoca_core::Span, +) -> Result, LowerError> { + let capacity = lhs.len().checked_add(rhs.len()).ok_or_else(|| { + lower_contract_violation( + format!("{context} overflows host index range"), + context_span, + ) + })?; + let mut merged = lower_vec_with_capacity(capacity, context, context_span)?; + merged.extend(lhs); + merged.extend(rhs); + merged.sort_unstable(); + merged.dedup(); + Ok(merged) +} + +pub(super) fn scalar_projection_block( + row: usize, + y_index: usize, + row_targets: &[Option], + projection_incidence: &ProjectionIncidence, + context_span: rumoca_core::Span, +) -> Result { + let mut rows = lower_vec_with_capacity( + 1, + "scalar algebraic projection block row count", + context_span, + )?; + rows.push(row); + let mut target_set = BTreeSet::from([y_index]); + if let Some(solve::ScalarSlot::Y { index, .. }) = row_targets.get(row).copied().flatten() + && projection_incidence.unknown_y_indices.contains(&index) + { + target_set.insert(index); + } + let mut y_indices = lower_vec_with_capacity( + target_set.len(), + "scalar algebraic projection block target count", + context_span, + )?; + y_indices.extend(target_set); + let causal_target = row_targets + .get(row) + .copied() + .flatten() + .and_then(|target| match target { + solve::ScalarSlot::Y { index, .. } => Some(index), + _ => None, + }) + .filter(|target| y_indices.contains(target)); + let causal_steps = if let Some(target) = causal_target { + vec![solve::AlgebraicProjectionStep { + row, + y_index: target, + }] + } else { + Vec::new() + }; + Ok(solve::AlgebraicProjectionBlock { + rows, + y_indices, + causal_steps, + }) +} + +pub(super) fn sorted_set_values( + values: BTreeSet, + context: &'static str, + context_span: rumoca_core::Span, +) -> Result, LowerError> { + let mut out = lower_vec_with_capacity(values.len(), context, context_span)?; + out.extend(values); + Ok(out) +} + +pub(super) fn collect_equation_rows( + equations: &[EquationRef], + context_span: rumoca_core::Span, +) -> Result, LowerError> { + let mut rows = lower_vec_with_capacity( + equations.len(), + "algebraic loop projection row count", + context_span, + )?; + for equation in equations { + rows.push(equation.0); + } + Ok(rows) +} + +pub(super) fn lower_algebraic_loop_projection_block( + equations: &[EquationRef], + unknowns: &[UnknownId], + row_targets: &[Option], + projection_incidence: &ProjectionIncidence, + context_span: rumoca_core::Span, +) -> Result, LowerError> { + let rows = collect_equation_rows(equations, context_span)?; + let y_indices = sorted_set_values( + loop_projection_target_set(unknowns, row_targets, &rows, projection_incidence), + "algebraic loop projection target count", + context_span, + )?; + if rows.is_empty() || y_indices.is_empty() { + return Ok(None); + } + Ok(Some(solve::AlgebraicProjectionBlock { + rows, + y_indices, + causal_steps: Vec::new(), + })) +} + +pub(super) fn loop_projection_target_set( + unknowns: &[UnknownId], + row_targets: &[Option], + rows: &[usize], + projection_incidence: &ProjectionIncidence, +) -> BTreeSet { + let mut y_indices = BTreeSet::new(); + for unknown in unknowns { + if let Some(index) = projection_y_index(unknown, projection_incidence) { + y_indices.insert(index); + } + } + for row in rows { + if let Some(solve::ScalarSlot::Y { index, .. }) = row_targets.get(*row).copied().flatten() + && projection_incidence.unknown_y_indices.contains(&index) + { + y_indices.insert(index); + } + } + y_indices +} diff --git a/crates/rumoca-phase-solve/src/residual_compute_block.rs b/crates/rumoca-phase-solve/src/residual_compute_block.rs index b92843b00..afe99ac76 100644 --- a/crates/rumoca-phase-solve/src/residual_compute_block.rs +++ b/crates/rumoca-phase-solve/src/residual_compute_block.rs @@ -4,15 +4,73 @@ use rumoca_ir_solve as solve; use crate::{lower, lower::LowerError, stencil}; +struct ResidualStructuredPartition<'a> { + equations: &'a [dae::StructuredEquationFamily], + source_equations: &'a [dae::Equation], + equation_index_offset: usize, +} + pub(crate) fn build_residual_compute_block( dae_model: &dae::Dae, layout: &solve::VarLayout, residual_rows: &[Vec], residual_targets: &[Option], residual_equations: &[(usize, &dae::Equation)], +) -> Result { + let partition = ResidualStructuredPartition { + equations: &dae_model.continuous.structured_equations, + source_equations: &dae_model.continuous.equations, + equation_index_offset: 0, + }; + build_residual_compute_block_with_structured_partition( + dae_model, + layout, + residual_rows, + residual_targets, + residual_equations, + &partition, + ) +} + +/// Initialization residuals share the generic ComputeBlock evaluator with the +/// continuous system, but their structured-family provenance lives in the DAE +/// initialization partition. Keeping that provenance here avoids scalarizing +/// an otherwise compact initial `for` family before the lean GPU settle path +/// can evaluate it. +pub(crate) fn build_initialization_residual_compute_block( + dae_model: &dae::Dae, + layout: &solve::VarLayout, + residual_rows: &[Vec], + residual_targets: &[Option], + residual_equations: &[(usize, &dae::Equation)], +) -> Result { + let continuous_len = dae_model.continuous.equations.len(); + let partition = ResidualStructuredPartition { + equations: &dae_model.initialization.structured_equations, + source_equations: &dae_model.initialization.equations, + equation_index_offset: continuous_len, + }; + build_residual_compute_block_with_structured_partition( + dae_model, + layout, + residual_rows, + residual_targets, + residual_equations, + &partition, + ) +} + +fn build_residual_compute_block_with_structured_partition( + dae_model: &dae::Dae, + layout: &solve::VarLayout, + residual_rows: &[Vec], + residual_targets: &[Option], + residual_equations: &[(usize, &dae::Equation)], + partition: &ResidualStructuredPartition<'_>, ) -> Result { let span = residual_context_span(dae_model, residual_equations); validate_residual_compute_block_contract( + dae_model, residual_rows.len(), residual_targets.len(), residual_equations, @@ -27,7 +85,8 @@ pub(crate) fn build_residual_compute_block( let structural_bindings = lower::structural_bindings_for_structured_access(dae_model)?; let mut residual_index = 0usize; for (equation_index, equation) in residual_equations { - let scalar_count = equation.scalar_count.max(1); + let scalar_count = + lower::residual_equation_effective_row_count(dae_model, equation)?.max(1); for row_offset in 0..scalar_count { let Some(ops) = residual_rows.get(residual_index).cloned() else { break; @@ -50,7 +109,7 @@ pub(crate) fn build_residual_compute_block( residual_index, equation.span, )?, - dae_equation_index: Some(*equation_index), + dae_equation_index: equation_index.checked_sub(partition.equation_index_offset), access_proof: residual_row_access_proof( layout, &structural_bindings, @@ -67,13 +126,14 @@ pub(crate) fn build_residual_compute_block( stencil::push_structured_programs( &mut block.nodes, &mut rows, - &dae_model.continuous.structured_equations, - &dae_model.continuous.equations, + partition.equations, + partition.source_equations, )?; Ok(block) } fn validate_residual_compute_block_contract( + dae_model: &dae::Dae, residual_row_count: usize, residual_target_count: usize, residual_equations: &[(usize, &dae::Equation)], @@ -84,17 +144,18 @@ fn validate_residual_compute_block_contract( format!( "residual target count {residual_target_count} does not match residual row count {residual_row_count}" ), - residual_contract_error_span(residual_target_count, residual_equations) + residual_contract_error_span(dae_model, residual_target_count, residual_equations) .or(fallback_span), )); } - let expected_rows = residual_equation_scalar_count(residual_equations)?; + let expected_rows = residual_equation_scalar_count(dae_model, residual_equations)?; if expected_rows != residual_row_count { return Err(residual_contract_error( format!( "residual equation scalar count {expected_rows} does not match residual row count {residual_row_count}" ), - residual_contract_error_span(residual_row_count, residual_equations).or(fallback_span), + residual_contract_error_span(dae_model, residual_row_count, residual_equations) + .or(fallback_span), )); } Ok(()) @@ -108,29 +169,34 @@ fn residual_contract_error(reason: String, span: Option) -> L } fn residual_equation_scalar_count( + dae_model: &dae::Dae, residual_equations: &[(usize, &dae::Equation)], ) -> Result { residual_equations .iter() .try_fold(0usize, |total, (_, equation)| { - total - .checked_add(equation.scalar_count.max(1)) - .ok_or_else(|| { - residual_contract_error( - "residual equation scalar count overflows usize".to_string(), - Some(equation.span), - ) - }) + let row_count = lower::residual_equation_effective_row_count(dae_model, equation)?; + total.checked_add(row_count.max(1)).ok_or_else(|| { + residual_contract_error( + "residual equation scalar count overflows usize".to_string(), + Some(equation.span), + ) + }) }) } fn residual_contract_error_span( + dae_model: &dae::Dae, row_count: usize, residual_equations: &[(usize, &dae::Equation)], ) -> Option { let mut row_start = 0usize; for (_, equation) in residual_equations { - let row_end = row_start.checked_add(equation.scalar_count.max(1))?; + let row_end = row_start.checked_add( + lower::residual_equation_effective_row_count(dae_model, equation) + .ok()? + .max(1), + )?; if row_count < row_end { return Some(equation.span); } @@ -274,10 +340,14 @@ fn collect_residual_var_ref_access_operands( owner_span: rumoca_core::Span, operands: &mut Vec, ) -> Result, LowerError> { - if subscripts - .iter() - .any(|subscript| matches!(subscript, rumoca_core::Subscript::Colon { .. })) - { + if subscripts.iter().any(|subscript| { + matches!(subscript, rumoca_core::Subscript::Colon { .. }) + || matches!( + subscript, + rumoca_core::Subscript::Expr { expr, .. } + if matches!(expr.as_ref(), rumoca_core::Expression::Range { .. }) + ) + }) { return Ok(None); } let indices = lower::compile_time_subscript_indices_for_structured_access( @@ -323,8 +393,15 @@ mod tests { 19, ); let equation = dae::Equation::residual_array(literal_zero(span), span, "eq", 2); - let err = validate_residual_compute_block_contract(2, 1, &[(0, &equation)], Some(span)) - .expect_err("target count mismatch should fail"); + let dae_model = dae::Dae::default(); + let err = validate_residual_compute_block_contract( + &dae_model, + 2, + 1, + &[(0, &equation)], + Some(span), + ) + .expect_err("target count mismatch should fail"); assert_eq!(err.source_span(), Some(span)); assert!( @@ -341,8 +418,15 @@ mod tests { 23, ); let equation = dae::Equation::residual_array(literal_zero(span), span, "eq", 2); - let err = validate_residual_compute_block_contract(1, 1, &[(0, &equation)], Some(span)) - .expect_err("row count mismatch should fail"); + let dae_model = dae::Dae::default(); + let err = validate_residual_compute_block_contract( + &dae_model, + 1, + 1, + &[(0, &equation)], + Some(span), + ) + .expect_err("row count mismatch should fail"); assert_eq!(err.source_span(), Some(span)); assert!( @@ -366,7 +450,9 @@ mod tests { let first = dae::Equation::residual_array(literal_zero(first_span), first_span, "eq1", 1); let second = dae::Equation::residual_array(literal_zero(second_span), second_span, "eq2", 1); + let dae_model = dae::Dae::default(); let err = validate_residual_compute_block_contract( + &dae_model, 1, 1, &[(0, &first), (1, &second)], @@ -379,7 +465,8 @@ mod tests { #[test] fn residual_compute_block_contract_does_not_fabricate_span_without_context() { - let err = validate_residual_compute_block_contract(1, 0, &[], None) + let dae_model = dae::Dae::default(); + let err = validate_residual_compute_block_contract(&dae_model, 1, 0, &[], None) .expect_err("unmatched residual rows without provenance should fail"); assert_eq!(err.source_span(), None); @@ -405,7 +492,8 @@ mod tests { 1, ); - let err = residual_equation_scalar_count(&[(0, &first), (1, &second)]) + let dae_model = dae::Dae::default(); + let err = residual_equation_scalar_count(&dae_model, &[(0, &first), (1, &second)]) .expect_err("oversized residual scalar count should fail"); assert_eq!(err.source_span(), None); diff --git a/crates/rumoca-phase-solve/src/solve_model.rs b/crates/rumoca-phase-solve/src/solve_model.rs index 2cb58e06f..c19b29090 100644 --- a/crates/rumoca-phase-solve/src/solve_model.rs +++ b/crates/rumoca-phase-solve/src/solve_model.rs @@ -15,13 +15,13 @@ use crate::initial_values::apply_initial_equations_to_start_values; use rumoca_eval_dae::build_runtime_parameter_tail_env_with_runtime; use rumoca_eval_dae::constant::eval_scalar_const_expr; use rumoca_eval_dae::eval::{ - EvalError, EvalRuntimeState, eval_matrix_values, eval_shaped_array_values, + EvalError, EvalRuntimeState, eval_shaped_array_values, external_table_data_for_parameter_values_in, }; use rumoca_eval_dae::{ build_partial_runtime_parameter_tail_env_with_declared_slots_and_runtime, build_runtime_parameter_tail_env_with_declared_slots_and_runtime, can_broadcast_start_value, - eval_array_values, eval_expr, start_expr_is_nonnumeric, + eval_expr, start_expr_is_nonnumeric, }; use std::collections::{HashMap, HashSet}; @@ -57,6 +57,14 @@ impl SolveModelLoweringProfile { fn needs_solve_artifacts(self) -> bool { self == Self::Runtime } + + /// GPU preparation receives an already-prepared direct DAE (or the output + /// of the structural funnel) and must preserve its state-only layout. + /// Runtime-only direct-state demotion is both redundant here and expands + /// regular state families one component at a time. + fn needs_structural_derivative_preparation(self) -> bool { + matches!(self, Self::Runtime | Self::RuntimeValueOnly) + } } #[derive(Debug)] @@ -354,12 +362,41 @@ fn lower_runtime_visible_outputs( lower_visible_observations(dae_model, layout, visible_expressions)?; crate::timing::log_stage("model.lower_visible_observations", timer); let timer = crate::timing::stage_start(); - let variable_meta = - build_variable_meta(metadata_dae_model.unwrap_or(dae_model), &visible_names)?; + let variable_meta = build_variable_meta( + metadata_dae_model.unwrap_or(dae_model), + dae_model, + &visible_names, + )?; crate::timing::log_stage("model.build_variable_meta", timer); Ok((visible_names, visible_value_rows, variable_meta)) } +fn prepare_structural_derivative_states( + dae_model: &mut dae::Dae, +) -> Result<(), SolveModelLowerError> { + let timer = crate::timing::stage_start(); + rumoca_phase_structural::dae_prepare::demote_direct_assigned_states(dae_model) + .map_err(|source| SolveModelLowerError::Structural { source })?; + rumoca_phase_structural::dae_prepare::reduce_constrained_dummy_derivatives(dae_model) + .map_err(|source| SolveModelLowerError::Structural { source })?; + crate::timing::log_stage("model.prepare_structural_derivative_states", timer); + Ok(()) +} + +fn runtime_visible_expressions( + dae_model: &dae::Dae, + visible_expressions: Option>, + profile: SolveModelLoweringProfile, +) -> Result, SolveModelLowerError> { + if !profile.needs_runtime_support() { + return Ok(Vec::new()); + } + match visible_expressions { + Some(visible_expressions) => Ok(visible_expressions), + None => visible_expressions_for_dae(dae_model).map_err(SolveModelLowerError::Lower), + } +} + fn lower_dae_to_solve_model_inner( mut dae_model: dae::Dae, visible_expressions: Option>, @@ -367,15 +404,12 @@ fn lower_dae_to_solve_model_inner( profile: SolveModelLoweringProfile, param_overrides: &HashMap, ) -> Result { + if profile.needs_structural_derivative_preparation() { + prepare_structural_derivative_states(&mut dae_model)?; + } let timer = crate::timing::stage_start(); - let visible_expressions = if profile.needs_runtime_support() { - match visible_expressions { - Some(visible_expressions) => visible_expressions, - None => visible_expressions_for_dae(&dae_model).map_err(SolveModelLowerError::Lower)?, - } - } else { - Vec::new() - }; + let visible_expressions = + runtime_visible_expressions(&dae_model, visible_expressions, profile)?; crate::timing::log_stage("model.visible_expressions", timer); let state_count = scalar_count(dae_model.variables.states.values())?; let eval_runtime = Arc::new(EvalRuntimeState::default()); @@ -428,15 +462,17 @@ fn lower_dae_to_solve_model_inner( eval_runtime.clone(), )?; crate::timing::log_stage("model.initial_solver_values", timer); - let timer = crate::timing::stage_start(); - apply_initial_equations_to_start_values( - &dae_model, - &problem.layout, - &mut parameters, - &mut initial_y, - eval_runtime.clone(), - )?; - crate::timing::log_stage("model.apply_initial_equations", timer); + if profile != SolveModelLoweringProfile::GpuPreparation { + let timer = crate::timing::stage_start(); + apply_initial_equations_to_start_values( + &dae_model, + &problem.layout, + &mut parameters, + &mut initial_y, + eval_runtime.clone(), + )?; + crate::timing::log_stage("model.apply_initial_equations", timer); + } let timer = crate::timing::stage_start(); let table_env = build_runtime_parameter_tail_env_with_declared_slots_and_runtime( &dae_model, @@ -445,7 +481,10 @@ fn lower_dae_to_solve_model_inner( eval_runtime, ) .map_err(|source| runtime_tail_error(&dae_model, source))?; - let external_tables = external_table_data_for_parameter_values_in(&table_env, ¶meters); + let external_tables = merged_external_tables( + crate::lower::external_table_data_for_dae(&dae_model)?, + external_table_data_for_parameter_values_in(&table_env, ¶meters), + ); crate::timing::log_stage("model.external_tables", timer); let (visible_names, visible_value_rows, variable_meta) = lower_runtime_visible_outputs( &dae_model, @@ -467,6 +506,38 @@ fn lower_dae_to_solve_model_inner( }) } +fn merged_external_tables( + mut tables: Vec, + parameter_tables: Vec, +) -> Vec { + tables.extend(parameter_tables); + let mut seen = HashSet::new(); + tables.retain(|table| seen.insert(table.id)); + tables.sort_by_key(|table| table.id); + tables +} + +#[cfg(test)] +#[test] +fn merged_external_tables_deduplicates_and_sorts_explicit_and_parameter_tables() { + let table = |id, marker| rumoca_core::ExternalTableData { + id, + data: vec![vec![marker]], + ..Default::default() + }; + + let merged = merged_external_tables( + vec![table(3, 30.0), table(1, 10.0), table(3, 31.0)], + vec![table(2, 20.0), table(3, 32.0)], + ); + + assert_eq!( + merged.iter().map(|table| table.id).collect::>(), + vec![1, 2, 3] + ); + assert_eq!(merged[2].data, vec![vec![30.0]]); +} + fn lower_contract_violation(reason: String, span: rumoca_core::Span) -> LowerError { if span.is_dummy() { LowerError::UnspannedContractViolation { reason } @@ -923,36 +994,23 @@ fn start_values( ) -> Result, SolveModelLowerError> { let default_start = default_start_value(dae_model, var); let size = solve_model_variable_size(var)?; + if size == 0 && !var.dims.is_empty() { + return Ok(Vec::new()); + } let Some(expr) = var.start.as_ref() else { return default_start_values_for_size(var, default_start, size); }; if start_expr_is_nonnumeric(expr, env) { return default_start_values_for_size(var, default_start, size); } - if size == 0 && !var.dims.is_empty() { - let raw = if var.dims.len() >= 2 { - match eval_matrix_values(expr, env) { - Ok(Some(matrix)) => flatten_start_matrix(matrix, var)?, - Ok(None) => { - eval_array_values::(expr, env).map_err(|err| eval_start_error(var, err))? - } - Err(err) => return Err(eval_start_error(var, err)), + if size <= 1 && var.dims.is_empty() { + let value = match eval_expr::(expr, env) { + Ok(value) => value, + Err(err) if start_eval_error_uses_default(&err) => { + return single_start_value(default_start, var); } - } else { - eval_array_values::(expr, env).map_err(|err| eval_start_error(var, err))? + Err(err) => return Err(eval_start_error(var, err)), }; - if raw.is_empty() { - return Err(eval_start_error( - var, - EvalError::UnsupportedExpression { - kind: "array start value", - }, - )); - } - return finite_start_values(raw, default_start, var); - } - if size <= 1 && var.dims.is_empty() { - let value = eval_expr::(expr, env).map_err(|err| eval_start_error(var, err))?; return single_start_value(finite_start_value(value, default_start), var); } let raw = match shaped_start_values(expr, env, size) { @@ -966,11 +1024,25 @@ fn start_values( expr.span().unwrap_or(var.source_span), )? } + Err(err) if start_eval_error_uses_default(&err) => { + return default_start_values_for_size(var, default_start, size); + } Err(err) => return Err(eval_start_error(var, err)), }; finite_start_values(raw, default_start, var) } +fn start_eval_error_uses_default(err: &EvalError) -> bool { + match err { + EvalError::UnsupportedExpression { + kind: "external table data" | "external table bounds", + } + | EvalError::UnsupportedExpression { kind: "empty" } => true, + EvalError::Spanned { source, .. } => start_eval_error_uses_default(source), + _ => false, + } +} + fn shaped_start_values( expr: &rumoca_core::Expression, env: &rumoca_eval_dae::VarEnv, @@ -1212,32 +1284,6 @@ fn seed_var_values( Ok(()) } -fn flatten_start_matrix( - matrix: Vec>, - var: &dae::Variable, -) -> Result, SolveModelLowerError> { - let mut len = 0usize; - for row in &matrix { - len = checked_solve_model_count_add( - len, - row.len(), - "matrix start value count", - var.source_span, - )?; - } - let mut values = solve_model_vec_with_capacity(len, "matrix start values", var.source_span)?; - for row in matrix { - reserve_solve_model_capacity( - &mut values, - row.len(), - "matrix start values", - var.source_span, - )?; - values.extend(row); - } - Ok(values) -} - fn single_start_value(value: f64, var: &dae::Variable) -> Result, SolveModelLowerError> { let mut values = solve_model_vec_with_capacity(1, "scalar start value", var.source_span)?; values.push(value); @@ -1643,46 +1689,149 @@ fn div_op() -> OpBinary { OpBinary::Div } +fn final_state_scalar_spans( + final_dae_model: &dae::Dae, +) -> Result, SolveModelLowerError> { + let mut scalars = IndexMap::new(); + for (name, var) in &final_dae_model.variables.states { + let names = scalar_names(name.as_str(), var)?; + reserve_solve_model_index_map_capacity( + &mut scalars, + names.len(), + "final state metadata scalar count", + var.source_span, + )?; + for name in names { + scalars.insert(name, var.source_span); + } + } + Ok(scalars) +} + +fn final_variable_meta_role( + source_role: &'static str, + source_is_state: bool, + scalar_name: &str, + final_state_scalars: &IndexMap, +) -> (&'static str, bool) { + if final_state_scalars.contains_key(scalar_name) { + ("state", true) + } else if source_role == "state" || source_role == "algebraic" { + ("algebraic", false) + } else { + (source_role, source_is_state) + } +} + +fn validate_visible_state_metadata( + final_dae_model: &dae::Dae, + final_state_scalars: &IndexMap, + visible_names: &[String], + by_scalar: &IndexMap, + meta: &[solve::SolveVariableMeta], +) -> Result<(), SolveModelLowerError> { + for (name, span) in final_state_scalars { + if !visible_names.contains(name) || by_scalar.contains_key(name) { + continue; + } + return Err(SolveModelLowerError::Lower(lower_contract_violation( + format!("final selected state `{name}` is missing from simulation metadata"), + *span, + ))); + } + let reported_state_count = meta.iter().filter(|item| item.is_state).count(); + let visible_final_state_count = visible_names + .iter() + .filter(|name| final_state_scalars.contains_key(*name)) + .count(); + if reported_state_count != visible_final_state_count { + return Err(SolveModelLowerError::Lower(lower_contract_violation( + format!( + "simulation metadata reports {reported_state_count} state scalars but {visible_final_state_count} visible scalars belong to the final solve state partition" + ), + dae_model_span(final_dae_model, "final state metadata scalar count")?, + ))); + } + Ok(()) +} + +fn build_scalar_variable_meta( + var: &dae::Variable, + scalar_name: String, + role: &str, + is_state: bool, + event_discontinuous_names: &IndexSet, +) -> solve::SolveVariableMeta { + let (value_type, variability, time_domain) = variable_meta_classification(role, is_state); + let time_domain = if continuous_real_role_is_event_discontinuous( + role, + &scalar_name, + event_discontinuous_names, + ) { + Some("event-discontinuous".to_string()) + } else { + time_domain + }; + solve::SolveVariableMeta { + name: scalar_name, + source_span: var.source_span, + role: role.to_string(), + is_state, + value_type, + variability, + time_domain, + unit: var.unit.clone(), + start: var.start.as_ref().map(|expr| format!("{expr:?}")), + min: var.min.as_ref().map(|expr| format!("{expr:?}")), + max: var.max.as_ref().map(|expr| format!("{expr:?}")), + nominal: var.nominal.as_ref().map(|expr| format!("{expr:?}")), + fixed: var.fixed, + description: var.description.clone(), + } +} + fn build_variable_meta( - dae_model: &dae::Dae, + metadata_dae_model: &dae::Dae, + final_dae_model: &dae::Dae, visible_names: &[String], ) -> Result, SolveModelLowerError> { - let event_discontinuous_names = event_discontinuous_scalar_names(dae_model)?; - let vars = dae_model + let event_discontinuous_names = event_discontinuous_scalar_names(metadata_dae_model)?; + let final_state_scalars = final_state_scalar_spans(final_dae_model)?; + let vars = metadata_dae_model .variables .states .iter() .map(|(name, var)| (name, var, "state", true)) .chain( - dae_model + metadata_dae_model .variables .algebraics .iter() .map(|(name, var)| (name, var, "algebraic", false)), ) .chain( - dae_model + metadata_dae_model .variables .outputs .iter() .map(|(name, var)| (name, var, "output", false)), ) .chain( - dae_model + metadata_dae_model .variables .inputs .iter() .map(|(name, var)| (name, var, "input", false)), ) .chain( - dae_model + metadata_dae_model .variables .discrete_reals .iter() .map(|(name, var)| (name, var, "discrete-real", false)), ) .chain( - dae_model + metadata_dae_model .variables .discrete_valued .iter() @@ -1690,11 +1839,6 @@ fn build_variable_meta( ); let mut by_scalar = IndexMap::new(); for (name, var, role, is_state) in vars { - let (value_type, variability, time_domain) = variable_meta_classification(role, is_state); - let start_text = var.start.as_ref().map(|expr| format!("{expr:?}")); - let min_text = var.min.as_ref().map(|expr| format!("{expr:?}")); - let max_text = var.max.as_ref().map(|expr| format!("{expr:?}")); - let nominal_text = var.nominal.as_ref().map(|expr| format!("{expr:?}")); let scalar_names = scalar_names(name.as_str(), var)?; reserve_solve_model_index_map_capacity( &mut by_scalar, @@ -1703,40 +1847,22 @@ fn build_variable_meta( var.source_span, )?; for scalar_name in scalar_names { - let time_domain = if continuous_real_role_is_event_discontinuous( + let (role, is_state) = + final_variable_meta_role(role, is_state, &scalar_name, &final_state_scalars); + let variable_meta = build_scalar_variable_meta( + var, + scalar_name.clone(), role, - &scalar_name, + is_state, &event_discontinuous_names, - ) { - Some("event-discontinuous".to_string()) - } else { - time_domain.clone() - }; - by_scalar.insert( - scalar_name.clone(), - solve::SolveVariableMeta { - name: scalar_name, - source_span: var.source_span, - role: role.to_string(), - is_state, - value_type: value_type.clone(), - variability: variability.clone(), - time_domain: time_domain.clone(), - unit: var.unit.clone(), - start: start_text.clone(), - min: min_text.clone(), - max: max_text.clone(), - nominal: nominal_text.clone(), - fixed: var.fixed, - description: var.description.clone(), - }, ); + by_scalar.insert(scalar_name, variable_meta); } } let mut meta = solve_model_vec_with_capacity( visible_names.len(), "visible variable metadata count", - dae_model_span(dae_model, "visible variable metadata count") + dae_model_span(metadata_dae_model, "visible variable metadata count") .map_err(SolveModelLowerError::Lower)?, )?; for name in visible_names { @@ -1744,6 +1870,13 @@ fn build_variable_meta( meta.push(variable_meta.clone()); } } + validate_visible_state_metadata( + final_dae_model, + &final_state_scalars, + visible_names, + &by_scalar, + &meta, + )?; Ok(meta) } diff --git a/crates/rumoca-phase-solve/src/solve_model/tests.rs b/crates/rumoca-phase-solve/src/solve_model/tests.rs index 2e82bed65..b700c3625 100644 --- a/crates/rumoca-phase-solve/src/solve_model/tests.rs +++ b/crates/rumoca-phase-solve/src/solve_model/tests.rs @@ -135,6 +135,17 @@ fn call_expr(name: &str, args: Vec) -> rumoca_core::Exp } } +fn builtin_call_expr( + function: rumoca_core::BuiltinFunction, + args: Vec, +) -> rumoca_core::Expression { + rumoca_core::Expression::BuiltinCall { + function, + args, + span: solve_model_test_span(), + } +} + fn constructor_call_expr( name: &str, args: Vec, @@ -471,6 +482,62 @@ fn default_start_values_for_size_preserves_scalar_default_start() { ); } +#[test] +fn start_values_skips_zero_length_array_expression_evaluation() { + let mut var = scalar_var("empty_c_start"); + var.dims = vec![0]; + var.start = Some(builtin_call_expr( + rumoca_core::BuiltinFunction::Fill, + vec![ + int_expr(0), + builtin_call_expr( + rumoca_core::BuiltinFunction::Size, + vec![ + builtin_call_expr( + rumoca_core::BuiltinFunction::Fill, + vec![string_expr(""), int_expr(0)], + ), + int_expr(1), + ], + ), + ], + )); + let env = rumoca_eval_dae::VarEnv::::new(); + + let values = start_values(&dae::Dae::default(), &var, &env) + .expect("zero-length array start should not evaluate element expressions"); + + assert!(values.is_empty()); +} + +#[test] +fn start_values_use_default_for_external_table_data_start_guess() { + let mut var = scalar_var("table_u_min"); + var.start = Some(call_expr( + "getTimeTableTmin", + vec![call_expr( + "ExternalCombiTimeTable", + vec![ + string_expr("NoName"), + string_expr("NoName"), + rumoca_core::Expression::Empty { + span: solve_model_test_span(), + }, + real_expr(0.0), + array_expr(vec![int_expr(2)], false), + int_expr(3), + int_expr(1), + ], + )], + )); + let env = rumoca_eval_dae::VarEnv::::new(); + + let values = start_values(&dae::Dae::default(), &var, &env) + .expect("external table metadata start should fall back to default guess"); + + assert_eq!(values, vec![0.0]); +} + #[test] fn identity_mass_matrix_reports_capacity_overflow() -> Result<(), SolveModelLowerError> { let span = rumoca_core::Span::from_offsets( @@ -760,6 +827,48 @@ fn lower_keeps_algebraic_binding_needed_by_derivative_rhs() { assert_eq!(prepared.initial_y.len(), 2); } +#[test] +fn lower_reduces_constrained_dummy_derivative_chain_before_solve_layout() { + let mut dae_model = dae::Dae::default(); + dae_model + .variables + .states + .insert("s".into(), scalar_var("s")); + dae_model + .variables + .states + .insert("sd".into(), scalar_var("sd")); + dae_model + .variables + .algebraics + .insert("sdd".into(), scalar_var("sdd")); + + dae_model.continuous.equations.push(dae::Equation::residual( + sub(var("s"), var("time")), + solve_model_test_span(), + "analytic position trajectory", + )); + dae_model.continuous.equations.push(dae::Equation::residual( + sub(var("sd"), der(var("s"))), + solve_model_test_span(), + "velocity dummy derivative", + )); + dae_model.continuous.equations.push(dae::Equation::residual( + sub(var("sdd"), der(var("sd"))), + solve_model_test_span(), + "acceleration dummy derivative", + )); + + let prepared = lower_dae_to_solve_model(&dae_model) + .expect("constrained derivative chain should lower without structural singularity"); + + assert_eq!(prepared.problem.solve_layout.state_scalar_count(), 0); + assert!( + prepared.problem.layout.binding("sdd").is_some(), + "demoted derivative alias remains available as an algebraic value" + ); +} + #[test] fn lower_replaces_nonfinite_start_guess_with_type_default() { let mut dae_model = dae::Dae::default(); @@ -848,6 +957,177 @@ fn lower_seeds_start_dependencies_removed_by_structural_metadata() { assert_eq!(prepared.initial_y, vec![0.0]); } +#[test] +fn lower_reports_final_partial_array_state_partition_in_metadata() { + let mut dae_model = dae::Dae::default(); + let mut p = scalar_var("p"); + p.dims = vec![3]; + dae_model + .variables + .states + .insert(rumoca_core::VarName::new("p"), p); + dae_model.continuous.equations.push(dae::Equation::residual( + sub(plain_var("p[1]", solve_model_test_span()), var("time")), + solve_model_test_span(), + "direct p[1] trajectory", + )); + for index in 2..=3 { + dae_model.continuous.equations.push(dae::Equation::residual( + sub( + der(plain_var( + format!("p[{index}]").as_str(), + solve_model_test_span(), + )), + real_expr(0.0), + ), + solve_model_test_span(), + format!("p[{index}] ODE"), + )); + } + let metadata_dae = dae_model.clone(); + let parent = dae_model + .variables + .states + .shift_remove(&rumoca_core::VarName::new("p")) + .expect("parent array state"); + for index in 1..=3 { + let name = rumoca_core::VarName::new(format!("p[{index}]")); + let mut scalar = parent.clone(); + scalar.name = name.clone(); + scalar.dims.clear(); + if index == 1 { + dae_model.variables.algebraics.insert(name, scalar); + } else { + dae_model.variables.states.insert(name, scalar); + } + } + let visible = (1..=3) + .map(|index| VisibleExpression { + name: format!("p[{index}]"), + expr: plain_var(format!("p[{index}]").as_str(), solve_model_test_span()), + }) + .collect(); + + let model = lower_dae_to_solve_model_owned_with_visible_expressions_and_metadata( + dae_model, + visible, + &metadata_dae, + ) + .expect("partial array state partition should lower"); + + let roles = model + .variable_meta + .iter() + .map(|meta| (meta.name.as_str(), meta.role.as_str())) + .collect::>(); + assert_eq!( + roles, + [("p[1]", "algebraic"), ("p[2]", "state"), ("p[3]", "state")] + ); + assert_eq!(model.state_scalar_count(), 2); + assert_eq!( + model + .variable_meta + .iter() + .filter(|meta| meta.is_state) + .count(), + 2 + ); +} + +#[test] +fn lower_reports_final_late_promote_and_demote_partition_in_metadata() { + let mut dae_model = dae::Dae::default(); + dae_model + .variables + .algebraics + .insert(rumoca_core::VarName::new("x"), scalar_var("x")); + let mut y = scalar_var("y"); + y.state_select = rumoca_core::StateSelect::Prefer; + dae_model + .variables + .states + .insert(rumoca_core::VarName::new("y"), y); + dae_model.continuous.equations.push(dae::Equation::residual( + sub(var("x"), var("y")), + solve_model_test_span(), + "direct x alias", + )); + dae_model.continuous.equations.push(dae::Equation::residual( + sub(der(var("y")), real_expr(1.0)), + solve_model_test_span(), + "y ODE", + )); + let mut metadata_dae = dae_model.clone(); + let x = metadata_dae + .variables + .algebraics + .shift_remove(&rumoca_core::VarName::new("x")) + .expect("x metadata"); + let y = metadata_dae + .variables + .states + .shift_remove(&rumoca_core::VarName::new("y")) + .expect("y metadata"); + metadata_dae + .variables + .states + .insert(rumoca_core::VarName::new("x"), x); + metadata_dae + .variables + .algebraics + .insert(rumoca_core::VarName::new("y"), y); + let visible = ["x", "y"] + .into_iter() + .map(|name| VisibleExpression { + name: name.to_string(), + expr: var(name), + }) + .collect(); + + let model = lower_dae_to_solve_model_owned_with_visible_expressions_and_metadata( + dae_model, + visible, + &metadata_dae, + ) + .expect("late state promotion and demotion should lower"); + + let roles = model + .variable_meta + .iter() + .map(|meta| (meta.name.as_str(), meta.role.as_str())) + .collect::>(); + assert_eq!(roles, [("x", "algebraic"), ("y", "state")]); + assert_eq!(model.state_scalar_count(), 1); + assert_eq!( + model + .variable_meta + .iter() + .filter(|meta| meta.is_state) + .count(), + 1 + ); +} + +#[test] +fn build_variable_meta_rejects_final_state_missing_from_metadata() { + let mut final_dae = dae::Dae::default(); + final_dae + .variables + .states + .insert(rumoca_core::VarName::new("x"), scalar_var("x")); + let mut metadata_dae = dae::Dae::default(); + metadata_dae + .variables + .algebraics + .insert(rumoca_core::VarName::new("dummy"), scalar_var("dummy")); + + let error = build_variable_meta(&metadata_dae, &final_dae, &["x".to_string()]) + .expect_err("missing final-state metadata must fail closed"); + + assert!(error.to_string().contains("x"), "got: {error}"); +} + #[test] fn lower_rejects_missing_binding_in_explicit_start_guess() { let start_span = rumoca_core::Span::from_offsets( @@ -1223,8 +1503,8 @@ fn variable_meta_marks_relation_driven_real_output_event_discontinuous() { scalar_count: 1, }); - let meta = - build_variable_meta(&dae_model, &["y".to_string()]).expect("valid meta should build"); + let meta = build_variable_meta(&dae_model, &dae_model, &["y".to_string()]) + .expect("valid meta should build"); assert_eq!(meta.len(), 1); assert_eq!(meta[0].variability.as_deref(), Some("continuous")); @@ -1262,8 +1542,8 @@ fn variable_meta_marks_guarded_residual_output_event_discontinuous() { scalar_count: 1, }); - let meta = - build_variable_meta(&dae_model, &["y".to_string()]).expect("valid meta should build"); + let meta = build_variable_meta(&dae_model, &dae_model, &["y".to_string()]) + .expect("valid meta should build"); assert_eq!(meta.len(), 1); assert_eq!(meta[0].variability.as_deref(), Some("continuous")); @@ -1312,8 +1592,8 @@ fn variable_meta_propagates_event_discontinuity_through_algebraic_dependency() { scalar_count: 1, }); - let meta = - build_variable_meta(&dae_model, &["y".to_string()]).expect("valid meta should build"); + let meta = build_variable_meta(&dae_model, &dae_model, &["y".to_string()]) + .expect("valid meta should build"); assert_eq!(meta.len(), 1); assert_eq!(meta[0].time_domain.as_deref(), Some("event-discontinuous")); @@ -1730,3 +2010,44 @@ fn lower_propagates_enum_parameter_binding_chain_to_runtime_start() { }; assert_eq!(prepared.parameters[index], 3.0); } + +#[test] +fn gpu_profile_skips_runtime_demotion_and_keeps_state_only_layout() { + assert!( + !SolveModelLoweringProfile::GpuPreparation.needs_structural_derivative_preparation(), + "the GPU profile must not run runtime-only direct-state demotion" + ); + + let mut dae_model = dae::Dae::default(); + let mut state = scalar_var("x"); + state.start = Some(int_expr(7)); + dae_model + .variables + .states + .insert(rumoca_core::VarName::new("x"), state); + dae_model.continuous.equations.push(dae::Equation::residual( + sub(der(var("x")), int_expr(0)), + solve_model_test_span(), + "der(x) = 0", + )); + let metadata = dae_model.clone(); + + let gpu = + lower_dae_to_solve_model_owned_for_gpu_preparation_with_metadata(dae_model, &metadata) + .expect("GPU preparation should preserve a direct state-only DAE"); + + assert_eq!(gpu.state_scalar_count(), 1); + assert_eq!( + gpu.problem.layout.binding("x"), + Some(solve::scalar_slot_y(0)) + ); + assert_eq!(gpu.initial_y, vec![7.0]); + assert!(gpu.problem.continuous.residual.is_empty()); + assert!( + gpu.problem + .continuous + .algebraic_projection_plan + .blocks + .is_empty() + ); +} diff --git a/crates/rumoca-phase-solve/src/stencil.rs b/crates/rumoca-phase-solve/src/stencil.rs index 052713f9b..b1a38c0cc 100644 --- a/crates/rumoca-phase-solve/src/stencil.rs +++ b/crates/rumoca-phase-solve/src/stencil.rs @@ -623,6 +623,14 @@ fn structured_dae_body_shapes_match( Ok(true) } +pub(crate) fn dae_equation_body_shapes_match( + first: &dae::Equation, + second: &dae::Equation, +) -> Result { + Ok(expression_body_shape(&first.rhs, first.span)? + == expression_body_shape(&second.rhs, second.span)?) +} + #[derive(Debug, Clone, PartialEq)] enum ExpressionBodyShape { Binary { @@ -837,7 +845,8 @@ fn max_structured_affine_domain( // A `regular` family is affine over its whole domain by construction (flatten // only classifies it regular when every cell shares one affine body), so the // entire candidate range is the stencil -- no need to search shrinking prefixes, - // and validation reads only the corner rows (base + one neighbor per binder). + // and validation reads only the corner rows (base + one neighbor per + // non-singleton binder). // Leaving the interior rows unread is the contract that lets flatten stop // materializing them. A non-regular family keeps the full prefix search below. if family.regular.is_some() { @@ -902,8 +911,8 @@ fn max_structured_affine_domain( /// Preserve a `regular` family's entire candidate range as one stencil domain, /// validated from corner rows only (body shape, strides, output map). Returns /// `None` to fall back to the prefix search when the corner model does not apply -/// (e.g. an extent-1 binder, or a corner that fails to validate) -- the same -/// conservative fallback the per-part corner helpers use. +/// (e.g. a corner fails to validate) -- the same conservative fallback the +/// per-part corner helpers use. fn regular_family_full_domain( rows: &[StructuredProgram], row_indices: &[usize], @@ -961,10 +970,10 @@ fn regular_family_full_domain( } /// Like [`structured_dae_body_shapes_match`] but reads only the family's corner -/// rows (base + one neighbor per binder). For a regular family every cell shares -/// one body, so the corners are representative -- this avoids reading the interior -/// rows' DAE bodies. Falls back to the full check when corners cannot be located -/// (e.g. an extent-1 binder). +/// rows (base + one neighbor per non-singleton binder). For a regular family every +/// cell shares one body, so the corners are representative -- this avoids reading +/// the interior rows' DAE bodies. Falls back to the full check when corners cannot +/// be located. Singleton binders need no neighbor because their stride is zero. fn corner_dae_body_shapes_match( rows: &[StructuredProgram], row_indices: &[usize], @@ -978,8 +987,8 @@ fn corner_dae_body_shapes_match( return Ok(false); } let Some(corner_positions) = corner_index_positions(&index_tuples, domain) else { - // No corners (e.g. an extent-1 binder): the full check reads every row's - // body, which is only valid when the interior bodies are real. + // The full check reads every row's body, which is only valid when the + // interior bodies are real. return if interiors_materialized { structured_dae_body_shapes_match(rows, row_indices, dae_equations) } else { @@ -1124,11 +1133,11 @@ fn affine_strides_from_representative_proofs( /// Load/const strides for a regular family, corner-first (P1). The strides depend /// only on the access pattern, so the family's corner rows (base + one neighbor -/// per dimension) yield the same result as diffing every row -- this is the +/// per non-singleton dimension) yield the same result as diffing every row -- this is the /// production path, so the interior rows are not read when the corner model /// applies. Falls back to the full-row derivation when the corners cannot be -/// isolated (e.g. an extent-1 dimension leaves a binder with no neighbor), keeping -/// behavior identical for families the corner model does not cover. In debug +/// isolated, keeping behavior identical for families the corner model does not +/// cover. In debug /// builds, when both succeed, asserts they agree -- the equivalence that lets /// flatten (P3) stop materializing the interior rows. fn affine_strides_for_family( @@ -1149,8 +1158,7 @@ fn affine_strides_for_family( // The corner path is row-blind (it reads only the corners), so if it builds a // stencil, diffing every row must build the same one -- otherwise a misclassified // family could slip a wrong stencil past it where the full scan would have declined - // to scalar. The reverse (corners decline, full succeeds) is the intended extent-1 - // fallback and is left to the `match` below. + // to scalar. The reverse remains a conservative full-row fallback. if cfg!(debug_assertions) && let Some(corner) = &corner { @@ -1202,13 +1210,13 @@ fn output_map_for_family( } /// Affine strides computed from only the family's CORNER rows: the base -/// iteration plus one neighbor per domain dimension. For a regular family (affine +/// iteration plus one neighbor per non-singleton domain dimension. For a regular family (affine /// accesses are guaranteed by construction) these O(ndim) rows determine the same /// strides the full-row inference produces -- this is the basis for building the /// stencil node without materializing every iteration. Handles `LoadP` /// (derived-parameter `c`) loads identically to `LoadY`, via the same per-tuple -/// index fit. Returns `None` (caller falls back to the full-row path) when the -/// corners cannot be located (e.g. an extent-1 dimension) or a proof is missing. +/// index fit. Singleton dimensions contribute zero stride without a neighbor. +/// Returns `None` (caller falls back) when a required corner or proof is missing. fn affine_strides_from_corner_rows( rows: &[StructuredProgram], row_indices: &[usize], @@ -1243,11 +1251,11 @@ fn affine_strides_from_corner_rows( proof_strides_for_base_ops(base_row.ops.as_slice(), base_proof, operand_deltas, span) } -/// The family's corner rows (base iteration + one neighbor per domain dimension) -/// paired with their index tuples, in `[base, +dim0, +dim1, ...]` order. This is -/// the shared selection both corner-row builders (strides and output map) consume. -/// `None` when the rows don't cover the full domain or a corner is absent (e.g. an -/// extent-1 dimension) -- callers then fall back to the full-row path. +/// The family's corner rows (base iteration + one neighbor per non-singleton domain +/// dimension) paired with their index tuples, in `[base, +dim0, +dim1, ...]` order. +/// This is the shared selection both corner-row builders (strides and output map) +/// consume. `None` when the rows don't cover the full domain or a required corner +/// is absent; callers then fall back to the full-row path. fn corner_rows_and_tuples<'a>( rows: &'a [StructuredProgram], row_indices: &[usize], @@ -1278,8 +1286,8 @@ fn corner_rows_and_tuples<'a>( /// Positions (into the family's row order) of the base iteration and one neighbor /// per domain dimension: position 0 is the first tuple (domain lower bounds), and /// for each dimension the tuple equal to base with that dimension advanced one -/// step. `None` when a dimension's neighbor is absent (an extent-1 dimension, -/// which cannot pin that dimension's stride -- the full-row path handles those). +/// step. Extent-1 dimensions are omitted because their stride is necessarily zero. +/// `None` means the domain is empty/invalid or a required neighbor is absent. fn corner_index_positions( index_tuples: &[Vec], domain: &rumoca_core::StructuredIndexDomain, @@ -1287,8 +1295,12 @@ fn corner_index_positions( let base = index_tuples.first()?; let mut positions = vec![0usize]; for (dimension, binder) in domain.binders.iter().enumerate() { + if binder.value_count().ok()? == 1 { + continue; + } let mut neighbor = base.clone(); - *neighbor.get_mut(dimension)? += binder.step; + let value = neighbor.get_mut(dimension)?; + *value = value.checked_add(binder.step)?; let position = index_tuples.iter().position(|tuple| tuple == &neighbor)?; positions.push(position); } @@ -1296,7 +1308,7 @@ fn corner_index_positions( } /// The tensor output map computed from only the family's corner rows (base + one -/// neighbor per dimension), mirroring `output_map_for_rows` on that subset. For a +/// neighbor per non-singleton dimension), mirroring `output_map_for_rows` on that subset. For a /// regular family the output index is affine in the binders, so the corners pin /// it exactly. `None` (caller falls back) when the corners cannot be located. fn output_map_from_corner_rows( diff --git a/crates/rumoca-phase-solve/src/stencil/access_proof.rs b/crates/rumoca-phase-solve/src/stencil/access_proof.rs index 00b02d429..0d6f492d2 100644 --- a/crates/rumoca-phase-solve/src/stencil/access_proof.rs +++ b/crates/rumoca-phase-solve/src/stencil/access_proof.rs @@ -108,6 +108,23 @@ where subscripts, span, } => collect_var_ref(name.as_str(), subscripts, *span, operands), + rumoca_core::Expression::Index { + base, + subscripts, + span, + } => { + let rumoca_core::Expression::VarRef { + name, + subscripts: base_subscripts, + .. + } = base.as_ref() + else { + return Ok(None); + }; + let mut combined = base_subscripts.clone(); + combined.extend(subscripts.iter().cloned()); + collect_var_ref(name.as_str(), &combined, *span, operands) + } rumoca_core::Expression::Binary { lhs, rhs, .. } => { let Some(()) = collect_structured_access_operands_result_inner(lhs, operands, collect_var_ref)? diff --git a/crates/rumoca-phase-solve/src/stencil/tests.rs b/crates/rumoca-phase-solve/src/stencil/tests.rs index 2e551d79b..bba1e6c72 100644 --- a/crates/rumoca-phase-solve/src/stencil/tests.rs +++ b/crates/rumoca-phase-solve/src/stencil/tests.rs @@ -221,6 +221,24 @@ fn test_domain_3d(d0: usize, d1: usize, d2: usize) -> rumoca_core::StructuredInd } } +fn test_domain_ranges(ranges: &[(i64, i64, i64)]) -> rumoca_core::StructuredIndexDomain { + rumoca_core::StructuredIndexDomain { + binders: ranges + .iter() + .enumerate() + .map( + |(id, &(lower, upper, step))| rumoca_core::StructuredIndexBinder { + id, + display_name: format!("i{id}"), + lower, + upper, + step, + }, + ) + .collect(), + } +} + fn family_2d(start: usize, rows: usize, cols: usize) -> dae::StructuredEquationFamily { dae::StructuredEquationFamily { domain: test_domain_2d(rows, cols), @@ -1153,21 +1171,54 @@ fn corner_rows_reproduce_full_row_affine_strides_for_2d_family() { } #[test] -fn corner_rows_decline_when_a_dimension_has_extent_one() { - // With j pinned to a single value there is no j-neighbor to pin its stride, so - // the corner path declines (returns None) and the caller falls back to the - // full-row inference. - let domain = test_domain_2d(3, 1); - let count = domain_scalar_count(&domain); - let rows: Vec = (0..count).map(|y| stencil_row(y, y)).collect(); - let row_indices: Vec = (0..count).collect(); +fn corner_rows_infer_zero_stride_for_singleton_dimensions() { + let cases = [ + (test_domain_2d(3, 1), vec![(0, 1)]), + (test_domain_2d(1, 3), vec![(1, 1)]), + (test_domain_3d(1, 3, 1), vec![(1, 1)]), + (test_domain_2d(1, 1), vec![]), + (test_domain_ranges(&[(3, 1, -1), (5, 5, -2)]), vec![(0, 1)]), + ]; + for (domain, expected_terms) in cases { + let count = domain_scalar_count(&domain); + let rows: Vec = (0..count).map(|y| stencil_row(y, y)).collect(); + let row_indices: Vec = (0..count).collect(); + let span = stencil_test_span(); + + let full = affine_strides_from_access_proofs(&rows, &row_indices, &domain, span) + .expect("full-row strides should compute") + .expect("full-row strides should be present"); + let corner = affine_strides_from_corner_rows(&rows, &row_indices, &domain, span) + .expect("corner-row strides should compute") + .expect("singleton dimensions should use zero stride"); + assert_eq!(corner, full); + for load in &corner.load_strides { + let terms = load + .terms + .iter() + .map(|term| (term.dimension, term.stride)) + .collect::>(); + assert_eq!(terms, expected_terms); + } + assert_eq!( + output_map_from_corner_rows(&rows, &row_indices, &domain, span) + .expect("corner output map should compute"), + output_map_for_rows(&rows, &row_indices, &domain, span) + .expect("full output map should compute") + ); + } +} + +#[test] +fn corner_positions_decline_for_empty_or_invalid_domains() { let span = stencil_test_span(); + let empty = test_domain_ranges(&[(1, 0, 1), (1, 3, 1)]); + let empty_tuples = structured_domain_index_tuples(&empty, span).expect("empty domain is valid"); + assert!(empty_tuples.is_empty()); + assert_eq!(corner_index_positions(&empty_tuples, &empty), None); - assert_eq!( - affine_strides_from_corner_rows(&rows, &row_indices, &domain, span) - .expect("corner-row strides should compute"), - None - ); + let invalid = test_domain_ranges(&[(1, 3, 0)]); + assert!(structured_domain_index_tuples(&invalid, span).is_err()); } #[test] @@ -1190,24 +1241,23 @@ fn affine_strides_for_family_uses_corners_then_falls_back_to_full_scan() { "with corners available the wrapper returns the corner result (== full)" ); - // 3x1: corners decline (no j-neighbor), so the wrapper must fall back to exactly - // the full-row result rather than declining to scalar. + // 3x1: the singleton j dimension contributes zero stride, so the corner path + // remains available even when interior rows are non-semantic placeholders. let domain = test_domain_2d(3, 1); let count = domain_scalar_count(&domain); let rows: Vec = (0..count).map(|y| stencil_row(y, y)).collect(); let row_indices: Vec = (0..count).collect(); + let full = affine_strides_from_access_proofs(&rows, &row_indices, &domain, span) + .expect("full-row strides should compute"); assert_eq!( affine_strides_from_corner_rows(&rows, &row_indices, &domain, span) .expect("corner-row strides should compute"), - None, - "precondition: corners decline on an extent-1 dimension" + full ); assert_eq!( - affine_strides_for_family(&rows, &row_indices, &domain, span, true) - .expect("family strides should compute"), - affine_strides_from_access_proofs(&rows, &row_indices, &domain, span) - .expect("full-row strides should compute"), - "with corners declined the wrapper returns the full-row fallback result" + affine_strides_for_family(&rows, &row_indices, &domain, span, false) + .expect("unmaterialized family strides should compute"), + full ); } diff --git a/crates/rumoca-phase-solve/src/tests.rs b/crates/rumoca-phase-solve/src/tests.rs index e195cef04..ebe3f4f05 100644 --- a/crates/rumoca-phase-solve/src/tests.rs +++ b/crates/rumoca-phase-solve/src/tests.rs @@ -1,8 +1,13 @@ +//! SPEC_0021 file-size exception: phase-solve integration tests still cover +//! cross-module lowering interactions. split plan: move remaining scalarization +//! and projection regression groups into focused test modules. + use super::*; use std::collections::BTreeSet; - mod derivative_row_tests; +mod direct_scalar_residual_projection; mod gpu_preparation; +mod projection_loop_matching; mod projection_plan_more; fn solve_test_span() -> rumoca_core::Span { @@ -37,6 +42,71 @@ fn scalar_program_block_fixture( .expect("valid solve-lowering fixture should scalarize") } +fn diagonal_linsolve_node(output_indices: Vec) -> rumoca_ir_solve::ComputeNode { + rumoca_ir_solve::ComputeNode::LinSolve { + setup_ops: vec![ + rumoca_ir_solve::LinearOp::Const { dst: 0, value: 2.0 }, + rumoca_ir_solve::LinearOp::Const { dst: 1, value: 0.0 }, + rumoca_ir_solve::LinearOp::Const { dst: 2, value: 0.0 }, + rumoca_ir_solve::LinearOp::Const { dst: 3, value: 4.0 }, + rumoca_ir_solve::LinearOp::Const { dst: 4, value: 8.0 }, + rumoca_ir_solve::LinearOp::Const { + dst: 5, + value: 20.0, + }, + ], + matrix_start: 0, + rhs_start: 4, + n: 2, + next_reg: 6, + output_indices, + metadata: Default::default(), + span: solve_test_span(), + } +} + +#[test] +fn implicit_rhs_remaps_explicit_linsolve_outputs_and_preserves_evaluation_slots() { + let residual = rumoca_ir_solve::ComputeBlock { + nodes: vec![diagonal_linsolve_node(vec![0, 2])], + }; + let remapped = super::implicit_rhs::remap_residual_compute_nodes( + &residual, + &[Some(1), None, Some(2)], + 1, + solve_test_span(), + ) + .expect("explicit LinSolve residual outputs should remap") + .expect("all LinSolve residual outputs are represented in the implicit block"); + let block = rumoca_ir_solve::ComputeBlock { nodes: remapped }; + + assert!(matches!( + block.nodes.as_slice(), + [rumoca_ir_solve::ComputeNode::LinSolve { output_indices, .. }] + if output_indices == &[1, 2] + )); + assert_eq!( + scalar_program_block_fixture(&block).output_indices, + vec![1, 2], + "scalarization must retain the translated implicit slots" + ); + + let prepared = rumoca_eval_solve::PreparedComputeBlock::new(&block) + .expect("remapped LinSolve should prepare"); + let mut output = vec![0.0; prepared.len()]; + prepared + .eval_with_context( + &[], + &[], + 0.0, + rumoca_eval_solve::RowEvalContext::default(), + &mut output, + ) + .expect("remapped LinSolve should scatter to translated slots"); + + assert_eq!(output, vec![0.0, 4.0, 5.0]); +} + fn array_var(name: &str, dims: &[i64]) -> dae::Variable { dae::Variable { name: rumoca_core::VarName::new(name), @@ -484,6 +554,47 @@ fn target_expr_scalar_name_accepts_spanned_index_base_ref() -> Result<(), LowerE Ok(()) } +#[test] +fn target_expr_scalar_name_folds_singleton_projection_of_known_scalar_target() +-> Result<(), LowerError> { + let mut dae_model = dae::Dae::new(); + dae_model.variables.algebraics.insert( + rumoca_core::VarName::new("volume.heatTransfer.heatPorts[1].T"), + scalar_var("volume.heatTransfer.heatPorts[1].T"), + ); + let span = solve_numbered_span(12, 5, 16); + let expr = rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("volume.heatTransfer.heatPorts[1].T"), + subscripts: vec![rumoca_core::Subscript::index(1, span)], + span, + }; + + let name = target_expr_scalar_name(&dae_model, &expr, 0, 1)?; + + assert_eq!(name.as_deref(), Some("volume.heatTransfer.heatPorts[1].T")); + Ok(()) +} + +#[test] +fn target_expr_scalar_name_preserves_real_array_target_subscript() -> Result<(), LowerError> { + let mut dae_model = dae::Dae::new(); + dae_model + .variables + .algebraics + .insert(rumoca_core::VarName::new("x"), array_var("x", &[2])); + let span = solve_numbered_span(13, 5, 16); + let expr = rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::new("x"), + subscripts: vec![rumoca_core::Subscript::index(1, span)], + span, + }; + + let name = target_expr_scalar_name(&dae_model, &expr, 0, 1)?; + + assert_eq!(name.as_deref(), Some("x[1]")); + Ok(()) +} + #[test] fn algebraic_projection_plan_uses_blt_scalar_blocks() { let rows = vec![ @@ -551,6 +662,60 @@ fn projection_incidence_uses_store_output_slice() { assert_eq!(y_indices, BTreeSet::from([12])); } +#[test] +fn projection_plan_row_selection_respects_explicit_target_and_identity_owner() { + let row = vec![ + solve::LinearOp::LoadY { dst: 0, index: 1 }, + solve::LinearOp::LoadY { dst: 1, index: 2 }, + solve::LinearOp::Binary { + dst: 2, + op: solve::BinaryOp::Add, + lhs: 0, + rhs: 1, + }, + solve::LinearOp::StoreOutput { src: 2 }, + ]; + let projection_set = BTreeSet::from([1, 2]); + let identity_projection_rows = std::collections::BTreeMap::from([(2, 1)]); + + assert_eq!( + super::projection_row_y_indices_for_plan( + &row, + None, + 0, + true, + &projection_set, + false, + &identity_projection_rows, + ), + Some(BTreeSet::from([1])) + ); + assert_eq!( + super::projection_row_y_indices_for_plan( + &row, + Some(solve::scalar_slot_y(2)), + 0, + true, + &projection_set, + false, + &identity_projection_rows, + ), + Some(BTreeSet::from([2])) + ); + assert_eq!( + super::projection_row_y_indices_for_plan( + &row, + None, + 0, + false, + &projection_set, + false, + &identity_projection_rows, + ), + None + ); +} + #[test] fn implicit_rhs_records_residual_row_placement() { let derivative = vec![constant_row(10.0)]; @@ -583,6 +748,56 @@ fn implicit_rhs_records_residual_row_placement() { assert_eq!(implicit.row_targets[2], Some(solve::scalar_slot_y(0))); } +#[test] +fn implicit_rhs_keeps_uncovered_non_state_rows_as_y_identity() { + let derivative = vec![constant_row(10.0)]; + + let implicit = super::build_implicit_rhs_rows(&derivative, &[], &[], 1, 3, solve_test_span()) + .expect("implicit rows should build"); + + assert_eq!(implicit.rows[0], derivative[0]); + assert_eq!( + implicit.rows[1], + vec![ + solve::LinearOp::LoadY { dst: 0, index: 1 }, + solve::LinearOp::StoreOutput { src: 0 }, + ] + ); + assert_eq!( + implicit.rows[2], + vec![ + solve::LinearOp::LoadY { dst: 0, index: 2 }, + solve::LinearOp::StoreOutput { src: 0 }, + ] + ); + assert_eq!( + implicit.residual_to_implicit_rows, + Vec::>::new() + ); +} + +#[test] +fn implicit_rhs_preserves_non_state_target_on_fallback_row() { + let derivative = vec![constant_row(10.0)]; + let residual = vec![constant_row(20.0), constant_row(30.0)]; + let residual_targets = vec![Some(solve::scalar_slot_y(2)), Some(solve::scalar_slot_y(2))]; + + let implicit = super::build_implicit_rhs_rows( + &derivative, + &residual, + &residual_targets, + 1, + 3, + solve_test_span(), + ) + .expect("implicit rows should build"); + + assert_eq!(implicit.rows[2], residual[0]); + assert_eq!(implicit.rows[1], residual[1]); + assert_eq!(implicit.row_targets[2], Some(solve::scalar_slot_y(2))); + assert_eq!(implicit.row_targets[1], Some(solve::scalar_slot_y(2))); +} + #[test] fn implicit_rhs_reports_solver_sized_buffer_overflow() -> Result<(), super::LowerError> { let err = match super::build_implicit_rhs_rows( @@ -751,6 +966,35 @@ fn algebraic_projection_loop_keeps_explicit_row_targets() -> Result<(), LowerErr Ok(()) } +#[test] +fn algebraic_projection_scalar_keeps_distinct_explicit_row_target() -> Result<(), LowerError> { + let projection_incidence = ProjectionIncidence { + incidence: Incidence::new( + vec![BTreeSet::from([0, 1]).into_iter().collect()], + vec![EquationRef(7)], + vec![UnknownId::SolverY(9), UnknownId::SolverY(10)], + ), + unknown_y_indices: vec![9, 10], + }; + let mut row_targets = vec![None; 8]; + row_targets[7] = Some(solve::scalar_slot_y(9)); + + let blocks = super::lower_blt_projection_blocks( + &[BltBlock::Scalar { + equation: EquationRef(7), + unknown: UnknownId::SolverY(10), + }], + &row_targets, + &projection_incidence, + solve_test_span(), + )?; + + assert_eq!(blocks.len(), 1); + assert_eq!(blocks[0].rows, vec![7]); + assert_eq!(blocks[0].y_indices, vec![9, 10]); + Ok(()) +} + #[test] fn algebraic_projection_plan_merges_blocks_that_share_row_targets() -> Result<(), LowerError> { let blocks = super::merge_overlapping_projection_blocks( @@ -788,6 +1032,38 @@ fn algebraic_projection_plan_merges_blocks_that_share_row_targets() -> Result<() Ok(()) } +#[test] +fn initial_row_targets_map_derivative_array_rows_to_state_slots() -> Result<(), LowerError> { + let mut dae_model = dae::Dae::default(); + dae_model + .variables + .states + .insert(rumoca_core::VarName::new("x"), array_var("x", &[2])); + let equation = dae::Equation::residual( + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Sub, + lhs: Box::new(der(var("x"))), + rhs: Box::new(rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Zeros, + args: vec![int_expr(2)], + span: solve_test_span(), + }), + span: solve_test_span(), + }, + solve_test_span(), + "derivative array initialization", + ); + let layout = layout::build_var_layout_with_solver_len(&dae_model, 2) + .expect("test DAE layout should build"); + let row_targets = lower_continuous_row_targets_for_equation(&dae_model, &equation, &layout, 2)?; + + assert_eq!( + row_targets, + vec![Some(solve::scalar_slot_y(0)), Some(solve::scalar_slot_y(1))] + ); + Ok(()) +} + fn var(name: &str) -> rumoca_core::Expression { rumoca_core::Expression::VarRef { name: rumoca_core::Reference::from_component_reference(source_component_ref_from_name( @@ -1070,6 +1346,7 @@ fn solve_appendix_b_validation_rejects_invalid_linsolve_operand_range() { rhs_start: 4, n: 2, next_reg: 5, + output_indices: Vec::new(), metadata: solve::TensorNodeMetadata::default(), span: solve_test_span(), }], @@ -1678,7 +1955,7 @@ fn solve_problem_preserves_scalar_program_source_spans() { } #[test] -fn solve_problem_zero_fills_missing_non_state_implicit_rows_for_templates() { +fn solve_problem_keeps_missing_non_state_implicit_rows_as_y_identity_for_templates() { let mut dae_model = dae::Dae::default(); dae_model .variables @@ -1707,7 +1984,13 @@ fn solve_problem_zero_fills_missing_non_state_implicit_rows_for_templates() { assert_eq!(problem.solve_layout.solver_scalar_count(), 2); let rhs = scalar_program_block_fixture(&problem.continuous.implicit_rhs); assert_eq!(rhs.programs.len(), 2); - assert_eq!(rhs.programs[1], zero_rhs_row()); + assert_eq!( + rhs.programs[1], + vec![ + solve::LinearOp::LoadY { dst: 0, index: 1 }, + solve::LinearOp::StoreOutput { src: 0 }, + ] + ); } #[test] diff --git a/crates/rumoca-phase-solve/src/tests/direct_scalar_residual_projection.rs b/crates/rumoca-phase-solve/src/tests/direct_scalar_residual_projection.rs new file mode 100644 index 000000000..34f04bf9f --- /dev/null +++ b/crates/rumoca-phase-solve/src/tests/direct_scalar_residual_projection.rs @@ -0,0 +1,168 @@ +use super::*; + +fn projected_pair_function() -> rumoca_core::Function { + let span = solve_test_span(); + let mut function = rumoca_core::Function::new("Test.projectedPair", span); + function + .inputs + .push(rumoca_core::FunctionParam::new("u", "Real", span)); + function + .outputs + .push(rumoca_core::FunctionParam::new("force", "Real", span).with_dims(vec![2])); + function.body.push(rumoca_core::Statement::Assignment { + comp: test_component_ref_from_name("force"), + value: rumoca_core::Expression::Array { + elements: vec![ + var("u"), + binary(rumoca_core::OpBinary::Add, var("u"), int_expr(1)), + ], + is_matrix: false, + span, + }, + span, + }); + function +} + +fn projected_pair_lane(index: i64) -> rumoca_core::Expression { + let span = solve_test_span(); + rumoca_core::Expression::Index { + base: Box::new(function_call("Test.projectedPair", vec![source_var("u")])), + subscripts: vec![rumoca_core::Subscript::index(index, span)], + span, + } +} + +fn projected_scalar_equation(owner: &str) -> dae::Equation { + dae::Equation { + lhs: Some(source_ref(owner)), + rhs: binary( + rumoca_core::OpBinary::Add, + source_var("u"), + projected_pair_lane(2), + ), + span: solve_test_span(), + origin: format!("projected scalar owner {owner}"), + scalar_count: 1, + } +} + +fn projected_scalar_dae(owner: &str) -> dae::Dae { + let mut dae_model = dae::Dae::default(); + dae_model.variables.algebraics.insert( + rumoca_core::VarName::new("y"), + source_array_var("y", &[2, 2]), + ); + dae_model + .variables + .parameters + .insert(rumoca_core::VarName::new("u"), source_scalar_var("u")); + dae_model.symbols.functions.insert( + rumoca_core::VarName::new("Test.projectedPair"), + projected_pair_function(), + ); + dae_model + .continuous + .equations + .push(projected_scalar_equation(owner)); + dae_model +} + +#[test] +fn solve_problem_uses_direct_scalar_owner_for_projected_function_residual() { + let dae_model = projected_scalar_dae("y[1,1]"); + + let problem = lower_solve_problem_with_solver_len(&dae_model, 4) + .expect("a rendered scalar owner with a direct layout binding should lower normally"); + let residual = scalar_program_block_fixture(&problem.continuous.residual); + let [row] = residual.programs.as_slice() else { + panic!("expected exactly one explicit residual row"); + }; + let owner = problem + .layout + .binding("y[1,1]") + .expect("the finite layout should bind the rendered scalar owner"); + + assert_eq!(owner, solve::scalar_slot_y(0)); + assert!( + row.iter() + .any(|op| matches!(op, solve::LinearOp::LoadY { index: 0, .. })) + ); + assert!( + row.iter() + .filter(|op| matches!( + op, + solve::LinearOp::Binary { + op: solve::BinaryOp::Add, + .. + } + )) + .count() + >= 2, + "the selected function-output computation and surrounding addition must remain present" + ); + assert_eq!(problem.continuous.implicit_row_targets[0], Some(owner)); +} + +#[test] +fn solve_problem_keeps_aggregate_projected_function_residual_projection() { + let mut dae_model = projected_scalar_dae("y[1,1]"); + dae_model.variables.algebraics.clear(); + dae_model + .variables + .algebraics + .insert(rumoca_core::VarName::new("y"), source_array_var("y", &[2])); + dae_model.continuous.equations.clear(); + dae_model + .continuous + .equations + .push(dae::Equation::explicit_with_scalar_count( + "y", + projected_pair_call(), + solve_test_span(), + "aggregate projected function residual", + 2, + )); + + let problem = lower_solve_problem_with_solver_len(&dae_model, 2) + .expect("aggregate array owners must continue through function projection"); + let residual = scalar_program_block_fixture(&problem.continuous.residual); + + assert_eq!(residual.programs.len(), 2); + assert_eq!( + problem.continuous.implicit_row_targets, + vec![ + problem.layout.binding("y[1]"), + problem.layout.binding("y[2]") + ] + ); +} + +fn projected_pair_call() -> rumoca_core::Expression { + function_call("Test.projectedPair", vec![source_var("u")]) +} + +#[test] +fn solve_problem_rejects_unbound_projected_scalar_owners() { + for (owner, solver_len) in [("missing", 4), ("y[2,2]", 3)] { + let dae_model = projected_scalar_dae(owner); + let error = lower_solve_problem_with_solver_len(&dae_model, solver_len) + .expect_err("missing and truncated owners must fail closed"); + + assert!( + error.is_missing_binding(), + "unexpected error for {owner}: {error}" + ); + assert!(error.reason().contains(owner), "unexpected error: {error}"); + } + + let dae_model = projected_scalar_dae("y[3,1]"); + let error = lower_solve_problem_with_solver_len(&dae_model, 4) + .expect_err("an out-of-range owner must fail closed"); + + assert_eq!(error.source_span(), Some(solve_test_span())); + assert!( + error.reason().contains("outside dimension bounds"), + "unexpected error: {error}" + ); +} diff --git a/crates/rumoca-phase-solve/src/tests/gpu_preparation.rs b/crates/rumoca-phase-solve/src/tests/gpu_preparation.rs index 135b0d2a2..fccbcebee 100644 --- a/crates/rumoca-phase-solve/src/tests/gpu_preparation.rs +++ b/crates/rumoca-phase-solve/src/tests/gpu_preparation.rs @@ -1,5 +1,112 @@ use super::*; +fn gpu_indexed_var(name: &str, index: i64, span: rumoca_core::Span) -> rumoca_core::Expression { + gpu_indexed_var_at(name, &[index], span) +} + +fn gpu_indexed_var_at( + name: &str, + indices: &[i64], + span: rumoca_core::Span, +) -> rumoca_core::Expression { + rumoca_core::Expression::VarRef { + name: source_ref(name), + subscripts: indices + .iter() + .map(|index| rumoca_core::Subscript::generated_index(*index, span)) + .collect(), + span, + } +} + +fn gpu_initial_family_fixture(values: &[i64], spans: &[rumoca_core::Span]) -> dae::Dae { + let indices = (1..=values.len() as i64) + .map(|index| vec![index]) + .collect::>(); + gpu_initial_family_fixture_at(values, spans, &[values.len() as i64], indices) +} + +fn gpu_initial_family_fixture_at( + values: &[i64], + spans: &[rumoca_core::Span], + shape: &[i64], + indices: Vec>, +) -> dae::Dae { + assert_eq!(values.len(), spans.len()); + assert_eq!(values.len(), indices.len()); + let mut dae_model = dae::Dae::default(); + dae_model + .variables + .states + .insert(rumoca_core::VarName::new("x"), array_var("x", shape)); + for ((&value, &span), indices) in values.iter().zip(spans).zip(&indices) { + dae_model.continuous.equations.push(dae::Equation::residual( + binary( + rumoca_core::OpBinary::Sub, + der(gpu_indexed_var_at("x", indices, span)), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(0), + span, + }, + ), + span, + "derivative row", + )); + dae_model + .initialization + .equations + .push(dae::Equation::residual( + binary( + rumoca_core::OpBinary::Sub, + gpu_indexed_var_at("x", indices, span), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + span, + }, + ), + span, + "structured initial row", + )); + dae_model + .initialization + .equation_provenance + .push(dae::InitializationEquationProvenance::User); + } + let template = dae_model.initialization.equations[0].rhs.clone(); + let binders = shape + .iter() + .enumerate() + .map(|(dimension, upper)| rumoca_core::StructuredIndexBinder { + id: dimension, + display_name: format!("i{dimension}"), + lower: 1, + upper: *upper, + step: 1, + }) + .collect(); + dae_model + .initialization + .structured_equations + .push(dae::StructuredEquationFamily { + domain: rumoca_core::StructuredIndexDomain { binders }, + first_equation_index: 0, + equation_counts: vec![1; values.len()], + span: spans[0], + origin: "structured initial fixture".to_string(), + regular: Some(rumoca_core::RegularForFamily { + binders: (0..shape.len()) + .map(|dimension| format!("i{dimension}")) + .collect(), + accesses: Vec::new(), + }), + template: Some(rumoca_core::ComprehensionTemplate { + body: vec![template], + }), + interiors_materialized: true, + }); + dae_model +} + #[test] fn gpu_preparation_inlines_input_driven_algebraic_in_derivative_rhs() { let span = solve_test_span(); @@ -76,3 +183,586 @@ fn gpu_preparation_inlines_input_driven_algebraic_in_derivative_rhs() { gpu_rhs.programs[0] ); } + +#[test] +fn gpu_preparation_rejects_nonstructured_initial_assignment_shape() { + let span = solve_test_span(); + let mut dae_model = dae::Dae::default(); + dae_model + .variables + .states + .insert(rumoca_core::VarName::new("x"), scalar_var("x")); + dae_model.continuous.equations.push(dae::Equation::residual( + binary(rumoca_core::OpBinary::Sub, der(var("x")), int_expr(0)), + span, + "der(x) = 0", + )); + dae_model + .initialization + .equations + .push(dae::Equation::residual( + binary(rumoca_core::OpBinary::Sub, var("x"), int_expr(7)), + span, + "x = 7", + )); + + let gpu = lower_solve_problem_with_solver_len_and_model_span_and_profile( + &dae_model, + 1, + Some(span), + SolveProblemLoweringProfile::GpuPreparation, + ) + .expect_err("GPU preparation must fail closed instead of scalarizing initialization rows"); + assert!(matches!( + &gpu, + crate::lower::LowerError::UnsupportedAt { .. } + )); + assert_eq!(gpu.source_span(), Some(span)); +} + +#[test] +fn gpu_preparation_ignores_automatic_fixed_start_rows() { + let span = solve_test_span(); + let mut dae_model = dae::Dae::default(); + dae_model + .variables + .states + .insert(rumoca_core::VarName::new("x"), scalar_var("x")); + dae_model.continuous.equations.push(dae::Equation::residual( + binary(rumoca_core::OpBinary::Sub, der(var("x")), int_expr(0)), + span, + "der(x) = 0", + )); + dae_model + .initialization + .equations + .push(dae::Equation::residual( + binary(rumoca_core::OpBinary::Sub, var("x"), int_expr(7)), + span, + "fixed start initialization for x", + )); + dae_model + .initialization + .equation_provenance + .push(dae::InitializationEquationProvenance::FixedStart); + + let gpu = lower_solve_problem_with_solver_len_and_model_span_and_profile( + &dae_model, + 1, + Some(span), + SolveProblemLoweringProfile::GpuPreparation, + ) + .expect("GPU preparation should retain declared fixed starts without scalar initialization"); + assert!(gpu.initialization.residual.is_empty()); + assert!(gpu.initialization.direct_families.is_empty()); + assert!(gpu.initialization.row_targets.is_empty()); +} + +#[test] +fn gpu_preparation_rejects_partial_fixed_start_target_coverage() { + let span = solve_test_span(); + let mut dae_model = dae::Dae::default(); + for name in ["x", "y"] { + dae_model + .variables + .states + .insert(rumoca_core::VarName::new(name), scalar_var(name)); + dae_model.continuous.equations.push(dae::Equation::residual( + binary(rumoca_core::OpBinary::Sub, der(var(name)), int_expr(0)), + span, + "derivative", + )); + } + dae_model + .initialization + .equations + .push(dae::Equation::residual( + binary(rumoca_core::OpBinary::Sub, var("x"), int_expr(7)), + span, + "fixed start initialization for x", + )); + dae_model + .initialization + .equation_provenance + .push(dae::InitializationEquationProvenance::FixedStart); + + let error = lower_solve_problem_with_solver_len_and_model_span_and_profile( + &dae_model, + 2, + Some(span), + SolveProblemLoweringProfile::GpuPreparation, + ) + .expect_err("partial GPU initialization coverage must fail closed"); + assert!(error.to_string().contains("cover every solver Y slot")); +} + +#[test] +fn gpu_preparation_rejects_overlapping_fixed_start_targets_at_conflicting_span() { + let first_span = solve_numbered_span(301, 10, 20); + let conflicting_span = solve_numbered_span(301, 30, 40); + let mut dae_model = dae::Dae::default(); + dae_model + .variables + .states + .insert(rumoca_core::VarName::new("x"), scalar_var("x")); + dae_model.continuous.equations.push(dae::Equation::residual( + binary(rumoca_core::OpBinary::Sub, der(var("x")), int_expr(0)), + first_span, + "derivative", + )); + for span in [first_span, conflicting_span] { + dae_model + .initialization + .equations + .push(dae::Equation::residual( + binary(rumoca_core::OpBinary::Sub, var("x"), int_expr(7)), + span, + "fixed start initialization for x", + )); + dae_model + .initialization + .equation_provenance + .push(dae::InitializationEquationProvenance::FixedStart); + } + + let error = lower_solve_problem_with_solver_len_and_model_span_and_profile( + &dae_model, + 1, + Some(first_span), + SolveProblemLoweringProfile::GpuPreparation, + ) + .expect_err("overlapping fixed-start ownership must fail closed"); + assert!(error.to_string().contains("overlap")); + assert_eq!(error.source_span(), Some(conflicting_span)); +} + +#[test] +fn gpu_preparation_emits_one_compact_fixed_range_for_array_target() { + let span = solve_numbered_span(302, 10, 20); + let mut dae_model = dae::Dae::default(); + dae_model + .variables + .states + .insert(rumoca_core::VarName::new("x"), array_var("x", &[128])); + let mut derivative = dae::Equation::residual( + binary(rumoca_core::OpBinary::Sub, der(var("x")), int_expr(0)), + span, + "array derivative", + ); + derivative.scalar_count = 128; + dae_model.continuous.equations.push(derivative); + let mut fixed = dae::Equation::residual( + binary(rumoca_core::OpBinary::Sub, var("x"), int_expr(7)), + span, + "fixed array start", + ); + fixed.scalar_count = 128; + dae_model.initialization.equations.push(fixed); + dae_model + .initialization + .equation_provenance + .push(dae::InitializationEquationProvenance::FixedStart); + + let gpu = lower_solve_problem_with_solver_len_and_model_span_and_profile( + &dae_model, + 128, + Some(span), + SolveProblemLoweringProfile::GpuPreparation, + ) + .expect("fixed array target should lower as one affine range"); + assert_eq!(gpu.initialization.fixed_target_ranges.len(), 1); + assert_eq!(gpu.initialization.fixed_target_ranges[0].start, 0); + assert_eq!(gpu.initialization.fixed_target_ranges[0].end, 128); +} + +#[test] +fn gpu_phase_range_validation_rejects_direct_fixed_overlap_at_later_span() { + let first_span = solve_numbered_span(303, 10, 20); + let conflicting_span = solve_numbered_span(303, 30, 40); + let error = crate::gpu_initialization::normalize_gpu_target_ranges( + vec![ + solve::InitializationTargetRange { + start: 0, + end: 2, + span: Some(first_span), + }, + solve::InitializationTargetRange { + start: 1, + end: 3, + span: Some(conflicting_span), + }, + ], + 3, + ) + .expect_err("phase validation must reject direct/fixed ownership overlap"); + assert!(error.to_string().contains("overlap")); + assert_eq!(error.source_span(), Some(conflicting_span)); +} + +#[test] +fn gpu_phase_range_validation_merges_adjacency_only() { + let span = solve_numbered_span(304, 10, 20); + let normalized = crate::gpu_initialization::normalize_gpu_target_ranges( + vec![ + solve::InitializationTargetRange { + start: 0, + end: 1, + span: Some(span), + }, + solve::InitializationTargetRange { + start: 1, + end: 3, + span: Some(span), + }, + ], + 3, + ) + .expect("adjacent phase ranges should merge"); + assert_eq!(normalized.len(), 1); + assert_eq!((normalized[0].start, normalized[0].end), (0, 3)); +} + +#[test] +fn gpu_corner_cell_index_rejects_request_for_singleton_dimension() { + let domain = rumoca_core::StructuredIndexDomain { + binders: vec![rumoca_core::StructuredIndexBinder { + id: 0, + display_name: "i".to_string(), + lower: 1, + upper: 1, + step: 1, + }], + }; + + let error = gpu_corner_cell_index(&domain, 0, solve_test_span()) + .expect_err("direct GPU initial projection must fail closed for one-cell binders"); + assert!(matches!( + error, + crate::lower::LowerError::UnsupportedAt { .. } + )); + assert!( + error + .to_string() + .contains("non-degenerate structured binder") + ); +} + +#[test] +fn gpu_preparation_accepts_mixed_singleton_structured_domain() { + let spans = [ + solve_numbered_span(310, 10, 20), + solve_numbered_span(310, 30, 40), + solve_numbered_span(310, 50, 60), + ]; + let dae_model = gpu_initial_family_fixture_at( + &[1, 2, 3], + &spans, + &[1, 3], + vec![vec![1, 1], vec![1, 2], vec![1, 3]], + ); + + let gpu = lower_solve_problem_with_solver_len_and_model_span_and_profile( + &dae_model, + 3, + Some(spans[0]), + SolveProblemLoweringProfile::GpuPreparation, + ) + .expect("singleton axis has no affine degree of freedom"); + assert!(gpu.initialization.row_targets.is_empty()); + assert_eq!(gpu.initialization.direct_families.len(), 1); + assert_eq!(gpu.initialization.required_target_ranges.len(), 1); + assert_eq!( + ( + gpu.initialization.required_target_ranges[0].start, + gpu.initialization.required_target_ranges[0].end, + ), + (0, 3) + ); + let solve::ComputeNode::Map { + domain, + load_strides, + const_strides, + .. + } = &gpu.initialization.residual.nodes[0] + else { + panic!("structured initialization should remain a compact Map") + }; + assert_eq!(domain.binders.len(), 2); + assert_eq!(domain.scalar_count(), Ok(3)); + assert_eq!( + load_strides + .iter() + .flat_map(|stride| &stride.terms) + .map(|term| term.dimension) + .collect::>(), + vec![1] + ); + assert_eq!( + const_strides + .iter() + .flat_map(|stride| &stride.terms) + .map(|term| term.dimension) + .collect::>(), + vec![1] + ); + let targets = &gpu.initialization.direct_families[0].targets; + assert_eq!( + targets, + &solve::TensorOutputMap::dense_contiguous(0, domain).unwrap() + ); + assert_eq!(targets.strides.len(), 1); + assert_eq!(targets.strides[0].dimension, 1); + assert_eq!(targets.strides[0].stride, 1); +} + +#[test] +fn gpu_preparation_accepts_all_singleton_structured_domain() { + let span = solve_numbered_span(311, 10, 20); + let dae_model = gpu_initial_family_fixture_at(&[7], &[span], &[1, 1], vec![vec![1, 1]]); + + let gpu = lower_solve_problem_with_solver_len_and_model_span_and_profile( + &dae_model, + 1, + Some(span), + SolveProblemLoweringProfile::GpuPreparation, + ) + .expect("all-singleton domain is one valid family cell"); + assert!(gpu.initialization.row_targets.is_empty()); + assert_eq!(gpu.initialization.direct_families.len(), 1); + assert_eq!(gpu.initialization.required_target_ranges.len(), 1); + assert_eq!( + ( + gpu.initialization.required_target_ranges[0].start, + gpu.initialization.required_target_ranges[0].end, + ), + (0, 1) + ); + let solve::ComputeNode::Map { + domain, + load_strides, + const_strides, + .. + } = &gpu.initialization.residual.nodes[0] + else { + panic!("structured initialization should remain a compact Map") + }; + assert_eq!(domain.binders.len(), 2); + assert_eq!(domain.scalar_count(), Ok(1)); + assert!(load_strides.is_empty()); + assert!(const_strides.is_empty()); + assert!( + gpu.initialization.direct_families[0] + .targets + .strides + .is_empty() + ); +} + +#[test] +fn gpu_initial_projection_handles_negative_binder_steps() { + // Source binder traversal may descend. Target maps remain canonical dense + // positive-stride maps; Solve-IR validation rejects negative target strides. + let domain = rumoca_core::StructuredIndexDomain { + binders: vec![rumoca_core::StructuredIndexBinder { + id: 0, + display_name: "i".to_string(), + lower: 3, + upper: 1, + step: -1, + }], + }; + + assert_eq!(gpu_corner_cell_index(&domain, 0, solve_test_span()), Ok(1)); +} + +#[test] +fn gpu_initial_uniformity_checks_destinations_nonloads_and_load_p() { + let span = solve_test_span(); + let base = vec![ + solve::LinearOp::LoadP { dst: 0, index: 4 }, + solve::LinearOp::Const { dst: 1, value: 2.0 }, + solve::LinearOp::Binary { + dst: 2, + op: solve::BinaryOp::Add, + lhs: 0, + rhs: 1, + }, + solve::LinearOp::StoreOutput { src: 2 }, + ]; + let mut corner = base.clone(); + corner[0] = solve::LinearOp::LoadP { dst: 0, index: 7 }; + let mut loads = Vec::new(); + let mut constants = Vec::new(); + append_gpu_corner_strides(&base, &corner, 0, &mut loads, &mut constants, span) + .expect("LoadP may vary affinely"); + assert_eq!(loads[0].terms[0].stride, 3); + + corner[2] = solve::LinearOp::Binary { + dst: 3, + op: solve::BinaryOp::Add, + lhs: 0, + rhs: 1, + }; + append_gpu_corner_strides(&base, &corner, 0, &mut Vec::new(), &mut Vec::new(), span) + .expect_err("destination register drift must fail closed"); +} + +#[test] +fn gpu_initial_projection_rejects_nonaffine_three_cell_constants_at_first_bad_row() { + let spans = [ + solve_numbered_span(305, 10, 20), + solve_numbered_span(305, 30, 40), + solve_numbered_span(305, 50, 60), + ]; + let dae_model = gpu_initial_family_fixture(&[1, 4, 9], &spans); + + let error = lower_solve_problem_with_solver_len_and_model_span_and_profile( + &dae_model, + 3, + Some(spans[0]), + SolveProblemLoweringProfile::GpuPreparation, + ) + .expect_err("non-affine constants must not be synthesized as 1, 4, 7"); + + assert!(error.to_string().contains("affine"), "{error}"); + assert_eq!(error.source_span(), Some(spans[2])); +} + +#[test] +fn gpu_initial_metadata_failures_preserve_family_and_equation_spans() { + let spans = [ + solve_numbered_span(306, 10, 20), + solve_numbered_span(306, 30, 40), + solve_numbered_span(306, 50, 60), + ]; + for missing_regular in [true, false] { + let mut dae_model = gpu_initial_family_fixture(&[1, 2, 3], &spans); + let family = &mut dae_model.initialization.structured_equations[0]; + if missing_regular { + family.regular = None; + } else { + family.template = None; + } + let error = lower_solve_problem_with_solver_len_and_model_span_and_profile( + &dae_model, + 3, + Some(spans[0]), + SolveProblemLoweringProfile::GpuPreparation, + ) + .expect_err("missing structured metadata must fail closed"); + assert_eq!(error.source_span(), Some(spans[0])); + } + + let mut dae_model = gpu_initial_family_fixture(&[1, 2, 3], &spans); + dae_model.initialization.equation_provenance.pop(); + let error = lower_solve_problem_with_solver_len_and_model_span_and_profile( + &dae_model, + 3, + Some(spans[0]), + SolveProblemLoweringProfile::GpuPreparation, + ) + .expect_err("missing equation provenance must fail closed"); + assert_eq!(error.source_span(), Some(spans[2])); +} + +#[test] +fn gpu_initial_coverage_failure_points_to_first_uncovered_user_equation() { + let spans = [ + solve_numbered_span(307, 10, 20), + solve_numbered_span(307, 30, 40), + solve_numbered_span(307, 50, 60), + ]; + let uncovered_span = solve_numbered_span(307, 70, 80); + let mut dae_model = gpu_initial_family_fixture(&[1, 2, 3], &spans); + dae_model + .initialization + .equations + .push(dae::Equation::residual( + binary( + rumoca_core::OpBinary::Sub, + gpu_indexed_var("x", 3, uncovered_span), + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(4), + span: uncovered_span, + }, + ), + uncovered_span, + "uncovered initial row", + )); + dae_model + .initialization + .equation_provenance + .push(dae::InitializationEquationProvenance::User); + + let error = lower_solve_problem_with_solver_len_and_model_span_and_profile( + &dae_model, + 3, + Some(spans[0]), + SolveProblemLoweringProfile::GpuPreparation, + ) + .expect_err("mixed structured and scalar initial rows must fail closed"); + assert_eq!(error.source_span(), Some(uncovered_span)); +} + +#[test] +fn gpu_initial_direct_and_body_shape_failures_preserve_user_spans() { + let spans = [ + solve_numbered_span(308, 10, 20), + solve_numbered_span(308, 30, 40), + solve_numbered_span(308, 50, 60), + ]; + let mut nondirect = gpu_initial_family_fixture(&[1, 2, 3], &spans); + nondirect.initialization.equations[0].rhs = binary( + rumoca_core::OpBinary::Add, + gpu_indexed_var("x", 1, spans[0]), + int_expr(1), + ); + let error = lower_solve_problem_with_solver_len_and_model_span_and_profile( + &nondirect, + 3, + Some(spans[0]), + SolveProblemLoweringProfile::GpuPreparation, + ) + .expect_err("non-direct base row must fail closed"); + assert_eq!(error.source_span(), Some(spans[0])); + + let mut mismatched = gpu_initial_family_fixture(&[1, 2, 3], &spans); + mismatched.initialization.equations[2].rhs = binary( + rumoca_core::OpBinary::Sub, + gpu_indexed_var("x", 3, spans[2]), + binary(rumoca_core::OpBinary::Add, int_expr(4), int_expr(5)), + ); + let error = lower_solve_problem_with_solver_len_and_model_span_and_profile( + &mismatched, + 3, + Some(spans[0]), + SolveProblemLoweringProfile::GpuPreparation, + ) + .expect_err("body-shape mismatch must fail closed"); + assert_eq!(error.source_span(), Some(spans[2])); + + let mut corner = vec![solve::LinearOp::Const { dst: 0, value: 1.0 }]; + corner.push(solve::LinearOp::StoreOutput { src: 0 }); + let error = append_gpu_corner_strides( + &[solve::LinearOp::Const { dst: 0, value: 1.0 }], + &corner, + 0, + &mut Vec::new(), + &mut Vec::new(), + spans[1], + ) + .expect_err("operation-shape mismatch must fail closed"); + assert_eq!(error.source_span(), Some(spans[1])); +} + +#[test] +fn gpu_initial_lowering_rejects_random_operations_with_source_span() { + let span = solve_numbered_span(309, 10, 20); + let error = reject_nondeterministic_gpu_initial_ops( + &[solve::LinearOp::ImpureRandomInit { dst: 1, seed: 0 }], + span, + ) + .expect_err("GPU initialization lowering must reject non-replayable operations"); + assert!(error.to_string().contains("random or impure")); + assert_eq!(error.source_span(), Some(span)); +} diff --git a/crates/rumoca-phase-solve/src/tests/observation_aliases.rs b/crates/rumoca-phase-solve/src/tests/observation_aliases.rs index 093458b43..9d272d02e 100644 --- a/crates/rumoca-phase-solve/src/tests/observation_aliases.rs +++ b/crates/rumoca-phase-solve/src/tests/observation_aliases.rs @@ -190,6 +190,115 @@ fn solve_problem_marks_event_relation_aliases_for_observation_refresh() { assert!(!refresh_for("sample.y")); } +#[test] +fn solve_problem_marks_lowered_change_pulse_aliases_for_observation_refresh() { + let mut dae_model = dae::Dae::default(); + for name in ["source", "changed", "trigger"] { + dae_model + .variables + .discrete_valued + .insert(rumoca_core::VarName::new(name), scalar_var(name)); + } + insert_pre_parameter(&mut dae_model, "source"); + dae_model + .clocks + .intervals + .insert("changed".to_string(), 0.5); + dae_model + .clocks + .intervals + .insert("trigger".to_string(), 0.5); + dae_model + .discrete + .valued_updates + .push(dae::Equation::explicit( + source_ref("changed"), + binary(rumoca_core::OpBinary::Neq, var("source"), pre_var("source")), + test_span(), + "changed = source <> pre(source)", + )); + dae_model + .discrete + .valued_updates + .push(dae::Equation::explicit( + source_ref("trigger"), + var("changed"), + test_span(), + "trigger = changed", + )); + + let problem = lower_solve_problem(&dae_model) + .expect("lowered change pulse and its observation alias should lower"); + + // `change(source)` is lowered to `source <> pre(source)`. Its result is an + // event-instant pulse, so observation settling must clear both the pulse + // and aliases after the event-entry history has been committed. + assert_eq!(problem.discrete.observation_refresh, vec![true, true]); +} + +#[test] +fn solve_problem_only_marks_matching_direct_lowered_change_relations() { + let mut dae_model = dae::Dae::default(); + for name in [ + "source", + "other", + "reversed_change", + "mismatched_pre", + "suppressed_change", + ] { + dae_model + .variables + .discrete_valued + .insert(rumoca_core::VarName::new(name), scalar_var(name)); + } + insert_pre_parameter(&mut dae_model, "source"); + insert_pre_parameter(&mut dae_model, "other"); + for name in ["reversed_change", "mismatched_pre", "suppressed_change"] { + dae_model.clocks.intervals.insert(name.to_string(), 0.5); + } + dae_model + .discrete + .valued_updates + .push(dae::Equation::explicit( + source_ref("reversed_change"), + binary(rumoca_core::OpBinary::Neq, pre_var("source"), var("source")), + test_span(), + "reversed_change = pre(source) <> source", + )); + dae_model + .discrete + .valued_updates + .push(dae::Equation::explicit( + source_ref("mismatched_pre"), + binary(rumoca_core::OpBinary::Neq, var("source"), pre_var("other")), + test_span(), + "mismatched_pre = source <> pre(other)", + )); + dae_model.discrete.valued_updates.push(dae::Equation { + lhs: Some(source_ref("suppressed_change")), + rhs: rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::NoEvent, + args: vec![binary( + rumoca_core::OpBinary::Neq, + var("source"), + pre_var("source"), + )], + span: test_span(), + }, + span: test_span(), + origin: "suppressed_change = noEvent(source <> pre(source))".to_string(), + scalar_count: 1, + }); + + let problem = lower_solve_problem(&dae_model) + .expect("direct lowered change relation classifications should lower"); + + assert_eq!( + problem.discrete.observation_refresh, + vec![true, false, false] + ); +} + #[test] fn solve_problem_does_not_observation_refresh_no_event_relation_aliases() { let mut dae_model = dae::Dae::default(); diff --git a/crates/rumoca-phase-solve/src/tests/projection_loop_matching.rs b/crates/rumoca-phase-solve/src/tests/projection_loop_matching.rs new file mode 100644 index 000000000..a2295f56f --- /dev/null +++ b/crates/rumoca-phase-solve/src/tests/projection_loop_matching.rs @@ -0,0 +1,59 @@ +use super::*; + +#[test] +fn algebraic_projection_loop_does_not_promote_structural_pairs_to_causal_steps() +-> Result<(), LowerError> { + let projection_incidence = ProjectionIncidence { + incidence: Incidence::new( + vec![ + BTreeSet::from([0, 1]).into_iter().collect(), + BTreeSet::from([0, 1]).into_iter().collect(), + ], + vec![EquationRef(3), EquationRef(4)], + vec![UnknownId::SolverY(20), UnknownId::SolverY(21)], + ), + unknown_y_indices: vec![20, 21], + }; + + let block = super::super::lower_algebraic_loop_projection_block( + &[EquationRef(3), EquationRef(4)], + &[UnknownId::SolverY(21), UnknownId::SolverY(20)], + &[], + &projection_incidence, + solve_test_span(), + )? + .expect("loop block should lower"); + + assert_eq!(block.rows, vec![3, 4]); + assert_eq!(block.y_indices, vec![20, 21]); + assert!(block.causal_steps.is_empty()); + Ok(()) +} + +#[test] +fn algebraic_projection_loop_keeps_assignment_targets_in_simultaneous_solve() +-> Result<(), LowerError> { + let projection_incidence = ProjectionIncidence { + incidence: Incidence::new( + vec![BTreeSet::from([0, 1]).into_iter().collect()], + vec![EquationRef(3)], + vec![UnknownId::SolverY(20), UnknownId::SolverY(21)], + ), + unknown_y_indices: vec![20, 21], + }; + let mut row_targets = vec![None; 4]; + row_targets[3] = Some(solve::scalar_slot_y(20)); + + let block = super::super::lower_algebraic_loop_projection_block( + &[EquationRef(3)], + &[UnknownId::SolverY(21)], + &row_targets, + &projection_incidence, + solve_test_span(), + )? + .expect("loop block should lower"); + + assert_eq!(block.y_indices, vec![20, 21]); + assert!(block.causal_steps.is_empty()); + Ok(()) +} diff --git a/crates/rumoca-phase-solve/src/tests/projection_plan_more.rs b/crates/rumoca-phase-solve/src/tests/projection_plan_more.rs index 9fc63950e..f93df7d4a 100644 --- a/crates/rumoca-phase-solve/src/tests/projection_plan_more.rs +++ b/crates/rumoca-phase-solve/src/tests/projection_plan_more.rs @@ -8,6 +8,97 @@ fn projection_plan_span() -> rumoca_core::Span { ) } +#[test] +fn solve_problem_projection_plan_covers_every_algebraic_tail_row() { + let span = projection_plan_span(); + let rows = vec![ + vec![ + solve::LinearOp::LoadY { dst: 0, index: 0 }, + solve::LinearOp::StoreOutput { src: 0 }, + ], + vec![ + solve::LinearOp::LoadY { dst: 0, index: 1 }, + solve::LinearOp::StoreOutput { src: 0 }, + ], + vec![ + solve::LinearOp::Const { dst: 0, value: 1.0 }, + solve::LinearOp::StoreOutput { src: 0 }, + ], + ]; + + let err = lower_algebraic_projection_plan(&rows, &[None; 3], 1, 3, span) + .expect_err("constant algebraic tail row must not be omitted from projection coverage"); + + assert!( + err.to_string() + .contains("algebraic projection plan omits implicit row 2") + ); + assert_eq!(err.source_span(), Some(span)); +} + +#[test] +fn solve_problem_leaves_dynamic_consistency_rows_to_full_tail_validation() { + let span = projection_plan_span(); + let rows = vec![ + vec![ + solve::LinearOp::LoadY { dst: 0, index: 0 }, + solve::LinearOp::StoreOutput { src: 0 }, + ], + vec![ + solve::LinearOp::LoadY { dst: 0, index: 1 }, + solve::LinearOp::StoreOutput { src: 0 }, + ], + vec![ + solve::LinearOp::LoadP { dst: 0, index: 0 }, + solve::LinearOp::StoreOutput { src: 0 }, + ], + ]; + + let plan = lower_algebraic_projection_plan(&rows, &[None; 3], 1, 3, span) + .expect("parameter-dependent consistency rows have no algebraic target to project"); + + assert_eq!(plan.blocks.len(), 1); + assert_eq!(plan.blocks[0].rows, vec![1]); + assert_eq!(plan.blocks[0].y_indices, vec![1]); +} + +#[test] +fn state_free_dynamic_residual_keeps_owner_when_fallback_identity_competes() { + let span = projection_plan_span(); + let residual = vec![vec![ + solve::LinearOp::LoadP { dst: 0, index: 0 }, + solve::LinearOp::LoadY { dst: 1, index: 1 }, + solve::LinearOp::Binary { + dst: 2, + op: solve::BinaryOp::Sub, + lhs: 0, + rhs: 1, + }, + solve::LinearOp::StoreOutput { src: 2 }, + ]]; + let implicit = build_implicit_rhs_rows(&[], &residual, &[None], 0, 2, span) + .expect("state-free implicit rows should build"); + + assert_eq!(implicit.residual_to_implicit_rows, vec![Some(1)]); + assert_eq!(implicit.rows[1], residual[0]); + assert_eq!(implicit.row_targets[1], Some(solve::scalar_slot_y(1))); + + let plan = lower_algebraic_projection_plan(&implicit.rows, &implicit.row_targets, 0, 2, span) + .expect("dynamic residual should replace its structurally matched fallback identity"); + let covered_rows = plan + .blocks + .iter() + .flat_map(|block| block.rows.iter().copied()) + .collect::>(); + + assert_eq!(covered_rows, BTreeSet::from([0, 1])); + assert!( + plan.blocks + .iter() + .any(|block| { block.rows == vec![1] && block.y_indices == vec![1] }) + ); +} + #[test] fn solve_problem_preserves_duplicate_target_residual_rows() { let mut dae_model = dae::Dae::default(); @@ -258,6 +349,97 @@ fn initialization_projection_plan_includes_unfixed_states() { assert!(projected.contains(&2), "algebraic z should be projected"); } +#[test] +fn initialization_projection_prefers_explicit_algebraic_row_target_over_state_dependency() { + let mut dae_model = dae::Dae::default(); + let mut x = scalar_var("x"); + x.fixed = Some(false); + dae_model.variables.states.insert("x".into(), x); + dae_model + .variables + .algebraics + .insert("z".into(), scalar_var("z")); + dae_model.continuous.equations.push(dae::Equation::residual( + binary(rumoca_core::OpBinary::Sub, der(var("x")), int_expr(0)), + solve_test_span(), + "derivative row for x", + )); + dae_model.continuous.equations.push(dae::Equation::residual( + binary(rumoca_core::OpBinary::Sub, var("z"), var("x")), + solve_test_span(), + "explicit algebraic target depends on state", + )); + + let problem = lower_solve_problem(&dae_model).expect("initialization plan should lower"); + let projected = problem + .initialization + .projection_plan + .blocks + .iter() + .flat_map(|block| block.y_indices.iter().copied()) + .collect::>(); + + assert!( + !projected.contains(&0), + "state x should remain a dependency when the row explicitly targets z" + ); + assert!(projected.contains(&1), "algebraic z should be projected"); +} + +#[test] +fn initialization_projection_does_not_infer_algebraic_targets_from_continuous_rows() { + let mut dae_model = dae::Dae::default(); + dae_model + .variables + .algebraics + .insert(rumoca_core::VarName::new("u"), scalar_var("u")); + dae_model + .variables + .algebraics + .insert(rumoca_core::VarName::new("z"), scalar_var("z")); + dae_model + .variables + .parameters + .insert(rumoca_core::VarName::new("p"), scalar_var("p")); + dae_model.continuous.equations.push(dae::Equation::residual( + rumoca_core::Expression::BuiltinCall { + function: rumoca_core::BuiltinFunction::Cos, + args: vec![var("u")], + span: solve_test_span(), + }, + solve_test_span(), + "continuous residual references u without an explicit target", + )); + dae_model + .initialization + .equations + .push(dae::Equation::residual( + binary(rumoca_core::OpBinary::Sub, var("z"), int_expr(3)), + solve_test_span(), + "initial equation drives z", + )); + + let problem = lower_solve_problem(&dae_model).expect("initialization plan should lower"); + let u_idx = problem.solve_layout.solver_maps.name_to_idx["u"]; + let z_idx = problem.solve_layout.solver_maps.name_to_idx["z"]; + let projected = problem + .initialization + .projection_plan + .blocks + .iter() + .flat_map(|block| block.y_indices.iter().copied()) + .collect::>(); + + assert!( + !projected.contains(&u_idx), + "continuous-derived rows without explicit targets are handled by continuous algebraic projection" + ); + assert!( + projected.contains(&z_idx), + "initialization-specific algebraic equations still project their algebraic unknowns" + ); +} + #[test] fn solve_problem_records_scaled_algebraic_residual_target() { let mut dae_model = dae::Dae::default(); @@ -524,6 +706,34 @@ fn solve_problem_skips_negated_additive_terms_when_recording_row_targets() { ); } +#[test] +fn solve_problem_records_rhs_target_for_source_minus_algebraic_residual() { + let mut dae_model = dae::Dae::default(); + dae_model + .variables + .parameters + .insert(rumoca_core::VarName::new("source"), scalar_var("source")); + dae_model + .variables + .algebraics + .insert(rumoca_core::VarName::new("u"), scalar_var("u")); + dae_model.continuous.equations.push(dae::Equation::residual( + binary(rumoca_core::OpBinary::Sub, var("source"), var("u")), + solve_test_span(), + "source expression drives algebraic input", + )); + + let problem = lower_solve_problem(&dae_model).expect("source-minus-target should lower"); + + assert!( + problem + .continuous + .implicit_row_targets + .iter() + .any(|target| matches!(target, Some(solve::ScalarSlot::Y { index: 0, .. }))) + ); +} + #[test] fn solve_problem_records_slice_assignment_row_targets() { let mut dae_model = dae::Dae::default(); diff --git a/crates/rumoca-phase-structural/src/dae_prepare/dae_prepare_demotion_tests.rs b/crates/rumoca-phase-structural/src/dae_prepare/dae_prepare_demotion_tests.rs index 7b8df7d51..486144911 100644 --- a/crates/rumoca-phase-structural/src/dae_prepare/dae_prepare_demotion_tests.rs +++ b/crates/rumoca-phase-structural/src/dae_prepare/dae_prepare_demotion_tests.rs @@ -1,3 +1,7 @@ +//! SPEC_0021 file-size exception: demotion tests keep matrix-state, alias, and +//! derivative ownership regressions together. split plan: split by demotion +//! category once the DAE prepare APIs settle. + use super::*; use indexmap::IndexSet; use rumoca_core::Span; @@ -54,6 +58,22 @@ fn var_sub(name: &str, subscript: Expression) -> Expression { } } +fn index_expr(base: Expression, idx: i64) -> Expression { + Expression::Index { + base: Box::new(base), + subscripts: vec![Subscript::generated_index(idx, Span::DUMMY)], + span: Span::DUMMY, + } +} + +fn field_expr(base: Expression, field: &str) -> Expression { + Expression::FieldAccess { + base: Box::new(base), + field: field.to_string(), + span: Span::DUMMY, + } +} + fn int(v: i64) -> Expression { Expression::Literal { value: Literal::Integer(v), @@ -138,6 +158,14 @@ fn no_event(expr: Expression) -> Expression { } } +fn builtin(function: BuiltinFunction, arg: Expression) -> Expression { + Expression::BuiltinCall { + function, + args: vec![arg], + span: rumoca_core::Span::DUMMY, + } +} + fn array(elements: Vec) -> Expression { Expression::Array { elements, @@ -159,6 +187,10 @@ fn call_with_span(name: &str, args: Vec, span: Span) -> Expression { } } +fn function_param(name: &str) -> rumoca_core::FunctionParam { + rumoca_core::FunctionParam::new(name, "Real", Span::DUMMY) +} + fn der(name: &str) -> Expression { Expression::BuiltinCall { function: BuiltinFunction::Der, @@ -476,6 +508,42 @@ fn test_demote_direct_assigned_states_keeps_state_defined_by_non_state_alias() { assert!(dae.variables.states.contains_key(&VarName::new("v"))); } +#[test] +fn test_demote_direct_assigned_states_allows_state_free_algebraic_closure() { + let mut dae = Dae::new(); + dae.variables + .states + .insert(VarName::new("x"), test_variable("x")); + dae.variables + .algebraics + .insert(VarName::new("a"), test_variable("a")); + dae.variables + .algebraics + .insert(VarName::new("v"), test_variable("v")); + + dae.continuous.equations.push(eq(sub(var("x"), var("a")))); + dae.continuous + .equations + .push(eq(sub(var("a"), var("time")))); + dae.continuous.equations.push(eq(sub(der("x"), var("v")))); + + let demoted = demote_direct_assigned_states(&mut dae).expect("direct demotion should succeed"); + + assert_eq!( + demoted, 1, + "a state-free algebraic closure can define a dummy trajectory state" + ); + assert!(!dae.variables.states.contains_key(&VarName::new("x"))); + assert!(dae.variables.algebraics.contains_key(&VarName::new("x"))); + assert!( + dae.continuous + .equations + .iter() + .all(|eq| !expr_contains_der_of(&eq.rhs, &VarName::new("x"))), + "demotion must rewrite derivative uses of the demoted state" + ); +} + #[test] fn test_demote_direct_assigned_states_allows_fixed_connection_alias() { let mut dae = Dae::new(); @@ -588,6 +656,57 @@ fn test_demote_direct_assigned_states_rejects_state_dependent_connection_alias() assert!(dae.variables.states.contains_key(&VarName::new("y"))); } +#[test] +fn test_demote_direct_assigned_input_state_from_indexed_state_signal() { + let mut dae = Dae::new(); + let mut x = test_variable("x"); + x.dims = vec![3]; + dae.variables.states.insert(VarName::new("x"), x); + let mut u1 = test_variable("u1"); + u1.causality = dae::VariableCausality::Input; + dae.variables.states.insert(VarName::new("u1"), u1); + let mut u2 = test_variable("u2"); + u2.causality = dae::VariableCausality::Input; + dae.variables.states.insert(VarName::new("u2"), u2); + let mut u3 = test_variable("u3"); + u3.causality = dae::VariableCausality::Input; + dae.variables.outputs.insert(VarName::new("u3"), u3); + + dae.continuous + .equations + .push(eq(sub(der_idx("x", 3), var("dx3")))); + dae.variables + .algebraics + .insert(VarName::new("dx3"), test_variable("dx3")); + dae.continuous + .equations + .push(eq(sub(var("u1"), var_idx("x", 3)))); + dae.continuous.equations.push(eq(sub(der("u1"), var("u2")))); + dae.continuous.equations.push(eq(sub(der("u2"), var("u3")))); + + let demoted = demote_direct_assigned_states(&mut dae).expect("direct demotion should succeed"); + + assert_eq!(demoted, 2); + assert!(!dae.variables.states.contains_key(&VarName::new("u1"))); + assert!(!dae.variables.states.contains_key(&VarName::new("u2"))); + assert!(dae.variables.algebraics.contains_key(&VarName::new("u1"))); + assert!(dae.variables.algebraics.contains_key(&VarName::new("u2"))); + assert!( + dae.continuous + .equations + .iter() + .all(|eq| !expr_contains_der_of(&eq.rhs, &VarName::new("u1"))), + "demotion should rewrite derivative uses of the input connector state" + ); + assert!( + dae.continuous + .equations + .iter() + .all(|eq| !expr_contains_der_of(&eq.rhs, &VarName::new("u2"))), + "demotion should rewrite derivative uses of the derived input connector state" + ); +} + #[test] fn test_demote_direct_assigned_states_allows_fixed_state_with_extra_value_ref() { let mut dae = Dae::new(); @@ -1229,6 +1348,242 @@ fn test_demote_direct_assigned_states_skips_unsliced_array_state_alias() { } } +#[test] +fn test_demote_direct_assigned_array_state_projects_indexed_derivative_reads() { + let mut dae = Dae::new(); + let mut psi = test_variable("psi"); + psi.dims = vec![2]; + dae.variables.states.insert(VarName::new("psi"), psi); + let mut v = test_variable("v"); + v.dims = vec![2]; + dae.variables.algebraics.insert(VarName::new("v"), v); + + dae.continuous + .equations + .push(eq(sub(var("psi"), array(vec![var("time"), var("time")])))); + dae.continuous + .equations + .push(eq(sub(var_idx("v", 1), der_idx("psi", 1)))); + dae.continuous + .equations + .push(eq(sub(var_idx("v", 2), der_idx("psi", 2)))); + + let demoted = demote_direct_assigned_states(&mut dae).expect("direct demotion should succeed"); + + assert_eq!(demoted, 1); + assert!(!dae.variables.states.contains_key(&VarName::new("psi"))); + assert!(dae.variables.algebraics.contains_key(&VarName::new("psi"))); + assert!( + dae.continuous + .equations + .iter() + .any(|eq| eq.rhs == sub(var_idx("v", 1), real(1.0))), + "indexed der(psi[1]) should be projected to the first derivative component" + ); + assert!( + dae.continuous + .equations + .iter() + .any(|eq| eq.rhs == sub(var_idx("v", 2), real(1.0))), + "indexed der(psi[2]) should be projected to the second derivative component" + ); +} + +#[test] +fn test_demote_direct_assigned_array_state_from_structured_scalar_slots() { + let mut dae = Dae::new(); + let mut x = test_variable("x"); + x.dims = vec![2]; + dae.variables.states.insert(VarName::new("x"), x); + let mut v = test_variable("v"); + v.dims = vec![2]; + dae.variables.algebraics.insert(VarName::new("v"), v); + + let span = test_span(); + dae.continuous + .equations + .push(eq(sub(var("x"), var_with_span("time", span)))); + dae.continuous.equations.push(eq(sub( + var("x"), + Expression::Binary { + op: OpBinary::Add, + lhs: Box::new(var_with_span("time", span)), + rhs: Box::new(var_with_span("time", span)), + span, + }, + ))); + dae.continuous + .equations + .push(eq(sub(der_idx("x", 1), var_idx("v", 1)))); + dae.continuous + .equations + .push(eq(sub(der_idx("x", 2), var_idx("v", 2)))); + dae.continuous.structured_equations = vec![dae::StructuredEquationFamily { + domain: rumoca_core::StructuredIndexDomain { + binders: vec![rumoca_core::StructuredIndexBinder { + id: 0, + display_name: "i".to_string(), + lower: 1, + upper: 2, + step: 1, + }], + }, + first_equation_index: 0, + equation_counts: vec![1, 1], + span: test_span(), + origin: "structured aggregate assignment".to_string(), + regular: None, + template: None, + interiors_materialized: true, + }]; + + let demoted = demote_direct_assigned_states(&mut dae).expect("direct demotion should succeed"); + + assert_eq!(demoted, 1); + assert!(!dae.variables.states.contains_key(&VarName::new("x"))); + assert!(dae.variables.algebraics.contains_key(&VarName::new("x"))); + assert!( + dae.continuous + .equations + .iter() + .all(|eq| !expr_contains_der_of(&eq.rhs, &VarName::new("x"))), + "structured scalar slots should be assembled into an array derivative replacement" + ); +} + +#[test] +fn test_seeded_relaxed_derivative_map_resolves_algebraic_alias_derivative() { + let mut dae = Dae::new(); + dae.variables + .states + .insert(VarName::new("q"), test_variable("q")); + dae.variables + .algebraics + .insert(VarName::new("i"), test_variable("i")); + dae.variables + .algebraics + .insert(VarName::new("u"), test_variable("u")); + + dae.continuous.equations.push(eq(sub(der("q"), var("u")))); + dae.continuous.equations.push(eq(sub(var("i"), var("q")))); + + let der_map = build_relaxed_derivative_map_for_exprs(&dae, &[var("i")]) + .expect("seeded relaxed derivative map should build"); + + assert_eq!(der_map.get("i"), Some(&var("u"))); +} + +#[test] +fn test_collect_rhs_var_refs_preserves_static_index_before_field_access() { + let refs = collect_rhs_var_refs(&field_expr(index_expr(var("mediums"), 1), "d")); + + assert!( + refs.contains(&VarName::new("mediums[1].d")), + "FieldAccess(Index(mediums, 1), d) must expose the scalar record-field reference" + ); + assert!( + refs.contains(&VarName::new("mediums[1]")), + "the indexed record reference is also needed for alias-definition closure" + ); +} + +#[test] +fn test_extract_unknown_defining_expr_solves_scaled_record_field_target() { + let target = VarName::new("mediums[1].Xi"); + let residual = sub( + var("mXi"), + mul(var("m"), field_expr(index_expr(var("mediums"), 1), "Xi")), + ); + + let defining_expr = extract_unknown_defining_expr(&residual, &target, Span::DUMMY) + .expect("scaled record-field target should be solved from residual"); + + assert_eq!(defining_expr, div(var("mXi"), var("m"))); +} + +#[test] +fn test_symbolic_derivative_keeps_preferred_algebraic_derivative_candidate() { + let mut dae = Dae::new(); + let mut p = test_variable("p"); + p.state_select = rumoca_core::StateSelect::Prefer; + dae.variables.algebraics.insert(VarName::new("p"), p); + + let derivative = symbolic_time_derivative(&var("p"), &dae, &build_der_value_map(&dae)) + .expect("preferred algebraics are valid state-selection derivative candidates"); + + assert!( + matches!( + derivative, + Expression::BuiltinCall { + function: BuiltinFunction::Der, + args, + .. + } if matches!( + args.as_slice(), + [Expression::VarRef { name, subscripts, .. }] + if name.as_str() == "p" && subscripts.is_empty() + ) + ), + "preferred algebraics should remain as der(p) candidates" + ); +} + +#[test] +fn test_symbolic_derivative_rejects_default_algebraic_without_derivative_map() { + let mut dae = Dae::new(); + dae.variables + .algebraics + .insert(VarName::new("h"), test_variable("h")); + + assert!( + symbolic_time_derivative(&var("h"), &dae, &build_der_value_map(&dae)).is_none(), + "default algebraics must not be silently promoted to state derivatives" + ); +} + +#[test] +fn test_symbolic_derivative_prefers_function_derivative_annotation() { + let mut dae = Dae::new(); + dae.variables + .states + .insert(VarName::new("x"), test_variable("x")); + dae.continuous + .equations + .push(eq(sub(der("x"), var("xdot")))); + + let mut function = rumoca_core::Function::new("f", Span::DUMMY); + function.inputs.push(function_param("x")); + function.outputs.push(function_param("y")); + function.body.push(rumoca_core::Statement::Assignment { + comp: rumoca_core::ComponentReference::from_flat_segments("y", Span::DUMMY, None), + value: mul(var("x"), var("x")), + span: Span::DUMMY, + }); + function + .derivatives + .push(rumoca_core::DerivativeAnnotation { + derivative_function: "f_der".to_string(), + order: 1, + zero_derivative: Vec::new(), + no_derivative: Vec::new(), + }); + dae.symbols.functions.insert(VarName::new("f"), function); + dae.symbols.functions.insert( + VarName::new("f_der"), + rumoca_core::Function::new("f_der", Span::DUMMY), + ); + + let derivative = + symbolic_time_derivative(&call("f", vec![var("x")]), &dae, &build_der_value_map(&dae)) + .expect("function call should differentiate through annotation"); + + let Expression::FunctionCall { name, args, .. } = derivative else { + panic!("expected derivative function call"); + }; + assert_eq!(name.as_str(), "f_der"); + assert_eq!(args, vec![var("x"), var("xdot")]); +} + #[test] fn test_index_reduction_differentiates_vector_function_constraint_with_structured_subscripts() { let constraint_span = Span::from_offsets( @@ -1460,6 +1815,57 @@ fn test_index_reduction_accepts_coupled_vector_constraint_rows() { ); } +#[test] +fn test_symbolic_time_derivative_handles_rotation_matrix_transpose() { + let mut dae = Dae::new(); + dae.variables + .states + .insert(VarName::new("theta"), test_variable("theta")); + dae.variables.states.insert(VarName::new("psi"), { + let mut psi = test_variable("psi"); + psi.dims = vec![2]; + psi + }); + dae.variables + .algebraics + .insert(VarName::new("omega"), test_variable("omega")); + + let rotation = array(vec![ + array(vec![ + builtin(BuiltinFunction::Cos, var("theta")), + neg(builtin(BuiltinFunction::Sin, var("theta"))), + ]), + array(vec![ + builtin(BuiltinFunction::Sin, var("theta")), + builtin(BuiltinFunction::Cos, var("theta")), + ]), + ]); + let expr = mul(builtin(BuiltinFunction::Transpose, rotation), var("psi")); + let der_map = HashMap::from([ + ("theta".to_string(), var("omega")), + ("psi".to_string(), array(vec![var("v1"), var("v2")])), + ]); + + let derivative = symbolic_time_derivative(&expr, &dae, &der_map) + .expect("rotation-frame vector derivative should be symbolic"); + + assert!( + !expr_contains_der_of_non_state( + &derivative, + &HashSet::from(["theta".to_string(), "psi".to_string()]) + ), + "derivative should not leave der(non-state) calls: {derivative:?}" + ); + assert!( + !expr_contains_der_of(&derivative, &VarName::new("psi")), + "indexed state derivatives should project through the derivative map: {derivative:?}" + ); + assert!( + format!("{derivative:?}").contains("Transpose"), + "transpose derivative should stay structured: {derivative:?}" + ); +} + #[test] fn test_exact_alias_component_rewrites_derivative_of_non_state_alias_to_canonical_state() { let mut dae = Dae::new(); @@ -1707,6 +2113,86 @@ fn test_constrained_dummy_derivative_reduction_reaches_fixed_point() { ); } +#[test] +fn test_constrained_dummy_reduction_differentiates_function_defined_position_constraint() { + let mut dae = Dae::new(); + dae.variables + .states + .insert(VarName::new("direct.phi"), test_variable("direct.phi")); + let mut inverse_phi = test_variable("inverse.phi"); + inverse_phi.state_select = rumoca_core::StateSelect::Prefer; + dae.variables + .states + .insert(VarName::new("inverse.phi"), inverse_phi); + for name in ["direct.w", "inverse.w"] { + dae.variables + .algebraics + .insert(VarName::new(name), test_variable(name)); + } + + let mut position = rumoca_core::Function::new("position", Span::DUMMY); + let mut q_qd_qdd = function_param("q_qd_qdd"); + q_qd_qdd.dims = vec![3]; + position.inputs.push(q_qd_qdd); + position.inputs.push(function_param("dummy")); + position.outputs.push(function_param("q")); + position.body.push(rumoca_core::Statement::Assignment { + comp: rumoca_core::ComponentReference::from_flat_segments("q", Span::DUMMY, None), + value: index_expr(var("q_qd_qdd"), 1), + span: Span::DUMMY, + }); + dae.symbols + .functions + .insert(VarName::new("position"), position); + + dae.continuous + .equations + .push(eq(sub(der("direct.phi"), var("direct.w")))); + dae.continuous + .equations + .push(eq(sub(der("inverse.phi"), var("inverse.w")))); + dae.continuous.equations.push(eq(sub( + var("inverse.phi"), + call( + "position", + vec![ + array(vec![var("direct.phi"), var("direct.w"), real(0.0)]), + var("time"), + ], + ), + ))); + + let demoted = reduce_constrained_dummy_derivatives(&mut dae) + .expect("function-defined position constraint should reduce"); + + assert_eq!(demoted, 1); + assert!( + !dae.variables + .states + .contains_key(&VarName::new("inverse.phi")), + "position-constrained dummy state should be demoted" + ); + assert!( + dae.variables + .algebraics + .contains_key(&VarName::new("inverse.phi")) + ); + assert!( + dae.continuous + .equations + .iter() + .all(|eq| !expr_contains_der_of(&eq.rhs, &VarName::new("inverse.phi"))), + "der(inverse.phi) should be replaced by the function-derived velocity" + ); + assert!( + dae.continuous.equations.iter().any(|eq| { + expr_contains_var(&eq.rhs, &VarName::new("direct.w")) + && expr_contains_var(&eq.rhs, &VarName::new("inverse.w")) + }), + "the derivative row should become the physical velocity constraint" + ); +} + #[test] fn test_constrained_dummy_state_names_maps_singleton_array_component_state() { let mut dae = Dae::new(); diff --git a/crates/rumoca-phase-structural/src/dae_prepare/direct_demotion.rs b/crates/rumoca-phase-structural/src/dae_prepare/direct_demotion.rs index 61189f663..c6e845080 100644 --- a/crates/rumoca-phase-structural/src/dae_prepare/direct_demotion.rs +++ b/crates/rumoca-phase-structural/src/dae_prepare/direct_demotion.rs @@ -65,6 +65,12 @@ pub(super) fn equation_defining_expr_for_unknown( } return Some(eq.rhs.clone()); } + if let Some(defining_expr) = extract_unknown_defining_expr(&eq.rhs, unknown_name, eq.span) { + if expression_contains_any_der_call(&defining_expr) { + return None; + } + return Some(defining_expr); + } if let Some((coef, remainder)) = split_linear_target(&eq.rhs, unknown_name, eq.span) { let defining_expr = match coef { 1 => sub_expr(zero_expr(eq.span), remainder, eq.span), @@ -169,7 +175,7 @@ fn defining_expr_references_unsafe_non_state_alias_closure( defining_expr: &Expression, state_name_set: &HashSet, non_state_unknown_names: &HashSet, - excluded_eq_index: usize, + excluded_eq_index: Option, alias_safety_cache: &mut AliasSafetyCache, ) -> bool { let mut visiting = HashSet::new(); @@ -178,12 +184,48 @@ fn defining_expr_references_unsafe_non_state_alias_closure( defining_expr, state_name_set, non_state_unknown_names, - Some(excluded_eq_index), + excluded_eq_index, &mut visiting, alias_safety_cache, ) } +fn defining_expr_references_unsafe_non_state_alias_closure_allowing_direct_state_refs( + definitions: &DefiningExprIndex, + defining_expr: &Expression, + state_name_set: &HashSet, + non_state_unknown_names: &HashSet, + excluded_eq_index: Option, + alias_safety_cache: &mut AliasSafetyCache, +) -> bool { + let mut refs = HashSet::new(); + defining_expr.collect_var_refs(&mut refs); + refs.into_iter().any(|ref_name| { + if state_name_set.contains(ref_name.as_str()) { + return false; + } + if !non_state_unknown_names.contains(ref_name.as_str()) { + return false; + } + !non_state_alias_closure_is_state_free( + definitions, + &ref_name, + state_name_set, + non_state_unknown_names, + excluded_eq_index, + &mut HashSet::new(), + alias_safety_cache, + ) + }) +} + +fn state_is_input_connector(dae: &Dae, state_name: &VarName) -> bool { + dae.variables + .states + .get(state_name) + .is_some_and(|var| matches!(var.causality, dae::VariableCausality::Input)) +} + fn apply_direct_demotion_plans( dae: &mut Dae, substitutions: &HashMap, @@ -198,10 +240,14 @@ fn apply_direct_demotion_plans( } pub(super) fn apply_direct_demotion_plan(dae: &mut Dae, plan: &DirectStateDemotionPlan) -> usize { + promote_plan_derivative_algebraics(dae, &plan.promote_der_algebraics); + ensure_scalar_state_partition_for_plan(dae, &plan.state_name); + let state_dims = variable_dims_for_direct_demotion(dae, &plan.state_name); for eq in &mut dae.continuous.equations { - eq.rhs = substitute_der_of_state(&eq.rhs, &plan.state_name, &plan.der_expr); + eq.rhs = substitute_der_of_state(&eq.rhs, &plan.state_name, &plan.der_expr, &state_dims); } - if let Some(var) = dae.variables.states.shift_remove(&plan.state_name) { + if let Some(mut var) = dae.variables.states.shift_remove(&plan.state_name) { + var.fixed = Some(false); dae.variables .algebraics .insert(plan.state_name.clone(), var); @@ -210,103 +256,131 @@ pub(super) fn apply_direct_demotion_plan(dae: &mut Dae, plan: &DirectStateDemoti 0 } +fn promote_plan_derivative_algebraics(dae: &mut Dae, names: &[VarName]) { + for name in names { + if dae.variables.states.contains_key(name) { + continue; + } + if let Some(var) = dae.variables.algebraics.shift_remove(name) { + dae.variables.states.insert(name.clone(), var); + } + } +} + +fn ensure_scalar_state_partition_for_plan(dae: &mut Dae, state_name: &VarName) { + if dae.variables.states.contains_key(state_name) { + return; + } + let Some(scalar) = rumoca_core::parse_scalar_name(state_name.as_str()) else { + return; + }; + let base_name = VarName::new(scalar.base); + let Some(base_var) = dae.variables.states.shift_remove(&base_name) else { + return; + }; + let size = base_var.size(); + if size <= 1 { + dae.variables.states.insert(base_name, base_var); + return; + } + for flat_index in 0..size { + let scalar_name = VarName::new(dae::scalar_name_text_for_flat_index( + base_name.as_str(), + &base_var.dims, + flat_index, + )); + let mut scalar_var = base_var.clone(); + scalar_var.name = scalar_name.clone(); + scalar_var.dims.clear(); + dae.variables.states.insert(scalar_name, scalar_var); + } +} + fn direct_demotion_plan_for_equation( round: &DirectDemotionRound<'_>, eq_index: usize, eq: &Equation, counters: &mut DirectDemotionCounters, alias_safety_cache: &mut AliasSafetyCache, -) -> Option { - let (state_name, defining_expr) = - extract_state_direct_assignment_equation(eq, &round.state_names, &round.state_name_set)?; - let defining_expr = if is_connection_equation_origin(&eq.origin) { - match connection_component_fixed_defining_expr( - round.dae, - &state_name, - &round.state_name_set, - ) { - Some(expr) => expr, - None => defining_expr, - } - } else { - defining_expr +) -> Result, StructuralError> { + let (state_name, defining_expr) = match direct_assignment_candidate(round, eq) { + Some(candidate) => candidate, + None => return Ok(None), }; counters.n_candidates += 1; if eq.origin.starts_with("flow sum equation:") { counters.n_skip_flow_sum_origin += 1; - return None; + return Ok(None); } log_direct_assignment_candidate(round.trace, counters, round.dae, eq, &state_name); if round.when_assigned_states.contains(state_name.as_str()) { counters.n_skip_when_assigned += 1; - return None; + return Ok(None); } - if expr_contains_der_of(&defining_expr, &state_name) { + if expr_contains_der_of_state_or_indexed(&defining_expr, &state_name) { counters.n_skip_self_der += 1; - return None; + return Ok(None); } - if !state_ders_in_expr_independently_defined(&defining_expr, &state_name, round) { + if !state_is_input_connector(round.dae, &state_name) + && !defining_expr_state_derivatives_are_demotable(round, &state_name, &defining_expr)? + { counters.n_skip_der_in_defining_expr += 1; - return None; + return Ok(None); } // `der(state)` links are substituted symbolically on demotion (gated by // `state_ders_in_expr_independently_defined` above and validated again in // `choose_derivative_replacement`), so mask them before scanning for value // dependencies on states or unsafe alias closures. let alias_scan_expr = mask_state_der_calls(&defining_expr, &round.state_name_set); - if defining_expr_references_unsafe_non_state_alias_closure( - &round.non_state_defining_exprs, - &alias_scan_expr, - &round.state_name_set, - &round.non_state_unknown_names, - eq_index, - alias_safety_cache, - ) { + let unsafe_alias_closure = if state_is_input_connector(round.dae, &state_name) { + defining_expr_references_unsafe_non_state_alias_closure_allowing_direct_state_refs( + &round.non_state_defining_exprs, + &alias_scan_expr, + &round.state_name_set, + &round.non_state_unknown_names, + Some(eq_index), + alias_safety_cache, + ) + } else { + defining_expr_references_unsafe_non_state_alias_closure( + &round.non_state_defining_exprs, + &alias_scan_expr, + &round.state_name_set, + &round.non_state_unknown_names, + Some(eq_index), + alias_safety_cache, + ) + }; + if unsafe_alias_closure { counters.n_skip_unsafe_non_state_alias += 1; - return None; + return Ok(None); } - if round - .dae - .variables - .states - .get(&state_name) - .is_some_and(|state| state.size() > 1) - || expr_contains_unsliced_vector_ref(&defining_expr, round.dae) - { + if !direct_assignment_shape_is_demotable(round.dae, &state_name, &defining_expr) { // MLS §10.1: array state shape is semantic IR. This path substitutes - // whole `der(state)` calls, so unsliced compound states stay intact. + // whole `der(state)` calls only when the defining expression has the + // same aggregate shape. Scalar states still reject unsliced vector refs. counters.n_skip_unsliced_vector_ref += 1; - return None; - } - let state_non_der_ref_rows = round - .dae - .continuous - .equations - .iter() - .filter(|row| { - expr_contains_var(&row.rhs, &state_name) && !expr_contains_der_of(&row.rhs, &state_name) - }) - .count(); - if state_non_der_ref_rows > 1 - && !expr_refs_only_parameters_constants_or_time(round.dae, &defining_expr) - { - counters.n_skip_extra_state_refs += 1; - return None; + return Ok(None); } + let der_map = + build_relaxed_derivative_map_for_exprs(round.dae, std::slice::from_ref(&defining_expr))?; let der_expr = choose_derivative_replacement( &defining_expr, &round.state_name_set, round.dae, - &round.der_map, + &der_map, counters, - )?; - if expr_contains_der_of(&der_expr, &state_name) { + ); + let Some(der_expr) = der_expr else { + return Ok(None); + }; + if expr_contains_der_of_state_or_indexed(&der_expr, &state_name) { counters.n_skip_self_der += 1; - return None; + return Ok(None); } if expr_contains_der_of_non_state(&der_expr, &round.state_name_set) { counters.n_skip_non_state_der += 1; - return None; + return Ok(None); } if round.trace && counters.n_trace_logged_candidates < 16 { crate::structural_trace!( @@ -316,10 +390,44 @@ fn direct_demotion_plan_for_equation( ); counters.n_trace_logged_candidates += 1; } - Some(DirectStateDemotionPlan { + Ok(Some(DirectStateDemotionPlan { state_name, der_expr, - }) + promote_der_algebraics: Vec::new(), + })) +} + +fn direct_assignment_candidate( + round: &DirectDemotionRound<'_>, + eq: &Equation, +) -> Option<(VarName, Expression)> { + let (state_name, defining_expr) = + extract_state_direct_assignment_equation(eq, &round.state_names, &round.state_name_set)?; + if !is_connection_equation_origin(&eq.origin) { + return Some((state_name, defining_expr)); + } + let defining_expr = + connection_component_fixed_defining_expr(round.dae, &state_name, &round.state_name_set) + .unwrap_or(defining_expr); + Some((state_name, defining_expr)) +} + +fn defining_expr_state_derivatives_are_demotable( + round: &DirectDemotionRound<'_>, + state_name: &VarName, + defining_expr: &Expression, +) -> Result { + if derivative_states_in_eq(defining_expr, &round.state_names).is_empty() { + return Ok(true); + } + let der_map = + build_relaxed_derivative_map_for_exprs(round.dae, std::slice::from_ref(defining_expr))?; + Ok(state_ders_in_expr_independently_defined( + defining_expr, + state_name, + round, + &der_map, + )) } /// `der(z)` links inside a defining expression are demotable only when `z`'s @@ -332,23 +440,25 @@ fn state_ders_in_expr_independently_defined( defining_expr: &Expression, candidate: &VarName, round: &DirectDemotionRound<'_>, + der_map: &HashMap, ) -> bool { derivative_states_in_eq(defining_expr, &round.state_names) .iter() .all(|inner_state| { - round - .der_map - .get(inner_state.as_str()) - .is_some_and(|value| { - !expr_contains_der_of(value, inner_state) - && !expr_contains_var(value, candidate) - }) + der_map.get(inner_state.as_str()).is_some_and(|value| { + !expr_contains_der_of(value, inner_state) && !expr_contains_var(value, candidate) + }) }) } -fn collect_direct_demotion_plans( +fn expr_contains_der_of_state_or_indexed(expr: &Expression, state_name: &VarName) -> bool { + expr_contains_der_of(expr, state_name) +} + +fn collect_direct_demotion_plans_with_boundary_substitutions( dae: &Dae, trace: bool, + boundary_substitutions: &[crate::eliminate::Substitution], ) -> Result, StructuralError> { let timer = structural_timing_start("direct_demotion.collect_round"); let Some(round) = DirectDemotionRound::new(dae, trace)? else { @@ -358,6 +468,16 @@ fn collect_direct_demotion_plans( let mut alias_safety_cache = AliasSafetyCache::new(); let mut substitutions = HashMap::new(); let mut counters = DirectDemotionCounters::default(); + for plan in collect_componentwise_direct_demotion_plans( + &round, + &mut counters, + boundary_substitutions, + &mut alias_safety_cache, + )? { + substitutions + .entry(plan.state_name.as_str().to_string()) + .or_insert(plan); + } let timer = structural_timing_start("direct_demotion.scan_equations"); for (eq_index, eq) in round.dae.continuous.equations.iter().enumerate() { @@ -367,7 +487,8 @@ fn collect_direct_demotion_plans( eq, &mut counters, &mut alias_safety_cache, - ) else { + )? + else { continue; }; substitutions @@ -380,6 +501,594 @@ fn collect_direct_demotion_plans( Ok(substitutions) } +fn collect_componentwise_direct_demotion_plans( + round: &DirectDemotionRound<'_>, + counters: &mut DirectDemotionCounters, + boundary_substitutions: &[crate::eliminate::Substitution], + alias_safety_cache: &mut AliasSafetyCache, +) -> Result, StructuralError> { + let by_state = collect_componentwise_direct_demotion_slots(round, boundary_substitutions)?; + + let mut plans = Vec::new(); + for (state_name_string, slots) in by_state { + let state_name = VarName::new(state_name_string.as_str()); + let Some(dims) = variable_dims_for_direct_demotion(round.dae, &state_name) else { + continue; + }; + let Some(component_exprs) = slots.iter().cloned().collect::>>() else { + plans.extend(componentwise_scalar_demotion_plans( + round, + counters, + &state_name, + &dims, + &slots, + alias_safety_cache, + )?); + continue; + }; + let Some(defining_expr) = array_expr_from_flat_values(component_exprs, &dims) else { + plans.extend(componentwise_scalar_demotion_plans( + round, + counters, + &state_name, + &dims, + &slots, + alias_safety_cache, + )?); + continue; + }; + let alias_scan_expr = mask_state_der_calls(&defining_expr, &round.state_name_set); + if defining_expr_references_unsafe_non_state_alias_closure( + &round.non_state_defining_exprs, + &alias_scan_expr, + &round.state_name_set, + &round.non_state_unknown_names, + None, + alias_safety_cache, + ) { + counters.n_skip_unsafe_non_state_alias += 1; + plans.extend(componentwise_scalar_demotion_plans( + round, + counters, + &state_name, + &dims, + &slots, + alias_safety_cache, + )?); + continue; + } + let der_map = build_relaxed_derivative_map_for_exprs( + round.dae, + std::slice::from_ref(&defining_expr), + )?; + let Some(der_expr) = choose_derivative_replacement( + &defining_expr, + &round.state_name_set, + round.dae, + &der_map, + counters, + ) else { + plans.extend(componentwise_scalar_demotion_plans( + round, + counters, + &state_name, + &dims, + &slots, + alias_safety_cache, + )?); + continue; + }; + if expr_contains_der_of_state_or_indexed(&der_expr, &state_name) + || expr_contains_der_of_non_state(&der_expr, &round.state_name_set) + { + plans.extend(componentwise_scalar_demotion_plans( + round, + counters, + &state_name, + &dims, + &slots, + alias_safety_cache, + )?); + continue; + } + plans.push(DirectStateDemotionPlan { + state_name, + der_expr, + promote_der_algebraics: Vec::new(), + }); + } + + Ok(plans) +} + +fn collect_componentwise_direct_demotion_slots( + round: &DirectDemotionRound<'_>, + boundary_substitutions: &[crate::eliminate::Substitution], +) -> Result>>, StructuralError> { + let mut by_state: IndexMap>> = IndexMap::new(); + for (eq_index, eq) in round.dae.continuous.equations.iter().enumerate() { + let component_assignment = extract_state_component_direct_assignment_equation( + round.dae, + eq_index, + eq, + &round.state_name_set, + ); + let aggregate_assignment = component_assignment + .is_none() + .then(|| aggregate_direct_assignment_component_slots(round, eq)); + let Some((state_name, slot_exprs)) = component_assignment + .map(|(state_name, flat_index, defining_expr)| { + (state_name, vec![(flat_index, defining_expr)]) + }) + .or_else(|| aggregate_assignment.flatten()) + else { + continue; + }; + if round.when_assigned_states.contains(state_name.as_str()) + || slot_exprs + .iter() + .any(|(_, expr)| expr_contains_der_of_state_or_indexed(expr, &state_name)) + { + continue; + } + if slot_exprs + .iter() + .any(|(_, expr)| !derivative_states_in_eq(expr, &round.state_names).is_empty()) + { + let seed_exprs = slot_exprs + .iter() + .map(|(_, expr)| expr.clone()) + .collect::>(); + let der_map = build_relaxed_derivative_map_for_exprs(round.dae, &seed_exprs)?; + if slot_exprs.iter().any(|(_, expr)| { + !state_ders_in_expr_independently_defined(expr, &state_name, round, &der_map) + }) { + continue; + } + } + let Some(size) = round + .dae + .variables + .states + .get(&state_name) + .map(Variable::size) + .filter(|size| *size > 1) + else { + continue; + }; + let slots = by_state + .entry(state_name.as_str().to_string()) + .or_insert_with(|| vec![None; size]); + for (flat_index, defining_expr) in slot_exprs { + if flat_index >= slots.len() || slots[flat_index].is_some() { + continue; + } + slots[flat_index] = Some(defining_expr); + } + } + collect_boundary_substitution_component_slots(round, &mut by_state, boundary_substitutions); + Ok(by_state) +} + +fn aggregate_direct_assignment_component_slots( + round: &DirectDemotionRound<'_>, + eq: &Equation, +) -> Option<(VarName, Vec<(usize, Expression)>)> { + let (state_name, defining_expr, dims) = direct_assignment_candidate(round, eq) + .and_then(|(state_name, defining_expr)| { + let dims = variable_dims_for_direct_demotion(round.dae, &state_name)?; + Some((state_name, defining_expr, dims)) + }) + .or_else(|| aggregate_scalar_state_family_assignment(round, eq))?; + if dims.is_empty() { + return None; + } + let size = dims.iter().try_fold(1usize, |acc, dim| { + (*dim > 0).then(|| acc.checked_mul(*dim as usize)).flatten() + })?; + let mut slots = Vec::with_capacity(size); + for flat_index in 0..size { + let expr = project_flat_index_with_span(&defining_expr, &dims, flat_index, None)?; + slots.push((flat_index, expr)); + } + Some((state_name, slots)) +} + +fn aggregate_scalar_state_family_assignment( + round: &DirectDemotionRound<'_>, + eq: &Equation, +) -> Option<(VarName, Expression, Vec)> { + if let Some(lhs) = &eq.lhs { + let (state_name, dims) = aggregate_scalar_state_family_dims(round.dae, lhs.var_name())?; + return Some((state_name, eq.rhs.clone(), dims)); + } + let Expression::Binary { + op: OpBinary::Sub, + lhs, + rhs, + .. + } = &eq.rhs + else { + return None; + }; + if let Some((state_name, dims)) = + aggregate_scalar_state_family_expr_dims(round.dae, lhs, &round.state_name_set) + && !expr_contains_var(rhs, &state_name) + { + return Some((state_name, *rhs.clone(), dims)); + } + if let Some((state_name, dims)) = + aggregate_scalar_state_family_expr_dims(round.dae, rhs, &round.state_name_set) + && !expr_contains_var(lhs, &state_name) + { + return Some((state_name, *lhs.clone(), dims)); + } + None +} + +fn aggregate_scalar_state_family_expr_dims( + dae: &Dae, + expr: &Expression, + state_name_set: &HashSet, +) -> Option<(VarName, Vec)> { + let Expression::VarRef { + name, subscripts, .. + } = expr + else { + return None; + }; + if !subscripts.is_empty() || state_name_set.contains(name.as_str()) { + return None; + } + aggregate_scalar_state_family_dims(dae, name.var_name()) +} + +fn aggregate_scalar_state_family_dims( + dae: &Dae, + state_name: &VarName, +) -> Option<(VarName, Vec)> { + if dae.variables.states.contains_key(state_name) { + return None; + } + let mut parsed_scalars = dae + .variables + .states + .keys() + .filter_map(|candidate| { + let scalar = rumoca_core::parse_scalar_name(candidate.as_str())?; + (scalar.base == state_name.as_str()).then_some(scalar) + }) + .collect::>(); + if parsed_scalars.is_empty() { + return None; + } + parsed_scalars.sort_by(|a, b| a.indices.cmp(&b.indices)); + let rank = parsed_scalars.first()?.indices.len(); + if rank == 0 + || parsed_scalars + .iter() + .any(|scalar| scalar.indices.len() != rank) + { + return None; + } + let mut dims = vec![0_i64; rank]; + for scalar in &parsed_scalars { + for (axis, index) in scalar.indices.iter().enumerate() { + dims[axis] = dims[axis].max(*index); + } + } + let size = dims.iter().try_fold(1usize, |acc, dim| { + (*dim > 0).then(|| acc.checked_mul(*dim as usize)).flatten() + })?; + if size != parsed_scalars.len() { + return None; + } + for flat_index in 0..size { + let scalar_name = VarName::new(dae::scalar_name_text_for_flat_index( + state_name.as_str(), + &dims, + flat_index, + )); + if !dae.variables.states.contains_key(&scalar_name) { + return None; + } + } + Some((state_name.clone(), dims)) +} + +fn componentwise_scalar_demotion_plans( + round: &DirectDemotionRound<'_>, + counters: &mut DirectDemotionCounters, + state_name: &VarName, + dims: &[i64], + slots: &[Option], + alias_safety_cache: &mut AliasSafetyCache, +) -> Result, StructuralError> { + let mut plans = Vec::new(); + for (flat_index, defining_expr) in slots.iter().enumerate() { + let Some(defining_expr) = defining_expr else { + continue; + }; + let scalar_state_name = VarName::new(dae::scalar_name_text_for_flat_index( + state_name.as_str(), + dims, + flat_index, + )); + let alias_scan_expr = mask_state_der_calls(defining_expr, &round.state_name_set); + if defining_expr_references_unsafe_non_state_alias_closure( + &round.non_state_defining_exprs, + &alias_scan_expr, + &round.state_name_set, + &round.non_state_unknown_names, + None, + alias_safety_cache, + ) { + counters.n_skip_unsafe_non_state_alias += 1; + continue; + } + let der_map = + build_relaxed_derivative_map_for_exprs(round.dae, std::slice::from_ref(defining_expr))?; + let Some(der_expr) = choose_derivative_replacement_allowing_preferred_promotions( + defining_expr, + &round.state_name_set, + round.dae, + &der_map, + counters, + )? + else { + continue; + }; + if expr_contains_der_of_state_or_indexed(&der_expr, &scalar_state_name) { + continue; + } + let promote_der_algebraics = preferred_derivative_algebraics(round.dae, &der_expr); + if expr_contains_der_of_unpromoted_non_state( + &der_expr, + &round.state_name_set, + &promote_der_algebraics, + ) { + continue; + } + plans.push(DirectStateDemotionPlan { + state_name: scalar_state_name, + der_expr, + promote_der_algebraics, + }); + } + Ok(plans) +} + +fn choose_derivative_replacement_allowing_preferred_promotions( + defining_expr: &Expression, + state_name_set: &HashSet, + dae: &Dae, + der_map: &HashMap, + counters: &mut DirectDemotionCounters, +) -> Result, StructuralError> { + let Some(symbolic) = symbolic_time_derivative(defining_expr, dae, der_map) else { + counters.n_skip_no_der_expr += 1; + return Ok(None); + }; + let promote_der_algebraics = preferred_derivative_algebraics(dae, &symbolic); + if expr_contains_der_of_unpromoted_non_state(&symbolic, state_name_set, &promote_der_algebraics) + { + counters.n_skip_non_state_der += 1; + return Ok(None); + } + Ok(Some(symbolic)) +} + +fn preferred_derivative_algebraics(dae: &Dae, expr: &Expression) -> Vec { + let mut names = Vec::new(); + collect_der_of_algebraics(expr, dae, &mut names); + let mut seen = HashSet::new(); + names.retain(|name| { + seen.insert(name.as_str().to_string()) && derivative_algebraic_is_preferred_state(dae, name) + }); + names +} + +fn derivative_algebraic_is_preferred_state(dae: &Dae, name: &VarName) -> bool { + dae.variables.algebraics.get(name).is_some_and(|var| { + state_select_rank(var.state_select) >= state_select_rank(rumoca_core::StateSelect::Prefer) + }) +} + +fn expr_contains_der_of_unpromoted_non_state( + expr: &Expression, + state_name_set: &HashSet, + promoted: &[VarName], +) -> bool { + let promoted = promoted + .iter() + .map(|name| name.as_str().to_string()) + .collect::>(); + let mut checker = UnpromotedNonStateDerivativeChecker { + state_name_set, + promoted: &promoted, + found: false, + }; + checker.visit_expression(expr); + checker.found +} + +struct UnpromotedNonStateDerivativeChecker<'a> { + state_name_set: &'a HashSet, + promoted: &'a HashSet, + found: bool, +} + +impl ExpressionVisitor for UnpromotedNonStateDerivativeChecker<'_> { + fn visit_builtin_call(&mut self, function: &BuiltinFunction, args: &[Expression]) { + if *function == BuiltinFunction::Der { + self.found = + der_arg_is_not_plain_state_or_promoted(args, self.state_name_set, self.promoted); + return; + } + for arg in args { + self.visit_expression(arg); + } + } +} + +fn der_arg_is_not_plain_state_or_promoted( + args: &[Expression], + state_name_set: &HashSet, + promoted: &HashSet, +) -> bool { + if args.len() != 1 { + return true; + } + match &args[0] { + Expression::VarRef { name, .. } => { + !state_name_set.contains(name.as_str()) && !promoted.contains(name.as_str()) + } + _ => true, + } +} + +fn collect_boundary_substitution_component_slots( + round: &DirectDemotionRound<'_>, + by_state: &mut IndexMap>>, + boundary_substitutions: &[crate::eliminate::Substitution], +) { + for substitution in boundary_substitutions { + let Some((state_name, flat_index)) = + boundary_substitution_state_component(round.dae, substitution, &round.state_name_set) + else { + continue; + }; + if round.when_assigned_states.contains(state_name.as_str()) + || expr_contains_der_of_state_or_indexed(&substitution.expr, &state_name) + { + continue; + } + let Some(size) = round + .dae + .variables + .states + .get(&state_name) + .map(Variable::size) + .filter(|size| *size > 1) + else { + continue; + }; + let slots = by_state + .entry(state_name.as_str().to_string()) + .or_insert_with(|| vec![None; size]); + if flat_index < slots.len() && slots[flat_index].is_none() { + slots[flat_index] = Some(substitution.expr.clone()); + } + } +} + +fn boundary_substitution_state_component( + dae: &Dae, + substitution: &crate::eliminate::Substitution, + state_name_set: &HashSet, +) -> Option<(VarName, usize)> { + let scalar = rumoca_core::parse_scalar_name(substitution.var_name.as_str())?; + if !state_name_set.contains(scalar.base) { + return None; + } + let state_name = VarName::new(scalar.base); + let dims = variable_dims_for_direct_demotion(dae, &state_name)?; + let flat_index = flat_index_from_indices(&dims, &scalar.indices)?; + Some((state_name, flat_index)) +} + +fn extract_state_component_direct_assignment_equation( + dae: &Dae, + eq_index: usize, + eq: &Equation, + state_name_set: &HashSet, +) -> Option<(VarName, usize, Expression)> { + let Expression::Binary { + op: OpBinary::Sub, + lhs, + rhs, + .. + } = &eq.rhs + else { + return None; + }; + if let Some((state_name, flat_index)) = state_component_ref_flat_index(dae, lhs, state_name_set) + && !expr_contains_var(rhs, &state_name) + { + return Some((state_name, flat_index, *rhs.clone())); + } + if let Some((state_name, flat_index)) = state_component_ref_flat_index(dae, rhs, state_name_set) + && !expr_contains_var(lhs, &state_name) + { + return Some((state_name, flat_index, *lhs.clone())); + } + if let Some((state_name, flat_index)) = + structured_scalar_slot_state_component(dae, eq_index, lhs, state_name_set) + && !expr_contains_var(rhs, &state_name) + { + return Some((state_name, flat_index, *rhs.clone())); + } + if let Some((state_name, flat_index)) = + structured_scalar_slot_state_component(dae, eq_index, rhs, state_name_set) + && !expr_contains_var(lhs, &state_name) + { + return Some((state_name, flat_index, *lhs.clone())); + } + None +} + +fn structured_scalar_slot_state_component( + dae: &Dae, + eq_index: usize, + expr: &Expression, + state_name_set: &HashSet, +) -> Option<(VarName, usize)> { + let Expression::VarRef { + name, subscripts, .. + } = expr + else { + return None; + }; + if !subscripts.is_empty() || !state_name_set.contains(name.as_str()) { + return None; + } + let state_name = name.var_name().clone(); + let state_size = dae.variables.states.get(&state_name)?.size(); + if state_size <= 1 { + return None; + } + let slot = dae::structured_equation_slot(&dae.continuous.structured_equations, eq_index)?; + if slot.equation_count != 1 || slot.equation_position != 0 { + return None; + } + let family = dae.continuous.structured_equations.get(slot.family_index)?; + if family.equation_counts.len() != state_size || slot.iteration_index >= state_size { + return None; + } + Some((state_name, slot.iteration_index)) +} + +fn state_component_ref_flat_index( + dae: &Dae, + expr: &Expression, + state_name_set: &HashSet, +) -> Option<(VarName, usize)> { + let Expression::VarRef { + name, subscripts, .. + } = expr + else { + return None; + }; + if subscripts.is_empty() || !state_name_set.contains(name.as_str()) { + return None; + } + let state_name = name.var_name().clone(); + let dims = variable_dims_for_direct_demotion(dae, &state_name)?; + let indices = static_subscript_indices(subscripts)?; + let flat_index = flat_index_from_indices(&dims, &indices)?; + Some((state_name, flat_index)) +} + /// Demote states that are explicitly defined by direct assignment equations /// (`state = expr`) and substitute `der(state)` with `d/dt(expr)` throughout /// the system. @@ -391,6 +1100,13 @@ fn collect_direct_demotion_plans( /// state is demoted. States assigned in `when` clauses are preserved, since /// they participate in event/reinit updates and must remain in the state vector. pub fn demote_direct_assigned_states(dae: &mut Dae) -> Result { + demote_direct_assigned_states_with_boundary_substitutions(dae, &[]) +} + +pub fn demote_direct_assigned_states_with_boundary_substitutions( + dae: &mut Dae, + boundary_substitutions: &[crate::eliminate::Substitution], +) -> Result { let max_rounds = dae.variables.states.len().clamp(1, 8); let mut total_demoted = 0usize; @@ -398,7 +1114,11 @@ pub fn demote_direct_assigned_states(dae: &mut Dae) -> Result Result bool { + let Some(state) = dae.variables.states.get(state_name) else { + return false; + }; + if state.size() <= 1 { + return expression_dims(defining_expr, dae).is_some_and(|dims| dims.is_empty()) + || !expr_contains_unsliced_vector_ref(defining_expr, dae); + } + let Some(state_dims) = variable_dims_for_direct_demotion(dae, state_name) else { + return false; + }; + expression_dims(defining_expr, dae).is_some_and(|expr_dims| expr_dims == state_dims) +} diff --git a/crates/rumoca-phase-structural/src/dae_prepare/direct_demotion_piecewise_tests.rs b/crates/rumoca-phase-structural/src/dae_prepare/direct_demotion_piecewise_tests.rs new file mode 100644 index 000000000..1eb1ec9bf --- /dev/null +++ b/crates/rumoca-phase-structural/src/dae_prepare/direct_demotion_piecewise_tests.rs @@ -0,0 +1,141 @@ +use super::*; + +fn test_variable(name: &str) -> Variable { + let mut variable = Variable::new(VarName::new(name), test_span()); + variable.source_span = test_span(); + variable +} + +fn test_span() -> Span { + Span::from_offsets( + rumoca_core::SourceId::from_source_name("direct_demotion_piecewise_tests.mo"), + 1, + 2, + ) +} + +fn eq(rhs: Expression) -> Equation { + Equation::residual(rhs, Span::DUMMY, "test") +} + +fn var(name: &str) -> Expression { + Expression::VarRef { + name: Reference::new(name), + subscripts: vec![], + span: Span::DUMMY, + } +} + +fn real(value: f64) -> Expression { + Expression::Literal { + value: Literal::Real(value), + span: Span::DUMMY, + } +} + +fn der(name: &str) -> Expression { + Expression::BuiltinCall { + function: BuiltinFunction::Der, + args: vec![var(name)], + span: Span::DUMMY, + } +} + +fn add(lhs: Expression, rhs: Expression) -> Expression { + Expression::Binary { + op: OpBinary::Add, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span: Span::DUMMY, + } +} + +fn sub(lhs: Expression, rhs: Expression) -> Expression { + Expression::Binary { + op: OpBinary::Sub, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span: Span::DUMMY, + } +} + +fn mul(lhs: Expression, rhs: Expression) -> Expression { + Expression::Binary { + op: OpBinary::Mul, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span: Span::DUMMY, + } +} + +fn lt(lhs: Expression, rhs: Expression) -> Expression { + Expression::Binary { + op: OpBinary::Lt, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span: Span::DUMMY, + } +} + +#[test] +fn demote_direct_assigned_states_extracts_piecewise_trajectory_definition() { + let mut dae = Dae::new(); + for name in ["s", "sd"] { + dae.variables + .states + .insert(VarName::new(name), test_variable(name)); + } + dae.variables + .algebraics + .insert(VarName::new("sdd"), test_variable("sdd")); + for name in ["q", "qd"] { + dae.variables + .outputs + .insert(VarName::new(name), test_variable(name)); + } + + dae.continuous.equations.push(eq(Expression::If { + branches: vec![(lt(var("time"), real(1.0)), sub(var("s"), real(0.0)))], + else_branch: Box::new(sub( + Expression::If { + branches: vec![( + lt(var("time"), real(2.0)), + sub(var("s"), mul(real(0.5), var("time"))), + )], + else_branch: Box::new(sub(var("s"), var("time"))), + span: Span::DUMMY, + }, + real(0.0), + )), + span: Span::DUMMY, + })); + dae.continuous.equations.push(eq(sub(var("sd"), der("s")))); + dae.continuous + .equations + .push(eq(sub(var("sdd"), der("sd")))); + dae.continuous + .equations + .push(eq(sub(var("q"), add(real(10.0), mul(real(2.0), var("s")))))); + dae.continuous + .equations + .push(eq(sub(var("qd"), mul(real(2.0), var("sd"))))); + + let demoted = + demote_direct_assigned_states(&mut dae).expect("piecewise direct trajectory should demote"); + + assert_eq!(demoted, 2); + assert!(dae.variables.states.is_empty()); + for name in ["s", "sd", "sdd"] { + assert!(dae.variables.algebraics.contains_key(&VarName::new(name))); + } + assert!( + dae.continuous + .equations + .iter() + .all( + |eq| !rumoca_ir_dae::expr_contains_der_of(&eq.rhs, &VarName::new("s")) + && !rumoca_ir_dae::expr_contains_der_of(&eq.rhs, &VarName::new("sd")) + ), + "demotion should replace the derivative chain with symbolic trajectory derivatives" + ); +} diff --git a/crates/rumoca-phase-structural/src/dae_prepare/dummy_state_metadata.rs b/crates/rumoca-phase-structural/src/dae_prepare/dummy_state_metadata.rs index a391f9e32..ae2453adb 100644 --- a/crates/rumoca-phase-structural/src/dae_prepare/dummy_state_metadata.rs +++ b/crates/rumoca-phase-structural/src/dae_prepare/dummy_state_metadata.rs @@ -142,9 +142,15 @@ fn defining_expr_references_non_alias_output( } fn output_is_direct_alias_of_state(dae: &Dae, output_name: &VarName, state_name: &VarName) -> bool { - super::find_defining_expr_candidates(dae, output_name) + let mut candidates = super::find_defining_expr_candidates(dae, output_name); + if candidates.is_empty() + && let Some(scalar) = rumoca_core::parse_scalar_name(output_name.as_str()) + { + candidates = super::find_defining_expr_candidates(dae, &VarName::new(scalar.base)); + } + candidates .into_iter() - .any(|expr| expression_is_plain_var_ref(&expr, state_name)) + .any(|expr| expression_is_state_or_component_alias(&expr, output_name, state_name)) } fn expression_is_plain_var_ref(expr: &Expression, expected: &VarName) -> bool { @@ -155,6 +161,40 @@ fn expression_is_plain_var_ref(expr: &Expression, expected: &VarName) -> bool { ) } +fn expression_is_state_or_component_alias( + expr: &Expression, + output_name: &VarName, + state_name: &VarName, +) -> bool { + if expression_is_plain_var_ref(expr, state_name) { + return true; + } + let Some(output_scalar) = rumoca_core::parse_scalar_name(output_name.as_str()) else { + return false; + }; + let Expression::VarRef { + name, subscripts, .. + } = expr + else { + return false; + }; + if let Some(state_scalar) = rumoca_core::parse_scalar_name(name.as_str()) { + return state_scalar.base == state_name.as_str() + && state_scalar.indices == output_scalar.indices + && subscripts.is_empty(); + } + name.var_name() == state_name + && subscripts + .iter() + .map(|subscript| match subscript { + Subscript::Index { value, .. } => Some(*value), + Subscript::Expr { expr, .. } => numeric_constant(expr).map(|value| value as i64), + Subscript::Colon { .. } => None, + }) + .collect::>>() + .is_some_and(|indices| indices == output_scalar.indices) +} + /// True when the constrained-dummy defining expression for `state_name` /// references another state, i.e. the defining constraint genuinely couples /// differential states (a high-index DAE such as `w1 = ratio * w2`). @@ -174,8 +214,95 @@ fn candidate_is_self_integrating_non_state_alias( definition: &ConstrainedDummyDefinition, state_names: &[VarName], ) -> bool { - state_has_standalone_der_equation(dae, state_name, state_names).unwrap_or(false) - && !defining_expr_couples_other_state(&definition.defining_expr, state_name, state_names) + let has_any_derivative_reference = dae + .continuous + .equations + .iter() + .any(|eq| super::expr_contains_der_of(&eq.rhs, state_name)); + if has_any_derivative_reference + && (defining_expr_is_plain_other_state_alias( + &definition.defining_expr, + state_name, + state_names, + ) || defining_expr_is_direct_output_alias_of_state( + dae, + &definition.defining_expr, + state_name, + ) || defining_expr_references_direct_output_alias_of_state( + dae, + &definition.defining_expr, + state_name, + )) + { + return true; + } + if !state_has_standalone_der_equation(dae, state_name, state_names).unwrap_or(false) { + return false; + } + !defining_expr_couples_other_state(&definition.defining_expr, state_name, state_names) + || defining_expr_is_plain_other_state_alias( + &definition.defining_expr, + state_name, + state_names, + ) + || defining_expr_is_direct_output_alias_of_state(dae, &definition.defining_expr, state_name) + || defining_expr_references_direct_output_alias_of_state( + dae, + &definition.defining_expr, + state_name, + ) +} + +fn defining_expr_is_plain_other_state_alias( + defining_expr: &Expression, + state_name: &VarName, + state_names: &[VarName], +) -> bool { + let Expression::VarRef { + name, subscripts, .. + } = defining_expr + else { + return false; + }; + subscripts.is_empty() + && name.var_name() != state_name + && state_names.iter().any(|state| state == name.var_name()) +} + +fn defining_expr_is_direct_output_alias_of_state( + dae: &Dae, + defining_expr: &Expression, + state_name: &VarName, +) -> bool { + let Some(row) = linear_terms(dae, defining_expr) else { + return false; + }; + if row.terms.len() != 1 { + return false; + } + row.terms.keys().any(|term| { + let term_name = VarName::new(term.clone()); + output_is_direct_alias_of_state(dae, &term_name, state_name) + || rumoca_core::parse_scalar_name(term).is_some_and(|scalar| { + let base = VarName::new(scalar.base); + output_is_direct_alias_of_state(dae, &base, state_name) + }) + }) +} + +fn defining_expr_references_direct_output_alias_of_state( + dae: &Dae, + defining_expr: &Expression, + state_name: &VarName, +) -> bool { + let mut refs = Vec::new(); + defining_expr.collect_var_refs(&mut refs); + refs.into_iter().any(|name| { + output_is_direct_alias_of_state(dae, &name, state_name) + || rumoca_core::parse_scalar_name(name.as_str()).is_some_and(|scalar| { + output_is_direct_alias_of_state(dae, &VarName::new(scalar.base), state_name) + }) + }) } fn numeric_constant(expr: &Expression) -> Option { diff --git a/crates/rumoca-phase-structural/src/dae_prepare/matrix_state_derivative_tests.rs b/crates/rumoca-phase-structural/src/dae_prepare/matrix_state_derivative_tests.rs new file mode 100644 index 000000000..59924f6e0 --- /dev/null +++ b/crates/rumoca-phase-structural/src/dae_prepare/matrix_state_derivative_tests.rs @@ -0,0 +1,335 @@ +use super::*; +use std::collections::{HashMap, HashSet}; + +fn test_span() -> Span { + Span::from_offsets( + rumoca_core::SourceId::from_source_name("matrix_state_derivative_test.mo"), + 1, + 2, + ) +} + +fn test_variable(name: &str) -> Variable { + let mut variable = Variable::new(VarName::new(name), test_span()); + variable.source_span = test_span(); + variable +} + +fn var(name: &str) -> Expression { + Expression::VarRef { + name: rumoca_core::Reference::new(name), + subscripts: vec![], + span: Span::DUMMY, + } +} + +fn var_idx2(name: &str, row: i64, col: i64) -> Expression { + Expression::VarRef { + name: rumoca_core::Reference::new(name), + subscripts: vec![ + Subscript::generated_index(row, Span::DUMMY), + Subscript::generated_index(col, Span::DUMMY), + ], + span: Span::DUMMY, + } +} + +fn var_idx1(name: &str, index: i64) -> Expression { + Expression::VarRef { + name: rumoca_core::Reference::new(name), + subscripts: vec![Subscript::generated_index(index, Span::DUMMY)], + span: Span::DUMMY, + } +} + +fn real(v: f64) -> Expression { + Expression::Literal { + value: Literal::Real(v), + span: Span::DUMMY, + } +} + +fn array(elements: Vec) -> Expression { + Expression::Array { + elements, + is_matrix: false, + span: Span::DUMMY, + } +} + +fn sub(lhs: Expression, rhs: Expression) -> Expression { + Expression::Binary { + op: OpBinary::Sub, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span: Span::DUMMY, + } +} + +fn mul(lhs: Expression, rhs: Expression) -> Expression { + Expression::Binary { + op: OpBinary::Mul, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span: Span::DUMMY, + } +} + +fn der(expr: Expression) -> Expression { + Expression::BuiltinCall { + function: BuiltinFunction::Der, + args: vec![expr], + span: test_span(), + } +} + +fn der_idx2(name: &str, row: i64, col: i64) -> Expression { + der(var_idx2(name, row, col)) +} + +fn der_idx1(name: &str, index: i64) -> Expression { + der(var_idx1(name, index)) +} + +fn eq(rhs: Expression) -> Equation { + Equation { + lhs: None, + rhs, + span: Span::DUMMY, + origin: "matrix state derivative regression".to_string(), + scalar_count: 1, + } +} + +fn matrix_derivative_product_dae() -> Dae { + let mut dae = Dae::new(); + let mut r = test_variable("R"); + r.dims = vec![3, 3]; + dae.variables.states.insert(VarName::new("R"), r); + let mut skew = test_variable("skew"); + skew.dims = vec![3, 3]; + dae.variables.algebraics.insert(VarName::new("skew"), skew); + + dae.continuous.equations.push(eq(sub( + var("skew"), + array(vec![ + array(vec![real(0.0), real(-1.0), real(0.0)]), + array(vec![real(1.0), real(0.0), real(0.0)]), + array(vec![real(0.0), real(0.0), real(0.0)]), + ]), + ))); + dae.continuous + .equations + .push(eq(sub(der(var("R")), mul(var("R"), var("skew"))))); + dae +} + +#[test] +fn test_exact_alias_component_rewrites_indexed_state_derivative_to_canonical_state() { + let mut dae = Dae::new(); + let mut position = test_variable("position[1]"); + position.state_select = rumoca_core::StateSelect::Always; + dae.variables + .states + .insert(VarName::new("position[1]"), position); + dae.variables + .states + .insert(VarName::new("p[1]"), test_variable("p[1]")); + dae.variables + .algebraics + .insert(VarName::new("v[1]"), test_variable("v[1]")); + + dae.continuous + .equations + .push(eq(sub(var_idx1("position", 1), var_idx1("p", 1)))); + dae.continuous + .equations + .push(eq(sub(der_idx1("p", 1), var_idx1("v", 1)))); + + let demoted = + demote_exact_alias_component_states(&mut dae).expect("exact alias demotion should run"); + + assert_eq!(demoted, 1); + assert!( + dae.variables + .states + .contains_key(&VarName::new("position[1]")) + ); + assert!(!dae.variables.states.contains_key(&VarName::new("p[1]"))); + assert!( + dae.continuous + .equations + .iter() + .any(|eq| expr_contains_der_of(&eq.rhs, &VarName::new("position[1]"))), + "indexed derivative alias should be rewritten to the canonical scalarized state" + ); + assert!( + !dae.continuous + .equations + .iter() + .any(|eq| expr_contains_der_of(&eq.rhs, &VarName::new("p[1]"))), + "demoted scalarized state derivative should not survive exact alias rewrite" + ); +} + +#[test] +fn test_assignable_derivative_rows_keep_matrix_state_component_rows() { + let mut dae = Dae::new(); + let mut r = test_variable("R"); + r.dims = vec![3, 3]; + dae.variables.states.insert(VarName::new("R"), r); + let mut rhs = test_variable("rhs"); + rhs.dims = vec![3, 3]; + dae.variables.algebraics.insert(VarName::new("rhs"), rhs); + + for row in 1..=3 { + for col in 1..=3 { + dae.continuous + .equations + .push(eq(sub(der_idx2("R", row, col), var_idx2("rhs", row, col)))); + } + } + + let demoted = demote_states_without_assignable_derivative_rows(&mut dae); + assert_eq!(demoted, 0); + assert!(dae.variables.states.contains_key(&VarName::new("R"))); +} + +#[test] +fn test_direct_demotion_keeps_matrix_state_component_ode_rows() { + let mut dae = Dae::new(); + let mut r = test_variable("R"); + r.dims = vec![3, 3]; + dae.variables.states.insert(VarName::new("R"), r); + let mut skew = test_variable("skew"); + skew.dims = vec![3, 3]; + dae.variables.algebraics.insert(VarName::new("skew"), skew); + + for row in 1..=3 { + for col in 1..=3 { + dae.continuous.equations.push(eq(sub( + der_idx2("R", row, col), + mul(var_idx2("R", row, 1), var_idx2("skew", 1, col)), + ))); + } + } + + let demoted = demote_direct_assigned_states(&mut dae).expect("direct demotion should succeed"); + assert_eq!(demoted, 0); + assert!(dae.variables.states.contains_key(&VarName::new("R"))); +} + +#[test] +fn test_direct_demotion_keeps_array_state_components_defined_by_output_aliases() { + let mut dae = Dae::new(); + let mut p = test_variable("p"); + p.dims = vec![3]; + dae.variables.states.insert(VarName::new("p"), p); + let mut position = test_variable("position"); + position.dims = vec![3]; + dae.variables + .outputs + .insert(VarName::new("position"), position); + let mut v = test_variable("v"); + v.dims = vec![3]; + dae.variables.algebraics.insert(VarName::new("v"), v); + + for index in 1..=3 { + dae.continuous + .equations + .push(eq(sub(var_idx1("position", index), var_idx1("p", index)))); + dae.continuous + .equations + .push(eq(sub(der_idx1("p", index), var_idx1("v", index)))); + } + + let demoted = demote_direct_assigned_states(&mut dae).expect("direct demotion should run"); + + assert_eq!(demoted, 0); + assert!(dae.variables.states.contains_key(&VarName::new("p"))); + for index in 1..=3 { + assert!( + dae.continuous + .equations + .iter() + .any(|eq| expr_contains_der_of(&eq.rhs, &VarName::new(format!("p[{index}]")))), + "state component p[{index}] should keep its derivative row" + ); + } +} + +#[test] +fn test_direct_demotion_keeps_scalarized_matrix_state_component_ode_rows() { + let mut dae = matrix_derivative_product_dae(); + + crate::scalarize::scalarize_equations(&mut dae).expect("scalarization should succeed"); + let demoted = demote_direct_assigned_states(&mut dae).expect("direct demotion should succeed"); + assert_eq!(demoted, 0); + assert!(dae.variables.states.contains_key(&VarName::new("R"))); +} + +#[test] +fn test_eliminate_trivial_accepts_scalarized_matrix_state_component_ode_rows() { + let mut dae = matrix_derivative_product_dae(); + + crate::scalarize::scalarize_equations(&mut dae).expect("scalarization should succeed"); + demote_direct_assigned_states(&mut dae).expect("direct demotion should succeed"); + let result = crate::eliminate::eliminate_trivial(&mut dae).expect("elimination should run"); + + assert!( + result.blt_error.is_none(), + "matrix ODE rows should stay structurally matchable: {:?}", + result.blt_error + ); + demote_states_without_retained_derivative_rows(&mut dae) + .expect("post-elimination retained derivative demotion should run"); + assert!(dae.variables.states.contains_key(&VarName::new("R"))); +} + +#[test] +fn test_prepare_retains_scalarized_matrix_state_component_ode_rows() { + let mut dae = matrix_derivative_product_dae(); + + crate::scalarize::scalarize_equations(&mut dae).expect("scalarization should succeed"); + assert_eq!(demote_direct_assigned_states(&mut dae).unwrap(), 0); + assert_eq!(reduce_constrained_dummy_derivatives(&mut dae).unwrap(), 0); + assert_eq!(index_reduce_missing_state_derivatives(&mut dae).unwrap(), 0); + assert_eq!( + demote_states_without_assignable_derivative_rows(&mut dae), + 0 + ); + eliminate_derivative_aliases(&mut dae).expect("derivative alias elimination should run"); + demote_states_without_retained_derivative_rows(&mut dae) + .expect("retained derivative demotion should run"); + + assert!(dae.variables.states.contains_key(&VarName::new("R"))); +} + +#[test] +fn test_expand_derivative_preserves_indexed_state_component_derivative() { + let mut dae = Dae::new(); + let mut r = test_variable("R"); + r.dims = vec![3, 3]; + dae.variables.states.insert(VarName::new("R"), r); + + let expanded = expand_der_in_expr_full( + &der_idx2("R", 1, 2), + &dae, + &HashMap::from([("R".to_string(), array(vec![real(0.0)]))]), + &HashSet::from(["R".to_string()]), + ); + + let Expression::BuiltinCall { function, args, .. } = expanded else { + panic!("expected indexed state derivative to remain a der() call"); + }; + assert_eq!(function, BuiltinFunction::Der); + assert_eq!(args.len(), 1); + let Expression::VarRef { + name, subscripts, .. + } = &args[0] + else { + panic!("expected indexed state derivative argument"); + }; + assert_eq!(name.as_str(), "R"); + assert_eq!(subscripts.len(), 2); +} diff --git a/crates/rumoca-phase-structural/src/dae_prepare/mod.rs b/crates/rumoca-phase-structural/src/dae_prepare/mod.rs index d26000888..7d5ee5480 100644 --- a/crates/rumoca-phase-structural/src/dae_prepare/mod.rs +++ b/crates/rumoca-phase-structural/src/dae_prepare/mod.rs @@ -31,6 +31,8 @@ type Variable = dae::Variable; type DefiningExprIndex = IndexMap>; type AliasSafetyCache = IndexMap<(String, Option), bool>; +const MAX_DIRECT_DEMOTION_DEFINING_EXPR_NODES: usize = 1024; + #[derive(Clone)] struct IndexedDefiningExpr { equation_index: usize, @@ -41,8 +43,9 @@ mod connection_alias; use connection_alias::connection_component_fixed_defining_expr; mod symbolic; use symbolic::{ - build_der_value_map, expand_der_in_expr_full, symbolic_time_derivative, truncate_debug, - try_extract_der_value, + array_expr_from_flat_values, build_der_value_map, expand_der_in_expr_full, expression_dims, + field_access_candidate_var_names, flat_index_from_indices, project_flat_index_with_span, + static_subscript_indices, symbolic_time_derivative, truncate_debug, try_extract_der_value, }; mod row_shape; use row_shape::{dae_variable_size, required_dae_variable_size, residual_scalar_width}; @@ -53,12 +56,15 @@ pub use dummy_state_metadata::{ }; mod direct_demotion; mod state_row_reduction; -pub use direct_demotion::demote_direct_assigned_states; use direct_demotion::{ collect_non_state_continuous_unknown_names, equation_defining_expr_for_unknown, expr_refs_only_parameters_constants_or_time, expression_contains_any_der_call, is_connection_equation_origin, }; +pub use direct_demotion::{ + demote_direct_assigned_states, demote_direct_assigned_states_with_boundary_substitutions, +}; +use state_row_reduction::expression_exact_name; pub use state_row_reduction::{ REGULARIZATION_LEVELS, demote_orphan_states_without_equation_refs, demote_states_without_assignable_derivative_rows, demote_states_without_derivative_refs, @@ -122,8 +128,8 @@ fn extract_scaled_target(expr: &Expression, target: &VarName) -> Option Option<(i32, Expression)> { let span = expr.span().unwrap_or(context_span); - if expr_refers_to_var(expr, target) { + if expression_is_linear_target_ref(expr, target) { return Some((1, zero_expr(span))); } @@ -179,18 +185,60 @@ fn split_linear_target( } } +fn expression_is_linear_target_ref(expr: &Expression, target: &VarName) -> bool { + matches!( + expr, + Expression::VarRef { .. } | Expression::Index { .. } | Expression::FieldAccess { .. } + ) && expression_exact_name(expr).is_some_and(|name| name == target.as_str()) +} + fn extract_defining_expr(eq: &Equation, alg_name: &VarName) -> Option { - let Expression::Binary { op, lhs, rhs, .. } = &eq.rhs else { + extract_unknown_defining_expr(&eq.rhs, alg_name, eq.span) +} + +fn extract_unknown_defining_expr( + residual: &Expression, + alg_name: &VarName, + context_span: Span, +) -> Option { + let Expression::Binary { op, lhs, rhs, .. } = residual else { + if let Expression::Unary { + op: OpUnary::Minus, + rhs, + .. + } = residual + { + return extract_unknown_defining_expr(rhs, alg_name, context_span); + } + if let Expression::If { + branches, + else_branch, + span, + } = residual + { + let mut defining_branches = Vec::with_capacity(branches.len()); + for (condition, branch_expr) in branches { + defining_branches.push(( + condition.clone(), + extract_unknown_defining_expr(branch_expr, alg_name, context_span)?, + )); + } + return Some(Expression::If { + branches: defining_branches, + else_branch: Box::new(extract_unknown_defining_expr( + else_branch, + alg_name, + context_span, + )?), + span: *span, + }); + } return None; }; if !matches!(op, OpBinary::Sub) { return None; } - - let is_var = |e: &Expression| -> bool { - matches!(e, Expression::VarRef { name, subscripts, .. } - if name.var_name() == alg_name && subscripts.is_empty()) - }; + let is_var = |e: &Expression| -> bool { expr_refers_to_var(e, alg_name) }; // 0 = var - expr → var = expr → return expr if is_var(lhs) { @@ -200,6 +248,12 @@ fn extract_defining_expr(eq: &Equation, alg_name: &VarName) -> Option Option x = rhs/coeff - return Some(div_expr(*rhs.clone(), coeff, eq.span)); + return Some(div_expr(*rhs.clone(), coeff, context_span)); } if rhs_has && let Some(coeff) = extract_scaled_target(rhs, alg_name) { // lhs - (coeff*x) = 0 => x = lhs/coeff - return Some(div_expr(*lhs.clone(), coeff, eq.span)); + return Some(div_expr(*lhs.clone(), coeff, context_span)); } - if lhs_has && let Some((coef, lhs_rem)) = split_linear_target(lhs, alg_name, eq.span) { + if lhs_has && let Some((coef, lhs_rem)) = split_linear_target(lhs, alg_name, context_span) { // (coef*x + lhs_rem) - rhs = 0 => x = (rhs - lhs_rem)/coef return Some(match coef { - 1 => sub_expr(*rhs.clone(), lhs_rem, eq.span), - -1 => sub_expr(lhs_rem, *rhs.clone(), eq.span), + 1 => sub_expr(*rhs.clone(), lhs_rem, context_span), + -1 => sub_expr(lhs_rem, *rhs.clone(), context_span), _ => return None, }); } - if rhs_has && let Some((coef, rhs_rem)) = split_linear_target(rhs, alg_name, eq.span) { + if rhs_has && let Some((coef, rhs_rem)) = split_linear_target(rhs, alg_name, context_span) { // lhs - (coef*x + rhs_rem) = 0 => x = (lhs - rhs_rem)/coef return Some(match coef { - 1 => sub_expr(*lhs.clone(), rhs_rem, eq.span), - -1 => sub_expr(rhs_rem, *lhs.clone(), eq.span), + 1 => sub_expr(*lhs.clone(), rhs_rem, context_span), + -1 => sub_expr(rhs_rem, *lhs.clone(), context_span), _ => return None, }); } @@ -258,9 +312,68 @@ fn push_indexed_defining_expr( fn collect_rhs_var_refs(expr: &Expression) -> IndexSet { let mut refs = IndexSet::new(); expr.collect_var_refs(&mut refs); + FieldAccessVarRefCollector { refs: &mut refs }.visit_expression(expr); refs } +struct FieldAccessVarRefCollector<'a> { + refs: &'a mut IndexSet, +} + +impl ExpressionVisitor for FieldAccessVarRefCollector<'_> { + fn visit_var_ref(&mut self, name: &Reference, subscripts: &[Subscript]) { + self.refs.insert(name.var_name().clone()); + if let Some(indices) = static_subscript_indices(subscripts) + && !indices.is_empty() + { + let index_text = indices + .iter() + .map(i64::to_string) + .collect::>() + .join(","); + self.refs + .insert(VarName::new(format!("{}[{}]", name.as_str(), index_text))); + } + for subscript in subscripts { + self.visit_subscript(subscript); + } + } + + fn visit_index(&mut self, base: &Expression, subscripts: &[Subscript]) { + if let Expression::VarRef { + name, + subscripts: base_subscripts, + .. + } = base + { + let mut combined = Vec::with_capacity(base_subscripts.len() + subscripts.len()); + combined.extend_from_slice(base_subscripts); + combined.extend_from_slice(subscripts); + if let Some(indices) = static_subscript_indices(&combined) + && !indices.is_empty() + { + let index_text = indices + .iter() + .map(i64::to_string) + .collect::>() + .join(","); + self.refs + .insert(VarName::new(format!("{}[{}]", name.as_str(), index_text))); + } + } + self.visit_expression(base); + for subscript in subscripts { + self.visit_subscript(subscript); + } + } + + fn visit_field_access(&mut self, base: &Expression, field: &str) { + self.refs + .extend(field_access_candidate_var_names(base, field)); + self.visit_expression(base); + } +} + fn collect_residual_defining_expr_index(dae: &Dae) -> DefiningExprIndex { let mut index = DefiningExprIndex::new(); for (equation_index, eq) in dae.continuous.equations.iter().enumerate() { @@ -276,6 +389,11 @@ fn collect_residual_defining_expr_index(dae: &Dae) -> DefiningExprIndex { fn collect_non_derivative_defining_expr_index(dae: &Dae) -> DefiningExprIndex { let mut index = DefiningExprIndex::new(); for (equation_index, eq) in dae.continuous.equations.iter().enumerate() { + if expression_node_count_exceeds(&eq.rhs, MAX_DIRECT_DEMOTION_DEFINING_EXPR_NODES) + .unwrap_or(true) + { + continue; + } let lhs_name = eq.lhs.as_ref().map(|lhs| lhs.var_name().clone()); if let Some(name) = &lhs_name && !expression_contains_any_der_call(&eq.rhs) @@ -295,6 +413,168 @@ fn collect_non_derivative_defining_expr_index(dae: &Dae) -> DefiningExprIndex { index } +fn expression_node_count_exceeds(expr: &Expression, limit: usize) -> Option { + let mut count = 0usize; + expression_node_count_visit(expr, limit, &mut count) +} + +fn expression_node_count_visit_all<'a>( + exprs: impl IntoIterator, + limit: usize, + count: &mut usize, +) -> Option { + for expr in exprs { + if expression_node_count_visit(expr, limit, count)? { + return Some(true); + } + } + Some(false) +} + +fn expression_node_count_visit_branches( + branches: &[(Expression, Expression)], + limit: usize, + count: &mut usize, +) -> Option { + for (condition, branch) in branches { + if expression_node_count_visit_all([condition, branch], limit, count)? { + return Some(true); + } + } + Some(false) +} + +fn expression_node_count_visit_subscripts( + subscripts: &[rumoca_core::Subscript], + limit: usize, + count: &mut usize, +) -> Option { + expression_node_count_visit_all( + subscripts.iter().filter_map(|subscript| match subscript { + rumoca_core::Subscript::Expr { expr, .. } => Some(expr.as_ref()), + _ => None, + }), + limit, + count, + ) +} + +fn expression_node_count_visit(expr: &Expression, limit: usize, count: &mut usize) -> Option { + *count = count.checked_add(1)?; + if *count > limit { + return Some(true); + } + match expr { + Expression::Binary { lhs, rhs, .. } => { + expression_node_count_visit_binary(lhs, rhs, limit, count) + } + Expression::Unary { rhs, .. } => expression_node_count_visit(rhs, limit, count), + Expression::BuiltinCall { args, .. } | Expression::FunctionCall { args, .. } => { + expression_node_count_visit_all(args, limit, count) + } + Expression::If { + branches, + else_branch, + .. + } => expression_node_count_visit_if(branches, else_branch, limit, count), + Expression::Array { elements, .. } | Expression::Tuple { elements, .. } => { + expression_node_count_visit_all(elements, limit, count) + } + Expression::Range { + start, step, end, .. + } => expression_node_count_visit_range(start, step.as_deref(), end, limit, count), + Expression::Index { + base, subscripts, .. + } => expression_node_count_visit_index(base, subscripts, limit, count), + Expression::FieldAccess { base, .. } => expression_node_count_visit(base, limit, count), + Expression::ArrayComprehension { + expr, + indices, + filter, + .. + } => expression_node_count_visit_comprehension( + expr, + indices, + filter.as_deref(), + limit, + count, + ), + _ => Some(false), + } +} + +fn expression_node_count_visit_binary( + lhs: &Expression, + rhs: &Expression, + limit: usize, + count: &mut usize, +) -> Option { + if expression_node_count_visit(lhs, limit, count)? { + return Some(true); + } + expression_node_count_visit(rhs, limit, count) +} + +fn expression_node_count_visit_if( + branches: &[(Expression, Expression)], + else_branch: &Expression, + limit: usize, + count: &mut usize, +) -> Option { + if expression_node_count_visit_branches(branches, limit, count)? { + return Some(true); + } + expression_node_count_visit(else_branch, limit, count) +} + +fn expression_node_count_visit_range( + start: &Expression, + step: Option<&Expression>, + end: &Expression, + limit: usize, + count: &mut usize, +) -> Option { + if expression_node_count_visit(start, limit, count)? { + return Some(true); + } + if let Some(step) = step + && expression_node_count_visit(step, limit, count)? + { + return Some(true); + } + expression_node_count_visit(end, limit, count) +} + +fn expression_node_count_visit_index( + base: &Expression, + subscripts: &[rumoca_core::Subscript], + limit: usize, + count: &mut usize, +) -> Option { + if expression_node_count_visit(base, limit, count)? { + return Some(true); + } + expression_node_count_visit_subscripts(subscripts, limit, count) +} + +fn expression_node_count_visit_comprehension( + expr: &Expression, + indices: &[rumoca_core::ComprehensionIndex], + filter: Option<&Expression>, + limit: usize, + count: &mut usize, +) -> Option { + if expression_node_count_visit(expr, limit, count)? { + return Some(true); + } + if expression_node_count_visit_all(indices.iter().map(|index| &index.range), limit, count)? { + return Some(true); + } + filter.map_or(Some(false), |filter| { + expression_node_count_visit(filter, limit, count) + }) +} + fn defining_expr_candidates<'a>( index: &'a DefiningExprIndex, name: &VarName, @@ -305,103 +585,197 @@ fn defining_expr_candidates<'a>( .flat_map(|candidates| candidates.iter().map(|candidate| &candidate.expr)) } -fn continuous_variable<'a>(dae: &'a Dae, name: &VarName) -> Option<&'a Variable> { - dae.variables - .states - .get(name) - .or_else(|| dae.variables.algebraics.get(name)) - .or_else(|| dae.variables.outputs.get(name)) - .or_else(|| dae.variables.inputs.get(name)) +fn build_relaxed_derivative_map_for_exprs( + dae: &Dae, + seed_exprs: &[Expression], +) -> Result, StructuralError> { + let defining_expr_index = collect_residual_defining_expr_index(dae); + build_relaxed_derivative_map_for_exprs_with_index(dae, &defining_expr_index, seed_exprs) } -fn insert_symbolic_derivative_fallback( +fn build_relaxed_derivative_map_for_exprs_with_index( dae: &Dae, - map: &mut HashMap, - name: &VarName, -) -> Result<(), StructuralError> { - let Some(variable) = continuous_variable(dae, name) else { - return Ok(()); - }; - match map.entry(name.as_str().to_string()) { - Entry::Occupied(_) => {} - Entry::Vacant(entry) => { - entry.insert(symbolic_der_var_ref_for_variable(variable)?); - } - } - Ok(()) + defining_expr_index: &DefiningExprIndex, + seed_exprs: &[Expression], +) -> Result, StructuralError> { + let mut map = build_der_value_map(dae); + let candidate_names = + collect_seeded_relaxation_candidates(dae, defining_expr_index, seed_exprs); + relax_algebraic_derivative_map_to_fixed_point( + dae, + defining_expr_index, + &mut map, + &candidate_names, + false, + ); + Ok(map) } -fn resolve_derivatives_for_expr( +fn collect_seeded_relaxation_candidates( dae: &Dae, defining_expr_index: &DefiningExprIndex, - map: &mut HashMap, - expr: &Expression, - visiting: &mut HashSet, -) -> Result<(), StructuralError> { - for ref_name in collect_rhs_var_refs(expr) { - resolve_derivative_for_var(dae, defining_expr_index, map, &ref_name, visiting)?; + seed_exprs: &[Expression], +) -> IndexSet { + let mut candidates = IndexSet::new(); + let mut stack: Vec = seed_exprs + .iter() + .flat_map(|expr| collect_rhs_var_refs(expr).into_iter()) + .collect(); + + while let Some(name) = stack.pop() { + if !(dae.variables.algebraics.contains_key(&name) + || dae.variables.outputs.contains_key(&name)) + { + continue; + } + if !candidates.insert(name.clone()) { + continue; + } + for defining_expr in defining_expr_candidates(defining_expr_index, &name) { + stack.extend(collect_rhs_var_refs(defining_expr)); + } } - Ok(()) + + candidates } -fn resolve_derivative_for_var( +fn collect_all_relaxation_candidates(dae: &Dae) -> IndexSet { + dae.variables + .algebraics + .keys() + .chain(dae.variables.outputs.keys()) + .cloned() + .collect() +} + +fn relax_algebraic_derivative_map_to_fixed_point( dae: &Dae, defining_expr_index: &DefiningExprIndex, - map: &mut HashMap, - name: &VarName, - visiting: &mut HashSet, -) -> Result<(), StructuralError> { - if name.as_str() == "time" - || dae.variables.parameters.contains_key(name) - || dae.variables.constants.contains_key(name) - || map.contains_key(name.as_str()) - { - return Ok(()); + der_map: &mut HashMap, + candidate_names: &IndexSet, + replace_existing: bool, +) { + let resolvable = candidate_names.len(); + let state_name_set = dae + .variables + .states + .keys() + .map(|name| name.as_str().to_string()) + .collect::>(); + for _ in 0..resolvable.max(1) { + let mut changed = false; + for alg_name in candidate_names { + if !replace_existing + && der_map + .get(alg_name.as_str()) + .is_some_and(|existing| !is_symbolic_derivative_of_var(existing, alg_name)) + { + continue; + } + let derivative = defining_expr_candidates(defining_expr_index, alg_name) + .filter_map(|expr| symbolic_time_derivative(expr, dae, der_map)) + .find(|derivative| { + !expr_contains_der_of(derivative, alg_name) + && !expr_contains_unrelaxed_derivative(derivative, dae, &state_name_set) + }); + let Some(derivative) = derivative else { + continue; + }; + if der_map.get(alg_name.as_str()) == Some(&derivative) { + continue; + } + der_map.insert(alg_name.as_str().to_string(), derivative); + changed = true; + } + if !changed { + break; + } } +} - if dae.variables.states.contains_key(name) { - return insert_symbolic_derivative_fallback(dae, map, name); - } +fn expr_contains_unrelaxed_derivative( + expr: &Expression, + dae: &Dae, + state_name_set: &HashSet, +) -> bool { + let mut checker = UnrelaxedDerivativeChecker { + dae, + state_name_set, + found: false, + }; + checker.visit_expression(expr); + checker.found +} - if !visiting.insert(name.as_str().to_string()) { - return insert_symbolic_derivative_fallback(dae, map, name); - } +struct UnrelaxedDerivativeChecker<'a> { + dae: &'a Dae, + state_name_set: &'a HashSet, + found: bool, +} - for defining_expr in defining_expr_candidates(defining_expr_index, name) { - resolve_derivatives_for_expr(dae, defining_expr_index, map, defining_expr, visiting)?; - let derivative = symbolic_time_derivative(defining_expr, dae, map); - if let Some(derivative) = derivative - && !expr_contains_der_of(&derivative, name) - { - map.insert(name.as_str().to_string(), derivative); - visiting.remove(name.as_str()); - return Ok(()); +impl ExpressionVisitor for UnrelaxedDerivativeChecker<'_> { + fn visit_expression(&mut self, expr: &Expression) { + if !self.found { + self.walk_expression(expr); } } - visiting.remove(name.as_str()); - insert_symbolic_derivative_fallback(dae, map, name) + fn visit_builtin_call(&mut self, function: &BuiltinFunction, args: &[Expression]) { + if *function == BuiltinFunction::Der { + self.found = + der_arg_is_not_state_or_preferred_algebraic(args, self.dae, self.state_name_set); + return; + } + for arg in args { + self.visit_expression(arg); + } + } } -fn build_relaxed_derivative_map_for_exprs( +fn der_arg_is_not_state_or_preferred_algebraic( + args: &[Expression], dae: &Dae, - seed_exprs: &[Expression], -) -> Result, StructuralError> { - let defining_expr_index = collect_residual_defining_expr_index(dae); - build_relaxed_derivative_map_for_exprs_with_index(dae, &defining_expr_index, seed_exprs) + state_name_set: &HashSet, +) -> bool { + if args.len() != 1 { + return true; + } + let Expression::VarRef { + name, subscripts, .. + } = &args[0] + else { + return true; + }; + if !subscripts.is_empty() { + return true; + } + if state_name_set.contains(name.as_str()) { + return false; + } + !dae.variables + .algebraics + .get(name.var_name()) + .is_some_and(|var| { + state_select_rank(var.state_select) + >= state_select_rank(rumoca_core::StateSelect::Prefer) + }) } -fn build_relaxed_derivative_map_for_exprs_with_index( - dae: &Dae, - defining_expr_index: &DefiningExprIndex, - seed_exprs: &[Expression], -) -> Result, StructuralError> { - let mut map = build_der_value_map(dae); - let mut visiting = HashSet::new(); - for expr in seed_exprs { - resolve_derivatives_for_expr(dae, defining_expr_index, &mut map, expr, &mut visiting)?; +fn is_symbolic_derivative_of_var(expr: &Expression, name: &VarName) -> bool { + let Expression::BuiltinCall { function, args, .. } = expr else { + return false; + }; + if *function != BuiltinFunction::Der || args.len() != 1 { + return false; } - Ok(map) + matches!( + &args[0], + Expression::VarRef { + name: ref_name, + subscripts, + .. + } if ref_name.var_name() == name && subscripts.is_empty() + ) } /// Iteratively resolve time derivatives for algebraic variables. @@ -845,36 +1219,14 @@ pub fn build_relaxed_derivative_map( } } - let resolvable = dae.variables.algebraics.len() + dae.variables.outputs.len(); - for _ in 0..resolvable.max(1) { - let mut changed = false; - // Outputs resolve like algebraics (see `compute_full_derivative_map`): - // their derivatives are needed to expand `der(output)` in differentiator - // chains such as `Modelica.Blocks.Continuous.Der`. - for alg_name in dae - .variables - .algebraics - .keys() - .chain(dae.variables.outputs.keys()) - { - let derivative = defining_expr_candidates(&defining_expr_index, alg_name) - .find_map(|expr| symbolic_time_derivative(expr, dae, &map)); - let Some(derivative) = derivative else { - continue; - }; - if expr_contains_der_of(&derivative, alg_name) { - continue; - } - if map.get(alg_name.as_str()) == Some(&derivative) { - continue; - } - map.insert(alg_name.as_str().to_string(), derivative); - changed = true; - } - if !changed { - break; - } - } + let candidate_names = collect_all_relaxation_candidates(dae); + relax_algebraic_derivative_map_to_fixed_point( + dae, + &defining_expr_index, + &mut map, + &candidate_names, + true, + ); Ok(map) } @@ -1019,26 +1371,17 @@ pub fn try_extract_state_alias_pair(rhs: &Expression) -> Option<(VarName, VarNam if !matches!(op, OpBinary::Sub) { return None; } - let Expression::VarRef { - name: lhs_name, - subscripts: lhs_subscripts, - span: _, - } = lhs.as_ref() - else { - return None; - }; - let Expression::VarRef { - name: rhs_name, - subscripts: rhs_subscripts, - span: _, - } = rhs.as_ref() - else { - return None; - }; - if !lhs_subscripts.is_empty() || !rhs_subscripts.is_empty() { - return None; + let lhs_name = expression_exact_name(lhs)?; + let rhs_name = expression_exact_name(rhs)?; + Some((VarName::new(lhs_name), VarName::new(rhs_name))) +} + +fn try_extract_state_alias_pair_from_equation(eq: &Equation) -> Option<(VarName, VarName)> { + if let Some(lhs) = eq.lhs.as_ref() { + let rhs_name = expression_exact_name(&eq.rhs)?; + return Some((VarName::new(lhs.as_str()), VarName::new(rhs_name))); } - Some((lhs_name.var_name().clone(), rhs_name.var_name().clone())) + try_extract_state_alias_pair(&eq.rhs) } fn state_select_rank(state_select: rumoca_core::StateSelect) -> u8 { @@ -1061,6 +1404,9 @@ fn choose_exact_alias_state_representative<'a>( .min_by_key(|(name, var)| { ( Reverse(state_select_rank(var.state_select)), + Reverse(u8::from(exact_alias_state_has_derivative_reference( + dae, name, + ))), Reverse(u8::from(var.fixed == Some(true))), Reverse(u8::from(var.start.is_some())), name.as_str().to_string(), @@ -1069,6 +1415,13 @@ fn choose_exact_alias_state_representative<'a>( .map(|(name, _)| name) } +fn exact_alias_state_has_derivative_reference(dae: &Dae, name: &VarName) -> bool { + dae.continuous + .equations + .iter() + .any(|eq| expr_contains_der_of(&eq.rhs, name)) +} + fn exact_alias_member_variable<'a>(dae: &'a Dae, name: &VarName) -> Option<&'a Variable> { dae.variables .states @@ -1116,8 +1469,9 @@ fn rewrite_component_member_derivatives_in_equations( member_name: &VarName, replacement: &Expression, ) { + let state_dims = None; for eq in equations { - eq.rhs = substitute_der_of_state(&eq.rhs, member_name, replacement); + eq.rhs = substitute_der_of_state(&eq.rhs, member_name, replacement, &state_dims); } } @@ -1126,8 +1480,9 @@ fn rewrite_component_member_derivatives_in_exprs( member_name: &VarName, replacement: &Expression, ) { + let state_dims = None; for expr in exprs { - *expr = substitute_der_of_state(expr, member_name, replacement); + *expr = substitute_der_of_state(expr, member_name, replacement, &state_dims); } } @@ -1220,7 +1575,7 @@ pub fn demote_exact_alias_component_states(dae: &mut Dae) -> Result Result extract_state_direct_assignment(rhs, state_name_set), + Expression::If { + branches, + else_branch, + span, + } => { + let mut state_name: Option = None; + let mut defining_branches = Vec::with_capacity(branches.len()); + for (condition, branch_expr) in branches { + let (branch_state, branch_defining_expr) = + extract_state_direct_assignment(branch_expr, state_name_set)?; + if state_name + .as_ref() + .is_some_and(|name| name != &branch_state) + { + return None; + } + state_name.get_or_insert(branch_state); + defining_branches.push((condition.clone(), branch_defining_expr)); + } + let (else_state, else_defining_expr) = + extract_state_direct_assignment(else_branch, state_name_set)?; + if state_name.as_ref().is_some_and(|name| name != &else_state) { + return None; + } + let state_name = state_name.unwrap_or(else_state); + Some(( + state_name, + Expression::If { + branches: defining_branches, + else_branch: Box::new(else_defining_expr), + span: *span, + }, + )) + } _ => None, } } +fn expression_is_zero_literal(expr: &Expression) -> bool { + match expr { + Expression::Literal { + value: Literal::Integer(0), + .. + } => true, + Expression::Literal { + value: Literal::Real(value), + .. + } => *value == 0.0, + _ => false, + } +} + fn extract_state_direct_assignment_equation( eq: &Equation, state_names: &[VarName], @@ -1495,7 +1904,7 @@ fn extract_state_direct_assignment_equation( // reject the unsafe cases (self-derivatives, derivative definitions that // feed back through the candidate). if let Some(pair) = extract_state_direct_assignment(&eq.rhs, state_name_set) - && !expr_contains_der_of(&pair.1, &pair.0) + && !expr_contains_der_of_state_or_component(&pair.1, &pair.0) { return Some(pair); } @@ -1507,6 +1916,9 @@ fn extract_state_direct_assignment_equation( if expr_contains_der_of(&eq.rhs, &state_name) { continue; } + if expr_contains_der_of_state_or_component(&eq.rhs, &state_name) { + continue; + } let Some((coef, remainder)) = split_linear_target(&eq.rhs, &state_name, eq.span) else { continue; }; @@ -1523,6 +1935,48 @@ fn extract_state_direct_assignment_equation( solved } +fn expr_contains_der_of_state_or_component(expr: &Expression, state_name: &VarName) -> bool { + let matcher = DerivativeNameMatcher::from_var_names(std::slice::from_ref(state_name)); + expr_contains_der_of_any(expr, &matcher) +} + +fn variable_dims_for_direct_demotion(dae: &Dae, state_name: &VarName) -> Option> { + dae.variables + .states + .get(state_name) + .map(|state| state.dims.clone()) + .filter(|dims| !dims.is_empty()) +} + +fn der_call_target_subscripts<'a>( + expr: &'a Expression, + state_name: &VarName, +) -> Option> { + let Expression::BuiltinCall { function, args, .. } = expr else { + return None; + }; + if *function != BuiltinFunction::Der || args.len() != 1 { + return None; + } + if expression_exact_name(&args[0]).as_deref() == Some(state_name.as_str()) { + return Some(None); + } + let Expression::VarRef { + name, subscripts, .. + } = &args[0] + else { + return expr_refers_to_var(&args[0], state_name).then_some(None); + }; + if name.var_name() != state_name { + return None; + } + if subscripts.is_empty() { + Some(None) + } else { + Some(Some(subscripts.as_slice())) + } +} + fn state_value_refs_outside_der(expr: &Expression, state_names: &[VarName]) -> Vec { let mut collector = StateValueRefCollector { state_names, @@ -1561,24 +2015,16 @@ impl ExpressionVisitor for StateValueRefCollector<'_> { } } -fn der_call_targets_state(expr: &Expression, state_name: &VarName) -> bool { - matches!( - expr, - Expression::BuiltinCall { function, args, .. } - if *function == BuiltinFunction::Der - && args.len() == 1 - && expr_refers_to_var(&args[0], state_name) - ) -} - fn substitute_der_of_state( expr: &Expression, state_name: &VarName, replacement: &Expression, + state_dims: &Option>, ) -> Expression { DerSubstitutionRewriter { state_name, replacement, + state_dims, } .rewrite_expression(expr) } @@ -1619,14 +2065,32 @@ fn mask_state_der_calls(expr: &Expression, state_name_set: &HashSet) -> struct DerSubstitutionRewriter<'a> { state_name: &'a VarName, replacement: &'a Expression, + state_dims: &'a Option>, } impl ExpressionRewriter for DerSubstitutionRewriter<'_> { fn rewrite_expression(&mut self, expr: &Expression) -> Expression { - if der_call_targets_state(expr, self.state_name) { - self.replacement.clone() - } else { - self.walk_expression(expr) + match der_call_target_subscripts(expr, self.state_name) { + Some(None) => self.replacement.clone(), + Some(Some(subscripts)) => { + let projected = self + .state_dims + .as_deref() + .and_then(|dims| static_subscript_indices(subscripts).zip(Some(dims))) + .and_then(|(indices, dims)| { + flat_index_from_indices(dims, &indices).zip(Some(dims)) + }) + .and_then(|(flat_index, dims)| { + project_flat_index_with_span( + self.replacement, + dims, + flat_index, + expr.span(), + ) + }); + projected.unwrap_or_else(|| self.walk_expression(expr)) + } + None => self.walk_expression(expr), } } } @@ -1635,6 +2099,7 @@ impl ExpressionRewriter for DerSubstitutionRewriter<'_> { struct DirectStateDemotionPlan { state_name: VarName, der_expr: Expression, + promote_der_algebraics: Vec, } #[derive(Default)] @@ -1659,7 +2124,6 @@ struct DirectDemotionRound<'a> { when_assigned_states: HashSet, non_state_unknown_names: HashSet, non_state_defining_exprs: DefiningExprIndex, - der_map: HashMap, trace: bool, } @@ -1676,10 +2140,6 @@ impl<'a> DirectDemotionRound<'a> { let non_state_unknown_names = collect_non_state_continuous_unknown_names(dae); let non_state_defining_exprs = collect_non_derivative_defining_expr_index(dae); structural_timing_done("direct_demotion.non_state_unknown_names", timer); - let timer = structural_timing_start("direct_demotion.build_relaxed_derivative_map"); - let seed_exprs = direct_demotion_derivative_seed_exprs(dae, &state_names, &state_name_set); - let der_map = build_relaxed_derivative_map_for_exprs(dae, &seed_exprs)?; - structural_timing_done("direct_demotion.build_relaxed_derivative_map", timer); Ok(Some(Self { dae, state_names, @@ -1687,7 +2147,6 @@ impl<'a> DirectDemotionRound<'a> { when_assigned_states, non_state_unknown_names, non_state_defining_exprs, - der_map, trace, })) } @@ -1768,26 +2227,6 @@ fn direct_demotion_round_context( Some((state_names, state_name_set, when_assigned_states)) } -fn direct_demotion_derivative_seed_exprs( - dae: &Dae, - state_names: &[VarName], - state_name_set: &HashSet, -) -> Vec { - dae.continuous - .equations - .iter() - .filter_map(|eq| { - let (state_name, defining_expr) = - extract_state_direct_assignment_equation(eq, state_names, state_name_set)?; - if !is_connection_equation_origin(&eq.origin) { - return Some(defining_expr); - } - connection_component_fixed_defining_expr(dae, &state_name, state_name_set) - .or(Some(defining_expr)) - }) - .collect() -} - /// Apply structural dummy-derivative reduction for constrained states. /// /// The source DAE initially marks every variable that appears under `der()` as @@ -1870,11 +2309,16 @@ fn constrained_dummy_derivative_plan( Some(DirectStateDemotionPlan { state_name: state_name.clone(), der_expr, + promote_der_algebraics: Vec::new(), }) } #[cfg(test)] mod dae_prepare_demotion_tests; +#[cfg(test)] +mod direct_demotion_piecewise_tests; +#[cfg(test)] +mod matrix_state_derivative_tests; /// Pin parameters whose compile-time values the constrained-dummy reduction /// baked into substituted derivative expressions: runtime tuning of them diff --git a/crates/rumoca-phase-structural/src/dae_prepare/state_row_reduction.rs b/crates/rumoca-phase-structural/src/dae_prepare/state_row_reduction.rs index 20fad8415..f6f16d7fc 100644 --- a/crates/rumoca-phase-structural/src/dae_prepare/state_row_reduction.rs +++ b/crates/rumoca-phase-structural/src/dae_prepare/state_row_reduction.rs @@ -41,6 +41,28 @@ fn try_match_state_to_row( fn states_with_assignable_derivative_rows(dae: &Dae, state_names: &[VarName]) -> HashSet { let bindings = structural_scalar_bindings(dae); + let mut matched_aggregate_states = HashSet::new(); + for (state_idx, state_name) in state_names.iter().enumerate() { + let Some(state) = dae.variables.states.get(state_name) else { + continue; + }; + if state.size() <= 1 { + continue; + } + let mut components = HashSet::new(); + for eq in &dae.continuous.equations { + if !state_derivative_row_is_assignable(eq, state_name, state_names, &bindings) { + continue; + } + components.extend(active_derivative_components_for_state( + &eq.rhs, state_name, &bindings, + )); + } + if components.len() == state.size() { + matched_aggregate_states.insert(state_idx); + } + } + let state_to_rows: Vec> = state_names .iter() .map(|state_name| { @@ -76,7 +98,9 @@ fn states_with_assignable_derivative_rows(dae: &Dae, state_names: &[VarName]) -> try_match_state_to_row(state_idx, &state_to_rows, &mut row_to_state, &mut seen_rows); } - row_to_state.into_iter().flatten().collect() + let mut matched_states: HashSet = row_to_state.into_iter().flatten().collect(); + matched_states.extend(matched_aggregate_states); + matched_states } fn state_derivative_row_is_assignable( @@ -158,6 +182,11 @@ impl ExpressionVisitor for ExactStateDerivativeChecker<'_> { } fn derivative_arg_matches_state(expr: &Expression, state_name: &VarName) -> bool { + if let Some(name) = derivative_arg_component_name(expr) { + return name == state_name.as_str() + || rumoca_core::parse_scalar_name(&name) + .is_some_and(|scalar| scalar.base == state_name.as_str()); + } let Some(exact_name) = expression_exact_name(expr) else { return false; }; @@ -166,6 +195,80 @@ fn derivative_arg_matches_state(expr: &Expression, state_name: &VarName) -> bool .is_some_and(|scalar| scalar.base == state_name.as_str()) } +fn derivative_arg_component_name(expr: &Expression) -> Option { + let Expression::VarRef { + name, subscripts, .. + } = expr + else { + return None; + }; + if subscripts.is_empty() { + return Some(name.as_str().to_string()); + } + let indices = static_subscript_indices(subscripts)?; + let indices: Vec = indices + .into_iter() + .map(usize::try_from) + .collect::>() + .ok()?; + Some(dae::format_subscript_key(name.as_str(), &indices)) +} + +fn active_derivative_components_for_state( + expr: &Expression, + state_name: &VarName, + bindings: &HashMap, +) -> HashSet { + let mut collector = ActiveDerivativeComponentCollector { + state_name, + bindings, + components: HashSet::new(), + }; + collector.visit_expression(expr); + collector.components +} + +struct ActiveDerivativeComponentCollector<'a> { + state_name: &'a VarName, + bindings: &'a HashMap, + components: HashSet, +} + +impl ExpressionVisitor for ActiveDerivativeComponentCollector<'_> { + fn visit_builtin_call(&mut self, function: &BuiltinFunction, args: &[Expression]) { + if *function == BuiltinFunction::Der { + if let Some(arg) = args.first() + && let Some(component_name) = derivative_arg_component_name(arg) + && rumoca_core::parse_scalar_name(&component_name) + .is_some_and(|scalar| scalar.base == self.state_name.as_str()) + { + self.components.insert(component_name); + } + return; + } + for arg in args { + self.visit_expression(arg); + } + } + + fn visit_if(&mut self, branches: &[(Expression, Expression)], else_branch: &Expression) { + for (condition, value) in branches { + match eval_static_bool(condition, self.bindings) { + Some(true) => { + self.visit_expression(value); + return; + } + Some(false) => continue, + None => { + self.visit_expression(condition); + self.visit_expression(value); + } + } + } + self.visit_expression(else_branch); + } +} + fn all_active_derivative_args_are_states( expr: &Expression, state_names: &[VarName], @@ -340,7 +443,7 @@ fn eval_static_number(expr: &Expression, bindings: &HashMap) -> Opt } } -fn expression_exact_name(expr: &Expression) -> Option { +pub(super) fn expression_exact_name(expr: &Expression) -> Option { match expr { Expression::VarRef { name, subscripts, .. @@ -533,6 +636,10 @@ pub fn index_reduce_missing_state_derivatives_once( if used_eq.contains(&idx) { return None; } + if !expr_dependency_closure_reaches_state(&eq.rhs, &defining_expr_index, state_name) + { + return None; + } if eq_contains_any_state_der_with_matcher(&eq.rhs, &state_derivative_matcher) { return None; } @@ -544,6 +651,12 @@ pub fn index_reduce_missing_state_derivatives_once( { return None; } + if is_explicit_non_state_definition(dae, eq) { + return None; + } + if is_non_state_direct_definition(dae, eq, idx, &defining_expr_index) { + return None; + } if is_indexed_state_component_alias_definition(eq, state_name) { return None; } @@ -591,6 +704,34 @@ pub fn index_reduce_missing_state_derivatives_once( Ok(changed) } +fn expr_dependency_closure_reaches_state( + expr: &Expression, + defining_expr_index: &DefiningExprIndex, + state_name: &VarName, +) -> bool { + if expr_contains_var(expr, state_name) { + return true; + } + + let mut seen = HashSet::new(); + let mut stack: Vec = collect_rhs_var_refs(expr).into_iter().collect(); + while let Some(name) = stack.pop() { + if name == *state_name { + return true; + } + if !seen.insert(name.as_str().to_string()) { + continue; + } + for defining_expr in defining_expr_candidates(defining_expr_index, &name) { + if expr_contains_var(defining_expr, state_name) { + return true; + } + stack.extend(collect_rhs_var_refs(defining_expr)); + } + } + false +} + fn is_unsliced_algebraic_definition(eq: &Equation, alg_name: &VarName) -> bool { let Expression::Binary { op, lhs, rhs, .. } = &eq.rhs else { return false; @@ -607,6 +748,84 @@ fn is_unsliced_algebraic_definition(eq: &Equation, alg_name: &VarName) -> bool { }) } +fn is_explicit_non_state_definition(dae: &Dae, eq: &Equation) -> bool { + let Some(lhs) = eq.lhs.as_ref() else { + return false; + }; + let exact_name = lhs.as_str(); + dae.variables + .algebraics + .keys() + .chain(dae.variables.outputs.keys()) + .any(|unknown| exact_name_matches_unknown_or_component(exact_name, unknown)) +} + +fn is_non_state_direct_definition( + dae: &Dae, + eq: &Equation, + equation_index: usize, + defining_expr_index: &DefiningExprIndex, +) -> bool { + let Expression::Binary { op, lhs, rhs, .. } = &eq.rhs else { + return false; + }; + if !matches!(op, OpBinary::Sub) { + return false; + } + [lhs.as_ref(), rhs.as_ref()] + .into_iter() + .filter_map(|expr| { + let exact_name = expression_non_state_unknown_exact_name(dae, expr)?; + let other = if std::ptr::eq(expr, lhs.as_ref()) { + rhs.as_ref() + } else { + lhs.as_ref() + }; + Some((exact_name, other)) + }) + .any(|(exact_name, other)| { + !non_state_unknown_has_independent_definition( + defining_expr_index, + &exact_name, + equation_index, + ) || !expr_refs_only_parameters_constants_or_time(dae, other) + }) +} + +fn expression_non_state_unknown_exact_name(dae: &Dae, expr: &Expression) -> Option { + let name = expression_exact_name(expr)?; + dae.variables + .algebraics + .keys() + .chain(dae.variables.outputs.keys()) + .any(|unknown| exact_name_matches_unknown_or_component(&name, unknown)) + .then_some(name) +} + +fn non_state_unknown_has_independent_definition( + defining_expr_index: &DefiningExprIndex, + exact_name: &str, + equation_index: usize, +) -> bool { + let base_name = rumoca_core::parse_scalar_name(exact_name) + .map(|scalar| scalar.base) + .unwrap_or(exact_name); + [exact_name, base_name].into_iter().any(|name| { + defining_expr_index.get(name).is_some_and(|candidates| { + candidates + .iter() + .any(|candidate| candidate.equation_index != equation_index) + }) + }) +} + +fn exact_name_matches_unknown_or_component(exact_name: &str, unknown: &VarName) -> bool { + exact_name == unknown.as_str() + || exact_name + .strip_prefix(unknown.as_str()) + .is_some_and(|suffix| suffix.starts_with('[')) +} + fn is_indexed_state_component_alias_definition(eq: &Equation, state_name: &VarName) -> bool { let Expression::Binary { op, lhs, rhs, .. } = &eq.rhs else { return false; @@ -804,7 +1023,7 @@ pub fn substitute_standalone_state_derivatives_in_non_ode_rows(dae: &mut Dae) -> if !expr_contains_der_of(&eq.rhs, state_name) { continue; } - eq.rhs = substitute_der_of_state(&eq.rhs, state_name, replacement); + eq.rhs = substitute_der_of_state(&eq.rhs, state_name, replacement, &None); rewritten = true; } rewritten_rows += usize::from(rewritten); @@ -842,6 +1061,20 @@ mod tests { } } + fn indexed_var_ref(name: &str, index: i64, span: Span) -> Expression { + Expression::VarRef { + name: Reference::from_var_name(VarName::new(name)), + subscripts: vec![Subscript::Index { value: index, span }], + span, + } + } + + fn test_variable(name: &str, dims: Vec) -> Variable { + let mut variable = Variable::new(VarName::new(name), test_span()); + variable.dims = dims; + variable + } + #[test] fn normalize_ode_equation_sign_uses_equation_span() { let span = test_span(); @@ -880,4 +1113,39 @@ mod tests { }; assert_eq!(actual, span); } + + #[test] + fn index_reduction_skips_indexed_algebraic_component_definition() { + let span = test_span(); + let mut dae = Dae::new(); + dae.variables + .states + .insert(VarName::new("x"), test_variable("x", vec![2])); + dae.variables + .algebraics + .insert(VarName::new("y"), test_variable("y", vec![2])); + let rhs = Expression::Binary { + op: OpBinary::Sub, + lhs: Box::new(indexed_var_ref("y", 2, span)), + rhs: Box::new(Expression::Binary { + op: OpBinary::Mul, + lhs: Box::new(Expression::Literal { + value: Literal::Real(3.0), + span, + }), + rhs: Box::new(indexed_var_ref("x", 2, span)), + span, + }), + span, + }; + dae.continuous + .equations + .push(Equation::residual(rhs.clone(), span, "test")); + + let changed = index_reduce_missing_state_derivatives_once(&mut dae) + .expect("index reduction should evaluate candidates"); + + assert_eq!(changed, 0); + assert_eq!(dae.continuous.equations[0].rhs, rhs); + } } diff --git a/crates/rumoca-phase-structural/src/dae_prepare/symbolic.rs b/crates/rumoca-phase-structural/src/dae_prepare/symbolic.rs index 4f09d2f3b..399f231d8 100644 --- a/crates/rumoca-phase-structural/src/dae_prepare/symbolic.rs +++ b/crates/rumoca-phase-structural/src/dae_prepare/symbolic.rs @@ -12,6 +12,10 @@ fn is_der_of_state(expr: &Expression, state_name: &VarName) -> bool { } fn make_binary(op: OpBinary, lhs: Expression, rhs: Expression, span: Span) -> Expression { + simplify_binary(op, lhs, rhs, span) +} + +fn make_binary_raw(op: OpBinary, lhs: Expression, rhs: Expression, span: Span) -> Expression { Expression::Binary { op, lhs: Box::new(lhs), @@ -21,6 +25,10 @@ fn make_binary(op: OpBinary, lhs: Expression, rhs: Expression, span: Span) -> Ex } fn make_unary(op: OpUnary, rhs: Expression, span: Span) -> Expression { + simplify_unary(op, rhs, span) +} + +fn make_unary_raw(op: OpUnary, rhs: Expression, span: Span) -> Expression { Expression::Unary { op, rhs: Box::new(rhs), @@ -35,6 +43,20 @@ fn real_literal(value: f64, span: Span) -> Expression { } } +fn literal_f64(expr: &Expression) -> Option { + match expr { + Expression::Literal { + value: Literal::Integer(value), + .. + } => Some(*value as f64), + Expression::Literal { + value: Literal::Real(value), + .. + } => Some(*value), + _ => None, + } +} + fn split_linear_der_target( expr: &Expression, state_name: &VarName, @@ -169,7 +191,10 @@ fn build_array_der_value(dae: &Dae, state_name: &VarName, dims: &[i64]) -> Optio array_expr_from_flat_values(values, dims) } -fn array_expr_from_flat_values(values: Vec, dims: &[i64]) -> Option { +pub(super) fn array_expr_from_flat_values( + values: Vec, + dims: &[i64], +) -> Option { match dims { [n] if *n >= 0 && *n as usize == values.len() => Some(Expression::Array { span: expression_sequence_span(&values)?, @@ -235,6 +260,73 @@ struct SymbolicDerivativeContext<'a> { } impl<'a> SymbolicDerivativeContext<'a> { + fn differentiate_builtin_call( + &self, + function: &BuiltinFunction, + args: &[Expression], + span: Span, + active_functions: &mut Vec, + ) -> Option { + if args.len() != 1 { + return None; + } + let arg = args.first()?; + let d_arg = self.differentiate(arg, active_functions)?; + if expression_is_zero_value(&d_arg) { + return match function { + BuiltinFunction::Max | BuiltinFunction::Min => Some(real_literal(0.0, span)), + _ => Some(d_arg), + }; + } + match function { + BuiltinFunction::Transpose => Some(Expression::BuiltinCall { + function: BuiltinFunction::Transpose, + args: vec![d_arg], + span, + }), + BuiltinFunction::Sin => Some(make_binary( + OpBinary::Mul, + Expression::BuiltinCall { + function: BuiltinFunction::Cos, + args: vec![arg.clone()], + span, + }, + d_arg, + span, + )), + BuiltinFunction::Cos => Some(make_unary( + OpUnary::Minus, + make_binary( + OpBinary::Mul, + Expression::BuiltinCall { + function: BuiltinFunction::Sin, + args: vec![arg.clone()], + span, + }, + d_arg, + span, + ), + span, + )), + BuiltinFunction::Sqrt => Some(make_binary( + OpBinary::Div, + d_arg, + make_binary( + OpBinary::Mul, + real_literal(2.0, span), + Expression::BuiltinCall { + function: BuiltinFunction::Sqrt, + args: vec![arg.clone()], + span, + }, + span, + ), + span, + )), + _ => None, + } + } + fn differentiate_variable( &self, name: &VarName, @@ -250,7 +342,20 @@ impl<'a> SymbolicDerivativeContext<'a> { return Some(real_literal(0.0, span)); } if !subscripts.is_empty() - && !self.dae.variables.states.contains_key(name) + && let Some(dims) = variable_dims_for_name(self.dae, name) + && let Some(indices) = static_subscript_indices(subscripts) + && let Some(flat_index) = flat_index_from_indices(&dims, &indices) + { + let scalar_name = VarName::new(dae::scalar_name_text_for_flat_index( + name.as_str(), + &dims, + flat_index, + )); + if let Some(derivative) = self.der_map.get(scalar_name.as_str()) { + return Some(derivative.clone().with_span(span)); + } + } + if !subscripts.is_empty() && let Some(derivative) = self.der_map.get(name.as_str()) && let Some(dims) = variable_dims_for_name(self.dae, name) && let Some(indices) = static_subscript_indices(subscripts) @@ -277,7 +382,26 @@ impl<'a> SymbolicDerivativeContext<'a> { span, }); } - self.der_map.get(name.as_str()).cloned() + if let Some(derivative) = self.der_map.get(name.as_str()) { + return Some(derivative.clone()); + } + if self.dae.variables.states.contains_key(name) + || self.dae.variables.algebraics.get(name).is_some_and(|var| { + state_select_rank(var.state_select) + >= state_select_rank(rumoca_core::StateSelect::Prefer) + }) + { + return Some(Expression::BuiltinCall { + function: BuiltinFunction::Der, + args: vec![Expression::VarRef { + name: rumoca_core::Reference::from_var_name(name.clone()), + subscripts: Vec::new(), + span, + }], + span, + }); + } + None } fn differentiate_binary( @@ -336,6 +460,29 @@ impl<'a> SymbolicDerivativeContext<'a> { let denom = make_binary(OpBinary::Mul, rhs.clone(), rhs.clone(), span); Some(make_binary(OpBinary::Div, numer, denom, span)) } + OpBinary::Exp | OpBinary::ExpElem => { + let exponent = literal_f64(rhs)?; + if exponent == 0.0 { + return Some(real_literal(0.0, span)); + } + if exponent == 1.0 { + return self.differentiate(lhs, active_functions); + } + let power = make_binary( + op.clone(), + lhs.clone(), + real_literal(exponent - 1.0, span), + span, + ); + let scaled_power = + make_binary(OpBinary::Mul, real_literal(exponent, span), power, span); + Some(make_binary( + OpBinary::Mul, + scaled_power, + self.differentiate(lhs, active_functions)?, + span, + )) + } _ => None, } } @@ -421,6 +568,7 @@ impl<'a> SymbolicDerivativeContext<'a> { name: &VarName, args: &[Expression], is_constructor: bool, + span: Span, active_functions: &mut Vec, ) -> Option { if is_constructor { @@ -433,6 +581,11 @@ impl<'a> SymbolicDerivativeContext<'a> { if !function.pure || function.external.is_some() || function.outputs.len() != 1 { return None; } + if let Some(derivative_call) = + self.differentiate_function_call_with_annotation(function, args, span, active_functions) + { + return Some(derivative_call); + } active_functions.push(name.clone()); let Some(output_expr) = function_output_expression(function, args) else { active_functions.pop(); @@ -443,6 +596,73 @@ impl<'a> SymbolicDerivativeContext<'a> { derivative } + fn differentiate_function_call_with_annotation( + &self, + function: &rumoca_core::Function, + args: &[Expression], + span: Span, + active_functions: &mut Vec, + ) -> Option { + let annotation = function + .derivatives + .iter() + .find(|annotation| annotation.order == 1)?; + let derivative_name = + self.resolve_derivative_function_name(&annotation.derivative_function)?; + let (named, positional) = split_named_and_positional_args(args)?; + let mut positional_idx = 0usize; + let mut actuals = Vec::with_capacity(function.inputs.len()); + for input in &function.inputs { + let actual = named.get(input.name.as_str()).cloned().or_else(|| { + let actual = positional.get(positional_idx).cloned(); + positional_idx += usize::from(actual.is_some()); + actual + }); + actuals.push(actual.or_else(|| input.default.clone())?); + } + + let mut derivative_args = actuals.clone(); + for (input, actual) in function.inputs.iter().zip(actuals.iter()) { + if annotation + .no_derivative + .iter() + .any(|name| name == &input.name) + { + continue; + } + let derivative = if annotation + .zero_derivative + .iter() + .any(|name| name == &input.name) + { + real_literal(0.0, actual.span().unwrap_or(span)) + } else { + self.differentiate(actual, active_functions)? + }; + derivative_args.push(derivative); + } + + Some(Expression::FunctionCall { + name: rumoca_core::Reference::from_var_name(derivative_name), + args: derivative_args, + is_constructor: false, + span, + }) + } + + fn resolve_derivative_function_name(&self, derivative_function: &str) -> Option { + let exact = VarName::new(derivative_function.to_string()); + if self.dae.symbols.functions.contains_key(&exact) { + return Some(exact); + } + self.dae + .symbols + .functions + .keys() + .find(|name| name.as_str().ends_with(derivative_function)) + .cloned() + } + fn differentiate( &self, expr: &Expression, @@ -478,17 +698,41 @@ impl<'a> SymbolicDerivativeContext<'a> { is_matrix: *is_matrix, span: *span, }), + Expression::Index { + base, + subscripts, + span, + } => self.differentiate_index(base, subscripts, *span, active_functions), + Expression::FieldAccess { base, field, span } => { + if let Some(name) = self.canonical_field_access_var_name(base, field) { + return self.differentiate_variable(&name, &[], *span); + } + if let Some(projected) = + self.project_field_expression(base, field, active_functions) + { + return self.differentiate(&projected, active_functions); + } + None + } Expression::FunctionCall { name, args, is_constructor, - .. + span, } => self.differentiate_function_call( name.var_name(), args, *is_constructor, + *span, active_functions, ), + Expression::BuiltinCall { + function, + args, + span, + } if *function != BuiltinFunction::Der => { + self.differentiate_builtin_call(function, args, *span, active_functions) + } // d/dt(der(X)) — a higher-order derivative (successive `Der` blocks, // or a relative acceleration `a = der(der(phi))`). `der(X)` is X's // first time-derivative; differentiate that expression to climb one @@ -515,6 +759,39 @@ impl<'a> SymbolicDerivativeContext<'a> { } } + fn differentiate_index( + &self, + base: &Expression, + subscripts: &[rumoca_core::Subscript], + span: rumoca_core::Span, + active_functions: &mut Vec, + ) -> Option { + if let Some(base_dims) = expression_dims(base, self.dae) + && let Some(indices) = static_subscript_indices(subscripts) + && let Some(flat_index) = flat_index_from_indices(&base_dims, &indices) + && let Some(projected) = project_flat_index(base, &base_dims, flat_index) + { + return self.differentiate(&projected, active_functions); + } + let d_base = self.differentiate(base, active_functions)?; + if expression_is_zero_value(&d_base) { + return Some(real_literal(0.0, span)); + } + let Some(base_dims) = expression_dims(&d_base, self.dae) else { + return Some(Expression::Index { + base: Box::new(d_base), + subscripts: subscripts.to_vec(), + span, + }); + }; + if base_dims.is_empty() { + return Some(d_base); + } + let indices = static_subscript_indices(subscripts)?; + let flat_index = flat_index_from_indices(&base_dims, &indices)?; + project_flat_index_with_span(&d_base, &base_dims, flat_index, Some(span)) + } + /// Differentiate `der(arg)` one order higher: take `arg`'s first derivative /// and differentiate it again. Bounded by `der_order` in the caller. fn differentiate_der_call( @@ -539,6 +816,150 @@ impl<'a> SymbolicDerivativeContext<'a> { }; self.differentiate(&first_derivative, active_functions) } + + fn canonical_field_access_var_name(&self, base: &Expression, field: &str) -> Option { + field_access_candidate_var_names(base, field) + .into_iter() + .find(|name| self.variable_or_derivative_exists(name)) + } + + fn variable_or_derivative_exists(&self, name: &VarName) -> bool { + self.der_map.contains_key(name.as_str()) + || self.dae.variables.states.contains_key(name) + || self.dae.variables.algebraics.contains_key(name) + || self.dae.variables.outputs.contains_key(name) + || self.dae.variables.parameters.contains_key(name) + || self.dae.variables.constants.contains_key(name) + } + + fn project_field_expression( + &self, + expr: &Expression, + field: &str, + active_functions: &mut Vec, + ) -> Option { + match expr { + Expression::If { + branches, + else_branch, + span, + } => { + let projected_branches = branches + .iter() + .map(|(cond, value)| { + Some(( + cond.clone(), + self.project_field_expression(value, field, active_functions)?, + )) + }) + .collect::>>()?; + Some(Expression::If { + branches: projected_branches, + else_branch: Box::new(self.project_field_expression( + else_branch, + field, + active_functions, + )?), + span: *span, + }) + } + Expression::FunctionCall { + name, + args, + is_constructor, + .. + } if *is_constructor => self.project_constructor_field(name.var_name(), args, field), + Expression::FunctionCall { + name, + args, + is_constructor, + .. + } => { + if *is_constructor + || active_functions + .iter() + .any(|active| active == name.var_name()) + { + return None; + } + let function = self.dae.symbols.functions.get(name.var_name())?; + if !function.pure || function.external.is_some() || function.outputs.len() != 1 { + return None; + } + active_functions.push(name.var_name().clone()); + let output_expr = function_output_expression(function, args); + let projected = output_expr.and_then(|output_expr| { + self.project_field_expression(&output_expr, field, active_functions) + }); + active_functions.pop(); + projected + } + _ => None, + } + } + + fn project_constructor_field( + &self, + constructor_name: &VarName, + args: &[Expression], + field: &str, + ) -> Option { + let constructor = self.dae.symbols.functions.get(constructor_name)?; + if !constructor.is_constructor { + return None; + } + constructor + .inputs + .iter() + .position(|input| input.name.as_str() == field) + .and_then(|idx| args.get(idx).cloned()) + } +} + +pub(super) fn field_access_candidate_var_names(base: &Expression, field: &str) -> Vec { + let Some((base_name, subscripts)) = field_access_base_var_ref(base) else { + return Vec::new(); + }; + let Some(indices) = static_subscript_indices(&subscripts) else { + return Vec::new(); + }; + if indices.is_empty() { + return vec![VarName::new(format!("{}.{}", base_name.as_str(), field))]; + } + let index_text = indices + .iter() + .map(i64::to_string) + .collect::>() + .join(","); + vec![ + VarName::new(format!("{}[{}].{}", base_name.as_str(), index_text, field)), + VarName::new(format!("{}.{}[{}]", base_name.as_str(), field, index_text)), + ] +} + +fn field_access_base_var_ref(base: &Expression) -> Option<(VarName, Vec)> { + match base { + Expression::VarRef { + name, subscripts, .. + } => Some((name.var_name().clone(), subscripts.clone())), + Expression::Index { + base, subscripts, .. + } => { + let Expression::VarRef { + name, + subscripts: base_subscripts, + .. + } = base.as_ref() + else { + return None; + }; + let mut combined = Vec::with_capacity(base_subscripts.len() + subscripts.len()); + combined.extend_from_slice(base_subscripts); + combined.extend_from_slice(subscripts); + Some((name.var_name().clone(), combined)) + } + _ => None, + } } pub(super) fn symbolic_time_derivative( @@ -546,12 +967,265 @@ pub(super) fn symbolic_time_derivative( dae: &Dae, der_map: &HashMap, ) -> Option { - SymbolicDerivativeContext { + let derivative = SymbolicDerivativeContext { dae, der_map, der_order: std::cell::Cell::new(0), } - .differentiate(expr, &mut Vec::new()) + .differentiate(expr, &mut Vec::new())?; + Some(simplify_symbolic_derivative(derivative)) +} + +fn expression_is_zero_value(expr: &Expression) -> bool { + match expr { + Expression::Unary { + op: OpUnary::Minus | OpUnary::DotMinus | OpUnary::Plus | OpUnary::DotPlus, + rhs, + .. + } => expression_is_zero_value(rhs), + Expression::Binary { op, lhs, rhs, .. } => match op { + OpBinary::Add | OpBinary::AddElem | OpBinary::Sub | OpBinary::SubElem => { + expression_is_zero_value(lhs) && expression_is_zero_value(rhs) + } + OpBinary::Mul | OpBinary::MulElem => { + expression_is_zero_value(lhs) || expression_is_zero_value(rhs) + } + OpBinary::Div | OpBinary::DivElem => expression_is_zero_value(lhs), + _ => false, + }, + Expression::Literal { + value: Literal::Integer(0), + .. + } => true, + Expression::Literal { + value: Literal::Real(value), + .. + } => *value == 0.0, + Expression::Array { elements, .. } => elements.iter().all(expression_is_zero_value), + Expression::Index { base, .. } => expression_is_zero_value(base), + _ => false, + } +} + +fn simplify_symbolic_derivative(expr: Expression) -> Expression { + match expr { + Expression::Binary { op, lhs, rhs, span } => simplify_binary( + op, + simplify_symbolic_derivative(*lhs), + simplify_symbolic_derivative(*rhs), + span, + ), + Expression::Unary { op, rhs, span } => { + simplify_unary(op, simplify_symbolic_derivative(*rhs), span) + } + Expression::If { + branches, + else_branch, + span, + } => { + let branches = branches + .into_iter() + .map(|(cond, value)| (cond, simplify_symbolic_derivative(value))) + .collect::>(); + let else_branch = simplify_symbolic_derivative(*else_branch); + if expression_is_zero_value(&else_branch) + && branches + .iter() + .all(|(_, value)| expression_is_zero_value(value)) + { + return real_literal(0.0, span); + } + if branches.iter().all(|(_, value)| value == &else_branch) { + return else_branch; + } + Expression::If { + branches, + else_branch: Box::new(else_branch), + span, + } + } + Expression::Array { + elements, + is_matrix, + span, + } => { + let elements = elements + .into_iter() + .map(simplify_symbolic_derivative) + .collect::>(); + Expression::Array { + elements, + is_matrix, + span, + } + } + Expression::Index { + base, + subscripts, + span, + } => simplify_symbolic_derivative_index(*base, subscripts, span), + Expression::BuiltinCall { + function, + args, + span, + } => { + let args = args + .into_iter() + .map(simplify_symbolic_derivative) + .collect::>(); + if matches!(function, BuiltinFunction::Max | BuiltinFunction::Min) + && args.iter().all(expression_is_zero_value) + { + return real_literal(0.0, span); + } + Expression::BuiltinCall { + function, + args, + span, + } + } + Expression::FunctionCall { + name, + args, + is_constructor, + span, + } => Expression::FunctionCall { + name, + args: args.into_iter().map(simplify_symbolic_derivative).collect(), + is_constructor, + span, + }, + _ => expr, + } +} + +fn simplify_symbolic_derivative_index( + base: Expression, + subscripts: Vec, + span: Span, +) -> Expression { + let base = simplify_symbolic_derivative(base); + if expression_is_zero_value(&base) { + return real_literal(0.0, span); + } + if syntactically_scalar_for_projection(&base) { + return base; + } + Expression::Index { + base: Box::new(base), + subscripts, + span, + } +} + +fn simplify_binary(op: OpBinary, lhs: Expression, rhs: Expression, span: Span) -> Expression { + match op { + OpBinary::Add | OpBinary::AddElem => { + if expression_is_zero_value(&lhs) { + return rhs; + } + if expression_is_zero_value(&rhs) { + return lhs; + } + } + OpBinary::Sub | OpBinary::SubElem => { + if expression_is_zero_value(&rhs) { + return lhs; + } + if expression_is_zero_value(&lhs) { + return make_unary(OpUnary::Minus, rhs, span); + } + } + OpBinary::Mul | OpBinary::MulElem => { + if expression_is_zero_value(&lhs) || expression_is_zero_value(&rhs) { + return real_literal(0.0, span); + } + if expression_is_one_value(&lhs) { + return rhs; + } + if expression_is_one_value(&rhs) { + return lhs; + } + } + OpBinary::Div | OpBinary::DivElem => { + if expression_is_zero_value(&lhs) { + return real_literal(0.0, span); + } + if expression_is_one_value(&rhs) { + return lhs; + } + } + OpBinary::Exp | OpBinary::ExpElem => { + if expression_is_one_value(&rhs) { + return lhs; + } + if expression_is_zero_value(&rhs) { + return real_literal(1.0, span); + } + } + _ => {} + } + if let (Some(lhs), Some(rhs)) = (literal_f64(&lhs), literal_f64(&rhs)) + && let Some(value) = fold_numeric_binary(&op, lhs, rhs) + { + return real_literal(value, span); + } + make_binary_raw(op, lhs, rhs, span) +} + +fn simplify_unary(op: OpUnary, rhs: Expression, span: Span) -> Expression { + match op { + OpUnary::Plus | OpUnary::DotPlus => return rhs, + OpUnary::Minus | OpUnary::DotMinus => { + if expression_is_zero_value(&rhs) { + return real_literal(0.0, span); + } + if let Some(value) = literal_f64(&rhs) { + return real_literal(-value, span); + } + if let Expression::Unary { + op: OpUnary::Minus | OpUnary::DotMinus, + rhs, + .. + } = rhs + { + return *rhs; + } + } + _ => {} + } + make_unary_raw(op, rhs, span) +} + +fn expression_is_one_value(expr: &Expression) -> bool { + match expr { + Expression::Unary { + op: OpUnary::Plus | OpUnary::DotPlus, + rhs, + .. + } => expression_is_one_value(rhs), + Expression::Literal { + value: Literal::Integer(1), + .. + } => true, + Expression::Literal { + value: Literal::Real(value), + .. + } => *value == 1.0, + _ => false, + } +} + +fn fold_numeric_binary(op: &OpBinary, lhs: f64, rhs: f64) -> Option { + match op { + OpBinary::Add | OpBinary::AddElem => Some(lhs + rhs), + OpBinary::Sub | OpBinary::SubElem => Some(lhs - rhs), + OpBinary::Mul | OpBinary::MulElem => Some(lhs * rhs), + OpBinary::Div | OpBinary::DivElem if rhs != 0.0 => Some(lhs / rhs), + OpBinary::Exp | OpBinary::ExpElem if !(lhs == 0.0 && rhs == 0.0) => Some(lhs.powf(rhs)), + _ => None, + } + .filter(|value| value.is_finite()) } fn function_output_expression( @@ -665,25 +1339,67 @@ fn scalar_array_element(expr: &Expression) -> Option { } } -fn expression_dims(expr: &Expression, dae: &Dae) -> Option> { +pub(super) fn expression_dims(expr: &Expression, dae: &Dae) -> Option> { match expr { Expression::VarRef { name, subscripts, .. - } if subscripts.is_empty() => variable_dims_for_name(dae, name.var_name()), + } if subscripts.is_empty() => variable_dims_for_name_including_scalar(dae, name.var_name()), + Expression::Literal { .. } => Some(Vec::new()), Expression::Array { elements, is_matrix, .. } => array_expression_dims(elements, *is_matrix), + Expression::Binary { op, lhs, rhs, .. } => { + let lhs_dims = expression_dims(lhs, dae)?; + let rhs_dims = expression_dims(rhs, dae)?; + combine_binary_expression_dims(op, &lhs_dims, &rhs_dims) + } Expression::BuiltinCall { function: BuiltinFunction::Der, args, .. } => args.first().and_then(|arg| expression_dims(arg, dae)), + Expression::BuiltinCall { + function: BuiltinFunction::Transpose, + args, + .. + } if args.len() == 1 => { + let dims = expression_dims(&args[0], dae)?; + match dims.as_slice() { + [rows, cols] => Some(vec![*cols, *rows]), + [n] => Some(vec![*n]), + [] => Some(Vec::new()), + _ => None, + } + } + Expression::Index { + base, subscripts, .. + } => { + let base_dims = expression_dims(base, dae)?; + let indices = static_subscript_indices(subscripts)?; + if base_dims.len() == indices.len() { + flat_index_from_indices(&base_dims, &indices)?; + return Some(Vec::new()); + } + None + } _ => None, } } +fn variable_dims_for_name_including_scalar(dae: &Dae, name: &VarName) -> Option> { + dae.variables + .states + .get(name) + .or_else(|| dae.variables.algebraics.get(name)) + .or_else(|| dae.variables.outputs.get(name)) + .or_else(|| dae.variables.inputs.get(name)) + .or_else(|| dae.variables.parameters.get(name)) + .or_else(|| dae.variables.constants.get(name)) + .map(|var| var.dims.clone()) +} + fn variable_dims_for_name(dae: &Dae, name: &VarName) -> Option> { dae.variables .states @@ -697,6 +1413,42 @@ fn variable_dims_for_name(dae: &Dae, name: &VarName) -> Option> { .filter(|dims| !dims.is_empty()) } +fn combine_binary_expression_dims(op: &OpBinary, lhs: &[i64], rhs: &[i64]) -> Option> { + match op { + OpBinary::Add | OpBinary::AddElem | OpBinary::Sub | OpBinary::SubElem => { + combine_additive_expression_dims(lhs, rhs) + } + OpBinary::Mul | OpBinary::MulElem => combine_multiplicative_expression_dims(lhs, rhs), + OpBinary::Div | OpBinary::DivElem if rhs.is_empty() => Some(lhs.to_vec()), + _ => None, + } +} + +fn combine_additive_expression_dims(lhs: &[i64], rhs: &[i64]) -> Option> { + if lhs == rhs { + return Some(lhs.to_vec()); + } + if lhs.is_empty() { + return Some(rhs.to_vec()); + } + if rhs.is_empty() { + return Some(lhs.to_vec()); + } + None +} + +fn combine_multiplicative_expression_dims(lhs: &[i64], rhs: &[i64]) -> Option> { + match (lhs, rhs) { + ([], []) => Some(Vec::new()), + ([], dims) | (dims, []) => Some(dims.to_vec()), + ([a], [b]) if a == b => Some(Vec::new()), + ([rows, cols], [n]) if cols == n => Some(vec![*rows]), + ([n], [rows, cols]) if n == rows => Some(vec![*cols]), + ([a_rows, a_cols], [b_rows, b_cols]) if a_cols == b_rows => Some(vec![*a_rows, *b_cols]), + _ => None, + } +} + fn array_expression_dims(elements: &[Expression], is_matrix: bool) -> Option> { if !is_matrix { return Some(vec![elements.len() as i64]); @@ -717,12 +1469,15 @@ fn projection_span(expr: &Expression, fallback_span: Option) -> Option, ) -> Option { + if syntactically_scalar_for_projection(expr) { + return Some(expr.clone()); + } match expr { Expression::VarRef { name, @@ -813,6 +1568,29 @@ fn project_flat_index_with_span( } } +fn syntactically_scalar_for_projection(expr: &Expression) -> bool { + match expr { + Expression::Literal { .. } => true, + Expression::VarRef { subscripts, .. } => !subscripts.is_empty(), + Expression::Index { .. } => true, + Expression::Unary { rhs, .. } => syntactically_scalar_for_projection(rhs), + Expression::Binary { lhs, rhs, .. } => { + syntactically_scalar_for_projection(lhs) && syntactically_scalar_for_projection(rhs) + } + Expression::BuiltinCall { + function: + BuiltinFunction::Sin + | BuiltinFunction::Cos + | BuiltinFunction::Sqrt + | BuiltinFunction::Max + | BuiltinFunction::Min, + args, + .. + } => args.iter().all(syntactically_scalar_for_projection), + _ => false, + } +} + fn generated_index_subscripts( indices: Vec, span: Span, @@ -830,7 +1608,7 @@ fn generated_index_subscripts( .collect() } -fn static_subscript_indices(subscripts: &[Subscript]) -> Option> { +pub(super) fn static_subscript_indices(subscripts: &[Subscript]) -> Option> { subscripts .iter() .map(|subscript| match subscript { @@ -851,7 +1629,7 @@ fn static_subscript_indices(subscripts: &[Subscript]) -> Option> { .collect() } -fn flat_index_from_indices(dims: &[i64], indices: &[i64]) -> Option { +pub(super) fn flat_index_from_indices(dims: &[i64], indices: &[i64]) -> Option { if dims.len() != indices.len() || dims.is_empty() { return None; } @@ -904,12 +1682,13 @@ pub(super) fn expand_der_in_expr_full( } if args.len() == 1 => { let arg = &args[0]; match arg { + Expression::VarRef { .. } if der_arg_refers_to_retained_state(arg, state_names) => { + expr.clone() + } Expression::VarRef { name, subscripts, .. } if subscripts.is_empty() => { - if state_names.contains(name.as_str()) { - expr.clone() - } else if let Some(deriv) = der_map.get(name.as_str()) { + if let Some(deriv) = der_map.get(name.as_str()) { deriv.clone() } else { expr.clone() @@ -1008,6 +1787,15 @@ pub(super) fn expand_der_in_expr_full( } } +fn der_arg_refers_to_retained_state(expr: &Expression, state_names: &HashSet) -> bool { + let Expression::VarRef { name, .. } = expr else { + return false; + }; + state_names.contains(name.as_str()) + || rumoca_core::parse_scalar_name(name.as_str()) + .is_some_and(|scalar| state_names.contains(scalar.base)) +} + pub(super) fn truncate_debug(s: &str, max_chars: usize) -> String { if s.chars().count() <= max_chars { return s.to_string(); @@ -1043,6 +1831,20 @@ mod tests { } } + fn var_ref_idx(name: &str, idx: i64, span: Span) -> Expression { + Expression::VarRef { + name: rumoca_core::Reference::new(name), + subscripts: vec![Subscript::Index { value: idx, span }], + span, + } + } + + fn test_variable(name: &str, dims: Vec) -> Variable { + let mut variable = Variable::new(VarName::new(name), test_span()); + variable.dims = dims; + variable + } + fn has_single_index_with_span(expr: &Expression, expected_span: Span) -> bool { let Expression::VarRef { subscripts, .. } = expr else { return false; @@ -1090,4 +1892,55 @@ mod tests { "projected binary should index both operands with the source span" ); } + + #[test] + fn symbolic_derivative_projects_indexed_vector_without_indexing_scalar_factor() { + let span = test_span(); + let mut dae = Dae::new(); + dae.variables + .states + .insert(VarName::new("omega"), test_variable("omega", vec![3])); + dae.variables + .algebraics + .insert(VarName::new("M_body"), test_variable("M_body", vec![3])); + + let mut der_map = HashMap::new(); + der_map.insert( + "omega".to_string(), + Expression::Array { + elements: vec![ + var_ref_idx("M_body", 1, span), + var_ref_idx("M_body", 2, span), + var_ref_idx("M_body", 3, span), + ], + is_matrix: false, + span, + }, + ); + let expr = Expression::Index { + base: Box::new(Expression::Binary { + op: OpBinary::MulElem, + lhs: Box::new(real_literal(-0.02, span)), + rhs: Box::new(var_ref("omega", span)), + span, + }), + subscripts: vec![Subscript::Index { value: 3, span }], + span, + }; + + let derivative = + symbolic_time_derivative(&expr, &dae, &der_map).expect("vector index derivative"); + + assert!( + matches!( + &derivative, + Expression::Binary { lhs, rhs, .. } + if matches!(lhs.as_ref(), Expression::Literal { value: Literal::Real(value), .. } if *value == -0.02) + && matches!(rhs.as_ref(), Expression::VarRef { name, subscripts, .. } + if name.as_str() == "M_body" + && matches!(subscripts.as_slice(), [Subscript::Index { value: 3, .. }])) + ), + "derivative should be -0.02 * M_body[3], got {derivative:?}" + ); + } } diff --git a/crates/rumoca-phase-structural/src/eliminate/aggregate_alias.rs b/crates/rumoca-phase-structural/src/eliminate/aggregate_alias.rs index 1188f9004..3ab6930f3 100644 --- a/crates/rumoca-phase-structural/src/eliminate/aggregate_alias.rs +++ b/crates/rumoca-phase-structural/src/eliminate/aggregate_alias.rs @@ -24,12 +24,46 @@ pub(super) fn aggregate_variable_fully_resolved( flat_index, )); if !resolved.contains(&scalar_name) { - return Ok(false); + return Ok(embedded_component_elements_fully_resolved( + name, + scalar_count, + resolved, + )); } } + if var.dims.len() == 1 { + return Ok(false); + } Ok(true) } +fn embedded_component_elements_fully_resolved( + name: &VarName, + scalar_count: usize, + resolved: &HashSet, +) -> bool { + let resolved_elements = resolved + .iter() + .filter(|resolved_name| { + scalarized_leaf_base_name(resolved_name).is_some_and(|base| base == name.as_str()) + }) + .count(); + resolved_elements >= scalar_count +} + +fn scalarized_leaf_base_name(name: &VarName) -> Option<&str> { + let text = name.as_str(); + let open = text.rfind('[')?; + if !text.ends_with(']') + || !text[open + 1..text.len() - 1] + .chars() + .all(|ch| ch.is_ascii_digit() || ch == ',') + { + return None; + } + Some(&text[..open]) +} + pub(super) fn is_scalarized_element_of_aggregate( dae: &Dae, name: &VarName, @@ -37,9 +71,6 @@ pub(super) fn is_scalarized_element_of_aggregate( let Some(scalar) = rumoca_core::parse_scalar_name(name.as_str()) else { return Ok(false); }; - if DaeVariableScope::new(dae).exact(name).is_some() { - return Ok(false); - } let base_name = VarName::new(scalar.base); let Some(base_var) = DaeVariableScope::new(dae).exact(&base_name) else { return Ok(false); diff --git a/crates/rumoca-phase-structural/src/eliminate/boundary_scan.rs b/crates/rumoca-phase-structural/src/eliminate/boundary_scan.rs index 860975a26..44fa347a9 100644 --- a/crates/rumoca-phase-structural/src/eliminate/boundary_scan.rs +++ b/crates/rumoca-phase-structural/src/eliminate/boundary_scan.rs @@ -2,7 +2,6 @@ use super::*; pub(super) struct BoundaryScanCtx<'a> { pub(super) dae: &'a Dae, - pub(super) state_names: &'a [VarName], pub(super) unknown_index: &'a BoundaryUnknownIndex<'a>, pub(super) state_derivative_matcher: &'a DerivativeNameMatcher, pub(super) runtime_protected_unknowns: &'a IndexSet, @@ -34,13 +33,48 @@ impl BoundaryScanState { var_name: VarName, solution: Expression, ) -> Result<(), StructuralError> { - self.substitutions - .push(substitution_for_var(dae, var_name.clone(), solution)?); + let substitution = substitution_for_var(dae, var_name.clone(), solution)?; + if self.substitution_would_overgrow_remaining_equations(dae, eq_idx, &substitution)? { + return Ok(()); + } + self.substitutions.push(substitution); self.eliminated_eq_indices.push(eq_idx); self.eliminated_eq_flags[eq_idx] = true; self.resolved.insert(var_name); Ok(()) } + + fn substitution_would_overgrow_remaining_equations( + &self, + dae: &Dae, + eliminated_eq_idx: usize, + substitution: &Substitution, + ) -> Result { + for (eq_idx, equation) in dae.continuous.equations.iter().enumerate() { + if eq_idx == eliminated_eq_idx || self.eliminated_eq_flags[eq_idx] { + continue; + } + let expr = equation_analysis_expr(equation); + if !expr_contains_var(&expr, &substitution.var_name) { + continue; + } + let Some(rhs) = + apply_substitutions_for_symbolic_candidate(&expr, &self.substitutions, dae)? + else { + return Ok(true); + }; + if apply_substitutions_for_symbolic_candidate( + &rhs, + std::slice::from_ref(substitution), + dae, + )? + .is_none() + { + return Ok(true); + } + } + Ok(false) + } } pub(super) fn scan_boundary_equations( @@ -64,9 +98,26 @@ fn scan_boundary_equation( } let equation = &ctx.dae.continuous.equations[eq_idx]; let expr = equation_analysis_expr(equation); - let eq_rhs = apply_substitutions_in_order(&expr, &state.substitutions)?; + let Some(eq_rhs) = + apply_substitutions_for_symbolic_candidate(&expr, &state.substitutions, ctx.dae)? + else { + return Ok(()); + }; let is_connection_eq = equation.origin.starts_with("connection equation:"); + if is_connection_eq && connection_refs_unanchored_scalarized_aggregate(ctx.dae, &eq_rhs)? { + return Ok(()); + } let live = find_live_scalar_unknowns(&eq_rhs, ctx.unknown_index, &state.resolved)?; + if is_connection_eq + && let Some((var_name, solution)) = scalar_connection_alias_for_elimination( + ctx.dae, + &eq_rhs, + ctx.runtime_protected_unknowns, + ctx.runtime_defined_discrete_targets, + )? + { + return state.push_solution(ctx.dae, eq_idx, var_name, solution); + } if should_skip_connection_equation( ctx.dae, &eq_rhs, @@ -76,7 +127,15 @@ fn scan_boundary_equation( ) { return Ok(()); } + if should_preserve_runtime_sensitive_continuous_assignment(ctx.dae, &eq_rhs) { + return Ok(()); + } let has_state_derivative = expr_contains_der_of_any(&eq_rhs, ctx.state_derivative_matcher); + let indexed_multiscalar_slice_equation = + expr_contains_indexed_multiscalar_slice_ref(&eq_rhs, ctx.dae)?; + if indexed_multiscalar_slice_equation { + return Ok(()); + } if let Some((var_name, solution)) = aggregate_alias_for_elimination( ctx.dae, &eq_rhs, @@ -88,7 +147,6 @@ fn scan_boundary_equation( if live.is_empty() { let mut zero_unknown_ctx = ZeroUnknownEliminationCtx { dae: ctx.dae, - state_names: ctx.state_names, unknown_index: ctx.unknown_index, resolved: &state.resolved, runtime_protected_unknowns: ctx.runtime_protected_unknowns, @@ -104,21 +162,37 @@ fn scan_boundary_equation( &mut zero_unknown_ctx, ); } - if is_flow_equation_origin(&equation.origin) - && expr_contains_indexed_multiscalar_ref(&eq_rhs, ctx.dae)? + let indexed_flow_equation = is_flow_equation_origin(&equation.origin) + && expr_contains_indexed_multiscalar_ref(&eq_rhs, ctx.dae)?; + let can_eliminate_pairwise_flow_alias = equation.origin.starts_with("flow sum equation:") + && can_eliminate_pairwise_flow_alias(ctx.dae, &eq_rhs, &live); + if (indexed_flow_equation || indexed_multiscalar_slice_equation) + && !can_eliminate_pairwise_flow_alias { return Ok(()); } - if let Some((var_name, solution)) = choose_solvable_unknown_for_elimination( - ctx.dae, + let choice_ctx = EliminationChoiceContext { + dae: ctx.dae, eq_idx, - &eq_rhs, - &live, has_state_derivative, - ctx.runtime_protected_unknowns, - ctx.direct_definitions, - )? { + runtime_protected_unknowns: ctx.runtime_protected_unknowns, + direct_definitions: ctx.direct_definitions, + allow_multi_live_trivial_alias: can_eliminate_pairwise_flow_alias, + }; + if let Some((var_name, solution)) = + choose_solvable_unknown_for_elimination(&choice_ctx, &eq_rhs, &live)? + { state.push_solution(ctx.dae, eq_idx, var_name, solution)?; } Ok(()) } + +fn can_eliminate_pairwise_flow_alias(dae: &Dae, eq_rhs: &Expression, live: &[VarName]) -> bool { + live.len() <= 2 + && live.iter().any(|candidate| { + try_solve_for_unknown_in_dae(dae, eq_rhs, candidate).is_some_and(|solution| { + !expr_contains_unknown_in_dae(dae, &solution, candidate) + && is_trivial_alias_in_dae(dae, &solution) + }) + }) +} diff --git a/crates/rumoca-phase-structural/src/eliminate/connection_policy.rs b/crates/rumoca-phase-structural/src/eliminate/connection_policy.rs index 1a08795b4..143f61495 100644 --- a/crates/rumoca-phase-structural/src/eliminate/connection_policy.rs +++ b/crates/rumoca-phase-structural/src/eliminate/connection_policy.rs @@ -1,8 +1,9 @@ use std::collections::HashSet; +use super::runtime_protection::assignment_var_ref_name; use super::{ - Dae, Expression, OpBinary, VarName, full_var_ref, runtime_partition_or_event_refs_var, - unknown_is_fixed, + Dae, Expression, OpBinary, VarName, collect_var_ref_nodes, exact_reference_expr_name_in_dae, + runtime_partition_or_event_refs_var, unknown_is_fixed, }; pub(super) fn should_skip_connection_equation( @@ -22,20 +23,36 @@ pub(super) fn should_skip_connection_equation( if touches_runtime_discrete_path { return true; } - if !connection_alias_refs_are_continuous_scalar_algebraics( + if live.len() == 1 + && can_eliminate_scalar_connection_alias_var( + dae, + &live[0], + runtime_defined_discrete_targets, + ) + { + return false; + } + let alias_refs_are_continuous = connection_alias_refs_are_continuous_scalar_algebraics( dae, eq_rhs, runtime_defined_discrete_targets, - ) { + ); + let single_live_has_known_boundary = single_live_connection_alias_has_known_boundary( + dae, + eq_rhs, + live, + runtime_defined_discrete_targets, + ); + let can_eliminate_multi_live = can_eliminate_continuous_connection_alias( + dae, + eq_rhs, + live, + runtime_defined_discrete_targets, + ); + if !alias_refs_are_continuous && !single_live_has_known_boundary { return true; } - live.len() > 1 - && !can_eliminate_continuous_connection_alias( - dae, - eq_rhs, - live, - runtime_defined_discrete_targets, - ) + live.len() > 1 && !can_eliminate_multi_live } fn connection_alias_refs_are_continuous_scalar_algebraics( @@ -52,14 +69,42 @@ fn connection_alias_refs_are_continuous_scalar_algebraics( else { return false; }; - let Some(lhs_ref) = full_var_ref(lhs) else { + let Some(lhs_name) = exact_reference_expr_name_in_dae(dae, lhs) else { return false; }; - let Some(rhs_ref) = full_var_ref(rhs) else { + let Some(rhs_name) = exact_reference_expr_name_in_dae(dae, rhs) else { return false; }; - [lhs_ref.var_name(), rhs_ref.var_name()].iter().all(|name| { + [&lhs_name, &rhs_name].iter().all(|name| { can_eliminate_scalar_connection_alias_var(dae, name, runtime_defined_discrete_targets) + || connection_alias_ref_is_structurally_known(dae, name) + }) +} + +fn single_live_connection_alias_has_known_boundary( + dae: &Dae, + eq_rhs: &Expression, + live: &[VarName], + runtime_defined_discrete_targets: &HashSet, +) -> bool { + let [live_name] = live else { + return false; + }; + can_eliminate_scalar_connection_alias_var(dae, live_name, runtime_defined_discrete_targets) + && connection_expr_refs_are_live_or_structurally_known(dae, eq_rhs, live_name) +} + +fn connection_expr_refs_are_live_or_structurally_known( + dae: &Dae, + eq_rhs: &Expression, + live_name: &VarName, +) -> bool { + let mut refs = Vec::new(); + collect_var_ref_nodes(eq_rhs, &mut refs); + refs.iter().all(|(name, _)| { + name.var_name() == live_name + || name.is_generated() + || connection_alias_ref_is_structurally_known(dae, name.var_name()) }) } @@ -78,20 +123,26 @@ fn can_eliminate_continuous_connection_alias( else { return false; }; - let Some(lhs_ref) = full_var_ref(lhs) else { + let Some(lhs_name) = exact_reference_expr_name_in_dae(dae, lhs) else { return false; }; - let Some(rhs_ref) = full_var_ref(rhs) else { + let Some(rhs_name) = exact_reference_expr_name_in_dae(dae, rhs) else { return false; }; - [lhs_ref.var_name(), rhs_ref.var_name()].iter().all(|name| { + let refs = [&lhs_name, &rhs_name]; + let has_eliminable_live = refs.iter().any(|name| { live.iter().any(|live_name| live_name == *name) && can_eliminate_scalar_connection_alias_var( dae, name, runtime_defined_discrete_targets, ) - }) + }); + has_eliminable_live + && refs.iter().all(|name| { + can_eliminate_scalar_connection_alias_var(dae, name, runtime_defined_discrete_targets) + || connection_alias_ref_is_structurally_known(dae, name) + }) } fn can_eliminate_scalar_connection_alias_var( @@ -99,9 +150,135 @@ fn can_eliminate_scalar_connection_alias_var( var_name: &VarName, runtime_defined_discrete_targets: &HashSet, ) -> bool { - !var_name.as_str().contains('[') + scalar_connection_alias_var_size(dae, var_name).is_some_and(|size| size == 1) && !unknown_is_fixed(dae, var_name) - && dae.variables.algebraics.contains_key(var_name) && !runtime_defined_discrete_targets.contains(var_name.as_str()) && !runtime_partition_or_event_refs_var(dae, var_name) } + +fn scalar_connection_alias_var_size(dae: &Dae, var_name: &VarName) -> Option { + if is_scalarized_aggregate_element(dae, var_name) + && !scalarized_element_has_non_connection_use(dae, var_name) + { + return None; + } + if is_scalarized_aggregate_element(dae, var_name) { + return Some(1); + } + dae.variables + .algebraics + .get(var_name) + .or_else(|| dae.variables.outputs.get(var_name)) + .map(|var| var.size()) +} + +fn is_scalarized_aggregate_element(dae: &Dae, var_name: &VarName) -> bool { + scalar_element_base_name(var_name) + .and_then(|base| { + dae.variables + .algebraics + .get(&base) + .or_else(|| dae.variables.outputs.get(&base)) + }) + .is_some_and(|var| var.size() > 1) +} + +fn scalar_element_base_name(var_name: &VarName) -> Option { + let text = var_name.as_str(); + if !text.ends_with(']') { + return None; + } + let open = text.rfind('[')?; + if text[open + 1..text.len() - 1] + .chars() + .all(|ch| ch.is_ascii_digit() || ch == ',') + { + Some(VarName::new(&text[..open])) + } else { + None + } +} + +fn scalarized_element_has_non_connection_use(dae: &Dae, var_name: &VarName) -> bool { + dae.continuous.equations.iter().any(|eq| { + !eq.origin.starts_with("connection equation:") + && expr_references_canonical_scalar(&eq.rhs, var_name) + }) +} + +fn expr_references_canonical_scalar(expr: &Expression, var_name: &VarName) -> bool { + let mut refs = Vec::new(); + super::collect_var_ref_nodes(expr, &mut refs); + refs.iter().any(|(name, subscripts)| { + assignment_var_ref_name(name.var_name(), subscripts) + .as_ref() + .is_some_and(|referenced| referenced == var_name) + }) +} + +fn connection_alias_ref_is_structurally_known(dae: &Dae, var_name: &VarName) -> bool { + let mut visiting = HashSet::new(); + connection_alias_ref_is_structurally_known_inner(dae, var_name, &mut visiting) +} + +fn connection_alias_ref_is_structurally_known_inner( + dae: &Dae, + var_name: &VarName, + visiting: &mut HashSet, +) -> bool { + if dae.variables.parameters.contains_key(var_name) + || dae.variables.constants.contains_key(var_name) + || var_name.as_str() == "time" + { + return true; + } + if !dae.variables.algebraics.contains_key(var_name) { + return false; + } + if !visiting.insert(var_name.as_str().to_string()) { + return false; + } + dae.continuous + .equations + .iter() + .filter_map(|eq| direct_definition_expr_for_var(dae, &eq.rhs, var_name)) + .any(|expr| connection_alias_expr_is_structurally_known(dae, expr, visiting)) +} + +fn direct_definition_expr_for_var<'a>( + dae: &Dae, + expr: &'a Expression, + var_name: &VarName, +) -> Option<&'a Expression> { + let Expression::Binary { + op: OpBinary::Sub, + lhs, + rhs, + .. + } = expr + else { + return None; + }; + let lhs_name = exact_reference_expr_name_in_dae(dae, lhs); + let rhs_name = exact_reference_expr_name_in_dae(dae, rhs); + if lhs_name.as_ref() == Some(var_name) { + Some(rhs) + } else if rhs_name.as_ref() == Some(var_name) { + Some(lhs) + } else { + None + } +} + +fn connection_alias_expr_is_structurally_known( + dae: &Dae, + expr: &Expression, + visiting: &mut HashSet, +) -> bool { + let mut refs = Vec::new(); + collect_var_ref_nodes(expr, &mut refs); + refs.iter().all(|(name, _)| { + name.is_generated() + || connection_alias_ref_is_structurally_known_inner(dae, name.var_name(), visiting) + }) +} diff --git a/crates/rumoca-phase-structural/src/eliminate/direct_definition_index.rs b/crates/rumoca-phase-structural/src/eliminate/direct_definition_index.rs index a8bfb16be..eb77740fb 100644 --- a/crates/rumoca-phase-structural/src/eliminate/direct_definition_index.rs +++ b/crates/rumoca-phase-structural/src/eliminate/direct_definition_index.rs @@ -46,6 +46,26 @@ impl DirectDefinitionIndex { .is_some_and(|targets| targets.iter().any(|target| target == candidate)); count > usize::from(current_has_definition) } + + pub(super) fn has_other_non_connection_direct_definition( + &self, + dae: &dae::Dae, + current_eq_idx: usize, + candidate: &VarName, + ) -> bool { + self.per_equation + .iter() + .enumerate() + .any(|(eq_idx, targets)| { + eq_idx != current_eq_idx + && targets.iter().any(|target| target == candidate) + && !dae + .continuous + .equations + .get(eq_idx) + .is_some_and(|eq| eq.origin.contains("connection")) + }) + } } fn direct_assignment_targets(expr: &Expression) -> IndexSet { diff --git a/crates/rumoca-phase-structural/src/eliminate/flow_policy.rs b/crates/rumoca-phase-structural/src/eliminate/flow_policy.rs index fb578a230..897d2b64e 100644 --- a/crates/rumoca-phase-structural/src/eliminate/flow_policy.rs +++ b/crates/rumoca-phase-structural/src/eliminate/flow_policy.rs @@ -1,4 +1,8 @@ -use super::{Dae, Expression, StructuralError, VarName, collect_var_ref_nodes}; +use super::scalar_shape::var_ref_is_scalar_after_subscripts; +use super::{ + Dae, Expression, StructuralError, VarName, collect_var_ref_nodes, + exact_reference_expr_name_in_dae, +}; use crate::variable_scope::{DaeVariableScope, DaeVariableShape, scalar_count_from_dims}; pub(super) fn is_flow_equation_origin(origin: &str) -> bool { @@ -9,31 +13,112 @@ pub(super) fn expr_contains_indexed_multiscalar_ref( expr: &Expression, dae: &Dae, ) -> Result { + let scope = DaeVariableScope::new(dae); + if let Some(exact_name) = exact_reference_expr_name_in_dae(dae, expr) + && let Some(var) = scope.exact(&exact_name) + && scalar_count_from_dims(&exact_name, &var.dims)? == 1 + { + return Ok(false); + } let mut refs = Vec::new(); collect_var_ref_nodes(expr, &mut refs); - let scope = DaeVariableScope::new(dae); for (name, subscripts) in refs { if name.as_str() == "time" { continue; } - if reference_has_embedded_multiscalar_index(dae, name.var_name())? { - return Ok(true); + if !reference_touches_continuous_unknown(dae, name.var_name()) { + continue; + } + match reference_has_embedded_multiscalar_index(dae, name.var_name()) { + Ok(true) => return Ok(true), + Ok(false) => {} + Err(StructuralError::ContractViolation { reason, .. }) + | Err(StructuralError::UnspannedContractViolation { reason }) + if reason.contains("missing DAE variable metadata") => + { + return Ok(true); + } + Err(err) => return Err(err), } if subscripts.is_empty() { continue; } - match scope.shape_for_reference(&name)? { - DaeVariableShape::Dimensions(dims) => { + match scope.shape_for_reference(&name) { + Err(StructuralError::ContractViolation { reason, .. }) + | Err(StructuralError::UnspannedContractViolation { reason }) + if reason.contains("missing DAE variable metadata") => + { + return Ok(true); + } + Err(err) => return Err(err), + Ok(DaeVariableShape::Dimensions(dims)) => { if scalar_count_from_dims(name.var_name(), &dims)? > 1 { return Ok(true); } } - DaeVariableShape::StructuredAggregate => return Ok(true), + Ok(DaeVariableShape::StructuredAggregate) => return Ok(true), } } Ok(false) } +pub(super) fn expr_contains_indexed_multiscalar_slice_ref( + expr: &Expression, + dae: &Dae, +) -> Result { + let mut refs = Vec::new(); + collect_var_ref_nodes(expr, &mut refs); + for (name, subscripts) in refs { + if subscripts.is_empty() || name.as_str() == "time" { + continue; + } + if !reference_touches_continuous_unknown(dae, name.var_name()) { + continue; + } + let reference_span = name + .span() + .or_else(|| subscripts.first().map(rumoca_core::Subscript::span)) + .ok_or_else(|| StructuralError::UnspannedContractViolation { + reason: format!( + "cannot classify indexed multiscalar reference `{}` without source provenance", + name.as_str() + ), + })?; + match var_ref_is_scalar_after_subscripts(&name, &subscripts, reference_span, dae) { + Ok(false) => return Ok(true), + Ok(true) => {} + Err(StructuralError::ContractViolation { reason, .. }) + | Err(StructuralError::UnspannedContractViolation { reason }) + if is_overspecified_indexed_reference_contract(&reason) => + { + return Ok(true); + } + Err(err) => return Err(err), + } + } + Ok(false) +} + +fn is_overspecified_indexed_reference_contract(reason: &str) -> bool { + reason.starts_with("indexed DAE reference ") + && reason.contains(" has ") + && reason.contains(" subscripts for dimensions ") +} + +fn reference_touches_continuous_unknown(dae: &Dae, name: &VarName) -> bool { + continuous_unknown_exists(dae, name) + || strip_embedded_subscripts(name.as_str()) + .is_some_and(|base| continuous_unknown_exists(dae, &VarName::new(base))) + || rumoca_ir_dae::split_complex_field_suffix(name.as_str()) + .is_some_and(|(base, _)| continuous_unknown_exists(dae, &VarName::new(base))) +} + +fn continuous_unknown_exists(dae: &Dae, name: &VarName) -> bool { + dae.variables.states.contains_key(name) + || dae.variables.algebraics.contains_key(name) + || dae.variables.outputs.contains_key(name) +} + fn reference_has_embedded_multiscalar_index( dae: &Dae, name: &VarName, @@ -70,3 +155,36 @@ fn strip_embedded_subscripts(name: &str) -> Option { } (depth == 0 && stripped != name).then_some(stripped) } + +#[cfg(test)] +mod tests { + use super::*; + use rumoca_core::{Reference, Span, Subscript}; + use rumoca_ir_dae as dae; + + #[test] + fn overspecified_vector_projection_is_classified_as_indexed_slice() { + let mut dae_model = dae::Dae::default(); + dae_model.variables.algebraics.insert( + VarName::new("x"), + dae::Variable { + name: VarName::new("x"), + dims: vec![2], + ..dae::Variable::empty_with_span(Span::DUMMY) + }, + ); + let expr = Expression::VarRef { + name: Reference::from_var_name(VarName::new("x")), + subscripts: vec![ + Subscript::index(1, Span::DUMMY), + Subscript::index(1, Span::DUMMY), + ], + span: Span::DUMMY, + }; + + assert!( + expr_contains_indexed_multiscalar_slice_ref(&expr, &dae_model) + .expect("slice classifier should not fatal on overspecified projection") + ); + } +} diff --git a/crates/rumoca-phase-structural/src/eliminate/mod.rs b/crates/rumoca-phase-structural/src/eliminate/mod.rs index d77bdefa6..da7a02332 100644 --- a/crates/rumoca-phase-structural/src/eliminate/mod.rs +++ b/crates/rumoca-phase-structural/src/eliminate/mod.rs @@ -1,5 +1,9 @@ //! Symbolic elimination of trivially solvable equations. //! +//! SPEC_0021 file-size exception: elimination still hosts boundary resolution, +//! substitution grouping, and provenance-preserving replacement helpers. split plan: +//! move substitution groups and scalar alias projection into submodules. +//! //! Two-phase pipeline: //! 1. **Boundary resolution** — removes redundant equations (0 unknowns) and //! resolves trivial single-unknown equations, making structurally singular @@ -45,24 +49,31 @@ use boundary_scan::{BoundaryScanCtx, BoundaryScanState, scan_boundary_equations} use connection_policy::should_skip_connection_equation; use diagnostics::trace_singular_reduced_rows; use direct_definition_index::DirectDefinitionIndex; -use flow_policy::{expr_contains_indexed_multiscalar_ref, is_flow_equation_origin}; +use flow_policy::{ + expr_contains_indexed_multiscalar_ref, expr_contains_indexed_multiscalar_slice_ref, + is_flow_equation_origin, +}; use orphan_unknowns::{drop_unreferenced_continuous_unknowns, output_partition_contains_unknown}; use profiling::{eliminate_profile_enabled, log_eliminate_profile}; use runtime_known::singular_rows_are_runtime_known_assignments; use runtime_protection::{ - assignment_target_name, expr_references_any_discrete_name, + assignment_target_name_in_dae, expr_references_any_discrete_name, expr_references_any_runtime_discrete_target, is_runtime_protected_unknown, runtime_defined_discrete_target_names, runtime_partition_or_event_refs_var, runtime_protected_unknown_names, should_preserve_runtime_known_assignment, }; use scalar_shape::expression_is_scalar_after_subscripts; pub use solve_for_unknown::try_solve_for_unknown; +use solve_for_unknown::{expr_contains_unknown_in_dae, try_solve_for_unknown_in_dae}; use substitution_application::{ - apply_substitutions_in_order, apply_substitutions_to_dae_partitions, - apply_substitutions_to_remaining_once, equation_analysis_expr, + apply_substitutions_to_dae_partitions, apply_substitutions_to_remaining_once, + canonicalize_exact_indexing_in_continuous_equations, equation_analysis_expr, }; use substitution_target::{ - expr_contains_derivative_substitution_target, expr_contains_substitution_target, + expr_contains_derivative_substitution_target, expr_contains_substitution_target_in_scope, + expression_is_exact_structured_substitution_target, + generated_scalar_reference_matches_exact_substitution_name, + substitution_requires_structured_identity, }; use tearing_elimination::tear_and_eliminate_loop_block; use unknown_index::{ @@ -71,14 +82,17 @@ use unknown_index::{ }; use crate::variable_scope::{DaeVariableScope, DaeVariableShape, scalar_count_from_dims}; -use crate::{BltBlock, EquationRef, StructuralError, UnknownId, sort_dae}; +use crate::{ + BltBlock, EquationRef, StructuralError, UnknownId, build_blt_from_incidence, + maximum_regular_subsystem, sort_dae, +}; use rumoca_core::ExpressionVisitor; #[cfg(test)] use rumoca_ir_dae::expr_contains_der_of; use rumoca_ir_dae::{ - DerivativeNameMatcher, expr_contains_der_of_any, expr_contains_var, split_complex_field_suffix, - subscripts_all_one, var_ref_matches_unknown, + DaeVisitor, DerivativeNameMatcher, expr_contains_der_of_any, expr_contains_var, + split_complex_field_suffix, subscripts_all_one, var_ref_matches_unknown, }; type Dae = dae::Dae; @@ -120,7 +134,6 @@ pub struct EliminationResult { struct ZeroUnknownEliminationCtx<'a> { dae: &'a Dae, - state_names: &'a [VarName], unknown_index: &'a BoundaryUnknownIndex<'a>, resolved: &'a HashSet, runtime_protected_unknowns: &'a IndexSet, @@ -191,11 +204,14 @@ pub fn eliminate_trivial(dae: &mut Dae) -> Result Some(sorted.blocks.clone()), Err(StructuralError::EmptySystem) => None, Err(err) if singular_rows_are_runtime_known_assignments(&sort_input, &err) => None, - Err(err) => { - trace_singular_reduced_rows(trace, &sort_input, &err); - blt_error = Some(err); - None - } + Err(err) => match regular_blt_blocks_for_fully_matched_rows(&sort_input, &err)? { + Some(blocks) => Some(blocks), + None => { + trace_singular_reduced_rows(trace, &sort_input, &err); + blt_error = Some(err); + None + } + }, }; log_eliminate_profile( profile, @@ -240,13 +256,129 @@ pub fn eliminate_trivial(dae: &mut Dae) -> Result Result>, StructuralError> { + let StructuralError::Singular { + n_equations, + n_unknowns, + n_matched, + unmatched_equations, + unmatched_unknowns, + .. + } = error + else { + return Ok(None); + }; + if unmatched_unknowns.is_empty() && n_matched == n_unknowns && n_unknowns < n_equations { + let incidence = crate::incidence::build_incidence(dae); + let regular = maximum_regular_subsystem(&incidence)?; + if regular.dropped_unknowns.is_empty() + && regular.dropped_equations.iter().all(|equation| { + dropped_equation_is_evaluation_assignment_row(dae, &incidence, equation) + }) + { + return build_blt_from_incidence(®ular.incidence).map(Some); + } + return Ok(None); + } + if n_matched != n_equations || !unmatched_equations.is_empty() || n_unknowns <= n_equations { + return Ok(None); + } + for unknown in unmatched_unknowns { + if !regular_subsystem_extra_unknown_is_direct_component_helper(dae, unknown)? { + return Ok(None); + } + } + let incidence = crate::incidence::build_incidence(dae); + let regular = maximum_regular_subsystem(&incidence)?; + if !regular.dropped_equations.is_empty() || regular.dropped_unknowns.is_empty() { + return Ok(None); + } + build_blt_from_incidence(®ular.incidence).map(Some) +} + +fn dropped_equation_is_evaluation_assignment_row( + dae: &Dae, + incidence: &crate::incidence::Incidence, + equation: &EquationRef, +) -> bool { + let Some(eq) = dae.continuous.equations.get(equation.0) else { + return false; + }; + let analysis_expr = equation_analysis_expr(eq); + let Some(target) = assignment_target_name_in_dae(dae, &analysis_expr) else { + return false; + }; + let Some(row_unknowns) = incidence + .equation_refs + .iter() + .position(|candidate| candidate == equation) + .and_then(|idx| incidence.eq_unknowns.get(idx)) + else { + return false; + }; + !row_unknowns + .iter() + .filter_map(|idx| incidence.unknown_names.get(*idx)) + .any(|unknown| unknown_matches_var_name(unknown, &target)) +} + +fn unknown_matches_var_name(unknown: &UnknownId, target: &VarName) -> bool { + matches!(unknown, UnknownId::Variable(name) if name == target) +} + +fn regular_subsystem_extra_unknown_is_direct_component_helper( + dae: &Dae, + unknown: &str, +) -> Result { + let var_name = VarName::new(unknown); + let is_scalarized_element = is_scalarized_element_of_aggregate(dae, &var_name)?; + let component_path_len = rumoca_core::ComponentPath::from_flat_path(unknown).len(); + let is_nested_component = component_path_len > 1; + if output_partition_contains_unknown(dae, &var_name) { + return Ok(true); + } + if !is_scalarized_element && component_path_len > 2 { + return Ok(true); + } + if !is_scalarized_element && !is_nested_component { + return Ok(false); + } + if is_scalarized_element { + return Ok(true); + } + direct_non_connection_definition_for_unknown(dae, &var_name) +} + +fn direct_non_connection_definition_for_unknown( + dae: &Dae, + var_name: &VarName, +) -> Result { + for (eq_idx, equation) in dae.continuous.equations.iter().enumerate() { + if equation.origin.starts_with("connection equation:") { + continue; + } + if !has_direct_assignment_form(dae, &equation.rhs, var_name) { + continue; + } + if !can_use_equation_for_elimination(dae, eq_idx) { + continue; + } + return Ok(true); + } + Ok(false) +} + +pub fn resolve_boundary_equations_to_fixpoint( dae: &mut Dae, ) -> Result { let mut result = EliminationResult::default(); loop { let pass = resolve_boundary_equations(dae)?; if pass.n_eliminated == 0 { + finalize_boundary_substitution_targets(dae, &result.substitutions)?; return Ok(result); } result.n_eliminated += pass.n_eliminated; @@ -254,6 +386,45 @@ fn resolve_boundary_equations_to_fixpoint( } } +/// Retire only variables reconstructed by this elimination result after their +/// final structural reference has disappeared. This deliberately does not scan +/// for arbitrary orphan variables: pre-existing unconstrained unknowns must +/// remain visible to singularity diagnostics. +fn finalize_boundary_substitution_targets( + dae: &mut Dae, + substitutions: &[Substitution], +) -> Result<(), StructuralError> { + apply_substitutions_to_dae_partitions(dae, substitutions)?; + for substitution in substitutions { + let name = &substitution.var_name; + if dae_expression_surfaces_ref_var(dae, name) { + continue; + } + dae.variables.algebraics.shift_remove(name); + dae.variables.outputs.shift_remove(name); + } + Ok(()) +} + +fn dae_expression_surfaces_ref_var(dae: &Dae, name: &VarName) -> bool { + let mut visitor = DaeReferenceVisitor { name, found: false }; + visitor.visit_dae(dae); + visitor.found +} + +struct DaeReferenceVisitor<'a> { + name: &'a VarName, + found: bool, +} + +impl DaeVisitor for DaeReferenceVisitor<'_> { + fn visit_expression(&mut self, expression: &Expression) { + if !self.found && expr_contains_var(expression, self.name) { + self.found = true; + } + } +} + fn resolve_boundary_and_direct_demotions_to_fixpoint( dae: &mut Dae, ) -> Result<(EliminationResult, usize), StructuralError> { @@ -275,7 +446,11 @@ fn resolve_boundary_and_direct_demotions_to_fixpoint( result.substitutions.extend(pass.substitutions); let p_demote = maybe_start_timer_if(profile); - let demoted = crate::dae_prepare::demote_direct_assigned_states(dae)?; + let demoted = + crate::dae_prepare::demote_direct_assigned_states_with_boundary_substitutions( + dae, + &result.substitutions, + )?; log_eliminate_profile(profile, "boundary_direct_demotion", p_demote, demoted); total_demoted += demoted; if eliminated == 0 && demoted == 0 { @@ -307,6 +482,7 @@ fn eliminate_trace_enabled() -> bool { /// /// ODE equations (containing `der(state)`) are always skipped. fn resolve_boundary_equations(dae: &mut Dae) -> Result { + canonicalize_exact_indexing_in_continuous_equations(dae)?; let profile = eliminate_profile_enabled(); let p_unknowns = maybe_start_timer_if(profile); let all_unknowns = collect_boundary_unknowns(dae)?; @@ -345,7 +521,6 @@ fn resolve_boundary_equations(dae: &mut Dae) -> Result Result Result= family.first_equation_index && idx < block_end); - if removed_inside_block { + let Some(block_end) = structured_family_block_end(family) else { + return false; + }; + if sorted_rows_touch_range(removed_sorted, family.first_equation_index, block_end) { return false; } - let shift = removed_sorted - .iter() - .filter(|&&idx| idx < family.first_equation_index) - .count(); - family.first_equation_index -= shift; + let shift = removed_sorted.partition_point(|idx| *idx < family.first_equation_index); + let Some(shifted_start) = family.first_equation_index.checked_sub(shift) else { + return false; + }; + family.first_equation_index = shifted_start; true }); } -/// Drop structured families whose row bodies were symbolically rewritten. +/// Drop structured families whose proof-carrying row bodies were symbolically rewritten. /// /// Substitutions can change a family row's LHS/RHS without changing row count. /// The old compact family metadata was proven for the pre-substitution body, so /// keeping it would let downstream Solve-IR rebuild a tensor node from stale row -/// provenance. Regenerating a family proof belongs in a dedicated pass; the safe -/// structural-elimination behavior is to scalarize rewritten families. +/// provenance. An unmaterialized regular family's interior rows are different: +/// they are explicitly non-semantic placeholders, and only its base/neighbor +/// corner rows carry the reconstruction proof. Rewriting an interior placeholder +/// therefore must not discard the authoritative family metadata. fn drop_structured_families_touching_equations(dae: &mut Dae, touched_sorted: &[usize]) { if touched_sorted.is_empty() { return; } dae.continuous.structured_equations.retain(|family| { - let total: usize = family.equation_counts.iter().sum(); - let block_end = family.first_equation_index + total; - !touched_sorted - .iter() - .any(|&idx| idx >= family.first_equation_index && idx < block_end) + let Some(block_end) = structured_family_block_end(family) else { + return false; + }; + let touches_family = + sorted_rows_touch_range(touched_sorted, family.first_equation_index, block_end); + if !touches_family { + return true; + } + if family.regular.is_some() && !family.interiors_materialized { + return unmaterialized_family_corner_is_touched(family, touched_sorted) == Some(false); + } + false }); } +fn structured_family_block_end(family: &dae::StructuredEquationFamily) -> Option { + family + .equation_counts + .iter() + .try_fold(family.first_equation_index, |end, count| { + end.checked_add(*count) + }) +} + +fn unmaterialized_family_corner_is_touched( + family: &dae::StructuredEquationFamily, + touched_sorted: &[usize], +) -> Option { + if family.domain.scalar_count().ok()? != family.equation_counts.len() { + return None; + } + let mut corner_iterations = vec![0usize]; + let mut iteration_stride = 1usize; + for binder in family.domain.binders.iter().rev() { + let extent = binder.value_count().ok()?; + if extent > 1 { + corner_iterations.push(iteration_stride); + } + iteration_stride = iteration_stride.checked_mul(extent)?; + } + corner_iterations.sort_unstable(); + corner_iterations.dedup(); + + let mut next_corner = corner_iterations.into_iter().peekable(); + let mut row_start = family.first_equation_index; + for (iteration, &row_count) in family.equation_counts.iter().enumerate() { + let row_end = row_start.checked_add(row_count)?; + if next_corner.peek() == Some(&iteration) { + if sorted_rows_touch_range(touched_sorted, row_start, row_end) { + return Some(true); + } + next_corner.next(); + } + row_start = row_end; + } + Some(false) +} + +fn sorted_rows_touch_range(rows: &[usize], start: usize, end: usize) -> bool { + let first_candidate = rows.partition_point(|row| *row < start); + rows.get(first_candidate).is_some_and(|row| *row < end) +} + fn finish_boundary_elimination( dae: &mut Dae, - substitutions: Vec, - eliminated_eq_flags: Vec, + mut substitutions: Vec, + mut eliminated_eq_flags: Vec, mut eliminated_eq_indices: Vec, resolved: &HashSet, + runtime_protected_unknowns: &IndexSet, + runtime_defined_discrete_targets: &HashSet, ) -> Result { + let mut resolved = resolved.clone(); + eliminated_eq_indices.retain(|&idx| { + let preserve = dae.continuous.equations.get(idx).is_some_and(|eq| { + should_preserve_runtime_sensitive_continuous_assignment( + dae, + &equation_analysis_expr(eq), + ) + }); + if preserve && let Some(flag) = eliminated_eq_flags.get_mut(idx) { + *flag = false; + } + if preserve + && let Some(target) = dae + .continuous + .equations + .get(idx) + .and_then(|eq| assignment_target_name_in_dae(dae, &equation_analysis_expr(eq))) + { + substitutions.retain(|sub| sub.var_name != target); + resolved.remove(&target); + } + !preserve + }); apply_substitutions_to_remaining_once(dae, &eliminated_eq_flags, &substitutions)?; let n_eliminated = eliminated_eq_indices.len(); eliminated_eq_indices.sort_unstable(); @@ -489,7 +745,13 @@ fn finish_boundary_elimination( dae.continuous.equations.remove(idx); } shift_structured_families_after_equation_removal(dae, &eliminated_eq_indices); - for name in fully_resolved_continuous_unknowns(dae, resolved)? { + apply_substitutions_to_dae_partitions(dae, &substitutions)?; + for name in fully_resolved_continuous_unknowns( + dae, + &resolved, + runtime_protected_unknowns, + runtime_defined_discrete_targets, + )? { dae.variables.algebraics.shift_remove(&name); dae.variables.outputs.shift_remove(&name); } @@ -503,6 +765,8 @@ fn finish_boundary_elimination( fn fully_resolved_continuous_unknowns( dae: &Dae, resolved: &HashSet, + runtime_protected_unknowns: &IndexSet, + runtime_defined_discrete_targets: &HashSet, ) -> Result, StructuralError> { let mut removable = IndexSet::new(); for (name, var) in dae @@ -511,6 +775,12 @@ fn fully_resolved_continuous_unknowns( .iter() .chain(dae.variables.outputs.iter()) { + if is_runtime_protected_unknown(name, runtime_protected_unknowns) + || runtime_defined_discrete_targets.contains(name.as_str()) + || dae_expression_surfaces_ref_var(dae, name) + { + continue; + } if resolved.contains(name) || aggregate_variable_fully_resolved(name, var, resolved)? { removable.insert(name.clone()); } @@ -590,8 +860,7 @@ fn can_eliminate_aggregate_alias_var( && !is_runtime_protected_unknown(var_name, runtime_protected_unknowns) && !runtime_defined_discrete_targets.contains(var_name.as_str()) && !runtime_partition_or_event_refs_var(dae, var_name) - && (dae.variables.algebraics.contains_key(var_name) - || dae.variables.outputs.contains_key(var_name)) + && continuous_algebraic_or_output_contains_unknown(dae, var_name) } fn preferred_aggregate_alias_candidate( @@ -608,21 +877,110 @@ fn aggregate_alias_rank(name: &VarName) -> (usize, usize) { (path.len(), name.as_str().len()) } +pub(super) fn scalar_connection_alias_for_elimination( + dae: &Dae, + rhs: &Expression, + runtime_protected_unknowns: &IndexSet, + runtime_defined_discrete_targets: &HashSet, +) -> Result, StructuralError> { + let Expression::Binary { + op: OpBinary::Sub, + lhs, + rhs: rhs_expr, + .. + } = rhs + else { + return Ok(None); + }; + let Some(lhs_name) = + exact_reference_expr_name_in_dae(dae, lhs).or_else(|| exact_reference_expr_name(lhs)) + else { + return Ok(None); + }; + let Some(rhs_name) = exact_reference_expr_name_in_dae(dae, rhs_expr) + .or_else(|| exact_reference_expr_name(rhs_expr)) + else { + return Ok(None); + }; + let lhs_rank = scalar_connection_alias_candidate_rank( + dae, + &lhs_name, + runtime_protected_unknowns, + runtime_defined_discrete_targets, + )?; + let rhs_rank = scalar_connection_alias_candidate_rank( + dae, + &rhs_name, + runtime_protected_unknowns, + runtime_defined_discrete_targets, + )?; + match (lhs_rank, rhs_rank) { + (Some(lhs_rank), Some(rhs_rank)) if lhs_rank <= rhs_rank => { + Ok(Some((lhs_name, rhs_expr.as_ref().clone()))) + } + (Some(_), Some(_)) => Ok(Some((rhs_name, lhs.as_ref().clone()))), + (Some(_), None) => Ok(Some((lhs_name, rhs_expr.as_ref().clone()))), + (None, Some(_)) => Ok(Some((rhs_name, lhs.as_ref().clone()))), + (None, None) => Ok(None), + } +} + +fn can_ignore_internal_output_definition_for_choice( + ctx: &EliminationChoiceContext<'_>, + candidate: &VarName, +) -> bool { + is_internal_component_output(ctx.dae, candidate) + && (is_local_component_output(ctx.dae, candidate) + || ctx + .direct_definitions + .has_other_non_connection_direct_definition(ctx.dae, ctx.eq_idx, candidate)) +} + +fn scalar_connection_alias_candidate_rank( + dae: &Dae, + var_name: &VarName, + runtime_protected_unknowns: &IndexSet, + runtime_defined_discrete_targets: &HashSet, +) -> Result, StructuralError> { + if !can_eliminate_aggregate_alias_var( + dae, + var_name, + runtime_protected_unknowns, + runtime_defined_discrete_targets, + ) || DaeVariableScope::new(dae).size(var_name)? != 1 + { + return Ok(None); + } + if is_scalarized_element_of_aggregate(dae, var_name)? + && !scalarized_element_has_non_connection_use(dae, var_name) + { + return Ok(None); + } + if dae.variables.algebraics.contains_key(var_name) { + return Ok(Some(0)); + } + if let Some(var) = dae.variables.outputs.get(var_name) { + return Ok(Some( + if matches!(var.causality, dae::VariableCausality::Input) { + 1 + } else { + 2 + }, + )); + } + Ok(None) +} + fn try_eliminate_zero_unknown_equation( eq_idx: usize, eq_rhs: &Expression, has_state_derivative: bool, ctx: &mut ZeroUnknownEliminationCtx<'_>, ) -> Result<(), StructuralError> { - let references_state_value = ctx - .state_names - .iter() - .any(|sn| expr_contains_var(eq_rhs, sn)); - if has_state_derivative - || references_state_value - || has_any_live_unknown(eq_rhs, ctx.unknown_index, ctx.resolved)? - || expr_contains_indexed_multiscalar_ref(eq_rhs, ctx.dae)? - { + if has_state_derivative || has_any_live_unknown(eq_rhs, ctx.unknown_index, ctx.resolved)? { + return Ok(()); + } + if zero_unknown_equation_preserves_state_value_constraint(ctx.dae, eq_rhs) { return Ok(()); } // MLS Appendix B / §8.3 / §16.5.1: a zero-unknown equation may still @@ -631,6 +989,10 @@ fn try_eliminate_zero_unknown_equation( if should_preserve_runtime_known_assignment(ctx.dae, eq_rhs) { return Ok(()); } + if should_preserve_runtime_sensitive_continuous_assignment(ctx.dae, eq_rhs) { + return Ok(()); + } + let assignment_target = assignment_target_name_in_dae(ctx.dae, eq_rhs); let n_subs_before = ctx.substitutions.len(); maybe_push_non_unknown_alias_substitution( ctx.dae, @@ -639,7 +1001,9 @@ fn try_eliminate_zero_unknown_equation( ctx.runtime_defined_discrete_targets, ctx.substitutions, )?; - if assignment_target_name(eq_rhs).is_some_and(|target| dae_var(ctx.dae, &target).is_some()) + if assignment_target + .as_ref() + .is_some_and(|target| ctx.dae.variables.states.contains_key(target)) && ctx.substitutions.len() == n_subs_before { return Ok(()); @@ -649,23 +1013,73 @@ fn try_eliminate_zero_unknown_equation( Ok(()) } +fn zero_unknown_equation_preserves_state_value_constraint(dae: &Dae, eq_rhs: &Expression) -> bool { + let assignment_target = assignment_target_name_in_dae(dae, eq_rhs); + dae.variables.states.keys().any(|state_name| { + assignment_target + .as_ref() + .is_none_or(|target| target != state_name) + && expr_contains_var(eq_rhs, state_name) + }) +} + +fn should_preserve_runtime_sensitive_continuous_assignment(dae: &Dae, eq_rhs: &Expression) -> bool { + if !expr_contains_runtime_sensitive_operator(eq_rhs) { + return false; + } + let Some(target) = assignment_target_name_in_dae(dae, eq_rhs) else { + return false; + }; + if is_internal_component_continuous_var(dae, &target) { + return false; + } + dae.variables.algebraics.contains_key(&target) + || dae.variables.outputs.contains_key(&target) + || rumoca_ir_dae::component_base_name(target.as_str()).is_some_and(|base| { + let base = VarName::new(base); + dae.variables.algebraics.contains_key(&base) + || dae.variables.outputs.contains_key(&base) + }) +} + +pub(in crate::eliminate) struct EliminationChoiceContext<'a> { + pub(in crate::eliminate) dae: &'a Dae, + pub(in crate::eliminate) eq_idx: usize, + pub(in crate::eliminate) has_state_derivative: bool, + pub(in crate::eliminate) runtime_protected_unknowns: &'a IndexSet, + pub(in crate::eliminate) direct_definitions: &'a DirectDefinitionIndex, + pub(in crate::eliminate) allow_multi_live_trivial_alias: bool, +} + fn choose_solvable_unknown_for_elimination( - dae: &Dae, - eq_idx: usize, + ctx: &EliminationChoiceContext<'_>, rhs: &Expression, live: &[VarName], - has_state_derivative: bool, - runtime_protected_unknowns: &IndexSet, - direct_definitions: &DirectDefinitionIndex, ) -> Result, StructuralError> { let mut candidates: Vec<&VarName> = live.iter().collect(); + let dae = ctx.dae; + let connection_rhs = connection_rhs_assignment_target(dae, ctx.eq_idx, rhs); candidates.sort_by(|a, b| { - let a_has_definition = direct_definitions.has_other_direct_definition(eq_idx, a); - let b_has_definition = direct_definitions.has_other_direct_definition(eq_idx, b); + let a_is_connection_rhs = + connection_rhs.is_some_and(|target| is_assignment_target(dae, target, a)); + let b_is_connection_rhs = + connection_rhs.is_some_and(|target| is_assignment_target(dae, target, b)); + let a_has_definition = ctx + .direct_definitions + .has_other_direct_definition(ctx.eq_idx, a) + && !can_ignore_internal_output_definition_for_choice(ctx, a); + let b_has_definition = ctx + .direct_definitions + .has_other_direct_definition(ctx.eq_idx, b) + && !can_ignore_internal_output_definition_for_choice(ctx, b); + let a_internal_defined_output = can_ignore_internal_output_definition_for_choice(ctx, a); + let b_internal_defined_output = can_ignore_internal_output_definition_for_choice(ctx, b); let a_is_output = output_partition_contains_unknown(dae, a); let b_is_output = output_partition_contains_unknown(dae, b); - a_has_definition - .cmp(&b_has_definition) + b_is_connection_rhs + .cmp(&a_is_connection_rhs) + .then_with(|| a_has_definition.cmp(&b_has_definition)) + .then_with(|| b_internal_defined_output.cmp(&a_internal_defined_output)) .then_with(|| b_is_output.cmp(&a_is_output)) .then_with(|| a.as_str().cmp(b.as_str())) }); @@ -680,50 +1094,23 @@ fn choose_solvable_unknown_for_elimination( if dae.variables.states.contains_key(candidate) { continue; } - if is_scalarized_element_of_aggregate(dae, candidate)? { - continue; - } - if is_runtime_protected_unknown(candidate, runtime_protected_unknowns) { - continue; - } let is_output = output_partition_contains_unknown(dae, candidate); - // Skip equations with state derivatives — unless the candidate is an - // output that forms a direct alias (e.g. `output y = der(x)`), which - // can be safely eliminated. - if has_state_derivative && !is_output { - continue; - } // Try the simple top-level Sub pattern first; fall back to the additive // solver so substitution residues like `x - (y - 0)` (which the simple // pattern can't see through) still resolve. The additive solver is gated // by `live` to avoid accidentally solving a multi-unknown equation. - let Some(solution) = try_solve_for_unknown(rhs, candidate) else { + let Some(solution) = try_solve_for_unknown_in_dae(dae, rhs, candidate) else { continue; }; - if expr_contains_var(&solution, candidate) { - continue; - } - let direct_assignment_solution = has_direct_assignment_form(rhs, candidate); - // Output variables exist for external callers — only eliminate them - // when the solution is a trivial alias (a single variable reference or - // its negation), since keeping non-trivial outputs enlarges the DAE and - // can hurt solver performance. - if is_output && !is_trivial_alias(&solution) { - continue; - } - if !direct_assignment_solution && !is_symbolically_stable_solution(&solution) { - continue; - } - if expr_contains_unsliced_multiscalar_ref(&solution, dae)? { - continue; - } - if expr_contains_indexed_multiscalar_ref(&solution, dae)? - && !(is_trivial_alias(&solution) - && expression_is_scalar_after_subscripts(&solution, dae)?) - { - continue; - } - if live.len() > 1 && !direct_assignment_solution { + let direct_assignment_solution = has_direct_assignment_form(dae, rhs, candidate); + if !elimination_solution_is_valid( + ctx, + candidate, + &solution, + is_output, + direct_assignment_solution, + live.len(), + )? { continue; } return Ok(Some((candidate.clone(), solution))); @@ -731,42 +1118,265 @@ fn choose_solvable_unknown_for_elimination( Ok(None) } -fn choose_solvable_non_unknown_alias_for_elimination( - dae: &Dae, - rhs: &Expression, - runtime_protected_unknowns: &IndexSet, - runtime_defined_discrete_targets: &HashSet, -) -> Result, StructuralError> { - let Expression::Binary { - op, lhs, rhs: r, .. - } = rhs - else { - return Ok(None); - }; - if !matches!(op, OpBinary::Sub) { - return Ok(None); +fn elimination_solution_is_valid( + ctx: &EliminationChoiceContext<'_>, + candidate: &VarName, + solution: &Expression, + is_output: bool, + direct_assignment_solution: bool, + live_len: usize, +) -> Result { + let dae = ctx.dae; + if expr_contains_unknown_in_dae(dae, solution, candidate) { + return Ok(false); } - - let mut candidates: Vec = Vec::with_capacity(2); - if let Expression::VarRef { - name, subscripts, .. - } = lhs.as_ref() - && subscripts.is_empty() + if !solution_matches_candidate_scalar_shape(dae, candidate, solution)? { + return Ok(false); + } + if !is_trivial_alias_in_dae(dae, solution) + && !solution_is_cheap_for_symbolic_substitution(solution) { - candidates.push(name.clone()); + return Ok(false); } - if let Expression::VarRef { - name, subscripts, .. - } = r.as_ref() - && subscripts.is_empty() - && !candidates - .iter() - .any(|existing| existing.var_name() == name.var_name()) + if !runtime_protected_elimination_is_valid(ctx, candidate, solution, direct_assignment_solution) { - candidates.push(name.clone()); + return Ok(false); } - - let scope = DaeVariableScope::new(dae); + if !scalarized_candidate_elimination_is_valid( + ctx, + candidate, + solution, + direct_assignment_solution, + )? { + return Ok(false); + } + if !state_derivative_elimination_is_valid( + ctx, + candidate, + solution, + is_output, + direct_assignment_solution, + ) { + return Ok(false); + } + if is_output + && !is_trivial_alias_in_dae(dae, solution) + && !is_internal_component_output(dae, candidate) + && !is_connection_rhs_boundary_input(ctx, candidate, direct_assignment_solution) + { + return Ok(false); + } + if !direct_assignment_solution && !is_symbolically_stable_solution(solution) { + return Ok(false); + } + if solution_has_blocking_unsliced_multiscalar_ref(solution, dae)? { + return Ok(false); + } + if indexed_multiscalar_slice_solution_is_blocked(dae, solution)? { + return Ok(false); + } + Ok(live_len <= 1 + || direct_assignment_solution + || (ctx.allow_multi_live_trivial_alias && is_trivial_alias_in_dae(dae, solution))) +} + +fn runtime_protected_elimination_is_valid( + ctx: &EliminationChoiceContext<'_>, + candidate: &VarName, + _solution: &Expression, + _direct_assignment_solution: bool, +) -> bool { + !is_runtime_protected_unknown(candidate, ctx.runtime_protected_unknowns) +} + +fn scalarized_candidate_elimination_is_valid( + ctx: &EliminationChoiceContext<'_>, + candidate: &VarName, + solution: &Expression, + direct_assignment_solution: bool, +) -> Result { + let dae = ctx.dae; + let is_scalarized_element = is_scalarized_element_of_aggregate(dae, candidate)?; + if !is_scalarized_element { + return Ok(true); + } + let is_trivial_alias = is_trivial_alias_in_dae(dae, solution); + if direct_assignment_solution + && !is_trivial_alias + && expr_contains_runtime_sensitive_operator(solution) + && scalarized_element_has_non_connection_use(dae, candidate) + { + return Ok(false); + } + if ctx.has_state_derivative { + return Ok(false); + } + if scalarized_element_has_coupled_derivative_use(dae, candidate) && !is_trivial_alias { + return Ok(false); + } + Ok(scalarized_element_has_non_connection_use(dae, candidate)) +} + +fn state_derivative_elimination_is_valid( + ctx: &EliminationChoiceContext<'_>, + candidate: &VarName, + solution: &Expression, + is_output: bool, + direct_assignment_solution: bool, +) -> bool { + !ctx.has_state_derivative + || (is_output && derivative_alias_has_other_equation(ctx, solution)) + || (direct_assignment_solution + && ctx + .direct_definitions + .has_other_non_connection_direct_definition(ctx.dae, ctx.eq_idx, candidate) + && is_derivative_alias_expr(solution)) +} + +fn derivative_alias_has_other_equation( + ctx: &EliminationChoiceContext<'_>, + solution: &Expression, +) -> bool { + let Expression::BuiltinCall { + function: BuiltinFunction::Der, + args, + .. + } = solution + else { + return false; + }; + let Some(state_name) = args + .first() + .and_then(|arg| exact_reference_expr_name_in_dae(ctx.dae, arg)) + else { + return false; + }; + let matcher = DerivativeNameMatcher::from_var_names([&state_name]); + ctx.dae + .continuous + .equations + .iter() + .enumerate() + .any(|(eq_idx, equation)| { + eq_idx != ctx.eq_idx + && expr_contains_der_of_any(&equation_analysis_expr(equation), &matcher) + }) +} + +fn is_connection_rhs_boundary_input( + ctx: &EliminationChoiceContext<'_>, + candidate: &VarName, + direct_assignment_solution: bool, +) -> bool { + direct_assignment_solution + && connection_rhs_assignment_target( + ctx.dae, + ctx.eq_idx, + &ctx.dae.continuous.equations[ctx.eq_idx].rhs, + ) + .is_some_and(|target| is_assignment_target(ctx.dae, target, candidate)) +} + +fn indexed_multiscalar_slice_solution_is_blocked( + dae: &Dae, + solution: &Expression, +) -> Result { + if !expr_contains_indexed_multiscalar_slice_ref(solution, dae)? { + return Ok(false); + } + Ok(!(is_scalar_reduction_solution_tree(solution) + || is_trivial_alias_in_dae(dae, solution) + && expression_is_scalar_after_subscripts(solution, dae)?)) +} + +fn solution_matches_candidate_scalar_shape( + dae: &Dae, + candidate: &VarName, + solution: &Expression, +) -> Result { + if DaeVariableScope::new(dae).size(candidate)? > 1 { + return Ok(true); + } + Ok(!expression_is_multiscalar_literal(solution)) +} + +fn expression_is_multiscalar_literal(expr: &Expression) -> bool { + match expr { + Expression::Array { elements, .. } | Expression::Tuple { elements, .. } => { + literal_scalar_count(elements) > 1 + } + _ => false, + } +} + +fn literal_scalar_count(elements: &[Expression]) -> usize { + elements + .iter() + .map(|element| match element { + Expression::Array { elements, .. } | Expression::Tuple { elements, .. } => { + literal_scalar_count(elements) + } + _ => 1, + }) + .sum() +} + +fn connection_rhs_assignment_target<'a>( + dae: &'a Dae, + eq_idx: usize, + rhs: &'a Expression, +) -> Option<&'a Expression> { + let eq = dae.continuous.equations.get(eq_idx)?; + if !eq.origin.starts_with("connection equation:") { + return None; + } + let Expression::Binary { + op: OpBinary::Sub, + rhs: target, + .. + } = rhs + else { + return None; + }; + Some(target.as_ref()) +} + +fn choose_solvable_non_unknown_alias_for_elimination( + dae: &Dae, + rhs: &Expression, + runtime_protected_unknowns: &IndexSet, + runtime_defined_discrete_targets: &HashSet, +) -> Result, StructuralError> { + let Expression::Binary { + op, lhs, rhs: r, .. + } = rhs + else { + return Ok(None); + }; + if !matches!(op, OpBinary::Sub) { + return Ok(None); + } + + let mut candidates: Vec = Vec::with_capacity(2); + if let Expression::VarRef { + name, subscripts, .. + } = lhs.as_ref() + && subscripts.is_empty() + { + candidates.push(name.clone()); + } + if let Expression::VarRef { + name, subscripts, .. + } = r.as_ref() + && subscripts.is_empty() + && !candidates + .iter() + .any(|existing| existing.var_name() == name.var_name()) + { + candidates.push(name.clone()); + } + + let scope = DaeVariableScope::new(dae); for candidate_ref in candidates { let candidate = candidate_ref.var_name().clone(); if candidate.as_str() == "time" { @@ -792,13 +1402,16 @@ fn choose_solvable_non_unknown_alias_for_elimination( None => continue, } - let Some(solution) = try_solve_for_unknown(rhs, &candidate) else { + let Some(solution) = try_solve_for_unknown_in_dae(dae, rhs, &candidate) else { continue; }; - if expr_contains_var(&solution, &candidate) { + if expr_contains_unknown_in_dae(dae, &solution, &candidate) { + continue; + } + if !solution_matches_candidate_scalar_shape(dae, &candidate, &solution)? { continue; } - if expr_contains_unsliced_multiscalar_ref(&solution, dae)? { + if solution_has_blocking_unsliced_multiscalar_ref(&solution, dae)? { continue; } if !is_symbolically_stable_solution(&solution) { @@ -831,126 +1444,1100 @@ fn maybe_push_non_unknown_alias_substitution( } fn unknown_is_fixed(dae: &Dae, name: &VarName) -> bool { + continuous_state_algebraic_or_output_var(dae, name) + .and_then(|var| var.fixed) + .unwrap_or(false) +} + +fn continuous_state_algebraic_or_output_var<'a>( + dae: &'a Dae, + name: &VarName, +) -> Option<&'a dae::Variable> { dae.variables .states .get(name) .or_else(|| dae.variables.algebraics.get(name)) .or_else(|| dae.variables.outputs.get(name)) - .and_then(|var| var.fixed) - .unwrap_or(false) + .or_else(|| { + rumoca_ir_dae::component_base_name(name.as_str()).and_then(|base| { + let base = VarName::new(base); + dae.variables + .states + .get(&base) + .or_else(|| dae.variables.algebraics.get(&base)) + .or_else(|| dae.variables.outputs.get(&base)) + }) + }) +} + +fn continuous_algebraic_or_output_contains_unknown(dae: &Dae, name: &VarName) -> bool { + dae.variables.algebraics.contains_key(name) + || output_partition_contains_unknown(dae, name) + || rumoca_ir_dae::component_base_name(name.as_str()) + .is_some_and(|base| dae.variables.algebraics.contains_key(&VarName::new(base))) +} + +fn has_direct_assignment_form(dae: &Dae, rhs: &Expression, candidate: &VarName) -> bool { + match rhs { + Expression::Binary { + op: OpBinary::Sub, + lhs, + rhs, + .. + } => is_assignment_target(dae, lhs, candidate) || is_assignment_target(dae, rhs, candidate), + Expression::Unary { + op: OpUnary::Minus, + rhs, + .. + } => has_direct_assignment_form(dae, rhs, candidate), + Expression::If { + branches, + else_branch, + .. + } => { + branches + .iter() + .all(|(_, branch)| has_direct_assignment_form(dae, branch, candidate)) + && has_direct_assignment_form(dae, else_branch, candidate) + } + _ => false, + } +} + +fn is_assignment_target(dae: &Dae, expr: &Expression, candidate: &VarName) -> bool { + if exact_reference_expr_name_in_dae(dae, expr).as_ref() == Some(candidate) { + return true; + } + match expr { + Expression::VarRef { + name, subscripts, .. + } => { + var_ref_matches_unknown(name, subscripts, candidate) + || assignment_slice_target_contains_scalarized_candidate( + dae, name, subscripts, candidate, + ) + || assignment_target_is_singleton_projection(dae, name, subscripts, candidate) + } + Expression::Index { + base, subscripts, .. + } => { + if let Expression::VarRef { + name, + subscripts: base_subscripts, + .. + } = base.as_ref() + { + let mut combined = Vec::with_capacity(base_subscripts.len() + subscripts.len()); + combined.extend_from_slice(base_subscripts); + combined.extend_from_slice(subscripts); + var_ref_matches_unknown(name, &combined, candidate) + || assignment_slice_target_contains_scalarized_candidate( + dae, name, &combined, candidate, + ) + || assignment_target_is_singleton_projection(dae, name, &combined, candidate) + } else { + false + } + } + _ => false, + } +} + +fn assignment_slice_target_contains_scalarized_candidate( + dae: &Dae, + name: &Reference, + subscripts: &[rumoca_core::Subscript], + candidate: &VarName, +) -> bool { + let Some(scalar) = rumoca_core::parse_scalar_name(candidate.as_str()) else { + return false; + }; + if name.as_str() != scalar.base || subscripts.len() != scalar.indices.len() { + return false; + } + subscripts + .iter() + .zip(scalar.indices.iter()) + .all(|(subscript, candidate_index)| match subscript { + rumoca_core::Subscript::Index { value, .. } => value == candidate_index, + rumoca_core::Subscript::Colon { .. } => true, + rumoca_core::Subscript::Expr { .. } => { + exact_subscript_index_in_dae(dae, subscript) == Some(*candidate_index) + } + }) +} + +fn assignment_target_is_singleton_projection( + dae: &Dae, + name: &Reference, + subscripts: &[rumoca_core::Subscript], + candidate: &VarName, +) -> bool { + if let Some((base, indices)) = scalar_var_ref_key_from_reference(name) + && base.as_str() == candidate.as_str() + && indices.iter().all(|index| *index == 1) + && DaeVariableScope::new(dae) + .exact(candidate) + .is_some_and(|var| var.dims.iter().all(|dim| *dim == 1)) + { + return true; + } + if subscripts.is_empty() || name.as_str() != candidate.as_str() { + return false; + } + let Some(var) = DaeVariableScope::new(dae).exact(candidate) else { + return false; + }; + var.dims.iter().all(|dim| *dim == 1) + && subscripts + .iter() + .all(|subscript| matches!(positive_usize_subscript(subscript), Some(1))) +} + +fn scalarized_element_has_non_connection_use(dae: &Dae, candidate: &VarName) -> bool { + dae.continuous.equations.iter().any(|eq| { + !eq.origin.starts_with("connection equation:") + && expr_references_var_for_presence(&eq.rhs, candidate) + }) +} + +pub(super) fn connection_refs_unanchored_scalarized_aggregate( + dae: &Dae, + expr: &Expression, +) -> Result { + let mut refs = Vec::new(); + collect_exact_reference_expr_names_in_dae(dae, expr, &mut refs); + for name in refs { + if is_scalarized_element_of_aggregate(dae, &name)? + && !scalarized_element_has_non_connection_use(dae, &name) + { + return Ok(true); + } + } + Ok(false) +} + +fn scalarized_element_has_coupled_derivative_use(dae: &Dae, candidate: &VarName) -> bool { + let mut exact_derivative_use_count = 0usize; + let has_aggregate_base_use = dae.continuous.equations.iter().any(|eq| { + if !expression_contains_der_call(&eq.rhs) { + return false; + } + let mut refs = Vec::new(); + collect_var_ref_nodes(&eq.rhs, &mut refs); + if refs + .iter() + .any(|(name, subscripts)| var_ref_matches_unknown(name, subscripts, candidate)) + { + exact_derivative_use_count += 1; + } + refs.iter().any(|(name, subscripts)| { + subscripts.is_empty() + && aggregate_ref_matches_scalarized_candidate(name, subscripts, candidate) + }) + }); + has_aggregate_base_use || exact_derivative_use_count > 1 +} + +fn expr_contains_runtime_sensitive_operator(expr: &Expression) -> bool { + struct Checker { + found: bool, + } + + impl ExpressionVisitor for Checker { + fn visit_expression(&mut self, expr: &Expression) { + if self.found { + return; + } + if let Expression::VarRef { name, .. } = expr + && name.as_str() == "time" + { + self.found = true; + return; + } + self.walk_expression(expr); + } + + fn visit_builtin_call(&mut self, function: &BuiltinFunction, args: &[Expression]) { + if matches!( + function, + BuiltinFunction::Pre + | BuiltinFunction::Sample + | BuiltinFunction::Initial + | BuiltinFunction::Terminal + | BuiltinFunction::Edge + | BuiltinFunction::Change + | BuiltinFunction::Reinit + ) { + self.found = true; + return; + } + for arg in args { + self.visit_expression(arg); + } + } + } + + let mut checker = Checker { found: false }; + checker.visit_expression(expr); + checker.found +} + +fn expression_contains_der_call(expr: &Expression) -> bool { + match expr { + Expression::BuiltinCall { + function: BuiltinFunction::Der, + .. + } => true, + Expression::BuiltinCall { args, .. } | Expression::FunctionCall { args, .. } => { + args.iter().any(expression_contains_der_call) + } + Expression::Binary { lhs, rhs, .. } => { + expression_contains_der_call(lhs) || expression_contains_der_call(rhs) + } + Expression::Unary { rhs, .. } => expression_contains_der_call(rhs), + Expression::If { + branches, + else_branch, + .. + } => { + branches.iter().any(|(condition, branch)| { + expression_contains_der_call(condition) || expression_contains_der_call(branch) + }) || expression_contains_der_call(else_branch) + } + Expression::Array { elements, .. } | Expression::Tuple { elements, .. } => { + elements.iter().any(expression_contains_der_call) + } + Expression::Index { + base, subscripts, .. + } => { + expression_contains_der_call(base) + || subscripts.iter().any(|subscript| match subscript { + rumoca_core::Subscript::Expr { expr, .. } => expression_contains_der_call(expr), + _ => false, + }) + } + Expression::FieldAccess { base, .. } => expression_contains_der_call(base), + Expression::ArrayComprehension { expr, filter, .. } => { + expression_contains_der_call(expr) + || filter + .as_ref() + .is_some_and(|filter| expression_contains_der_call(filter)) + } + Expression::Range { + start, step, end, .. + } => { + expression_contains_der_call(start) + || step + .as_ref() + .is_some_and(|step| expression_contains_der_call(step)) + || expression_contains_der_call(end) + } + Expression::VarRef { .. } | Expression::Literal { .. } | Expression::Empty { .. } => false, + } +} + +fn expr_references_var_for_presence(expr: &Expression, candidate: &VarName) -> bool { + let mut refs = Vec::new(); + collect_var_ref_nodes(expr, &mut refs); + refs.iter().any(|(name, subscripts)| { + var_ref_matches_unknown(name, subscripts, candidate) + || aggregate_ref_matches_scalarized_candidate(name, subscripts, candidate) + }) +} + +fn aggregate_ref_matches_scalarized_candidate( + name: &Reference, + subscripts: &[rumoca_core::Subscript], + candidate: &VarName, +) -> bool { + if !subscripts.is_empty() { + return false; + } + rumoca_core::parse_scalar_name(candidate.as_str()) + .is_some_and(|scalar| name.var_name().as_str() == scalar.base) +} + +/// Returns true if the expression is a single variable reference or its +/// negation — i.e., a trivial alias like `x` or `-x`. +fn is_trivial_alias(expr: &Expression) -> bool { + if exact_reference_expr_name(expr).is_some() { + return true; + } + match expr { + Expression::VarRef { .. } => true, + Expression::Unary { + op: OpUnary::Minus, + rhs, + .. + } => is_trivial_alias(rhs), + Expression::BuiltinCall { + function: BuiltinFunction::Der, + args, + .. + } => args.len() == 1 && matches!(&args[0], Expression::VarRef { .. }), + _ => false, + } +} + +pub(super) fn is_trivial_alias_in_dae(dae: &Dae, expr: &Expression) -> bool { + if exact_reference_expr_name_in_dae(dae, expr).is_some() { + return true; + } + match expr { + Expression::Unary { + op: OpUnary::Minus, + rhs, + .. + } => is_trivial_alias_in_dae(dae, rhs), + Expression::BuiltinCall { + function: BuiltinFunction::Der, + args, + .. + } => args.len() == 1 && exact_reference_expr_name_in_dae(dae, &args[0]).is_some(), + _ => is_trivial_alias(expr), + } +} + +fn is_symbolically_stable_solution(expr: &Expression) -> bool { + match expr { + Expression::If { + branches, + else_branch, + .. + } => { + branches.iter().all(|(condition, branch)| { + is_symbolically_stable_solution(condition) + && is_symbolically_stable_solution(branch) + }) && is_symbolically_stable_solution(else_branch) + } + Expression::BuiltinCall { function, args, .. } => { + !matches!( + function, + rumoca_core::BuiltinFunction::Smooth + | rumoca_core::BuiltinFunction::NoEvent + | rumoca_core::BuiltinFunction::Homotopy + ) && args.iter().all(is_symbolically_stable_solution) + } + Expression::Binary { lhs, rhs, .. } => { + is_symbolically_stable_solution(lhs) && is_symbolically_stable_solution(rhs) + } + Expression::Unary { rhs, .. } => is_symbolically_stable_solution(rhs), + Expression::FunctionCall { args, .. } => args.iter().all(is_symbolically_stable_solution), + Expression::Array { elements, .. } | Expression::Tuple { elements, .. } => { + elements.iter().all(is_symbolically_stable_solution) + } + Expression::Range { + start, step, end, .. + } => { + is_symbolically_stable_solution(start) + && step.as_deref().is_none_or(is_symbolically_stable_solution) + && is_symbolically_stable_solution(end) + } + Expression::ArrayComprehension { expr, filter, .. } => { + is_symbolically_stable_solution(expr) + && filter + .as_deref() + .is_none_or(is_symbolically_stable_solution) + } + Expression::Index { + base, subscripts, .. + } => { + is_symbolically_stable_solution(base) + && subscripts.iter().all(|sub| match sub { + rumoca_core::Subscript::Expr { expr, .. } => { + is_symbolically_stable_solution(expr) + } + _ => true, + }) + } + Expression::FieldAccess { base, .. } => is_symbolically_stable_solution(base), + Expression::VarRef { .. } + | Expression::Literal { value: _, .. } + | Expression::Empty { .. } => true, + } +} + +fn solution_has_blocking_unsliced_multiscalar_ref( + expr: &Expression, + dae: &Dae, +) -> Result { + let scope = DaeVariableScope::new(dae); + solution_has_blocking_unsliced_multiscalar_ref_inner(expr, &scope) +} + +const MAX_SYMBOLIC_SUBSTITUTION_NODES: usize = 64; +const MAX_BLT_SYMBOLIC_CANDIDATE_NODES: usize = 256; + +pub(super) fn solution_is_cheap_for_symbolic_substitution(expr: &Expression) -> bool { + expression_node_count_exceeds(expr, MAX_SYMBOLIC_SUBSTITUTION_NODES) + .is_some_and(|exceeds| !exceeds) +} + +fn expression_is_within_symbolic_candidate_budget(expr: &Expression) -> bool { + expression_node_count_exceeds(expr, MAX_BLT_SYMBOLIC_CANDIDATE_NODES) + .is_some_and(|exceeds| !exceeds) +} + +fn expression_node_count_exceeds(expr: &Expression, limit: usize) -> Option { + let mut count = 0usize; + expression_node_count_visit(expr, limit, &mut count) +} + +fn expression_node_count_visit_all<'a>( + exprs: impl IntoIterator, + limit: usize, + count: &mut usize, +) -> Option { + for expr in exprs { + if expression_node_count_visit(expr, limit, count)? { + return Some(true); + } + } + Some(false) +} + +fn expression_node_count_visit_branches( + branches: &[(Expression, Expression)], + limit: usize, + count: &mut usize, +) -> Option { + for (condition, branch) in branches { + if expression_node_count_visit_all([condition, branch], limit, count)? { + return Some(true); + } + } + Some(false) +} + +fn expression_node_count_visit_subscripts( + subscripts: &[rumoca_core::Subscript], + limit: usize, + count: &mut usize, +) -> Option { + expression_node_count_visit_all( + subscripts.iter().filter_map(|subscript| match subscript { + rumoca_core::Subscript::Expr { expr, .. } => Some(expr.as_ref()), + _ => None, + }), + limit, + count, + ) +} + +fn expression_node_count_visit(expr: &Expression, limit: usize, count: &mut usize) -> Option { + *count = count.checked_add(1)?; + if *count > limit { + return Some(true); + } + match expr { + Expression::Binary { lhs, rhs, .. } => { + expression_node_count_visit_binary(lhs, rhs, limit, count) + } + Expression::Unary { rhs, .. } => expression_node_count_visit(rhs, limit, count), + Expression::BuiltinCall { args, .. } | Expression::FunctionCall { args, .. } => { + expression_node_count_visit_all(args, limit, count) + } + Expression::If { + branches, + else_branch, + .. + } => expression_node_count_visit_if(branches, else_branch, limit, count), + Expression::Array { elements, .. } | Expression::Tuple { elements, .. } => { + expression_node_count_visit_all(elements, limit, count) + } + Expression::Range { + start, step, end, .. + } => expression_node_count_visit_range(start, step.as_deref(), end, limit, count), + Expression::Index { + base, subscripts, .. + } => expression_node_count_visit_index(base, subscripts, limit, count), + Expression::FieldAccess { base, .. } => expression_node_count_visit(base, limit, count), + Expression::ArrayComprehension { + expr, + indices, + filter, + .. + } => expression_node_count_visit_comprehension( + expr, + indices, + filter.as_deref(), + limit, + count, + ), + _ => Some(false), + } +} + +fn expression_node_count_visit_binary( + lhs: &Expression, + rhs: &Expression, + limit: usize, + count: &mut usize, +) -> Option { + if expression_node_count_visit(lhs, limit, count)? { + return Some(true); + } + expression_node_count_visit(rhs, limit, count) +} + +fn expression_node_count_visit_if( + branches: &[(Expression, Expression)], + else_branch: &Expression, + limit: usize, + count: &mut usize, +) -> Option { + if expression_node_count_visit_branches(branches, limit, count)? { + return Some(true); + } + expression_node_count_visit(else_branch, limit, count) +} + +fn expression_node_count_visit_range( + start: &Expression, + step: Option<&Expression>, + end: &Expression, + limit: usize, + count: &mut usize, +) -> Option { + if expression_node_count_visit(start, limit, count)? { + return Some(true); + } + if let Some(step) = step + && expression_node_count_visit(step, limit, count)? + { + return Some(true); + } + expression_node_count_visit(end, limit, count) +} + +fn expression_node_count_visit_index( + base: &Expression, + subscripts: &[rumoca_core::Subscript], + limit: usize, + count: &mut usize, +) -> Option { + if expression_node_count_visit(base, limit, count)? { + return Some(true); + } + expression_node_count_visit_subscripts(subscripts, limit, count) +} + +fn expression_node_count_visit_comprehension( + expr: &Expression, + indices: &[rumoca_core::ComprehensionIndex], + filter: Option<&Expression>, + limit: usize, + count: &mut usize, +) -> Option { + if expression_node_count_visit(expr, limit, count)? { + return Some(true); + } + if expression_node_count_visit_all(indices.iter().map(|index| &index.range), limit, count)? { + return Some(true); + } + filter.map_or(Some(false), |filter| { + expression_node_count_visit(filter, limit, count) + }) +} + +fn is_scalar_reduction_solution_tree(expr: &Expression) -> bool { + match expr { + Expression::BuiltinCall { + function: BuiltinFunction::Sum | BuiltinFunction::Product, + args, + .. + } => args.len() == 1, + Expression::If { + branches, + else_branch, + .. + } => { + branches + .iter() + .all(|(_, branch)| is_scalar_reduction_solution_tree(branch)) + && is_scalar_reduction_solution_tree(else_branch) + } + Expression::Literal { .. } => true, + _ => false, + } +} + +fn solution_has_blocking_unsliced_multiscalar_ref_inner( + expr: &Expression, + scope: &DaeVariableScope<'_>, +) -> Result { + Ok(match expr { + Expression::VarRef { + name, subscripts, .. + } => { + if !subscripts.is_empty() || name.as_str() == "time" { + false + } else { + match scope.shape_for_reference(name) { + Ok(DaeVariableShape::Dimensions(dims)) => { + scalar_count_from_dims(name.var_name(), &dims)? > 1 + } + Ok(DaeVariableShape::StructuredAggregate) => true, + Err(StructuralError::ContractViolation { reason, .. }) + | Err(StructuralError::UnspannedContractViolation { reason }) + if reason.contains("missing DAE variable metadata") + && reference_has_scalar_indices(name) => + { + true + } + Err(err) => return Err(err), + } + } + } + Expression::BuiltinCall { + function: BuiltinFunction::Sum | BuiltinFunction::Product, + args, + .. + } if args.len() == 1 => false, + Expression::BuiltinCall { args, .. } | Expression::FunctionCall { args, .. } => { + exprs_have_blocking_unsliced_multiscalar_ref(args, scope)? + } + Expression::Binary { lhs, rhs, .. } => { + solution_has_blocking_unsliced_multiscalar_ref_inner(lhs, scope)? + || solution_has_blocking_unsliced_multiscalar_ref_inner(rhs, scope)? + } + Expression::Unary { rhs, .. } => { + solution_has_blocking_unsliced_multiscalar_ref_inner(rhs, scope)? + } + Expression::If { + branches, + else_branch, + .. + } => { + let mut blocked = false; + for (condition, branch) in branches { + blocked |= solution_has_blocking_unsliced_multiscalar_ref_inner(condition, scope)?; + blocked |= solution_has_blocking_unsliced_multiscalar_ref_inner(branch, scope)?; + } + blocked || solution_has_blocking_unsliced_multiscalar_ref_inner(else_branch, scope)? + } + Expression::Array { elements, .. } | Expression::Tuple { elements, .. } => { + exprs_have_blocking_unsliced_multiscalar_ref(elements, scope)? + } + Expression::Range { + start, step, end, .. + } => { + let step_blocked = match step.as_deref() { + Some(step) => solution_has_blocking_unsliced_multiscalar_ref_inner(step, scope)?, + None => false, + }; + solution_has_blocking_unsliced_multiscalar_ref_inner(start, scope)? + || step_blocked + || solution_has_blocking_unsliced_multiscalar_ref_inner(end, scope)? + } + Expression::ArrayComprehension { expr, filter, .. } => { + let filter_blocked = match filter.as_deref() { + Some(filter) => { + solution_has_blocking_unsliced_multiscalar_ref_inner(filter, scope)? + } + None => false, + }; + solution_has_blocking_unsliced_multiscalar_ref_inner(expr, scope)? || filter_blocked + } + Expression::Index { + base, subscripts, .. + } => { + if solution_has_blocking_unsliced_multiscalar_ref_inner(base, scope)? { + true + } else { + subscripts_have_blocking_unsliced_multiscalar_ref(subscripts, scope)? + } + } + Expression::FieldAccess { base, .. } => { + solution_has_blocking_unsliced_multiscalar_ref_inner(base, scope)? + } + Expression::Literal { .. } | Expression::Empty { .. } => false, + }) +} + +fn subscripts_have_blocking_unsliced_multiscalar_ref( + subscripts: &[rumoca_core::Subscript], + scope: &DaeVariableScope<'_>, +) -> Result { + let mut blocked = false; + for sub in subscripts { + if let rumoca_core::Subscript::Expr { expr, .. } = sub { + blocked |= solution_has_blocking_unsliced_multiscalar_ref_inner(expr, scope)?; + } + } + Ok(blocked) +} + +fn exprs_have_blocking_unsliced_multiscalar_ref( + exprs: &[Expression], + scope: &DaeVariableScope<'_>, +) -> Result { + for expr in exprs { + if solution_has_blocking_unsliced_multiscalar_ref_inner(expr, scope)? { + return Ok(true); + } + } + Ok(false) +} + +pub(super) fn collect_var_ref_nodes( + expr: &Expression, + out: &mut Vec<(Reference, Vec)>, +) { + struct Collector<'out> { + out: &'out mut Vec<(Reference, Vec)>, + } + + impl ExpressionVisitor for Collector<'_> { + fn visit_var_ref(&mut self, name: &Reference, subscripts: &[rumoca_core::Subscript]) { + self.out.push((name.clone(), subscripts.to_vec())); + self.walk_var_ref(name, subscripts); + } + } + + Collector { out }.visit_expression(expr); +} + +pub(super) fn collect_exact_reference_expr_names_in_dae( + dae: &Dae, + expr: &Expression, + out: &mut Vec, +) { + match expr { + Expression::VarRef { .. } => { + if let Some(name) = exact_reference_expr_name_in_dae(dae, expr) { + out.push(name); + } + } + Expression::Index { base, .. } | Expression::FieldAccess { base, .. } => { + if let Some(name) = exact_reference_expr_name_in_dae(dae, expr) { + out.push(name); + } else { + collect_exact_reference_expr_names_in_dae(dae, base, out); + } + } + Expression::Binary { lhs, rhs, .. } => { + collect_exact_reference_expr_names_in_dae(dae, lhs, out); + collect_exact_reference_expr_names_in_dae(dae, rhs, out); + } + Expression::Unary { rhs, .. } => collect_exact_reference_expr_names_in_dae(dae, rhs, out), + Expression::BuiltinCall { args, .. } | Expression::FunctionCall { args, .. } => { + for arg in args { + collect_exact_reference_expr_names_in_dae(dae, arg, out); + } + } + Expression::If { + branches, + else_branch, + .. + } => { + for (condition, value) in branches { + collect_exact_reference_expr_names_in_dae(dae, condition, out); + collect_exact_reference_expr_names_in_dae(dae, value, out); + } + collect_exact_reference_expr_names_in_dae(dae, else_branch, out); + } + Expression::Array { elements, .. } | Expression::Tuple { elements, .. } => { + for element in elements { + collect_exact_reference_expr_names_in_dae(dae, element, out); + } + } + Expression::Range { + start, step, end, .. + } => { + collect_exact_reference_expr_names_in_dae(dae, start, out); + if let Some(step) = step { + collect_exact_reference_expr_names_in_dae(dae, step, out); + } + collect_exact_reference_expr_names_in_dae(dae, end, out); + } + Expression::ArrayComprehension { expr, filter, .. } => { + collect_exact_reference_expr_names_in_dae(dae, expr, out); + if let Some(filter) = filter { + collect_exact_reference_expr_names_in_dae(dae, filter, out); + } + } + Expression::Literal { .. } | Expression::Empty { .. } => {} + } +} + +pub(super) fn exact_reference_expr_name(expr: &Expression) -> Option { + match expr { + Expression::VarRef { + name, subscripts, .. + } => exact_name_with_subscripts(name.var_name().as_str(), subscripts).map(VarName::new), + Expression::Index { + base, subscripts, .. + } => { + let base_name = exact_reference_expr_name(base)?; + exact_name_with_subscripts(base_name.as_str(), subscripts).map(VarName::new) + } + Expression::FieldAccess { base, field, .. } => { + let base_name = exact_reference_expr_name(base)?; + Some(VarName::new(format!("{}.{field}", base_name.as_str()))) + } + _ => None, + } +} + +pub(super) fn exact_reference_expr_name_in_dae(dae: &Dae, expr: &Expression) -> Option { + match expr { + Expression::VarRef { + name, subscripts, .. + } => exact_name_with_subscripts_in_dae(dae, name.var_name().as_str(), subscripts) + .map(VarName::new), + Expression::Index { + base, subscripts, .. + } => { + let base_name = exact_reference_expr_name_in_dae(dae, base)?; + exact_name_with_subscripts_in_dae(dae, base_name.as_str(), subscripts).map(VarName::new) + } + Expression::FieldAccess { base, field, .. } => { + let base_name = exact_reference_expr_name_in_dae(dae, base)?; + Some(VarName::new(format!("{}.{field}", base_name.as_str()))) + } + _ => None, + } +} + +fn exact_name_with_subscripts(base: &str, subscripts: &[rumoca_core::Subscript]) -> Option { + if subscripts.is_empty() { + return Some(base.to_string()); + } + let mut indices = Vec::with_capacity(subscripts.len()); + for subscript in subscripts { + indices.push(exact_subscript_index(subscript)?.to_string()); + } + Some(format!("{base}[{}]", indices.join(","))) +} + +fn exact_name_with_subscripts_in_dae( + dae: &Dae, + base: &str, + subscripts: &[rumoca_core::Subscript], +) -> Option { + if subscripts.is_empty() { + return Some(base.to_string()); + } + let mut indices = Vec::with_capacity(subscripts.len()); + for subscript in subscripts { + indices.push(exact_subscript_index_in_dae(dae, subscript)?.to_string()); + } + Some(format!("{base}[{}]", indices.join(","))) +} + +fn exact_subscript_index(subscript: &rumoca_core::Subscript) -> Option { + match subscript { + rumoca_core::Subscript::Index { value, .. } => Some(*value), + rumoca_core::Subscript::Expr { expr, .. } => exact_index_expr(expr), + rumoca_core::Subscript::Colon { .. } => None, + } +} + +pub(super) fn exact_subscript_index_in_dae( + dae: &Dae, + subscript: &rumoca_core::Subscript, +) -> Option { + match subscript { + rumoca_core::Subscript::Index { value, .. } => Some(*value), + rumoca_core::Subscript::Expr { expr, .. } => exact_index_expr_in_dae(dae, expr), + rumoca_core::Subscript::Colon { .. } => None, + } +} + +fn exact_index_expr_in_dae(dae: &Dae, expr: &Expression) -> Option { + if let Some(value) = exact_index_expr(expr) { + return Some(value); + } + match expr { + Expression::VarRef { + name, subscripts, .. + } if subscripts.is_empty() => fixed_integer_parameter_start(dae, name.var_name()), + Expression::Binary { op, lhs, rhs, .. } => { + let lhs = exact_index_expr_in_dae(dae, lhs)?; + let rhs = exact_index_expr_in_dae(dae, rhs)?; + exact_binary_index_value(op, lhs, rhs) + } + Expression::BuiltinCall { function, args, .. } => { + exact_builtin_index_value_in_dae(dae, function, args) + } + _ => None, + } +} + +fn exact_index_expr(expr: &Expression) -> Option { + match expr { + Expression::Literal { + value: rumoca_core::Literal::Integer(value), + .. + } => Some(*value), + Expression::Literal { + value: rumoca_core::Literal::Real(value), + .. + } if value.is_finite() && value.fract() == 0.0 => Some(*value as i64), + Expression::Unary { op, rhs, .. } => match op { + OpUnary::Plus => exact_index_expr(rhs), + OpUnary::Minus => exact_index_expr(rhs).and_then(|value| value.checked_neg()), + _ => None, + }, + Expression::Binary { op, lhs, rhs, .. } => { + let lhs = exact_index_expr(lhs)?; + let rhs = exact_index_expr(rhs)?; + exact_binary_index_value(op, lhs, rhs) + } + Expression::BuiltinCall { function, args, .. } => exact_builtin_index_value(function, args), + _ => None, + } +} + +fn exact_binary_index_value(op: &OpBinary, lhs: i64, rhs: i64) -> Option { + match op { + OpBinary::Add | OpBinary::AddElem => lhs.checked_add(rhs), + OpBinary::Sub | OpBinary::SubElem => lhs.checked_sub(rhs), + OpBinary::Mul | OpBinary::MulElem => lhs.checked_mul(rhs), + OpBinary::Div | OpBinary::DivElem if rhs != 0 && lhs % rhs == 0 => Some(lhs / rhs), + _ => None, + } +} + +fn exact_builtin_index_value(function: &BuiltinFunction, args: &[Expression]) -> Option { + match function { + BuiltinFunction::Integer if args.len() == 1 => { + exact_real_expr(&args[0]).and_then(finite_floor_i64) + } + BuiltinFunction::Mod if args.len() == 2 => { + let lhs = exact_index_expr(&args[0])?; + let rhs = exact_index_expr(&args[1])?; + (rhs != 0).then(|| lhs.rem_euclid(rhs)) + } + _ => None, + } } -fn has_direct_assignment_form(rhs: &Expression, candidate: &VarName) -> bool { - match rhs { - Expression::Binary { - op: OpBinary::Sub, - lhs, - rhs, - .. - } => is_assignment_target(lhs, candidate) || is_assignment_target(rhs, candidate), - Expression::Unary { - op: OpUnary::Minus, - rhs, - .. - } => has_direct_assignment_form(rhs, candidate), - _ => false, +fn exact_builtin_index_value_in_dae( + dae: &Dae, + function: &BuiltinFunction, + args: &[Expression], +) -> Option { + match function { + BuiltinFunction::Integer if args.len() == 1 => { + exact_real_expr_in_dae(dae, &args[0]).and_then(finite_floor_i64) + } + BuiltinFunction::Mod if args.len() == 2 => { + let lhs = exact_index_expr_in_dae(dae, &args[0])?; + let rhs = exact_index_expr_in_dae(dae, &args[1])?; + (rhs != 0).then(|| lhs.rem_euclid(rhs)) + } + _ => exact_builtin_index_value(function, args), } } -fn is_assignment_target(expr: &Expression, candidate: &VarName) -> bool { +fn exact_real_expr_in_dae(dae: &Dae, expr: &Expression) -> Option { + if let Some(value) = exact_real_expr(expr) { + return Some(value); + } match expr { Expression::VarRef { name, subscripts, .. - } => var_ref_matches_unknown(name, subscripts, candidate), - _ => false, + } if subscripts.is_empty() => { + fixed_integer_parameter_start(dae, name.var_name()).map(|value| value as f64) + } + Expression::Binary { op, lhs, rhs, .. } => { + let lhs = exact_real_expr_in_dae(dae, lhs)?; + let rhs = exact_real_expr_in_dae(dae, rhs)?; + exact_binary_real_value(op, lhs, rhs) + } + Expression::BuiltinCall { function, args, .. } => { + exact_builtin_real_value_in_dae(dae, function, args) + } + _ => None, } } -/// Returns true if the expression is a single variable reference or its -/// negation — i.e., a trivial alias like `x` or `-x`. -fn is_trivial_alias(expr: &Expression) -> bool { +fn exact_real_expr(expr: &Expression) -> Option { match expr { - Expression::VarRef { .. } => true, - Expression::Unary { - op: OpUnary::Minus, - rhs, + Expression::Literal { + value: rumoca_core::Literal::Integer(value), .. - } => is_trivial_alias(rhs), - Expression::BuiltinCall { - function: BuiltinFunction::Der, - args, + } => Some(*value as f64), + Expression::Literal { + value: rumoca_core::Literal::Real(value), .. - } => args.len() == 1 && matches!(&args[0], Expression::VarRef { .. }), - _ => false, + } if value.is_finite() => Some(*value), + Expression::Unary { op, rhs, .. } => match op { + OpUnary::Plus => exact_real_expr(rhs), + OpUnary::Minus => exact_real_expr(rhs).map(|value| -value), + _ => None, + }, + Expression::Binary { op, lhs, rhs, .. } => { + let lhs = exact_real_expr(lhs)?; + let rhs = exact_real_expr(rhs)?; + exact_binary_real_value(op, lhs, rhs) + } + Expression::BuiltinCall { function, args, .. } => exact_builtin_real_value(function, args), + _ => None, } } -fn is_symbolically_stable_solution(expr: &Expression) -> bool { - match expr { - Expression::If { .. } => false, - Expression::BuiltinCall { function, args, .. } => { - !matches!( - function, - rumoca_core::BuiltinFunction::Smooth - | rumoca_core::BuiltinFunction::NoEvent - | rumoca_core::BuiltinFunction::Homotopy - ) && args.iter().all(is_symbolically_stable_solution) - } - Expression::Binary { lhs, rhs, .. } => { - is_symbolically_stable_solution(lhs) && is_symbolically_stable_solution(rhs) - } - Expression::Unary { rhs, .. } => is_symbolically_stable_solution(rhs), - Expression::FunctionCall { args, .. } => args.iter().all(is_symbolically_stable_solution), - Expression::Array { elements, .. } | Expression::Tuple { elements, .. } => { - elements.iter().all(is_symbolically_stable_solution) - } - Expression::Range { - start, step, end, .. - } => { - is_symbolically_stable_solution(start) - && step.as_deref().is_none_or(is_symbolically_stable_solution) - && is_symbolically_stable_solution(end) - } - Expression::ArrayComprehension { expr, filter, .. } => { - is_symbolically_stable_solution(expr) - && filter - .as_deref() - .is_none_or(is_symbolically_stable_solution) - } - Expression::Index { - base, subscripts, .. - } => { - is_symbolically_stable_solution(base) - && subscripts.iter().all(|sub| match sub { - rumoca_core::Subscript::Expr { expr, .. } => { - is_symbolically_stable_solution(expr) - } - _ => true, - }) - } - Expression::FieldAccess { base, .. } => is_symbolically_stable_solution(base), - Expression::VarRef { .. } - | Expression::Literal { value: _, .. } - | Expression::Empty { .. } => true, +fn exact_binary_real_value(op: &OpBinary, lhs: f64, rhs: f64) -> Option { + match op { + OpBinary::Add | OpBinary::AddElem => Some(lhs + rhs), + OpBinary::Sub | OpBinary::SubElem => Some(lhs - rhs), + OpBinary::Mul | OpBinary::MulElem => Some(lhs * rhs), + OpBinary::Div | OpBinary::DivElem if rhs != 0.0 => Some(lhs / rhs), + _ => None, } + .filter(|value| value.is_finite()) } -pub(super) fn collect_var_ref_nodes( - expr: &Expression, - out: &mut Vec<(Reference, Vec)>, -) { - struct Collector<'out> { - out: &'out mut Vec<(Reference, Vec)>, +fn exact_builtin_real_value(function: &BuiltinFunction, args: &[Expression]) -> Option { + match function { + BuiltinFunction::Integer if args.len() == 1 => exact_real_expr(&args[0]).map(f64::floor), + BuiltinFunction::Mod if args.len() == 2 => { + let lhs = exact_real_expr(&args[0])?; + let rhs = exact_real_expr(&args[1])?; + (rhs != 0.0).then(|| lhs - (lhs / rhs).floor() * rhs) + } + _ => None, } + .filter(|value| value.is_finite()) +} - impl ExpressionVisitor for Collector<'_> { - fn visit_var_ref(&mut self, name: &Reference, subscripts: &[rumoca_core::Subscript]) { - self.out.push((name.clone(), subscripts.to_vec())); - self.walk_var_ref(name, subscripts); +fn exact_builtin_real_value_in_dae( + dae: &Dae, + function: &BuiltinFunction, + args: &[Expression], +) -> Option { + match function { + BuiltinFunction::Integer if args.len() == 1 => { + exact_real_expr_in_dae(dae, &args[0]).map(f64::floor) + } + BuiltinFunction::Mod if args.len() == 2 => { + let lhs = exact_real_expr_in_dae(dae, &args[0])?; + let rhs = exact_real_expr_in_dae(dae, &args[1])?; + (rhs != 0.0).then(|| lhs - (lhs / rhs).floor() * rhs) } + _ => exact_builtin_real_value(function, args), } + .filter(|value| value.is_finite()) +} - Collector { out }.visit_expression(expr); +fn finite_floor_i64(value: f64) -> Option { + let value = value.floor(); + (value.is_finite() && value >= i64::MIN as f64 && value <= i64::MAX as f64) + .then_some(value as i64) +} + +fn fixed_integer_parameter_start(dae: &Dae, name: &VarName) -> Option { + let var = dae + .variables + .parameters + .get(name) + .filter(|var| !var.is_tunable) + .or_else(|| dae.variables.constants.get(name))?; + let start = var.start.as_ref()?; + exact_index_expr_in_dae(dae, start) } fn dae_var_size(dae: &Dae, name: &VarName) -> Result { @@ -971,9 +2558,18 @@ pub(super) fn substitution_for_var( expr: Expression, ) -> Result { let scope = DaeVariableScope::new(dae); + let expr = project_scalarized_unknown_solution(&var_name, expr); + let var_dims = scope.dims(&var_name)?; + let replacement_dims = replacement_expr_dims(dae, &expr).or_else(|err| { + if var_dims.is_empty() && is_scalar_external_alias_expr(&expr, &err) { + Ok(Vec::new()) + } else { + Err(err) + } + })?; Ok(Substitution { - var_dims: scope.dims(&var_name)?, - replacement_dims: replacement_expr_dims(dae, &expr)?, + var_dims, + replacement_dims, env_keys: vec![var_name.as_str().to_string()], var_ref: scope .exact(&var_name) @@ -984,6 +2580,39 @@ pub(super) fn substitution_for_var( }) } +fn is_scalar_external_alias_expr(expr: &Expression, err: &StructuralError) -> bool { + matches!( + err, + StructuralError::ContractViolation { reason, .. } + | StructuralError::UnspannedContractViolation { reason } + if reason.contains("missing DAE variable metadata") + ) && matches!(expr, Expression::VarRef { subscripts, .. } if subscripts.is_empty()) +} + +fn project_scalarized_unknown_solution(var_name: &VarName, expr: Expression) -> Expression { + let Some(scalar_name) = rumoca_core::parse_scalar_name(var_name.as_str()) else { + return expr; + }; + let Some(span) = expr.span() else { + return expr; + }; + let Ok(subscripts) = scalar_name + .indices + .iter() + .map(|index| { + rumoca_core::Subscript::try_generated_index( + *index, + span, + "scalarized unknown solution projection", + ) + }) + .collect::, _>>() + else { + return expr; + }; + project_replacement_expr_with_subscripts(&expr, &subscripts, span).unwrap_or(expr) +} + fn replacement_expr_dims(dae: &Dae, expr: &Expression) -> Result, StructuralError> { Ok(match expr { Expression::VarRef { @@ -1131,19 +2760,35 @@ fn scalar_blt_solution( return Ok(None); }; let var_name = raw_var_name.clone(); - if !can_eliminate_scalar_unknown(dae, &var_name, runtime_protected_unknowns)? { - return Ok(None); - } - let eq_idx = equation.0; let is_output = output_partition_contains_unknown(dae, &var_name); let has_state_derivative = equation_has_state_derivative(dae, eq_idx, state_derivative_matcher); + if !can_eliminate_scalar_unknown( + dae, + &var_name, + runtime_protected_unknowns, + has_state_derivative, + )? { + return Ok(None); + } if has_state_derivative && !is_output { return Ok(None); } - let eq_rhs = - apply_substitutions_in_order(&dae.continuous.equations[eq_idx].rhs, substitutions)?; + let Some(eq_rhs) = apply_substitutions_for_symbolic_candidate( + &dae.continuous.equations[eq_idx].rhs, + substitutions, + dae, + )? + else { + return Ok(None); + }; + if should_preserve_runtime_sensitive_continuous_assignment(dae, &eq_rhs) { + return Ok(None); + } + if expr_contains_indexed_multiscalar_slice_ref(&eq_rhs, dae)? { + return Ok(None); + } if is_flow_equation_origin(&dae.continuous.equations[eq_idx].origin) && expr_contains_indexed_multiscalar_ref(&eq_rhs, dae)? { @@ -1161,12 +2806,95 @@ fn scalar_blt_solution( let Some(solution) = stable_solution_for_unknown(dae, &eq_rhs, &var_name)? else { return Ok(None); }; - if is_output && !is_trivial_alias(&solution) { + if !solution_matches_candidate_scalar_shape(dae, &var_name, &solution)? { + return Ok(None); + } + if scalar_blt_solution_would_break_aggregate_element(dae, &var_name, &solution)? { + return Ok(None); + } + if is_output + && !is_trivial_alias_in_dae(dae, &solution) + && !is_internal_component_output(dae, &var_name) + { + return Ok(None); + } + if !is_trivial_alias_in_dae(dae, &solution) + && !solution_is_cheap_for_symbolic_substitution(&solution) + { return Ok(None); } Ok(Some((eq_idx, var_name, solution))) } +pub(super) fn apply_substitutions_for_symbolic_candidate( + expr: &Expression, + substitutions: &[Substitution], + dae: &Dae, +) -> Result, StructuralError> { + let dae_scope = DaeVariableScope::new(dae); + let mut out = apply_record_field_aggregate_substitutions(expr, substitutions, Some(dae)); + for sub in substitutions { + if !expr_contains_substitution_target_in_scope(&out, sub, Some(&dae_scope), Some(dae)) { + continue; + } + if !is_trivial_alias(&sub.expr) && !solution_is_cheap_for_symbolic_substitution(&sub.expr) { + return Ok(None); + } + out = SubstituteVarRewriter { + substitution: sub, + replacement: &sub.expr, + replacement_dims: &sub.replacement_dims, + derivative_replacement: None, + dae_scope: Some(&dae_scope), + dae_context: Some(dae), + } + .rewrite_expression(&out)?; + if !expression_is_within_symbolic_candidate_budget(&out) { + return Ok(None); + } + } + Ok(Some(out)) +} + +fn is_internal_component_output(dae: &Dae, var_name: &VarName) -> bool { + let Some(var) = dae.variables.outputs.get(var_name).or_else(|| { + rumoca_ir_dae::component_base_name(var_name.as_str()) + .and_then(|base| dae.variables.outputs.get(&VarName::new(base))) + }) else { + return false; + }; + var.component_ref + .as_ref() + .is_some_and(|component_ref| component_ref.parts.len() > 1) + || rumoca_core::component_reference_from_flat_name(var_name, var.source_span) + .is_some_and(|component_ref| component_ref.parts.len() > 1) +} + +fn is_local_component_output(dae: &Dae, var_name: &VarName) -> bool { + dae.variables + .outputs + .get(var_name) + .or_else(|| { + rumoca_ir_dae::component_base_name(var_name.as_str()) + .and_then(|base| dae.variables.outputs.get(&VarName::new(base))) + }) + .is_some_and(|var| matches!(var.causality, dae::VariableCausality::Local)) +} + +fn is_internal_component_continuous_var(dae: &Dae, var_name: &VarName) -> bool { + let Some(var) = dae_var(dae, var_name).or_else(|| { + rumoca_ir_dae::component_base_name(var_name.as_str()) + .and_then(|base| dae_var(dae, &VarName::new(base))) + }) else { + return false; + }; + var.component_ref + .as_ref() + .is_some_and(|component_ref| component_ref.parts.len() > 1) + || rumoca_core::component_reference_from_flat_name(var_name, var.source_span) + .is_some_and(|component_ref| component_ref.parts.len() > 1) +} + fn algebraic_or_output_unknown(unknown: &UnknownId) -> Option<&VarName> { match unknown { UnknownId::Variable(name) => Some(name), @@ -1174,18 +2902,50 @@ fn algebraic_or_output_unknown(unknown: &UnknownId) -> Option<&VarName> { } } +fn is_derivative_alias_expr(expr: &Expression) -> bool { + matches!( + expr, + Expression::BuiltinCall { + function: BuiltinFunction::Der, + args, + .. + } if args.len() == 1 + ) +} + fn can_eliminate_scalar_unknown( dae: &Dae, var_name: &VarName, runtime_protected_unknowns: &IndexSet, + has_state_derivative: bool, ) -> Result { - Ok( - !is_runtime_protected_unknown(var_name, runtime_protected_unknowns) - && !unknown_is_fixed(dae, var_name) - && !dae.variables.states.contains_key(var_name) - && !is_scalarized_element_of_aggregate(dae, var_name)? - && dae_var_size(dae, var_name)? == 1, - ) + if is_runtime_protected_unknown(var_name, runtime_protected_unknowns) + || unknown_is_fixed(dae, var_name) + || dae.variables.states.contains_key(var_name) + || dae_var_size(dae, var_name)? != 1 + { + return Ok(false); + } + if !is_scalarized_element_of_aggregate(dae, var_name)? { + return Ok(true); + } + Ok(!has_state_derivative && scalarized_element_has_non_connection_use(dae, var_name)) +} + +pub(super) fn scalar_blt_solution_would_break_aggregate_element( + dae: &Dae, + var_name: &VarName, + solution: &Expression, +) -> Result { + if !is_scalarized_element_of_aggregate(dae, var_name)? { + return Ok(false); + } + if is_trivial_alias_in_dae(dae, solution) { + return Ok(false); + } + Ok(scalarized_element_has_coupled_derivative_use(dae, var_name) + || (expr_contains_runtime_sensitive_operator(solution) + && scalarized_element_has_non_connection_use(dae, var_name))) } fn can_use_equation_for_elimination(dae: &Dae, eq_idx: usize) -> bool { @@ -1230,11 +2990,11 @@ fn stable_solution_for_unknown( rhs: &Expression, var_name: &VarName, ) -> Result, StructuralError> { - let Some(solution) = try_solve_for_unknown(rhs, var_name) else { + let Some(solution) = try_solve_for_unknown_in_dae(dae, rhs, var_name) else { return Ok(None); }; if expr_contains_var(&solution, var_name) - || expr_contains_unsliced_multiscalar_ref(&solution, dae)? + || solution_has_blocking_unsliced_multiscalar_ref(&solution, dae)? || !is_symbolically_stable_solution(&solution) { return Ok(None); @@ -1253,13 +3013,33 @@ pub fn apply_substitutions_to_expr( pub(crate) fn apply_substitutions_to_expr_with_derivatives( expr: &Expression, substitutions: &[Substitution], + derivative_replacement_for: impl FnMut(&Substitution) -> Result, StructuralError>, +) -> Result { + apply_substitutions_to_expr_with_derivatives_and_dae( + expr, + substitutions, + None, + derivative_replacement_for, + ) +} + +pub(crate) fn apply_substitutions_to_expr_with_derivatives_and_dae( + expr: &Expression, + substitutions: &[Substitution], + dae_context: Option<&Dae>, mut derivative_replacement_for: impl FnMut( &Substitution, ) -> Result, StructuralError>, ) -> Result { - let mut out = apply_record_field_aggregate_substitutions(expr, substitutions); + let substitution_scope = dae_context.map(DaeVariableScope::new); + let mut out = apply_record_field_aggregate_substitutions(expr, substitutions, dae_context); for sub in substitutions { - if expr_contains_substitution_target(&out, sub) { + if expr_contains_substitution_target_in_scope( + &out, + sub, + substitution_scope.as_ref(), + dae_context, + ) { let derivative_replacement = if expr_contains_derivative_substitution_target(&out, sub) { derivative_replacement_for(sub)? @@ -1271,6 +3051,8 @@ pub(crate) fn apply_substitutions_to_expr_with_derivatives( replacement: &sub.expr, replacement_dims: &sub.replacement_dims, derivative_replacement: derivative_replacement.as_ref(), + dae_scope: substitution_scope.as_ref(), + dae_context, } .rewrite_expression(&out)?; } @@ -1281,8 +3063,11 @@ pub(crate) fn apply_substitutions_to_expr_with_derivatives( fn apply_record_field_aggregate_substitutions( expr: &Expression, substitutions: &[Substitution], + dae_context: Option<&Dae>, ) -> Expression { - let aggregate_alias_groups = aggregate_alias_substitution_groups(substitutions); + let dae_scope = dae_context.map(DaeVariableScope::new); + let aggregate_alias_groups = + aggregate_alias_substitution_groups(substitutions, dae_scope.as_ref()); let complex_groups = complex_field_substitution_groups(substitutions); if aggregate_alias_groups.is_empty() && complex_groups.is_empty() { return expr.clone(); @@ -1290,6 +3075,7 @@ fn apply_record_field_aggregate_substitutions( RecordFieldAggregateRewriter { aggregate_alias_groups, complex_groups, + dae_scope, } .rewrite_expression(expr) } @@ -1302,7 +3088,7 @@ struct AggregateAliasSubstitutionGroup { } impl AggregateAliasSubstitutionGroup { - fn insert(&mut self, indices: Vec, expr: Expression) { + fn insert(&mut self, indices: Vec, expr: Expression, replacement_dims: &[i64]) { if indices.len() > self.dims.len() { self.dims.resize(indices.len(), 0); } @@ -1310,7 +3096,10 @@ impl AggregateAliasSubstitutionGroup { self.dims[idx] = self.dims[idx].max(*value); } self.replacement_base = replacement_aggregate_base(&expr, &indices, &self.replacement_base); - self.values.insert(indices, expr); + self.values.insert( + indices.clone(), + project_scalar_alias_replacement(expr, &indices, replacement_dims), + ); } fn to_replacement_expr(&self, span: rumoca_core::Span) -> Option { @@ -1344,6 +3133,22 @@ impl AggregateAliasSubstitutionGroup { }) } + fn covers_dims(&self, dims: &[usize]) -> bool { + !dims.is_empty() && self.dims == dims && self.values.len() == dims.iter().product::() + } + + fn to_partial_replacement_expr( + &self, + name: &Reference, + dims: &[usize], + span: rumoca_core::Span, + ) -> Option { + if dims.is_empty() || dims.iter().product::() <= 1 || self.values.is_empty() { + return None; + } + self.partial_array_expr_at_depth(name, dims, 0, &mut Vec::new(), span) + } + fn expected_len(&self) -> usize { self.dims.iter().product() } @@ -1363,14 +3168,124 @@ impl AggregateAliasSubstitutionGroup { elements.push(self.array_expr_at_depth(depth + 1, current, span)?); current.pop(); } - Some(Expression::Array { - elements, - is_matrix: depth == 0 && self.dims.len() == 2, + Some(Expression::Array { + elements, + is_matrix: depth == 0 && self.dims.len() == 2, + span, + }) + } + + fn partial_array_expr_at_depth( + &self, + name: &Reference, + dims: &[usize], + depth: usize, + current: &mut Vec, + span: rumoca_core::Span, + ) -> Option { + if depth >= dims.len() { + return self + .values + .get(current) + .cloned() + .or_else(|| scalar_ref_for_indices(name, current, span)); + } + let mut elements = Vec::with_capacity(dims[depth]); + for index in 1..=dims[depth] { + current.push(index); + elements.push(self.partial_array_expr_at_depth( + name, + dims, + depth + 1, + current, + span, + )?); + current.pop(); + } + Some(Expression::Array { + elements, + is_matrix: depth == 0 && dims.len() == 2, + span, + }) + } +} + +fn project_scalar_alias_replacement( + expr: Expression, + indices: &[usize], + replacement_dims: &[i64], +) -> Expression { + if replacement_dims.is_empty() || indices.is_empty() { + return expr; + } + let replacement_rank = replacement_dims.len(); + let start = indices.len().saturating_sub(replacement_rank); + let projected_indices = &indices[start..]; + let Some(span) = expr.span() else { + return expr; + }; + let Ok(owner) = span.require_provenance("scalar alias replacement projection") else { + return expr; + }; + let projected_subscripts = projected_indices + .iter() + .map(|index| rumoca_core::Subscript::generated_index_with_provenance(*index as i64, owner)) + .collect::>(); + match expr { + Expression::VarRef { + name, + mut subscripts, + span, + } => { + let existing_rank = subscripts.len(); + if replacement_rank <= existing_rank { + return Expression::VarRef { + name, + subscripts, + span, + }; + } + let missing_rank = replacement_rank - existing_rank; + let start = projected_subscripts.len().saturating_sub(missing_rank); + subscripts.extend(projected_subscripts.into_iter().skip(start)); + Expression::VarRef { + name, + subscripts, + span, + } + } + _ => Expression::Index { + base: Box::new(expr), + subscripts: projected_subscripts, span, - }) + }, } } +fn scalar_ref_for_indices( + name: &Reference, + indices: &[usize], + span: rumoca_core::Span, +) -> Option { + let Ok(owner) = span.require_provenance("scalar alias replacement reference") else { + return Some(Expression::VarRef { + name: name.clone(), + subscripts: Vec::new(), + span, + }); + }; + Some(Expression::VarRef { + name: name.clone(), + subscripts: indices + .iter() + .map(|index| { + rumoca_core::Subscript::generated_index_with_provenance(*index as i64, owner) + }) + .collect(), + span, + }) +} + fn replacement_aggregate_base( expr: &Expression, expected_indices: &[usize], @@ -1386,26 +3301,45 @@ fn replacement_aggregate_base( fn aggregate_alias_substitution_groups( substitutions: &[Substitution], + dae_scope: Option<&DaeVariableScope<'_>>, ) -> IndexMap { let mut groups = IndexMap::new(); for substitution in substitutions { - let Some((base, indices)) = scalar_substitution_target_key(substitution) else { + let target = scalar_substitution_target_key_in_scope(substitution, dae_scope); + let Some((base, indices)) = target else { continue; }; groups .entry(base.into_var_name()) .or_insert_with(AggregateAliasSubstitutionGroup::default) - .insert(indices, substitution.expr.clone()); + .insert( + indices, + substitution.expr.clone(), + &substitution.replacement_dims, + ); } groups } +fn scalar_substitution_target_key_in_scope( + substitution: &Substitution, + dae_scope: Option<&DaeVariableScope<'_>>, +) -> Option<(Reference, Vec)> { + match dae_scope { + Some(scope) => scope.scalarized_aggregate_target(&substitution.var_name), + None => scalar_substitution_target_key(substitution), + } +} + fn scalar_substitution_target_key(substitution: &Substitution) -> Option<(Reference, Vec)> { if let Some(var_ref) = &substitution.var_ref && let Some(key) = scalar_var_ref_key_from_reference(var_ref) { return Some(key); } + if let Some(key) = embedded_indexed_path_key(substitution.var_name.as_str()) { + return Some(key); + } let scalar = rumoca_core::parse_scalar_name(substitution.var_name.as_str())?; let indices = scalar .indices @@ -1421,6 +3355,36 @@ fn scalar_substitution_target_key(substitution: &Substitution) -> Option<(Refere } } +fn embedded_indexed_path_key(raw: &str) -> Option<(Reference, Vec)> { + let mut base = String::with_capacity(raw.len()); + let mut indices = Vec::new(); + let mut chars = raw.char_indices().peekable(); + while let Some((_, ch)) = chars.next() { + if ch != '[' { + base.push(ch); + continue; + } + let mut value = String::new(); + let mut closed = false; + for (_, inner) in chars.by_ref() { + if inner == ']' { + closed = true; + break; + } + value.push(inner); + } + if !closed || value.is_empty() || !value.chars().all(|c| c.is_ascii_digit()) { + return None; + } + let index = value.parse::().ok()?; + if index == 0 { + return None; + } + indices.push(index); + } + (!indices.is_empty() && base != raw).then_some((Reference::new(base), indices)) +} + fn scalar_var_ref_key(expr: &Expression) -> Option<(Reference, Vec)> { let Expression::VarRef { name, subscripts, .. @@ -1439,6 +3403,9 @@ fn scalar_var_ref_key_from_reference(reference: &Reference) -> Option<(Reference let mut base = component_ref.clone(); let mut indices = Vec::new(); for part in &mut base.parts { + if part.subs.is_empty() { + continue; + } indices.extend(positive_usize_subscripts(&part.subs)?); part.subs.clear(); } @@ -1483,7 +3450,17 @@ fn references_same_base(lhs: &Reference, rhs: &Reference) -> bool { } fn reference_has_scalar_indices(reference: &Reference) -> bool { - scalar_var_ref_key_from_reference(reference).is_some() + reference + .component_ref() + .is_some_and(component_ref_has_scalar_indices) +} + +fn component_ref_has_scalar_indices(component_ref: &rumoca_core::ComponentReference) -> bool { + component_ref + .parts + .iter() + .flat_map(|part| &part.subs) + .any(|subscript| positive_usize_subscript(subscript).is_some()) } fn substitution_has_scalar_indices(substitution: &Substitution) -> bool { @@ -1559,12 +3536,13 @@ fn complex_field_substitution_groups( groups } -struct RecordFieldAggregateRewriter { +struct RecordFieldAggregateRewriter<'a> { aggregate_alias_groups: IndexMap, complex_groups: IndexMap, + dae_scope: Option>, } -impl ExpressionRewriter for RecordFieldAggregateRewriter { +impl ExpressionRewriter for RecordFieldAggregateRewriter<'_> { fn rewrite_var_ref_expression( &mut self, name: &Reference, @@ -1573,20 +3551,28 @@ impl ExpressionRewriter for RecordFieldAggregateRewriter { ) -> Expression { if !subscripts.is_empty() && !subscripts_are_static_scalar_indices(subscripts) - && let Some(replacement) = self - .aggregate_alias_groups - .get(name.var_name()) - .and_then(|group| group.to_indexed_replacement_expr(subscripts, span)) + && let Some(group) = self.aggregate_alias_groups.get(name.var_name()) + && self.group_covers_reference_dims(name, group) + && let Some(replacement) = group.to_indexed_replacement_expr(subscripts, span) { return replacement; } if !subscripts.is_empty() { return self.walk_var_ref_expression(name, subscripts, span); } - if let Some(replacement) = self - .aggregate_alias_groups - .get(name.var_name()) - .and_then(|group| group.to_replacement_expr(span)) + if let Some(group) = self.aggregate_alias_groups.get(name.var_name()) + && self.group_covers_reference_dims(name, group) + && let Some(replacement) = group.to_replacement_expr(span) + { + return replacement; + } + if let Some(replacement) = + self.aggregate_alias_groups + .get(name.var_name()) + .and_then(|group| { + self.aggregate_dims_for_reference(name) + .and_then(|dims| group.to_partial_replacement_expr(name, &dims, span)) + }) { return replacement; } @@ -1597,6 +3583,26 @@ impl ExpressionRewriter for RecordFieldAggregateRewriter { } } +impl RecordFieldAggregateRewriter<'_> { + fn group_covers_reference_dims( + &self, + name: &Reference, + group: &AggregateAliasSubstitutionGroup, + ) -> bool { + self.aggregate_dims_for_reference(name) + .is_none_or(|dims| group.covers_dims(&dims)) + } + + fn aggregate_dims_for_reference(&self, name: &Reference) -> Option> { + let dims = self.dae_scope.as_ref()?.dims_for_reference(name).ok()??; + dims.into_iter() + .map(usize::try_from) + .collect::, _>>() + .ok() + .filter(|dims| dims.iter().all(|dim| *dim > 0)) + } +} + pub fn resolve_substitutions_in_expr( expr: &Expression, substitutions: &[Substitution], @@ -1612,29 +3618,6 @@ pub fn resolve_substitutions_in_expr( Ok(out) } -fn expr_contains_unsliced_multiscalar_ref( - expr: &Expression, - dae: &Dae, -) -> Result { - let mut refs = Vec::new(); - collect_var_ref_nodes(expr, &mut refs); - let scope = DaeVariableScope::new(dae); - for (name, subscripts) in refs { - if !subscripts.is_empty() || name.as_str() == "time" { - continue; - } - match scope.shape_for_reference(&name)? { - DaeVariableShape::Dimensions(dims) => { - if scalar_count_from_dims(name.var_name(), &dims)? > 1 { - return Ok(true); - } - } - DaeVariableShape::StructuredAggregate => {} - } - } - Ok(false) -} - pub(super) fn embedded_alias_indices_for_substitution( name: &Reference, subscripts: &[rumoca_core::Subscript], @@ -1661,6 +3644,7 @@ pub(super) fn embedded_alias_indices_for_substitution( fn index_replacement_expr( replacement: &Expression, indices: &[i64], + replacement_dims: &[i64], fallback_span: rumoca_core::Span, ) -> Result { let provenance = projection_owner_span(replacement, fallback_span)?; @@ -1668,6 +3652,9 @@ fn index_replacement_expr( if indices.is_empty() { return Ok(replacement.clone().with_span(span)); } + if replacement_dims.is_empty() && !unindexed_var_ref_replacement(replacement) { + return Ok(replacement.clone().with_span(span)); + } let extra_subscripts = indices .iter() .copied() @@ -1677,8 +3664,29 @@ fn index_replacement_expr( Expression::VarRef { name, subscripts, .. } => { + let replacement_rank = replacement_dims.len(); + if replacement_rank > 0 && replacement_rank <= subscripts.len() { + return Ok(Expression::VarRef { + name: name.clone(), + subscripts: subscripts.clone(), + span, + }); + } + if replacement_rank == 0 && !subscripts.is_empty() { + return Ok(Expression::VarRef { + name: name.clone(), + subscripts: subscripts.clone(), + span, + }); + } + let missing_rank = if replacement_rank == 0 { + extra_subscripts.len() + } else { + replacement_rank - subscripts.len() + }; + let start = extra_subscripts.len().saturating_sub(missing_rank); let mut projected_subscripts = subscripts.clone(); - projected_subscripts.extend(extra_subscripts); + projected_subscripts.extend(extra_subscripts.into_iter().skip(start)); Expression::VarRef { name: name.clone(), subscripts: projected_subscripts, @@ -1696,32 +3704,312 @@ fn index_replacement_expr( fn index_replacement_expr_with_subscripts( replacement: &Expression, subscripts: &[rumoca_core::Subscript], + replacement_dims: &[i64], fallback_span: rumoca_core::Span, ) -> Result { let span = projection_owner_span(replacement, fallback_span)?.span(); if subscripts.is_empty() { return Ok(replacement.clone().with_span(span)); } + if replacement_dims.is_empty() && !unindexed_var_ref_replacement(replacement) { + return Ok(replacement.clone().with_span(span)); + } Ok(match replacement { Expression::VarRef { name, subscripts: replacement_subscripts, .. } => { + let replacement_rank = replacement_dims.len(); + if replacement_rank > 0 && replacement_rank <= replacement_subscripts.len() { + return Ok(Expression::VarRef { + name: name.clone(), + subscripts: replacement_subscripts.clone(), + span, + }); + } + if replacement_rank == 0 && !replacement_subscripts.is_empty() { + return Ok(Expression::VarRef { + name: name.clone(), + subscripts: replacement_subscripts.clone(), + span, + }); + } + let missing_rank = if replacement_rank == 0 { + subscripts.len() + } else { + replacement_rank - replacement_subscripts.len() + }; + let start = subscripts.len().saturating_sub(missing_rank); let mut projected_subscripts = replacement_subscripts.clone(); - projected_subscripts.extend(subscripts.iter().cloned()); + projected_subscripts.extend(subscripts.iter().skip(start).cloned()); Expression::VarRef { name: name.clone(), subscripts: projected_subscripts, span, } } - _ => Expression::Index { - base: Box::new(replacement.clone()), - subscripts: subscripts.to_vec(), + _ => project_replacement_expr_with_subscripts(replacement, subscripts, span) + .unwrap_or_else(|| Expression::Index { + base: Box::new(replacement.clone()), + subscripts: subscripts.to_vec(), + span, + }), + }) +} + +fn unindexed_var_ref_replacement(replacement: &Expression) -> bool { + matches!( + replacement, + Expression::VarRef { subscripts, .. } if subscripts.is_empty() + ) +} + +fn project_replacement_expr_with_subscripts( + replacement: &Expression, + subscripts: &[rumoca_core::Subscript], + span: rumoca_core::Span, +) -> Option { + let (first, rest) = subscripts.split_first()?; + match first { + rumoca_core::Subscript::Index { value, .. } => { + project_replacement_expr_index(replacement, *value, rest, span) + } + rumoca_core::Subscript::Expr { expr, .. } => { + if let Some(value) = literal_integer_value(expr) { + project_replacement_expr_index(replacement, value, rest, span) + } else { + project_replacement_expr_symbolic_index(replacement, expr, rest, span) + } + } + rumoca_core::Subscript::Colon { .. } => None, + } +} + +pub(super) fn project_index_expr_with_exact_subscripts_in_dae( + dae: &Dae, + base: &Expression, + subscripts: &[rumoca_core::Subscript], + span: rumoca_core::Span, +) -> Option { + if subscripts.is_empty() { + return Some(base.clone().with_span(span)); + } + let mut exact_subscripts = Vec::with_capacity(subscripts.len()); + for subscript in subscripts { + let value = exact_subscript_index_in_dae(dae, subscript)?; + exact_subscripts.push(rumoca_core::Subscript::Index { value, span }); + } + project_replacement_expr_with_subscripts(base, &exact_subscripts, span) +} + +fn project_replacement_expr_index( + replacement: &Expression, + one_based_index: i64, + rest: &[rumoca_core::Subscript], + span: rumoca_core::Span, +) -> Option { + match replacement { + Expression::Array { elements, .. } => { + let zero_based = usize::try_from(one_based_index.checked_sub(1)?).ok()?; + let element = elements.get(zero_based)?.clone().with_span(span); + if rest.is_empty() { + Some(element) + } else { + project_replacement_expr_with_subscripts(&element, rest, span) + } + } + Expression::ArrayComprehension { + expr, + indices, + filter, + .. + } => { + if filter.is_some() || indices.len() != 1 { + return None; + } + let value = comprehension_index_value(&indices[0].range, one_based_index)?; + let selected = + substitute_comprehension_index_literal(expr, &indices[0].name, value, span); + if rest.is_empty() { + Some(selected) + } else { + project_replacement_expr_with_subscripts(&selected, rest, span) + } + } + Expression::Binary { op, lhs, rhs, .. } => { + let lhs_selected = project_replacement_expr_index(lhs, one_based_index, rest, span); + let rhs_selected = project_replacement_expr_index(rhs, one_based_index, rest, span); + project_binary_selection(op.clone(), lhs, rhs, lhs_selected, rhs_selected, span) + } + _ => None, + } +} + +fn project_replacement_expr_symbolic_index( + replacement: &Expression, + index: &Expression, + rest: &[rumoca_core::Subscript], + span: rumoca_core::Span, +) -> Option { + match replacement { + Expression::ArrayComprehension { + expr, + indices, + filter, + .. + } => { + if filter.is_some() || indices.len() != 1 { + return None; + } + let selected = substitute_comprehension_index_expr(expr, &indices[0].name, index, span); + if rest.is_empty() { + Some(selected) + } else { + project_replacement_expr_with_subscripts(&selected, rest, span) + } + } + Expression::Binary { op, lhs, rhs, .. } => { + let lhs_selected = project_replacement_expr_symbolic_index(lhs, index, rest, span); + let rhs_selected = project_replacement_expr_symbolic_index(rhs, index, rest, span); + project_binary_selection(op.clone(), lhs, rhs, lhs_selected, rhs_selected, span) + } + _ => None, + } +} + +fn project_binary_selection( + op: OpBinary, + lhs: &Expression, + rhs: &Expression, + lhs_selected: Option, + rhs_selected: Option, + span: rumoca_core::Span, +) -> Option { + match (lhs_selected, rhs_selected) { + (Some(lhs), Some(rhs)) => Some(Expression::Binary { + op, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span, + }), + (Some(lhs), None) => Some(Expression::Binary { + op, + lhs: Box::new(lhs), + rhs: Box::new(rhs.clone().with_span(span)), + span, + }), + (None, Some(rhs)) => Some(Expression::Binary { + op, + lhs: Box::new(lhs.clone().with_span(span)), + rhs: Box::new(rhs), + span, + }), + (None, None) => None, + } +} + +fn comprehension_index_value(range: &Expression, one_based_index: i64) -> Option { + let Expression::Range { + start, step, end, .. + } = range + else { + return None; + }; + let start = literal_integer_value(start)?; + let step = match step.as_deref() { + Some(step) => literal_integer_value(step)?, + None => 1, + }; + let value = start + (one_based_index.checked_sub(1)?) * step; + let end = literal_integer_value(end)?; + ((step > 0 && value <= end) || (step < 0 && value >= end) || (step == 0 && value == start)) + .then_some(value) +} + +fn literal_integer_value(expr: &Expression) -> Option { + match expr { + Expression::Literal { + value: rumoca_core::Literal::Integer(value), + .. + } => Some(*value), + _ => None, + } +} + +fn substitute_comprehension_index_literal( + expr: &Expression, + name: &str, + value: i64, + span: rumoca_core::Span, +) -> Expression { + substitute_comprehension_index_expr( + expr, + name, + &Expression::Literal { + value: rumoca_core::Literal::Integer(value), span, }, - }) + span, + ) +} + +fn substitute_comprehension_index_expr( + expr: &Expression, + name: &str, + replacement: &Expression, + span: rumoca_core::Span, +) -> Expression { + struct Substituter<'a> { + name: &'a str, + replacement: &'a Expression, + span: rumoca_core::Span, + } + + impl rumoca_core::ExpressionRewriter for Substituter<'_> { + fn walk_var_ref_expression( + &mut self, + name: &Reference, + subscripts: &[rumoca_core::Subscript], + span: rumoca_core::Span, + ) -> Expression { + if name.as_str() == self.name && subscripts.is_empty() { + return self.replacement.clone().with_span(self.span); + } + Expression::VarRef { + name: name.clone(), + subscripts: self.rewrite_subscripts(subscripts), + span, + } + } + + fn walk_array_comprehension_expression( + &mut self, + expr: &Expression, + indices: &[rumoca_core::ComprehensionIndex], + filter: Option<&Expression>, + span: rumoca_core::Span, + ) -> Expression { + if indices.iter().any(|index| index.name == self.name) { + return Expression::ArrayComprehension { + expr: Box::new(expr.clone()), + indices: indices.to_vec(), + filter: filter.cloned().map(Box::new), + span, + }; + } + rumoca_core::ExpressionRewriter::walk_array_comprehension_expression( + self, expr, indices, filter, span, + ) + } + } + + let mut substituter = Substituter { + name, + replacement, + span, + }; + substituter.rewrite_expression(expr) } fn projection_owner_span( @@ -1769,10 +4057,26 @@ pub(super) fn var_ref_matches_unknown_for_substitution( subscripts: &[rumoca_core::Subscript], substitution: &Substitution, ) -> bool { + var_ref_matches_unknown_for_substitution_in_scope(name, subscripts, substitution, None) +} + +pub(super) fn var_ref_matches_unknown_for_substitution_in_scope( + name: &Reference, + subscripts: &[rumoca_core::Subscript], + substitution: &Substitution, + dae_scope: Option<&DaeVariableScope<'_>>, +) -> bool { + if indexed_component_ref_mismatch_for_substitution(name, subscripts, substitution, dae_scope) { + return false; + } if name.var_name().id() == substitution.var_name.id() { return subscripts.is_empty() || subscripts_all_one(subscripts); } + if scalar_subscript_ref_matches_substitution(name, subscripts, substitution, dae_scope) { + return true; + } + if subscripts.is_empty() && substitution_indexed_base_matches(name, substitution) { return false; } @@ -1797,33 +4101,102 @@ pub(super) fn var_ref_matches_unknown_for_substitution( var_ref_matches_unknown(name, subscripts, &substitution.var_name) } +fn scalar_subscript_ref_matches_substitution( + name: &Reference, + subscripts: &[rumoca_core::Subscript], + substitution: &Substitution, + dae_scope: Option<&DaeVariableScope<'_>>, +) -> bool { + let Some((base, expected_indices)) = + scalar_substitution_target_key_in_scope(substitution, dae_scope) + else { + return false; + }; + if !references_same_base(name, &base) || subscripts.len() != expected_indices.len() { + return false; + } + subscripts + .iter() + .zip(expected_indices) + .all(|(subscript, expected)| { + exact_subscript_index(subscript).and_then(|index| usize::try_from(index).ok()) + == Some(expected) + }) +} + pub(super) fn aggregate_subscript_ref_matches_var( name: &Reference, subscripts: &[rumoca_core::Subscript], substitution: &Substitution, ) -> bool { + if indexed_component_ref_mismatch_for_substitution(name, subscripts, substitution, None) { + return false; + } !substitution.var_dims.is_empty() && !subscripts.is_empty() && name.var_name().id() == substitution.var_name.id() } +fn indexed_component_ref_mismatch_for_substitution( + name: &Reference, + _subscripts: &[rumoca_core::Subscript], + substitution: &Substitution, + dae_scope: Option<&DaeVariableScope<'_>>, +) -> bool { + if generated_scalar_reference_matches_exact_substitution_name(name, substitution) { + return false; + } + if !reference_has_scalar_indices(name) || !substitution_has_scalar_indices(substitution) { + return false; + } + if let Some((base, _)) = + dae_scope.and_then(|scope| scope.scalarized_aggregate_target(&substitution.var_name)) + { + return !references_same_base(name, &base); + } + substitution_exact_reference_name(substitution) + .is_some_and(|substitution_name| name.as_str() != substitution_name) +} + +fn substitution_exact_reference_name(substitution: &Substitution) -> Option<&str> { + substitution + .var_ref + .as_ref() + .map(Reference::as_str) + .or(Some(substitution.var_name.as_str())) +} + struct SubstituteVarRewriter<'a> { substitution: &'a Substitution, replacement: &'a Expression, replacement_dims: &'a [i64], derivative_replacement: Option<&'a Expression>, + dae_scope: Option<&'a DaeVariableScope<'a>>, + dae_context: Option<&'a Dae>, } impl FallibleExpressionRewriter for SubstituteVarRewriter<'_> { type Error = StructuralError; fn rewrite_expression(&mut self, expr: &Expression) -> Result { + let exact_target = match self.dae_context { + Some(dae) if self.requires_structured_identity() => { + expression_is_exact_structured_substitution_target(dae, expr, self.substitution) + } + _ => exact_reference_expr_name(expr).as_ref() == Some(&self.substitution.var_name), + }; + if exact_target { + return Ok(expr.span().map_or_else( + || self.replacement.clone(), + |span| replacement_with_owner_span(self.replacement, span), + )); + } match expr { Expression::BuiltinCall { function: BuiltinFunction::Der, args, .. - } if self.der_call_matches_scalar_substitution(args) => self + } if self.der_call_matches_substitution(args) => self .derivative_replacement .cloned() .map_or_else(|| self.walk_expression(expr), Ok), @@ -1849,6 +4222,27 @@ impl FallibleExpressionRewriter for SubstituteVarRewriter<'_> { .transpose()?, span: *span, }), + Expression::Index { + base, + subscripts, + span, + } => { + if !self.requires_structured_identity() + && let Expression::VarRef { + name, + subscripts: base_subscripts, + .. + } = base.as_ref() + && self.indexed_var_ref_matches_substitution(name, base_subscripts, subscripts) + { + let mut combined_subscripts = + Vec::with_capacity(base_subscripts.len() + subscripts.len()); + combined_subscripts.extend_from_slice(base_subscripts); + combined_subscripts.extend_from_slice(subscripts); + return self.rewrite_var_ref_expression(name, &combined_subscripts, *span); + } + self.walk_expression(expr) + } _ => self.walk_expression(expr), } } @@ -1859,6 +4253,9 @@ impl FallibleExpressionRewriter for SubstituteVarRewriter<'_> { subscripts: &[rumoca_core::Subscript], span: rumoca_core::Span, ) -> Result { + if self.requires_structured_identity() { + return self.walk_var_ref_expression(name, subscripts, span); + } if let Some(indices) = embedded_alias_indices_for_substitution(name, subscripts, self.substitution) { @@ -1870,12 +4267,32 @@ impl FallibleExpressionRewriter for SubstituteVarRewriter<'_> { &self.substitution.var_dims, self.replacement_dims, ); - index_replacement_expr(self.replacement, &replacement_indices, span) + index_replacement_expr( + self.replacement, + &replacement_indices, + self.replacement_dims, + span, + ) } else if aggregate_subscript_ref_matches_var(name, subscripts, self.substitution) { - index_replacement_expr_with_subscripts(self.replacement, subscripts, span) - } else if var_ref_matches_unknown_for_substitution(name, subscripts, self.substitution) { - if !subscripts.is_empty() && !self.substitution.var_dims.is_empty() { - return index_replacement_expr_with_subscripts(self.replacement, subscripts, span); + index_replacement_expr_with_subscripts( + self.replacement, + subscripts, + self.replacement_dims, + span, + ) + } else if var_ref_matches_unknown_for_substitution_in_scope( + name, + subscripts, + self.substitution, + self.dae_scope, + ) { + if !subscripts.is_empty() { + return index_replacement_expr_with_subscripts( + self.replacement, + subscripts, + self.replacement_dims, + span, + ); } Ok(replacement_with_owner_span(self.replacement, span)) } else { @@ -1884,6 +4301,37 @@ impl FallibleExpressionRewriter for SubstituteVarRewriter<'_> { } } +impl SubstituteVarRewriter<'_> { + fn requires_structured_identity(&self) -> bool { + self.dae_context + .is_some_and(|dae| substitution_requires_structured_identity(dae, self.substitution)) + } + + fn indexed_var_ref_matches_substitution( + &self, + name: &Reference, + base_subscripts: &[rumoca_core::Subscript], + subscripts: &[rumoca_core::Subscript], + ) -> bool { + let mut combined_subscripts = Vec::with_capacity(base_subscripts.len() + subscripts.len()); + combined_subscripts.extend_from_slice(base_subscripts); + combined_subscripts.extend_from_slice(subscripts); + aggregate_subscript_ref_matches_var(name, &combined_subscripts, self.substitution) + || var_ref_matches_unknown_for_substitution_in_scope( + name, + &combined_subscripts, + self.substitution, + self.dae_scope, + ) + || embedded_alias_indices_for_substitution( + name, + &combined_subscripts, + self.substitution, + ) + .is_some() + } +} + fn replacement_with_owner_span( replacement: &Expression, owner_span: rumoca_core::Span, @@ -1896,7 +4344,15 @@ fn replacement_with_owner_span( } impl SubstituteVarRewriter<'_> { - fn der_call_matches_scalar_substitution(&self, args: &[Expression]) -> bool { + fn der_call_matches_substitution(&self, args: &[Expression]) -> bool { + if self.requires_structured_identity() { + let [arg] = args else { + return false; + }; + return self.dae_context.is_some_and(|dae| { + expression_is_exact_structured_substitution_target(dae, arg, self.substitution) + }); + } der_call_matches_scalar_substitution(args, self.substitution) } } diff --git a/crates/rumoca-phase-structural/src/eliminate/orphan_unknowns.rs b/crates/rumoca-phase-structural/src/eliminate/orphan_unknowns.rs index 9cd522d2b..d99d461a4 100644 --- a/crates/rumoca-phase-structural/src/eliminate/orphan_unknowns.rs +++ b/crates/rumoca-phase-structural/src/eliminate/orphan_unknowns.rs @@ -1,13 +1,13 @@ -use rumoca_core::VarName; +use rumoca_core::{Expression, Reference, Subscript, VarName}; use rumoca_ir_dae::{Dae, expr_contains_var}; +use super::{ + collect_exact_reference_expr_names_in_dae, exact_subscript_index_in_dae, scalar_count_from_dims, +}; + pub(super) fn drop_unreferenced_continuous_unknowns(dae: &mut Dae) { - let referenced = |name: &VarName| { - dae.continuous - .equations - .iter() - .any(|eq| expr_contains_var(&eq.rhs, name)) - }; + let exact_references = collect_continuous_exact_references(dae); + let referenced = |name: &VarName| exact_reference_keeps_unknown(dae, &exact_references, name); let algebraics = dae .variables .algebraics @@ -30,6 +30,153 @@ pub(super) fn drop_unreferenced_continuous_unknowns(dae: &mut Dae) { } } +fn collect_continuous_exact_references(dae: &Dae) -> Vec { + let mut refs = Vec::new(); + for equation in &dae.continuous.equations { + if let Some(lhs) = &equation.lhs { + collect_exact_reference_expr_names_in_dae( + dae, + &Expression::VarRef { + name: lhs.clone(), + subscripts: Vec::new(), + span: equation.span, + }, + &mut refs, + ); + collect_scalarized_lhs_owners(dae, lhs, equation.scalar_count, &mut refs); + } + collect_exact_reference_expr_names_in_dae(dae, &equation.rhs, &mut refs); + } + refs.sort(); + refs.dedup(); + refs +} + +fn collect_scalarized_lhs_owners( + dae: &Dae, + lhs: &Reference, + scalar_count: usize, + out: &mut Vec, +) { + if scalar_count == 0 { + return; + } + let Some((base, selectors)) = lhs_base_and_selectors(lhs) else { + return; + }; + let Some(base_var) = crate::variable_scope::DaeVariableScope::new(dae).exact(&base) else { + return; + }; + let Some(owned) = + projected_lhs_owner_names(dae, &base, &base_var.dims, selectors, scalar_count) + else { + return; + }; + out.extend(owned); +} + +fn lhs_base_and_selectors(lhs: &Reference) -> Option<(VarName, &[Subscript])> { + let Some(component_ref) = lhs.component_ref() else { + return Some((lhs.var_name().clone(), &[])); + }; + let selectors = component_ref.parts.last()?.subs.as_slice(); + let mut base_ref = component_ref.clone(); + base_ref.parts.last_mut()?.subs.clear(); + Some((base_ref.to_var_name(), selectors)) +} + +fn projected_lhs_owner_names( + dae: &Dae, + base: &VarName, + dims: &[i64], + selectors: &[Subscript], + scalar_count: usize, +) -> Option> { + if dims.is_empty() || (!selectors.is_empty() && selectors.len() != dims.len()) { + return None; + } + let fixed_indices = if selectors.is_empty() { + vec![None; dims.len()] + } else { + selectors + .iter() + .zip(dims) + .map(|(selector, dim)| match selector { + Subscript::Colon { .. } if *dim > 0 => Some(None), + Subscript::Index { value, .. } if *value > 0 && *value <= *dim => { + usize::try_from(*value).ok().map(Some) + } + Subscript::Index { .. } => None, + Subscript::Expr { .. } => exact_subscript_index_in_dae(dae, selector) + .filter(|value| *value > 0 && *value <= *dim) + .and_then(|value| usize::try_from(value).ok()) + .map(Some), + Subscript::Colon { .. } => None, + }) + .collect::>>()? + }; + let projected_dims = fixed_indices + .iter() + .zip(dims) + .filter_map(|(index, dim)| index.is_none().then_some(*dim)) + .collect::>(); + if scalar_count_from_dims(base, &projected_dims).ok()? != scalar_count { + return None; + } + + let owned = (0..scalar_count) + .map(|flat_index| { + let projected = if projected_dims.is_empty() { + Vec::new() + } else { + rumoca_ir_dae::flat_index_to_subscripts(&projected_dims, flat_index)? + }; + let mut projected = projected.into_iter(); + let indices = fixed_indices + .iter() + .map(|index| index.or_else(|| projected.next())) + .collect::>>()?; + Some(VarName::new(rumoca_ir_dae::format_subscript_key( + base.as_str(), + &indices, + ))) + }) + .collect::>>()?; + owned + .iter() + .all(|name| { + dae.variables.algebraics.contains_key(name) || dae.variables.outputs.contains_key(name) + }) + .then_some(owned) +} + +fn exact_reference_keeps_unknown(dae: &Dae, exact_refs: &[VarName], name: &VarName) -> bool { + if exact_refs.binary_search(name).is_ok() { + return true; + } + if exact_refs.iter().any(|exact_ref| { + rumoca_core::parse_scalar_name(exact_ref.as_str()) + .is_some_and(|scalar| scalar.base == name.as_str()) + }) { + return true; + } + if continuous_unknown_is_scalar(dae, name) { + return false; + } + dae.continuous + .equations + .iter() + .any(|eq| expr_contains_var(&eq.rhs, name)) +} + +fn continuous_unknown_is_scalar(dae: &Dae, name: &VarName) -> bool { + dae.variables + .algebraics + .get(name) + .or_else(|| dae.variables.outputs.get(name)) + .is_none_or(|var| var.dims.iter().all(|dim| *dim == 1)) +} + pub(super) fn output_partition_contains_unknown(dae: &Dae, name: &VarName) -> bool { dae.variables.outputs.contains_key(name) || rumoca_ir_dae::component_base_name(name.as_str()) diff --git a/crates/rumoca-phase-structural/src/eliminate/runtime_protection.rs b/crates/rumoca-phase-structural/src/eliminate/runtime_protection.rs index 636133829..2193f11d0 100644 --- a/crates/rumoca-phase-structural/src/eliminate/runtime_protection.rs +++ b/crates/rumoca-phase-structural/src/eliminate/runtime_protection.rs @@ -4,12 +4,27 @@ use indexmap::IndexSet; use rumoca_core::{BuiltinFunction, Expression, ExpressionVisitor, OpBinary, VarName}; use rumoca_ir_dae::{self as dae, Dae}; -use super::{equation_analysis_expr, expr_contains_var}; +use super::{ + collect_var_ref_nodes, equation_analysis_expr, exact_reference_expr_name_in_dae, + expr_contains_var, +}; pub(super) fn runtime_protected_unknown_names(dae: &Dae) -> IndexSet { let mut protected = crate::runtime_defined::runtime_defined_continuous_unknown_names(dae); protected.extend(branch_local_analog_protected_unknown_names(dae)); protected.extend(clocked_value_source_protected_unknown_names(dae)); + protected.extend(pre_snapshot_source_protected_unknown_names(dae)); + protected +} + +fn pre_snapshot_source_protected_unknown_names(dae: &Dae) -> HashSet { + let mut protected = HashSet::new(); + for name in dae.variables.parameters.keys() { + let Some(source_name) = rumoca_core::pre_slot_base(name.as_str()) else { + continue; + }; + maybe_protect_branch_local_unknown(dae, &mut protected, &VarName::new(source_name)); + } protected } @@ -203,6 +218,18 @@ pub(super) fn assignment_target_name(expr: &Expression) -> Option { None } +pub(super) fn assignment_target_name_in_dae(dae: &Dae, expr: &Expression) -> Option { + let Expression::Binary { op, lhs, rhs, .. } = expr else { + return None; + }; + if !matches!(op, OpBinary::Sub) { + return None; + } + exact_reference_expr_name_in_dae(dae, lhs) + .or_else(|| exact_reference_expr_name_in_dae(dae, rhs)) + .or_else(|| assignment_target_name(expr)) +} + pub(super) fn assignment_var_ref_name( name: &VarName, subscripts: &[rumoca_core::Subscript], @@ -269,7 +296,7 @@ pub(super) fn runtime_partition_or_event_refs_var(dae: &Dae, var_name: &VarName) } pub(super) fn should_preserve_runtime_known_assignment(dae: &Dae, eq_rhs: &Expression) -> bool { - let Some(target) = assignment_target_name(eq_rhs) else { + let Some(target) = assignment_target_name_in_dae(dae, eq_rhs) else { return false; }; dae.variables.discrete_reals.contains_key(&target) @@ -285,9 +312,12 @@ pub(super) fn expr_references_any_runtime_discrete_target( return false; } - let mut refs: HashSet = HashSet::new(); - expr.collect_var_refs(&mut refs); - refs.iter().any(|name| { + let mut refs = Vec::new(); + collect_var_ref_nodes(expr, &mut refs); + refs.iter().any(|(name, _)| { + if name.is_generated() { + return false; + } let raw = name.as_str(); runtime_defined_discrete_targets.contains(raw) || dae::component_base_name(raw) @@ -296,11 +326,15 @@ pub(super) fn expr_references_any_runtime_discrete_target( } pub(super) fn expr_references_any_discrete_name(dae: &Dae, expr: &Expression) -> bool { - let mut refs: HashSet = HashSet::new(); - expr.collect_var_refs(&mut refs); - refs.iter().any(|name| { - dae.variables.discrete_reals.contains_key(name) - || dae.variables.discrete_valued.contains_key(name) + let mut refs = Vec::new(); + collect_var_ref_nodes(expr, &mut refs); + refs.iter().any(|(name, _)| { + if name.is_generated() { + return false; + } + let var_name = name.var_name(); + dae.variables.discrete_reals.contains_key(var_name) + || dae.variables.discrete_valued.contains_key(var_name) || dae::component_base_name(name.as_str()).is_some_and(|base| { let base = VarName::new(base.as_str()); dae.variables.discrete_reals.contains_key(&base) diff --git a/crates/rumoca-phase-structural/src/eliminate/scalar_shape.rs b/crates/rumoca-phase-structural/src/eliminate/scalar_shape.rs index e8b716a58..988530801 100644 --- a/crates/rumoca-phase-structural/src/eliminate/scalar_shape.rs +++ b/crates/rumoca-phase-structural/src/eliminate/scalar_shape.rs @@ -1,13 +1,20 @@ -use rumoca_core::{Expression, Reference, Subscript, VarName}; +use rumoca_core::{ComponentReference, Expression, Reference, Subscript, VarName}; use rumoca_ir_dae as dae; use crate::StructuralError; use crate::variable_scope::{DaeVariableScope, scalar_count_from_dims}; +use super::exact_reference_expr_name_in_dae; + pub(super) fn expression_is_scalar_after_subscripts( expr: &Expression, dae: &dae::Dae, ) -> Result { + if let Some(exact_name) = exact_reference_expr_name_in_dae(dae, expr) + && let Some(var) = DaeVariableScope::new(dae).exact(&exact_name) + { + return Ok(scalar_count_from_dims(&exact_name, &var.dims)? == 1); + } match expr { Expression::Literal { .. } | Expression::Empty { .. } => Ok(true), Expression::VarRef { @@ -45,20 +52,84 @@ pub(super) fn expression_is_scalar_after_subscripts( } } -fn var_ref_is_scalar_after_subscripts( +pub(super) fn var_ref_is_scalar_after_subscripts( name: &Reference, subscripts: &[Subscript], reference_span: rumoca_core::Span, dae: &dae::Dae, ) -> Result { let scope = DaeVariableScope::new(dae); - let Some(dims) = scope.dims_for_reference(name)? else { - return Ok(false); + if !subscripts.is_empty() { + let expr = Expression::VarRef { + name: name.clone(), + subscripts: subscripts.to_vec(), + span: reference_span, + }; + if let Some(exact_name) = exact_reference_expr_name_in_dae(dae, &expr) + && let Some(var) = DaeVariableScope::new(dae).exact(&exact_name) + { + return Ok(scalar_count_from_dims(&exact_name, &var.dims)? == 1); + } + } + let dims = match scope.dims_for_reference(name) { + Ok(Some(dims)) => dims, + Ok(None) => return Ok(false), + Err(StructuralError::ContractViolation { reason, .. }) + | Err(StructuralError::UnspannedContractViolation { reason }) + if reason.contains("missing DAE variable metadata") => + { + return Ok(false); + } + Err(err) => return Err(err), }; - let remaining_dims = dims_after_subscripts(name.var_name(), &dims, subscripts, reference_span)?; + let subscript_offset = if subscripts.len() > dims.len() { + component_scalar_selection_count(name).min(subscripts.len()) + } else { + 0 + }; + let remaining_dims = dims_after_subscripts( + name.var_name(), + &dims, + &subscripts[subscript_offset..], + reference_span, + )?; Ok(scalar_count_from_dims(name.var_name(), &remaining_dims)? == 1) } +fn component_scalar_selection_count(name: &Reference) -> usize { + name.component_ref() + .map(component_ref_scalar_selection_count) + .unwrap_or(0) +} + +fn component_ref_scalar_selection_count(component_ref: &ComponentReference) -> usize { + component_ref + .parts + .iter() + .flat_map(|part| &part.subs) + .filter(|subscript| positive_subscript_index(subscript).is_some()) + .count() +} + +fn positive_subscript_index(subscript: &Subscript) -> Option { + match subscript { + Subscript::Index { value, .. } if *value > 0 => Some(*value), + Subscript::Expr { expr, .. } => match expr.as_ref() { + Expression::Literal { + value: rumoca_core::Literal::Integer(value), + .. + } if *value > 0 => Some(*value), + Expression::Literal { + value: rumoca_core::Literal::Real(value), + .. + } if value.is_finite() && value.fract() == 0.0 && *value > 0.0 => Some(*value as i64), + _ => None, + }, + Subscript::Colon { .. } => None, + _ => None, + } +} + fn dims_after_subscripts( name: &VarName, dims: &[i64], @@ -103,3 +174,102 @@ fn dims_after_subscripts( remaining.extend_from_slice(&dims[subscripts.len()..]); Ok(remaining) } + +#[cfg(test)] +mod tests { + use super::*; + use rumoca_core::{ComponentRefPart, ComponentReference, Literal, Span}; + + fn index(value: i64) -> Subscript { + Subscript::Index { + value, + span: Span::DUMMY, + } + } + + fn component_ref(parts: Vec) -> ComponentReference { + ComponentReference { + local: false, + span: Span::DUMMY, + parts, + def_id: None, + } + } + + fn part(ident: &str, subs: Vec) -> ComponentRefPart { + ComponentRefPart { + ident: ident.to_string(), + span: Span::DUMMY, + subs, + } + } + + #[test] + fn component_selected_scalar_reference_consumes_outer_projection_subscripts() { + let mut dae_model = dae::Dae::default(); + dae_model.variables.algebraics.insert( + VarName::new("machine.inertia[1,1].w"), + dae::Variable { + name: VarName::new("machine.inertia[1,1].w"), + dims: Vec::new(), + component_ref: Some(component_ref(vec![ + part("machine", Vec::new()), + part("inertia", vec![index(1), index(1)]), + part("w", Vec::new()), + ])), + ..dae::Variable::empty_with_span(Span::DUMMY) + }, + ); + let reference = Reference::from_component_reference(component_ref(vec![ + part("machine", Vec::new()), + part("inertia", vec![index(1), index(1)]), + part("w", Vec::new()), + ])); + + assert!( + var_ref_is_scalar_after_subscripts( + &reference, + &[index(1), index(1)], + Span::DUMMY, + &dae_model + ) + .expect("component-selected scalar should remain scalar") + ); + } + + #[test] + fn plain_scalar_reference_still_rejects_extra_subscripts() { + let mut dae_model = dae::Dae::default(); + dae_model.variables.algebraics.insert( + VarName::new("x"), + dae::Variable { + name: VarName::new("x"), + dims: Vec::new(), + ..dae::Variable::empty_with_span(Span::DUMMY) + }, + ); + let reference = Reference::from_var_name(VarName::new("x")); + + assert!(matches!( + var_ref_is_scalar_after_subscripts(&reference, &[index(1)], Span::DUMMY, &dae_model), + Err(StructuralError::ContractViolation { reason, .. }) + if reason.contains("indexed DAE reference `x` has 1 subscripts for dimensions []") + )); + } + + #[test] + fn real_literal_component_subscript_counts_as_scalar_selection() { + let reference = Reference::from_component_reference(component_ref(vec![part( + "x", + vec![Subscript::Expr { + expr: Box::new(Expression::Literal { + value: Literal::Real(1.0), + span: Span::DUMMY, + }), + span: Span::DUMMY, + }], + )])); + + assert_eq!(component_scalar_selection_count(&reference), 1); + } +} diff --git a/crates/rumoca-phase-structural/src/eliminate/solve_for_unknown.rs b/crates/rumoca-phase-structural/src/eliminate/solve_for_unknown.rs index 8b95be130..1193ac937 100644 --- a/crates/rumoca-phase-structural/src/eliminate/solve_for_unknown.rs +++ b/crates/rumoca-phase-structural/src/eliminate/solve_for_unknown.rs @@ -9,9 +9,33 @@ use super::*; /// - `0 = -(z - expr)` -> `z = expr` /// - `0 = -(expr - z)` -> `z = expr` pub fn try_solve_for_unknown(rhs: &Expression, unknown: &VarName) -> Option { + try_solve_for_unknown_with_context(rhs, unknown, None) +} + +pub(super) fn try_solve_for_unknown_in_dae( + dae: &Dae, + rhs: &Expression, + unknown: &VarName, +) -> Option { + try_solve_for_unknown_with_context(rhs, unknown, Some(dae)) +} + +pub(super) fn expr_contains_unknown_in_dae( + dae: &Dae, + expr: &Expression, + unknown: &VarName, +) -> bool { + expr_contains_unknown(expr, unknown, Some(dae)) +} + +fn try_solve_for_unknown_with_context( + rhs: &Expression, + unknown: &VarName, + dae: Option<&Dae>, +) -> Option { match rhs { // Pattern: 0 = z -> z = 0 - Expression::VarRef { span, .. } if is_symbolic_solve_target(rhs, unknown) => { + Expression::VarRef { span, .. } if is_symbolic_solve_target(rhs, unknown, dae) => { Some(Expression::Literal { value: rumoca_core::Literal::Real(0.0), span: *span, @@ -25,25 +49,29 @@ pub fn try_solve_for_unknown(rhs: &Expression, unknown: &VarName) -> Option { // 0 = z - expr -> z = expr - if is_symbolic_solve_target(lhs, unknown) && !expr_contains_var(rhs_inner, unknown) { + if is_symbolic_solve_target(lhs, unknown, dae) + && !expr_contains_unknown(rhs_inner, unknown, dae) + { return Some(*rhs_inner.clone()); } // 0 = expr - z -> z = expr - if is_symbolic_solve_target(rhs_inner, unknown) && !expr_contains_var(lhs, unknown) { + if is_symbolic_solve_target(rhs_inner, unknown, dae) + && !expr_contains_unknown(lhs, unknown, dae) + { return Some(*lhs.clone()); } - solve_unit_affine_residual(rhs, unknown, *span) + solve_unit_affine_residual(rhs, unknown, *span, dae) } Expression::Binary { op: OpBinary::Add, span, .. - } => solve_unit_affine_residual(rhs, unknown, *span), + } => solve_unit_affine_residual(rhs, unknown, *span, dae), Expression::Unary { op: OpUnary::Plus, rhs: inner, .. - } => try_solve_for_unknown(inner, unknown), + } => try_solve_for_unknown_with_context(inner, unknown, dae), Expression::Unary { op: OpUnary::Minus, rhs: inner, @@ -51,18 +79,47 @@ pub fn try_solve_for_unknown(rhs: &Expression, unknown: &VarName) -> Option { // Recurse into the negated expression. // -(z - expr) has the same solutions as (z - expr). - try_solve_for_unknown(inner, unknown) + try_solve_for_unknown_with_context(inner, unknown, dae) } + Expression::If { + branches, + else_branch, + span, + } => solve_if_residual(branches, else_branch, unknown, *span), _ => None, } } +fn solve_if_residual( + branches: &[(Expression, Expression)], + else_branch: &Expression, + unknown: &VarName, + span: rumoca_core::Span, +) -> Option { + let mut solved_branches = Vec::with_capacity(branches.len()); + for (condition, branch_residual) in branches { + if expr_contains_var(condition, unknown) { + return None; + } + solved_branches.push(( + condition.clone(), + try_solve_for_unknown(branch_residual, unknown)?, + )); + } + Some(Expression::If { + branches: solved_branches, + else_branch: Box::new(try_solve_for_unknown(else_branch, unknown)?), + span, + }) +} + fn solve_unit_affine_residual( rhs: &Expression, unknown: &VarName, residual_span: rumoca_core::Span, + dae: Option<&Dae>, ) -> Option { - let (coef, remainder) = split_unit_affine_residual(rhs, unknown, residual_span)?; + let (coef, remainder) = split_unit_affine_residual(rhs, unknown, residual_span, dae)?; match coef { 1 => Some(negate_expr(remainder, residual_span)), -1 => Some(remainder), @@ -74,8 +131,9 @@ fn split_unit_affine_residual( expr: &Expression, unknown: &VarName, span: rumoca_core::Span, + dae: Option<&Dae>, ) -> Option<(i32, Expression)> { - if is_symbolic_solve_target(expr, unknown) { + if is_symbolic_solve_target(expr, unknown, dae) { return Some(( 1, Expression::Literal { @@ -89,26 +147,26 @@ fn split_unit_affine_residual( }; match op { OpBinary::Add | OpBinary::AddElem => { - if let Some((coef, rem)) = split_unit_affine_residual(lhs, unknown, span) - && !expr_contains_var(rhs, unknown) + if let Some((coef, rem)) = split_unit_affine_residual(lhs, unknown, span, dae) + && !expr_contains_unknown(rhs, unknown, dae) { return Some((coef, add_expr(rem, *rhs.clone(), span))); } - if let Some((coef, rem)) = split_unit_affine_residual(rhs, unknown, span) - && !expr_contains_var(lhs, unknown) + if let Some((coef, rem)) = split_unit_affine_residual(rhs, unknown, span, dae) + && !expr_contains_unknown(lhs, unknown, dae) { return Some((coef, add_expr(*lhs.clone(), rem, span))); } None } OpBinary::Sub | OpBinary::SubElem => { - if let Some((coef, rem)) = split_unit_affine_residual(lhs, unknown, span) - && !expr_contains_var(rhs, unknown) + if let Some((coef, rem)) = split_unit_affine_residual(lhs, unknown, span, dae) + && !expr_contains_unknown(rhs, unknown, dae) { return Some((coef, sub_expr(rem, *rhs.clone(), span))); } - if let Some((coef, rem)) = split_unit_affine_residual(rhs, unknown, span) - && !expr_contains_var(lhs, unknown) + if let Some((coef, rem)) = split_unit_affine_residual(rhs, unknown, span, dae) + && !expr_contains_unknown(lhs, unknown, dae) { return Some((-coef, sub_expr(*lhs.clone(), rem, span))); } @@ -177,7 +235,43 @@ fn is_zero_literal(expr: &Expression) -> bool { ) } -fn is_symbolic_solve_target(expr: &Expression, unknown: &VarName) -> bool { +fn expr_contains_unknown(expr: &Expression, unknown: &VarName, dae: Option<&Dae>) -> bool { + if let Some(dae) = dae + && dae_unknown_is_exact_scalar(dae, unknown) + { + let mut exact_names = Vec::new(); + collect_exact_reference_expr_names_in_dae(dae, expr, &mut exact_names); + return exact_names.iter().any(|name| name == unknown); + } + if expr_contains_var(expr, unknown) { + return true; + } + let Some(dae) = dae else { + return false; + }; + let mut exact_names = Vec::new(); + collect_exact_reference_expr_names_in_dae(dae, expr, &mut exact_names); + exact_names.iter().any(|name| name == unknown) +} + +fn dae_unknown_is_exact_scalar(dae: &Dae, unknown: &VarName) -> bool { + dae.variables + .states + .get(unknown) + .or_else(|| dae.variables.algebraics.get(unknown)) + .or_else(|| dae.variables.outputs.get(unknown)) + .is_some_and(|var| var.dims.iter().all(|dim| *dim == 1)) +} + +fn is_symbolic_solve_target(expr: &Expression, unknown: &VarName, dae: Option<&Dae>) -> bool { + let exact_name = dae.and_then(|dae| exact_reference_expr_name_in_dae(dae, expr)); + if exact_name + .or_else(|| exact_reference_expr_name(expr)) + .as_ref() + == Some(unknown) + { + return true; + } match expr { Expression::VarRef { name, subscripts, .. diff --git a/crates/rumoca-phase-structural/src/eliminate/substitution_application.rs b/crates/rumoca-phase-structural/src/eliminate/substitution_application.rs index 1bb900bef..f59766292 100644 --- a/crates/rumoca-phase-structural/src/eliminate/substitution_application.rs +++ b/crates/rumoca-phase-structural/src/eliminate/substitution_application.rs @@ -1,10 +1,12 @@ use std::collections::HashMap; use rumoca_core::{Literal, OpUnary}; +use rumoca_ir_dae::{DaeEquationPartition, TryDaeExpressionRewriter}; use super::{ - Dae, Expression, OpBinary, Substitution, VarName, apply_substitutions_to_expr, - apply_substitutions_to_expr_with_derivatives, + Dae, Expression, OpBinary, Substitution, VarName, + apply_substitutions_to_expr_with_derivatives_and_dae, exact_subscript_index_in_dae, + project_index_expr_with_exact_subscripts_in_dae, }; use crate::StructuralError; @@ -23,6 +25,22 @@ pub(super) fn equation_analysis_expr(eq: &rumoca_ir_dae::Equation) -> Expression ) } +pub(super) fn canonicalize_exact_indexing_in_continuous_equations( + dae: &mut Dae, +) -> Result<(), StructuralError> { + let dae_context = dae.clone(); + let mut touched_equations = Vec::new(); + for (index, equation) in dae.continuous.equations.iter_mut().enumerate() { + let original_rhs = equation.rhs.clone(); + equation.rhs = simplify_after_substitution(original_rhs.clone(), &dae_context); + if equation.rhs != original_rhs { + touched_equations.push(index); + } + } + super::drop_structured_families_touching_equations(dae, &touched_equations); + Ok(()) +} + pub(super) fn apply_substitutions_to_remaining_once( dae: &mut Dae, eliminated_eq_flags: &[bool], @@ -49,9 +67,10 @@ pub(super) fn apply_substitutions_to_remaining_once( } let original_lhs = eq.lhs.clone(); let original_rhs = eq.rhs.clone(); - let rhs = apply_substitutions_in_order_with_derivatives( + let rhs = apply_substitutions_in_order_with_derivatives_and_dae( &eq.rhs, substitutions, + &derivative_source, &mut derivative_replacements, )?; let Some(lhs) = eq.lhs.as_ref() else { @@ -66,9 +85,10 @@ pub(super) fn apply_substitutions_to_remaining_once( subscripts: Vec::new(), span: eq.span, }; - let substituted_lhs = apply_substitutions_in_order_with_derivatives( + let substituted_lhs = apply_substitutions_in_order_with_derivatives_and_dae( &lhs_expr, substitutions, + &derivative_source, &mut derivative_replacements, )?; if substituted_lhs == lhs_expr { @@ -96,27 +116,26 @@ pub(super) fn apply_substitutions_to_dae_partitions( let mut rewriter = SubstitutionDaeRewriter { substitutions, derivative_replacements: DerivativeReplacementCache::new(&derivative_source), + touched_continuous_equations: Vec::new(), }; - rewriter.rewrite_dae(dae) -} - -pub(super) fn apply_substitutions_in_order( - expr: &Expression, - substitutions: &[Substitution], -) -> Result { - let substituted = apply_substitutions_to_expr(expr, substitutions)?; - Ok(simplify_arithmetic_identities(substituted)) + rewriter.try_rewrite_dae(dae)?; + super::drop_structured_families_touching_equations(dae, &rewriter.touched_continuous_equations); + Ok(()) } -fn apply_substitutions_in_order_with_derivatives( +fn apply_substitutions_in_order_with_derivatives_and_dae( expr: &Expression, substitutions: &[Substitution], + dae_context: &Dae, derivative_replacements: &mut DerivativeReplacementCache<'_>, ) -> Result { - let substituted = apply_substitutions_to_expr_with_derivatives(expr, substitutions, |sub| { - derivative_replacements.replacement_for(sub) - })?; - Ok(simplify_arithmetic_identities(substituted)) + let substituted = apply_substitutions_to_expr_with_derivatives_and_dae( + expr, + substitutions, + Some(dae_context), + |sub| derivative_replacements.replacement_for(sub), + )?; + Ok(simplify_after_substitution(substituted, dae_context)) } /// Fold exact arithmetic identities introduced by substitution. @@ -129,11 +148,11 @@ fn apply_substitutions_in_order_with_derivatives( /// /// Only handles identities that are mathematically exact across all numeric /// types — division-by-zero and `0^0` are intentionally not folded. -fn simplify_arithmetic_identities(expr: Expression) -> Expression { +fn simplify_after_substitution(expr: Expression, dae_context: &Dae) -> Expression { match expr { Expression::Binary { op, lhs, rhs, span } => { - let lhs = simplify_arithmetic_identities(*lhs); - let rhs = simplify_arithmetic_identities(*rhs); + let lhs = simplify_after_substitution(*lhs, dae_context); + let rhs = simplify_after_substitution(*rhs, dae_context); match op { OpBinary::Add => { if is_numeric_zero(&lhs) { @@ -175,7 +194,7 @@ fn simplify_arithmetic_identities(expr: Expression) -> Expression { } } Expression::Unary { op, rhs, span } => { - let inner = simplify_arithmetic_identities(*rhs); + let inner = simplify_after_substitution(*rhs, dae_context); if matches!(op, OpUnary::Minus) { // -(-x) → x if let Expression::Unary { @@ -197,10 +216,60 @@ fn simplify_arithmetic_identities(expr: Expression) -> Expression { span, } } + Expression::Index { + base, + subscripts, + span, + } => { + let base = simplify_after_substitution(*base, dae_context); + let subscripts = subscripts + .into_iter() + .map(|subscript| simplify_subscript_after_substitution(subscript, dae_context)) + .collect::>(); + project_index_expr_with_exact_subscripts_in_dae(dae_context, &base, &subscripts, span) + .unwrap_or_else(|| Expression::Index { + base: Box::new(base), + subscripts, + span, + }) + } + Expression::VarRef { + name, + subscripts, + span, + } => Expression::VarRef { + name, + subscripts: subscripts + .into_iter() + .map(|subscript| simplify_subscript_after_substitution(subscript, dae_context)) + .collect(), + span, + }, _ => expr, } } +fn simplify_subscript_after_substitution( + subscript: rumoca_core::Subscript, + dae_context: &Dae, +) -> rumoca_core::Subscript { + match subscript { + rumoca_core::Subscript::Expr { expr, span } => { + let expr = simplify_after_substitution(*expr, dae_context); + let subscript = rumoca_core::Subscript::Expr { + expr: Box::new(expr), + span, + }; + if let Some(value) = exact_subscript_index_in_dae(dae_context, &subscript) { + rumoca_core::Subscript::Index { value, span } + } else { + subscript + } + } + other => other, + } +} + fn is_numeric_zero(expr: &Expression) -> bool { matches!( expr, @@ -253,9 +322,10 @@ fn apply_substitutions_to_equation( substitutions: &[Substitution], derivative_replacements: &mut DerivativeReplacementCache<'_>, ) -> Result<(), StructuralError> { - let rhs = apply_substitutions_in_order_with_derivatives( + let rhs = apply_substitutions_in_order_with_derivatives_and_dae( &eq.rhs, substitutions, + derivative_replacements.dae, derivative_replacements, )?; let Some(lhs) = eq.lhs.as_ref() else { @@ -267,9 +337,10 @@ fn apply_substitutions_to_equation( subscripts: Vec::new(), span: eq.span, }; - let substituted_lhs = apply_substitutions_in_order_with_derivatives( + let substituted_lhs = apply_substitutions_in_order_with_derivatives_and_dae( &lhs_expr, substitutions, + derivative_replacements.dae, derivative_replacements, )?; if substituted_lhs == lhs_expr { @@ -284,53 +355,33 @@ fn apply_substitutions_to_equation( struct SubstitutionDaeRewriter<'a> { substitutions: &'a [Substitution], derivative_replacements: DerivativeReplacementCache<'a>, + touched_continuous_equations: Vec, } -impl SubstitutionDaeRewriter<'_> { - fn rewrite_dae(&mut self, dae: &mut Dae) -> Result<(), StructuralError> { - let touched_equations = - self.rewrite_equations_collecting_changes(&mut dae.continuous.equations)?; - super::drop_structured_families_touching_equations(dae, &touched_equations); - self.rewrite_equations(&mut dae.initialization.equations)?; - self.rewrite_equations(&mut dae.discrete.real_updates)?; - self.rewrite_equations(&mut dae.discrete.valued_updates)?; - self.rewrite_equations(&mut dae.conditions.equations)?; - self.rewrite_expression_slots(&mut dae.conditions.relations)?; - self.rewrite_expression_slots(&mut dae.events.synthetic_root_conditions)?; - self.rewrite_event_actions(&mut dae.events.event_actions)?; - self.rewrite_expression_slots(&mut dae.clocks.constructor_exprs)?; - self.rewrite_expression_slots(&mut dae.clocks.triggered_conditions)?; - Ok(()) - } +impl TryDaeExpressionRewriter for SubstitutionDaeRewriter<'_> { + type Error = StructuralError; - fn rewrite_equations( + fn try_rewrite_equations( &mut self, + partition: DaeEquationPartition, equations: &mut [rumoca_ir_dae::Equation], ) -> Result<(), StructuralError> { - for equation in equations { - self.rewrite_equation(equation)?; - } - Ok(()) - } - - fn rewrite_equations_collecting_changes( - &mut self, - equations: &mut [rumoca_ir_dae::Equation], - ) -> Result, StructuralError> { - let mut touched = Vec::new(); for (index, equation) in equations.iter_mut().enumerate() { let original_lhs = equation.lhs.clone(); let original_rhs = equation.rhs.clone(); - self.rewrite_equation(equation)?; - if equation.lhs != original_lhs || equation.rhs != original_rhs { - touched.push(index); + self.try_rewrite_equation(partition, equation)?; + if partition == DaeEquationPartition::Continuous + && (equation.lhs != original_lhs || equation.rhs != original_rhs) + { + self.touched_continuous_equations.push(index); } } - Ok(touched) + Ok(()) } - fn rewrite_equation( + fn try_rewrite_equation( &mut self, + _partition: DaeEquationPartition, equation: &mut rumoca_ir_dae::Equation, ) -> Result<(), StructuralError> { apply_substitutions_to_equation( @@ -340,33 +391,14 @@ impl SubstitutionDaeRewriter<'_> { ) } - fn rewrite_expression(&mut self, expr: &Expression) -> Result { - apply_substitutions_in_order_with_derivatives( + fn try_rewrite_expression(&mut self, expr: &Expression) -> Result { + apply_substitutions_in_order_with_derivatives_and_dae( expr, self.substitutions, + self.derivative_replacements.dae, &mut self.derivative_replacements, ) } - - fn rewrite_expression_slots( - &mut self, - expressions: &mut [Expression], - ) -> Result<(), StructuralError> { - for expression in expressions { - *expression = self.rewrite_expression(expression)?; - } - Ok(()) - } - - fn rewrite_event_actions( - &mut self, - actions: &mut [rumoca_ir_dae::DaeEventAction], - ) -> Result<(), StructuralError> { - for action in actions { - action.condition = self.rewrite_expression(&action.condition)?; - } - Ok(()) - } } struct DerivativeReplacementCache<'a> { diff --git a/crates/rumoca-phase-structural/src/eliminate/substitution_target.rs b/crates/rumoca-phase-structural/src/eliminate/substitution_target.rs index f223f11e1..1f7fd1149 100644 --- a/crates/rumoca-phase-structural/src/eliminate/substitution_target.rs +++ b/crates/rumoca-phase-structural/src/eliminate/substitution_target.rs @@ -1,16 +1,202 @@ -use rumoca_core::{BuiltinFunction, Expression, ExpressionVisitor, Reference}; +use rumoca_core::{ + BuiltinFunction, ComponentRefPart, ComponentReference, Expression, ExpressionVisitor, + Reference, Subscript, component_reference_from_flat_name, +}; use super::{ - Substitution, aggregate_subscript_ref_matches_var, der_call_matches_scalar_substitution, - embedded_alias_indices_for_substitution, var_ref_matches_unknown_for_substitution, + Dae, Substitution, aggregate_subscript_ref_matches_var, der_call_matches_scalar_substitution, + embedded_alias_indices_for_substitution, var_ref_matches_unknown_for_substitution_in_scope, }; +use crate::variable_scope::DaeVariableScope; -pub(super) fn expr_contains_substitution_target( +pub(super) fn expression_is_exact_structured_substitution_target( + dae: &Dae, expr: &Expression, substitution: &Substitution, +) -> bool { + let Some(target) = DaeVariableScope::new(dae).exact(&substitution.var_name) else { + return false; + }; + let Some(target_ref) = target.component_ref.as_ref() else { + return false; + }; + let Some(reference) = structured_reference_from_expr(expr) else { + return false; + }; + if !component_reference_path_matches(&reference.component_ref, target_ref) + || unique_structured_target_name(dae, target_ref).as_ref() != Some(&substitution.var_name) + { + return false; + } + !reference.terminal_identity_complete + || match (reference.component_ref.def_id, target_ref.def_id) { + (Some(expr_def_id), Some(target_def_id)) => expr_def_id == target_def_id, + (None, _) => true, + (Some(_), None) => false, + } +} + +pub(super) fn substitution_requires_structured_identity( + dae: &Dae, + substitution: &Substitution, +) -> bool { + substitution.var_dims.is_empty() + && DaeVariableScope::new(dae) + .exact(&substitution.var_name) + .is_some_and(|var| var.component_ref.is_some()) +} + +pub(super) fn generated_scalar_reference_matches_exact_substitution_name( + name: &Reference, + substitution: &Substitution, +) -> bool { + if substitution.var_ref.is_some() || name.var_name().id() != substitution.var_name.id() { + return false; + } + let Some(actual) = name.component_ref() else { + return false; + }; + if actual.local || actual.def_id.is_some() { + return false; + } + component_reference_from_flat_name(&substitution.var_name, actual.span) + .is_some_and(|canonical| component_reference_path_matches(actual, &canonical)) +} + +struct StructuredExpressionReference { + component_ref: ComponentReference, + terminal_identity_complete: bool, +} + +fn structured_reference_from_expr(expr: &Expression) -> Option { + match expr { + Expression::VarRef { + name, + subscripts, + span, + } => { + let mut component_ref = name + .component_ref() + .cloned() + .or_else(|| component_reference_from_flat_name(name.var_name(), *span))?; + append_exact_subscripts(&mut component_ref, subscripts)?; + Some(StructuredExpressionReference { + component_ref, + terminal_identity_complete: subscripts.is_empty(), + }) + } + Expression::Index { + base, subscripts, .. + } => { + let mut reference = structured_reference_from_expr(base)?; + append_exact_subscripts(&mut reference.component_ref, subscripts)?; + reference.terminal_identity_complete = false; + Some(reference) + } + Expression::FieldAccess { base, field, span } => { + let mut reference = structured_reference_from_expr(base)?; + reference.component_ref.parts.push(ComponentRefPart { + ident: field.clone(), + span: *span, + subs: Vec::new(), + }); + reference.terminal_identity_complete = false; + Some(reference) + } + _ => None, + } +} + +fn append_exact_subscripts( + component_ref: &mut ComponentReference, + subscripts: &[Subscript], +) -> Option<()> { + let part = component_ref.parts.last_mut()?; + for subscript in subscripts { + let Subscript::Index { value, span } = subscript else { + return None; + }; + part.subs.push(Subscript::Index { + value: *value, + span: *span, + }); + } + Some(()) +} + +fn component_reference_path_matches( + expression: &ComponentReference, + target: &ComponentReference, +) -> bool { + expression.local == target.local + && expression.parts.len() == target.parts.len() + && expression + .parts + .iter() + .zip(&target.parts) + .all(|(expression, target)| component_ref_part_matches(expression, target)) +} + +fn component_ref_part_matches(expression: &ComponentRefPart, target: &ComponentRefPart) -> bool { + if expression.ident != target.ident { + return false; + } + let Some(expression_subscripts) = exact_subscript_values(&expression.subs) else { + return false; + }; + let Some(target_subscripts) = exact_subscript_values(&target.subs) else { + return false; + }; + expression_subscripts == target_subscripts +} + +fn exact_subscript_values(subscripts: &[Subscript]) -> Option> { + subscripts + .iter() + .map(|subscript| match subscript { + Subscript::Index { value, .. } => Some(*value), + Subscript::Colon { .. } | Subscript::Expr { .. } => None, + }) + .collect() +} + +fn unique_structured_target_name( + dae: &Dae, + target: &ComponentReference, +) -> Option { + let variables = &dae.variables; + let mut matches = variables + .states + .values() + .chain(variables.algebraics.values()) + .chain(variables.inputs.values()) + .chain(variables.outputs.values()) + .chain(variables.parameters.values()) + .chain(variables.constants.values()) + .chain(variables.discrete_reals.values()) + .chain(variables.discrete_valued.values()) + .filter(|variable| { + variable + .component_ref + .as_ref() + .is_some_and(|candidate| component_reference_path_matches(candidate, target)) + }); + let name = matches.next()?.name.clone(); + matches.next().is_none().then_some(name) +} + +pub(super) fn expr_contains_substitution_target_in_scope( + expr: &Expression, + substitution: &Substitution, + dae_scope: Option<&DaeVariableScope<'_>>, + dae_context: Option<&Dae>, ) -> bool { let mut checker = SubstitutionTargetChecker { substitution, + dae_scope, + dae_context, + structured_identity_required: dae_context + .is_some_and(|dae| substitution_requires_structured_identity(dae, substitution)), found: false, }; checker.visit_expression(expr); @@ -19,19 +205,40 @@ pub(super) fn expr_contains_substitution_target( struct SubstitutionTargetChecker<'a> { substitution: &'a Substitution, + dae_scope: Option<&'a DaeVariableScope<'a>>, + dae_context: Option<&'a Dae>, + structured_identity_required: bool, found: bool, } impl ExpressionVisitor for SubstitutionTargetChecker<'_> { fn visit_expression(&mut self, expr: &Expression) { - if !self.found { - self.walk_expression(expr); + if self.found { + return; } + if self.dae_context.is_some_and(|dae| { + expression_is_exact_structured_substitution_target(dae, expr, self.substitution) + }) { + self.found = true; + return; + } + self.walk_expression(expr); } fn visit_var_ref(&mut self, name: &Reference, subscripts: &[rumoca_core::Subscript]) { + if self.structured_identity_required { + for subscript in subscripts { + self.visit_subscript(subscript); + } + return; + } if aggregate_subscript_ref_matches_var(name, subscripts, self.substitution) - || var_ref_matches_unknown_for_substitution(name, subscripts, self.substitution) + || var_ref_matches_unknown_for_substitution_in_scope( + name, + subscripts, + self.substitution, + self.dae_scope, + ) || embedded_alias_indices_for_substitution(name, subscripts, self.substitution) .is_some() { diff --git a/crates/rumoca-phase-structural/src/eliminate/tearing_elimination.rs b/crates/rumoca-phase-structural/src/eliminate/tearing_elimination.rs index ed4ea3e55..f4b552eb6 100644 --- a/crates/rumoca-phase-structural/src/eliminate/tearing_elimination.rs +++ b/crates/rumoca-phase-structural/src/eliminate/tearing_elimination.rs @@ -6,9 +6,11 @@ use crate::{EquationRef, StructuralError, UnknownId, tear_algebraic_loop}; use super::{ Dae, DerivativeNameMatcher, Substitution, VarName, algebraic_or_output_unknown, - apply_substitutions_in_order, can_eliminate_scalar_unknown, can_use_equation_for_elimination, - equation_has_state_derivative, expr_contains_var, stable_solution_for_unknown, - substitution_for_var, + apply_substitutions_for_symbolic_candidate, can_eliminate_scalar_unknown, + can_use_equation_for_elimination, equation_has_state_derivative, expr_contains_var, + is_trivial_alias_in_dae, scalar_blt_solution_would_break_aggregate_element, + should_preserve_runtime_sensitive_continuous_assignment, + solution_is_cheap_for_symbolic_substitution, stable_solution_for_unknown, substitution_for_var, }; #[allow(clippy::too_many_arguments)] @@ -50,13 +52,28 @@ pub(super) fn tear_and_eliminate_loop_block( return Ok(()); } let var_name = var_names[local_var].clone(); - let eq_rhs = apply_substitutions_in_order( + let Some(eq_rhs) = apply_substitutions_for_symbolic_candidate( &dae.continuous.equations[eq_idx].rhs, &trial_substitutions, - )?; + dae, + )? + else { + return Ok(()); + }; + if should_preserve_runtime_sensitive_continuous_assignment(dae, &eq_rhs) { + return Ok(()); + } let Some(solution) = stable_solution_for_unknown(dae, &eq_rhs, &var_name)? else { return Ok(()); }; + if !is_trivial_alias_in_dae(dae, &solution) + && !solution_is_cheap_for_symbolic_substitution(&solution) + { + return Ok(()); + } + if scalar_blt_solution_would_break_aggregate_element(dae, &var_name, &solution)? { + return Ok(()); + } trial_substitutions.push(substitution_for_var( dae, var_name.clone(), @@ -85,7 +102,7 @@ fn loop_elimination_unknowns( return Ok(None); }; let var_name = raw_var_name.clone(); - if !can_eliminate_scalar_unknown(dae, &var_name, runtime_protected_unknowns)? { + if !can_eliminate_scalar_unknown(dae, &var_name, runtime_protected_unknowns, false)? { return Ok(None); } var_names.push(var_name); @@ -115,8 +132,14 @@ fn loop_local_incidence( eq_indices .iter() .map(|&eq_idx| { - let rhs = - apply_substitutions_in_order(&dae.continuous.equations[eq_idx].rhs, substitutions)?; + let Some(rhs) = apply_substitutions_for_symbolic_candidate( + &dae.continuous.equations[eq_idx].rhs, + substitutions, + dae, + )? + else { + return Ok(HashSet::new()); + }; Ok(var_names .iter() .enumerate() diff --git a/crates/rumoca-phase-structural/src/eliminate/tests.rs b/crates/rumoca-phase-structural/src/eliminate/tests.rs index 55f81d780..8ebcce84e 100644 --- a/crates/rumoca-phase-structural/src/eliminate/tests.rs +++ b/crates/rumoca-phase-structural/src/eliminate/tests.rs @@ -1,3 +1,6 @@ +// SPEC_0021 file-size exception: elimination regression tests still share DAE +// builders and substitution fixtures. split plan: split array-boundary, +// complex-field, and scalar-alias replacement cases into focused modules. use super::*; use rumoca_core::Span; use rumoca_ir_dae as dae; @@ -42,6 +45,25 @@ fn test_dae_variable(path: &str) -> dae::Variable { var } +#[test] +fn pre_snapshot_source_survives_boundary_alias_elimination() { + let mut dae = dae::Dae::default(); + dae.variables + .algebraics + .insert("sample.u".into(), test_dae_variable("sample.u")); + dae.variables.parameters.insert( + rumoca_core::pre_slot_name("sample.u"), + test_dae_variable("__pre__.sample.u"), + ); + + let protected = runtime_protected_unknown_names(&dae); + + assert!( + protected.contains("sample.u"), + "pre-lowering must not hide an inferred-clock sample source from structural protection" + ); +} + fn substitute_var(expr: &Expression, var: &VarName, replacement: &Expression) -> Expression { let substitution = test_substitution(var.as_str(), replacement.clone()); structural_ok( @@ -50,6 +72,8 @@ fn substitute_var(expr: &Expression, var: &VarName, replacement: &Expression) -> replacement, replacement_dims: &substitution.replacement_dims, derivative_replacement: None, + dae_scope: None, + dae_context: None, } .rewrite_expression(expr), ) @@ -83,6 +107,17 @@ fn var_ref_idx(name: &str, idx: i64) -> Expression { } } +fn var_ref_with_subscript_expr(name: &str, expr: Expression) -> Expression { + Expression::VarRef { + name: reference(name), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(expr), + span: test_span(), + }], + span: test_span(), + } +} + fn real(value: f64) -> Expression { Expression::Literal { value: Literal::Real(value), @@ -106,6 +141,14 @@ fn der(expr: Expression) -> Expression { } } +fn builtin(function: BuiltinFunction, args: Vec) -> Expression { + Expression::BuiltinCall { + function, + args, + span: test_span(), + } +} + fn binary(op: OpBinary, lhs: Expression, rhs: Expression) -> Expression { Expression::Binary { op, @@ -453,6 +496,35 @@ fn test_try_solve_negated() { ); } +#[test] +fn test_try_solve_if_residual_solves_each_branch() { + let rhs = Expression::If { + branches: vec![( + var_ref("c"), + Expression::Binary { + op: sub_op(), + lhs: Box::new(var_ref("z")), + rhs: Box::new(var_ref("u")), + span: rumoca_core::Span::DUMMY, + }, + )], + else_branch: Box::new(Expression::Binary { + op: sub_op(), + lhs: Box::new(var_ref("z")), + rhs: Box::new(lit(0.0)), + span: rumoca_core::Span::DUMMY, + }), + span: rumoca_core::Span::DUMMY, + }; + + let result = try_solve_for_unknown(&rhs, &VarName::new("z")); + + assert!( + matches!(&result, Some(Expression::If { branches, .. }) if branches.len() == 1), + "expected branch-wise If solution, got {result:?}" + ); +} + #[test] fn test_try_solve_sub_lhs_with_unity_subscript_alias_matches() { let rhs = Expression::Binary { @@ -631,6 +703,288 @@ fn test_substitute_var_projects_embedded_array_alias_component() { ); } +#[test] +fn substitution_with_dae_context_rewrites_indexed_component_field_reference() { + let mut dae = dae::Dae::default(); + let target = "stack.cell[1,1].cell.ocv_soc.u"; + let mut target_var = component_var(target); + target_var.component_ref.as_mut().unwrap().def_id = Some(rumoca_core::DefId(42)); + dae.variables + .outputs + .insert(VarName::new(target), target_var); + let indexed_cell = Expression::Index { + base: Box::new(var_ref("stack.cell")), + subscripts: vec![ + rumoca_core::Subscript::generated_index(1, test_span()), + rumoca_core::Subscript::generated_index(1, test_span()), + ], + span: test_span(), + }; + let expr = Expression::FieldAccess { + base: Box::new(Expression::FieldAccess { + base: Box::new(Expression::FieldAccess { + base: Box::new(indexed_cell), + field: "cell".to_string(), + span: test_span(), + }), + field: "ocv_soc".to_string(), + span: test_span(), + }), + field: "u".to_string(), + span: test_span(), + }; + let substitution = test_substitution(target, var_ref("replacement")); + + let result = structural_ok(apply_substitutions_to_expr_with_derivatives_and_dae( + &expr, + &[substitution], + Some(&dae), + |_| Ok(None), + )); + + assert!( + matches!(result, Expression::VarRef { ref name, .. } if name.as_str() == "replacement"), + "nested indexed component field reference should be rewritten as one scalar target: {result:?}" + ); +} + +fn structured_ref_with_identity( + path: &str, + local: bool, + def_id: Option, +) -> rumoca_core::ComponentReference { + let mut component_ref = component_ref(path); + component_ref.local = local; + component_ref.def_id = def_id.map(rumoca_core::DefId); + component_ref +} + +fn structured_var_ref(path: &str, local: bool, def_id: Option) -> Expression { + Expression::VarRef { + name: Reference::from_component_reference(structured_ref_with_identity( + path, local, def_id, + )), + subscripts: Vec::new(), + span: test_span(), + } +} + +fn structured_substitution_fixture( + target: &str, + target_local: bool, + target_def_id: Option, +) -> (dae::Dae, Substitution) { + let mut dae = dae::Dae::default(); + let mut target_var = test_dae_variable(target); + target_var.component_ref = Some(structured_ref_with_identity( + target, + target_local, + target_def_id, + )); + dae.variables + .outputs + .insert(VarName::new(target), target_var); + (dae, test_substitution(target, var_ref("replacement"))) +} + +fn apply_dae_substitution( + dae: &dae::Dae, + expr: &Expression, + substitution: Substitution, +) -> Expression { + structural_ok(apply_substitutions_to_expr_with_derivatives_and_dae( + expr, + &[substitution], + Some(dae), + |_| Ok(None), + )) +} + +#[test] +fn structured_substitution_accepts_matching_terminal_def_id() { + let target = "plant.sensor.u"; + let (dae, substitution) = structured_substitution_fixture(target, false, Some(42)); + let expr = structured_var_ref(target, false, Some(42)); + + let result = apply_dae_substitution(&dae, &expr, substitution); + + assert!( + matches!(result, Expression::VarRef { ref name, .. } if name.as_str() == "replacement") + ); +} + +#[test] +fn structured_substitution_rejects_mismatched_def_id_and_local_scope() { + let target = "plant.sensor.u"; + let (dae, substitution) = structured_substitution_fixture(target, false, Some(42)); + for expr in [ + structured_var_ref(target, false, Some(43)), + structured_var_ref(target, true, Some(42)), + ] { + let result = apply_dae_substitution(&dae, &expr, substitution.clone()); + assert_eq!(result, expr, "identity mismatch must fail closed"); + } +} + +#[test] +fn generated_scalar_substitution_accepts_only_the_exact_canonical_owner() { + let mut dae = dae::Dae::default(); + let mut aggregate = component_var("product.u"); + aggregate.dims = vec![2]; + dae.variables + .outputs + .insert(VarName::new("product.u"), aggregate); + let substitution = Substitution { + var_name: VarName::new("product.u[2]"), + var_ref: None, + expr: var_ref("physical_source"), + var_dims: Vec::new(), + replacement_dims: Vec::new(), + env_keys: vec!["product.u[2]".to_string()], + }; + let exact = structured_var_ref("product.u[2]", false, None); + let result = apply_dae_substitution(&dae, &exact, substitution.clone()); + assert!( + matches!(result, Expression::VarRef { ref name, .. } if name.as_str() == "physical_source"), + "the generated exact scalar owner must receive its physical producer" + ); + + for other in [ + structured_var_ref("product.u[1]", false, None), + structured_var_ref("product.u[2]", true, None), + structured_var_ref("product.u[2]", false, Some(41)), + ] { + let result = apply_dae_substitution(&dae, &other, substitution.clone()); + assert_eq!( + result, other, + "sibling index and mismatched local/definition owners must fail closed" + ); + } +} + +#[test] +fn structured_substitution_rejects_neighbor_partial_and_dynamic_references() { + let target = "stack.cell[1,1].cell.ocv_soc.u"; + let (dae, substitution) = structured_substitution_fixture(target, false, Some(42)); + let dynamic_index = Expression::Index { + base: Box::new(structured_var_ref("stack.cell", false, Some(7))), + subscripts: vec![rumoca_core::Subscript::colon(test_span())], + span: test_span(), + }; + let expression_index = Expression::Index { + base: Box::new(structured_var_ref("stack.cell", false, Some(7))), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(var_ref("i")), + span: test_span(), + }], + span: test_span(), + }; + let cases = [ + structured_var_ref("stack.cell[1,2].cell.ocv_soc.u", false, Some(42)), + structured_var_ref("stack.cell[1,1].cell.ocv_soc.y", false, Some(42)), + structured_var_ref("stack.cell[1,1].cell.ocv_soc", false, Some(42)), + structured_var_ref("stack.cell[1,1]", false, Some(42)), + Expression::FieldAccess { + base: Box::new(dynamic_index), + field: "u".to_string(), + span: test_span(), + }, + Expression::FieldAccess { + base: Box::new(expression_index), + field: "u".to_string(), + span: test_span(), + }, + ]; + for expr in cases { + let result = apply_dae_substitution(&dae, &expr, substitution.clone()); + assert_eq!( + result, expr, + "non-exact structured reference must not match" + ); + } +} + +#[test] +fn structured_substitution_requires_unique_dae_target_identity() { + let target = "plant.sensor.u"; + let (mut dae, substitution) = structured_substitution_fixture(target, false, None); + let mut shadow = test_dae_variable("shadow"); + shadow.component_ref = Some(structured_ref_with_identity(target, false, None)); + dae.variables + .algebraics + .insert(VarName::new("shadow"), shadow); + let expr = structured_var_ref(target, false, None); + + let result = apply_dae_substitution(&dae, &expr, substitution); + + assert_eq!(result, expr, "ambiguous structured target must fail closed"); +} + +#[test] +fn structured_substitution_rejects_non_concrete_component_ref_part_subscripts() { + let target = "plant.sensor.u"; + let expression_subscript = |name: &str| rumoca_core::Subscript::Expr { + expr: Box::new(var_ref(name)), + span: test_span(), + }; + let cases = [ + (expression_subscript("i"), expression_subscript("i")), + (expression_subscript("i"), expression_subscript("j")), + ( + rumoca_core::Subscript::colon(test_span()), + rumoca_core::Subscript::colon(test_span()), + ), + ]; + for (target_subscript, expression_subscript) in cases { + let mut target_ref = structured_ref_with_identity(target, false, Some(42)); + target_ref.parts.last_mut().unwrap().subs = vec![target_subscript]; + let mut dae = dae::Dae::default(); + let mut target_var = test_dae_variable(target); + target_var.component_ref = Some(target_ref); + dae.variables + .outputs + .insert(VarName::new(target), target_var); + + let mut expression_ref = structured_ref_with_identity(target, false, Some(42)); + expression_ref.parts.last_mut().unwrap().subs = vec![expression_subscript]; + let expr = Expression::VarRef { + name: Reference::from_component_reference(expression_ref), + subscripts: Vec::new(), + span: test_span(), + }; + let substitution = test_substitution(target, var_ref("replacement")); + + let result = apply_dae_substitution(&dae, &expr, substitution); + + assert_eq!( + result, expr, + "non-concrete component_ref part subscript must fail closed" + ); + } +} + +#[test] +fn structured_substitution_preserves_pre_edge_change_arguments() { + let target = "plant.sensor.u"; + let (dae, substitution) = structured_substitution_fixture(target, false, Some(42)); + for function in [ + BuiltinFunction::Pre, + BuiltinFunction::Edge, + BuiltinFunction::Change, + ] { + let expr = Expression::BuiltinCall { + function, + args: vec![structured_var_ref(target, false, Some(42))], + span: test_span(), + }; + let result = apply_dae_substitution(&dae, &expr, substitution.clone()); + assert_eq!( + result, expr, + "event operator argument must remain untouched" + ); + } +} + #[test] fn substitute_var_subscripted_replacement_uses_reference_span_when_replacement_is_unspanned() { let span = Span::from_offsets( @@ -666,6 +1020,56 @@ fn substitute_var_subscripted_replacement_uses_reference_span_when_replacement_i ); } +#[test] +fn substitute_var_does_not_double_project_already_indexed_alias_replacement() { + let expr = var_ref("u[1]"); + let substitution = Substitution { + var_name: VarName::new("u[1]"), + var_ref: Some(reference("u[1]")), + expr: var_ref_idx("y", 1), + var_dims: Vec::new(), + replacement_dims: vec![2], + env_keys: vec!["u[1]".to_string()], + }; + + let result = structural_ok(apply_substitutions_to_expr(&expr, &[substitution])); + + assert!( + matches!( + result, + Expression::VarRef { name, subscripts, .. } + if name.as_str() == "y" + && matches!(subscripts.as_slice(), [rumoca_core::Subscript::Index { value: 1, .. }]) + ), + "already-indexed scalar alias replacement must not gain a second subscript" + ); +} + +#[test] +fn substitute_var_does_not_double_project_indexed_replacement_for_aggregate_use_site() { + let expr = var_ref("u[1]"); + let substitution = Substitution { + var_name: VarName::new("u"), + var_ref: Some(reference("u")), + expr: var_ref_idx("y", 1), + var_dims: vec![1], + replacement_dims: vec![2], + env_keys: vec!["u".to_string()], + }; + + let result = structural_ok(apply_substitutions_to_expr(&expr, &[substitution])); + + assert!( + matches!( + &result, + Expression::VarRef { name, subscripts, .. } + if name.as_str() == "y" + && matches!(subscripts.as_slice(), [rumoca_core::Subscript::Index { value: 1, .. }]) + ), + "aggregate use-site indexing must not double-project an already-indexed replacement: {result:?}" + ); +} + #[test] fn substitute_var_rejects_unspanned_embedded_array_alias_projection() { let expr = Expression::VarRef { @@ -733,6 +1137,226 @@ fn test_substitute_var_does_not_project_scalarized_alias_as_aggregate() { ); } +#[test] +fn test_substitute_var_does_not_rewrite_sibling_indexed_component_ref() { + let expr = var_ref("analysatorAC.iH1[2].u"); + let result = substitute_var( + &expr, + &VarName::new("analysatorAC.iH1[6].u"), + &var_ref("analysatorAC.multiSensorAC.i[6]"), + ); + + assert!( + matches!( + result, + Expression::VarRef { name, subscripts, .. } + if name.as_str() == "analysatorAC.iH1[2].u" && subscripts.is_empty() + ), + "a scalarized component substitution must not rewrite sibling indexed component instances" + ); +} + +#[test] +fn test_aggregate_subscript_substitution_does_not_rewrite_sibling_indexed_component_ref() { + let expr = var_ref_idx("analysatorAC.iH1[2].product2.u", 1); + let substitutions = [Substitution { + var_name: VarName::new("analysatorAC.iH1[6].product2.u"), + var_ref: Some(reference("analysatorAC.iH1[6].product2.u")), + expr: var_ref("analysatorAC.multiSensorAC.i"), + var_dims: vec![2], + replacement_dims: vec![6], + env_keys: Vec::new(), + }]; + let result = structural_ok(apply_substitutions_to_expr(&expr, &substitutions)); + + assert!( + matches!( + result, + Expression::VarRef { name, subscripts, .. } + if name.as_str() == "analysatorAC.iH1[2].product2.u" + && matches!(subscripts.as_slice(), [rumoca_core::Subscript::Index { value: 1, .. }]) + ), + "aggregate alias substitution must keep sibling indexed component instances distinct" + ); +} + +#[test] +fn aggregate_dims_use_structured_indexed_parent_reference_instead_of_leaf_display_name() { + let mut dae = Dae::new(); + let aggregate_name = "analysatorAC.iH1[1].product2.u"; + let mut aggregate = component_var(aggregate_name); + aggregate.dims = vec![2]; + dae.variables + .algebraics + .insert(VarName::new(aggregate_name), aggregate); + + let scalar_component_ref = component_ref("analysatorAC.iH1[1].product2.u[2]"); + let scalar_leaf_reference = Reference::with_component_reference("u[2]", scalar_component_ref); + let rewriter = RecordFieldAggregateRewriter { + aggregate_alias_groups: IndexMap::new(), + complex_groups: IndexMap::new(), + dae_scope: Some(DaeVariableScope::new(&dae)), + }; + + assert_eq!( + rewriter.aggregate_dims_for_reference(&scalar_leaf_reference), + Some(Vec::new()), + "the structured reference must resolve the indexed parent vector before projecting leaf index 2" + ); +} + +#[test] +fn scoped_scalar_target_does_not_fall_back_when_dae_identity_is_exact() { + let mut dae = Dae::new(); + let exact_name = "plant.parent[1].u[1]"; + dae.variables + .algebraics + .insert(VarName::new(exact_name), component_var(exact_name)); + let substitution = Substitution { + var_name: VarName::new(exact_name), + var_ref: None, + expr: var_ref("source"), + var_dims: Vec::new(), + replacement_dims: Vec::new(), + env_keys: Vec::new(), + }; + let scope = DaeVariableScope::new(&dae); + + assert!( + scalar_substitution_target_key_in_scope(&substitution, Some(&scope)).is_none(), + "an exact DAE scalar must fail closed instead of falling back to a display-string aggregate key" + ); +} + +#[test] +fn elimination_substitution_pipeline_keeps_indexed_parent_vector_leaves_distinct() { + let mut dae = indexed_parent_elimination_dae(); + let elimination = eliminate_trivial(&mut dae) + .expect("trivial elimination should preserve structured indexed-parent references"); + for parent in [1, 2] { + assert!( + elimination.substitutions.iter().any(|substitution| { + substitution.var_name.as_str() == format!("plant.parent[{parent}].u[1]") + }), + "parent {parent} scalar leaf should be eliminated through the real boundary pipeline" + ); + } + assert_eq!( + dae.continuous.equations.len(), + 2, + "remaining origins: {:?}", + dae.continuous + .equations + .iter() + .map(|equation| &equation.origin) + .collect::>() + ); + + for (index, equation) in dae.continuous.equations.iter().enumerate() { + if index >= 2 { + break; + } + let Expression::Binary { rhs, .. } = &equation.rhs else { + panic!("expected parent aggregate product residual"); + }; + assert!( + matches!( + rhs.as_ref(), + Expression::BuiltinCall { function: BuiltinFunction::Product, args, .. } + if matches!( + args.as_slice(), + [Expression::Array { elements, .. }] + if matches!( + elements.as_slice(), + [Expression::VarRef { name: source, subscripts: source_subscripts, .. }, + Expression::VarRef { name: leaf, subscripts: leaf_subscripts, .. }] + if source.as_str() == format!("source{}", index + 1) + && source_subscripts.is_empty() + && leaf.as_str() == format!("plant.parent[{}].u", index + 1) + && matches!(leaf_subscripts.as_slice(), [rumoca_core::Subscript::Index { value: 2, .. }]) + ) + ) + ), + "parent {} vector aggregate must materialize its eliminated leaf: {rhs:?}", + index + 1 + ); + } + assert_indexed_parent_incidence(&dae); +} + +fn indexed_parent_elimination_dae() -> Dae { + let mut dae = Dae::new(); + for parent in [1, 2] { + let name = format!("plant.parent[{parent}].u"); + let mut variable = component_var(&name); + variable.dims = vec![2]; + dae.variables + .algebraics + .insert(VarName::new(name), variable); + let source = format!("source{parent}"); + dae.variables + .parameters + .insert(VarName::new(&source), component_var(&source)); + for output in [format!("y{parent}"), format!("scalar_y{parent}")] { + dae.variables + .states + .insert(VarName::new(&output), component_var(&output)); + } + let aggregate_name = format!("plant.parent[{parent}].u"); + dae.continuous.equations.push(residual( + var_ref(&format!("y{parent}")), + builtin( + BuiltinFunction::Product, + vec![Expression::VarRef { + name: Reference::from_component_reference(component_ref(&aggregate_name)), + subscripts: Vec::new(), + span: test_span(), + }], + ), + 1, + &format!("parent {parent} aggregate use"), + )); + let scalar_use = Expression::VarRef { + name: Reference::from_component_reference(component_ref(&aggregate_name)), + subscripts: vec![rumoca_core::Subscript::Index { + value: 1, + span: test_span(), + }], + span: test_span(), + }; + dae.continuous.equations.push(residual( + var_ref(&format!("scalar_y{parent}")), + scalar_use.clone(), + 1, + &format!("parent {parent} scalar use"), + )); + dae.continuous.equations.push(residual( + scalar_use, + var_ref(&format!("source{parent}")), + 1, + &format!("connection equation: plant.parent[{parent}].u[1] = source{parent}"), + )); + } + dae +} + +fn assert_indexed_parent_incidence(dae: &Dae) { + let incidence = crate::incidence::build_incidence(dae); + let unknown_index = |name: &str| { + incidence + .unknown_names + .iter() + .position(|unknown| unknown == &UnknownId::Variable(VarName::new(name))) + .unwrap_or_else(|| panic!("missing incidence unknown {name}")) + }; + let parent1_leaf2 = unknown_index("plant.parent[1].u[2]"); + let parent2_leaf2 = unknown_index("plant.parent[2].u[2]"); + assert!(incidence.eq_unknowns[0].contains(&parent1_leaf2)); + assert!(!incidence.eq_unknowns[0].contains(&parent2_leaf2)); + assert!(incidence.eq_unknowns[1].contains(&parent2_leaf2)); + assert!(!incidence.eq_unknowns[1].contains(&parent1_leaf2)); +} + #[test] fn test_substitute_var_rewrites_exact_indexed_component_without_use_site_component_ref() { let expr = Expression::VarRef { @@ -842,6 +1466,74 @@ fn test_complete_scalar_alias_group_rewrites_aggregate_function_argument() { ); } +#[test] +fn test_partial_scalar_alias_group_rewrites_dae_aggregate_function_argument() { + let mut dae = Dae::new(); + let mut u = component_var("block.u"); + u.dims = vec![2]; + dae.variables.algebraics.insert(VarName::new("block.u"), u); + dae.variables + .algebraics + .insert(VarName::new("x"), test_dae_variable("x")); + dae.variables + .algebraics + .insert(VarName::new("y"), test_dae_variable("y")); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: sub_op(), + lhs: Box::new(var_ref("y")), + rhs: Box::new(Expression::BuiltinCall { + function: BuiltinFunction::Product, + args: vec![var_ref("block.u")], + span: test_span(), + }), + span: test_span(), + }, + span: test_span(), + origin: "product aggregate use".to_string(), + scalar_count: 1, + }); + let substitutions = [Substitution { + var_name: VarName::new("block.u[1]"), + var_ref: Some(reference("block.u[1]")), + expr: var_ref("x"), + var_dims: Vec::new(), + replacement_dims: Vec::new(), + env_keys: Vec::new(), + }]; + + apply_substitutions_to_dae_partitions(&mut dae, &substitutions) + .expect("DAE substitutions should succeed"); + + let Expression::Binary { rhs, .. } = &dae.continuous.equations[0].rhs else { + panic!("expected product equation residual"); + }; + assert!( + matches!( + rhs.as_ref(), + Expression::BuiltinCall { function: BuiltinFunction::Product, args, .. } + if matches!( + args.as_slice(), + [Expression::Array { elements, .. }] + if elements.len() == 2 + && matches!( + &elements[0], + Expression::VarRef { name, subscripts, .. } + if name.as_str() == "x" && subscripts.is_empty() + ) + && matches!( + &elements[1], + Expression::VarRef { name, subscripts, .. } + if name.as_str() == "block.u" + && matches!(subscripts.as_slice(), [rumoca_core::Subscript::Index { value: 2, .. }]) + ) + ) + ), + "partial scalar alias groups should materialize aggregate function arguments with original fallback elements" + ); +} + #[test] fn test_complete_scalar_alias_group_without_exact_var_ref_rewrites_aggregate_function_argument() { let expr = var_ref("vehicle.attitude.omega"); @@ -1421,6 +2113,98 @@ fn test_eliminate_trivial_preserves_scalarized_matrix_derivative_rows() { } } +#[test] +fn test_eliminate_trivial_preserves_runtime_sensitive_aggregate_driver_rows() { + let mut dae = Dae::new(); + let mut n_latent = test_dae_variable("nLatent"); + n_latent.start = Some(Expression::Literal { + value: Literal::Integer(8), + span: test_span(), + }); + dae.variables + .parameters + .insert(VarName::new("nLatent"), n_latent); + let mut features = test_dae_variable("features"); + features.dims = vec![10]; + features.source_span = test_span(); + dae.variables + .algebraics + .insert(VarName::new("features"), features); + dae.variables + .algebraics + .insert(VarName::new("features[9]"), component_var("features[9]")); + dae.variables + .algebraics + .insert(VarName::new("a1"), component_var("a1")); + + dae.continuous.equations.push(residual( + var_ref_with_subscript_expr( + "features", + binary( + OpBinary::Add, + var_ref("nLatent"), + Expression::Literal { + value: Literal::Integer(1), + span: test_span(), + }, + ), + ), + builtin(BuiltinFunction::Sin, vec![var_ref("time")]), + 1, + "binding equation for features[9]", + )); + dae.continuous + .equations + .push(dae::Equation::explicit_with_scalar_count( + VarName::new("features[10]"), + builtin( + BuiltinFunction::Cos, + vec![binary(OpBinary::Mul, real(0.35), var_ref("time"))], + ), + test_span(), + "explicit binding equation for features[10]", + 1, + )); + dae.continuous.equations.push(residual( + var_ref("a1"), + var_ref_idx("features", 9), + 1, + "consumer equation for a1", + )); + + let result = eliminate_trivial(&mut dae).expect("elimination should not fail structurally"); + + assert!( + result.blt_error.is_none(), + "runtime-sensitive aggregate driver rows must remain matchable: {:?}", + result.blt_error + ); + assert!( + dae.variables + .algebraics + .contains_key(&VarName::new("features[9]")), + "features[9] must stay as a refreshable continuous producer" + ); + assert!( + dae.continuous + .equations + .iter() + .any(|eq| assignment_target_name_in_dae(&dae, &eq.rhs).as_ref() + == Some(&VarName::new("features[9]")) + && expr_contains_runtime_sensitive_operator(&eq.rhs)), + "features[9] = sin(time) producer row was eliminated" + ); + assert!( + dae.continuous.equations.iter().any(|eq| { + eq.lhs + .as_ref() + .is_some_and(|lhs| lhs.as_str() == "features[10]") + && expr_contains_runtime_sensitive_operator(&eq.rhs) + }), + "explicit features[10] = cos(0.35*time) producer row was eliminated" + ); +} + #[test] fn test_eliminate_trivial_reports_missing_replacement_metadata() { let mut dae = Dae::new(); @@ -1789,6 +2573,33 @@ fn shift_structured_families_drops_family_with_removed_interior_row() { ); } +#[test] +fn shift_structured_families_drops_overflowing_row_ranges() { + for (name, first_equation_index, equation_counts) in [ + ("row-count sum overflow", 0, vec![usize::MAX, 1]), + ("first-index overflow", usize::MAX, vec![1]), + ] { + let mut dae = Dae::new(); + dae.continuous.structured_equations = vec![dae::StructuredEquationFamily { + domain: rumoca_core::StructuredIndexDomain { binders: vec![] }, + first_equation_index, + equation_counts, + span: test_span(), + origin: name.to_string(), + regular: None, + template: None, + interiors_materialized: true, + }]; + + shift_structured_families_after_equation_removal(&mut dae, &[0]); + + assert!( + dae.continuous.structured_equations.is_empty(), + "malformed family must be dropped: {name}" + ); + } +} + /// A substitution can rewrite a structured family's row bodies while leaving the /// row count unchanged. The original family proof no longer applies, so the /// family must be dropped and lowered as scalar rows. @@ -1824,6 +2635,184 @@ fn drop_structured_families_touching_equations_drops_rewritten_family() { ); } +#[test] +fn drop_structured_families_handles_compact_corners_and_malformed_ranges() { + for case in compact_family_cases() + .into_iter() + .chain(malformed_family_cases()) + { + assert_compact_family_retention(case); + } +} + +type CompactFamilyCase = ( + &'static str, + Vec, + usize, + Vec, + Vec, + bool, +); + +fn compact_family_cases() -> Vec { + vec![ + ( + "1d positive non-unit base corner", + vec![compact_family_binder(0, 1, 7, 2)], + 10, + vec![1; 4], + vec![10], + false, + ), + ( + "1d positive non-unit corner", + vec![compact_family_binder(0, 1, 7, 2)], + 10, + vec![1; 4], + vec![11], + false, + ), + ( + "1d positive non-unit interior", + vec![compact_family_binder(0, 1, 7, 2)], + 10, + vec![1; 4], + vec![12], + true, + ), + ( + "1d negative non-unit corner", + vec![compact_family_binder(0, 7, 1, -2)], + 20, + vec![1; 4], + vec![21], + false, + ), + ( + "1d negative non-unit interior", + vec![compact_family_binder(0, 7, 1, -2)], + 20, + vec![1; 4], + vec![22], + true, + ), + ( + "2d corner row range", + vec![ + compact_family_binder(0, 1, 3, 1), + compact_family_binder(1, 10, 14, 2), + ], + 100, + vec![1, 2, 1, 1, 1, 1, 1, 1, 1], + vec![102], + false, + ), + ( + "2d interior row", + vec![ + compact_family_binder(0, 1, 3, 1), + compact_family_binder(1, 10, 14, 2), + ], + 100, + vec![1, 2, 1, 1, 1, 1, 1, 1, 1], + vec![103], + true, + ), + ( + "2d outer-dimension corner", + vec![ + compact_family_binder(0, 1, 3, 1), + compact_family_binder(1, 10, 14, 2), + ], + 100, + vec![1, 2, 1, 1, 1, 1, 1, 1, 1], + vec![104], + false, + ), + ] +} + +fn malformed_family_cases() -> Vec { + vec![ + ( + "empty domain", + vec![compact_family_binder(0, 1, 0, 1)], + 400, + vec![], + vec![400], + true, + ), + ( + "domain count mismatch", + vec![compact_family_binder(0, 1, 3, 1)], + 500, + vec![1, 1], + vec![500], + false, + ), + ( + "row count overflow", + vec![compact_family_binder(0, 1, 2, 1)], + 0, + vec![usize::MAX, 1], + vec![0], + false, + ), + ( + "first row overflow", + vec![compact_family_binder(0, 1, 1, 1)], + usize::MAX, + vec![1], + vec![usize::MAX], + false, + ), + ] +} + +fn compact_family_binder( + id: usize, + lower: i64, + upper: i64, + step: i64, +) -> rumoca_core::StructuredIndexBinder { + rumoca_core::StructuredIndexBinder { + id, + display_name: format!("i{id}"), + lower, + upper, + step, + } +} + +fn assert_compact_family_retention(case: CompactFamilyCase) { + let (name, binders, first_equation_index, equation_counts, touched, retained) = case; + let regular_binders = binders + .iter() + .map(|binder| binder.display_name.clone()) + .collect(); + let mut dae = Dae::new(); + dae.continuous.structured_equations = vec![dae::StructuredEquationFamily { + domain: rumoca_core::StructuredIndexDomain { binders }, + first_equation_index, + equation_counts, + span: test_span(), + origin: name.to_string(), + regular: Some(rumoca_core::RegularForFamily { + binders: regular_binders, + accesses: vec![], + }), + template: None, + interiors_materialized: false, + }]; + + drop_structured_families_touching_equations(&mut dae, &touched); + assert_eq!( + !dae.continuous.structured_equations.is_empty(), + retained, + "{name}" + ); +} + /// Residual rows have no `lhs`, so substitution used to rewrite the RHS and /// return before recording the touched row. Structured metadata for that row /// must still be dropped because its compact body proof is now stale. diff --git a/crates/rumoca-phase-structural/src/eliminate/tests/array_boundary.rs b/crates/rumoca-phase-structural/src/eliminate/tests/array_boundary.rs index 7a4644720..2b2f38a4c 100644 --- a/crates/rumoca-phase-structural/src/eliminate/tests/array_boundary.rs +++ b/crates/rumoca-phase-structural/src/eliminate/tests/array_boundary.rs @@ -36,3 +36,54 @@ fn test_boundary_keeps_array_slice_unknowns_before_scalarization() { "array slice row must remain for scalarization" ); } + +#[test] +fn test_boundary_resolves_constant_expression_subscript_to_scalarized_unknown() { + let mut dae = Dae::new(); + let mut n = component_var("N"); + n.start = Some(Expression::Literal { + value: Literal::Integer(10), + span: test_span(), + }); + dae.variables.parameters.insert(VarName::new("N"), n); + dae.variables + .algebraics + .insert(VarName::new("Q[11]"), component_var("Q[11]")); + + let subscript_expr = Expression::Binary { + op: OpBinary::Add, + lhs: Box::new(var_ref("N")), + rhs: Box::new(Expression::Literal { + value: Literal::Integer(1), + span: test_span(), + }), + span: test_span(), + }; + dae.continuous.equations.push(residual( + var_ref_with_subscript_expr("Q", subscript_expr), + real(0.0), + 1, + "Q boundary", + )); + + let result = eliminate_trivial(&mut dae).expect("structural elimination should succeed"); + + assert_eq!(result.n_eliminated, 1); + assert!(dae.continuous.equations.is_empty()); + assert!( + !dae.variables + .algebraics + .contains_key(&VarName::new("Q[11]")) + ); +} + +fn var_ref_with_subscript_expr(name: &str, expr: Expression) -> Expression { + Expression::VarRef { + name: reference(name), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(expr), + span: test_span(), + }], + span: test_span(), + } +} diff --git a/crates/rumoca-phase-structural/src/eliminate/tests/boundary_cases.rs b/crates/rumoca-phase-structural/src/eliminate/tests/boundary_cases.rs index f5f2e265c..17bf934b7 100644 --- a/crates/rumoca-phase-structural/src/eliminate/tests/boundary_cases.rs +++ b/crates/rumoca-phase-structural/src/eliminate/tests/boundary_cases.rs @@ -204,6 +204,85 @@ fn test_boundary_cascade_resolution() { assert!(dae.variables.algebraics.is_empty()); } +#[test] +fn test_boundary_eliminates_indexed_scalar_connection_alias() { + let mut dae = Dae::new(); + dae.variables + .algebraics + .insert(VarName::new("pin[2].v"), test_dae_variable("pin[2].v")); + dae.variables + .algebraics + .insert(VarName::new("node.v"), test_dae_variable("node.v")); + + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: sub_op(), + lhs: Box::new(var_ref("pin[2].v")), + rhs: Box::new(var_ref("node.v")), + span: rumoca_core::Span::DUMMY, + }, + span: Span::DUMMY, + origin: "connection equation: pin[2].v = node.v".to_string(), + scalar_count: 1, + }); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: sub_op(), + lhs: Box::new(var_ref("node.v")), + rhs: Box::new(lit(0.0)), + span: rumoca_core::Span::DUMMY, + }, + span: Span::DUMMY, + origin: "ground equation".to_string(), + scalar_count: 1, + }); + + let result = eliminate_trivial(&mut dae).expect("structural elimination should succeed"); + + assert!(result.blt_error.is_none()); + assert_eq!(result.n_eliminated, 2); + assert!(dae.continuous.equations.is_empty()); + assert!(dae.variables.algebraics.is_empty()); +} + +#[test] +fn test_boundary_eliminates_scalar_output_connection_alias() { + let mut dae = Dae::new(); + dae.variables + .outputs + .insert(VarName::new("sensor.y"), test_dae_variable("sensor.y")); + dae.variables + .parameters + .insert(VarName::new("source.y"), test_dae_variable("source.y")); + + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: sub_op(), + lhs: Box::new(var_ref("sensor.y")), + rhs: Box::new(var_ref("source.y")), + span: rumoca_core::Span::DUMMY, + }, + span: Span::DUMMY, + origin: "connection equation: sensor.y = source.y".to_string(), + scalar_count: 1, + }); + + let result = eliminate_trivial(&mut dae).expect("structural elimination should succeed"); + + assert!(result.blt_error.is_none()); + assert_eq!(result.n_eliminated, 1); + assert!(dae.continuous.equations.is_empty()); + assert!( + !dae.variables + .outputs + .contains_key(&VarName::new("sensor.y")) + ); + assert_eq!(result.substitutions[0].var_name.as_str(), "sensor.y"); +} + #[test] fn test_boundary_eliminates_negated_additive_single_unknown() { let mut dae = Dae::new(); @@ -508,19 +587,667 @@ fn test_boundary_eliminates_single_unknown_connection_after_substitution() { }); let result = eliminate_trivial(&mut dae).expect("structural elimination should succeed"); - // y is a non-trivial output (if-expression) — preserved in the DAE. - // u cannot be eliminated because y also remains live in the connection - // equation, keeping both unknowns alive. - assert_eq!(result.n_eliminated, 0); - assert_eq!(dae.continuous.equations.len(), 2); + // y is a non-trivial output (if-expression), so preserve y and its source + // equation. The connection alias can still eliminate the scalar sink u. + assert_eq!(result.n_eliminated, 1); + assert_eq!(dae.continuous.equations.len(), 1); assert!( dae.variables.outputs.contains_key(&VarName::new("y")), "output y should remain (non-trivial expression)" ); assert!( - dae.variables.algebraics.contains_key(&VarName::new("u")), - "u should remain (y not eliminated, connection eq still has two unknowns)" + !dae.variables.algebraics.contains_key(&VarName::new("u")), + "connection sink u should be eliminated as an alias of y" + ); +} + +#[test] +fn test_boundary_connection_alias_prefers_rhs_sink_element() { + let mut dae = Dae::new(); + dae.variables.algebraics.insert( + VarName::new("analysatorAC.voltageLine2Line[6].product1.u[2]"), + test_dae_variable("analysatorAC.voltageLine2Line[6].product1.u[2]"), + ); + dae.variables.algebraics.insert( + VarName::new("analysatorAC.voltageLine2Line[6].product2.u[1]"), + test_dae_variable("analysatorAC.voltageLine2Line[6].product2.u[1]"), + ); + + let connection_rhs = Expression::Binary { + op: sub_op(), + lhs: Box::new(var_ref("analysatorAC.voltageLine2Line[6].product1.u[2]")), + rhs: Box::new(var_ref("analysatorAC.voltageLine2Line[6].product2.u[1]")), + span: rumoca_core::Span::DUMMY, + }; + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: connection_rhs.clone(), + span: Span::DUMMY, + origin: "connection equation: analysatorAC.voltageLine2Line[6].product1.u[2] = analysatorAC.voltageLine2Line[6].product2.u[1]".to_string(), + scalar_count: 1, + }); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: OpBinary::Add, + lhs: Box::new(var_ref("analysatorAC.voltageLine2Line[6].product1.u[2]")), + rhs: Box::new(var_ref("analysatorAC.voltageLine2Line[6].product2.u[1]")), + span: rumoca_core::Span::DUMMY, + }, + span: Span::DUMMY, + origin: "product input use".to_string(), + scalar_count: 1, + }); + + let live = vec![ + VarName::new("analysatorAC.voltageLine2Line[6].product1.u[2]"), + VarName::new("analysatorAC.voltageLine2Line[6].product2.u[1]"), + ]; + let direct_definitions = DirectDefinitionIndex::build(&dae); + let protected = IndexSet::new(); + let choice_ctx = EliminationChoiceContext { + dae: &dae, + eq_idx: 0, + has_state_derivative: false, + runtime_protected_unknowns: &protected, + direct_definitions: &direct_definitions, + allow_multi_live_trivial_alias: true, + }; + let (var_name, solution) = + choose_solvable_unknown_for_elimination(&choice_ctx, &connection_rhs, &live) + .expect("candidate selection should not fail") + .expect("connection alias should be solvable"); + + assert_eq!( + var_name.as_str(), + "analysatorAC.voltageLine2Line[6].product2.u[1]" + ); + assert!( + matches!( + solution, + Expression::VarRef { ref name, ref subscripts, .. } + if name.as_str() == "analysatorAC.voltageLine2Line[6].product1.u[2]" + && subscripts.is_empty() + ), + "sink input should resolve to source input, got {:?}", + solution + ); +} + +#[test] +fn test_boundary_connection_policy_accepts_scalar_element_of_aggregate_var() { + let mut dae = Dae::new(); + let mut product1_u = test_dae_variable("analysatorAC.voltageLine2Line[6].product1.u"); + product1_u.dims = vec![2]; + dae.variables.algebraics.insert( + VarName::new("analysatorAC.voltageLine2Line[6].product1.u"), + product1_u, + ); + let mut product2_u = test_dae_variable("analysatorAC.voltageLine2Line[6].product2.u"); + product2_u.dims = vec![2]; + dae.variables.algebraics.insert( + VarName::new("analysatorAC.voltageLine2Line[6].product2.u"), + product2_u, + ); + + let connection_rhs = Expression::Binary { + op: sub_op(), + lhs: Box::new(var_ref_idx( + "analysatorAC.voltageLine2Line[6].product1.u", + 2, + )), + rhs: Box::new(var_ref_idx( + "analysatorAC.voltageLine2Line[6].product2.u", + 1, + )), + span: rumoca_core::Span::DUMMY, + }; + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: connection_rhs.clone(), + span: Span::DUMMY, + origin: "connection equation: analysatorAC.voltageLine2Line[6].product1.u[2] = analysatorAC.voltageLine2Line[6].product2.u[1]".to_string(), + scalar_count: 1, + }); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: OpBinary::Add, + lhs: Box::new(var_ref_idx( + "analysatorAC.voltageLine2Line[6].product1.u", + 2, + )), + rhs: Box::new(var_ref_idx( + "analysatorAC.voltageLine2Line[6].product2.u", + 1, + )), + span: rumoca_core::Span::DUMMY, + }, + span: Span::DUMMY, + origin: "product input use".to_string(), + scalar_count: 1, + }); + + let live = vec![ + VarName::new("analysatorAC.voltageLine2Line[6].product1.u[2]"), + VarName::new("analysatorAC.voltageLine2Line[6].product2.u[1]"), + ]; + + assert!( + !should_skip_connection_equation(&dae, &connection_rhs, true, &live, &HashSet::new(),), + "scalar element aliases of aggregate variables should be eligible for connection elimination" + ); +} + +#[test] +fn test_boundary_connection_policy_preserves_aggregate_only_scalar_elements() { + let mut dae = Dae::new(); + for name in ["block.product1.u", "block.product2.u"] { + let mut variable = test_dae_variable(name); + variable.dims = vec![2]; + dae.variables + .algebraics + .insert(VarName::new(name), variable); + } + + let connection_rhs = Expression::Binary { + op: sub_op(), + lhs: Box::new(var_ref_idx("block.product1.u", 2)), + rhs: Box::new(var_ref_idx("block.product2.u", 1)), + span: test_span(), + }; + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: connection_rhs.clone(), + span: test_span(), + origin: "connection equation: block.product1.u[2] = block.product2.u[1]".to_string(), + scalar_count: 1, + }); + for (output, aggregate) in [ + ("product1.y", "block.product1.u"), + ("product2.y", "block.product2.u"), + ] { + dae.variables + .algebraics + .insert(VarName::new(output), test_dae_variable(output)); + dae.continuous.equations.push(residual( + var_ref(output), + builtin(BuiltinFunction::Product, vec![var_ref(aggregate)]), + 1, + "aggregate-only product input use", + )); + } + + let live = vec![ + VarName::new("block.product1.u[2]"), + VarName::new("block.product2.u[1]"), + ]; + + assert!( + should_skip_connection_equation(&dae, &connection_rhs, true, &live, &HashSet::new()), + "a scalar aggregate leaf whose only non-connection use is the materializable aggregate must keep its connection equation because the base vector storage cannot remove one leaf" + ); +} + +#[test] +fn test_boundary_eliminates_scalarized_element_with_boundary_known_derivative_use() { + let mut dae = Dae::new(); + dae.variables + .states + .insert(VarName::new("x"), test_dae_variable("x")); + let mut v = test_dae_variable("v"); + v.dims = vec![3]; + dae.variables.algebraics.insert(VarName::new("v"), v); + + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: sub_op(), + lhs: Box::new(var_ref_idx("v", 2)), + rhs: Box::new(lit(20.0)), + span: rumoca_core::Span::DUMMY, + }, + span: Span::DUMMY, + origin: "binding equation for v[2]".to_string(), + scalar_count: 1, + }); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: sub_op(), + lhs: Box::new(Expression::BuiltinCall { + function: BuiltinFunction::Der, + args: vec![var_ref("x")], + span: rumoca_core::Span::DUMMY, + }), + rhs: Box::new(var_ref_idx("v", 2)), + span: rumoca_core::Span::DUMMY, + }, + span: Span::DUMMY, + origin: "ode".to_string(), + scalar_count: 1, + }); + + let result = eliminate_trivial(&mut dae).expect("structural elimination should succeed"); + + assert!( + result + .substitutions + .iter() + .any(|sub| sub.var_name.as_str() == "v[2]"), + "boundary-known scalarized element should be substituted into derivative use" + ); + assert_eq!(dae.continuous.equations.len(), 1); + let Expression::Binary { rhs, .. } = &dae.continuous.equations[0].rhs else { + panic!("remaining ODE should stay as a binary residual"); + }; + assert!( + matches!( + rhs.as_ref(), + Expression::Literal { + value: Literal::Real(value), + .. + } if *value == 20.0 + ), + "derivative row should receive the boundary-known scalar value, got {:?}", + dae.continuous.equations[0].rhs + ); +} + +#[test] +fn test_boundary_rejects_derivative_alias_backed_only_by_connection_definition() { + let mut dae = Dae::new(); + dae.variables + .states + .insert(VarName::new("x"), test_dae_variable("x")); + for (name, causality) in [ + ("mean.u", dae::VariableCausality::Input), + ("sensor.v", dae::VariableCausality::Output), + ] { + let mut variable = test_dae_variable(name); + variable.causality = causality; + dae.variables.outputs.insert(VarName::new(name), variable); + } + dae.continuous + .equations + .push(residual(der(var_ref("x")), var_ref("mean.u"), 1, "ode")); + dae.continuous.equations.push(residual( + var_ref("sensor.v"), + var_ref("mean.u"), + 1, + "connection equation: sensor.v = mean.u", + )); + + let direct_definitions = DirectDefinitionIndex::build(&dae); + let runtime_protected_unknowns = IndexSet::new(); + let ctx = EliminationChoiceContext { + dae: &dae, + eq_idx: 0, + has_state_derivative: true, + runtime_protected_unknowns: &runtime_protected_unknowns, + direct_definitions: &direct_definitions, + allow_multi_live_trivial_alias: false, + }; + let choice = choose_solvable_unknown_for_elimination( + &ctx, + &dae.continuous.equations[0].rhs, + &[VarName::new("mean.u")], + ) + .expect("candidate analysis should succeed"); + + assert!( + choice.is_none(), + "a connection-only future definition is not a proven surviving producer" + ); +} + +fn add_derivative_aggregate_variables(dae: &mut Dae) { + dae.variables + .states + .insert(VarName::new("x"), test_dae_variable("x")); + dae.variables + .states + .insert(VarName::new("rms"), test_dae_variable("rms")); + for name in ["a", "b"] { + dae.variables + .algebraics + .insert(VarName::new(name), test_dae_variable(name)); + } + dae.continuous.equations.push(residual( + binary(OpBinary::Mul, var_ref("a"), var_ref("b")), + lit(1.0), + 1, + "coupled source constraint", + )); + dae.continuous.equations.push(residual( + binary(OpBinary::Add, var_ref("a"), var_ref("b")), + lit(2.0), + 1, + "coupled source constraint", + )); + let mut inputs = component_var("product.u"); + inputs.dims = vec![2]; + inputs.causality = dae::VariableCausality::Input; + dae.variables + .outputs + .insert(VarName::new("product.u"), inputs); + for (name, causality) in [ + ("product.y", dae::VariableCausality::Output), + ("sensor.v", dae::VariableCausality::Output), + ("mean.u", dae::VariableCausality::Input), + ("rms.u", dae::VariableCausality::Input), + ("rms.mean.u", dae::VariableCausality::Input), + ] { + let mut variable = component_var(name); + variable.causality = causality; + dae.variables.outputs.insert(VarName::new(name), variable); + } +} + +fn indexed_product_input(index: usize) -> Expression { + let name = format!("product.u[{index}]"); + Expression::VarRef { + name: Reference::with_component_reference(&name, component_ref(&name)), + subscripts: Vec::new(), + span: test_span(), + } +} + +fn add_derivative_aggregate_equations(dae: &mut Dae) { + dae.continuous.equations.push(residual( + var_ref("product.y"), + builtin(BuiltinFunction::Product, vec![var_ref("product.u")]), + 1, + "aggregate product", + )); + dae.continuous.equations.push(residual( + der(var_ref("rms")), + var_ref("rms.mean.u"), + 1, + "aggregate consumer ode", + )); + dae.continuous + .equations + .push(residual(der(var_ref("x")), var_ref("mean.u"), 1, "ode")); + dae.continuous.equations.push(residual( + var_ref("sensor.v"), + binary(OpBinary::Sub, var_ref("a"), var_ref("b")), + 1, + "source definition", + )); + dae.continuous.equations.push(residual( + var_ref("product.y"), + var_ref("rms.mean.u"), + 1, + "connection equation: product.y = rms.mean.u", + )); + dae.continuous.equations.push(residual( + var_ref("rms.u"), + indexed_product_input(1), + 1, + "connection equation: rms.u = product.u[1]", + )); + dae.continuous.equations.push(residual( + indexed_product_input(2), + var_ref("sensor.v"), + 1, + "connection equation: product.u[2] = sensor.v", + )); + dae.continuous.equations.push(residual( + var_ref("sensor.v"), + var_ref("mean.u"), + 1, + "connection equation: sensor.v = mean.u", + )); + dae.continuous.equations.push(residual( + indexed_product_input(1), + indexed_product_input(2), + 1, + "connection equation: product.u[1] = product.u[2]", + )); +} + +#[test] +fn test_boundary_keeps_producer_for_derivative_dependent_aggregate_leaf() { + let mut dae = Dae::new(); + add_derivative_aggregate_variables(&mut dae); + add_derivative_aggregate_equations(&mut dae); + + eliminate_trivial(&mut dae).expect("structural elimination should succeed"); + + let state_names = dae.variables.states.keys().cloned().collect::>(); + let state_derivatives = DerivativeNameMatcher::from_var_names(state_names.iter()); + let ode = dae + .continuous + .equations + .iter() + .find(|eq| { + expr_contains_der_of_any(&eq.rhs, &state_derivatives) + && expr_contains_var(&eq.rhs, &VarName::new("x")) + }) + .expect("state derivative row must remain"); + assert!( + expr_contains_var(&ode.rhs, &VarName::new("a")) + && !expr_contains_var(&ode.rhs, &VarName::new("mean.u")) + && !expr_contains_var(&ode.rhs, &VarName::new("sensor.v")) + && !expr_contains_var(&ode.rhs, &VarName::new("product.u[2]")), + "the derivative must resolve to the canonical physical source, got {:?}", + ode.rhs + ); +} + +#[test] +fn test_boundary_eliminates_scalarized_element_with_static_index_expression_derivative_use() { + let mut dae = Dae::new(); + dae.variables + .states + .insert(VarName::new("x"), test_dae_variable("x")); + let mut v = test_dae_variable("v"); + v.dims = vec![3]; + dae.variables.algebraics.insert(VarName::new("v"), v); + + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: sub_op(), + lhs: Box::new(var_ref_idx("v", 2)), + rhs: Box::new(lit(20.0)), + span: rumoca_core::Span::DUMMY, + }, + span: Span::DUMMY, + origin: "binding equation for v[2]".to_string(), + scalar_count: 1, + }); + let static_index = Expression::Binary { + op: OpBinary::Add, + lhs: Box::new(Expression::BuiltinCall { + function: BuiltinFunction::Mod, + args: vec![lit(3.0), lit(2.0)], + span: rumoca_core::Span::DUMMY, + }), + rhs: Box::new(lit(1.0)), + span: rumoca_core::Span::DUMMY, + }; + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: sub_op(), + lhs: Box::new(Expression::BuiltinCall { + function: BuiltinFunction::Der, + args: vec![var_ref("x")], + span: rumoca_core::Span::DUMMY, + }), + rhs: Box::new(var_ref_with_subscript_expr("v", static_index)), + span: rumoca_core::Span::DUMMY, + }, + span: Span::DUMMY, + origin: "ode with static subscript expression".to_string(), + scalar_count: 1, + }); + + let result = eliminate_trivial(&mut dae).expect("structural elimination should succeed"); + + assert!( + result + .substitutions + .iter() + .any(|sub| sub.var_name.as_str() == "v[2]"), + "boundary-known scalarized element should be substituted into derivative use" + ); + assert_eq!(dae.continuous.equations.len(), 1); + let Expression::Binary { rhs, .. } = &dae.continuous.equations[0].rhs else { + panic!("remaining ODE should stay as a binary residual"); + }; + assert!( + matches!( + rhs.as_ref(), + Expression::Literal { + value: Literal::Real(value), + .. + } if *value == 20.0 + ), + "static-index derivative row should receive the boundary-known scalar value, got {:?}", + dae.continuous.equations[0].rhs + ); +} + +#[test] +fn test_boundary_eliminates_internal_output_nontrivial_definition() { + let mut dae = Dae::new(); + let mut y = test_dae_variable("block.y"); + y.component_ref = + rumoca_core::component_reference_from_flat_name(&VarName::new("block.y"), Span::DUMMY); + dae.variables.outputs.insert(VarName::new("block.y"), y); + dae.variables + .algebraics + .insert(VarName::new("u"), test_dae_variable("u")); + dae.variables + .parameters + .insert(VarName::new("p"), test_dae_variable("p")); + + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: sub_op(), + lhs: Box::new(var_ref("block.y")), + rhs: Box::new(Expression::Binary { + op: OpBinary::Add, + lhs: Box::new(var_ref("p")), + rhs: Box::new(lit(1.0)), + span: rumoca_core::Span::DUMMY, + }), + span: rumoca_core::Span::DUMMY, + }, + span: Span::DUMMY, + origin: "internal output definition".to_string(), + scalar_count: 1, + }); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: sub_op(), + lhs: Box::new(var_ref("block.y")), + rhs: Box::new(var_ref("u")), + span: rumoca_core::Span::DUMMY, + }, + span: Span::DUMMY, + origin: "connection equation: block.y = u".to_string(), + scalar_count: 1, + }); + + let result = eliminate_trivial(&mut dae).expect("structural elimination should succeed"); + + assert!(result.blt_error.is_none()); + assert_eq!(result.n_eliminated, 2); + assert!(dae.continuous.equations.is_empty()); + assert!(!dae.variables.outputs.contains_key(&VarName::new("block.y"))); + assert!(!dae.variables.algebraics.contains_key(&VarName::new("u"))); +} + +#[test] +fn test_boundary_eliminates_internal_output_scalar_reduction_definition() { + let mut dae = Dae::new(); + let mut y = test_dae_variable("block.y"); + y.component_ref = + rumoca_core::component_reference_from_flat_name(&VarName::new("block.y"), Span::DUMMY); + dae.variables.outputs.insert(VarName::new("block.y"), y); + + let mut u = test_dae_variable("block.u"); + u.dims = vec![2]; + dae.variables.algebraics.insert(VarName::new("block.u"), u); + dae.variables + .algebraics + .insert(VarName::new("sink.u"), test_dae_variable("sink.u")); + dae.variables + .parameters + .insert(VarName::new("p1"), test_dae_variable("p1")); + dae.variables + .parameters + .insert(VarName::new("p2"), test_dae_variable("p2")); + + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: sub_op(), + lhs: Box::new(var_ref("block.y")), + rhs: Box::new(Expression::BuiltinCall { + function: BuiltinFunction::Product, + args: vec![var_ref("block.u")], + span: rumoca_core::Span::DUMMY, + }), + span: rumoca_core::Span::DUMMY, + }, + span: Span::DUMMY, + origin: "internal reduction output definition".to_string(), + scalar_count: 1, + }); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: sub_op(), + lhs: Box::new(var_ref("block.y")), + rhs: Box::new(var_ref("sink.u")), + span: rumoca_core::Span::DUMMY, + }, + span: Span::DUMMY, + origin: "connection equation: block.y = sink.u".to_string(), + scalar_count: 1, + }); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: sub_op(), + lhs: Box::new(var_ref("block.u[1]")), + rhs: Box::new(var_ref("p1")), + span: rumoca_core::Span::DUMMY, + }, + span: Span::DUMMY, + origin: "block.u[1] = p1".to_string(), + scalar_count: 1, + }); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: sub_op(), + lhs: Box::new(var_ref("block.u[2]")), + rhs: Box::new(var_ref("p2")), + span: rumoca_core::Span::DUMMY, + }, + span: Span::DUMMY, + origin: "block.u[2] = p2".to_string(), + scalar_count: 1, + }); + + let result = eliminate_trivial(&mut dae).expect("structural elimination should succeed"); + + assert!(result.blt_error.is_none()); + assert!( + result + .substitutions + .iter() + .any(|sub| sub.var_name.as_str() == "block.y") ); + assert!(!dae.variables.outputs.contains_key(&VarName::new("block.y"))); } #[test] @@ -669,6 +1396,179 @@ fn test_eliminate_trivial_accepts_runtime_known_assignment_tail_after_output_ali assert!(dae.variables.outputs.is_empty()); } +#[test] +fn test_boundary_eliminates_connection_rhs_output_from_source_expression() { + let mut dae = Dae::new(); + let mut state = test_dae_variable("x"); + state.start = Some(lit(0.0)); + dae.variables.states.insert(VarName::new("x"), state); + dae.variables + .outputs + .insert(VarName::new("sink.u"), test_dae_variable("sink.u")); + dae.variables.parameters.insert( + VarName::new("source.offset"), + test_dae_variable("source.offset"), + ); + dae.variables + .discrete_valued + .insert(VarName::new("c"), test_dae_variable("c")); + + dae.continuous + .equations + .push(residual(der(var_ref("x")), var_ref("sink.u"), 1, "ode")); + dae.continuous.equations.push(residual( + binary( + OpBinary::Add, + var_ref("source.offset"), + Expression::If { + branches: vec![( + Expression::VarRef { + name: rumoca_core::Reference::generated("c"), + subscripts: vec![rumoca_core::Subscript::generated_index(1, test_span())], + span: test_span(), + }, + lit(0.0), + )], + else_branch: Box::new(lit(1.0)), + span: test_span(), + }, + ), + var_ref("sink.u"), + 1, + "connection equation: source.y = sink.u", + )); + + let result = eliminate_trivial(&mut dae).expect("structural elimination should succeed"); + + assert_eq!(result.n_eliminated, 1); + assert!( + result + .substitutions + .iter() + .any(|sub| sub.var_name.as_str() == "sink.u"), + "source-driven input port should be substituted" + ); + assert_eq!(dae.continuous.equations.len(), 1); + assert!( + !expr_contains_var(&dae.continuous.equations[0].rhs, &VarName::new("sink.u")), + "remaining ODE should reference the source expression directly" + ); + assert!(!dae.variables.outputs.contains_key(&VarName::new("sink.u"))); +} + +#[test] +fn test_boundary_eliminates_input_connected_to_non_dae_source_alias() { + let mut dae = Dae::new(); + let mut state = test_dae_variable("x"); + state.start = Some(lit(0.0)); + dae.variables.states.insert(VarName::new("x"), state); + let mut sink_u = test_dae_variable("sink.u"); + sink_u.causality = dae::VariableCausality::Input; + dae.variables.outputs.insert(VarName::new("sink.u"), sink_u); + + dae.continuous.equations.push(residual( + var_ref("source.y"), + var_ref("sink.u"), + 1, + "connection equation: source.y = sink.u", + )); + dae.continuous + .equations + .push(residual(der(var_ref("x")), var_ref("sink.u"), 1, "ode")); + + let result = eliminate_trivial(&mut dae).expect("structural elimination should succeed"); + let sink_sub = result + .substitutions + .iter() + .find(|sub| sub.var_name.as_str() == "sink.u") + .expect("input connected to a source alias should be eliminated"); + + assert!( + matches!( + sink_sub.expr, + Expression::VarRef { ref name, ref subscripts, .. } + if name.as_str() == "source.y" && subscripts.is_empty() + ), + "input should resolve to the non-DAE source alias, got {:?}", + sink_sub.expr + ); +} + +#[test] +fn test_boundary_prefers_connection_input_over_source_output() { + let mut dae = Dae::new(); + let mut source_y = test_dae_variable("source.y"); + source_y.causality = dae::VariableCausality::Output; + dae.variables + .outputs + .insert(VarName::new("source.y"), source_y); + let mut sink_u = test_dae_variable("sink.u"); + sink_u.causality = dae::VariableCausality::Input; + dae.variables.outputs.insert(VarName::new("sink.u"), sink_u); + dae.variables + .algebraics + .insert(VarName::new("consumer.u"), test_dae_variable("consumer.u")); + + dae.continuous.equations.push(residual( + var_ref("source.y"), + var_ref("sink.u"), + 1, + "connection equation: source.y = sink.u", + )); + dae.continuous.equations.push(residual( + var_ref("consumer.u"), + var_ref("sink.u"), + 1, + "consumer equation", + )); + + let result = eliminate_trivial(&mut dae).expect("structural elimination should succeed"); + let sink_sub = result + .substitutions + .iter() + .find(|sub| sub.var_name.as_str() == "sink.u") + .expect("connection input should be eliminated"); + + assert!( + matches!( + sink_sub.expr, + Expression::VarRef { ref name, ref subscripts, .. } + if name.as_str() == "source.y" && subscripts.is_empty() + ), + "connection should substitute sink input from source output, got {:?}", + sink_sub.expr + ); + assert!( + result + .substitutions + .iter() + .all(|sub| sub.var_name.as_str() != "source.y"), + "connection alias should preserve the source output" + ); +} + +#[test] +fn test_exact_subscript_index_in_dae_resolves_builtin_integer_expression() { + let dae = Dae::new(); + let subscript = rumoca_core::Subscript::Expr { + expr: Box::new(binary( + OpBinary::Add, + binary( + OpBinary::Add, + builtin(BuiltinFunction::Mod, vec![lit(3.0), lit(2.0)]), + builtin( + BuiltinFunction::Integer, + vec![binary(OpBinary::Div, lit(3.0), lit(2.0))], + ), + ), + lit(1.0), + )), + span: test_span(), + }; + + assert_eq!(exact_subscript_index_in_dae(&dae, &subscript), Some(3)); +} + #[test] fn test_eliminate_trivial_keeps_sampled_value_source_unknown() { let mut dae = Dae::new(); diff --git a/crates/rumoca-phase-structural/src/eliminate/tests/boundary_extra.rs b/crates/rumoca-phase-structural/src/eliminate/tests/boundary_extra.rs index b5e8aad7f..35d43bbb6 100644 --- a/crates/rumoca-phase-structural/src/eliminate/tests/boundary_extra.rs +++ b/crates/rumoca-phase-structural/src/eliminate/tests/boundary_extra.rs @@ -170,6 +170,444 @@ fn test_boundary_preserves_indexed_array_connection_constraints() { ); } +#[test] +fn test_boundary_eliminates_indexed_scalar_algebraic_connection_alias() { + let mut dae = Dae::new(); + + dae.variables.algebraics.insert( + VarName::new("plug.pin[1].v.re"), + component_var("plug.pin[1].v.re"), + ); + dae.variables.algebraics.insert( + VarName::new("adapter.pin[1].v.re"), + component_var("adapter.pin[1].v.re"), + ); + + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: sub_op(), + lhs: Box::new(var_ref("plug.pin[1].v.re")), + rhs: Box::new(var_ref("adapter.pin[1].v.re")), + span: rumoca_core::Span::DUMMY, + }, + span: Span::DUMMY, + origin: "connection equation: plug.pin[1].v.re = adapter.pin[1].v.re".to_string(), + scalar_count: 1, + }); + + let result = eliminate_trivial(&mut dae).expect("structural elimination should succeed"); + + assert_eq!(result.n_eliminated, 1); + assert_eq!(dae.continuous.equations.len(), 0); + assert!( + result + .substitutions + .iter() + .any(|sub| sub.var_name.as_str() == "plug.pin[1].v.re" + || sub.var_name.as_str() == "adapter.pin[1].v.re"), + "indexed scalar algebraic connection alias should produce a substitution" + ); +} + +#[test] +fn test_boundary_eliminates_nested_index_field_connection_alias() { + let mut dae = Dae::new(); + + dae.variables.algebraics.insert( + VarName::new("plug.pin[1].v.re"), + component_var("plug.pin[1].v.re"), + ); + dae.variables.algebraics.insert( + VarName::new("adapter.pin[1].v.re"), + component_var("adapter.pin[1].v.re"), + ); + + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: sub_op(), + lhs: Box::new(field_access( + field_access(index_access(field_access(var_ref("plug"), "pin"), 1), "v"), + "re", + )), + rhs: Box::new(field_access( + field_access( + index_access(field_access(var_ref("adapter"), "pin"), 1), + "v", + ), + "re", + )), + span: rumoca_core::Span::DUMMY, + }, + span: Span::DUMMY, + origin: "connection equation: plug.pin[1].v.re = adapter.pin[1].v.re".to_string(), + scalar_count: 1, + }); + + let result = eliminate_trivial(&mut dae).expect("structural elimination should succeed"); + + assert_eq!(result.n_eliminated, 1); + assert_eq!(dae.continuous.equations.len(), 0); + assert!( + result + .substitutions + .iter() + .any(|sub| sub.var_name.as_str() == "plug.pin[1].v.re" + || sub.var_name.as_str() == "adapter.pin[1].v.re"), + "nested index/field connection alias should produce a scalar substitution" + ); +} + +#[test] +fn test_orphan_drop_does_not_keep_scalarized_unknown_by_base_alias_only() { + let mut dae = Dae::new(); + + dae.variables.algebraics.insert( + VarName::new("resistor.plug_p.pin[2].v.im"), + test_dae_variable("resistor.plug_p.pin[2].v.im"), + ); + dae.variables.inputs.insert( + VarName::new("resistor.plug_p.pin"), + test_dae_variable("resistor.plug_p.pin"), + ); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: var_ref("resistor.plug_p.pin"), + span: Span::DUMMY, + origin: "metadata-only aggregate reference".to_string(), + scalar_count: 1, + }); + + drop_unreferenced_continuous_unknowns(&mut dae); + + assert!( + !dae.variables + .algebraics + .contains_key(&VarName::new("resistor.plug_p.pin[2].v.im")), + "scalarized algebraic unknowns require an exact live reference" + ); +} + +#[test] +fn test_orphan_drop_keeps_exact_scalarized_lhs_owner() { + let mut dae = Dae::new(); + + dae.variables.algebraics.insert( + VarName::new("resistor.plug_p.pin[2].v.im"), + test_dae_variable("resistor.plug_p.pin[2].v.im"), + ); + dae.continuous + .equations + .push(dae::Equation::explicit_with_scalar_count( + VarName::new("resistor.plug_p.pin[2].v.im"), + lit(0.0), + Span::DUMMY, + "exact scalarized lhs", + 1, + )); + + drop_unreferenced_continuous_unknowns(&mut dae); + + let sorted = crate::sort_dae(&dae) + .expect("the retained explicit scalarized lhs must remain structurally matchable"); + assert!( + dae.variables + .algebraics + .contains_key(&VarName::new("resistor.plug_p.pin[2].v.im")), + "an exact scalarized lhs must keep its owning unknown live" + ); + assert_eq!(dae.continuous.equations.len(), 1); + assert_eq!( + sorted.matching.len(), + 1, + "the retained equation must match its one exact scalarized unknown" + ); +} + +#[test] +fn test_orphan_drop_keeps_exact_scalarized_slice_lhs_owners() { + let mut dae = Dae::new(); + let span = test_span(); + + let mut aggregate_metadata = component_var("leg_force_w"); + aggregate_metadata.dims = vec![3, 2]; + dae.variables + .inputs + .insert(VarName::new("leg_force_w"), aggregate_metadata); + for row in 1..=3 { + let name = format!("leg_force_w[{row},1]"); + dae.variables + .algebraics + .insert(VarName::new(&name), component_var(&name)); + } + dae.variables.algebraics.insert( + VarName::new("leg_force_w[1,2]"), + component_var("leg_force_w[1,2]"), + ); + + let lhs = Reference::with_component_reference( + "leg_force_w", + rumoca_core::ComponentReference { + local: false, + span, + parts: vec![rumoca_core::ComponentRefPart { + ident: "leg_force_w".to_string(), + span, + subs: vec![ + rumoca_core::Subscript::Colon { span }, + rumoca_core::Subscript::Index { value: 1, span }, + ], + }], + def_id: None, + }, + ); + dae.continuous + .equations + .push(dae::Equation::explicit_with_scalar_count( + lhs, + array(vec![lit(0.0), lit(0.0), lit(0.0)]), + span, + "three-row slice lhs", + 3, + )); + + drop_unreferenced_continuous_unknowns(&mut dae); + + for row in 1..=3 { + let name = VarName::new(format!("leg_force_w[{row},1]")); + assert!( + dae.variables.algebraics.contains_key(&name), + "slice lhs must keep exact owner `{}` live", + name.as_str() + ); + } + assert!( + !dae.variables + .algebraics + .contains_key(&VarName::new("leg_force_w[1,2]")), + "slice lhs must not keep an unrelated scalar leaf by base alias" + ); + let mut mismatched = dae.clone(); + mismatched.continuous.equations[0].scalar_count = 2; + drop_unreferenced_continuous_unknowns(&mut mismatched); + assert!( + mismatched.variables.algebraics.is_empty(), + "a slice whose DAE shape disagrees with scalar_count must fail closed" + ); + let resolver = crate::incidence::ScalarUnknownResolver::from_entries( + (1..=3).map(|row| (format!("leg_force_w[{row},1]"), row - 1)), + ); + let mut lhs_columns = std::collections::HashSet::new(); + crate::incidence::collect_equation_lhs_unknown( + dae.continuous.equations[0].lhs.as_ref(), + &resolver, + &mut lhs_columns, + ); + assert_eq!( + lhs_columns.len(), + 3, + "the retained slice lhs must expose all three exact owners to structural incidence" + ); +} + +#[test] +fn test_orphan_drop_keeps_structured_fixed_singleton_lhs_owner() { + let mut dae = Dae::new(); + let span = test_span(); + + let mut aggregate_metadata = component_var("fixed_target"); + aggregate_metadata.dims = vec![2, 2]; + dae.variables + .inputs + .insert(VarName::new("fixed_target"), aggregate_metadata); + for name in ["fixed_target[1,1]", "fixed_target[2,1]"] { + dae.variables + .algebraics + .insert(VarName::new(name), component_var(name)); + } + + let lhs = Reference::with_component_reference( + "fixed_target", + rumoca_core::ComponentReference { + local: false, + span, + parts: vec![rumoca_core::ComponentRefPart { + ident: "fixed_target".to_string(), + span, + subs: vec![ + rumoca_core::Subscript::Index { value: 1, span }, + rumoca_core::Subscript::Index { value: 1, span }, + ], + }], + def_id: None, + }, + ); + dae.continuous + .equations + .push(dae::Equation::explicit_with_scalar_count( + lhs, + lit(0.0), + span, + "one exact structured lhs owner", + 1, + )); + + drop_unreferenced_continuous_unknowns(&mut dae); + + assert!( + dae.variables + .algebraics + .contains_key(&VarName::new("fixed_target[1,1]")), + "a structured fixed lhs with scalar_count=1 must keep its exact leaf owner" + ); + assert!( + !dae.variables + .algebraics + .contains_key(&VarName::new("fixed_target[2,1]")), + "a structured fixed lhs must not keep an unrelated leaf" + ); +} + +#[test] +fn test_orphan_drop_keeps_exact_scalarized_unknown_reference() { + let mut dae = Dae::new(); + + dae.variables.algebraics.insert( + VarName::new("resistor.plug_p.pin[2].v.im"), + test_dae_variable("resistor.plug_p.pin[2].v.im"), + ); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: var_ref("resistor.plug_p.pin[2].v.im"), + span: Span::DUMMY, + origin: "exact scalarized reference".to_string(), + scalar_count: 1, + }); + + drop_unreferenced_continuous_unknowns(&mut dae); + + assert!( + dae.variables + .algebraics + .contains_key(&VarName::new("resistor.plug_p.pin[2].v.im")), + "exact scalarized algebraic references must keep the unknown live" + ); +} + +#[test] +fn test_boundary_eliminates_single_live_indexed_flow_alias() { + let mut dae = Dae::new(); + + dae.variables.algebraics.insert( + VarName::new("star.plugToPin[2].pin_p.i.im"), + test_dae_variable("star.plugToPin[2].pin_p.i.im"), + ); + dae.variables.parameters.insert( + VarName::new("star.pin_p[2].i.im"), + test_dae_variable("star.pin_p[2].i.im"), + ); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: OpBinary::Add, + lhs: Box::new(var_ref("star.plugToPin[2].pin_p.i.im")), + rhs: Box::new(Expression::Unary { + op: OpUnary::Minus, + rhs: Box::new(var_ref("star.pin_p[2].i.im")), + span: Span::DUMMY, + }), + span: Span::DUMMY, + }, + span: Span::DUMMY, + origin: "flow sum equation: star.plugToPin[2].pin_p.i.im + -star.pin_p[2].i.im = 0" + .to_string(), + scalar_count: 1, + }); + + let result = eliminate_trivial(&mut dae).expect("structural elimination should succeed"); + + assert_eq!(result.n_eliminated, 1); + assert!( + result + .substitutions + .iter() + .any(|sub| sub.var_name.as_str() == "star.plugToPin[2].pin_p.i.im"), + "single-live indexed flow aliases should be eliminated" + ); +} + +#[test] +fn test_boundary_eliminates_pairwise_indexed_flow_alias() { + let mut dae = Dae::new(); + + for name in [ + "star.plugToPin[2].pin_p.i.im", + "star.pin_p[2].i.im", + "star.pin_n.i.im", + ] { + dae.variables + .algebraics + .insert(VarName::new(name), test_dae_variable(name)); + } + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: OpBinary::Add, + lhs: Box::new(var_ref("star.plugToPin[2].pin_p.i.im")), + rhs: Box::new(Expression::Unary { + op: OpUnary::Minus, + rhs: Box::new(var_ref("star.pin_p[2].i.im")), + span: Span::DUMMY, + }), + span: Span::DUMMY, + }, + span: Span::DUMMY, + origin: "flow sum equation: star.plugToPin[2].pin_p.i.im + -star.pin_p[2].i.im = 0" + .to_string(), + scalar_count: 1, + }); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: OpBinary::Add, + lhs: Box::new(var_ref("star.pin_p[2].i.im")), + rhs: Box::new(var_ref("star.pin_n.i.im")), + span: Span::DUMMY, + }, + span: Span::DUMMY, + origin: "flow sum equation: star.pin_p[2].i.im + star.pin_n.i.im = 0".to_string(), + scalar_count: 1, + }); + + let all_unknowns = collect_boundary_unknowns(&dae).expect("boundary unknown collection works"); + let unknown_index = + BoundaryUnknownIndex::build(&dae, &all_unknowns).expect("boundary index builds"); + let live = find_live_scalar_unknowns( + &dae.continuous.equations[0].rhs, + &unknown_index, + &std::collections::HashSet::new(), + ) + .expect("live unknown scan works"); + assert_eq!( + live, + vec![ + VarName::new("star.plugToPin[2].pin_p.i.im"), + VarName::new("star.pin_p[2].i.im"), + ], + "pairwise flow alias must expose only its two scalar current unknowns" + ); + + let result = eliminate_trivial(&mut dae).expect("structural elimination should succeed"); + + assert!( + result.substitutions.iter().any(|sub| sub.var_name.as_str() + == "star.plugToPin[2].pin_p.i.im" + || sub.var_name.as_str() == "star.pin_p[2].i.im"), + "pairwise indexed flow aliases should be eliminated before KCL rows" + ); +} + #[test] fn test_boundary_keeps_internal_discrete_connection_chain_for_runtime_alias_paths() { let mut dae = Dae::new(); @@ -268,3 +706,186 @@ fn test_boundary_keeps_internal_discrete_connection_chain_for_runtime_alias_path "internal discrete connector aliases must remain live after boundary elimination" ); } + +fn index_access(base: Expression, idx: i64) -> Expression { + Expression::Index { + base: Box::new(base), + subscripts: vec![rumoca_core::Subscript::generated_index(idx, test_span())], + span: test_span(), + } +} + +fn field_access(base: Expression, field: &str) -> Expression { + Expression::FieldAccess { + base: Box::new(base), + field: field.to_string(), + span: test_span(), + } +} + +fn boundary_alias_fixture() -> Dae { + let mut dae = Dae::new(); + for name in ["alias", "preexisting_orphan"] { + dae.variables + .algebraics + .insert(VarName::new(name), test_dae_variable(name)); + } + dae.variables.parameters.insert( + VarName::new("external_source"), + test_dae_variable("external_source"), + ); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: sub_op(), + lhs: Box::new(var_ref("external_source")), + rhs: Box::new(var_ref("alias")), + span: Span::DUMMY, + }, + span: Span::DUMMY, + origin: "alias definition".to_string(), + scalar_count: 1, + }); + dae +} + +fn runtime_equation(origin: &str, rhs: Expression) -> dae::Equation { + dae::Equation { + lhs: None, + rhs, + span: Span::DUMMY, + origin: origin.to_string(), + scalar_count: 1, + } +} + +fn builtin(function: BuiltinFunction, arg: Expression) -> Expression { + Expression::BuiltinCall { + function, + args: vec![arg], + span: Span::DUMMY, + } +} + +fn event_message(action: &dae::DaeEventAction) -> &Expression { + match &action.kind { + dae::DaeEventActionKind::Assert { message } + | dae::DaeEventActionKind::Terminate { message } => message, + } +} + +#[test] +fn boundary_fixpoint_rewrites_every_plain_dae_surface_before_retiring_target() { + let mut dae = boundary_alias_fixture(); + dae.initialization + .equations + .push(runtime_equation("initial", var_ref("alias"))); + dae.events.synthetic_root_conditions.push(var_ref("alias")); + dae.clocks.constructor_exprs.push(var_ref("alias")); + dae.clocks.triggered_conditions.push(var_ref("alias")); + for kind in [ + dae::DaeEventActionKind::Assert { + message: var_ref("alias"), + }, + dae::DaeEventActionKind::Terminate { + message: var_ref("alias"), + }, + ] { + dae.events.event_actions.push(dae::DaeEventAction { + condition: var_ref("alias"), + kind, + span: Span::DUMMY, + origin: "event action".to_string(), + }); + } + + let result = resolve_boundary_equations_to_fixpoint(&mut dae).unwrap(); + let alias = VarName::new("alias"); + + assert!( + result.substitutions.iter().any(|sub| sub.var_name == alias), + "substitutions: {:?}", + result + .substitutions + .iter() + .map(|sub| sub.var_name.as_str()) + .collect::>() + ); + assert!(!dae.variables.algebraics.contains_key(&alias)); + assert!( + dae.variables + .algebraics + .contains_key(&VarName::new("preexisting_orphan")) + ); + assert!( + dae.initialization + .equations + .iter() + .all(|eq| !expr_contains_var(&eq.rhs, &alias)) + ); + assert!( + dae.events + .synthetic_root_conditions + .iter() + .all(|expr| !expr_contains_var(expr, &alias)) + ); + assert!( + dae.clocks + .constructor_exprs + .iter() + .all(|expr| !expr_contains_var(expr, &alias)) + ); + assert!( + dae.clocks + .triggered_conditions + .iter() + .all(|expr| !expr_contains_var(expr, &alias)) + ); + assert!(dae.events.event_actions.iter().all(|action| { + !expr_contains_var(&action.condition, &alias) + && !expr_contains_var(event_message(action), &alias) + })); +} + +#[test] +fn boundary_fixpoint_keeps_target_referenced_by_pre_edge_change_surfaces() { + let mut dae = boundary_alias_fixture(); + dae.initialization.equations.push(runtime_equation( + "initial pre", + builtin(BuiltinFunction::Pre, var_ref("alias")), + )); + dae.clocks + .triggered_conditions + .push(builtin(BuiltinFunction::Edge, var_ref("alias"))); + dae.events.event_actions.push(dae::DaeEventAction { + condition: lit(1.0), + kind: dae::DaeEventActionKind::Terminate { + message: builtin(BuiltinFunction::Change, var_ref("alias")), + }, + span: Span::DUMMY, + origin: "terminate".to_string(), + }); + + let result = resolve_boundary_equations_to_fixpoint(&mut dae).unwrap(); + let alias = VarName::new("alias"); + + assert!(result.substitutions.iter().any(|sub| sub.var_name == alias)); + assert!(dae.variables.algebraics.contains_key(&alias)); + assert!(expr_contains_var( + &dae.initialization.equations[0].rhs, + &alias + )); + assert!(expr_contains_var( + &dae.clocks.triggered_conditions[0], + &alias + )); + assert!(expr_contains_var( + event_message(&dae.events.event_actions[0]), + &alias + )); + assert!( + dae.variables + .algebraics + .contains_key(&VarName::new("preexisting_orphan")) + ); +} diff --git a/crates/rumoca-phase-structural/src/eliminate/tests/substitution_more.rs b/crates/rumoca-phase-structural/src/eliminate/tests/substitution_more.rs index 76128d0f2..37dadc8ed 100644 --- a/crates/rumoca-phase-structural/src/eliminate/tests/substitution_more.rs +++ b/crates/rumoca-phase-structural/src/eliminate/tests/substitution_more.rs @@ -1,5 +1,294 @@ use super::*; +#[test] +fn test_substitute_indexed_use_site_keeps_projected_scalar_replacement() { + let mut dae = Dae::new(); + let mut u = test_dae_variable("block.u"); + u.dims = vec![1]; + dae.variables.algebraics.insert(VarName::new("block.u"), u); + let mut y = test_dae_variable("block.y"); + y.dims = vec![3]; + dae.variables.outputs.insert(VarName::new("block.y"), y); + + let substitution = + substitution_for_var(&dae, VarName::new("block.u"), var_ref_idx("block.y", 2)) + .expect("projected vector output should be a valid scalar replacement"); + + let rewritten = resolve_substitutions_in_expr(&var_ref_idx("block.u", 1), &[substitution]) + .expect("substitution should rewrite indexed use site"); + + let Expression::VarRef { + name, subscripts, .. + } = rewritten + else { + panic!("rewritten expression should remain a VarRef"); + }; + assert_eq!(name.as_str(), "block.y"); + assert_eq!(subscripts.len(), 1, "replacement must not become y[2,1]"); + assert!( + matches!( + &subscripts[0], + rumoca_core::Subscript::Index { value: 2, .. } + ), + "replacement should preserve the projected physical element" + ); +} + +#[test] +fn test_substitute_indexed_use_site_keeps_index_expression_scalar_replacement() { + let substitution = Substitution { + var_name: VarName::new("block.u"), + var_ref: None, + expr: Expression::Index { + base: Box::new(var_ref("block.y")), + subscripts: vec![rumoca_core::Subscript::Index { + value: 2, + span: Span::DUMMY, + }], + span: Span::DUMMY, + }, + var_dims: vec![1], + replacement_dims: Vec::new(), + env_keys: vec!["block.u".to_string()], + }; + + let rewritten = resolve_substitutions_in_expr(&var_ref_idx("block.u", 1), &[substitution]) + .expect("substitution should rewrite indexed use site"); + + let Expression::Index { + base, subscripts, .. + } = rewritten + else { + panic!("rewritten expression should remain the projected scalar Index"); + }; + let Expression::VarRef { + name, + subscripts: base_subscripts, + .. + } = base.as_ref() + else { + panic!("index base should remain a VarRef"); + }; + assert_eq!(name.as_str(), "block.y"); + assert!(base_subscripts.is_empty()); + assert_eq!(subscripts.len(), 1, "replacement must not become y[2,1]"); + assert!(matches!( + &subscripts[0], + rumoca_core::Subscript::Index { value: 2, .. } + )); +} + +#[test] +fn test_eliminate_trivial_rejects_multiscalar_array_solution_for_scalar_target() { + let mut dae = Dae::new(); + let span = Span::from_offsets( + rumoca_core::SourceId::from_source_name("scalar_array_solution.mo"), + 1, + 12, + ); + let mut y = test_dae_variable("y"); + y.dims = vec![3]; + dae.variables.outputs.insert(VarName::new("y"), y); + for name in ["a", "b", "c"] { + dae.variables + .parameters + .insert(VarName::new(name), test_dae_variable(name)); + } + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: sub_op(), + lhs: Box::new(var_ref_idx("y", 2)), + rhs: Box::new(Expression::Array { + elements: vec![var_ref("a"), var_ref("b"), var_ref("c")], + is_matrix: false, + span, + }), + span, + }, + span, + origin: "scalar target cannot equal vector".to_string(), + scalar_count: 1, + }); + + let result = eliminate_trivial(&mut dae).expect("elimination should reject invalid solution"); + + assert_eq!(result.n_eliminated, 0); + assert_eq!(dae.continuous.equations.len(), 1); + assert!(dae.variables.outputs.contains_key(&VarName::new("y"))); +} + +#[test] +fn test_eliminate_trivial_resolves_scalarized_internal_output_element_alias() { + let mut dae = Dae::new(); + let mut y = test_dae_variable("block.multiplex.y"); + y.dims = vec![3]; + dae.variables + .outputs + .insert(VarName::new("block.multiplex.y"), y); + dae.variables.algebraics.insert( + VarName::new("block.multiplex.u3"), + test_dae_variable("block.multiplex.u3"), + ); + dae.variables + .states + .insert(VarName::new("body.w"), test_dae_variable("body.w")); + dae.variables + .states + .insert(VarName::new("body.phi"), test_dae_variable("body.phi")); + + dae.continuous.equations.push(residual( + var_ref_idx("block.multiplex.y", 3), + var_ref("block.multiplex.u3"), + 1, + "equation from block.multiplex", + )); + dae.continuous.equations.push(residual( + var_ref("body.phi"), + Expression::FunctionCall { + name: reference("Move.position"), + args: vec![ + array(vec![ + var_ref("body.phi"), + var_ref("body.w"), + var_ref_idx("block.multiplex.y", 3), + ]), + var_ref("time"), + ], + is_constructor: false, + span: Span::DUMMY, + }, + 1, + "equation from block.move", + )); + dae.continuous.equations.push(residual( + der(var_ref("body.w")), + var_ref_idx("block.multiplex.y", 3), + 1, + "connection equation: body.a = block.multiplex.u3[1]", + )); + + let result = eliminate_trivial(&mut dae) + .expect("scalarized internal output element should be eligible for substitution"); + + assert!( + result + .substitutions + .iter() + .any(|sub| sub.var_name.as_str() == "block.multiplex.y[3]"), + "expected scalarized output element substitution, got {:?}", + result.substitutions + ); + assert!( + dae.continuous + .equations + .iter() + .all(|eq| !expr_contains_var(&eq.rhs, &VarName::new("block.multiplex.y[3]"))), + "remaining equations should not retain the eliminated scalarized output element" + ); +} + +#[test] +fn test_eliminate_trivial_resolves_internal_output_torque_definition() { + let mut dae = Dae::new(); + dae.variables.outputs.insert( + VarName::new("adapter.tau2"), + test_dae_variable("adapter.tau2"), + ); + for name in [ + "spring.c", + "spring.phi_rel", + "spring.w_rel", + "body.J", + "body.a", + ] { + dae.variables + .algebraics + .insert(VarName::new(name), test_dae_variable(name)); + } + + let spring_torque = binary( + OpBinary::Add, + binary( + OpBinary::Mul, + var_ref("spring.c"), + var_ref("spring.phi_rel"), + ), + binary(OpBinary::Mul, real(100.0), var_ref("spring.w_rel")), + ); + dae.continuous.equations.push(residual( + var_ref("adapter.tau2"), + spring_torque, + 1, + "equation from adapter.torqueSensor", + )); + dae.continuous.equations.push(residual( + binary(OpBinary::Mul, var_ref("body.J"), var_ref("body.a")), + Expression::Unary { + op: OpUnary::Minus, + rhs: Box::new(var_ref("adapter.tau2")), + span: Span::DUMMY, + }, + 1, + "equation from body", + )); + + let result = + eliminate_trivial(&mut dae).expect("internal output torque should be substitutable"); + + assert!( + result + .substitutions + .iter() + .any(|sub| sub.var_name.as_str() == "adapter.tau2"), + "expected adapter.tau2 substitution, got {:?}", + result.substitutions + ); + assert!( + dae.continuous + .equations + .iter() + .all(|eq| !expr_contains_var(&eq.rhs, &VarName::new("adapter.tau2"))), + "remaining equations should not retain adapter.tau2" + ); +} + +#[test] +fn test_scalar_blt_solution_rejects_multiscalar_array_solution_for_scalar_target() { + let mut dae = Dae::new(); + dae.variables + .algebraics + .insert(VarName::new("y"), test_dae_variable("y")); + for name in ["a", "b", "c"] { + dae.variables + .parameters + .insert(VarName::new(name), test_dae_variable(name)); + } + dae.continuous.equations.push(residual( + var_ref("y"), + array(vec![var_ref("a"), var_ref("b"), var_ref("c")]), + 1, + "scalar target cannot equal vector through BLT", + )); + + let state_names = Vec::::new(); + let state_derivative_matcher = DerivativeNameMatcher::from_var_names(state_names.iter()); + let result = structural_ok(scalar_blt_solution( + &dae, + &EquationRef(0), + &UnknownId::Variable(VarName::new("y")), + &IndexSet::new(), + &HashSet::new(), + &state_derivative_matcher, + &[], + )); + + assert!( + result.is_none(), + "BLT scalar elimination must reject scalar-to-vector substitutions" + ); +} + #[test] fn test_elimination_substitution_differentiates_scalar_alias_in_derivative_call() { let mut dae = Dae::new(); @@ -1145,6 +1434,201 @@ fn test_eliminate_trivial_skips_substitution_to_unsliced_multiscalar_solution() ); } +#[test] +fn test_scalarized_substitution_projects_array_comprehension_solution() { + let mut dae = Dae::new(); + dae.variables.algebraics.insert( + VarName::new("ductOut.fluidVolumes[1]"), + test_dae_variable("ductOut.fluidVolumes[1]"), + ); + let span = test_span(); + let index_ref = || Expression::VarRef { + name: reference("i"), + subscripts: Vec::new(), + span, + }; + let indexed_ref = |name: &str| Expression::VarRef { + name: reference(name), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(index_ref()), + span, + }], + span, + }; + let comprehension = Expression::ArrayComprehension { + expr: Box::new(binary( + mul_op(), + indexed_ref("ductOut.crossAreas"), + indexed_ref("ductOut.lengths"), + )), + indices: vec![rumoca_core::ComprehensionIndex { + name: "i".to_string(), + range: Expression::Range { + start: Box::new(Expression::Literal { + value: Literal::Integer(1), + span, + }), + step: None, + end: Box::new(Expression::Literal { + value: Literal::Integer(2), + span, + }), + span, + }, + }], + filter: None, + span, + }; + let aggregate_solution = binary(mul_op(), comprehension, lit(1.0)); + + let substitution = substitution_for_var( + &dae, + VarName::new("ductOut.fluidVolumes[1]"), + aggregate_solution, + ) + .expect("scalarized substitution should be constructed"); + + assert!( + !contains_array_comprehension_expr(&substitution.expr), + "scalarized substitution should project the aggregate RHS: {:?}", + substitution.expr + ); + assert!( + contains_indexed_literal_ref(&substitution.expr, "ductOut.crossAreas", 1) + && contains_indexed_literal_ref(&substitution.expr, "ductOut.lengths", 1), + "projected substitution should select the first physical duct volume element: {:?}", + substitution.expr + ); +} + +fn contains_array_comprehension_expr(expr: &Expression) -> bool { + match expr { + Expression::ArrayComprehension { .. } => true, + Expression::Binary { lhs, rhs, .. } => { + contains_array_comprehension_expr(lhs) || contains_array_comprehension_expr(rhs) + } + Expression::Unary { rhs, .. } => contains_array_comprehension_expr(rhs), + Expression::BuiltinCall { args, .. } | Expression::FunctionCall { args, .. } => { + args.iter().any(contains_array_comprehension_expr) + } + Expression::If { + branches, + else_branch, + .. + } => { + branches.iter().any(|(condition, value)| { + contains_array_comprehension_expr(condition) + || contains_array_comprehension_expr(value) + }) || contains_array_comprehension_expr(else_branch) + } + Expression::Array { elements, .. } | Expression::Tuple { elements, .. } => { + elements.iter().any(contains_array_comprehension_expr) + } + Expression::Range { + start, step, end, .. + } => { + contains_array_comprehension_expr(start) + || step + .as_ref() + .is_some_and(|step| contains_array_comprehension_expr(step)) + || contains_array_comprehension_expr(end) + } + Expression::Index { + base, subscripts, .. + } => { + contains_array_comprehension_expr(base) + || subscripts.iter().any(|subscript| match subscript { + rumoca_core::Subscript::Expr { expr, .. } => { + contains_array_comprehension_expr(expr) + } + _ => false, + }) + } + Expression::FieldAccess { base, .. } => contains_array_comprehension_expr(base), + Expression::VarRef { .. } | Expression::Literal { .. } | Expression::Empty { .. } => false, + } +} + +fn contains_indexed_literal_ref(expr: &Expression, needle: &str, index: i64) -> bool { + match expr { + Expression::VarRef { + name, subscripts, .. + } => { + name.as_str() == needle + && matches!( + subscripts.as_slice(), + [rumoca_core::Subscript::Expr { + expr, + .. + }] if matches!( + expr.as_ref(), + Expression::Literal { + value: Literal::Integer(value), + .. + } if *value == index + ) + ) + } + Expression::Binary { lhs, rhs, .. } => { + contains_indexed_literal_ref(lhs, needle, index) + || contains_indexed_literal_ref(rhs, needle, index) + } + Expression::Unary { rhs, .. } => contains_indexed_literal_ref(rhs, needle, index), + Expression::BuiltinCall { args, .. } | Expression::FunctionCall { args, .. } => args + .iter() + .any(|arg| contains_indexed_literal_ref(arg, needle, index)), + Expression::If { + branches, + else_branch, + .. + } => { + branches.iter().any(|(condition, value)| { + contains_indexed_literal_ref(condition, needle, index) + || contains_indexed_literal_ref(value, needle, index) + }) || contains_indexed_literal_ref(else_branch, needle, index) + } + Expression::Array { elements, .. } | Expression::Tuple { elements, .. } => elements + .iter() + .any(|element| contains_indexed_literal_ref(element, needle, index)), + Expression::Range { + start, step, end, .. + } => { + contains_indexed_literal_ref(start, needle, index) + || step + .as_ref() + .is_some_and(|step| contains_indexed_literal_ref(step, needle, index)) + || contains_indexed_literal_ref(end, needle, index) + } + Expression::ArrayComprehension { + expr, + indices, + filter, + .. + } => { + contains_indexed_literal_ref(expr, needle, index) + || indices + .iter() + .any(|index_def| contains_indexed_literal_ref(&index_def.range, needle, index)) + || filter + .as_ref() + .is_some_and(|filter| contains_indexed_literal_ref(filter, needle, index)) + } + Expression::Index { + base, subscripts, .. + } => { + contains_indexed_literal_ref(base, needle, index) + || subscripts.iter().any(|subscript| match subscript { + rumoca_core::Subscript::Expr { expr, .. } => { + contains_indexed_literal_ref(expr, needle, index) + } + _ => false, + }) + } + Expression::FieldAccess { base, .. } => contains_indexed_literal_ref(base, needle, index), + Expression::Literal { .. } | Expression::Empty { .. } => false, + } +} + #[test] fn test_eliminate_trivial_rewrites_eliminated_indexed_record_field_aggregate() { let expr = Expression::BuiltinCall { @@ -1210,9 +1694,15 @@ fn test_eliminate_trivial_rewrites_eliminated_complex_field_parent_ref() { ); } -#[test] -fn test_apply_elimination_substitutions_rewrites_dae_runtime_partitions() { +fn dae_with_alias_in_runtime_partitions() -> Dae { let mut dae = Dae::new(); + dae.initialization.equations.push(dae::Equation { + lhs: None, + rhs: var_ref("alias"), + span: Span::DUMMY, + origin: "initial".to_string(), + scalar_count: 1, + }); dae.discrete.real_updates.push(dae::Equation { lhs: Some(VarName::new("z").into()), rhs: var_ref("alias"), @@ -1220,6 +1710,13 @@ fn test_apply_elimination_substitutions_rewrites_dae_runtime_partitions() { origin: "f_z".to_string(), scalar_count: 1, }); + dae.discrete.valued_updates.push(dae::Equation { + lhs: Some(VarName::new("m").into()), + rhs: var_ref("alias"), + span: Span::DUMMY, + origin: "f_m".to_string(), + scalar_count: 1, + }); dae.conditions.equations.push(dae::Equation { lhs: Some(VarName::new("c").into()), rhs: var_ref("alias"), @@ -1234,15 +1731,25 @@ fn test_apply_elimination_substitutions_rewrites_dae_runtime_partitions() { dae.events.event_actions.push(dae::DaeEventAction { condition: var_ref("alias"), kind: dae::DaeEventActionKind::Terminate { - message: rumoca_core::Expression::Literal { - value: rumoca_core::Literal::String("stop".to_string()), - span: Span::DUMMY, - }, + message: var_ref("alias"), + }, + span: Span::DUMMY, + origin: "assert".to_string(), + }); + dae.events.event_actions.push(dae::DaeEventAction { + condition: var_ref("alias"), + kind: dae::DaeEventActionKind::Assert { + message: var_ref("alias"), }, span: Span::DUMMY, origin: "assert".to_string(), }); + dae +} +#[test] +fn test_apply_elimination_substitutions_rewrites_dae_runtime_partitions() { + let mut dae = dae_with_alias_in_runtime_partitions(); let substitutions = [Substitution { var_name: VarName::new("alias"), var_ref: Some(reference("alias")), @@ -1256,10 +1763,18 @@ fn test_apply_elimination_substitutions_rewrites_dae_runtime_partitions() { &substitutions, )); + assert!(contains_exact_var_ref( + &dae.initialization.equations[0].rhs, + "source" + )); assert!(contains_exact_var_ref( &dae.discrete.real_updates[0].rhs, "source" )); + assert!(contains_exact_var_ref( + &dae.discrete.valued_updates[0].rhs, + "source" + )); assert!(contains_exact_var_ref( &dae.conditions.equations[0].rhs, "source" @@ -1280,10 +1795,14 @@ fn test_apply_elimination_substitutions_rewrites_dae_runtime_partitions() { &dae.clocks.triggered_conditions[0], "source" )); - assert!(contains_exact_var_ref( - &dae.events.event_actions[0].condition, - "source" - )); + assert!(dae.events.event_actions.iter().all(|action| { + let message = match &action.kind { + dae::DaeEventActionKind::Assert { message } + | dae::DaeEventActionKind::Terminate { message } => message, + }; + contains_exact_var_ref(&action.condition, "source") + && contains_exact_var_ref(message, "source") + })); } #[test] diff --git a/crates/rumoca-phase-structural/src/eliminate/unknown_index.rs b/crates/rumoca-phase-structural/src/eliminate/unknown_index.rs index 51ccd5168..6dbd311a7 100644 --- a/crates/rumoca-phase-structural/src/eliminate/unknown_index.rs +++ b/crates/rumoca-phase-structural/src/eliminate/unknown_index.rs @@ -10,12 +10,16 @@ use std::collections::HashSet; use indexmap::IndexMap; +use rumoca_ir_dae as dae; use rumoca_ir_dae::{ component_base_name, parse_embedded_subscripts, split_complex_field_suffix, var_ref_matches_unknown, }; -use super::{Dae, Expression, Reference, VarName, collect_var_ref_nodes}; +use super::{ + Dae, Expression, Reference, VarName, collect_exact_reference_expr_names_in_dae, + collect_var_ref_nodes, exact_subscript_index_in_dae, +}; use crate::StructuralError; use crate::variable_scope::DaeVariableScope; @@ -162,7 +166,7 @@ impl<'a> BoundaryUnknownIndex<'a> { if let Some(bucket) = self.all_one_embedded.get(base) { out.extend_from_slice(bucket); } - } else if let Some(indices) = literal_subscript_indices(subscripts) + } else if let Some(indices) = dae_subscript_indices(self.dae, subscripts) && let Some(bucket) = self.by_base_and_trailing.get(&(base.to_string(), indices)) { out.extend_from_slice(bucket); @@ -186,16 +190,25 @@ impl<'a> BoundaryUnknownIndex<'a> { ) -> Result, StructuralError> { let mut var_refs = Vec::new(); collect_var_ref_nodes(expr, &mut var_refs); + let mut exact_names = Vec::new(); + collect_exact_reference_expr_names_in_dae(self.dae, expr, &mut exact_names); let mut candidates = Vec::new(); for (name, subscripts) in &var_refs { self.append_candidate_positions(name, subscripts, &mut candidates); } + for exact_name in &exact_names { + if let Some(&position) = self.by_name.get(exact_name.as_str()) { + candidates.push(position); + } + } candidates.sort_unstable(); candidates.dedup(); let mut live = Vec::new(); for position in candidates { let unknown = &self.all_unknowns[position]; - if !resolved.contains(unknown) && refs_contain_unknown(&var_refs, unknown, self.dae)? { + if !resolved.contains(unknown) + && refs_contain_unknown(&var_refs, &exact_names, unknown, self.dae)? + { live.push(position); } } @@ -214,22 +227,10 @@ fn component_ref_ident_path(component_ref: &rumoca_core::ComponentReference) -> path } -/// Literal integer values of subscripts, mirroring the acceptance rules of -/// `subscripts_match_indices` (plain indices and literal integer expressions). -fn literal_subscript_indices(subscripts: &[rumoca_core::Subscript]) -> Option> { +fn dae_subscript_indices(dae: &Dae, subscripts: &[rumoca_core::Subscript]) -> Option> { subscripts .iter() - .map(|sub| match sub { - rumoca_core::Subscript::Index { value, .. } => Some(*value), - rumoca_core::Subscript::Expr { expr, .. } => match expr.as_ref() { - Expression::Literal { - value: rumoca_core::Literal::Integer(i), - .. - } => Some(*i), - _ => None, - }, - rumoca_core::Subscript::Colon { .. } => None, - }) + .map(|subscript| exact_subscript_index_in_dae(dae, subscript)) .collect() } @@ -270,9 +271,21 @@ pub(super) fn find_live_scalar_unknowns( fn refs_contain_unknown( refs: &[(Reference, Vec)], + exact_names: &[VarName], unknown: &VarName, dae: &Dae, ) -> Result { + if exact_names.iter().any(|name| name == unknown) { + return Ok(true); + } + for (name, subscripts) in refs { + if var_ref_matches_unknown(name, subscripts.as_slice(), unknown) { + return Ok(true); + } + } + if exact_scalar_unknown_exists(dae, unknown)? { + return Ok(false); + } for (name, subscripts) in refs { if var_ref_mentions_unknown_for_presence(name, subscripts.as_slice(), unknown, dae)? { return Ok(true); @@ -281,6 +294,11 @@ fn refs_contain_unknown( Ok(false) } +fn exact_scalar_unknown_exists(dae: &Dae, unknown: &VarName) -> Result { + let scope = DaeVariableScope::new(dae); + Ok(scope.exact(unknown).is_some() && scope.size(unknown)? == 1) +} + fn unknown_scalar_size(dae: &Dae, unknown: &VarName) -> Result { DaeVariableScope::new(dae).size(unknown) } @@ -291,6 +309,18 @@ fn var_ref_mentions_unknown_for_presence( unknown: &VarName, dae: &Dae, ) -> Result { + if name.component_ref().is_some() + && !subscripts.is_empty() + && let Some(indices) = dae_subscript_indices(dae, subscripts) + { + let indices = indices + .into_iter() + .map(|idx| usize::try_from(idx).ok()) + .collect::>>(); + if let Some(indices) = indices { + return Ok(dae::format_subscript_key(name.as_str(), &indices) == unknown.as_str()); + } + } if var_ref_matches_unknown(name, subscripts, unknown) { return Ok(true); } diff --git a/crates/rumoca-phase-structural/src/incidence.rs b/crates/rumoca-phase-structural/src/incidence.rs index 3dbad5dee..553b7b23f 100644 --- a/crates/rumoca-phase-structural/src/incidence.rs +++ b/crates/rumoca-phase-structural/src/incidence.rs @@ -1,8 +1,13 @@ //! Incidence matrix construction for DAE structural analysis. +//! +//! SPEC_0021 file-size exception: incidence currently keeps unknown collection, +//! exact subscript evaluation, and reference traversal in one place. split plan: +//! move exact index/subscript evaluation and its regressions into a dedicated +//! submodule once boundary-elimination and incidence share the same evaluator. use std::collections::{HashMap, HashSet}; -use rumoca_core::ExpressionVisitor; +use rumoca_core::{ExpressionVisitor, Literal}; use rumoca_ir_dae as dae; use crate::types::{EquationRef, UnknownId}; @@ -50,6 +55,7 @@ impl Incidence { pub(crate) fn build_incidence(dae: &dae::Dae) -> Incidence { let (_unknown_map, unknown_names, unknown_spans) = build_unknown_map(dae); let (der_resolver, variable_resolver) = build_unknown_resolvers(&unknown_names); + let constants = ConstantEvalContext::from_dae(dae); let mut equation_refs = Vec::new(); let mut equations = Vec::new(); @@ -63,11 +69,17 @@ pub(crate) fn build_incidence(dae: &dae::Dae) -> Incidence { let mut eq_unknowns: Vec> = equations .iter() - .map(|eq| collect_equation_unknowns(eq, &der_resolver, &variable_resolver)) + .map(|eq| collect_equation_unknowns(eq, &der_resolver, &variable_resolver, &constants)) .collect(); apply_regular_family_corner_incidence(dae, &mut eq_unknowns); debug_check_regular_family_incidence(dae, &eq_unknowns); + let (unknown_names, unknown_spans) = prune_zero_degree_scalarized_aggregate_variables( + dae, + &mut eq_unknowns, + unknown_names, + unknown_spans, + ); Incidence { n_eq, @@ -79,6 +91,70 @@ pub(crate) fn build_incidence(dae: &dae::Dae) -> Incidence { } } +fn prune_zero_degree_scalarized_aggregate_variables( + dae: &dae::Dae, + eq_unknowns: &mut [HashSet], + unknown_names: Vec, + unknown_spans: Vec>, +) -> (Vec, Vec>) { + let mut has_incidence = vec![false; unknown_names.len()]; + for row in eq_unknowns.iter() { + for &idx in row { + if let Some(slot) = has_incidence.get_mut(idx) { + *slot = true; + } + } + } + + let mut remap = vec![None; unknown_names.len()]; + let mut pruned_names = Vec::with_capacity(unknown_names.len()); + let mut pruned_spans = Vec::with_capacity(unknown_spans.len()); + for (old_idx, (unknown, span)) in unknown_names.into_iter().zip(unknown_spans).enumerate() { + if !has_incidence[old_idx] && is_scalarized_aggregate_variable_unknown(dae, &unknown) { + continue; + } + let new_idx = pruned_names.len(); + remap[old_idx] = Some(new_idx); + pruned_names.push(unknown); + pruned_spans.push(span); + } + + for row in eq_unknowns.iter_mut() { + let remapped = row + .iter() + .filter_map(|&idx| remap.get(idx).and_then(|mapped| *mapped)) + .collect(); + *row = remapped; + } + + (pruned_names, pruned_spans) +} + +fn is_scalarized_aggregate_variable_unknown(dae: &dae::Dae, unknown: &UnknownId) -> bool { + let UnknownId::Variable(name) = unknown else { + return false; + }; + let Some(scalar) = rumoca_core::parse_scalar_name(name.as_str()) else { + return false; + }; + let base = rumoca_core::VarName::new(scalar.base); + if dae + .variables + .algebraics + .get(&base) + .or_else(|| dae.variables.outputs.get(&base)) + .is_some_and(|var| var.size() > 1) + { + return true; + } + + dae.variables + .algebraics + .get(name) + .or_else(|| dae.variables.outputs.get(name)) + .is_some() +} + /// Debug-only invariant guarding the P3 family-native lowering contract: every /// `regular` family's materialized incidence equals the incidence /// [`synthesize_regular_family_incidence`] reconstructs from its corner rows alone. @@ -385,6 +461,7 @@ fn collect_equation_unknowns( eq: &dae::Equation, der_resolver: &ScalarUnknownResolver, variable_resolver: &ScalarUnknownResolver, + constants: &ConstantEvalContext, ) -> HashSet { let mut result = HashSet::new(); @@ -398,18 +475,20 @@ fn collect_equation_unknowns( let mut der_collector = DerOperandCollector::default(); der_collector.visit_expression(&eq.rhs); for (name, subscripts) in der_collector.operands { - for idx in der_resolver.resolve_var_ref_all(&name, &subscripts) { + for idx in der_resolver.resolve_var_ref_all_with_constants(&name, &subscripts, constants) { result.insert(idx); } } collect_equation_lhs_unknown(eq.lhs.as_ref(), variable_resolver, &mut result); - collect_expression_unknowns(&eq.rhs, variable_resolver, &mut result); + collect_expression_unknowns_with_constants(&eq.rhs, variable_resolver, &mut result, constants); if !result.is_empty() && let Some(target) = direct_residual_definition_target(&eq.rhs) && equation_contains_derivative(&eq.rhs) { - for idx in variable_resolver.resolve_var_ref_all(target.0, target.1) { + for idx in + variable_resolver.resolve_var_ref_all_with_constants(target.0, target.1, constants) + { result.remove(&idx); } } @@ -606,6 +685,12 @@ impl ScalarUnknownResolver { if let Some(base) = dae::component_base_name(name) { base_all.entry(base).or_default().push(idx); } + if let Some(base) = record_field_parent_name(name) { + base_all.entry(base.clone()).or_default().push(idx); + if let Some(unsubscripted_base) = dae::component_base_name(&base) { + base_all.entry(unsubscripted_base).or_default().push(idx); + } + } } fn resolve_name(&self, name: &str) -> Option { @@ -616,7 +701,13 @@ impl ScalarUnknownResolver { pub(crate) fn resolve_name_all(&self, name: &str) -> Vec { if let Some(idx) = self.resolve_name(name) { - return vec![idx]; + let mut resolved = vec![idx]; + if let Some(expanded) = self.base_all.get(name) { + resolved.extend(expanded.iter().copied()); + resolved.sort_unstable(); + resolved.dedup(); + } + return resolved; } dae::component_base_name(name) .and_then(|base| self.base_all.get(&base).cloned()) @@ -636,6 +727,36 @@ impl ScalarUnknownResolver { } self.resolve_name_all(name.as_str()) } + + fn resolve_var_ref_all_with_constants( + &self, + name: &rumoca_core::Reference, + subscripts: &[rumoca_core::Subscript], + constants: &ConstantEvalContext, + ) -> Vec { + if let Some(canonical) = canonical_var_ref_key_with_constants(name, subscripts, constants) { + let resolved = self.resolve_name_all(&canonical); + if !resolved.is_empty() { + return resolved; + } + } + self.resolve_var_ref_all(name, subscripts) + } +} + +fn record_field_parent_name(name: &str) -> Option { + let mut bracket_depth = 0usize; + for (idx, ch) in name.char_indices().rev() { + match ch { + ']' => bracket_depth = bracket_depth.saturating_add(1), + '[' => bracket_depth = bracket_depth.saturating_sub(1), + '.' if bracket_depth == 0 => { + return (idx > 0).then(|| name[..idx].to_string()); + } + _ => {} + } + } + None } fn subscript_index_value(sub: &rumoca_core::Subscript) -> Option { @@ -675,18 +796,49 @@ fn canonical_var_ref_key( Some(dae::format_subscript_key(name.as_str(), &indices)) } +fn canonical_var_ref_key_with_constants( + name: &rumoca_core::Reference, + subscripts: &[rumoca_core::Subscript], + constants: &ConstantEvalContext, +) -> Option { + if subscripts.is_empty() { + return Some(name.as_str().to_string()); + } + + let mut indices = Vec::with_capacity(subscripts.len()); + for sub in subscripts { + indices.push(subscript_index_value_with_constants(sub, constants)?); + } + Some(dae::format_subscript_key(name.as_str(), &indices)) +} + +fn subscript_index_value_with_constants( + sub: &rumoca_core::Subscript, + constants: &ConstantEvalContext, +) -> Option { + subscript_index_value(sub).or_else(|| match sub { + rumoca_core::Subscript::Expr { expr, .. } => constants.eval_positive_integer(expr), + _ => None, + }) +} + pub(crate) fn collect_expression_unknowns( expr: &rumoca_core::Expression, resolver: &ScalarUnknownResolver, cols: &mut HashSet, ) { - let mut collector = ExpressionUnknownCollector { resolver, cols }; + let mut collector = ExpressionUnknownCollector { + resolver, + cols, + constants: None, + }; collector.visit_expression(expr); } struct ExpressionUnknownCollector<'a> { resolver: &'a ScalarUnknownResolver, cols: &'a mut HashSet, + constants: Option<&'a ConstantEvalContext>, } impl ExpressionVisitor for ExpressionUnknownCollector<'_> { @@ -695,7 +847,14 @@ impl ExpressionVisitor for ExpressionUnknownCollector<'_> { name: &rumoca_core::Reference, subscripts: &[rumoca_core::Subscript], ) { - for idx in self.resolver.resolve_var_ref_all(name, subscripts) { + let resolved = self + .constants + .map(|constants| { + self.resolver + .resolve_var_ref_all_with_constants(name, subscripts, constants) + }) + .unwrap_or_else(|| self.resolver.resolve_var_ref_all(name, subscripts)); + for idx in resolved { self.cols.insert(idx); } for subscript in subscripts { @@ -711,14 +870,20 @@ impl ExpressionVisitor for ExpressionUnknownCollector<'_> { if let rumoca_core::Expression::VarRef { name, subscripts: base_subscripts, - span, + span: _, } = base - && span.is_dummy() { let mut combined = Vec::with_capacity(base_subscripts.len() + subscripts.len()); combined.extend_from_slice(base_subscripts); combined.extend_from_slice(subscripts); - for idx in self.resolver.resolve_var_ref_all(name, &combined) { + let resolved = self + .constants + .map(|constants| { + self.resolver + .resolve_var_ref_all_with_constants(name, &combined, constants) + }) + .unwrap_or_else(|| self.resolver.resolve_var_ref_all(name, &combined)); + for idx in resolved { self.cols.insert(idx); } for subscript in base_subscripts { @@ -737,17 +902,10 @@ impl ExpressionVisitor for ExpressionUnknownCollector<'_> { } fn visit_field_access(&mut self, base: &rumoca_core::Expression, field: &str) { - if let Some((name, subscripts)) = indexed_field_access_var_ref_key(base, field) { - let reference = rumoca_core::Reference::new(&name); - for idx in self.resolver.resolve_var_ref_all(&reference, &subscripts) { - self.cols.insert(idx); - } - for subscript in &subscripts { - self.visit_subscript(subscript); - } + if self.collect_indexed_field_access_unknowns(base, field) { return; } - self.visit_expression(base); + self.visit_projected_field_expression(base, field); } fn visit_builtin_call( @@ -764,14 +922,430 @@ impl ExpressionVisitor for ExpressionUnknownCollector<'_> { } } -fn indexed_field_access_var_ref_key( +impl ExpressionUnknownCollector<'_> { + fn collect_indexed_field_access_unknowns( + &mut self, + base: &rumoca_core::Expression, + field: &str, + ) -> bool { + let Some((candidates, traversal_subscripts)) = + indexed_field_access_var_ref_keys(base, field) + else { + return false; + }; + + for (name, subscripts) in candidates { + let reference = rumoca_core::Reference::new(&name); + let resolved = self + .constants + .map(|constants| { + self.resolver.resolve_var_ref_all_with_constants( + &reference, + &subscripts, + constants, + ) + }) + .unwrap_or_else(|| self.resolver.resolve_var_ref_all(&reference, &subscripts)); + if resolved.is_empty() { + continue; + } + for idx in resolved { + self.cols.insert(idx); + } + break; + } + for subscript in &traversal_subscripts { + self.visit_subscript(subscript); + } + true + } + + fn visit_projected_field_expression(&mut self, expr: &rumoca_core::Expression, field: &str) { + if self.collect_indexed_field_access_unknowns(expr, field) { + return; + } + + match expr { + rumoca_core::Expression::Binary { lhs, rhs, .. } => { + self.visit_projected_field_expression(lhs, field); + self.visit_projected_field_expression(rhs, field); + } + rumoca_core::Expression::Unary { rhs, .. } => { + self.visit_projected_field_expression(rhs, field); + } + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => { + if !self.visit_projected_if_branches(branches, field) { + self.visit_projected_field_expression(else_branch, field); + } + } + rumoca_core::Expression::Array { elements, .. } + | rumoca_core::Expression::Tuple { elements, .. } => { + for element in elements { + self.visit_projected_field_expression(element, field); + } + } + rumoca_core::Expression::FunctionCall { name, args, .. } + if name.last_segment() == "Complex" => + { + if let Some(projected) = complex_constructor_field_arg(args, field) { + self.visit_expression(projected); + } else { + self.visit_expression_args(args); + } + } + rumoca_core::Expression::BuiltinCall { args, .. } + | rumoca_core::Expression::FunctionCall { args, .. } => { + self.visit_expression_args(args); + } + rumoca_core::Expression::Index { + base, subscripts, .. + } => { + self.visit_projected_field_expression(base, field); + for subscript in subscripts { + self.visit_subscript(subscript); + } + } + rumoca_core::Expression::Range { + start, step, end, .. + } => { + self.visit_expression(start); + if let Some(step) = step { + self.visit_expression(step); + } + self.visit_expression(end); + } + rumoca_core::Expression::ArrayComprehension { + expr, + indices, + filter, + .. + } => { + for index in indices { + self.visit_expression(&index.range); + } + self.visit_projected_field_expression(expr, field); + if let Some(filter) = filter { + self.visit_expression(filter); + } + } + _ => self.visit_expression(expr), + } + } + + fn visit_projected_if_branches( + &mut self, + branches: &[(rumoca_core::Expression, rumoca_core::Expression)], + field: &str, + ) -> bool { + let mut has_unresolved_condition = false; + for (condition, value) in branches { + match self + .constants + .and_then(|constants| constants.eval_bool(condition)) + { + Some(true) => { + self.visit_projected_field_expression(value, field); + return true; + } + Some(false) => continue, + None => {} + } + has_unresolved_condition = true; + self.visit_expression(condition); + self.visit_projected_field_expression(value, field); + } + !has_unresolved_condition && !branches.is_empty() + } + + fn visit_expression_args(&mut self, args: &[rumoca_core::Expression]) { + for arg in args { + self.visit_expression(arg); + } + } +} + +fn collect_expression_unknowns_with_constants( + expr: &rumoca_core::Expression, + resolver: &ScalarUnknownResolver, + cols: &mut HashSet, + constants: &ConstantEvalContext, +) { + let mut collector = ExpressionUnknownCollector { + resolver, + cols, + constants: Some(constants), + }; + collector.visit_expression(expr); +} + +fn complex_constructor_field_arg<'a>( + args: &'a [rumoca_core::Expression], + field: &str, +) -> Option<&'a rumoca_core::Expression> { + match field { + "re" => args.first(), + "im" => args.get(1), + _ => None, + } +} + +#[derive(Default)] +struct ConstantEvalContext { + scalars: HashMap, + booleans: HashMap, +} + +impl ConstantEvalContext { + fn from_dae(dae: &dae::Dae) -> Self { + let mut ctx = Self::default(); + let variables = dae + .variables + .parameters + .iter() + .chain(dae.variables.constants.iter()) + .collect::>(); + for _ in 0..variables.len().max(1) { + let before = ctx.scalars.len() + ctx.booleans.len(); + for (name, variable) in &variables { + ctx.insert_variable(name, variable); + } + if ctx.scalars.len() + ctx.booleans.len() == before { + break; + } + } + ctx + } + + fn insert_variable(&mut self, name: &rumoca_core::VarName, variable: &dae::Variable) { + let Some(start) = &variable.start else { + return; + }; + let keys = self.variable_keys(name, variable); + if let Some(value) = self.eval_scalar(start) { + for key in keys { + self.scalars.insert(key, value); + } + return; + } + if let Some(value) = literal_boolean(start) { + for key in keys { + self.booleans.insert(key, value); + } + } + } + + fn variable_keys(&self, name: &rumoca_core::VarName, variable: &dae::Variable) -> Vec { + let mut keys = vec![name.as_str().to_string()]; + if let Some(component_ref) = &variable.component_ref { + keys.push(component_ref_flat_name(component_ref)); + } + keys.sort(); + keys.dedup(); + keys + } + + fn eval_bool(&self, expr: &rumoca_core::Expression) -> Option { + match expr { + rumoca_core::Expression::Literal { + value: Literal::Boolean(value), + .. + } => Some(*value), + rumoca_core::Expression::VarRef { + name, subscripts, .. + } if subscripts.is_empty() => self.booleans.get(name.as_str()).copied(), + rumoca_core::Expression::Unary { + op: rumoca_core::OpUnary::Not, + rhs, + .. + } => self.eval_bool(rhs).map(|value| !value), + rumoca_core::Expression::Binary { op, lhs, rhs, .. } => { + self.eval_bool_binary(op.clone(), lhs, rhs) + } + _ => None, + } + } + + fn eval_bool_binary( + &self, + op: rumoca_core::OpBinary, + lhs: &rumoca_core::Expression, + rhs: &rumoca_core::Expression, + ) -> Option { + match op { + rumoca_core::OpBinary::Eq | rumoca_core::OpBinary::Neq => { + let equal = if let (Some(lhs), Some(rhs)) = + (self.eval_scalar(lhs), self.eval_scalar(rhs)) + { + scalar_almost_eq(lhs, rhs) + } else { + self.eval_bool(lhs)? == self.eval_bool(rhs)? + }; + Some(if matches!(op, rumoca_core::OpBinary::Eq) { + equal + } else { + !equal + }) + } + rumoca_core::OpBinary::Lt => Some(self.eval_scalar(lhs)? < self.eval_scalar(rhs)?), + rumoca_core::OpBinary::Le => Some(self.eval_scalar(lhs)? <= self.eval_scalar(rhs)?), + rumoca_core::OpBinary::Gt => Some(self.eval_scalar(lhs)? > self.eval_scalar(rhs)?), + rumoca_core::OpBinary::Ge => Some(self.eval_scalar(lhs)? >= self.eval_scalar(rhs)?), + rumoca_core::OpBinary::And => Some(self.eval_bool(lhs)? && self.eval_bool(rhs)?), + rumoca_core::OpBinary::Or => Some(self.eval_bool(lhs)? || self.eval_bool(rhs)?), + _ => None, + } + } + + fn eval_scalar(&self, expr: &rumoca_core::Expression) -> Option { + match expr { + rumoca_core::Expression::Literal { .. } => literal_scalar(expr), + rumoca_core::Expression::VarRef { + name, subscripts, .. + } if subscripts.is_empty() => self.lookup_scalar(name), + rumoca_core::Expression::Unary { + op: rumoca_core::OpUnary::Minus, + rhs, + .. + } => self.eval_scalar(rhs).map(|value| -value), + rumoca_core::Expression::Unary { + op: rumoca_core::OpUnary::Plus, + rhs, + .. + } => self.eval_scalar(rhs), + rumoca_core::Expression::Binary { op, lhs, rhs, .. } => { + self.eval_scalar_binary(op.clone(), lhs, rhs) + } + rumoca_core::Expression::BuiltinCall { function, args, .. } => { + self.eval_scalar_builtin(*function, args) + } + _ => None, + } + } + + fn eval_scalar_binary( + &self, + op: rumoca_core::OpBinary, + lhs: &rumoca_core::Expression, + rhs: &rumoca_core::Expression, + ) -> Option { + let lhs = self.eval_scalar(lhs)?; + let rhs = self.eval_scalar(rhs)?; + match op { + rumoca_core::OpBinary::Add | rumoca_core::OpBinary::AddElem => Some(lhs + rhs), + rumoca_core::OpBinary::Sub | rumoca_core::OpBinary::SubElem => Some(lhs - rhs), + rumoca_core::OpBinary::Mul | rumoca_core::OpBinary::MulElem => Some(lhs * rhs), + rumoca_core::OpBinary::Div | rumoca_core::OpBinary::DivElem if rhs != 0.0 => { + Some(lhs / rhs) + } + rumoca_core::OpBinary::Exp | rumoca_core::OpBinary::ExpElem => Some(lhs.powf(rhs)), + _ => None, + } + } + + fn eval_scalar_builtin( + &self, + function: rumoca_core::BuiltinFunction, + args: &[rumoca_core::Expression], + ) -> Option { + match (function, args) { + (rumoca_core::BuiltinFunction::Abs, [arg]) => Some(self.eval_scalar(arg)?.abs()), + (rumoca_core::BuiltinFunction::Floor, [arg]) + | (rumoca_core::BuiltinFunction::Integer, [arg]) => { + Some(self.eval_scalar(arg)?.floor()) + } + (rumoca_core::BuiltinFunction::Ceil, [arg]) => Some(self.eval_scalar(arg)?.ceil()), + (rumoca_core::BuiltinFunction::Min, [lhs, rhs]) => { + Some(self.eval_scalar(lhs)?.min(self.eval_scalar(rhs)?)) + } + (rumoca_core::BuiltinFunction::Max, [lhs, rhs]) => { + Some(self.eval_scalar(lhs)?.max(self.eval_scalar(rhs)?)) + } + (rumoca_core::BuiltinFunction::Div, [lhs, rhs]) => { + let lhs = self.eval_scalar(lhs)?; + let rhs = self.eval_scalar(rhs)?; + (rhs != 0.0).then_some((lhs / rhs).floor()) + } + (rumoca_core::BuiltinFunction::Mod, [lhs, rhs]) => { + let lhs = self.eval_scalar(lhs)?; + let rhs = self.eval_scalar(rhs)?; + (rhs != 0.0).then(|| lhs - (lhs / rhs).floor() * rhs) + } + (rumoca_core::BuiltinFunction::Rem, [lhs, rhs]) => { + let lhs = self.eval_scalar(lhs)?; + let rhs = self.eval_scalar(rhs)?; + (rhs != 0.0).then(|| lhs - (lhs / rhs).trunc() * rhs) + } + _ => None, + } + } + + fn eval_positive_integer(&self, expr: &rumoca_core::Expression) -> Option { + let value = self.eval_scalar(expr)?; + (value.is_finite() && value.fract() == 0.0) + .then(|| usize::try_from(value as i64).ok()) + .flatten() + .filter(|index| *index > 0) + } + + fn lookup_scalar(&self, name: &rumoca_core::Reference) -> Option { + self.scalars.get(name.as_str()).copied().or_else(|| { + name.component_ref() + .and_then(|component_ref| self.scalars.get(&component_ref_flat_name(component_ref))) + .copied() + }) + } +} + +fn component_ref_flat_name(component_ref: &rumoca_core::ComponentReference) -> String { + component_ref.to_var_name().as_str().to_string() +} + +fn literal_scalar(expr: &rumoca_core::Expression) -> Option { + match expr { + rumoca_core::Expression::Literal { + value: Literal::Integer(value), + .. + } => Some(*value as f64), + rumoca_core::Expression::Literal { + value: Literal::Real(value), + .. + } => Some(*value), + _ => None, + } +} + +fn literal_boolean(expr: &rumoca_core::Expression) -> Option { + match expr { + rumoca_core::Expression::Literal { + value: Literal::Boolean(value), + .. + } => Some(*value), + _ => None, + } +} + +fn scalar_almost_eq(lhs: f64, rhs: f64) -> bool { + (lhs - rhs).abs() <= 1.0e-12 * (1.0 + lhs.abs().max(rhs.abs())) +} + +type IndexedFieldCandidate = (String, Vec); +type IndexedFieldAccessKeys = (Vec, Vec); + +fn indexed_field_access_var_ref_keys( base: &rumoca_core::Expression, field: &str, -) -> Option<(String, Vec)> { +) -> Option { match base { rumoca_core::Expression::VarRef { name, subscripts, .. - } => Some((format!("{}.{}", name.as_str(), field), subscripts.clone())), + } => Some(( + vec![(format!("{}.{}", name.as_str(), field), subscripts.clone())], + subscripts.clone(), + )), rumoca_core::Expression::Index { base, subscripts, .. } => { @@ -786,7 +1360,12 @@ fn indexed_field_access_var_ref_key( let mut combined = Vec::with_capacity(base_subscripts.len() + subscripts.len()); combined.extend_from_slice(base_subscripts); combined.extend_from_slice(subscripts); - Some((format!("{}.{}", name.as_str(), field), combined)) + let mut candidates = Vec::new(); + candidates.push((format!("{}.{}", name.as_str(), field), combined.clone())); + if let Some(indexed_base) = canonical_var_ref_key(name, &combined) { + candidates.push((format!("{indexed_base}.{field}"), Vec::new())); + } + Some((candidates, combined)) } _ => None, } @@ -800,11 +1379,12 @@ fn indexed_field_access_var_ref_key( /// This is intended to be called after equation reordering for solver use. pub fn build_solver_sparsity_triplets(dae: &dae::Dae) -> Vec<(usize, usize)> { let resolver = ScalarUnknownResolver::from_dae(dae); + let constants = ConstantEvalContext::from_dae(dae); let mut triplets = Vec::new(); for (row, eq) in dae.continuous.equations.iter().enumerate() { let mut cols = HashSet::new(); - collect_expression_unknowns(&eq.rhs, &resolver, &mut cols); + collect_expression_unknowns_with_constants(&eq.rhs, &resolver, &mut cols, &constants); let mut cols_sorted: Vec = cols.into_iter().collect(); cols_sorted.sort_unstable(); triplets.extend(cols_sorted.into_iter().map(|col| (row, col))); @@ -870,6 +1450,21 @@ mod tests { } } + fn index_expr( + expr: rumoca_core::Expression, + subscript: rumoca_core::Expression, + ) -> rumoca_core::Expression { + let span = test_span(); + rumoca_core::Expression::Index { + base: Box::new(expr), + subscripts: vec![rumoca_core::Subscript::Expr { + expr: Box::new(subscript), + span, + }], + span, + } + } + fn field(expr: rumoca_core::Expression, name: &str) -> rumoca_core::Expression { let span = test_span(); rumoca_core::Expression::FieldAccess { @@ -887,6 +1482,52 @@ mod tests { } } + fn int_lit(v: i64) -> rumoca_core::Expression { + let span = test_span(); + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(v), + span, + } + } + + fn bin( + op: rumoca_core::OpBinary, + lhs: rumoca_core::Expression, + rhs: rumoca_core::Expression, + ) -> rumoca_core::Expression { + let span = test_span(); + rumoca_core::Expression::Binary { + op, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span, + } + } + + fn complex( + re: rumoca_core::Expression, + im: rumoca_core::Expression, + ) -> rumoca_core::Expression { + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Complex"), + args: vec![re, im], + is_constructor: true, + span: test_span(), + } + } + + fn component_ref(name: &str) -> rumoca_core::ComponentReference { + rumoca_core::ComponentReference::from_flat_segments(name, test_span(), None) + } + + fn structured_var(name: &str) -> rumoca_core::Expression { + rumoca_core::Expression::VarRef { + name: rumoca_core::Reference::from_component_reference(component_ref(name)), + subscripts: vec![], + span: test_span(), + } + } + fn sub(lhs: rumoca_core::Expression, rhs: rumoca_core::Expression) -> rumoca_core::Expression { let span = test_span(); rumoca_core::Expression::Binary { @@ -907,6 +1548,27 @@ mod tests { } } + fn div(lhs: rumoca_core::Expression, rhs: rumoca_core::Expression) -> rumoca_core::Expression { + let span = test_span(); + rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Div, + lhs: Box::new(lhs), + rhs: Box::new(rhs), + span, + } + } + + fn builtin( + function: rumoca_core::BuiltinFunction, + args: Vec, + ) -> rumoca_core::Expression { + rumoca_core::Expression::BuiltinCall { + function, + args, + span: test_span(), + } + } + fn eq(rhs: rumoca_core::Expression) -> dae::Equation { let span = test_span(); dae::Equation { @@ -1047,10 +1709,219 @@ mod tests { assert_eq!(triplets, vec![(0, 2)]); } + #[test] + fn test_build_solver_sparsity_triplets_resolves_array_element_record_fields() { + let mut dae = dae::Dae::new(); + + dae.variables.algebraics.insert( + rumoca_core::VarName::new("mediums[2].Xi"), + dae::Variable::new(rumoca_core::VarName::new("mediums[2].Xi"), test_span()), + ); + dae.continuous + .equations + .push(eq(sub(field(index(var("mediums"), 2), "Xi"), lit(0.0)))); + + let triplets = build_solver_sparsity_triplets(&dae); + assert_eq!(triplets, vec![(0, 0)]); + } + + #[test] + fn field_projection_over_binary_collects_projected_scalar_fields() { + let resolver = ScalarUnknownResolver::from_entries([ + ("a".to_string(), 0), + ("a.re".to_string(), 1), + ("a.im".to_string(), 2), + ("b".to_string(), 3), + ("b.re".to_string(), 4), + ("b.im".to_string(), 5), + ]); + let expr = field(sub(var("a"), var("b")), "re"); + + let mut cols = HashSet::new(); + collect_expression_unknowns(&expr, &resolver, &mut cols); + + assert_eq!(cols, set(&[1, 4])); + } + + #[test] + fn incidence_projects_constant_selected_complex_if_branch() { + let mut dae = dae::Dae::new(); + for name in ["a", "b", "c", "d"] { + dae.variables.algebraics.insert( + rumoca_core::VarName::new(name), + dae::Variable::new(rumoca_core::VarName::new(name), test_span()), + ); + } + let mut k = dae::Variable::new(rumoca_core::VarName::new("k"), test_span()); + k.start = Some(lit(3.0)); + dae.variables + .parameters + .insert(rumoca_core::VarName::new("k"), k); + + dae.continuous.equations.push(eq(field( + rumoca_core::Expression::If { + branches: vec![( + bin(rumoca_core::OpBinary::Eq, int_lit(3), var("k")), + complex(var("a"), var("b")), + )], + else_branch: Box::new(complex(var("c"), var("d"))), + span: test_span(), + }, + "im", + ))); + + let triplets = build_solver_sparsity_triplets(&dae); + assert_eq!(triplets, vec![(0, 1)]); + } + + #[test] + fn incidence_preserves_indexed_component_ref_constant_keys() { + let mut dae = dae::Dae::new(); + for name in ["a1", "a2", "b1", "b2"] { + dae.variables.algebraics.insert( + rumoca_core::VarName::new(name), + dae::Variable::new(rumoca_core::VarName::new(name), test_span()), + ); + } + for (index, value) in [(1, 1.0), (2, 2.0)] { + let flat_name = format!("plugToPin[{index}].k"); + let mut k = dae::Variable::new( + rumoca_core::VarName::new(format!("internal_k_{index}")), + test_span(), + ); + k.component_ref = Some(component_ref(&flat_name)); + k.start = Some(lit(value)); + dae.variables + .parameters + .insert(rumoca_core::VarName::new(format!("internal_k_{index}")), k); + } + + dae.continuous.equations.push(eq(field( + rumoca_core::Expression::If { + branches: vec![( + bin( + rumoca_core::OpBinary::Eq, + int_lit(2), + structured_var("plugToPin[2].k"), + ), + complex(var("a1"), var("a2")), + )], + else_branch: Box::new(complex(var("b1"), var("b2"))), + span: test_span(), + }, + "im", + ))); + + let triplets = build_solver_sparsity_triplets(&dae); + assert_eq!(triplets, vec![(0, 1)]); + } + + #[test] + fn incidence_resolves_parameter_expression_subscripts_to_scalar_unknowns() { + let mut dae = dae::Dae::new(); + + let mut features = dae::Variable::new(rumoca_core::VarName::new("features"), test_span()); + features.dims = vec![10]; + dae.variables + .algebraics + .insert(rumoca_core::VarName::new("features"), features); + + let mut n_latent = dae::Variable::new(rumoca_core::VarName::new("nLatent"), test_span()); + n_latent.start = Some(int_lit(8)); + dae.variables + .parameters + .insert(rumoca_core::VarName::new("nLatent"), n_latent); + + let mut n_feature = dae::Variable::new(rumoca_core::VarName::new("nFeature"), test_span()); + n_feature.start = Some(add(var("nLatent"), int_lit(2))); + dae.variables + .parameters + .insert(rumoca_core::VarName::new("nFeature"), n_feature); + + dae.continuous.equations.push(eq(sub( + index_expr(var("features"), var("nFeature")), + lit(0.0), + ))); + + let triplets = build_solver_sparsity_triplets(&dae); + assert_eq!(triplets, vec![(0, 9)]); + } + + #[test] + fn incidence_resolves_builtin_expression_subscripts_to_scalar_unknowns() { + let mut dae = dae::Dae::new(); + + let mut uu = dae::Variable::new(rumoca_core::VarName::new("uu"), test_span()); + uu.dims = vec![3]; + dae.variables + .algebraics + .insert(rumoca_core::VarName::new("uu"), uu); + + let subscript = add( + add( + builtin( + rumoca_core::BuiltinFunction::Mod, + vec![int_lit(3), int_lit(2)], + ), + builtin( + rumoca_core::BuiltinFunction::Integer, + vec![div(int_lit(3), int_lit(2))], + ), + ), + int_lit(1), + ); + dae.continuous + .equations + .push(eq(sub(index_expr(var("uu"), subscript), lit(0.0)))); + + let triplets = build_solver_sparsity_triplets(&dae); + assert_eq!(triplets, vec![(0, 2)]); + } + + #[test] + fn aggregate_refs_include_expanded_scalar_descendants() { + let resolver = ScalarUnknownResolver::from_entries([ + ("z".to_string(), 0), + ("z.re".to_string(), 1), + ("z.im".to_string(), 2), + ]); + + assert_eq!(resolver.resolve_name_all("z"), vec![0, 1, 2]); + } + fn set(values: &[usize]) -> HashSet { values.iter().copied().collect() } + fn dae_var(name: &str, dims: Vec) -> dae::Variable { + let mut var = dae::Variable::new(rumoca_core::VarName::new(name), test_span()); + var.dims = dims; + var + } + + #[test] + fn prunes_zero_degree_exact_scalarized_aggregate_element_unknowns() { + let mut dae = dae::Dae::new(); + dae.variables.algebraics.insert( + rumoca_core::VarName::new("ductOut.statesFM[3].X[2]"), + dae_var("ductOut.statesFM[3].X[2]", Vec::new()), + ); + + let mut eq_unknowns = vec![HashSet::new()]; + let (names, spans) = prune_zero_degree_scalarized_aggregate_variables( + &dae, + &mut eq_unknowns, + vec![UnknownId::Variable(rumoca_core::VarName::new( + "ductOut.statesFM[3].X[2]", + ))], + vec![Some(test_span())], + ); + + assert!(names.is_empty()); + assert!(spans.is_empty()); + assert!(eq_unknowns[0].is_empty()); + } + #[test] fn uniform_translation_detects_constant_shift() { // A stencil cell {der(X[i]), X[i-1], X[i], X[i+1]} stepped by one position diff --git a/crates/rumoca-phase-structural/src/scalarize.rs b/crates/rumoca-phase-structural/src/scalarize.rs index bbda5d14b..7c5d38ac7 100644 --- a/crates/rumoca-phase-structural/src/scalarize.rs +++ b/crates/rumoca-phase-structural/src/scalarize.rs @@ -1,5 +1,6 @@ use std::collections::HashMap; +use rumoca_core::ExpressionRewriter; use rumoca_ir_dae as dae; use crate::projection_maps::{ @@ -11,8 +12,8 @@ mod shape; #[cfg(test)] use projection::is_complex_field_scalar_name; use projection::{ - ScalarProjectionContext, lower_scalar_linear_algebra_exprs, project_rhs_for_scalar_target, - scalarized_equation_lhs, + ScalarProjectionContext, lower_scalar_linear_algebra_exprs, + project_explicit_rhs_for_scalar_target, project_rhs_for_scalar_target, scalarized_equation_lhs, }; use shape::*; @@ -252,6 +253,70 @@ pub fn build_complex_field_map(dae: &Dae) -> HashMap; 2] /// - `FunctionCall { is_constructor: true }` → project positional constructor arg `i` /// - `Binary/Unary/BuiltinCall/If/FunctionCall/Index` → recurse into children /// - Scalars (Literal, etc.) → broadcast unchanged +fn comprehension_index_value_for_scalar(range: &Expression, one_based_index: usize) -> Option { + let Expression::Range { + start, step, end, .. + } = range + else { + return None; + }; + let start = integer_literal_value(start)?; + let step = match step.as_deref() { + Some(step) => integer_literal_value(step)?, + None => 1, + }; + let offset = i64::try_from(one_based_index.checked_sub(1)?).ok()?; + let value = start.checked_add(offset.checked_mul(step)?)?; + let end = integer_literal_value(end)?; + ((step > 0 && value <= end) || (step < 0 && value >= end) || (step == 0 && value == start)) + .then_some(value) +} + +struct ComprehensionIndexSubstitution<'a> { + name: &'a str, + value: i64, + span: Span, +} + +impl ExpressionRewriter for ComprehensionIndexSubstitution<'_> { + fn walk_var_ref_expression( + &mut self, + name: &Reference, + subscripts: &[Subscript], + span: Span, + ) -> Expression { + if name.as_str() == self.name && subscripts.is_empty() { + return Expression::Literal { + value: Literal::Integer(self.value), + span: self.span, + }; + } + Expression::VarRef { + name: name.clone(), + subscripts: self.rewrite_subscripts(subscripts), + span, + } + } + + fn walk_array_comprehension_expression( + &mut self, + expr: &Expression, + indices: &[rumoca_core::ComprehensionIndex], + filter: Option<&Expression>, + span: Span, + ) -> Expression { + if indices.iter().any(|index| index.name == self.name) { + return Expression::ArrayComprehension { + expr: Box::new(expr.clone()), + indices: indices.to_vec(), + filter: filter.cloned().map(Box::new), + span, + }; + } + ExpressionRewriter::walk_array_comprehension_expression(self, expr, indices, filter, span) + } +} + pub struct IndexProjectionContext<'a> { i: usize, context_span: Option, @@ -743,22 +808,16 @@ impl<'a> IndexProjectionContext<'a> { .. } => Ok(project_array_literal_scalar(elements, *is_matrix, self.i) .unwrap_or_else(|| expr.clone())), + Expression::ArrayComprehension { + expr: inner, + indices, + filter, + span, + } => self.project_array_comprehension(inner, indices, filter.as_deref(), *span), Expression::VarRef { name, subscripts, .. } => self.project_var_ref(name, subscripts, expr), - Expression::Binary { op, lhs, rhs, span } => { - if matches!(op, OpBinary::Mul) - && let Some(projected) = self.project_matrix_mul(lhs, rhs, *span)? - { - return Ok(projected); - } - Ok(Expression::Binary { - op: op.clone(), - lhs: Box::new(self.project(lhs)?), - rhs: Box::new(self.project(rhs)?), - span: *span, - }) - } + Expression::Binary { op, lhs, rhs, span } => self.project_binary(op, lhs, rhs, *span), Expression::Unary { op, rhs, span } => Ok(Expression::Unary { op: op.clone(), rhs: Box::new(self.project(rhs)?), @@ -807,11 +866,36 @@ impl<'a> IndexProjectionContext<'a> { base, subscripts, span, - } => Ok(Expression::Index { - base: Box::new(self.project(base)?), - subscripts: subscripts.clone(), - span: *span, - }), + } => { + if let Expression::VarRef { + name, + subscripts: base_subscripts, + .. + } = base.as_ref() + && base_subscripts.is_empty() + && let Some(dims) = self.var_dims.get(name.as_str()) + { + return project_dimmed_var_ref(name, dims, subscripts, expr, self); + } + if let [Subscript::Index { value, .. }] = subscripts.as_slice() + && let Ok(index) = usize::try_from(*value) + && index > 0 + { + let selected_base = self.project_at(base, index)?; + if self + .expression_dims(&selected_base) + .is_some_and(|dims| dims.is_empty()) + { + return Ok(selected_base); + } + } + let projected_base = self.project(base)?; + Ok(Expression::Index { + base: Box::new(projected_base), + subscripts: subscripts.clone(), + span: *span, + }) + } Expression::FieldAccess { base, field, .. } => { if let Some(projected) = self.project_record_array_member_slice(base, field)? { return Ok(projected); @@ -822,6 +906,58 @@ impl<'a> IndexProjectionContext<'a> { } } + fn project_binary( + &self, + op: &OpBinary, + lhs: &Expression, + rhs: &Expression, + span: Span, + ) -> Result { + if matches!(op, OpBinary::Mul) + && let Some(projected) = self.project_matrix_mul(lhs, rhs, span)? + { + return Ok(projected); + } + Ok(Expression::Binary { + op: op.clone(), + lhs: Box::new(self.project(lhs)?), + rhs: Box::new(self.project(rhs)?), + span, + }) + } + + fn project_array_comprehension( + &self, + inner: &Expression, + indices: &[rumoca_core::ComprehensionIndex], + filter: Option<&Expression>, + span: Span, + ) -> Result { + if filter.is_some() || indices.len() != 1 { + return Ok(Expression::ArrayComprehension { + expr: Box::new(inner.clone()), + indices: indices.to_vec(), + filter: filter.cloned().map(Box::new), + span, + }); + } + let Some(value) = comprehension_index_value_for_scalar(&indices[0].range, self.i) else { + return Ok(Expression::ArrayComprehension { + expr: Box::new(inner.clone()), + indices: indices.to_vec(), + filter: None, + span, + }); + }; + let mut substitution = ComprehensionIndexSubstitution { + name: &indices[0].name, + value, + span: indices[0].range.span().unwrap_or(span), + }; + let selected = substitution.rewrite_expression(inner); + self.project(&selected) + } + fn project_function_call( &self, name: &Reference, @@ -992,10 +1128,7 @@ fn project_subscripted_dims( return Ok(None); }; let scalar_count = output_scalar_count(&projected_dims, span)?; - if scalar_count == 0 - || linear_index > scalar_count - || (scalar_count == 1 && projected_dims.is_empty()) - { + if scalar_count == 0 || linear_index > scalar_count { return Ok(None); } @@ -1054,8 +1187,18 @@ fn project_subscripted_dims( )?); dim_idx += 1; } - Subscript::Expr { expr: _, .. } => { - projected_subscripts.push(subscript.clone()); + Subscript::Expr { expr, .. } => { + if let Some(index) = + eval_structural_int_expr(expr, structural_values, &HashMap::new()) + { + projected_subscripts.push(positive_generated_index_subscript( + index, + span, + "structural fixed-expression subscript", + )?); + } else { + projected_subscripts.push(subscript.clone()); + } dim_idx += 1; } Subscript::Colon { .. } => { @@ -1422,22 +1565,17 @@ pub fn scalarization_var_ref_name( pub fn residual_lhs_target_name(expr: &Expression) -> Option { let Expression::Binary { op: OpBinary::Sub, - lhs, + lhs: _, .. } = expr else { return None; }; - if let Expression::VarRef { - name, subscripts, .. - } = lhs.as_ref() - { - return scalarization_var_ref_name(name, subscripts); - } - None + residual_lhs_var_ref(expr) + .and_then(|(name, subscripts)| scalarization_var_ref_name(name, &subscripts)) } -fn residual_lhs_var_ref(expr: &Expression) -> Option<(&Reference, &[Subscript])> { +fn residual_lhs_var_ref(expr: &Expression) -> Option<(&Reference, Vec)> { let Expression::Binary { op: OpBinary::Sub, lhs, @@ -1446,13 +1584,36 @@ fn residual_lhs_var_ref(expr: &Expression) -> Option<(&Reference, &[Subscript])> else { return None; }; - if let Expression::VarRef { - name, subscripts, .. - } = lhs.as_ref() - { - return Some((name, subscripts)); + residual_lhs_var_ref_from_expr(lhs.as_ref()) +} + +fn residual_lhs_var_ref_from_expr(expr: &Expression) -> Option<(&Reference, Vec)> { + match expr { + Expression::VarRef { + name, subscripts, .. + } => Some((name, subscripts.clone())), + Expression::Array { elements, .. } | Expression::Tuple { elements, .. } + if elements.len() == 1 => + { + residual_lhs_var_ref_from_expr(elements.first()?) + } + Expression::Index { + base, subscripts, .. + } => { + let Expression::VarRef { + name, + subscripts: base_subscripts, + .. + } = base.as_ref() + else { + return None; + }; + base_subscripts + .is_empty() + .then(|| (name, subscripts.clone())) + } + _ => None, } - None } fn residual_lhs_scalar_targets( @@ -1471,7 +1632,8 @@ fn residual_lhs_scalar_targets( return Ok(Vec::new()); } - let Some(projected_dims) = apply_subscripts_to_dims(dims, subscripts, structural_values) else { + let Some(projected_dims) = apply_subscripts_to_dims(dims, &subscripts, structural_values) + else { return Ok(Vec::new()); }; let scalar_count = output_scalar_count(&projected_dims, span)?; @@ -1479,7 +1641,7 @@ fn residual_lhs_scalar_targets( let mut targets = Vec::new(); for idx in 1..=scalar_count { let Some(projected_subscripts) = - project_subscripted_dims(dims, subscripts, idx, span, structural_values)? + project_subscripted_dims(dims, &subscripts, idx, span, structural_values)? else { continue; }; @@ -1679,66 +1841,15 @@ pub fn scalarize_equations(dae: &mut Dae) -> Result<(), StructuralError> { let mut spans: Vec<(usize, usize)> = Vec::with_capacity(dae.continuous.equations.len()); for eq in &dae.continuous.equations { let new_start = expanded.len(); - let eq_projection = projection.with_context_span(eq.span); - let scalarization_target = eq - .lhs - .as_ref() - .map(|lhs| lhs.as_str().to_string()) - .or_else(|| residual_lhs_target_name(&eq.rhs)); - let residual_lhs_targets = - residual_lhs_scalar_targets(&eq.rhs, eq.span, &var_dims, &structural_values)?; - let (lhs_targets, has_residual_lhs_targets) = scalar_lhs_targets_for_equation( - residual_lhs_targets, - scalarization_target.as_deref(), - eq.lhs.as_ref(), - eq.span, + expand_scalarized_equation( + eq, + &projection, &scalar_names, &var_dims, &var_spans, + &structural_values, + &mut expanded, )?; - let rhs_shape_count = shape_scalar_count(eq_projection.expression_shape(&eq.rhs)); - let scalar_count = if has_residual_lhs_targets { - lhs_targets.len().max(1) - } else if let Some(rhs_count) = rhs_shape_count.filter(|count| *count > 1) { - // MLS §10.6 / SPEC_0019: array equations represent one scalar - // equation per array element. Prefer the expression IR shape over - // stale scalar_count metadata for residuals such as - // `J * der(omega) - M_body`. - rhs_count.max(lhs_targets.len()) - } else { - eq.scalar_count.max(lhs_targets.len()).max(1) - }; - if scalar_count <= 1 { - let mut lowered = eq.clone(); - let rhs = if projection - .expression_shape(&lowered.rhs) - .is_singleton_array() - { - eq_projection.project_index(&lowered.rhs, 1)? - } else { - lowered.rhs.clone() - }; - lowered.rhs = eq_projection.lower_scalar_linear_algebra(&rhs)?; - expanded.push(lowered); - } else { - for i in 1..=scalar_count { - let target = lhs_targets.get(i - 1); - expanded.push(Equation { - lhs: scalarized_equation_lhs(eq, target, i, eq.span)?, - rhs: project_rhs_for_scalar_target( - &eq.rhs, - i, - scalarization_target.as_deref(), - target, - eq.span, - &eq_projection, - )?, - span: eq.span, - origin: eq.origin.clone(), - scalar_count: 1, - }); - } - } spans.push((new_start, expanded.len() - new_start)); } dae.continuous.equations = expanded; @@ -1754,5 +1865,121 @@ pub fn scalarize_equations(dae: &mut Dae) -> Result<(), StructuralError> { Ok(()) } +fn expand_scalarized_equation( + eq: &Equation, + projection: &ScalarProjectionContext<'_>, + scalar_names: &[String], + var_dims: &HashMap>, + var_spans: &HashMap, + structural_values: &HashMap, + expanded: &mut Vec, +) -> Result<(), StructuralError> { + let eq_projection = projection.with_context_span(eq.span); + let scalarization_target = eq + .lhs + .as_ref() + .map(|lhs| lhs.as_str().to_string()) + .or_else(|| residual_lhs_target_name(&eq.rhs)); + let residual_lhs_targets = + residual_lhs_scalar_targets(&eq.rhs, eq.span, var_dims, structural_values)?; + let (lhs_targets, has_residual_lhs_targets) = scalar_lhs_targets_for_equation( + residual_lhs_targets, + scalarization_target.as_deref(), + eq.lhs.as_ref(), + eq.span, + scalar_names, + var_dims, + var_spans, + )?; + let scalar_count = + equation_scalar_count(eq, &eq_projection, &lhs_targets, has_residual_lhs_targets); + if scalar_count <= 1 { + return expand_single_scalar_equation( + eq, + &eq_projection, + scalarization_target.as_deref(), + &lhs_targets, + has_residual_lhs_targets, + expanded, + ); + } + expand_multi_scalar_equation( + eq, + &eq_projection, + scalarization_target.as_deref(), + &lhs_targets, + scalar_count, + expanded, + ) +} + +fn expand_single_scalar_equation( + eq: &Equation, + projection: &ScalarProjectionContext<'_>, + scalarization_target: Option<&str>, + lhs_targets: &[ScalarizedLhsTarget], + has_residual_lhs_targets: bool, + expanded: &mut Vec, +) -> Result<(), StructuralError> { + if has_residual_lhs_targets + && let Some(target) = lhs_targets.first() + && let Some(rhs) = project_explicit_rhs_for_scalar_target( + &eq.rhs, + scalarization_target, + target, + projection, + )? + { + expanded.push(Equation::explicit_with_scalar_count( + VarName::new(target.name.clone()), + projection.lower_scalar_linear_algebra(&rhs)?, + eq.span, + eq.origin.clone(), + 1, + )); + return Ok(()); + } + let mut lowered = eq.clone(); + let rhs = if projection + .expression_shape(&lowered.rhs) + .is_singleton_array() + { + projection.project_index(&lowered.rhs, 1)? + } else { + lowered.rhs.clone() + }; + lowered.rhs = projection.lower_scalar_linear_algebra(&rhs)?; + expanded.push(lowered); + Ok(()) +} + +fn expand_multi_scalar_equation( + eq: &Equation, + projection: &ScalarProjectionContext<'_>, + scalarization_target: Option<&str>, + lhs_targets: &[ScalarizedLhsTarget], + scalar_count: usize, + expanded: &mut Vec, +) -> Result<(), StructuralError> { + for i in 1..=scalar_count { + let target = lhs_targets.get(i - 1); + expanded.push(Equation { + lhs: scalarized_equation_lhs(eq, target, i, eq.span)?, + rhs: project_rhs_for_scalar_target( + &eq.rhs, + i, + scalarization_target, + target, + eq.span, + projection, + )?, + span: eq.span, + origin: eq.origin.clone(), + scalar_count: 1, + }); + } + Ok(()) +} + #[cfg(test)] mod tests; diff --git a/crates/rumoca-phase-structural/src/scalarize/projection.rs b/crates/rumoca-phase-structural/src/scalarize/projection.rs index 11619aa31..53f281dba 100644 --- a/crates/rumoca-phase-structural/src/scalarize/projection.rs +++ b/crates/rumoca-phase-structural/src/scalarize/projection.rs @@ -440,10 +440,7 @@ pub(super) fn project_rhs_for_scalar_target( .. } = rhs && matches!(op, OpBinary::Sub) - && let Expression::VarRef { - name, subscripts, .. - } = lhs.as_ref() - && scalarization_var_ref_name(name, subscripts) + && expression_scalarization_name(lhs.as_ref()) .as_deref() .is_some_and(|lhs_row_name| lhs_row_name == lhs_name) { @@ -475,6 +472,71 @@ pub(super) fn project_rhs_for_scalar_target( projection.project_index(rhs, scalar_idx) } +pub(super) fn project_explicit_rhs_for_scalar_target( + rhs: &Expression, + lhs_target: Option<&str>, + target: &ScalarizedLhsTarget, + projection: &ScalarProjectionContext<'_>, +) -> Result, StructuralError> { + let Expression::Binary { + op, + lhs, + rhs: row_rhs, + .. + } = rhs + else { + return Ok(None); + }; + if !matches!(op, OpBinary::Sub) + || !lhs_target.is_some_and(|lhs_name| { + expression_scalarization_name(lhs.as_ref()) + .as_deref() + .is_some_and(|lhs_row_name| lhs_row_name == lhs_name) + }) + { + return Ok(None); + } + + let mut projected_rhs = (*row_rhs.clone()).clone(); + if let Some(idx) = target.array_selector { + projected_rhs = projection.project_index(&projected_rhs, idx)?; + } + if let Some(field_idx) = target.field_selector { + projected_rhs = project_complex_component(&projected_rhs, field_idx, projection)?; + } + Ok(Some(projected_rhs)) +} + +fn expression_scalarization_name(expr: &Expression) -> Option { + match expr { + Expression::VarRef { + name, subscripts, .. + } => scalarization_var_ref_name(name, subscripts), + Expression::Array { elements, .. } | Expression::Tuple { elements, .. } + if elements.len() == 1 => + { + expression_scalarization_name(elements.first()?) + } + Expression::Index { + base, subscripts, .. + } => { + let Expression::VarRef { + name, + subscripts: base_subscripts, + .. + } = base.as_ref() + else { + return None; + }; + base_subscripts + .is_empty() + .then(|| scalarization_var_ref_name(name, subscripts)) + .flatten() + } + _ => None, + } +} + pub(super) fn scalarized_equation_lhs( eq: &Equation, target: Option<&ScalarizedLhsTarget>, @@ -485,10 +547,11 @@ pub(super) fn scalarized_equation_lhs( return Ok(None); }; if let Some(name) = target { - return Ok(Some(structured_scalar_target( - &rumoca_core::VarName::new(name.name.clone()), - span, - ))); + let target_name = rumoca_core::VarName::new(name.name.clone()); + if let Some(reference) = scalarized_target_from_lhs(lhs, &target_name, span)? { + return Ok(Some(reference)); + } + return Ok(Some(structured_scalar_target(&target_name, span))); } let indexed = lhs.with_appended_index( scalar_idx as i64, @@ -503,6 +566,42 @@ pub(super) fn scalarized_equation_lhs( ))) } +fn scalarized_target_from_lhs( + lhs: &rumoca_core::Reference, + target: &rumoca_core::VarName, + equation_span: Span, +) -> Result, StructuralError> { + let Some(component_ref) = lhs.component_ref() else { + return Ok(None); + }; + let Some(scalar) = rumoca_core::parse_scalar_name(target.as_str()) else { + return Ok(None); + }; + if scalar.base != lhs.as_str() { + return Ok(None); + } + let owner_span = scalarized_lhs_index_owner_span(lhs, equation_span)?; + let mut component_ref = component_ref.clone(); + let Some(last) = component_ref.parts.last_mut() else { + return Ok(None); + }; + for index in scalar.indices { + let subscript = rumoca_core::Subscript::try_generated_index( + index, + owner_span.into(), + "scalarized equation LHS", + ) + .map_err(|err| StructuralError::UnspannedContractViolation { + reason: err.to_string(), + })?; + last.subs.push(subscript); + } + Ok(Some(rumoca_core::Reference::with_component_reference( + target.as_str(), + component_ref, + ))) +} + fn scalarized_lhs_index_owner_span( lhs: &rumoca_core::Reference, equation_span: Span, diff --git a/crates/rumoca-phase-structural/src/scalarize/shape.rs b/crates/rumoca-phase-structural/src/scalarize/shape.rs index 42b1683e6..29d73467d 100644 --- a/crates/rumoca-phase-structural/src/scalarize/shape.rs +++ b/crates/rumoca-phase-structural/src/scalarize/shape.rs @@ -144,6 +144,24 @@ pub(super) fn shape_scalar_count(shape: ExpressionShape) -> Option { } } +pub(super) fn equation_scalar_count( + eq: &Equation, + projection: &ScalarProjectionContext<'_>, + lhs_targets: &[ScalarizedLhsTarget], + has_residual_lhs_targets: bool, +) -> usize { + if has_residual_lhs_targets { + return lhs_targets.len().max(1); + } + let rhs_shape_count = shape_scalar_count(projection.expression_shape(&eq.rhs)); + if let Some(rhs_count) = rhs_shape_count.filter(|count| *count > 1) { + // MLS 10.6 / SPEC_0019: array equations represent one scalar equation + // per array element. Prefer expression IR shape over stale metadata. + return rhs_count.max(lhs_targets.len()); + } + eq.scalar_count.max(lhs_targets.len()).max(1) +} + pub(super) fn combine_additive_shapes( lhs: ExpressionShape, rhs: ExpressionShape, diff --git a/crates/rumoca-phase-structural/src/scalarize/tests.rs b/crates/rumoca-phase-structural/src/scalarize/tests.rs index f559e9660..57c1d2b07 100644 --- a/crates/rumoca-phase-structural/src/scalarize/tests.rs +++ b/crates/rumoca-phase-structural/src/scalarize/tests.rs @@ -527,6 +527,27 @@ fn expr_contains_der_var_idx(expr: &Expression, target: &str, idx: i64) -> bool } } +fn assert_scalar_assignment(expr: &Expression, target: &str, indices: &[i64], expected: f64) { + let Expression::Binary { + op: OpBinary::Sub, + lhs, + rhs, + .. + } = expr + else { + panic!("expected scalar residual assignment, got {expr:?}"); + }; + assert_eq!(lhs.as_ref(), &var_idx(target, indices)); + let Expression::Literal { + value: Literal::Real(actual), + .. + } = rhs.as_ref() + else { + panic!("expected real RHS literal, got {rhs:?}"); + }; + assert_eq!(*actual, expected); +} + #[test] fn build_output_names_orders_states_algebraics_outputs_and_expands_arrays() { let mut dae_model = dae::Dae::default(); @@ -762,6 +783,63 @@ fn scalarize_matrix_vector_product_uses_row_dot_product() { ); } +#[test] +fn scalarize_vector_residual_projects_single_column_matrix_rhs_lanes() { + let mut dae_model = dae::Dae::default(); + dae_model + .variables + .outputs + .insert(VarName::new("y"), variable("y", &[3])); + for name in ["u1", "u2", "u3"] { + dae_model + .variables + .algebraics + .insert(VarName::new(name), variable(name, &[1])); + } + let lhs = Expression::Array { + elements: vec![var("y")], + is_matrix: true, + span: test_span(), + }; + let rhs = Expression::Array { + elements: vec![ + Expression::Array { + elements: vec![var("u1")], + is_matrix: true, + span: test_span(), + }, + Expression::Array { + elements: vec![var("u2")], + is_matrix: true, + span: test_span(), + }, + Expression::Array { + elements: vec![var("u3")], + is_matrix: true, + span: test_span(), + }, + ], + is_matrix: true, + span: test_span(), + }; + dae_model + .continuous + .equations + .push(residual_with_binary_span(lhs, rhs, 3, test_span())); + + scalarize_equations(&mut dae_model).unwrap(); + + assert_eq!(dae_model.continuous.equations.len(), 3); + for (idx, expected_rhs) in ["u1", "u2", "u3"].into_iter().enumerate() { + let eq = &dae_model.continuous.equations[idx]; + assert_eq!(residual_lhs_ref(eq), Some(("y", vec![(idx + 1) as i64]))); + let Expression::Binary { rhs, .. } = &eq.rhs else { + panic!("expected residual subtraction"); + }; + assert_eq!(rhs.as_ref(), &var(expected_rhs)); + } +} + #[test] fn scalarize_projected_function_output_keeps_array_argument_whole() { let mut dae_model = dae::Dae::default(); @@ -1414,6 +1492,103 @@ fn scalarize_residual_column_slice_lhs_infers_targets_from_array_ir() { ); } +#[test] +fn scalarize_residual_index_column_slice_lhs_infers_targets_from_array_ir() { + let mut dae_model = dae::Dae::default(); + dae_model + .variables + .algebraics + .insert(VarName::new("M"), variable("M", &[3, 4])); + dae_model + .variables + .parameters + .insert(VarName::new("r"), variable("r", &[3, 4])); + dae_model + .variables + .parameters + .insert(VarName::new("f"), variable("f", &[3, 4])); + let column = |name: &str| Expression::Index { + base: Box::new(var(name)), + subscripts: vec![ + Subscript::generated_colon(test_span()), + Subscript::generated_index(2, test_span()), + ], + span: test_span(), + }; + dae_model.continuous.equations.push(Equation { + lhs: None, + rhs: sub_expr(column("M"), cross(column("r"), column("f"))), + span: test_span(), + origin: "index slice residual".to_string(), + scalar_count: 3, + }); + + scalarize_equations(&mut dae_model).unwrap(); + + assert_eq!(dae_model.continuous.equations.len(), 3); + assert_eq!( + dae_model.continuous.equations[0].rhs, + sub_expr( + var_idx("M", &[1, 2]), + sub_expr( + mul_expr(var_idx("r", &[2, 2]), var_idx("f", &[3, 2])), + mul_expr(var_idx("r", &[3, 2]), var_idx("f", &[2, 2])), + ), + ) + ); + assert_eq!( + dae_model.continuous.equations[1].rhs, + sub_expr( + var_idx("M", &[2, 2]), + sub_expr( + mul_expr(var_idx("r", &[3, 2]), var_idx("f", &[1, 2])), + mul_expr(var_idx("r", &[1, 2]), var_idx("f", &[3, 2])), + ), + ) + ); + assert_eq!( + dae_model.continuous.equations[2].rhs, + sub_expr( + var_idx("M", &[3, 2]), + sub_expr( + mul_expr(var_idx("r", &[1, 2]), var_idx("f", &[2, 2])), + mul_expr(var_idx("r", &[2, 2]), var_idx("f", &[1, 2])), + ), + ) + ); +} + +#[test] +fn scalarize_consumes_static_cross_component_index_after_reprojection() { + let mut dae_model = dae::Dae::default(); + for name in ["r", "f"] { + dae_model + .variables + .parameters + .insert(VarName::new(name), variable(name, &[3])); + } + dae_model + .variables + .algebraics + .insert(VarName::new("m"), variable("m", &[])); + dae_model + .continuous + .equations + .push(eq("m", index(cross(var("r"), var("f")), &[2]), 1)); + + scalarize_equations(&mut dae_model).expect("indexed cross component should scalarize"); + + assert_eq!(dae_model.continuous.equations.len(), 1); + assert!( + !matches!( + dae_model.continuous.equations[0].rhs, + Expression::Index { .. } + ), + "a scalarized cross component must not retain a stale outer Index: {:#?}", + dae_model.continuous.equations[0].rhs + ); +} + #[test] fn scalarize_matrix_vector_derivative_residual_uses_expression_shape() { let mut dae_model = dae::Dae::default(); @@ -1545,6 +1720,42 @@ fn scalarize_matrix_matrix_derivative_residual_preserves_derivative_lhs_rows() { } } +#[test] +fn scalarize_matrix_literal_assignment_projects_whole_array_lhs() { + let mut dae_model = dae::Dae::default(); + dae_model + .variables + .algebraics + .insert(VarName::new("skew"), variable("skew", &[3, 3])); + + dae_model.continuous.equations.push(Equation { + lhs: None, + rhs: sub_expr( + var("skew"), + array(vec![ + array(vec![real(0.0), real(-1.0), real(0.0)]), + array(vec![real(1.0), real(0.0), real(0.0)]), + array(vec![real(0.0), real(0.0), real(0.0)]), + ]), + ), + span: Span::DUMMY, + origin: "matrix literal residual".to_string(), + scalar_count: 1, + }); + + scalarize_equations(&mut dae_model).unwrap(); + + assert_eq!(dae_model.continuous.equations.len(), 9); + assert_scalar_assignment(&dae_model.continuous.equations[0].rhs, "skew", &[1, 1], 0.0); + assert_scalar_assignment( + &dae_model.continuous.equations[1].rhs, + "skew", + &[1, 2], + -1.0, + ); + assert_scalar_assignment(&dae_model.continuous.equations[3].rhs, "skew", &[2, 1], 1.0); +} + #[test] fn scalarize_lowers_indexed_matrix_vector_product_base() { let mut dae_model = dae::Dae::default(); diff --git a/crates/rumoca-phase-structural/src/variable_scope.rs b/crates/rumoca-phase-structural/src/variable_scope.rs index bc83e6c42..263508e52 100644 --- a/crates/rumoca-phase-structural/src/variable_scope.rs +++ b/crates/rumoca-phase-structural/src/variable_scope.rs @@ -55,6 +55,9 @@ impl<'a> DaeVariableScope<'a> { validate_scalarized_indices(name, &scalar.indices, &base_var.dims, None)?; return Ok(base_var.dims[scalar.indices.len()..].to_vec()); } + if let Some(dims) = self.indexed_descendant_aggregate_dims(name) { + return Ok(dims); + } Err(missing_dae_variable_metadata(name, None)) } @@ -66,13 +69,21 @@ impl<'a> DaeVariableScope<'a> { &self, name: &Reference, ) -> Result { - if let Some(var) = self.exact(name.var_name()) { + if let Some(var) = self.exact_reference(name) { return Ok(DaeVariableShape::Dimensions(var.dims.clone())); } if let Some(dims) = self.indexed_reference_dims(name)? { return Ok(DaeVariableShape::Dimensions(dims)); } - if name.as_str() == "time" || self.has_descendant_reference(name) { + if let Some(dims) = self.indexed_descendant_aggregate_dims(name.var_name()) { + return Ok(DaeVariableShape::Dimensions(dims)); + } + if name.as_str() == "time" + || self.has_descendant_reference(name) + || name + .component_ref() + .is_some_and(|component_ref| component_ref.parts.len() > 1) + { return Ok(DaeVariableShape::StructuredAggregate); } Err(missing_dae_variable_metadata(name.var_name(), name.span())) @@ -97,6 +108,30 @@ impl<'a> DaeVariableScope<'a> { .transpose() } + pub(crate) fn scalarized_aggregate_target( + &self, + name: &VarName, + ) -> Option<(Reference, Vec)> { + if self.exact(name).is_some() { + return None; + } + let scalar = rumoca_core::parse_scalar_name(name.as_str())?; + let base_name = VarName::new(scalar.base); + let base_var = self.exact(&base_name)?; + if base_var.dims.is_empty() || scalar.indices.len() != base_var.dims.len() { + return None; + } + validate_scalarized_indices(name, &scalar.indices, &base_var.dims, None).ok()?; + let indices = scalar + .indices + .into_iter() + .map(usize::try_from) + .collect::, _>>() + .ok()?; + let base = Reference::from_component_reference(base_var.component_ref.clone()?); + Some((base, indices)) + } + pub(crate) fn is_indexed_component_variable(&self, name: &VarName) -> bool { self.exact(name) .and_then(|var| var.component_ref.as_ref()) @@ -117,7 +152,11 @@ impl<'a> DaeVariableScope<'a> { else { return false; }; - component_refs_share_base(name_ref, unknown_ref) + if component_ref_has_scalar_subscript(name_ref) { + component_refs_match_exactly(name_ref, unknown_ref) + } else { + component_refs_share_base(name_ref, unknown_ref) + } } pub(crate) fn has_descendant_reference(&self, prefix: &Reference) -> bool { @@ -194,6 +233,55 @@ impl<'a> DaeVariableScope<'a> { .chain(self.dae.variables.discrete_valued.values()) } + fn indexed_descendant_aggregate_dims(&self, name: &VarName) -> Option> { + let mut max_indices: Vec = Vec::new(); + let mut descendant_dims: Option> = None; + let mut matched = false; + for variable in self.all_variables() { + let Some(component_ref) = &variable.component_ref else { + continue; + }; + let Some((stripped_name, indices)) = + stripped_component_ref_name_and_indices(component_ref) + else { + continue; + }; + if stripped_name != name.as_str() { + continue; + } + matched = true; + if max_indices.len() < indices.len() { + max_indices.resize(indices.len(), 0); + } + for (slot, index) in max_indices.iter_mut().zip(indices) { + *slot = (*slot).max(index); + } + match &descendant_dims { + Some(dims) if dims != &variable.dims => return None, + Some(_) => {} + None => descendant_dims = Some(variable.dims.clone()), + } + } + if !matched || max_indices.is_empty() { + return None; + } + max_indices.extend(descendant_dims.unwrap_or_default()); + Some(max_indices) + } + + fn exact_reference(&self, name: &Reference) -> Option<&'a dae::Variable> { + if name + .component_ref() + .is_some_and(component_ref_has_scalar_subscript) + { + self.exact(&VarName::new(name.as_str())) + .or_else(|| self.exact(name.var_name())) + } else { + self.exact(name.var_name()) + .or_else(|| self.exact(&VarName::new(name.as_str()))) + } + } + fn indexed_reference_dims( &self, name: &Reference, @@ -202,11 +290,15 @@ impl<'a> DaeVariableScope<'a> { return Ok(None); }; let mut base_ref = component_ref.clone(); - let mut scalar_indices = Vec::new(); - for part in &mut base_ref.parts { - scalar_indices.extend(part.subs.iter().filter_map(scalar_subscript_index)); - part.subs.clear(); - } + let Some(leaf) = base_ref.parts.last_mut() else { + return Ok(None); + }; + let scalar_indices = leaf + .subs + .iter() + .filter_map(scalar_subscript_index) + .collect::>(); + leaf.subs.clear(); if scalar_indices.is_empty() { return Ok(None); } @@ -313,6 +405,15 @@ fn component_refs_share_base(lhs: &ComponentReference, rhs: &ComponentReference) .all(|(lhs, rhs)| lhs.ident == rhs.ident) } +fn component_refs_match_exactly(lhs: &ComponentReference, rhs: &ComponentReference) -> bool { + lhs.parts.len() == rhs.parts.len() + && lhs + .parts + .iter() + .zip(&rhs.parts) + .all(|(lhs, rhs)| part_matches_without_span(lhs, rhs)) +} + fn component_ref_has_scalar_subscript(reference: &ComponentReference) -> bool { reference .parts @@ -394,6 +495,23 @@ fn index_key(component_ref: &ComponentReference, indexed_part: usize) -> Option< .collect() } +fn stripped_component_ref_name_and_indices( + component_ref: &ComponentReference, +) -> Option<(String, Vec)> { + let mut stripped = component_ref.clone(); + let mut indices = Vec::new(); + for part in &mut stripped.parts { + if part.subs.is_empty() { + continue; + } + for subscript in &part.subs { + indices.push(scalar_subscript_index(subscript)?); + } + part.subs.clear(); + } + (!indices.is_empty()).then(|| (stripped.to_var_name().as_str().to_string(), indices)) +} + fn scalar_subscript_index(subscript: &Subscript) -> Option { match subscript { Subscript::Index { value, .. } if *value > 0 => Some(*value), @@ -544,4 +662,170 @@ mod tests { }) if reason.contains("missing DAE variable metadata for `missing`") && actual == span )); } + + #[test] + fn hierarchical_structured_reference_without_leaf_metadata_is_aggregate() { + let span = test_span(); + let reference = Reference::from_component_reference(component_ref_with_span( + vec![ + part("machine", Vec::new()), + part("plug", Vec::new()), + part("pin", Vec::new()), + part("i", Vec::new()), + ], + span, + )); + let dae_model = dae::Dae::default(); + let scope = DaeVariableScope::new(&dae_model); + + assert_eq!( + scope + .shape_for_reference(&reference) + .expect("hierarchical source reference should retain aggregate shape"), + DaeVariableShape::StructuredAggregate + ); + } + + #[test] + fn indexed_component_reference_prefers_scalar_variable_metadata() { + let span = test_span(); + let mut dae_model = dae::Dae::default(); + dae_model.variables.algebraics.insert( + VarName::new("machine.plug.pin[1].i"), + dae::Variable { + name: VarName::new("machine.plug.pin[1].i"), + dims: Vec::new(), + component_ref: Some(component_ref_with_span( + vec![ + part("machine", Vec::new()), + part("plug", Vec::new()), + part("pin", vec![index(1)]), + part("i", Vec::new()), + ], + span, + )), + ..rumoca_ir_dae::Variable::empty_with_span(span) + }, + ); + let reference = Reference::from_component_reference(component_ref_with_span( + vec![ + part("machine", Vec::new()), + part("plug", Vec::new()), + part("pin", vec![index(1)]), + part("i", Vec::new()), + ], + span, + )); + let scope = DaeVariableScope::new(&dae_model); + + assert_eq!( + scope + .dims_for_reference(&reference) + .expect("indexed scalar reference should resolve"), + Some(Vec::new()) + ); + } + + #[test] + fn aggregate_dims_can_be_inferred_from_indexed_scalar_descendants() { + let span = test_span(); + let mut dae_model = dae::Dae::default(); + for pin_index in [1, 2, 3] { + dae_model.variables.algebraics.insert( + VarName::new(format!("machine.plug.pin[{pin_index}].i")), + dae::Variable { + name: VarName::new(format!("machine.plug.pin[{pin_index}].i")), + dims: Vec::new(), + component_ref: Some(component_ref_with_span( + vec![ + part("machine", Vec::new()), + part("plug", Vec::new()), + part("pin", vec![index(pin_index)]), + part("i", Vec::new()), + ], + span, + )), + ..rumoca_ir_dae::Variable::empty_with_span(span) + }, + ); + } + let scope = DaeVariableScope::new(&dae_model); + + assert_eq!( + scope + .dims(&VarName::new("machine.plug.pin.i")) + .expect("aggregate shape should be inferred from scalar descendants"), + vec![3] + ); + } + + #[test] + fn aggregate_dims_append_descendant_variable_dimensions() { + let span = test_span(); + let mut dae_model = dae::Dae::default(); + for component_index in [1, 2] { + dae_model.variables.algebraics.insert( + VarName::new(format!( + "springDamper.angleToTorque1.move_w[{component_index}].u" + )), + dae::Variable { + name: VarName::new(format!( + "springDamper.angleToTorque1.move_w[{component_index}].u" + )), + dims: vec![2], + component_ref: Some(component_ref_with_span( + vec![ + part("springDamper", Vec::new()), + part("angleToTorque1", Vec::new()), + part("move_w", vec![index(component_index)]), + part("u", Vec::new()), + ], + span, + )), + ..rumoca_ir_dae::Variable::empty_with_span(span) + }, + ); + } + let scope = DaeVariableScope::new(&dae_model); + + assert_eq!( + scope + .dims(&VarName::new("springDamper.angleToTorque1.move_w.u")) + .expect("aggregate shape should include component and leaf variable dimensions"), + vec![2, 2] + ); + } + + #[test] + fn indexed_reference_dims_preserves_parent_component_indices() { + let mut dae_model = dae::Dae::default(); + dae_model.variables.algebraics.insert( + VarName::new("analysatorAC.iH1[1].product2.u"), + dae::Variable { + name: VarName::new("analysatorAC.iH1[1].product2.u"), + dims: vec![2], + component_ref: Some(component_ref(vec![ + part("analysatorAC", Vec::new()), + part("iH1", vec![index(1)]), + part("product2", Vec::new()), + part("u", Vec::new()), + ])), + ..rumoca_ir_dae::Variable::empty_with_span(test_span()) + }, + ); + let reference = Reference::from_component_reference(component_ref(vec![ + part("analysatorAC", Vec::new()), + part("iH1", vec![index(1)]), + part("product2", Vec::new()), + part("u", vec![index(2)]), + ])); + let scope = DaeVariableScope::new(&dae_model); + + assert_eq!( + scope + .shape_for_reference(&reference) + .expect("indexed component element should resolve"), + DaeVariableShape::Dimensions(Vec::new()) + ); + } } diff --git a/crates/rumoca-phase-typecheck/src/instanced.rs b/crates/rumoca-phase-typecheck/src/instanced.rs index 5c5c1c00d..e373dcfb9 100644 --- a/crates/rumoca-phase-typecheck/src/instanced.rs +++ b/crates/rumoca-phase-typecheck/src/instanced.rs @@ -43,6 +43,7 @@ impl TypeChecker { .iter() .map(|(def_id, name)| (*def_id, name.clone())) .collect(); + self.expandable_connector_defs = Self::collect_expandable_connector_defs(tree); self.eval_ctx = rumoca_eval_ast::eval::TypeCheckEvalContext::new(); let (type_table, type_ids_by_def_id) = match self.build_type_context(tree) { Ok(context) => context, @@ -134,9 +135,10 @@ impl TypeChecker { // Uses scope-aware evaluation so `nout = max(size(deltaq,1))` in component // `kinematicPTP` can resolve `deltaq` as `kinematicPTP.deltaq`. for instance_data in overlay.components.values() { - let path = instance_data.qualified_name.to_component_path(); - let name = path.to_flat_string(); - let scope = path + let name = instance_data.qualified_name.to_flat_string(); + let scope = instance_data + .qualified_name + .to_component_path() .parent() .map(|path| path.to_flat_string()) .unwrap_or_default(); @@ -266,6 +268,7 @@ impl TypeChecker { let prev_scope_types = std::mem::take(&mut self.current_component_types); let prev_scope_shapes = std::mem::take(&mut self.current_component_shapes); + let prev_expandable_member_surfaces = std::mem::take(&mut self.expandable_member_surfaces); let (full_prefix, short_model) = Self::instanced_scope_prefixes(model_name); self.current_component_types = Self::build_instanced_component_type_scope(overlay, &full_prefix, &short_model); @@ -273,6 +276,7 @@ impl TypeChecker { Self::build_instanced_component_shape_scope(overlay, &full_prefix, &short_model); self.check_component_modifier_types_in_class(model_class, type_table); + self.collect_expandable_member_surfaces(model_class, type_table); walk_equations(self, &model_class.equations, type_table); walk_equations(self, &model_class.initial_equations, type_table); // Note: component *bindings* are not walked here. Binding @@ -283,5 +287,6 @@ impl TypeChecker { self.current_component_types = prev_scope_types; self.current_component_shapes = prev_scope_shapes; + self.expandable_member_surfaces = prev_expandable_member_surfaces; } } diff --git a/crates/rumoca-phase-typecheck/src/lib.rs b/crates/rumoca-phase-typecheck/src/lib.rs index 0d1f79993..ae80829ab 100644 --- a/crates/rumoca-phase-typecheck/src/lib.rs +++ b/crates/rumoca-phase-typecheck/src/lib.rs @@ -283,6 +283,17 @@ pub struct TypeChecker { /// Keys are class DefIds; values map component member names to their TypeIds /// (including inherited members, with extends `break` names removed). component_modifier_member_types: HashMap>, + /// DefIds of connector classes declared with `expandable` (MLS §9.1.3). + /// + /// Expandable connector members are inferred from connect-equations and + /// component references, so missing members are not ordinary static + /// member errors. + expandable_connector_defs: HashSet, + /// Expandable connector member references visible in the current class. + /// + /// This is collected before validating connect equations so equation order + /// does not decide whether a dynamic bus member is considered synthesized. + expandable_member_surfaces: HashSet, /// Type aliases whose targets could not be resolved during type-table /// construction (e.g. an MSL alias into a library that is not loaded). /// @@ -292,6 +303,15 @@ pub struct TypeChecker { deferred_alias_errors: HashMap, } +struct ClassOverrideAliasContext<'a> { + tree: &'a ClassTree, + component_index: &'a HashMap, + comp_scope: &'a str, + active_alias: Option<&'a str>, + eval_ctx: &'a mut rumoca_eval_ast::eval::TypeCheckEvalContext, + cleared_alias_scopes: &'a mut HashSet, +} + impl TypeChecker { /// Create a new type checker. pub fn new() -> Self { @@ -307,6 +327,8 @@ impl TypeChecker { current_component_shapes: HashMap::new(), component_modifier_targets: HashMap::new(), component_modifier_member_types: HashMap::new(), + expandable_connector_defs: HashSet::new(), + expandable_member_surfaces: HashSet::new(), deferred_alias_errors: HashMap::new(), } } @@ -344,6 +366,7 @@ impl TypeChecker { .iter() .map(|(def_id, name)| (*def_id, name.clone())) .collect(); + self.expandable_connector_defs = Self::collect_expandable_connector_defs(tree); let (type_table, type_ids_by_def_id) = match self.build_type_context(tree) { Ok(context) => context, Err(error) => { @@ -395,12 +418,19 @@ impl TypeChecker { .collect(); const MAX_PASSES: usize = 5; + let mut cleared_alias_scopes = HashSet::::new(); for _ in 0..MAX_PASSES { let prev = ctx.integers.len() + ctx.dimensions.len() + ctx.reals.len() + ctx.booleans.len(); for data in overlay.components.values() { - Self::apply_instance_class_overrides(tree, &component_index, data, ctx); + Self::apply_instance_class_overrides( + tree, + &component_index, + data, + ctx, + &mut cleared_alias_scopes, + ); } let new = @@ -417,11 +447,23 @@ impl TypeChecker { } } + fn collect_expandable_connector_defs(tree: &ClassTree) -> HashSet { + tree.name_map + .values() + .filter_map(|def_id| tree.get_class_by_def_id(*def_id)) + .filter(|class| { + class.expandable && matches!(class.class_type, rumoca_core::ClassType::Connector) + }) + .filter_map(|class| class.def_id) + .collect() + } + fn apply_instance_class_overrides( tree: &ClassTree, component_index: &HashMap, data: &rumoca_ir_ast::InstanceData, ctx: &mut rumoca_eval_ast::eval::TypeCheckEvalContext, + cleared_alias_scopes: &mut HashSet, ) { if data.class_overrides.is_empty() { return; @@ -433,51 +475,63 @@ impl TypeChecker { let active_alias = Self::component_active_alias(data); for class_override in data.class_overrides.values() { - Self::apply_class_override_alias( + let mut alias_ctx = ClassOverrideAliasContext { tree, component_index, - &comp_scope, - active_alias.as_deref(), + comp_scope: &comp_scope, + active_alias: active_alias.as_deref(), + eval_ctx: ctx, + cleared_alias_scopes, + }; + Self::apply_class_override_alias( + &mut alias_ctx, &class_override.alias, class_override.target_def_id, - ctx, ); } } fn apply_class_override_alias( - tree: &ClassTree, - component_index: &HashMap, - comp_scope: &str, - active_alias: Option<&str>, + alias_ctx: &mut ClassOverrideAliasContext<'_>, alias: &str, def_id: DefId, - ctx: &mut rumoca_eval_ast::eval::TypeCheckEvalContext, ) { if Self::try_apply_forwarded_parent_alias_constants( - tree, - component_index, - comp_scope, - active_alias, + alias_ctx.tree, + alias_ctx.component_index, + alias_ctx.comp_scope, + alias_ctx.active_alias, alias, def_id, - ctx, + alias_ctx.eval_ctx, ) { return; } - let is_active_alias = active_alias == Some(alias); + let is_active_alias = alias_ctx.active_alias == Some(alias); - let alias_scope = format!("{comp_scope}.{alias}"); + let alias_scope = format!("{}.{alias}", alias_ctx.comp_scope); // MLS §7.3: instance-level redeclare overrides must replace inherited/default // package constants in the local alias scope. - Self::clear_alias_scope_values(ctx, &alias_scope); - Self::extract_override_class_constants(tree, &alias_scope, def_id, ctx); + if alias_ctx.cleared_alias_scopes.insert(alias_scope.clone()) { + Self::clear_alias_scope_values(alias_ctx.eval_ctx, &alias_scope); + } + Self::extract_override_class_constants( + alias_ctx.tree, + &alias_scope, + def_id, + alias_ctx.eval_ctx, + ); // For declarations like `Medium.BaseProperties medium`, expose // unqualified constants (`medium.nX`) from the active alias only. if is_active_alias { - Self::extract_override_class_constants(tree, comp_scope, def_id, ctx); + Self::extract_override_class_constants( + alias_ctx.tree, + alias_ctx.comp_scope, + def_id, + alias_ctx.eval_ctx, + ); } } @@ -1067,6 +1121,7 @@ impl TypeChecker { fn is_connector_alias_wrapper(class: &ClassDef) -> bool { matches!(class.class_type, rumoca_core::ClassType::Connector) + && !class.expandable && class.extends.len() == 1 && class.classes.is_empty() && class.components.is_empty() diff --git a/crates/rumoca-phase-typecheck/src/tests.rs b/crates/rumoca-phase-typecheck/src/tests.rs index 1d33993e3..99b5efe30 100644 --- a/crates/rumoca-phase-typecheck/src/tests.rs +++ b/crates/rumoca-phase-typecheck/src/tests.rs @@ -96,6 +96,268 @@ fn integer_fold_overflow_emits_warning() { ); } +#[test] +fn expandable_connector_dynamic_member_reference_typechecks() { + let diagnostics = typecheck_diagnostics( + r#" + connector SignalBus + end SignalBus; + + expandable connector WeatherBus + extends SignalBus; + end WeatherBus; + + model Test + WeatherBus bus; + equation + bus.TDryBul = 293.15; + end Test; + "#, + ); + + assert!( + diagnostics + .iter() + .all(|diag| diag.code.as_deref() != Some("ET001")), + "expandable connector dynamic members must not be rejected as ordinary unknown members: {diagnostics:?}" + ); +} + +#[test] +fn expandable_connector_connect_uses_class_visible_dynamic_member_surface() { + let diagnostics = typecheck_diagnostics( + r#" + connector RealInput = input Real; + + connector SignalBus + end SignalBus; + + expandable connector WeatherBus + extends SignalBus; + end WeatherBus; + + model OutdoorAirSource + RealInput T_in; + end OutdoorAirSource; + + model Test + WeatherBus bus; + OutdoorAirSource source; + equation + bus.TDryBul = 293.15; + connect(bus.TDryBul, source.T_in); + end Test; + "#, + ); + + assert!( + diagnostics + .iter() + .all(|diag| diag.code.as_deref() != Some("ET001")), + "connect should accept expandable members already visible in the class surface: {diagnostics:?}" + ); +} + +#[test] +fn expandable_connector_connect_unknown_member_still_fails_typecheck() { + let diagnostics = typecheck_diagnostics( + r#" + connector RealInput = input Real; + + connector SignalBus + end SignalBus; + + expandable connector WeatherBus + extends SignalBus; + end WeatherBus; + + model OutdoorAirSource + RealInput T_in; + end OutdoorAirSource; + + model Test + WeatherBus bus; + OutdoorAirSource source; + equation + connect(bus.TDryBul, source.T_in); + end Test; + "#, + ); + + assert!( + diagnostics.iter().any(|diag| { + diag.code.as_deref() == Some("ET001") + && diag.message.contains("unknown member `TDryBul`") + }), + "connect should still reject expandable members with no visible synthesis source: {diagnostics:?}" + ); +} + +#[test] +fn instanced_expandable_connector_connect_uses_visible_dynamic_member_surface() { + let source = r#" + connector RealInput = input Real; + + connector SignalBus + end SignalBus; + + expandable connector WeatherBus + extends SignalBus; + end WeatherBus; + + model OutdoorAirSource + RealInput T_in; + end OutdoorAirSource; + + model Test + WeatherBus bus; + OutdoorAirSource source; + equation + bus.TDryBul = 293.15; + connect(bus.TDryBul, source.T_in); + end Test; + "#; + + let parsed = parse(source); + let resolved = resolve(parsed).expect("resolve should succeed"); + let tree = resolved.into_inner(); + let test = tree + .definitions + .classes + .get("Test") + .expect("Test class should exist"); + let source_class = tree + .definitions + .classes + .get("OutdoorAirSource") + .expect("OutdoorAirSource class should exist"); + + let mut overlay = InstanceOverlay::new(); + add_test_instance( + &mut overlay, + "Test.bus", + test.components.get("bus").expect("bus declaration"), + None, + ); + add_test_instance( + &mut overlay, + "Test.source", + test.components.get("source").expect("source declaration"), + None, + ); + add_test_instance( + &mut overlay, + "Test.source.T_in", + source_class + .components + .get("T_in") + .expect("T_in declaration"), + None, + ); + + typecheck_instanced(&tree, &mut overlay, "Test") + .expect("instanced expandable connect should use visible dynamic member surface"); +} + +#[test] +fn instanced_expandable_connector_connect_unknown_member_still_fails_typecheck() { + let source = r#" + connector RealInput = input Real; + + connector SignalBus + end SignalBus; + + expandable connector WeatherBus + extends SignalBus; + end WeatherBus; + + model OutdoorAirSource + RealInput T_in; + end OutdoorAirSource; + + model Test + WeatherBus bus; + OutdoorAirSource source; + equation + connect(bus.TDryBul, source.T_in); + end Test; + "#; + + let parsed = parse(source); + let resolved = resolve(parsed).expect("resolve should succeed"); + let tree = resolved.into_inner(); + let test = tree + .definitions + .classes + .get("Test") + .expect("Test class should exist"); + let source_class = tree + .definitions + .classes + .get("OutdoorAirSource") + .expect("OutdoorAirSource class should exist"); + + let mut overlay = InstanceOverlay::new(); + add_test_instance( + &mut overlay, + "Test.bus", + test.components.get("bus").expect("bus declaration"), + None, + ); + add_test_instance( + &mut overlay, + "Test.source", + test.components.get("source").expect("source declaration"), + None, + ); + add_test_instance( + &mut overlay, + "Test.source.T_in", + source_class + .components + .get("T_in") + .expect("T_in declaration"), + None, + ); + + let err = typecheck_instanced(&tree, &mut overlay, "Test") + .expect_err("instanced expandable connect should reject unsynthesized members"); + assert!( + err.iter().any(|diag| { + diag.code.as_deref() == Some("ET001") + && diag.message.contains("unknown member `TDryBul`") + }), + "expected ET001 for unsynthesized expandable connect member, got: {err:?}" + ); +} + +#[test] +fn ordinary_connector_unknown_member_still_fails_typecheck() { + let diagnostics = typecheck_diagnostics( + r#" + connector SignalBus + end SignalBus; + + connector FixedBus + extends SignalBus; + end FixedBus; + + model Test + FixedBus bus; + equation + bus.TDryBul = 293.15; + end Test; + "#, + ); + + assert!( + diagnostics + .iter() + .any(|diag| diag.code.as_deref() == Some("ET001")), + "ordinary connector unknown members should remain strict: {diagnostics:?}" + ); +} + #[test] fn model_qualified_nested_package_constants_evaluate_dimensions() { let diagnostics = typecheck_diagnostics( diff --git a/crates/rumoca-phase-typecheck/src/tests/record_constructor_alias_tests.rs b/crates/rumoca-phase-typecheck/src/tests/record_constructor_alias_tests.rs index bbd315967..e0576a387 100644 --- a/crates/rumoca-phase-typecheck/src/tests/record_constructor_alias_tests.rs +++ b/crates/rumoca-phase-typecheck/src/tests/record_constructor_alias_tests.rs @@ -116,6 +116,17 @@ fn test_propagate_alias_map_copies_root_and_prefixed_fields() { assert_eq!(values.get("dst2.nX"), None); } +#[test] +fn test_alias_head_ignores_dots_inside_subscripts() { + assert_eq!(TypeChecker::alias_head("src.stackData.cellData"), "src"); + assert_eq!( + TypeChecker::alias_head("plug[data.medium].port"), + "plug[data.medium]" + ); + assert_eq!(TypeChecker::alias_head("ch[1].chi.vol2"), "ch[1]"); + assert_eq!(TypeChecker::alias_head("standalone"), "standalone"); +} + #[test] fn test_extract_simple_path_preserves_subscripted_component_refs() { let expr = Expression::ComponentReference(ComponentReference { diff --git a/crates/rumoca-phase-typecheck/src/typechecker/late_methods.rs b/crates/rumoca-phase-typecheck/src/typechecker/late_methods.rs index 2fd962f9d..ce6be3f46 100644 --- a/crates/rumoca-phase-typecheck/src/typechecker/late_methods.rs +++ b/crates/rumoca-phase-typecheck/src/typechecker/late_methods.rs @@ -3,6 +3,13 @@ use super::traversal_adapter::{ }; use super::*; use rumoca_core::ComponentPath; +use rumoca_ir_ast::Visitor; +use std::ops::ControlFlow; + +// SPEC_0021 file-size exception: late typecheck owns cross-cutting checks over +// resolved class structure; split plan is to move connector membership and +// modifier validation into sibling typechecker modules as the current FMU +// correctness path stabilizes. #[path = "equation_shape.rs"] mod equation_shape; @@ -52,6 +59,16 @@ impl TypeCheckTraversalCallbacks for TypeChecker { self.check_power_operators(rhs, type_table); } + fn on_connect_equation( + &mut self, + lhs: &rumoca_ir_ast::ComponentReference, + rhs: &rumoca_ir_ast::ComponentReference, + type_table: &TypeTable, + ) { + self.validate_expandable_connect_reference(lhs, type_table); + self.validate_expandable_connect_reference(rhs, type_table); + } + fn on_expression_function_call( &mut self, comp: &rumoca_ir_ast::ComponentReference, @@ -156,6 +173,7 @@ impl TypeChecker { } } + #[allow(dead_code)] pub(crate) fn alias_field_key_range<'a>( sorted_keys: &'a [String], target_prefix: &str, @@ -554,7 +572,24 @@ impl TypeChecker { let name = name_path.to_flat_string(); let scope = Self::enclosing_scope_name(&name_path); - // Try to infer from binding first + // MLS §7.2.4: modification expressions are evaluated in the lexical + // scope where the modifier was written. For array-component modifiers + // the resolved binding may be an element-indexed value (`row[1]`), while + // the preserved source expression (`row`) carries the modified field's + // array shape. + if instance_data.binding_from_modification + && let Some(dims) = self.infer_modifier_source_dims(instance_data) + { + return Self::apply_inferred_instance_dims( + instance_data, + &mut self.eval_ctx, + &name, + dims, + ); + } + + // Try to infer from binding first for declaration bindings and + // modification bindings without usable source metadata. if let Some(ref binding) = instance_data.binding && let Some(dims) = rumoca_eval_ast::eval::infer_dimensions_from_binding_with_scope( binding, @@ -569,6 +604,14 @@ impl TypeChecker { dims, ); } + if let Some(dims) = self.infer_modifier_source_dims(instance_data) { + return Self::apply_inferred_instance_dims( + instance_data, + &mut self.eval_ctx, + &name, + dims, + ); + } // Fallback: try to infer from start value of a record element binding if let Some(ref start) = instance_data.start @@ -589,6 +632,23 @@ impl TypeChecker { false } + fn infer_modifier_source_dims( + &self, + instance_data: &rumoca_ir_ast::InstanceData, + ) -> Option> { + let source = instance_data.binding_source.as_ref()?; + let source_scope = instance_data + .binding_source_scope + .as_ref()? + .to_flat_string(); + let dims = rumoca_eval_ast::eval::infer_dimensions_from_binding_with_scope( + source, + &self.eval_ctx, + &source_scope, + )?; + (!dims.is_empty()).then_some(dims) + } + fn apply_inferred_instance_dims( instance_data: &mut rumoca_ir_ast::InstanceData, eval_ctx: &mut rumoca_eval_ast::eval::TypeCheckEvalContext, @@ -759,6 +819,7 @@ impl TypeChecker { ); } self.current_component_types = scope_types; + let prev_expandable_member_surfaces = std::mem::take(&mut self.expandable_member_surfaces); let mut scope_shapes = HashMap::new(); for (name, comp) in &class.components { let shape = if !comp.shape.is_empty() { @@ -785,6 +846,11 @@ impl TypeChecker { // Mark structural parameters (MLS §18.3) self.mark_structural_parameters(class); + // Expandable connector membership is a class-level equation surface: + // equation order must not decide whether a later connect-reference is + // considered already synthesized. + self.collect_expandable_member_surfaces(class, type_table); + // Type check equations walk_equations(self, &class.equations, type_table); walk_equations(self, &class.initial_equations, type_table); @@ -804,6 +870,7 @@ impl TypeChecker { // Restore parent class scope. self.current_component_types = prev_scope_types; + self.expandable_member_surfaces = prev_expandable_member_surfaces; } /// Type check a component declaration. @@ -1603,7 +1670,7 @@ impl TypeChecker { { return; } - if !Self::is_strict_component_member_owner(type_table, base_type) { + if !self.is_strict_component_member_owner(type_table, base_type) { return; } let Some(location) = base.get_location() else { @@ -1623,6 +1690,73 @@ impl TypeChecker { ); } + pub(crate) fn collect_expandable_member_surfaces( + &mut self, + class: &ClassDef, + type_table: &TypeTable, + ) { + let mut collector = ExpandableMemberSurfaceCollector { + checker: self, + type_table, + }; + let _ = collector.visit_each(&class.equations, rumoca_ir_ast::Visitor::visit_equation); + let _ = collector.visit_each( + &class.initial_equations, + rumoca_ir_ast::Visitor::visit_equation, + ); + for statements in &class.algorithms { + let _ = collector.visit_each(statements, rumoca_ir_ast::Visitor::visit_statement); + } + for statements in &class.initial_algorithms { + let _ = collector.visit_each(statements, rumoca_ir_ast::Visitor::visit_statement); + } + } + + fn record_expandable_member_surface( + &mut self, + comp: &rumoca_ir_ast::ComponentReference, + type_table: &TypeTable, + ) { + let Some((mut current_type, prefix_len)) = self.find_component_ref_prefix_type(comp) else { + return; + }; + for (idx, part) in comp.parts.iter().enumerate().skip(prefix_len) { + let member_name = part.ident.text.as_ref(); + if let Some(next_type) = + self.lookup_component_member_type(current_type, member_name, type_table) + { + current_type = next_type; + continue; + } + if self.is_expandable_connector_type(type_table, current_type) { + self.expandable_member_surfaces + .insert(component_reference_member_surface_key(comp, idx)); + } + return; + } + } + + fn record_expandable_field_access_surface( + &mut self, + base: &Expression, + field: &str, + type_table: &TypeTable, + ) { + let Some(base_type) = self.infer_expression_type(base, type_table) else { + return; + }; + if self + .lookup_component_member_type(base_type, field, type_table) + .is_some() + { + return; + } + if self.is_expandable_connector_type(type_table, base_type) { + self.expandable_member_surfaces + .insert(format!("{base}.{field}")); + } + } + pub(in crate::typechecker) fn resolve_component_reference_type( &self, comp: &rumoca_ir_ast::ComponentReference, @@ -1639,7 +1773,7 @@ impl TypeChecker { let member_name = part.ident.text.to_string(); match self.lookup_component_member_type(current_type, &member_name, type_table) { Some(next_type) => current_type = next_type, - None if !Self::is_strict_component_member_owner(type_table, current_type) => { + None if !self.is_strict_component_member_owner(type_table, current_type) => { return Ok(TypeId::UNKNOWN); } None => { @@ -1659,6 +1793,54 @@ impl TypeChecker { Ok(current_type) } + fn validate_expandable_connect_reference( + &mut self, + comp: &rumoca_ir_ast::ComponentReference, + type_table: &TypeTable, + ) { + let Some((mut current_type, prefix_len)) = self.find_component_ref_prefix_type(comp) else { + return; + }; + if prefix_len == comp.parts.len() { + return; + } + + for (idx, part) in comp.parts.iter().enumerate().skip(prefix_len) { + let member_name = part.ident.text.to_string(); + if let Some(next_type) = + self.lookup_component_member_type(current_type, &member_name, type_table) + { + current_type = next_type; + continue; + } + if !self.is_expandable_connector_type(type_table, current_type) { + return; + } + if self + .expandable_member_surfaces + .contains(&component_reference_member_surface_key(comp, idx)) + { + return; + } + match self.component_reference_member_span(&part.ident.location) { + Ok(span) => self.emit_unknown_component_member( + MissingComponentMember { + owner_type: current_type, + member_name, + reference: comp.to_string(), + span, + }, + type_table, + ), + Err(ComponentReferenceTypeError::MissingSourceContext(error)) => { + self.emit_typecheck_error(error); + } + Err(ComponentReferenceTypeError::MissingMember(_)) => {} + } + return; + } + } + fn component_reference_member_span( &self, location: &rumoca_core::Location, @@ -1746,23 +1928,37 @@ impl TypeChecker { names } - fn is_strict_component_member_owner(type_table: &TypeTable, owner_type: TypeId) -> bool { + fn is_strict_component_member_owner(&self, type_table: &TypeTable, owner_type: TypeId) -> bool { + let owner_root = self.resolve_type_root(type_table, owner_type); + let Some(Type::Class(class_type)) = type_table.get(owner_root) else { + return false; + }; + if class_type.kind == ClassKind::Connector + && self.expandable_connector_defs.contains(&class_type.def_id) + { + return false; + } matches!( - type_table.get(Self::resolve_alias_root(type_table, owner_type)), - Some(Type::Class(class_type)) - if matches!( - class_type.kind, - ClassKind::Class - | ClassKind::Model - | ClassKind::Block - | ClassKind::Record - | ClassKind::Connector - | ClassKind::Type - | ClassKind::Operator - ) + class_type.kind, + ClassKind::Class + | ClassKind::Model + | ClassKind::Block + | ClassKind::Record + | ClassKind::Connector + | ClassKind::Type + | ClassKind::Operator ) } + fn is_expandable_connector_type(&self, type_table: &TypeTable, owner_type: TypeId) -> bool { + let owner_root = self.resolve_type_root(type_table, owner_type); + let Some(Type::Class(class_type)) = type_table.get(owner_root) else { + return false; + }; + class_type.kind == ClassKind::Connector + && self.expandable_connector_defs.contains(&class_type.def_id) + } + pub(in crate::typechecker) fn emit_unknown_component_member( &mut self, missing: MissingComponentMember, @@ -1926,3 +2122,62 @@ impl TypeChecker { self.diagnostics } } + +struct ExpandableMemberSurfaceCollector<'a> { + checker: &'a mut TypeChecker, + type_table: &'a TypeTable, +} + +impl ExpandableMemberSurfaceCollector<'_> { + fn record_expression_surface(&mut self, expression: &Expression) { + let Expression::FieldAccess { base, field, .. } = expression else { + return; + }; + self.checker + .record_expandable_field_access_surface(base, field, self.type_table); + } + + fn record_reference_surface( + &mut self, + comp: &rumoca_ir_ast::ComponentReference, + ctx: rumoca_ir_ast::ComponentReferenceContext, + ) { + if matches!( + ctx, + rumoca_ir_ast::ComponentReferenceContext::EquationConnectLhs + | rumoca_ir_ast::ComponentReferenceContext::EquationConnectRhs + ) { + return; + } + self.checker + .record_expandable_member_surface(comp, self.type_table); + } +} + +impl rumoca_ir_ast::Visitor for ExpandableMemberSurfaceCollector<'_> { + fn visit_expression(&mut self, expression: &Expression) -> ControlFlow<()> { + self.record_expression_surface(expression); + rumoca_ir_ast::walk_expression_default(self, expression) + } + + fn visit_component_reference_ctx( + &mut self, + comp: &rumoca_ir_ast::ComponentReference, + ctx: rumoca_ir_ast::ComponentReferenceContext, + ) -> ControlFlow<()> { + self.record_reference_surface(comp, ctx); + rumoca_ir_ast::walk_component_reference_default(self, comp) + } +} + +fn component_reference_member_surface_key( + comp: &rumoca_ir_ast::ComponentReference, + member_idx: usize, +) -> String { + comp.parts + .iter() + .take(member_idx + 1) + .map(ToString::to_string) + .collect::>() + .join(".") +} diff --git a/crates/rumoca-phase-typecheck/src/typechecker/record_aliases.rs b/crates/rumoca-phase-typecheck/src/typechecker/record_aliases.rs index b5eab99cb..56ecf6cfa 100644 --- a/crates/rumoca-phase-typecheck/src/typechecker/record_aliases.rs +++ b/crates/rumoca-phase-typecheck/src/typechecker/record_aliases.rs @@ -184,22 +184,28 @@ impl TypeChecker { if record_aliases.is_empty() || values.is_empty() { return false; } - let mut sorted_keys: Vec = values.keys().cloned().collect(); - sorted_keys.sort_unstable(); let mut updates: rustc_hash::FxHashMap = rustc_hash::FxHashMap::default(); for (alias_source, alias_target) in record_aliases { Self::queue_alias_root_update(alias_source, alias_target, values, &mut updates); + } + + let mut alias_prefixes_by_head: rustc_hash::FxHashMap> = + rustc_hash::FxHashMap::default(); + for (alias_source, alias_target) in record_aliases { let target_prefix = format!("{alias_target}."); - for field_name in Self::alias_field_key_range(&sorted_keys, &target_prefix) { - Self::queue_alias_field_update( - alias_source, - &target_prefix, - field_name, - values, - &mut updates, - ); - } + alias_prefixes_by_head + .entry(Self::alias_head(&target_prefix).to_string()) + .or_default() + .push((alias_source.as_str(), target_prefix)); + } + let value_keys = values.keys().cloned().collect::>(); + for field_name in value_keys { + let Some(alias_prefixes) = alias_prefixes_by_head.get(Self::alias_head(&field_name)) + else { + continue; + }; + Self::queue_alias_field_updates(alias_prefixes, &field_name, values, &mut updates); } let mut progress = false; @@ -211,4 +217,38 @@ impl TypeChecker { } progress } + + fn queue_alias_field_updates( + alias_prefixes: &[(&str, String)], + field_name: &str, + values: &rustc_hash::FxHashMap, + updates: &mut rustc_hash::FxHashMap, + ) { + for (alias_source, target_prefix) in alias_prefixes { + if !field_name.starts_with(target_prefix) { + continue; + } + Self::queue_alias_field_update( + alias_source, + target_prefix, + field_name, + values, + updates, + ); + } + } + + #[allow(dead_code)] + pub(crate) fn alias_head(path: &str) -> &str { + let mut bracket_depth = 0usize; + for (index, byte) in path.bytes().enumerate() { + match byte { + b'[' => bracket_depth += 1, + b']' => bracket_depth = bracket_depth.saturating_sub(1), + b'.' if bracket_depth == 0 => return &path[..index], + _ => {} + } + } + path + } } diff --git a/crates/rumoca-phase-typecheck/src/typechecker/traversal_adapter.rs b/crates/rumoca-phase-typecheck/src/typechecker/traversal_adapter.rs index 93c93369e..9599e5dfd 100644 --- a/crates/rumoca-phase-typecheck/src/typechecker/traversal_adapter.rs +++ b/crates/rumoca-phase-typecheck/src/typechecker/traversal_adapter.rs @@ -23,6 +23,15 @@ pub(crate) trait TypeCheckTraversalCallbacks { /// Called after both sides of a simple equation are traversed. fn on_simple_equation(&mut self, lhs: &Expression, rhs: &Expression, type_table: &TypeTable); + /// Called for a connect equation after both references are traversed. + fn on_connect_equation( + &mut self, + _lhs: &ComponentReference, + _rhs: &ComponentReference, + _type_table: &TypeTable, + ) { + } + /// Called after an expression-form function call and all arguments are traversed. fn on_expression_function_call( &mut self, @@ -70,6 +79,24 @@ impl Visitor for TypeCheckTraversal<'_, C> { ast::visitor::walk_equation_default(self, equation) } + fn visit_connect( + &mut self, + lhs: &ComponentReference, + rhs: &ComponentReference, + ) -> ControlFlow<()> { + self.visit_component_reference_ctx( + lhs, + ast::ComponentReferenceContext::EquationConnectLhs, + )?; + self.visit_component_reference_ctx( + rhs, + ast::ComponentReferenceContext::EquationConnectRhs, + )?; + self.callbacks + .on_connect_equation(lhs, rhs, self.type_table); + ControlFlow::Continue(()) + } + fn visit_simple_equation(&mut self, lhs: &Expression, rhs: &Expression) -> ControlFlow<()> { self.visit_expression(lhs)?; self.visit_expression(rhs)?; diff --git a/crates/rumoca-sim/Cargo.toml b/crates/rumoca-sim/Cargo.toml index 226dd7db7..47e9d5c6a 100644 --- a/crates/rumoca-sim/Cargo.toml +++ b/crates/rumoca-sim/Cargo.toml @@ -80,3 +80,6 @@ rumoca-input-gamepad = { workspace = true, optional = true } [lints] workspace = true + +[dev-dependencies] +bincode = { workspace = true } diff --git a/crates/rumoca-sim/src/gpu_initialization.rs b/crates/rumoca-sim/src/gpu_initialization.rs new file mode 100644 index 000000000..018d52980 --- /dev/null +++ b/crates/rumoca-sim/src/gpu_initialization.rs @@ -0,0 +1,1033 @@ +//! Lean, backend-neutral initial-condition settlement for GPU preparation. +//! +//! This path intentionally accepts only lowering-proven direct assignments. +//! General nonlinear or coupled initialization remains unsupported here rather +//! than silently using a finite-difference CPU projection. + +use rumoca_ir_solve as solve; + +const INITIAL_RESIDUAL_TOLERANCE: f64 = 1.0e-9; + +#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)] +pub struct GpuInitializationMetrics { + pub residual_evaluations: usize, + pub passes: usize, + pub temporary_values: usize, +} + +#[derive(Debug, thiserror::Error)] +pub enum GpuInitializationError { + #[error("GPU initial projection does not support {feature} (row={row}, span={span:?})")] + Unsupported { + feature: &'static str, + row: usize, + span: Option, + }, + #[error("GPU initial projection is malformed: {message} (row={row}, span={span:?})")] + Malformed { + message: String, + row: usize, + span: Option, + }, + #[error( + "GPU initial projection {kind} did not settle (row={row}, value={value:.6e}, span={span:?})" + )] + NonConverged { + kind: &'static str, + row: usize, + value: f64, + span: Option, + }, + #[error("GPU initial projection evaluation failed: {message} (span={span:?})")] + Evaluation { + message: String, + span: Option, + }, +} + +#[derive(Debug)] +pub struct GpuInitializationResult { + pub y0: Vec, + pub p0: Vec, + pub metrics: GpuInitializationMetrics, +} + +/// Settle a GPU-prepared model without introducing a continuous solver/JVP +/// payload. The artifact is complete-or-error: input vectors are cloned and +/// never exposed after a failed evaluation or residual check. +pub fn settle_gpu_initial_conditions( + model: &solve::SolveModel, + t_start: f64, +) -> Result { + let initialization = &model.problem.initialization; + reject_unsupported_runtime_features(model)?; + let mut y0 = model.initial_y.clone(); + let p0 = model.parameters.clone(); + ensure_finite(&y0, "initial y", None)?; + ensure_finite(&p0, "initial p", None)?; + validate_assignment_shape(initialization, model.initial_y.len())?; + if initialization.residual.is_empty() { + return Ok(GpuInitializationResult { + y0, + p0, + metrics: GpuInitializationMetrics::default(), + }); + } + let runtime_state = rumoca_eval_solve::SimulationRuntimeState::new(); + let eval_context = rumoca_eval_solve::RowEvalContext { + external_tables: Some(model.external_tables.as_slice()), + runtime_state: Some(&runtime_state), + ..Default::default() + }; + let mut worst = (0usize, 0.0f64, None); + let mut native_metrics = rumoca_eval_solve::MapEvaluationMetrics::default(); + for family in &initialization.direct_families { + execute_direct_family( + family, + DirectFamilyExecution { + initialization, + y: &mut y0, + p: &p0, + t: t_start, + context: eval_context, + apply: true, + worst: &mut worst, + metrics: &mut native_metrics, + }, + )?; + } + ensure_finite(&y0, "settled y", None)?; + worst = (0usize, 0.0f64, None); + for family in &initialization.direct_families { + execute_direct_family( + family, + DirectFamilyExecution { + initialization, + y: &mut y0, + p: &p0, + t: t_start, + context: eval_context, + apply: false, + worst: &mut worst, + metrics: &mut native_metrics, + }, + )?; + } + if !worst.1.is_finite() || worst.1.abs() > INITIAL_RESIDUAL_TOLERANCE { + return Err(GpuInitializationError::NonConverged { + kind: "residual", + row: worst.0, + value: worst.1, + span: worst.2, + }); + } + Ok(GpuInitializationResult { + y0, + p0, + metrics: GpuInitializationMetrics { + residual_evaluations: 2, + passes: 1, + temporary_values: native_metrics + .temporary_values + .saturating_add(initialization.direct_families.len()), + }, + }) +} + +fn validate_assignment_shape( + initialization: &solve::InitializationSolveSystem, + y_len: usize, +) -> Result<(), GpuInitializationError> { + let required = normalize_target_ranges(&initialization.required_target_ranges, y_len)?; + let fixed = normalize_target_ranges(&initialization.fixed_target_ranges, y_len)?; + let expected = if y_len == 0 { + Vec::new() + } else { + vec![solve::InitializationTargetRange { + start: 0, + end: y_len, + span: required.first().and_then(|range| range.span), + }] + }; + if initialization.residual.is_empty() { + return validate_empty_assignment_shape(initialization, &required, &fixed, &expected); + } + if initialization.direct_families.is_empty() { + return Err(GpuInitializationError::Unsupported { + feature: "non-direct or incomplete initial residual system", + row: 0, + span: None, + }); + } + if !initialization.row_targets.is_empty() { + return Err(GpuInitializationError::Malformed { + message: "compact GPU initialization must not materialize scalar row targets" + .to_string(), + row: 0, + span: None, + }); + } + let mut actual_ranges = fixed; + actual_ranges.extend(validate_direct_node_ownership(initialization)?); + let actual = normalize_target_ranges(&actual_ranges, y_len)?; + if !same_target_coverage(&required, &expected) || !same_target_coverage(&actual, &required) { + return Err(GpuInitializationError::Malformed { + message: "incomplete direct plus fixed-start target union".to_string(), + row: 0, + span: actual + .first() + .and_then(|range| range.span) + .or_else(|| required.first().and_then(|range| range.span)), + }); + } + Ok(()) +} + +fn validate_empty_assignment_shape( + initialization: &solve::InitializationSolveSystem, + required: &[solve::InitializationTargetRange], + fixed: &[solve::InitializationTargetRange], + expected: &[solve::InitializationTargetRange], +) -> Result<(), GpuInitializationError> { + if !initialization.direct_families.is_empty() || !initialization.row_targets.is_empty() { + return Err(GpuInitializationError::Malformed { + message: "empty initial residual cannot own assignment rows".to_string(), + row: 0, + span: initialization + .direct_families + .first() + .map(|family| family.span), + }); + } + if required.is_empty() && fixed.is_empty() { + return Ok(()); + } + if !same_target_coverage(required, expected) || !same_target_coverage(fixed, required) { + return Err(GpuInitializationError::Malformed { + message: "incomplete fixed-start target union".to_string(), + row: 0, + span: fixed + .first() + .and_then(|range| range.span) + .or_else(|| required.first().and_then(|range| range.span)), + }); + } + Ok(()) +} + +fn validate_direct_node_ownership( + initialization: &solve::InitializationSolveSystem, +) -> Result, GpuInitializationError> { + let mut ranges = Vec::with_capacity(initialization.direct_families.len()); + let mut node_owners = vec![None; initialization.residual.nodes.len()]; + for family in &initialization.direct_families { + let Some(owner) = node_owners.get_mut(family.node_index) else { + return Err(GpuInitializationError::Malformed { + message: "direct initial family references a missing residual node".to_string(), + row: family.node_index, + span: Some(family.span), + }); + }; + if owner.replace(family.span).is_some() { + return Err(GpuInitializationError::Malformed { + message: "duplicate direct ownership of one residual node".to_string(), + row: family.node_index, + span: Some(family.span), + }); + } + if !matches!(family.residual_sign, -1 | 1) { + return Err(GpuInitializationError::Malformed { + message: "direct initial family must have a unit residual sign".to_string(), + row: 0, + span: Some(family.span), + }); + } + let Some(solve::ComputeNode::Map { + domain, base_ops, .. + }) = initialization.residual.nodes.get(family.node_index) + else { + return Err(GpuInitializationError::Unsupported { + feature: "non-Map direct initial family", + row: 0, + span: Some(family.span), + }); + }; + if has_random_or_impure_ops(base_ops) { + return Err(GpuInitializationError::Unsupported { + feature: "random or impure direct initial operations", + row: 0, + span: Some(family.span), + }); + } + let dense = solve::TensorOutputMap::dense_contiguous(family.targets.start, domain) + .map_err(|error| GpuInitializationError::Malformed { + message: format!("invalid direct target map: {error:?}"), + row: 0, + span: Some(family.span), + })?; + if family.targets.strides != dense.strides { + return Err(GpuInitializationError::Malformed { + message: "direct target map must be dense and contiguous".to_string(), + row: 0, + span: Some(family.span), + }); + } + let count = domain + .scalar_count() + .map_err(|error| GpuInitializationError::Malformed { + message: format!("invalid direct target domain: {error}"), + row: 0, + span: Some(family.span), + })?; + let end = family.targets.start.checked_add(count).ok_or_else(|| { + GpuInitializationError::Malformed { + message: "direct target range overflow".to_string(), + row: 0, + span: Some(family.span), + } + })?; + ranges.push(solve::InitializationTargetRange { + start: family.targets.start, + end, + span: Some(family.span), + }); + } + if let Some(node_index) = node_owners.iter().position(Option::is_none) { + let span = initialization + .residual + .nodes + .get(node_index) + .and_then(|node| match node { + solve::ComputeNode::Map { span, .. } + | solve::ComputeNode::AffineStencil { span, .. } + | solve::ComputeNode::MatMul { span, .. } + | solve::ComputeNode::LinSolve { span, .. } => Some(*span), + solve::ComputeNode::ScalarPrograms(block) => block.first_source_span(), + }); + return Err(GpuInitializationError::Malformed { + message: "direct initial families must own every residual node".to_string(), + row: node_index, + span, + }); + } + Ok(ranges) +} + +fn normalize_target_ranges( + ranges: &[solve::InitializationTargetRange], + upper_bound: usize, +) -> Result, GpuInitializationError> { + let mut ranges = ranges.to_vec(); + ranges.sort_unstable_by_key(|range| (range.start, range.end)); + let mut normalized: Vec = Vec::with_capacity(ranges.len()); + for range in ranges { + if range.start >= range.end || range.end > upper_bound { + return Err(GpuInitializationError::Malformed { + message: "initial target range is empty or out of bounds".to_string(), + row: 0, + span: range.span, + }); + } + if let Some(last) = normalized.last_mut() { + if range.start < last.end { + return Err(GpuInitializationError::Malformed { + message: "initial target ranges overlap".to_string(), + row: 0, + span: range.span.or(last.span), + }); + } + if range.start == last.end { + last.end = range.end; + continue; + } + } + normalized.push(range); + } + Ok(normalized) +} + +fn same_target_coverage( + left: &[solve::InitializationTargetRange], + right: &[solve::InitializationTargetRange], +) -> bool { + left.len() == right.len() + && left + .iter() + .zip(right) + .all(|(left, right)| left.start == right.start && left.end == right.end) +} + +struct DirectFamilyExecution<'a> { + initialization: &'a solve::InitializationSolveSystem, + y: &'a mut [f64], + p: &'a [f64], + t: f64, + context: rumoca_eval_solve::RowEvalContext<'a>, + apply: bool, + worst: &'a mut (usize, f64, Option), + metrics: &'a mut rumoca_eval_solve::MapEvaluationMetrics, +} + +fn execute_direct_family( + family: &solve::InitializationDirectFamily, + execution: DirectFamilyExecution<'_>, +) -> Result<(), GpuInitializationError> { + let Some(node @ solve::ComputeNode::Map { .. }) = execution + .initialization + .residual + .nodes + .get(family.node_index) + else { + return Err(GpuInitializationError::Unsupported { + feature: "non-Map direct initial family", + row: 0, + span: Some(family.span), + }); + }; + let evaluation = rumoca_eval_solve::eval_map_elements_with_context( + node, + execution.y, + execution.p, + execution.t, + execution.context, + |ordinal, value, y| { + let row = direct_map_index(&family.targets, ordinal, family.span).map_err(|error| { + rumoca_eval_solve::EvalSolveError::InvalidRow { + message: error.to_string(), + span: Some(family.span), + } + })?; + if !value.is_finite() { + return Err(rumoca_eval_solve::EvalSolveError::InvalidRow { + message: format!("non-finite direct initial residual at y[{row}]"), + span: Some(family.span), + }); + } + if execution.apply { + *y.get_mut(row).ok_or_else(|| { + rumoca_eval_solve::EvalSolveError::InvalidRow { + message: format!("direct target y[{row}] is outside the state vector"), + span: Some(family.span), + } + })? -= f64::from(family.residual_sign) * value; + } + if value.abs() > execution.worst.1.abs() { + *execution.worst = (row, value, Some(family.span)); + } + Ok(()) + }, + ) + .map_err(|error| GpuInitializationError::Evaluation { + message: error.to_string(), + span: error.source_span().or(Some(family.span)), + })?; + execution.metrics.elements = execution + .metrics + .elements + .saturating_add(evaluation.elements); + execution.metrics.temporary_values = execution + .metrics + .temporary_values + .max(evaluation.temporary_values); + Ok(()) +} + +fn has_random_or_impure_ops(ops: &[solve::LinearOp]) -> bool { + ops.iter().any(|op| { + matches!( + op, + solve::LinearOp::RandomInitialState { .. } + | solve::LinearOp::RandomResult { .. } + | solve::LinearOp::RandomState { .. } + | solve::LinearOp::ImpureRandomInit { .. } + | solve::LinearOp::ImpureRandom { .. } + | solve::LinearOp::ImpureRandomInteger { .. } + ) + }) +} + +fn direct_map_index( + map: &solve::TensorOutputMap, + ordinal: &[usize], + span: rumoca_core::Span, +) -> Result { + let offset = map + .strides + .iter() + .try_fold(0isize, |total, term| { + total.checked_add( + term.stride + .checked_mul(isize::try_from(*ordinal.get(term.dimension)?).ok()?)?, + ) + }) + .ok_or_else(|| GpuInitializationError::Malformed { + message: "direct target map overflow".to_string(), + row: 0, + span: Some(span), + })?; + map.start + .checked_add_signed(offset) + .ok_or_else(|| GpuInitializationError::Malformed { + message: "direct target map overflow".to_string(), + row: 0, + span: Some(span), + }) +} + +fn reject_unsupported_runtime_features( + model: &solve::SolveModel, +) -> Result<(), GpuInitializationError> { + let problem = &model.problem; + let has_events = !problem.events.root_conditions.is_empty() + || !problem.events.root_relation_memory_targets.is_empty() + || !problem.events.scheduled_root_conditions.is_empty() + || !problem.events.scheduled_time_events.is_empty() + || !problem.events.dynamic_time_event_names.is_empty() + || !problem.events.dynamic_time_event_rhs.is_empty() + || !problem.events.action_conditions.is_empty() + || !problem.events.actions.is_empty(); + let has_discrete = !problem.discrete.runtime_assignment_rhs.is_empty() + || !problem.discrete.rhs.is_empty() + || !problem.discrete.update_targets.is_empty() + || !problem.discrete.pre_modes.is_empty(); + let has_memory = !problem + .solve_layout + .relation_memory_parameter_indices + .is_empty() + || !problem.solve_layout.pre_param_bindings.is_empty(); + if has_events + || has_discrete + || has_memory + || !problem.clocks.periodic_event_schedules.is_empty() + { + return Err(GpuInitializationError::Unsupported { + feature: "event, discrete, pre, relation-memory, or clock initialization", + row: 0, + span: None, + }); + } + Ok(()) +} + +fn ensure_finite( + values: &[f64], + kind: &'static str, + span: Option, +) -> Result<(), GpuInitializationError> { + if let Some((row, value)) = values + .iter() + .copied() + .enumerate() + .find(|(_, value)| !value.is_finite()) + { + return Err(GpuInitializationError::NonConverged { + kind, + row, + value, + span, + }); + } + let _ = kind; + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use rumoca_ir_solve::{ + AffineStencilIndexStrideTerm, BinaryOp, ComputeBlock, ComputeNode, LinearOp, + TensorNodeMetadata, TensorOutputMap, + }; + + fn span() -> rumoca_core::Span { + rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("gpu_initialization_test.mo"), + 1, + 2, + ) + } + + fn direct_model() -> solve::SolveModel { + let span = span(); + let mut rows = vec![ + vec![ + LinearOp::LoadY { dst: 0, index: 0 }, + LinearOp::Const { dst: 1, value: 2.0 }, + LinearOp::Binary { + dst: 2, + op: BinaryOp::Sub, + lhs: 0, + rhs: 1, + }, + LinearOp::StoreOutput { src: 2 }, + ], + vec![ + LinearOp::LoadY { dst: 0, index: 1 }, + LinearOp::Const { dst: 1, value: 0.0 }, + LinearOp::Binary { + dst: 2, + op: BinaryOp::Sub, + lhs: 0, + rhs: 1, + }, + LinearOp::StoreOutput { src: 2 }, + ], + ]; + let domain = rumoca_core::StructuredIndexDomain { + binders: vec![rumoca_core::StructuredIndexBinder { + id: 0, + display_name: "i".to_string(), + lower: 1, + upper: 2, + step: 1, + }], + }; + let residual = ComputeNode::Map { + domain: domain.clone(), + output_map: TensorOutputMap::dense_contiguous(0, &domain).unwrap(), + base_ops: rows.remove(0), + load_strides: vec![rumoca_ir_solve::AffineStencilLoadStride { + op_position: 0, + terms: vec![AffineStencilIndexStrideTerm { + dimension: 0, + stride: 1, + }], + }], + const_strides: vec![rumoca_ir_solve::AffineStencilConstStride { + op_position: 1, + terms: vec![rumoca_ir_solve::AffineStencilConstStrideTerm { + dimension: 0, + stride: -2.0, + }], + }], + metadata: TensorNodeMetadata::default(), + span, + }; + let initialization = solve::InitializationSolveSystem { + residual: ComputeBlock { + nodes: vec![residual.clone()], + }, + direct_families: vec![solve::InitializationDirectFamily { + node_index: 0, + targets: TensorOutputMap::dense_contiguous(0, &domain).unwrap(), + residual_sign: 1, + span, + }], + required_target_ranges: vec![solve::InitializationTargetRange { + start: 0, + end: 2, + span: Some(span), + }], + ..Default::default() + }; + solve::SolveModel { + problem: solve::SolveProblem { + initialization, + ..Default::default() + }, + initial_y: vec![0.0, 0.0], + ..Default::default() + } + } + + fn singleton_domain(active_upper: i64) -> rumoca_core::StructuredIndexDomain { + rumoca_core::StructuredIndexDomain { + binders: vec![ + rumoca_core::StructuredIndexBinder { + id: 0, + display_name: "i".to_string(), + lower: 1, + upper: 1, + step: 1, + }, + rumoca_core::StructuredIndexBinder { + id: 1, + display_name: "j".to_string(), + lower: 1, + upper: active_upper, + step: 1, + }, + ], + } + } + + #[test] + fn direct_initial_assignment_is_one_pass_with_linear_temporary_storage() { + let result = settle_gpu_initial_conditions(&direct_model(), 0.0) + .expect("proven direct rows should settle"); + assert_eq!(result.y0, vec![2.0, 0.0]); + assert_eq!(result.metrics.residual_evaluations, 2); + assert_eq!(result.metrics.passes, 1); + assert!(result.metrics.temporary_values <= result.y0.len() * 3); + } + + #[test] + fn settlement_executes_mixed_singleton_domain_without_scalar_rows() { + let mut model = direct_model(); + let domain = singleton_domain(2); + let solve::ComputeNode::Map { + domain: residual_domain, + output_map, + load_strides, + const_strides, + .. + } = &mut model.problem.initialization.residual.nodes[0] + else { + unreachable!() + }; + *residual_domain = domain.clone(); + *output_map = TensorOutputMap::dense_contiguous(0, &domain).unwrap(); + load_strides[0].terms[0].dimension = 1; + const_strides[0].terms[0].dimension = 1; + model.problem.initialization.direct_families[0].targets = + TensorOutputMap::dense_contiguous(0, &domain).unwrap(); + + assert!(model.problem.initialization.row_targets.is_empty()); + assert_eq!( + model.problem.initialization.direct_families[0] + .targets + .strides, + vec![AffineStencilIndexStrideTerm { + dimension: 1, + stride: 1, + }] + ); + let result = settle_gpu_initial_conditions(&model, 0.0) + .expect("mixed-singleton direct rows should settle without scalar fallback"); + assert_eq!(result.y0, vec![2.0, 0.0]); + assert_eq!(result.metrics.residual_evaluations, 2); + assert_eq!(result.metrics.passes, 1); + } + + #[test] + fn settlement_executes_all_singleton_domain_with_empty_strides() { + let mut model = direct_model(); + let domain = singleton_domain(1); + let solve::ComputeNode::Map { + domain: residual_domain, + output_map, + base_ops, + load_strides, + const_strides, + .. + } = &mut model.problem.initialization.residual.nodes[0] + else { + unreachable!() + }; + *residual_domain = domain.clone(); + *output_map = TensorOutputMap::dense_contiguous(0, &domain).unwrap(); + let LinearOp::Const { value, .. } = &mut base_ops[1] else { + unreachable!() + }; + *value = 7.0; + load_strides.clear(); + const_strides.clear(); + model.problem.initialization.direct_families[0].targets = + TensorOutputMap::dense_contiguous(0, &domain).unwrap(); + model.problem.initialization.required_target_ranges[0].end = 1; + model.initial_y.truncate(1); + + assert!(model.problem.initialization.row_targets.is_empty()); + assert!( + model.problem.initialization.direct_families[0] + .targets + .strides + .is_empty() + ); + let result = settle_gpu_initial_conditions(&model, 0.0) + .expect("all-singleton direct row should settle with empty strides"); + assert_eq!(result.y0, vec![7.0]); + assert_eq!(result.metrics.residual_evaluations, 2); + assert_eq!(result.metrics.passes, 1); + } + + #[test] + fn direct_initial_assignment_uses_model_external_tables_for_apply_and_verify() { + let mut model = direct_model(); + let table_id = 515_151_u64; + model.external_tables = solve::ExternalTables::new(vec![rumoca_core::ExternalTableData { + id: table_id, + data: vec![vec![1.0, 10.0], vec![3.0, 30.0]], + columns: vec![2], + smoothness: 3, + extrapolation: 1, + }]); + let solve::ComputeNode::Map { + base_ops, + const_strides, + .. + } = &mut model.problem.initialization.residual.nodes[0] + else { + unreachable!() + }; + *base_ops = vec![ + LinearOp::LoadY { dst: 0, index: 0 }, + LinearOp::Const { + dst: 1, + value: table_id as f64, + }, + LinearOp::Const { dst: 2, value: 1.0 }, + LinearOp::Const { dst: 3, value: 2.0 }, + LinearOp::TableLookup { + dst: 4, + table_id: 1, + column: 2, + input: 3, + }, + LinearOp::Binary { + dst: 5, + op: BinaryOp::Sub, + lhs: 0, + rhs: 4, + }, + LinearOp::StoreOutput { src: 5 }, + ]; + const_strides.clear(); + + let settled = settle_gpu_initial_conditions(&model, 0.0) + .expect("table-backed direct initialization should settle and verify"); + assert_eq!(settled.y0, vec![10.0, 10.0]); + } + + #[test] + fn settlement_admission_rejects_random_direct_initial_operations_with_span() { + let mut model = direct_model(); + let solve::ComputeNode::Map { base_ops, .. } = + &mut model.problem.initialization.residual.nodes[0] + else { + unreachable!() + }; + base_ops.insert(0, LinearOp::ImpureRandomInit { dst: 6, seed: 1 }); + + let error = settle_gpu_initial_conditions(&model, 0.0) + .expect_err("random initialization cannot be replayed for residual verification"); + assert!(matches!( + error, + GpuInitializationError::Unsupported { + feature: "random or impure direct initial operations", + span: Some(actual), + .. + } if actual == span() + )); + } + + #[test] + fn settlement_rejects_partial_required_target_union_independently() { + let mut model = direct_model(); + model.problem.initialization.required_target_ranges[0].end = 1; + let error = settle_gpu_initial_conditions(&model, 0.0) + .expect_err("settlement must reject a partial hand-built artifact"); + assert!(error.to_string().contains("incomplete")); + } + + #[test] + fn settlement_rejects_partial_fixed_only_metadata_before_empty_residual_return() { + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("fixed_only_partial.mo"), + 20, + 30, + ); + let mut model = solve::SolveModel { + initial_y: vec![1.0, 2.0], + ..Default::default() + }; + model.problem.initialization.required_target_ranges = + vec![solve::InitializationTargetRange { + start: 0, + end: 2, + span: Some(span), + }]; + model.problem.initialization.fixed_target_ranges = vec![solve::InitializationTargetRange { + start: 0, + end: 1, + span: Some(span), + }]; + + let error = settle_gpu_initial_conditions(&model, 0.0) + .expect_err("partial fixed-only metadata must be validated before early return"); + assert!(error.to_string().contains("incomplete")); + assert!(matches!( + error, + GpuInitializationError::Malformed { span: Some(actual), .. } if actual == span + )); + } + + #[test] + fn settlement_rejects_duplicate_direct_node_ownership() { + let mut model = direct_model(); + let node = model.problem.initialization.residual.nodes[0].clone(); + model.problem.initialization.residual.nodes.push(node); + let mut duplicate = model.problem.initialization.direct_families[0].clone(); + duplicate.targets.start = 2; + model.problem.initialization.direct_families.push(duplicate); + model.problem.initialization.required_target_ranges[0].end = 4; + model.initial_y.resize(4, 0.0); + + let error = settle_gpu_initial_conditions(&model, 0.0) + .expect_err("every residual node must have one unique direct owner"); + assert!(error.to_string().contains("duplicate")); + } + + #[test] + fn settlement_rejects_unowned_residual_node() { + let mut model = direct_model(); + model + .problem + .initialization + .residual + .nodes + .push(model.problem.initialization.residual.nodes[0].clone()); + + let error = settle_gpu_initial_conditions(&model, 0.0) + .expect_err("every residual node must have a direct owner"); + assert!(error.to_string().contains("own every residual node")); + } + + #[test] + fn settlement_rejects_direct_fixed_overlap_at_fixed_span() { + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("direct_fixed_overlap.mo"), + 40, + 50, + ); + let mut model = direct_model(); + model.problem.initialization.fixed_target_ranges = vec![solve::InitializationTargetRange { + start: 1, + end: 2, + span: Some(span), + }]; + + let error = settle_gpu_initial_conditions(&model, 0.0) + .expect_err("simulation must reject direct/fixed ownership overlap"); + assert!(error.to_string().contains("overlap")); + assert!(matches!( + error, + GpuInitializationError::Malformed { span: Some(actual), .. } if actual == span + )); + } + + #[test] + fn settlement_accepts_adjacent_complete_fixed_only_ranges() { + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("adjacent_fixed.mo"), + 10, + 20, + ); + let mut model = solve::SolveModel { + initial_y: vec![-0.0, 2.0], + ..Default::default() + }; + model.problem.initialization.required_target_ranges = + vec![solve::InitializationTargetRange { + start: 0, + end: 2, + span: Some(span), + }]; + model.problem.initialization.fixed_target_ranges = vec![ + solve::InitializationTargetRange { + start: 0, + end: 1, + span: Some(span), + }, + solve::InitializationTargetRange { + start: 1, + end: 2, + span: Some(span), + }, + ]; + + let settled = settle_gpu_initial_conditions(&model, 0.0) + .expect("adjacent fixed ranges form exact complete coverage"); + assert_eq!(settled.y0[0].to_bits(), (-0.0f64).to_bits()); + assert_eq!(settled.y0[1], 2.0); + } + + #[test] + fn settlement_invalid_range_reports_span_after_json_and_bincode() { + let span = rumoca_core::Span::from_offsets( + rumoca_core::SourceId::from_source_name("invalid_range_roundtrip.mo"), + 70, + 80, + ); + let range = solve::InitializationTargetRange { + start: 0, + end: 3, + span: Some(span), + }; + let json = serde_json::to_string(&range).expect("serialize invalid range JSON"); + let from_json: solve::InitializationTargetRange = + serde_json::from_str(&json).expect("deserialize invalid range JSON"); + let bytes = bincode::serialize(&range).expect("serialize invalid range bincode"); + let from_bincode: solve::InitializationTargetRange = + bincode::deserialize(&bytes).expect("deserialize invalid range bincode"); + + for decoded in [from_json, from_bincode] { + let mut model = solve::SolveModel { + initial_y: vec![0.0, 0.0], + ..Default::default() + }; + model.problem.initialization.required_target_ranges = vec![decoded]; + let error = settle_gpu_initial_conditions(&model, 0.0) + .expect_err("out-of-bounds range must fail closed"); + assert!(matches!( + error, + GpuInitializationError::Malformed { span: Some(actual), .. } if actual == span + )); + } + } + + #[test] + fn event_system_is_rejected_before_returning_initial_vectors() { + let mut model = solve::SolveModel::default(); + model.problem.events.scheduled_time_events.push(0.0); + let error = settle_gpu_initial_conditions(&model, 0.0) + .expect_err("GPU preparation must reject event systems"); + assert!(matches!(error, GpuInitializationError::Unsupported { .. })); + } + + #[test] + fn nonfinite_initial_vector_is_rejected_without_partial_settlement() { + let mut model = direct_model(); + model.initial_y[0] = f64::NAN; + + let error = settle_gpu_initial_conditions(&model, 0.0) + .expect_err("non-finite GPU initialization input must fail closed"); + assert!(matches!(error, GpuInitializationError::NonConverged { .. })); + assert!(error.to_string().contains("initial y")); + } + + #[test] + fn no_initial_equations_preserve_finite_vectors_exactly() { + let model = solve::SolveModel { + initial_y: vec![-0.0, 3.25], + parameters: vec![7.5, -2.0], + ..Default::default() + }; + let result = settle_gpu_initial_conditions(&model, 0.0).expect("finite vectors pass"); + assert_eq!( + result.y0.iter().map(|v| v.to_bits()).collect::>(), + model + .initial_y + .iter() + .map(|v| v.to_bits()) + .collect::>() + ); + assert_eq!( + result.p0.iter().map(|v| v.to_bits()).collect::>(), + model + .parameters + .iter() + .map(|v| v.to_bits()) + .collect::>() + ); + } + + #[test] + fn no_initial_equations_still_reject_nonfinite_vectors() { + let model = solve::SolveModel { + initial_y: vec![f64::INFINITY], + ..Default::default() + }; + assert!(settle_gpu_initial_conditions(&model, 0.0).is_err()); + } +} diff --git a/crates/rumoca-sim/src/lib.rs b/crates/rumoca-sim/src/lib.rs index 769490616..f1d6420a8 100644 --- a/crates/rumoca-sim/src/lib.rs +++ b/crates/rumoca-sim/src/lib.rs @@ -13,7 +13,9 @@ use serde::{Deserialize, Serialize}; /// [`rumoca_eval_solve::nan_trace`]. pub use rumoca_eval_solve::nan_trace; use rumoca_ir_dae as dae; -pub use rumoca_phase_solve::{lower_solve_artifacts, lower_solve_problem}; +pub use rumoca_phase_solve::{ + lower_dae_to_solve_model_owned, lower_solve_artifacts, lower_solve_problem, +}; pub use rumoca_solver::{ BackendState, DiffsolMethod, LoopStats, RuntimeProgressSnapshot, RuntimeStopSchedule, RuntimeTraceContext, SimBackend, SimOptions, SimPacingMode, SimResult, SimSolverMode, @@ -28,6 +30,7 @@ pub use rumoca_solver::{ mod build_timing; pub mod bulk; +mod gpu_initialization; pub mod row_eval_trace; pub mod sim_trace_compare; #[cfg(any(feature = "solver-diffsol", feature = "solver-rk45"))] @@ -47,6 +50,10 @@ pub use diffsol::{ build_simulation_with_stage_timing_and_solve_model, check_initialization, check_prepared_initialization, run_prepared_simulation, simulate, simulate_dae, }; +pub use gpu_initialization::{ + GpuInitializationError, GpuInitializationMetrics, GpuInitializationResult, + settle_gpu_initial_conditions, +}; #[cfg(any(feature = "solver-diffsol", feature = "solver-rk45"))] pub use prepared_vectors::{PreparedVectorError, refresh_prepared_vectors}; #[cfg(any(feature = "solver-diffsol", feature = "solver-rk45"))] @@ -61,12 +68,14 @@ pub use solve_lowering::{ ObjectiveGradientProbe, ParameterJacobianProbe, SimulationDiagnosticError, SingularityDiagnosis, StateAndParameterJacobianProbe, SteadyStateSensitivityProbe, StructuralReport, TearingReport, UnmatchedEquationDiagnosis, UnmatchedUnknownDiagnosis, - diagnose_structural_singularity, eval_dae_at, jacobian_for_dae, lower_dae_for_gpu_preparation, - lower_dae_for_simulation, lower_for_differentiation_with_overrides, - lower_for_simulation_with_overrides, parameter_jacobian_for_dae, - state_and_parameter_jacobian_for_dae, steady_state_adjoint_objective_gradient_for_dae, - steady_state_objective_gradient_for_dae, steady_state_parameter_sensitivity_for_dae, - structural_report_for_dae, structurally_lowered_dae_for_simulation_artifact, + boundary_reduced_dae_for_simulation_artifact, diagnose_structural_singularity, eval_dae_at, + jacobian_for_dae, lower_dae_for_gpu_preparation, lower_dae_for_simulation, + lower_for_differentiation_with_overrides, lower_for_simulation_with_overrides, + parameter_jacobian_for_dae, state_and_parameter_jacobian_for_dae, + steady_state_adjoint_objective_gradient_for_dae, steady_state_objective_gradient_for_dae, + steady_state_parameter_sensitivity_for_dae, structural_report_for_dae, + structurally_lowered_dae_for_simulation_artifact, + structurally_prepared_dae_for_simulation_artifact, }; #[cfg(feature = "scenario-config")] diff --git a/crates/rumoca-sim/src/sim_trace_compare.rs b/crates/rumoca-sim/src/sim_trace_compare.rs index 40f90b10d..ebeed0d33 100644 --- a/crates/rumoca-sim/src/sim_trace_compare.rs +++ b/crates/rumoca-sim/src/sim_trace_compare.rs @@ -24,6 +24,8 @@ pub const SEVERE_CHANNEL_MAX_THRESHOLD: f64 = 0.80; pub struct SimTrace { #[serde(default)] pub model_name: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub n_states: Option, pub times: Vec, pub names: Vec, pub data: Vec>>, @@ -246,11 +248,7 @@ pub fn compare_model_traces( return Err(TraceCompareError::NoComparableSamples); } - channels.sort_by(|a, b| { - b.bounded_normalized_l1_error - .partial_cmp(&a.bounded_normalized_l1_error) - .unwrap_or(std::cmp::Ordering::Equal) - }); + channels.sort_by(compare_channel_deviation_desc); let compared_variables = channels.len(); let samples_compared = channels.iter().map(|m| m.samples).sum::(); @@ -267,7 +265,7 @@ pub fn compare_model_traces( .iter() .map(|m| m.bounded_normalized_l1_error) .collect::>(); - channel_scores.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)); + channel_scores.sort_by(f64::total_cmp); let bounded_normalized_l1_score = median_of_sorted(&channel_scores).unwrap_or(0.0); let channel_counts = count_channel_agreement_bands_default( channels @@ -621,14 +619,25 @@ fn compare_channel( omc_values: &[Option], use_step_hold: bool, ) -> Option { - if rumoca_times.len() < 2 - || omc_times.len() < 2 + if rumoca_times.is_empty() + || omc_times.is_empty() || rumoca_times.len() != rumoca_values.len() || omc_times.len() != omc_values.len() { return None; } + if rumoca_times.len() < 2 || omc_times.len() < 2 { + return compare_shared_initial_sample( + name, + rumoca_times, + rumoca_values, + omc_times, + omc_values, + use_step_hold, + ); + } + let deduped_grid = channel_comparison_grid(rumoca_times, omc_times)?; if deduped_grid.len() < 2 { return None; @@ -657,12 +666,12 @@ fn compare_channel( let (Some(r0), Some(o0), Some(r1), Some(o1)) = (r0, o0, r1, o1) else { continue; }; - let d0 = r0 - o0; - let d1 = r1 - o1; - let e0 = d0.abs(); - let e1 = d1.abs(); + let e0 = finite_abs_error(r0, o0); + let e1 = finite_abs_error(r1, o1); - integral_abs_error += 0.5 * (e0 + e1) * dt; + integral_abs_error = finite_non_negative_metric( + integral_abs_error + finite_non_negative_metric((0.5 * e0 + 0.5 * e1) * dt), + ); integral_duration += dt; max_abs_error = max_abs_error.max(e0).max(e1); ref_samples.push(o0); @@ -675,19 +684,20 @@ fn compare_channel( return None; } - let mean_abs_error = integral_abs_error / integral_duration; + let integral_abs_error = finite_non_negative_metric(integral_abs_error); + let mean_abs_error = finite_non_negative_metric(integral_abs_error / integral_duration); let reference_range = robust_reference_range(&ref_samples).unwrap_or(0.0); let normalization_scale = if use_step_hold { reference_range.max(1.0).max(NORMALIZATION_SCALE_EPS) } else { reference_range.max(NORMALIZATION_SCALE_EPS) }; - let normalized_l1_error = mean_abs_error / normalization_scale; - let bounded_normalized_l1_error = normalized_l1_error / (1.0 + normalized_l1_error); + let normalized_l1_error = finite_non_negative_metric(mean_abs_error / normalization_scale); + let bounded_normalized_l1_error = bounded_normalized_error(normalized_l1_error); let initial_abs_error = initial_abs_error(&samples); let initial_bounded_normalized_error = initial_abs_error.map(|error| { let normalized = error / normalization_scale; - normalized / (1.0 + normalized) + bounded_normalized_error(normalized) }); let shape = classify_channel_deviation_shape( &paired_samples, @@ -708,19 +718,101 @@ fn compare_channel( normalization_scale, normalized_l1_error, bounded_normalized_l1_error, - normalized_max_abs_error: max_abs_error / normalization_scale, + normalized_max_abs_error: finite_non_negative_metric(max_abs_error / normalization_scale), initial_abs_error, initial_bounded_normalized_error, }) } +fn compare_shared_initial_sample( + name: &str, + rumoca_times: &[f64], + rumoca_values: &[Option], + omc_times: &[f64], + omc_values: &[Option], + use_step_hold: bool, +) -> Option { + if (rumoca_times[0] - omc_times[0]).abs() > GRID_DEDUP_EPS { + return None; + } + let (Some(rumoca), Some(omc)) = (rumoca_values[0], omc_values[0]) else { + return None; + }; + if !rumoca.is_finite() || !omc.is_finite() { + return None; + } + + let initial_abs_error = finite_abs_error(rumoca, omc); + let normalization_scale = if use_step_hold { + omc.abs().max(1.0).max(NORMALIZATION_SCALE_EPS) + } else { + omc.abs().max(NORMALIZATION_SCALE_EPS) + }; + let normalized_error = finite_non_negative_metric(initial_abs_error / normalization_scale); + let bounded_error = bounded_normalized_error(normalized_error); + let shape = if bounded_error <= HIGH_AGREEMENT_CHANNEL_THRESHOLD { + TraceDeviationShape::WithinTolerance + } else { + TraceDeviationShape::WrongInitialValueOnly + }; + + Some(ChannelDeviationMetric { + name: name.to_string(), + shape, + samples: 1, + integral_duration: 0.0, + integral_abs_error: 0.0, + mean_abs_error: initial_abs_error, + normalization_scale, + normalized_l1_error: normalized_error, + bounded_normalized_l1_error: bounded_error, + normalized_max_abs_error: normalized_error, + initial_abs_error: Some(initial_abs_error), + initial_bounded_normalized_error: Some(bounded_error), + }) +} + +fn compare_channel_deviation_desc( + lhs: &ChannelDeviationMetric, + rhs: &ChannelDeviationMetric, +) -> std::cmp::Ordering { + rhs.bounded_normalized_l1_error + .total_cmp(&lhs.bounded_normalized_l1_error) + .then_with(|| lhs.name.cmp(&rhs.name)) +} + +fn bounded_normalized_error(normalized: f64) -> f64 { + if !normalized.is_finite() { + return if normalized.is_sign_negative() { + 0.0 + } else { + 1.0 + }; + } + let non_negative = normalized.max(0.0); + non_negative / (1.0 + non_negative) +} + +fn finite_abs_error(lhs: f64, rhs: f64) -> f64 { + finite_non_negative_metric((lhs - rhs).abs()) +} + +fn finite_non_negative_metric(value: f64) -> f64 { + if value.is_finite() { + value.max(0.0) + } else if value.is_sign_negative() { + 0.0 + } else { + f64::MAX + } +} + fn initial_abs_error(samples: &[(f64, Option, Option)]) -> Option { samples.iter().find_map(|(_, rumoca, omc)| { let (Some(rumoca), Some(omc)) = (*rumoca, *omc) else { return None; }; - let error = (rumoca - omc).abs(); - error.is_finite().then_some(error) + Some(finite_abs_error(rumoca, omc)) }) } @@ -1058,6 +1150,7 @@ mod tests { fn trace(model_name: &str, times: Vec, names: Vec<&str>, data: Vec>) -> SimTrace { SimTrace { model_name: Some(model_name.to_string()), + n_states: None, times, names: names.into_iter().map(ToOwned::to_owned).collect(), data: data @@ -1103,6 +1196,43 @@ mod tests { assert!((metric.bounded_normalized_l1_error - 1.0).abs() < 1.0e-10); } + #[test] + fn model_trace_compare_bounds_infinite_channel_error_without_sort_panic() { + let rumoca = trace( + "M", + vec![0.0, 1.0], + vec!["x", "y"], + vec![vec![f64::MAX, f64::MAX], vec![f64::MAX, f64::MAX]], + ); + let omc = trace( + "M", + vec![0.0, 1.0], + vec!["x", "y"], + vec![vec![-f64::MAX, -f64::MAX], vec![0.0, 0.0]], + ); + + let metric = compare_model_traces("M", &rumoca, &omc).expect("model compare"); + + assert_eq!(metric.compared_variables, 2); + assert_eq!(metric.max_channel_bounded_normalized_l1, 1.0); + assert_eq!(metric.worst_variables[0].bounded_normalized_l1_error, 1.0); + assert_eq!(metric.worst_variables[0].name, "x"); + assert_eq!(metric.worst_variables[1].name, "y"); + assert!(metric.worst_variables[0].integral_abs_error.is_finite()); + assert!(metric.worst_variables[0].mean_abs_error.is_finite()); + assert!(metric.worst_variables[0].normalized_l1_error.is_finite()); + assert!( + metric.worst_variables[0] + .normalized_max_abs_error + .is_finite() + ); + + let encoded = serde_json::to_value(&metric).expect("serialize metric"); + let decoded = + serde_json::from_value::(encoded).expect("deserialize metric"); + assert_eq!(decoded.max_channel_bounded_normalized_l1, 1.0); + } + #[test] fn channel_mean_abs_error_uses_time_weighted_integration() { let metric = compare_channel( @@ -1234,6 +1364,35 @@ mod tests { ); } + #[test] + fn compare_trace_uses_single_shared_initial_sample_after_duplicate_collapse() { + let rumoca = trace("M", vec![0.0], vec!["x"], vec![vec![0.0]]); + let mut omc = trace( + "M", + vec![0.0, 0.0], + vec!["x"], + vec![vec![0.0, 0.5_f64.asin()]], + ); + normalize_trace(&mut omc); + + let metric = compare_model_traces("M", &rumoca, &omc) + .expect("a legal shared initial sample should be comparable"); + + assert_eq!(metric.compared_variables, 1); + assert_eq!(metric.samples_compared, 1); + assert_eq!(metric.initial_condition.channels_compared, 1); + assert_eq!(metric.worst_variables[0].samples, 1); + assert_eq!(metric.worst_variables[0].integral_duration, 0.0); + assert!( + (metric.worst_variables[0] + .initial_abs_error + .expect("initial error") + - 0.5_f64.asin()) + .abs() + <= 1.0e-12 + ); + } + #[test] fn discrete_channel_uses_step_hold_interpolation() { let metric = compare_channel( @@ -1255,6 +1414,7 @@ mod tests { fn event_discontinuous_real_channel_uses_step_hold_interpolation() { let rumoca = SimTrace { model_name: Some("M".to_string()), + n_states: None, times: vec![0.0, 1.0], names: vec!["y".to_string()], data: vec![vec![Some(0.0), Some(1.0)]], @@ -1268,6 +1428,7 @@ mod tests { }; let omc = SimTrace { model_name: Some("M".to_string()), + n_states: None, times: vec![0.0, 0.5, 1.0], names: vec!["y".to_string()], data: vec![vec![Some(0.0), Some(0.0), Some(1.0)]], @@ -1286,6 +1447,7 @@ mod tests { fn discrete_only_model_traces_contribute_to_metrics() { let rumoca = SimTrace { model_name: Some("M".to_string()), + n_states: None, times: vec![0.0, 1.0], names: vec!["q".to_string()], data: vec![vec![Some(0.0), Some(1.0)]], @@ -1299,6 +1461,7 @@ mod tests { }; let omc = SimTrace { model_name: Some("M".to_string()), + n_states: None, times: vec![0.0, 0.5, 1.0], names: vec!["q".to_string()], data: vec![vec![Some(0.0), Some(0.0), Some(1.0)]], diff --git a/crates/rumoca-sim/src/solve_lowering.rs b/crates/rumoca-sim/src/solve_lowering.rs index 7eaf6063d..dc414b14e 100644 --- a/crates/rumoca-sim/src/solve_lowering.rs +++ b/crates/rumoca-sim/src/solve_lowering.rs @@ -29,8 +29,9 @@ pub use rumoca_phase_structural::{BlockReport, StructuralReport, TearingReport}; pub use diagnostics::SimulationDiagnosticError; pub use entry::{ - lower_dae_for_gpu_preparation, lower_dae_for_simulation, - structurally_lowered_dae_for_simulation_artifact, + boundary_reduced_dae_for_simulation_artifact, lower_dae_for_gpu_preparation, + lower_dae_for_simulation, structurally_lowered_dae_for_simulation_artifact, + structurally_prepared_dae_for_simulation_artifact, }; pub use probe::{ EvalAtProbe, JacobianProbe, ObjectiveGradientProbe, ParameterJacobianProbe, diff --git a/crates/rumoca-sim/src/solve_lowering/direct.rs b/crates/rumoca-sim/src/solve_lowering/direct.rs index c3d64fd28..e803b9904 100644 --- a/crates/rumoca-sim/src/solve_lowering/direct.rs +++ b/crates/rumoca-sim/src/solve_lowering/direct.rs @@ -62,19 +62,110 @@ pub(super) fn lower_direct_dae_for_simulation( pub(super) fn lower_direct_dae_for_gpu_preparation( dae_model: &dae::Dae, -) -> Result, rumoca_phase_solve::SolveModelLowerError> { +) -> Result { + validate_gpu_dae_admission(dae_model)?; let metadata_dae = attach_reference_metadata(dae_model)?; let lowered = metadata_dae.clone(); - match rumoca_phase_solve::lower_dae_to_solve_model_owned_for_gpu_preparation_with_metadata( - lowered, - &metadata_dae, - ) { - Ok(solve_model) => Ok(Some(solve_model)), - Err(err) => { - trace_direct_rejection(format!("direct GPU-preparation lowering failed: {err}")); - Ok(None) - } + let solve_model = + rumoca_phase_solve::lower_dae_to_solve_model_owned_for_gpu_preparation_with_metadata( + lowered, + &metadata_dae, + )?; + Ok(solve_model) +} + +/// GPU preparation has no event/discrete runtime payload. Reject these DAE +/// forms before the direct path attempts lowering; falling through to the +/// structural path would otherwise hide a semantic admission failure. +fn validate_gpu_dae_admission( + dae_model: &dae::Dae, +) -> Result<(), rumoca_phase_solve::SolveModelLowerError> { + let rejection = |reason: &str, span: Option| { + let Some(span) = span else { + return rumoca_phase_solve::SolveModelLowerError::Lower( + rumoca_phase_solve::LowerError::Unsupported { + reason: format!("GPU preparation rejects {reason} without source provenance"), + }, + ); + }; + rumoca_phase_solve::SolveModelLowerError::Lower( + rumoca_phase_solve::LowerError::UnsupportedAt { + reason: format!("GPU preparation rejects {reason}"), + contexts: vec!["GPU DAE admission".to_string()], + span, + }, + ) + }; + if let Some(equation) = dae_model.discrete.real_updates.first() { + return Err(rejection("discrete real updates", Some(equation.span))); + } + if let Some(equation) = dae_model.discrete.valued_updates.first() { + return Err(rejection("discrete-valued updates", Some(equation.span))); + } + if let Some(variable) = dae_model.variables.discrete_reals.values().next() { + return Err(rejection( + "discrete Real variables", + Some(variable.source_span), + )); + } + if let Some(variable) = dae_model.variables.discrete_valued.values().next() { + return Err(rejection( + "discrete-valued variables", + Some(variable.source_span), + )); + } + if let Some(equation) = dae_model.conditions.equations.first() { + return Err(rejection("condition equations", Some(equation.span))); + } + if let Some(relation) = dae_model.conditions.relations.first() { + return Err(rejection("relation memory", relation.span())); + } + if let Some(expression) = dae_model.events.synthetic_root_conditions.first() { + return Err(rejection("root conditions", expression.span())); + } + if let Some(event) = dae_model.events.scheduled_time_events.first() { + return Err(rejection("scheduled time events", event.source_span)); + } + if let Some(event) = dae_model.events.scheduled_root_conditions.first() { + let span = dae_model + .conditions + .relations + .get(event.root_index) + .and_then(rumoca_core::Expression::span); + return Err(rejection("scheduled root conditions", span)); + } + if let Some(action) = dae_model.events.event_actions.first() { + return Err(rejection("event actions", Some(action.span))); + } + if let Some(schedule) = dae_model.clocks.schedules.first() { + return Err(rejection("clock schedules", Some(schedule.source_span))); + } + if let Some(expression) = dae_model.clocks.constructor_exprs.first() { + return Err(rejection("clock constructors", expression.span())); + } + if let Some(expression) = dae_model.clocks.triggered_conditions.first() { + return Err(rejection("triggered clock conditions", expression.span())); + } + if let Some(variable) = dae_model + .variables + .parameters + .values() + .find(|variable| rumoca_core::pre_slot_base(variable.name.as_str()).is_some()) + { + return Err(rejection("pre-state memory", Some(variable.source_span))); + } + if let Some(equation) = dae_model.initialization.equations.iter().find(|equation| { + equation.lhs.as_ref().is_some_and(|lhs| { + let name = lhs.var_name(); + dae_model.variables.parameters.contains_key(name) + || dae_model.variables.inputs.contains_key(name) + || dae_model.variables.discrete_reals.contains_key(name) + || dae_model.variables.discrete_valued.contains_key(name) + }) + }) { + return Err(rejection("initial P-slot target", Some(equation.span))); } + Ok(()) } fn attach_reference_metadata( @@ -303,3 +394,104 @@ fn trace_direct_rejection(reason: impl AsRef) { reason.as_ref() ); } + +#[cfg(test)] +mod tests { + use super::validate_gpu_dae_admission; + use rumoca_core::{Expression, Literal, SourceId, Span, VarName}; + use rumoca_ir_dae as dae; + + fn span() -> Span { + Span::from_offsets(SourceId::from_source_name("gpu_admission.mo"), 1, 2) + } + + fn zero() -> Expression { + Expression::Literal { + value: Literal::Real(0.0), + span: span(), + } + } + + fn equation() -> dae::Equation { + dae::Equation::residual(zero(), span(), "GPU admission fixture") + } + + #[test] + fn gpu_admission_rejects_discrete_and_condition_partitions_with_source_span() { + let mut discrete = dae::Dae::default(); + discrete.discrete.real_updates.push(equation()); + let error = validate_gpu_dae_admission(&discrete).expect_err("discrete must reject"); + assert!(error.to_string().contains("discrete real updates")); + + let mut conditions = dae::Dae::default(); + conditions.conditions.equations.push(equation()); + let error = validate_gpu_dae_admission(&conditions).expect_err("conditions must reject"); + assert!(error.to_string().contains("condition equations")); + + let mut bare_discrete = dae::Dae::default(); + let name = VarName::new("z"); + bare_discrete + .variables + .discrete_reals + .insert(name.clone(), dae::Variable::empty_with_span(span())); + let error = validate_gpu_dae_admission(&bare_discrete) + .expect_err("a discrete variable without updates must reject"); + assert!(error.to_string().contains("discrete Real variables")); + } + + #[test] + fn gpu_admission_rejects_event_and_clock_metadata_with_source_span() { + let mut events = dae::Dae::default(); + events.events.synthetic_root_conditions.push(zero()); + let error = validate_gpu_dae_admission(&events).expect_err("events must reject"); + assert!(error.to_string().contains("root conditions")); + + let mut scheduled = dae::Dae::default(); + let event_span = Span::from_offsets(SourceId::from_source_name("event-only.mo"), 40, 55); + scheduled + .events + .scheduled_time_events + .push(dae::DaeScheduledTimeEvent { + time: 1.0, + source_span: Some(event_span), + }); + let error = validate_gpu_dae_admission(&scheduled) + .expect_err("scheduled time events must reject before fast lowering"); + assert!(error.to_string().contains("scheduled time events")); + assert_eq!(error.source_span(), Some(event_span)); + + let mut clocks = dae::Dae::default(); + clocks.clocks.constructor_exprs.push(zero()); + let error = validate_gpu_dae_admission(&clocks).expect_err("clocks must reject"); + assert!(error.to_string().contains("clock constructors")); + } + + #[test] + fn gpu_admission_rejects_typed_pre_slot_memory() { + let mut dae_model = dae::Dae::default(); + let name = rumoca_core::pre_slot_name("x"); + let mut variable = dae::Variable::empty_with_span(span()); + variable.name = name.clone(); + dae_model.variables.parameters.insert(name, variable); + let error = validate_gpu_dae_admission(&dae_model).expect_err("pre slot must reject"); + assert!(error.to_string().contains("pre-state memory")); + } + + #[test] + fn gpu_admission_rejects_initial_parameter_target_before_fast_lowering() { + let mut dae_model = dae::Dae::default(); + let name = VarName::new("p"); + dae_model + .variables + .parameters + .insert(name.clone(), dae::Variable::empty_with_span(span())); + dae_model + .initialization + .equations + .push(dae::Equation::explicit(name, zero(), span(), "P target")); + + let error = validate_gpu_dae_admission(&dae_model) + .expect_err("initial parameter target must reject"); + assert!(error.to_string().contains("initial P-slot target")); + } +} diff --git a/crates/rumoca-sim/src/solve_lowering/entry.rs b/crates/rumoca-sim/src/solve_lowering/entry.rs index 9133f1512..0d0fb31ce 100644 --- a/crates/rumoca-sim/src/solve_lowering/entry.rs +++ b/crates/rumoca-sim/src/solve_lowering/entry.rs @@ -22,16 +22,9 @@ pub fn lower_dae_for_simulation( pub fn lower_dae_for_gpu_preparation( dae_model: &dae::Dae, - opts: &SimOptions, + _opts: &SimOptions, ) -> Result { - if let Some(solve_model) = lower_direct_dae_for_gpu_preparation(dae_model)? { - return Ok(solve_model); - } - let structurally_lowered = structurally_lower_dae_for_simulation(dae_model, opts)?; - rumoca_phase_solve::lower_dae_to_solve_model_owned_for_gpu_preparation_with_metadata( - structurally_lowered.dae, - &structurally_lowered.metadata_dae, - ) + lower_direct_dae_for_gpu_preparation(dae_model) } /// Lower for simulation while applying tunable scalar-parameter overrides during @@ -57,6 +50,9 @@ pub(crate) fn lower_dae_for_differentiation_with_param_overrides( param_overrides: &std::collections::HashMap, ) -> Result { let structurally_lowered = structurally_lower_dae_for_simulation(dae_model, opts)?; + // Differentiation always requires the full Jacobian/sensitivity artifact + // set. `solver_mode` selects the ordinary simulation backend and must not + // downgrade optimization lowering to the RK value-only profile. rumoca_phase_solve::lower_dae_to_solve_model_owned_with_visible_expressions_and_metadata_and_overrides( structurally_lowered.dae, structurally_lowered.visible_expressions, @@ -212,3 +208,17 @@ pub fn structurally_lowered_dae_for_simulation_artifact( ) -> Result { structurally_lower_dae_for_simulation(dae_model, opts).map(|lowered| lowered.dae) } + +pub fn structurally_prepared_dae_for_simulation_artifact( + dae_model: &dae::Dae, + opts: &SimOptions, +) -> Result { + super::structural_lowering::prepare_structural_dae_for_simulation_artifact(dae_model, opts) +} + +pub fn boundary_reduced_dae_for_simulation_artifact( + dae_model: &dae::Dae, + opts: &SimOptions, +) -> Result { + super::structural_lowering::boundary_reduced_dae_for_simulation_artifact(dae_model, opts) +} diff --git a/crates/rumoca-sim/src/solve_lowering/structural_lowering.rs b/crates/rumoca-sim/src/solve_lowering/structural_lowering.rs index 98404943f..c3add1db7 100644 --- a/crates/rumoca-sim/src/solve_lowering/structural_lowering.rs +++ b/crates/rumoca-sim/src/solve_lowering/structural_lowering.rs @@ -21,6 +21,11 @@ pub(super) fn prepare_dae_for_structural_analysis( lowered: &mut dae::Dae, opts: &SimOptions, ) -> Result<(), rumoca_phase_solve::SolveModelLowerError> { + log_solve_lowering_start("prepare.scalarize_vector_member_slices"); + let timer = stage_timer_start(); + rumoca_phase_dae::scalarize_phantom_vector_equations(lowered) + .map_err(vector_scalarization_lower_error)?; + log_solve_lowering_done("prepare.scalarize_vector_member_slices", timer); if opts.scalarize { log_solve_lowering_start("prepare.scalarize_equations"); let timer = stage_timer_start(); @@ -91,6 +96,18 @@ pub(super) fn prepare_dae_for_structural_analysis( "prepare.substitute_standalone_state_derivatives_in_non_ode_rows", timer, ); + log_solve_lowering_start("prepare.demote_states_after_standalone_derivative_substitution"); + let timer = stage_timer_start(); + rumoca_phase_structural::dae_prepare::demote_states_without_retained_derivative_rows(lowered) + .map_err(|source| rumoca_phase_solve::SolveModelLowerError::Structural { source })?; + log_solve_lowering_done( + "prepare.demote_states_after_standalone_derivative_substitution", + timer, + ); + log_solve_lowering_start("prepare.remove_nonnumeric_continuous_equations"); + let timer = stage_timer_start(); + remove_nonnumeric_continuous_equations(lowered); + log_solve_lowering_done("prepare.remove_nonnumeric_continuous_equations", timer); if tracing::enabled!(target: "rumoca_phase_structural", tracing::Level::DEBUG) { for (index, eq) in lowered.continuous.equations.iter().enumerate() { let summary = format!("{}{}", equation_lhs_prefix(eq), debug_render_expr(&eq.rhs)); @@ -105,6 +122,311 @@ pub(super) fn prepare_dae_for_structural_analysis( Ok(()) } +/// Boundary elimination can expose state constraints that were hidden behind +/// aliases or simple connector equalities. Re-run the state/dummy preparation +/// steps that depend on those exposed equations before the final BLT match. +pub(super) fn prepare_dae_after_boundary_elimination( + lowered: &mut dae::Dae, + boundary_substitutions: &[rumoca_phase_structural::eliminate::Substitution], +) -> Result { + log_solve_lowering_start("structural.post_boundary.demote_exact_alias_component_states"); + let timer = stage_timer_start(); + let exact_alias_demoted = + rumoca_phase_structural::dae_prepare::demote_exact_alias_component_states(lowered) + .map_err(|source| rumoca_phase_solve::SolveModelLowerError::Structural { source })?; + log_solve_lowering_done( + "structural.post_boundary.demote_exact_alias_component_states", + timer, + ); + + log_solve_lowering_start("structural.post_boundary.demote_direct_assigned_states"); + let timer = stage_timer_start(); + let direct_demoted = + rumoca_phase_structural::dae_prepare::demote_direct_assigned_states_with_boundary_substitutions( + lowered, + boundary_substitutions, + ) + .map_err(|source| rumoca_phase_solve::SolveModelLowerError::Structural { source })?; + log_solve_lowering_done( + "structural.post_boundary.demote_direct_assigned_states", + timer, + ); + + log_solve_lowering_start("structural.post_boundary.reduce_constrained_dummy_derivatives"); + let timer = stage_timer_start(); + let dummy_reduced = + rumoca_phase_structural::dae_prepare::reduce_constrained_dummy_derivatives(lowered) + .map_err(|source| rumoca_phase_solve::SolveModelLowerError::Structural { source })?; + log_solve_lowering_done( + "structural.post_boundary.reduce_constrained_dummy_derivatives", + timer, + ); + + log_solve_lowering_start( + "structural.post_boundary.demote_states_without_retained_derivative_rows", + ); + let timer = stage_timer_start(); + let (states_demoted, unassignable_demoted) = + rumoca_phase_structural::dae_prepare::demote_states_without_retained_derivative_rows( + lowered, + ) + .map_err(|source| rumoca_phase_solve::SolveModelLowerError::Structural { source })?; + log_solve_lowering_done( + "structural.post_boundary.demote_states_without_retained_derivative_rows", + timer, + ); + + Ok(exact_alias_demoted > 0 + || direct_demoted > 0 + || dummy_reduced > 0 + || states_demoted > 0 + || unassignable_demoted > 0) +} + +fn remove_nonnumeric_continuous_equations(dae: &mut dae::Dae) { + let removed = dae + .continuous + .equations + .iter() + .enumerate() + .filter_map(|(idx, equation)| { + expression_contains_nonnumeric_metadata(&equation.rhs).then_some(idx) + }) + .collect::>(); + if removed.is_empty() { + return; + } + let mut next_removed = 0usize; + let mut next_idx = 0usize; + dae.continuous.equations.retain(|_| { + let remove = removed + .get(next_removed) + .is_some_and(|idx| *idx == next_idx); + next_idx += 1; + if remove { + next_removed += 1; + } + !remove + }); + shift_structured_families_after_equation_removal( + &mut dae.continuous.structured_equations, + &removed, + ); +} + +fn expression_contains_nonnumeric_metadata(expr: &rumoca_core::Expression) -> bool { + struct Visitor { + found: bool, + } + + impl rumoca_core::ExpressionVisitor for Visitor { + fn visit_function_call( + &mut self, + name: &rumoca_core::Reference, + args: &[rumoca_core::Expression], + _is_constructor: bool, + ) { + // Table intrinsics are numeric even though their constructor + // metadata contains String fields. Skip only this call's argument + // subtree; nonnumeric siblings and wrappers must still be seen. + if external_table_numeric_intrinsic_reference(name) { + return; + } + for arg in args { + self.visit_expression(arg); + } + } + + fn visit_literal(&mut self, value: &rumoca_core::Literal) { + if matches!(value, rumoca_core::Literal::String(_)) { + self.found = true; + } + } + } + + let mut visitor = Visitor { found: false }; + rumoca_core::ExpressionVisitor::visit_expression(&mut visitor, expr); + visitor.found +} + +fn external_table_numeric_intrinsic_reference(name: &rumoca_core::Reference) -> bool { + if external_table_numeric_intrinsic(name.last_segment()) { + return true; + } + + // Scalarization represents a scalar function result as an exact output + // projection (`function.y`). Match that structured shape explicitly; a + // generic suffix match would accidentally preserve unrelated String-valued + // calls whose names merely contain an intrinsic name. + let segments = name.segments(); + segments.last() == Some(&"y") + && segments + .get(segments.len().saturating_sub(2)) + .is_some_and(|function| external_table_numeric_intrinsic(function)) +} + +fn external_table_numeric_intrinsic(short_name: &str) -> bool { + matches!( + short_name, + "getTimeTableTmin" + | "getTimeTableTmax" + | "getNextTimeEvent" + | "getTimeTableValueNoDer" + | "getTimeTableValueNoDer2" + | "getTimeTableValue" + | "getTable1DAbscissaUmin" + | "getTable1DAbscissaUmax" + | "getTable1DValueNoDer" + | "getTable1DValueNoDer2" + | "getTable1DValue" + ) +} + +#[cfg(test)] +fn projected_external_table_test_call() -> rumoca_core::Expression { + let span = rumoca_core::Span::DUMMY; + rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new( + "Modelica.Blocks.Tables.Internal.getTimeTableValueNoDer.y", + ), + args: vec![rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("NoName".to_string()), + span, + }], + is_constructor: false, + span, + } +} + +#[cfg(test)] +#[test] +fn projected_external_table_numeric_intrinsic_survives_metadata_pruning() { + let span = rumoca_core::Span::DUMMY; + let mut model = dae::Dae::default(); + model.continuous.equations.push(dae::Equation { + lhs: None, + rhs: projected_external_table_test_call(), + span, + origin: "scalarized external table lookup".to_string(), + scalar_count: 1, + }); + + remove_nonnumeric_continuous_equations(&mut model); + + assert_eq!(model.continuous.equations.len(), 1); +} + +#[cfg(test)] +#[test] +fn projected_external_table_call_does_not_hide_string_sibling_metadata() { + let span = rumoca_core::Span::DUMMY; + let expr = rumoca_core::Expression::Binary { + op: rumoca_core::OpBinary::Add, + lhs: Box::new(projected_external_table_test_call()), + rhs: Box::new(rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("metadata".to_string()), + span, + }), + span, + }; + + assert!(expression_contains_nonnumeric_metadata(&expr)); +} + +#[cfg(test)] +#[test] +fn numeric_named_argument_wrapper_is_not_nonnumeric_metadata() { + let span = rumoca_core::Span::DUMMY; + let named = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new(format!("{}A", rumoca_core::NAMED_FUNCTION_ARG_PREFIX)), + args: vec![rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(1.0), + span, + }], + is_constructor: false, + span, + }; + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.numericFunction"), + args: vec![named], + is_constructor: false, + span, + }; + + assert!(!expression_contains_nonnumeric_metadata(&expr)); +} + +#[cfg(test)] +#[test] +fn named_argument_wrapper_does_not_hide_nonnumeric_actual() { + let span = rumoca_core::Span::DUMMY; + let named = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new(format!( + "{}fileName", + rumoca_core::NAMED_FUNCTION_ARG_PREFIX + )), + args: vec![rumoca_core::Expression::Literal { + value: rumoca_core::Literal::String("table.csv".to_string()), + span, + }], + is_constructor: false, + span, + }; + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.numericFunction"), + args: vec![named], + is_constructor: false, + span, + }; + + assert!(expression_contains_nonnumeric_metadata(&expr)); +} + +#[cfg(test)] +#[test] +fn projected_external_table_call_remains_numeric_inside_named_argument() { + let span = rumoca_core::Span::DUMMY; + let named = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new(format!( + "{}tableValue", + rumoca_core::NAMED_FUNCTION_ARG_PREFIX + )), + args: vec![projected_external_table_test_call()], + is_constructor: false, + span, + }; + let expr = rumoca_core::Expression::FunctionCall { + name: rumoca_core::Reference::new("Pkg.numericFunction"), + args: vec![named], + is_constructor: false, + span, + }; + + assert!(!expression_contains_nonnumeric_metadata(&expr)); +} + +fn shift_structured_families_after_equation_removal( + families: &mut Vec, + removed_sorted: &[usize], +) { + families.retain_mut(|family| { + let total: usize = family.equation_counts.iter().sum(); + let block_end = family.first_equation_index + total; + if removed_sorted + .iter() + .any(|&idx| idx >= family.first_equation_index && idx < block_end) + { + return false; + } + let shift = removed_sorted + .iter() + .filter(|&&idx| idx < family.first_equation_index) + .count(); + family.first_equation_index -= shift; + true + }); +} + pub(super) struct StructurallyLoweredDae { pub(super) dae: dae::Dae, pub(super) metadata_dae: dae::Dae, @@ -124,24 +446,61 @@ pub(super) fn structurally_lower_dae_for_simulation( let PreparedStructuralDaes { source_dae, mut lowered, - mut metadata_dae, + metadata_dae, } = prepare_structural_daes(dae_model, opts)?; - + if dae_model.variables.states.is_empty() { + let mut residual_shape_dae = source_dae.clone(); + remove_nonnumeric_continuous_equations(&mut residual_shape_dae); + validate_residual_shapes_for_simulation(&residual_shape_dae)?; + } log_solve_lowering_start("structural.eliminate_trivial"); let timer = stage_timer_start(); - let elimination = rumoca_phase_structural::eliminate::eliminate_trivial(&mut lowered) + let mut elimination = rumoca_phase_structural::eliminate::eliminate_trivial(&mut lowered) .map_err(|source| rumoca_phase_solve::SolveModelLowerError::Structural { source })?; log_solve_lowering_done("structural.eliminate_trivial", timer); + let first_blt_error = elimination.blt_error.take(); + if first_blt_error.is_some() + && prepare_dae_after_boundary_elimination(&mut lowered, &elimination.substitutions)? + { + let mut substitutions = elimination.substitutions; + log_solve_lowering_start("structural.eliminate_trivial_after_boundary_state_prep"); + let timer = stage_timer_start(); + elimination = rumoca_phase_structural::eliminate::eliminate_trivial(&mut lowered) + .map_err(|source| rumoca_phase_solve::SolveModelLowerError::Structural { source })?; + log_solve_lowering_done( + "structural.eliminate_trivial_after_boundary_state_prep", + timer, + ); + substitutions.extend(elimination.substitutions); + if elimination.blt_error.is_some() { + reprepare_dae_after_boundary_elimination(&mut lowered, opts)?; + log_solve_lowering_start("structural.eliminate_trivial_after_boundary_reprepare"); + let timer = stage_timer_start(); + elimination = rumoca_phase_structural::eliminate::eliminate_trivial(&mut lowered) + .map_err( + |source| rumoca_phase_solve::SolveModelLowerError::Structural { source }, + )?; + log_solve_lowering_done( + "structural.eliminate_trivial_after_boundary_reprepare", + timer, + ); + substitutions.extend(elimination.substitutions); + } + elimination.substitutions = substitutions; + } else if let Some(source) = first_blt_error { + if dae_model.variables.states.is_empty() { + validate_residual_shapes_for_simulation(dae_model)?; + } + return Err(rumoca_phase_solve::SolveModelLowerError::Structural { source }); + } if let Some(source) = elimination.blt_error { if dae_model.variables.states.is_empty() { validate_residual_shapes_for_simulation(dae_model)?; } return Err(rumoca_phase_solve::SolveModelLowerError::Structural { source }); } - apply_simulation_elimination(&mut lowered, &elimination.substitutions)?; trace_simulation_elimination(&lowered, &elimination.substitutions); - mark_state_selection_metadata(&mut metadata_dae, &elimination.substitutions)?; let visible_expressions = visible_expressions_after_elimination(&source_dae, &elimination.substitutions, opts)?; @@ -183,6 +542,47 @@ fn prepare_structural_daes( }) } +fn reprepare_dae_after_boundary_elimination( + lowered: &mut dae::Dae, + opts: &SimOptions, +) -> Result<(), rumoca_phase_solve::SolveModelLowerError> { + log_solve_lowering_start("structural.reprepare_after_boundary_elimination"); + let timer = stage_timer_start(); + prepare_dae_for_structural_analysis(lowered, opts)?; + remove_duplicate_continuous_equations(lowered); + log_solve_lowering_done("structural.reprepare_after_boundary_elimination", timer); + Ok(()) +} + +pub(super) fn prepare_structural_dae_for_simulation_artifact( + dae_model: &dae::Dae, + opts: &SimOptions, +) -> Result { + prepare_structural_daes(dae_model, opts).map(|prepared| prepared.lowered) +} + +pub(super) fn boundary_reduced_dae_for_simulation_artifact( + dae_model: &dae::Dae, + opts: &SimOptions, +) -> Result { + let mut lowered = prepare_structural_daes(dae_model, opts)?.lowered; + let mut elimination = rumoca_phase_structural::eliminate::eliminate_trivial(&mut lowered) + .map_err(|source| rumoca_phase_solve::SolveModelLowerError::Structural { source })?; + if elimination.blt_error.is_some() + && prepare_dae_after_boundary_elimination(&mut lowered, &elimination.substitutions)? + { + elimination = rumoca_phase_structural::eliminate::eliminate_trivial(&mut lowered) + .map_err(|source| rumoca_phase_solve::SolveModelLowerError::Structural { source })?; + if elimination.blt_error.is_some() { + reprepare_dae_after_boundary_elimination(&mut lowered, opts)?; + rumoca_phase_structural::eliminate::eliminate_trivial(&mut lowered).map_err( + |source| rumoca_phase_solve::SolveModelLowerError::Structural { source }, + )?; + } + } + Ok(lowered) +} + fn apply_simulation_elimination( lowered: &mut dae::Dae, substitutions: &[rumoca_phase_structural::eliminate::Substitution], @@ -231,39 +631,6 @@ fn trace_simulation_elimination( } } -fn mark_state_selection_metadata( - metadata_dae: &mut dae::Dae, - substitutions: &[rumoca_phase_structural::eliminate::Substitution], -) -> Result<(), rumoca_phase_solve::SolveModelLowerError> { - log_solve_lowering_start("structural.clone_state_selection_dae"); - let timer = stage_timer_start(); - let mut state_selection_dae = metadata_dae.clone(); - log_solve_lowering_done("structural.clone_state_selection_dae", timer); - log_solve_lowering_start("structural.apply_state_selection_substitutions"); - let timer = stage_timer_start(); - rumoca_phase_structural::eliminate::apply_elimination_substitutions_to_dae( - &mut state_selection_dae, - substitutions, - ) - .map_err(|source| rumoca_phase_solve::SolveModelLowerError::Structural { source })?; - log_solve_lowering_done("structural.apply_state_selection_substitutions", timer); - log_solve_lowering_start("structural.demote_state_selection_dae"); - let timer = stage_timer_start(); - rumoca_phase_structural::dae_prepare::demote_states_without_retained_derivative_rows( - &mut state_selection_dae, - ) - .map_err(|source| rumoca_phase_solve::SolveModelLowerError::Structural { source })?; - log_solve_lowering_done("structural.demote_state_selection_dae", timer); - log_solve_lowering_start("structural.mark_constrained_dummy_states_in_metadata"); - let timer = stage_timer_start(); - mark_constrained_dummy_states_in_metadata(&state_selection_dae, metadata_dae); - log_solve_lowering_done( - "structural.mark_constrained_dummy_states_in_metadata", - timer, - ); - Ok(()) -} - fn visible_expressions_after_elimination( source_dae: &dae::Dae, substitutions: &[rumoca_phase_structural::eliminate::Substitution], @@ -322,6 +689,16 @@ pub(super) fn metadata_attachment_lower_error( )) } +fn vector_scalarization_lower_error( + err: rumoca_phase_dae::ToDaeError, +) -> rumoca_phase_solve::SolveModelLowerError { + let reason = format!("DAE vector scalarization failed: {err}"); + rumoca_phase_solve::SolveModelLowerError::Lower(lower_contract_error_from_optional_span( + reason, + err.source_span(), + )) +} + fn lower_contract_error_from_optional_span( reason: String, span: Option, @@ -343,17 +720,3 @@ fn validate_residual_shapes_for_simulation( rumoca_phase_solve::lower::lower_residual(dae_model, &layout)?; Ok(()) } - -fn mark_constrained_dummy_states_in_metadata( - structural_dae: &dae::Dae, - metadata_dae: &mut dae::Dae, -) { - for state_name in - rumoca_phase_structural::dae_prepare::constrained_dummy_state_names(structural_dae) - { - let name = rumoca_core::VarName::new(state_name); - if let Some(var) = metadata_dae.variables.states.shift_remove(&name) { - metadata_dae.variables.algebraics.insert(name, var); - } - } -} diff --git a/crates/rumoca-sim/src/solve_lowering/structure_report.rs b/crates/rumoca-sim/src/solve_lowering/structure_report.rs index 6d8240e17..c8d2eca34 100644 --- a/crates/rumoca-sim/src/solve_lowering/structure_report.rs +++ b/crates/rumoca-sim/src/solve_lowering/structure_report.rs @@ -7,7 +7,32 @@ use rumoca_solver::SimOptions; use super::diagnostics::SimulationDiagnosticError; use super::expr_util::equation_lhs_prefix; -use super::structural_lowering::prepare_dae_for_structural_analysis; +use super::structural_lowering::{ + prepare_dae_after_boundary_elimination, prepare_dae_for_structural_analysis, +}; + +fn prepare_dae_for_structure_report( + prepared: &mut dae::Dae, + opts: &SimOptions, +) -> Result<(), SimulationDiagnosticError> { + prepare_dae_for_structural_analysis(prepared, opts) + .map_err(SimulationDiagnosticError::SolveLowering)?; + loop { + let boundary = + rumoca_phase_structural::eliminate::resolve_boundary_equations_to_fixpoint(prepared) + .map_err(|error| { + SimulationDiagnosticError::Solver(format!( + "structural boundary preparation failed: {error}" + )) + })?; + let changed = prepare_dae_after_boundary_elimination(prepared, &boundary.substitutions) + .map_err(SimulationDiagnosticError::SolveLowering)?; + if boundary.n_eliminated == 0 && !changed { + break; + } + } + Ok(()) +} /// Structurally analyze the model and return a named report of the matching, /// BLT blocks, coupled SCCs, and tearing. The analysis runs on the flattened @@ -24,8 +49,7 @@ pub fn structural_report_for_dae( // Report on the same prepared system the simulator matches, so the analysis // reflects reality (e.g. `der(x)` references in non-ODE rows are resolved). let mut prepared = dae_model.clone(); - prepare_dae_for_structural_analysis(&mut prepared, opts) - .map_err(SimulationDiagnosticError::SolveLowering)?; + prepare_dae_for_structure_report(&mut prepared, opts)?; rumoca_phase_structural::build_structural_report(&prepared).map_err(|error| { SimulationDiagnosticError::Solver(format!("structural analysis failed: {error}")) }) @@ -68,8 +92,7 @@ pub fn diagnose_structural_singularity( opts: &SimOptions, ) -> Result, SimulationDiagnosticError> { let mut prepared = dae_model.clone(); - prepare_dae_for_structural_analysis(&mut prepared, opts) - .map_err(SimulationDiagnosticError::SolveLowering)?; + prepare_dae_for_structure_report(&mut prepared, opts)?; let error = match rumoca_phase_structural::build_structural_report(&prepared) { Ok(_) => return Ok(None), diff --git a/crates/rumoca-sim/src/solve_lowering/tests.rs b/crates/rumoca-sim/src/solve_lowering/tests.rs index a0d4abe40..5851533aa 100644 --- a/crates/rumoca-sim/src/solve_lowering/tests.rs +++ b/crates/rumoca-sim/src/solve_lowering/tests.rs @@ -10,7 +10,8 @@ use super::diagnostics::SimulationDiagnosticError; use super::entry::{lower_dae_for_simulation, lower_dae_for_simulation_with_stage_timing}; use super::probe::{eval_dae_at, jacobian_for_dae}; use super::structural_lowering::{ - metadata_attachment_lower_error, structurally_lower_dae_for_simulation, + metadata_attachment_lower_error, prepare_dae_after_boundary_elimination, + structurally_lower_dae_for_simulation, }; fn sim_source_span(source: u64, start: usize, end: usize) -> Span { @@ -72,6 +73,56 @@ fn simulation_structural_lowering_reports_blt_singularity() { assert!(err.to_string().contains("structurally singular system")); } +#[test] +fn simulation_structural_lowering_drops_nonnumeric_continuous_metadata_rows() { + let mut dae = dae::Dae::new(); + dae.variables.algebraics.insert( + VarName::new("p"), + dae::Variable::new(VarName::new("p"), fixture_span()), + ); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: Expression::Binary { + op: OpBinary::Mul, + lhs: Box::new(Expression::Literal { + value: rumoca_core::Literal::String("Air".to_string()), + span: fixture_span(), + }), + rhs: Box::new(var("p")), + span: fixture_span(), + }, + span: fixture_span(), + origin: "medium metadata".to_string(), + scalar_count: 1, + }); + + let lowered = structurally_lower_dae_for_simulation(&dae, &SimOptions::default()) + .expect("nonnumeric metadata rows should not enter structural matching"); + + assert!(lowered.dae.continuous.equations.is_empty()); +} + +#[test] +fn simulation_structural_lowering_preserves_numeric_named_argument_rows() { + let mut dae = dae::Dae::new(); + dae.variables.algebraics.insert( + VarName::new("p"), + dae::Variable::new(VarName::new("p"), fixture_span()), + ); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: mul(named_arg("MM", vec![var("p")]), int(2)), + span: fixture_span(), + origin: "medium record metadata".to_string(), + scalar_count: 1, + }); + + let lowered = structurally_lower_dae_for_simulation(&dae, &SimOptions::default()) + .expect("numeric named arguments should remain available to structural matching"); + + assert_eq!(lowered.dae.continuous.equations.len(), 1); +} + #[test] fn metadata_attachment_lower_error_preserves_dae_source_span() { let span = sim_source_span(9, 21, 34); @@ -261,6 +312,75 @@ fn simulation_structural_lowering_demotes_vector_state_with_only_alias_rows() { ); } +#[test] +fn post_boundary_prep_demotes_newly_exposed_exact_state_alias() { + let mut dae = dae::Dae::new(); + dae.variables.states.insert( + VarName::new("direct.w"), + dae::Variable::new(VarName::new("direct.w"), fixture_span()), + ); + dae.variables.states.insert( + VarName::new("inverse.w"), + dae::Variable::new(VarName::new("inverse.w"), fixture_span()), + ); + dae.variables.algebraics.insert( + VarName::new("drive"), + dae::Variable::new(VarName::new("drive"), fixture_span()), + ); + dae.variables.algebraics.insert( + VarName::new("inverse.a"), + dae::Variable::new(VarName::new("inverse.a"), fixture_span()), + ); + + dae.continuous + .equations + .push(eq(sub(der(var("direct.w")), var("drive")))); + dae.continuous.equations.push(dae::Equation { + lhs: None, + rhs: sub(var("inverse.w"), var("direct.w")), + span: fixture_span(), + origin: "post-boundary state alias".to_string(), + scalar_count: 1, + }); + dae.continuous.equations.push(dae::Equation { + lhs: Some(reference("inverse.a")), + rhs: der(var("inverse.w")), + span: fixture_span(), + origin: "post-boundary derivative alias".to_string(), + scalar_count: 1, + }); + + let changed = prepare_dae_after_boundary_elimination(&mut dae, &[]) + .expect("post-boundary exact aliases should demote duplicate states"); + + assert!(changed); + assert!(dae.variables.states.contains_key(&VarName::new("direct.w"))); + assert!( + !dae.variables + .states + .contains_key(&VarName::new("inverse.w")) + ); + assert!( + dae.variables + .algebraics + .contains_key(&VarName::new("inverse.w")) + ); + assert!( + dae.continuous + .equations + .iter() + .all(|eq| !rumoca_ir_dae::expr_contains_der_of(&eq.rhs, &VarName::new("inverse.w"))), + "post-boundary prep must rewrite derivative users to the canonical state" + ); + assert!( + dae.continuous + .equations + .iter() + .any(|eq| rumoca_ir_dae::expr_contains_der_of(&eq.rhs, &VarName::new("direct.w"))), + "canonical state derivative should remain represented after alias rewrite" + ); +} + #[test] fn simulation_structural_lowering_differentiates_vector_function_constraint_for_coupled_state() { let dae = quaternion_constraint_dae(); @@ -296,9 +416,11 @@ fn simulation_structural_lowering_reports_state_metadata_before_elimination() { .expect("y should remain visible"); let selected_count = usize::from(x_meta.is_state) + usize::from(y_meta.is_state); + assert_eq!(model.state_scalar_count(), 1); assert_eq!( - selected_count, 1, - "exact alias component should report one selected state" + selected_count, + model.state_scalar_count(), + "reported state metadata must match the final solve layout" ); } @@ -777,6 +899,19 @@ fn call(name: &str, args: Vec) -> Expression { } } +fn named_arg(name: &str, args: Vec) -> Expression { + Expression::FunctionCall { + name: rumoca_core::Reference::from(format!( + "{}{}", + rumoca_core::NAMED_FUNCTION_ARG_PREFIX, + name + )), + args, + is_constructor: true, + span: fixture_span(), + } +} + fn component_ref(name: &str) -> rumoca_core::ComponentReference { let span = fixture_span(); rumoca_core::ComponentReference { @@ -937,6 +1072,49 @@ fn jacobian_for_dae_assembles_named_matrix_and_flags_zero_pivots() { assert!(report.error.is_none()); } +#[test] +fn lower_dae_for_simulation_preserves_matrix_derivative_state_slots() { + let mut dae = dae::Dae::new(); + let mut r = dae::Variable::new(VarName::new("R"), fixture_span()); + r.dims = vec![3, 3]; + dae.variables.states.insert(VarName::new("R"), r); + let mut skew = dae::Variable::new(VarName::new("skew"), fixture_span()); + skew.dims = vec![3, 3]; + dae.variables.algebraics.insert(VarName::new("skew"), skew); + + dae.continuous.equations.push(eq(sub( + var("skew"), + array(vec![ + array(vec![real(0.0), real(-1.0), real(0.0)]), + array(vec![real(1.0), real(0.0), real(0.0)]), + array(vec![real(0.0), real(0.0), real(0.0)]), + ]), + ))); + dae.continuous + .equations + .push(eq(sub(der(var("R")), mul(var("R"), var("skew"))))); + + let structurally_lowered = structurally_lower_dae_for_simulation(&dae, &SimOptions::default()) + .expect("matrix derivative state should structurally lower"); + assert!( + structurally_lowered + .dae + .variables + .states + .contains_key(&VarName::new("R")), + "structural lowering must retain matrix state R" + ); + + let model = lower_dae_for_simulation(&dae, &SimOptions::default()) + .expect("matrix derivative state should lower to solve IR"); + assert_eq!(model.problem.solve_layout.state_scalar_count, 9); + assert!( + model.problem.solve_layout.solver_maps.names[..9] + .iter() + .any(|name| name == "R[1,2]") + ); +} + #[test] fn eval_dae_at_rejects_unknown_state_name() { let err = eval_dae_at( diff --git a/crates/rumoca-solver-diffsol/src/bdf.rs b/crates/rumoca-solver-diffsol/src/bdf.rs index 945fdc462..291dc9058 100644 --- a/crates/rumoca-solver-diffsol/src/bdf.rs +++ b/crates/rumoca-solver-diffsol/src/bdf.rs @@ -1,9 +1,4 @@ -use std::{ - any::Any, - collections::{BTreeMap, BTreeSet}, - fmt::Display, - sync::Mutex, -}; +use std::{any::Any, collections::BTreeSet, fmt::Display, sync::Mutex}; use diffsol::{ BdfState, MatrixCommon, OdeEquations, OdeEquationsImplicit, OdeSolverMethod, OdeSolverProblem, @@ -289,6 +284,7 @@ where S: OdeSolverMethod<'a, Eqn>, { runtime_params.borrow_mut().copy_from_slice(params); + let previous_h = solver.state().h; let problem = solver.problem(); let mut fresh_state = S::State::new_without_initialise(problem) .map_err(|err| SimError::SolverError(format!("solver state reset: {err}")))?; @@ -299,24 +295,62 @@ where *state.t = t; } fresh_state.set_step_size(problem.h0, &problem.atol, problem.rtol, &problem.eqn, 1); + apply_event_restart_step_size::(&mut fresh_state, previous_h, h_cap); solver.set_state(fresh_state); - cap_solver_step_size(solver, h_cap); + mark_solver_state_modified_for_reinit(solver); Ok(()) } -fn cap_solver_step_size<'a, Eqn, S>(solver: &mut S, h_cap: f64) +fn apply_event_restart_step_size<'a, Eqn, S>(state: &mut S::State, previous_h: f64, h_cap: f64) where Eqn: OdeEquations + 'a, Eqn::V: VectorHost, S: OdeSolverMethod<'a, Eqn>, { - let state = solver.state_mut(); - if *state.h > h_cap { - *state.h = h_cap; + if let Some(restart_h) = event_restart_step_size(previous_h, h_cap) { + *state.as_mut().h = restart_h; + } +} + +fn event_restart_step_size(previous_h: f64, h_cap: f64) -> Option { + if !previous_h.is_finite() || previous_h == 0.0 || !h_cap.is_finite() || h_cap <= 0.0 { + return None; } + let magnitude = previous_h.abs().min(h_cap); + (magnitude > 0.0).then(|| magnitude.copysign(previous_h)) +} + +fn mark_solver_state_modified_for_reinit<'a, Eqn, S>(solver: &mut S) +where + Eqn: OdeEquations + 'a, + Eqn::V: VectorHost, + S: OdeSolverMethod<'a, Eqn>, +{ + let _state = solver.state_mut(); } pub(crate) fn can_use_state_only_bdf(model: &solve::SolveModel) -> Result { + match state_only_bdf_eligibility(model)? { + StateOnlyBdfEligibility::Eligible => Ok(true), + StateOnlyBdfEligibility::Ineligible(reason) => { + tracing::debug!( + target: "rumoca_solver_diffsol::bdf_path", + reason = reason.as_str(), + "state-only BDF ineligible" + ); + Ok(false) + } + } +} + +enum StateOnlyBdfEligibility { + Eligible, + Ineligible(String), +} + +fn state_only_bdf_eligibility( + model: &solve::SolveModel, +) -> Result { let state_count = model.state_scalar_count(); let derivative_rhs_len = model .problem @@ -325,12 +359,21 @@ pub(crate) fn can_use_state_only_bdf(model: &solve::SolveModel) -> Result, -) -> Result { +) -> Result, SimError> { let state_count = model.state_scalar_count(); let solver_count = model.solver_scalar_count(); let implicit_rows = solve_eval::to_scalar_program_block(&model.problem.continuous.implicit_rhs)?; - // The projection plan references residual rows by OUTPUT index; a program may - // now emit several outputs, so the producer map is bounded by output count - // and resolved to its producing program. - let Some(producer_rows) = projection_producer_rows(model, implicit_rows.output_count()) else { - return Ok(false); + let producer_programs = match solve_eval::algebraic_projection_producer_programs(model) { + Ok(producer_rows) => producer_rows, + Err(error) => return Ok(Some(error.to_string())), }; let mut needed = BTreeSet::new(); let mut stack = direct_deps.into_iter().collect::>(); @@ -367,49 +408,32 @@ fn projection_plan_covers_non_state_loads( if index < state_count || !needed.insert(index) { continue; } - let Some(output_idx) = producer_rows.get(&index).copied() else { - return Ok(false); - }; - let Some(program_idx) = implicit_rows.program_index_for_output(output_idx) else { - return Ok(false); + let Some(program_idx) = producer_programs.get(&index).copied() else { + return Ok(Some(format!( + "missing projection producer for solver_y[{index}] ({})", + solver_name(model, index) + ))); }; let Some(row) = implicit_rows.programs.get(program_idx) else { - return Ok(false); + return Ok(Some(format!( + "missing implicit producer program {program_idx} for solver_y[{index}] ({})", + solver_name(model, index) + ))); }; stack.extend(non_state_y_loads(row, state_count, solver_count)); } - Ok(true) + Ok(None) } -fn projection_producer_rows( - model: &solve::SolveModel, - implicit_row_count: usize, -) -> Option> { - let mut producer_rows = BTreeMap::new(); - for block in &model.problem.continuous.algebraic_projection_plan.blocks { - for (row_idx, target_index) in block - .rows - .iter() - .copied() - .zip(block.y_indices.iter().copied()) - .chain( - block - .causal_steps - .iter() - .map(|step| (step.row, step.y_index)), - ) - { - if row_idx >= implicit_row_count { - return None; - } - if let Some(previous_row) = producer_rows.insert(target_index, row_idx) - && previous_row != row_idx - { - return None; - } - } - } - Some(producer_rows) +fn solver_name(model: &solve::SolveModel, index: usize) -> &str { + model + .problem + .solve_layout + .solver_maps + .names + .get(index) + .map(String::as_str) + .unwrap_or("") } fn non_state_y_loads( @@ -432,3 +456,23 @@ fn non_state_y_loads( loads.dedup(); loads } + +#[cfg(test)] +mod tests { + use super::event_restart_step_size; + + #[test] + fn event_restart_preserves_accepted_scale_direction_and_cap() { + assert_eq!(event_restart_step_size(1.0e-4, 1.0e-3), Some(1.0e-4)); + assert_eq!(event_restart_step_size(2.0e-3, 1.0e-3), Some(1.0e-3)); + assert_eq!(event_restart_step_size(-2.0e-3, 1.0e-3), Some(-1.0e-3)); + } + + #[test] + fn event_restart_rejects_invalid_scale_without_overwriting_fresh_state() { + assert_eq!(event_restart_step_size(0.0, 1.0e-3), None); + assert_eq!(event_restart_step_size(f64::NAN, 1.0e-3), None); + assert_eq!(event_restart_step_size(1.0e-4, f64::INFINITY), None); + assert_eq!(event_restart_step_size(1.0e-4, 0.0), None); + } +} diff --git a/crates/rumoca-solver-diffsol/src/init_projection.rs b/crates/rumoca-solver-diffsol/src/init_projection.rs index 77e80a0fb..a7a7d279a 100644 --- a/crates/rumoca-solver-diffsol/src/init_projection.rs +++ b/crates/rumoca-solver-diffsol/src/init_projection.rs @@ -50,6 +50,7 @@ pub(crate) fn initialize_state_runtime_values( model, opts, runtime, + equilibrium_model, current_y, params, current_t: t_start, @@ -133,6 +134,7 @@ struct StateInitialEventUpdates<'a> { model: &'a solve::SolveModel, opts: &'a SimOptions, runtime: &'a SolveRuntime, + equilibrium_model: &'a OdeModel, current_y: &'a mut [f64], params: &'a mut [f64], current_t: f64, @@ -148,6 +150,7 @@ fn apply_state_initial_event_updates( model, opts, runtime, + equilibrium_model, current_y, params, current_t, @@ -168,7 +171,16 @@ fn apply_state_initial_event_updates( dynamic_event, apply_without_initial_event: true, }, - |y, p, t| refresh_algebraics_and_detect_changes(runtime, y, p, t, tol), + |y, p, t| { + project_algebraics_and_detect_changes( + equilibrium_model, + y, + p, + t, + equilibrium_model.state_count_for_projection(), + tol, + ) + }, )?; commit_pre_params_after_event(model, current_y, params, tol); Ok(outcome) @@ -280,14 +292,12 @@ impl RuntimeEventBoundaryHandler for EventObservation<'_> { type Error = SimError; fn on_event_time(&mut self, event_t: f64, _event: RuntimeEventStop) -> Result<(), Self::Error> { - refresh_observation_rows_and_relation_memory( + let (event_y, event_p) = event_time_observation_values( self.model, - self.runtime, - self.equilibrium_model, self.y, self.params, - event_t, - self.tol, + self.event_pre_y, + self.event_pre_p, )?; let mut samples = SampleRecorder { runtime: Some(self.runtime), @@ -298,8 +308,8 @@ impl RuntimeEventBoundaryHandler for EventObservation<'_> { record_sample_if_new( &mut samples, SamplePoint { - y: self.y, - params: self.params, + y: &event_y, + params: &event_p, t: event_t, }, )?; @@ -313,12 +323,14 @@ impl RuntimeEventBoundaryHandler for EventObservation<'_> { ) -> Result<(), Self::Error> { apply_event_updates_with_event_pre(EventUpdateInput { runtime: self.runtime, + ode_model: self.equilibrium_model, y: self.y, p: self.params, t: right_t, tol: self.tol, event_pre_y: self.event_pre_y, event_pre_p: self.event_pre_p, + root_relation_overrides: &[], })?; refresh_observation_rows_and_relation_memory( self.model, @@ -329,14 +341,10 @@ impl RuntimeEventBoundaryHandler for EventObservation<'_> { right_t, self.tol, )?; - let mut samples = SampleRecorder { - runtime: Some(self.runtime), - model: self.model, - recorded_times: &mut *self.recorded_times, - data: &mut *self.data, - }; - record_sample_if_new( - &mut samples, + crate::record_runtime_sample_at_distinct_time( + self.runtime, + &mut *self.recorded_times, + &mut *self.data, SamplePoint { y: self.y, params: self.params, @@ -346,3 +354,47 @@ impl RuntimeEventBoundaryHandler for EventObservation<'_> { Ok(()) } } + +fn event_time_observation_values( + model: &solve::SolveModel, + post_y: &[f64], + post_p: &[f64], + event_pre_y: &[f64], + event_pre_p: &[f64], +) -> Result<(Vec, Vec), SimError> { + // Continuous/algebraic equations (including table-backed runtime aliases) + // retain their event-entry value at the discontinuity sample. Discrete + // equations, however, take their newly settled value at that same event + // instant. The right-limit callback subsequently refreshes every runtime + // alias and algebraic using the next representable time. + let mut observed_y = event_pre_y.to_vec(); + let mut observed_p = event_pre_p.to_vec(); + for target in &model.problem.discrete.update_targets { + match *target { + solve::ScalarSlot::Y { index, .. } => { + let value = post_y.get(index).copied().ok_or_else(|| { + SimError::SolveIr(format!( + "event observation update target y[{index}] is out of bounds" + )) + })?; + let slot = observed_y.get_mut(index).ok_or_else(|| { + SimError::SolveIr(format!("event observation pre snapshot omits y[{index}]")) + })?; + *slot = value; + } + solve::ScalarSlot::P { index, .. } => { + let value = post_p.get(index).copied().ok_or_else(|| { + SimError::SolveIr(format!( + "event observation update target p[{index}] is out of bounds" + )) + })?; + let slot = observed_p.get_mut(index).ok_or_else(|| { + SimError::SolveIr(format!("event observation pre snapshot omits p[{index}]")) + })?; + *slot = value; + } + solve::ScalarSlot::Time | solve::ScalarSlot::Constant(_) => {} + } + } + Ok((observed_y, observed_p)) +} diff --git a/crates/rumoca-solver-diffsol/src/lib.rs b/crates/rumoca-solver-diffsol/src/lib.rs index b51606709..0518f3532 100644 --- a/crates/rumoca-solver-diffsol/src/lib.rs +++ b/crates/rumoca-solver-diffsol/src/lib.rs @@ -16,7 +16,11 @@ mod prepared; mod runtime; pub mod session; -use std::{cell::RefCell, rc::Rc, sync::Arc}; +use std::{ + cell::{Cell, RefCell}, + rc::Rc, + sync::Arc, +}; use bdf::can_use_state_only_bdf; pub(crate) use bdf::{ @@ -30,7 +34,8 @@ use diffsol::{ }; use init_projection::{EventObservation, initialize_state_runtime_values}; use rumoca_eval_solve::sim_driver::{ - SimDriverError, SolverAdvanceBackend, StateTrajectory, StepOutcome, simulate_state_targets, + RootStartBoundary, SimDriverError, SolverAdvanceBackend, StateTrajectory, StepOutcome, + simulate_state_targets, }; use rumoca_eval_solve::{ self as solve_eval, RowEvalContext, SolveRuntime, current_dynamic_time_event_stop, @@ -44,10 +49,12 @@ use rumoca_solver::{ replace_last_visible_values, runtime_event_horizon, runtime_root_event_application_time, timeline::sample_time_match_with_tol, }; +#[cfg(test)] +pub(crate) use runtime::apply_event_updates; pub(crate) use runtime::{ - EventUpdateInput, apply_event_updates, apply_event_updates_with_event_pre, - apply_initialization_updates, refresh_algebraics_and_detect_changes, + EventUpdateInput, apply_event_updates_with_event_pre, apply_initialization_updates, seed_initial_discrete_values, settle_algebraics_and_relation_memory, + settle_algebraics_and_relation_memory_with_overrides, }; use runtime::{check_no_state_initialization, simulate_no_state_solve_ir}; @@ -56,6 +63,40 @@ type Vector = ::V; type Scalar = ::T; pub(crate) type LinearSolver = FaerSparseLU; pub(crate) type RuntimeParameters = Rc>>; +#[derive(Clone)] +pub(crate) struct AcceptedSolverSeeds { + derivative: Vec, + root: Vec, + observation: Vec, +} +pub(crate) type AcceptedSolverY = Rc>; +pub(crate) type RootStartTime = Rc>; +#[derive(Clone, Debug, PartialEq)] +pub(crate) enum RootStartMode { + Initial, + Root(Vec<(usize, f64)>), + Scheduled, +} +pub(crate) type RootStartModeHandle = Rc>; + +fn with_root_start_boundary( + time: &RootStartTime, + mode: &RootStartModeHandle, + new_time: f64, + new_mode: RootStartMode, + reset: impl FnOnce() -> Result, +) -> Result { + let old_time = time.get(); + let old_mode = mode.borrow().clone(); + time.set(new_time); + *mode.borrow_mut() = new_mode; + let result = reset(); + if result.is_err() { + time.set(old_time); + *mode.borrow_mut() = old_mode; + } + result +} pub use error::SimError; pub(crate) use ode::{ OdeModel, build_ode_problem_with_runtime_params_and_initial, @@ -143,10 +184,14 @@ pub fn check_initialization(model: &solve::SolveModel, opts: &SimOptions) -> Res &mut current_t, )?; let runtime_params: RuntimeParameters = Rc::new(RefCell::new(params.clone())); + let root_start_time = Rc::new(Cell::new(current_t)); + let root_start_mode = Rc::new(RefCell::new(RootStartMode::Initial)); let problem = build_ode_problem_with_runtime_params_and_initial( model, opts, runtime_params, + root_start_time, + root_start_mode, current_t, current_y.clone(), equilibrium_model.clone(), @@ -229,12 +274,16 @@ fn simulate_with_states( // Shared runtime params captured by ODE closures and updated by event handlers. let runtime_params: RuntimeParameters = Rc::new(RefCell::new(params.clone())); + let root_start_time = Rc::new(Cell::new(current_t)); + let root_start_mode = Rc::new(RefCell::new(RootStartMode::Initial)); // Build the ODE problem once — the persistent BDF solver borrows it for the // full simulation lifetime. let problem = build_ode_problem_with_runtime_params_and_initial( model, opts, runtime_params.clone(), + root_start_time.clone(), + root_start_mode.clone(), current_t, current_y.clone(), equilibrium_model.clone(), @@ -246,6 +295,8 @@ fn simulate_with_states( equilibrium_model, runtime, runtime_params: runtime_params.clone(), + root_start_time, + root_start_mode, problem: &problem, current_y: ¤t_y, params: ¶ms, @@ -292,6 +343,8 @@ where equilibrium_model: &'a Arc, runtime: &'a Arc, runtime_params: RuntimeParameters, + root_start_time: RootStartTime, + root_start_mode: RootStartModeHandle, problem: &'a OdeSolverProblem, current_y: &'b [f64], params: &'b [f64], @@ -311,6 +364,8 @@ where equilibrium_model, runtime, runtime_params, + root_start_time, + root_start_mode, problem, current_y, params, @@ -338,6 +393,9 @@ where equilibrium_model: equilibrium_model.as_ref(), runtime: runtime.as_ref(), runtime_params, + root_start_time, + root_start_mode, + accepted_solver_y: None, opts, mode: DiffsolMode::General, }, @@ -362,6 +420,9 @@ where equilibrium_model: equilibrium_model.as_ref(), runtime: runtime.as_ref(), runtime_params, + root_start_time, + root_start_mode, + accepted_solver_y: None, opts, mode: DiffsolMode::General, }, @@ -469,13 +530,18 @@ fn simulate_state_only_bdf( }, &initial_observations, )?; - let runtime_params: RuntimeParameters = Rc::new(RefCell::new(params.clone())); + let accepted_solver_y = initial_accepted_solver_y(¤t_y); + let root_start_time = Rc::new(Cell::new(current_t)); + let root_start_mode = Rc::new(RefCell::new(RootStartMode::Initial)); let eval_counters = new_bdf_eval_counters(); let problem = build_state_ode_problem_with_runtime_params_and_initial( model, opts, runtime_params.clone(), + root_start_time.clone(), + root_start_mode.clone(), + accepted_solver_y.clone(), current_t, current_state.clone(), eval_counters.clone(), @@ -493,13 +559,13 @@ fn simulate_state_only_bdf( equilibrium_model, runtime, runtime_params: runtime_params.clone(), + root_start_time, + root_start_mode, + accepted_solver_y: Some(accepted_solver_y), opts, mode: DiffsolMode::StateOnly, }); - // Drive the reduced state-only solver through the *same* backend-neutral - // output / event / root loop as the general path; `DiffsolMode::StateOnly` - // (inside the backend) projects the reduced state to the full solver_y. let result = simulate_state_targets( model, opts, @@ -535,6 +601,14 @@ fn simulate_state_only_bdf( ) } +fn initial_accepted_solver_y(current_y: &[f64]) -> AcceptedSolverY { + Rc::new(RefCell::new(AcceptedSolverSeeds { + derivative: current_y.to_vec(), + root: current_y.to_vec(), + observation: current_y.to_vec(), + })) +} + fn initial_state_only_bdf_state( runtime: &SolveRuntime, problem: &diffsol::OdeSolverProblem, @@ -603,6 +677,10 @@ struct DiffsolAdvanceBackend<'a, Eqn, S> { equilibrium_model: &'a OdeModel, runtime: &'a SolveRuntime, runtime_params: RuntimeParameters, + root_start_time: RootStartTime, + root_start_mode: RootStartModeHandle, + accepted_solver_y: Option, + pending_reset_seed: RefCell>, opts: &'a SimOptions, mode: DiffsolMode, _eqn: std::marker::PhantomData Eqn>, @@ -614,6 +692,9 @@ struct DiffsolAdvanceBackendInputs<'a, S> { equilibrium_model: &'a OdeModel, runtime: &'a SolveRuntime, runtime_params: RuntimeParameters, + root_start_time: RootStartTime, + root_start_mode: RootStartModeHandle, + accepted_solver_y: Option, opts: &'a SimOptions, mode: DiffsolMode, } @@ -631,6 +712,10 @@ where equilibrium_model: inputs.equilibrium_model, runtime: inputs.runtime, runtime_params: inputs.runtime_params, + root_start_time: inputs.root_start_time, + root_start_mode: inputs.root_start_mode, + accepted_solver_y: inputs.accepted_solver_y, + pending_reset_seed: RefCell::new(None), opts: inputs.opts, mode: inputs.mode, _eqn: std::marker::PhantomData, @@ -640,6 +725,80 @@ where fn tol(&self) -> f64 { self.opts.atol.max(1.0e-10) } + + fn project_observation_and_commit( + &self, + native: &[f64], + t: f64, + params: &[f64], + ) -> Result, SimDriverError> { + let Some(accepted) = &self.accepted_solver_y else { + return Ok(native.to_vec()); + }; + let state_count = self.model.state_scalar_count().min(native.len()); + let mut seeds = accepted.borrow().clone(); + let mut observation = seeds.observation; + self.runtime.full_solver_y_with_guess( + t, + &native[..state_count], + params, + &mut observation, + self.tol(), + EVENT_UPDATE_MAX_ITERS, + )?; + seeds.observation = observation.clone(); + *accepted.borrow_mut() = seeds; + Ok(observation) + } + + fn project_root_and_commit( + &self, + native: &[f64], + t: f64, + params: &[f64], + ) -> Result<(), SimDriverError> { + let Some(accepted) = &self.accepted_solver_y else { + return Ok(()); + }; + let state_count = self.model.state_scalar_count().min(native.len()); + let mut root = accepted.borrow().root.clone(); + let mut out = vec![0.0; self.model.problem.events.root_conditions.len().max(1)]; + self.runtime + .eval_root_search_conditions_with_guess_into(solve_eval::RootSearchInput { + t, + state: &native[..state_count], + params, + guess: &mut root, + tol: self.tol(), + max_iters: EVENT_UPDATE_MAX_ITERS, + out: &mut out, + })?; + accepted.borrow_mut().root = root; + Ok(()) + } + + fn project_derivative_and_commit( + &self, + native: &[f64], + t: f64, + params: &[f64], + ) -> Result<(), SimDriverError> { + let Some(accepted) = &self.accepted_solver_y else { + return Ok(()); + }; + let state_count = self.model.state_scalar_count().min(native.len()); + let mut derivative = accepted.borrow().derivative.clone(); + let _ = self.runtime.eval_state_derivatives_with_guess( + t, + &native[..state_count], + params, + &mut derivative, + self.tol(), + EVENT_UPDATE_MAX_ITERS, + )?; + accepted.borrow_mut().derivative = derivative; + Ok(()) + } } impl<'a, Eqn, S> SolverAdvanceBackend for DiffsolAdvanceBackend<'a, Eqn, S> @@ -657,10 +816,25 @@ where } fn step(&mut self) -> Result { - match solver_call("BDF step", || self.solver.step()).map_err(sim_to_driver)? { + let reason = solver_call("BDF step", || self.solver.step()).map_err(sim_to_driver)?; + if matches!( + reason, + OdeSolverStopReason::TstopReached | OdeSolverStopReason::InternalTimestep + ) && self.accepted_solver_y.is_some() + { + let t = self.solver.state().t; + let native = self.solver.state().y.as_slice().to_vec(); + let params = self.runtime_params.borrow().clone(); + self.project_derivative_and_commit(&native, t, ¶ms)?; + self.project_root_and_commit(&native, t, ¶ms)?; + } + match reason { OdeSolverStopReason::TstopReached => Ok(StepOutcome::Stop), OdeSolverStopReason::InternalTimestep => Ok(StepOutcome::Internal), - OdeSolverStopReason::RootFound(t_root, _) => Ok(StepOutcome::Root { t_root }), + OdeSolverStopReason::RootFound(t_root, root_index) => Ok(StepOutcome::Root { + t_root, + root_indices: vec![root_index], + }), } } @@ -689,16 +863,7 @@ where ) -> Result, SimDriverError> { match self.mode { DiffsolMode::General => Ok(native.to_vec()), - DiffsolMode::StateOnly => { - let state_count = self.model.state_scalar_count().min(native.len()); - Ok(self.runtime.full_solver_y( - t, - &native[..state_count], - params, - self.tol(), - EVENT_UPDATE_MAX_ITERS, - )?) - } + DiffsolMode::StateOnly => self.project_observation_and_commit(native, t, params), } } @@ -718,13 +883,43 @@ where DiffsolMode::StateOnly => { let state_count = self.model.state_scalar_count().min(current_y.len()); let native = current_y[..state_count].to_vec(); - let dy = self.runtime.eval_state_derivatives( + let mut derivative = current_y.to_vec(); + let dy = self.runtime.eval_state_derivatives_with_guess( t, &native, params, + &mut derivative, self.tol(), EVENT_UPDATE_MAX_ITERS, )?; + let mut root = current_y.to_vec(); + let mut root_out = + vec![0.0; self.model.problem.events.root_conditions.len().max(1)]; + self.runtime.eval_root_search_conditions_with_guess_into( + solve_eval::RootSearchInput { + t, + state: &native, + params, + guess: &mut root, + tol: self.tol(), + max_iters: EVENT_UPDATE_MAX_ITERS, + out: &mut root_out, + }, + )?; + let mut observation = current_y.to_vec(); + self.runtime.full_solver_y_with_guess( + t, + &native, + params, + &mut observation, + self.tol(), + EVENT_UPDATE_MAX_ITERS, + )?; + *self.pending_reset_seed.borrow_mut() = Some(AcceptedSolverSeeds { + derivative, + root, + observation, + }); Ok((native, dy)) } } @@ -737,21 +932,52 @@ where params: &[f64], t: f64, h_cap: f64, + boundary: RootStartBoundary<'_>, ) -> Result<(), SimDriverError> { - reset_solver_state( - &mut self.solver, - &self.runtime_params, - native_y, - native_dy, - params, + let new_mode = match boundary { + RootStartBoundary::Scheduled => RootStartMode::Scheduled, + RootStartBoundary::Root(overrides) => RootStartMode::Root(overrides.to_vec()), + }; + let result = with_root_start_boundary( + &self.root_start_time, + &self.root_start_mode, t, - h_cap, - ) - .map_err(sim_to_driver) + new_mode, + || { + reset_solver_state( + &mut self.solver, + &self.runtime_params, + native_y, + native_dy, + params, + t, + h_cap, + ) + .map_err(sim_to_driver) + }, + ); + match result { + Ok(()) => { + if let (Some(accepted), Some(seed)) = ( + self.accepted_solver_y.as_ref(), + self.pending_reset_seed.borrow_mut().take(), + ) { + *accepted.borrow_mut() = seed; + } + Ok(()) + } + Err(error) => { + self.pending_reset_seed.borrow_mut().take(); + Err(error) + } + } } fn prefer_exact_output_steps(&self) -> bool { - self.mode == DiffsolMode::StateOnly && !model_is_event_free(self.model) + // The shared driver already stops at runtime event boundaries. Forcing + // BDF to land on every output sample compresses its multistep history on + // dense output grids and can collapse otherwise smooth state-only runs. + false } fn project_algebraics( @@ -761,27 +987,14 @@ where t: f64, tol: f64, ) -> Result { - match self.mode { - DiffsolMode::General => project_algebraics_and_detect_changes( - self.equilibrium_model, - y, - p, - t, - self.equilibrium_model.state_count_for_projection(), - tol, - ), - DiffsolMode::StateOnly => { - let before = y.to_vec(); - self.runtime.refresh_algebraic_and_output_slots( - t, - y, - p, - tol, - EVENT_UPDATE_MAX_ITERS, - )?; - Ok(values_changed(&before, y, tol)) - } - } + project_algebraics_and_detect_changes( + self.equilibrium_model, + y, + p, + t, + self.equilibrium_model.state_count_for_projection(), + tol, + ) } fn derivative_guess(&self, y: &[f64], p: &[f64], t: f64) -> Result, SimDriverError> { @@ -904,6 +1117,27 @@ pub(crate) fn record_sample_if_new( Ok(()) } +pub(crate) fn record_runtime_sample_at_distinct_time( + runtime: &SolveRuntime, + recorded_times: &mut Vec, + data: &mut [Vec], + sample: SamplePoint<'_>, +) -> Result<(), SimError> { + let values = runtime + .visible_values(sample.y, sample.params, sample.t) + .map_err(|err| SimError::SolveIr(err.to_string()))?; + if recorded_times + .last() + .is_some_and(|last| last.to_bits() == sample.t.to_bits()) + { + replace_last_visible_values(data, &values)?; + return Ok(()); + } + recorded_times.push(sample.t); + push_visible_values(data, &values)?; + Ok(()) +} + fn record_initial_samples( recorder: &mut SampleRecorder<'_>, runtime: &SolveRuntime, @@ -1003,25 +1237,6 @@ fn visible_values( .map_err(|err| SimError::SolveIr(err.to_string())) } -fn values_changed(before: &[f64], after: &[f64], tol: f64) -> bool { - before - .iter() - .zip(after.iter()) - .any(|(before, after)| (*before - *after).abs() > tol) -} - -/// True when the model has no discontinuities (zero-crossing roots, scheduled -/// time events, or discrete `when` updates), so the BDF solution is smooth and -/// safe to dense-output / interpolate at arbitrary times. -fn model_is_event_free(model: &solve::SolveModel) -> bool { - let events = &model.problem.events; - let discrete = &model.problem.discrete; - events.root_conditions.is_empty() - && events.scheduled_time_events.is_empty() - && discrete.update_targets.is_empty() - && discrete.runtime_assignment_targets.is_empty() -} - fn trace_bdf_step_failure( equilibrium_model: &OdeModel, y: &[f64], diff --git a/crates/rumoca-solver-diffsol/src/ode.rs b/crates/rumoca-solver-diffsol/src/ode.rs index 9b7f8486b..82107b01e 100644 --- a/crates/rumoca-solver-diffsol/src/ode.rs +++ b/crates/rumoca-solver-diffsol/src/ode.rs @@ -12,7 +12,10 @@ use rumoca_eval_solve::{ use rumoca_ir_solve as solve; use rumoca_solver::{AlgebraicProjectionModel, PreparedMassMatrix, RuntimeSolveError, SimOptions}; -use crate::{EVENT_UPDATE_MAX_ITERS, Matrix, RuntimeParameters, Scalar, SimError, Vector}; +use crate::{ + AcceptedSolverY, EVENT_UPDATE_MAX_ITERS, Matrix, RootStartMode, RootStartModeHandle, + RootStartTime, RuntimeParameters, Scalar, SimError, Vector, +}; #[derive(Debug, Default)] pub(crate) struct BdfEvalCounters { @@ -273,9 +276,18 @@ impl AlgebraicProjectionModel for OdeModel { p: &[f64], t: f64, ) -> Result, RuntimeSolveError> { + let Some((program_index, output_offset)) = self + .implicit_scalar_rhs + .program_position_for_output_index(row_idx) + else { + return Ok(None); + }; + if output_offset != 0 { + return Ok(None); + } self.implicit_scalar_rhs .eval_target_assignment_row_unchecked_with_context( - row_idx, + program_index, target_y_index, y, p, @@ -342,10 +354,16 @@ pub(crate) fn validate_model(model: &solve::SolveModel) -> Result<(), SimError> Ok(()) } +#[expect( + clippy::too_many_arguments, + reason = "the general ODE builder forwards independently owned solver callback state" +)] pub(crate) fn build_ode_problem_with_runtime_params_and_initial( model: &solve::SolveModel, opts: &SimOptions, runtime_params: RuntimeParameters, + root_start_time: RootStartTime, + root_start_mode: RootStartModeHandle, t_start: f64, initial_y: Vec, ode_model: Arc, @@ -363,15 +381,25 @@ pub(crate) fn build_ode_problem_with_runtime_params_and_initial( t_start, initial_y, Some(runtime_params), + root_start_time, + root_start_mode, ode_model, root_runtime, ) } +#[expect( + clippy::too_many_arguments, + clippy::too_many_lines, + reason = "the reduced ODE builder assembles three stateful diffsol callbacks" +)] pub(crate) fn build_state_ode_problem_with_runtime_params_and_initial( model: &solve::SolveModel, opts: &SimOptions, runtime_params: RuntimeParameters, + root_start_time: RootStartTime, + root_start_mode: RootStartModeHandle, + accepted_solver_y: AcceptedSolverY, t_start: f64, initial_state: Vec, eval_counters: Option>, @@ -394,13 +422,26 @@ pub(crate) fn build_state_ode_problem_with_runtime_params_and_initial( let rhs_params = Some(runtime_params.clone()); let jac_params = Some(runtime_params.clone()); let root_params = Some(runtime_params); + let root_start_time_for_eval = root_start_time; + let root_start_mode_for_eval = root_start_mode; + let rhs_accepted = accepted_solver_y.clone(); + let jac_accepted = accepted_solver_y.clone(); + let root_accepted = accepted_solver_y; let tol = opts.atol.max(1.0e-10); - let rhs_fn = move |y: &Vector, p: &Vector, t: Scalar, out: &mut Vector| { let start = rhs_counters.as_ref().map(|_| Instant::now()); with_runtime_params(&rhs_params, p.as_slice(), |params| { + let mut trial = rhs_accepted.borrow().derivative.clone(); if rhs_runtime - .eval_state_derivatives_into(t, y.as_slice(), params, tol, 256, out.as_mut_slice()) + .eval_state_derivatives_with_guess_into( + t, + y.as_slice(), + params, + &mut trial, + tol, + 256, + out.as_mut_slice(), + ) .is_err() { fill_eval_error(out.as_mut_slice()); @@ -413,8 +454,9 @@ pub(crate) fn build_state_ode_problem_with_runtime_params_and_initial( let jac_fn = move |y: &Vector, p: &Vector, t: Scalar, v: &Vector, out: &mut Vector| { let start = jac_counters.as_ref().map(|_| Instant::now()); with_runtime_params(&jac_params, p.as_slice(), |params| { + let mut trial = jac_accepted.borrow().derivative.clone(); if jac_runtime - .eval_state_jacobian_v_ad_into( + .eval_state_jacobian_v_ad_with_guess_into( solve_eval::AlgebraicLinearization { t, params, @@ -425,6 +467,7 @@ pub(crate) fn build_state_ode_problem_with_runtime_params_and_initial( }, y.as_slice(), v.as_slice(), + &mut trial, out.as_mut_slice(), ) .is_err() @@ -439,17 +482,31 @@ pub(crate) fn build_state_ode_problem_with_runtime_params_and_initial( let root_fn = move |y: &Vector, p: &Vector, t: Scalar, out: &mut Vector| { let start = root_counters.as_ref().map(|_| Instant::now()); with_runtime_params(&root_params, p.as_slice(), |params| { - if root_runtime - .eval_root_conditions_into( + let mut trial = root_accepted.borrow().root.clone(); + let evaluated = root_runtime.eval_root_search_conditions_with_guess_into( + solve_eval::RootSearchInput { t, - y.as_slice(), + state: y.as_slice(), + params, + guess: &mut trial, + tol, + max_iters: EVENT_UPDATE_MAX_ITERS, + out: out.as_mut_slice(), + }, + ); + let root_start = root_start_time_for_eval.get(); + let initialized = evaluated.and_then(|()| { + initialize_root_search_values_at_time( + &root_runtime, params, tol, - EVENT_UPDATE_MAX_ITERS, + &root_start_mode_for_eval.borrow(), out.as_mut_slice(), + t, + root_start, ) - .is_err() - { + }); + if initialized.is_err() { fill_eval_error(out.as_mut_slice()); } }); @@ -476,12 +533,18 @@ pub(crate) fn build_state_ode_problem_with_runtime_params_and_initial( .map_err(|err| SimError::SolverError(format!("ODE problem builder failed: {err}"))) } +#[expect( + clippy::too_many_arguments, + reason = "the internal ODE builder owns independent root and runtime callback state" +)] fn build_ode_problem_with_initial( model: &solve::SolveModel, opts: &SimOptions, t_start: f64, initial_y: Vec, runtime_params: Option, + root_start_time: RootStartTime, + root_start_mode: RootStartModeHandle, ode_model: Arc, root_runtime: Arc, ) -> Result< @@ -505,6 +568,8 @@ fn build_ode_problem_with_initial( let rhs_runtime_params = runtime_params.clone(); let jac_runtime_params = runtime_params.clone(); let root_runtime_params = runtime_params.clone(); + let root_start_time_for_eval = root_start_time; + let root_start_mode_for_eval = root_start_mode; let jac_model = ode_model.clone(); let tol = opts.atol.max(1.0e-10); let jac_fn = move |y: &Vector, p: &Vector, t: Scalar, v: &Vector, out: &mut Vector| { @@ -519,17 +584,26 @@ fn build_ode_problem_with_initial( }; let root_fn = move |y: &Vector, p: &Vector, t: Scalar, out: &mut Vector| { with_runtime_params(&root_runtime_params, p.as_slice(), |params| { - if root_runtime - .eval_root_conditions_into( - t, - y.as_slice(), + let evaluated = root_runtime.eval_root_search_conditions_into( + t, + y.as_slice(), + params, + tol, + EVENT_UPDATE_MAX_ITERS, + out.as_mut_slice(), + ); + let initialized = evaluated.and_then(|()| { + initialize_root_search_values_at_time( + &root_runtime, params, tol, - EVENT_UPDATE_MAX_ITERS, + &root_start_mode_for_eval.borrow(), out.as_mut_slice(), + t, + root_start_time_for_eval.get(), ) - .is_err() - { + }); + if initialized.is_err() { fill_eval_error(out.as_mut_slice()); } }); @@ -570,6 +644,27 @@ fn build_ode_problem_with_initial( .map_err(|err| SimError::SolverError(format!("ODE problem builder failed: {err}"))) } +fn initialize_root_search_values_at_time( + runtime: &SolveRuntime, + params: &[f64], + tol: f64, + mode: &RootStartMode, + out: &mut [f64], + t: f64, + root_start: f64, +) -> Result<(), RuntimeSolveError> { + if t != root_start { + return Ok(()); + } + match mode { + RootStartMode::Initial => runtime.neutralize_initial_root_search_values(params, tol, out), + RootStartMode::Root(overrides) => { + runtime.apply_consumed_root_search_overrides(params, tol, overrides, out) + } + RootStartMode::Scheduled => Ok(()), + } +} + fn fill_eval_error(out: &mut [f64]) { out.fill(f64::NAN); } diff --git a/crates/rumoca-solver-diffsol/src/runtime.rs b/crates/rumoca-solver-diffsol/src/runtime.rs index ec9886eea..6b1cf9cfc 100644 --- a/crates/rumoca-solver-diffsol/src/runtime.rs +++ b/crates/rumoca-solver-diffsol/src/runtime.rs @@ -1,6 +1,7 @@ use super::*; use rumoca_eval_solve::{ - EventUpdateRowFilter, ProjectedEventUpdateInput, apply_discrete_slot_values, + EventUpdateRowFilter, ProjectedEventUpdateInput, ProjectedRuntimeSettleInput, + apply_discrete_slot_values, }; use rumoca_solver::{ EventActionOutcome, EventPreMode, NoStateEventStep, NoStateOrchestrationBackend, @@ -10,11 +11,11 @@ use rumoca_solver::{ pub(crate) fn settle_algebraics_and_relation_memory( runtime: &SolveRuntime, - _model: &OdeModel, + model: &OdeModel, y: &mut [f64], p: &mut [f64], t: f64, - _state_count: usize, + state_count: usize, tol: f64, ) -> Result<(), SimError> { runtime @@ -24,33 +25,43 @@ pub(crate) fn settle_algebraics_and_relation_memory( t, tol, EVENT_UPDATE_MAX_ITERS, - move |y, p| refresh_algebraics_and_detect_changes(runtime, y, p, t, tol), + move |y, p| project_algebraics_and_detect_changes(model, y, p, t, state_count, tol), ) .map_err(Into::into) } -pub(crate) fn refresh_algebraics_and_detect_changes( +#[expect( + clippy::too_many_arguments, + reason = "event projection requires both solve models, mutable stores, and root overrides" +)] +pub(crate) fn settle_algebraics_and_relation_memory_with_overrides( runtime: &SolveRuntime, + model: &OdeModel, y: &mut [f64], p: &mut [f64], t: f64, + state_count: usize, tol: f64, -) -> Result { - let before = y.to_vec(); - runtime.refresh_algebraic_and_output_slots(t, y, p, tol, EVENT_UPDATE_MAX_ITERS)?; - Ok(values_changed(&before, y, tol)) -} - -fn values_changed(before: &[f64], after: &[f64], tol: f64) -> bool { - before - .iter() - .zip(after.iter()) - .any(|(before, after)| (*before - *after).abs() > tol) + root_relation_overrides: &[(usize, f64)], +) -> Result<(), SimError> { + runtime + .settle_projected_runtime_and_relation_memory_with_overrides( + ProjectedRuntimeSettleInput { + y, + p, + t, + tol, + max_iters: EVENT_UPDATE_MAX_ITERS, + root_relation_overrides, + }, + move |y, p| project_algebraics_and_detect_changes(model, y, p, t, state_count, tol), + ) + .map_err(Into::into) } pub(crate) fn apply_event_updates( runtime: &SolveRuntime, - _ode_model: &OdeModel, + ode_model: &OdeModel, y: &mut [f64], p: &mut [f64], t: f64, @@ -60,23 +71,27 @@ pub(crate) fn apply_event_updates( let event_pre_p = p.to_vec(); apply_event_updates_with_event_pre(EventUpdateInput { runtime, + ode_model, y, p, t, tol, event_pre_y: &event_pre_y, event_pre_p: &event_pre_p, + root_relation_overrides: &[], }) } pub(crate) struct EventUpdateInput<'a> { pub(crate) runtime: &'a SolveRuntime, + pub(crate) ode_model: &'a OdeModel, pub(crate) y: &'a mut [f64], pub(crate) p: &'a mut [f64], pub(crate) t: f64, pub(crate) tol: f64, pub(crate) event_pre_y: &'a [f64], pub(crate) event_pre_p: &'a [f64], + pub(crate) root_relation_overrides: &'a [(usize, f64)], } pub(crate) fn apply_event_updates_with_event_pre( @@ -91,12 +106,14 @@ fn apply_event_updates_with_filter( ) -> Result<(), SimError> { let EventUpdateInput { runtime, + ode_model, y, p, t, tol, event_pre_y, event_pre_p, + root_relation_overrides, } = input; let outcome = runtime.apply_projected_event_update( ProjectedEventUpdateInput { @@ -108,9 +125,9 @@ fn apply_event_updates_with_filter( event_pre_p, max_iters: EVENT_UPDATE_MAX_ITERS, row_filter, - root_relation_overrides: &[], + root_relation_overrides, }, - project_algebraics_callback(runtime, t, tol), + project_algebraics_callback(ode_model, t, tol), )?; event_action_outcome_to_result(outcome, t) } @@ -141,11 +158,20 @@ pub(crate) fn apply_initialization_updates( } fn project_algebraics_callback( - runtime: &SolveRuntime, + model: &OdeModel, t: f64, tol: f64, ) -> impl FnMut(&mut [f64], &mut [f64]) -> Result + '_ { - move |y, p| refresh_algebraics_and_detect_changes(runtime, y, p, t, tol) + move |y, p| { + project_algebraics_and_detect_changes( + model, + y, + p, + t, + model.state_count_for_projection(), + tol, + ) + } } fn event_action_outcome_to_result( @@ -643,9 +669,9 @@ fn first_root_crossing_time( eval_refreshed_roots(runtime, y, p, t_end, tol, &mut end)?; let mut crossing = None; - for (a, b) in start.iter().zip(end.iter()) { + for (root_index, (a, b)) in start.iter().zip(end.iter()).enumerate() { if root_surface_crossed_or_near(*a, *b, tol) { - let root = bisect_first_root(runtime, model, y, p, t_start, t_end, tol)?; + let root = bisect_first_root(runtime, model, y, p, (t_start, t_end), root_index, tol)?; crossing = Some(crossing.map_or(root, |current| f64::min(current, root))); } } @@ -678,21 +704,20 @@ fn bisect_first_root( model: &OdeModel, y: &[f64], p: &[f64], - mut lo: f64, - mut hi: f64, + interval: (f64, f64), + root_index: usize, tol: f64, ) -> Result { + let (mut lo, mut hi) = interval; let mut lo_roots = vec![0.0; model.root_conditions.len()]; eval_refreshed_roots(runtime, y, p, lo, tol, &mut lo_roots)?; for _ in 0..ROOT_BISECTION_ITERS { let mid = lo + 0.5 * (hi - lo); let mut mid_roots = vec![0.0; model.root_conditions.len()]; eval_refreshed_roots(runtime, y, p, mid, tol, &mut mid_roots)?; - if lo_roots - .iter() - .zip(mid_roots.iter()) - .any(|(a, b)| a.signum() != b.signum() || root_surface_near_zero(*b, tol)) - { + let lo_root = lo_roots.get(root_index).copied().unwrap_or(0.0); + let mid_root = mid_roots.get(root_index).copied().unwrap_or(0.0); + if lo_root.signum() != mid_root.signum() || root_surface_near_zero(mid_root, 0.0) { hi = mid; } else { lo = mid; diff --git a/crates/rumoca-solver-diffsol/src/session.rs b/crates/rumoca-solver-diffsol/src/session.rs index 9224574c9..dae1cdf39 100644 --- a/crates/rumoca-solver-diffsol/src/session.rs +++ b/crates/rumoca-solver-diffsol/src/session.rs @@ -4,12 +4,14 @@ use diffsol::{OdeSolverMethod, VectorHost}; use indexmap::IndexMap; -use rumoca_eval_solve::{self as solve_eval, SolveRuntime}; +use rumoca_eval_solve::{ + self as solve_eval, SolveRuntime, sim_driver::post_root_relation_overrides, +}; use rumoca_ir_solve as solve; use rumoca_solver::{ SimOptions, event_solver_step_cap, runtime_root_event_application_time, time_match_with_tol, }; -use std::cell::RefCell; +use std::cell::{Cell, RefCell}; use std::rc::Rc; use std::sync::Arc; @@ -18,13 +20,17 @@ use crate::runtime::{ initialize_no_state_runtime, }; use crate::{ - LinearSolver, OdeModel, RuntimeParameters, SimError, apply_event_updates, bdf_derivative_guess, + EventUpdateInput, LinearSolver, OdeModel, RuntimeParameters, SimError, + apply_event_updates_with_event_pre, bdf_derivative_guess, build_ode_problem_with_runtime_params_and_initial, initial_bdf_state, reset_solver_state, - settle_algebraics_and_relation_memory, solver_call, validate_model, write_state_to_solver, + settle_algebraics_and_relation_memory, settle_algebraics_and_relation_memory_with_overrides, + solver_call, validate_model, write_state_to_solver, }; type StepFn = Box Result>; type ResetFn = Box Result<(), SimError>>; +type EventResetFn = Box Result<(), SimError>>; +type ProjectFn = Box Result, SimError>>; const SESSION_ADVANCE_EVENT_LIMIT: usize = 256; pub struct SimulationSession { @@ -42,15 +48,16 @@ struct BdfSession { step_fn: StepFn, time_fn: Box f64>, y_fn: Box Vec>, - event_reset_fn: Box Result<(), SimError>>, + event_reset_fn: EventResetFn, reset_fn: ResetFn, refresh_input_fn: Box Result<(), SimError>>, - project_fn: Box Result<(), SimError>>, + project_fn: ProjectFn, runtime: SolveRuntime, runtime_params: RuntimeParameters, reset_snapshot: BdfResetSnapshot, input_values: IndexMap, inputs_dirty: bool, + pending_root: Option, } #[derive(Clone)] @@ -60,9 +67,15 @@ struct BdfResetSnapshot { params: Vec, } -#[derive(Debug, Clone, Copy, Default)] +#[derive(Debug, Clone, Default)] struct StepAdvance { - hit_root: bool, + root: Option, +} + +#[derive(Debug, Clone)] +struct SessionPendingRoot { + t: f64, + indices: Vec, } struct RuntimeOnlyDriver { @@ -318,6 +331,8 @@ impl BdfSession { model, &opts, runtime_params.clone(), + Rc::new(Cell::new(opts.t_start)), + Rc::new(RefCell::new(crate::RootStartMode::Initial)), opts.t_start, initial_y.clone(), ode_model.clone(), @@ -360,7 +375,7 @@ impl BdfSession { let event_reset_model = model.clone(); let event_reset_opts = opts.clone(); let event_reset_params = runtime_params.clone(); - let event_reset_fn = Box::new(move |t_start: f64| { + let event_reset_fn = Box::new(move |t_start: f64, overrides: &[(usize, f64)]| { let initial_y = { let solver = event_reset_solver.borrow(); solver.state().y.as_slice().to_vec() @@ -368,19 +383,12 @@ impl BdfSession { let ode_model = Arc::new(OdeModel::new(&event_reset_model)?); let reset_runtime = SolveRuntime::new(&event_reset_model)?; let root_runtime = Arc::new(reset_runtime.clone()); - let initial_y = settled_problem_y( - &event_reset_model, - &reset_runtime, - &ode_model, - &event_reset_opts, - &event_reset_params, - t_start, - initial_y, - )?; let problem = build_ode_problem_with_runtime_params_and_initial( &event_reset_model, &event_reset_opts, event_reset_params.clone(), + Rc::new(Cell::new(t_start)), + Rc::new(RefCell::new(crate::RootStartMode::Root(overrides.to_vec()))), t_start, initial_y.clone(), ode_model.clone(), @@ -445,6 +453,7 @@ impl BdfSession { reset_snapshot, input_values: IndexMap::new(), inputs_dirty: false, + pending_root: None, }) } @@ -481,17 +490,27 @@ impl BdfSession { if target_time <= current_time { return Ok(()); } + if let Some(root) = self.pending_root.take() { + let reset_time = runtime_root_event_application_time(root.t, target_time); + let overrides = (self.project_fn)(reset_time, &root.indices)?; + (self.event_reset_fn)(reset_time, &overrides)?; + continue; + } if self.inputs_dirty { (self.refresh_input_fn)()?; self.inputs_dirty = false; } let advance = (self.step_fn)(target_time - current_time)?; - if !advance.hit_root { + let Some(root) = advance.root else { + return Ok(()); + }; + if time_match_with_tol(root.t, target_time) { + self.pending_root = Some(root); return Ok(()); } - (self.project_fn)()?; - let reset_time = runtime_root_event_application_time(self.time(), target_time); - (self.event_reset_fn)(reset_time)?; + let reset_time = runtime_root_event_application_time(root.t, target_time); + let overrides = (self.project_fn)(reset_time, &root.indices)?; + (self.event_reset_fn)(reset_time, &overrides)?; } Err(SimError::SolverError(format!( "event processing did not settle before t={target_time}" @@ -501,6 +520,7 @@ impl BdfSession { fn reset(&mut self, t_start: f64) -> Result<(), SimError> { self.input_values.clear(); self.inputs_dirty = false; + self.pending_root = None; (self.reset_fn)(t_start, &self.reset_snapshot) } @@ -704,7 +724,7 @@ fn make_project_fn( runtime: SolveRuntime, params: RuntimeParameters, opts: &SimOptions, -) -> Result Result<(), SimError>>, SimError> +) -> Result where Eqn: diffsol::OdeEquations + 'static, Eqn::V: VectorHost, @@ -714,7 +734,7 @@ where let event_model = model.clone(); let state_count = model.state_scalar_count(); let tol = opts.atol.max(1.0e-10); - Ok(Box::new(move || { + Ok(Box::new(move |event_t, root_indices| { project_session_algebraics( &solver, &event_model, @@ -723,44 +743,65 @@ where ¶ms, state_count, tol, + event_t, + root_indices, ) })) } +#[expect( + clippy::too_many_arguments, + reason = "session event projection bridges solver state, runtime state, and root metadata" +)] fn project_session_algebraics( solver: &Rc>, - _solve_model: &solve::SolveModel, + solve_model: &solve::SolveModel, runtime: &SolveRuntime, model: &OdeModel, params: &RuntimeParameters, state_count: usize, tol: f64, -) -> Result<(), SimError> + event_t: f64, + root_indices: &[usize], +) -> Result, SimError> where Eqn: diffsol::OdeEquations + 'static, Eqn::V: VectorHost, S: OdeSolverMethod<'static, Eqn>, { let mut solver = solver.borrow_mut(); - let t = solver.state().t; let mut y = solver.state().y.as_slice().to_vec(); + let event_pre_y = y.clone(); + let event_pre_p = params.borrow().clone(); + let overrides = post_root_relation_overrides(solve_model, root_indices, &event_pre_p, tol)?; { let mut params = params.borrow_mut(); - settle_algebraics_and_relation_memory( + settle_algebraics_and_relation_memory_with_overrides( runtime, model, &mut y, params.as_mut_slice(), - t, + event_t, state_count, tol, + &overrides, )?; - apply_event_updates(runtime, model, &mut y, params.as_mut_slice(), t, tol)?; + apply_event_updates_with_event_pre(EventUpdateInput { + runtime, + ode_model: model, + y: &mut y, + p: params.as_mut_slice(), + t: event_t, + tol, + event_pre_y: &event_pre_y, + event_pre_p: &event_pre_p, + root_relation_overrides: &overrides, + })?; } solver.state_mut().y.as_mut_slice().copy_from_slice(&y); let state = solver.state_clone(); solver.set_state(state); - Ok(()) + Ok(overrides) } fn step_solver_by( @@ -816,7 +857,7 @@ where diffsol::OdeSolverStopReason::TstopReached | diffsol::OdeSolverStopReason::InternalTimestep, ) => continue, - Ok(diffsol::OdeSolverStopReason::RootFound(t_root, _)) => { + Ok(diffsol::OdeSolverStopReason::RootFound(t_root, root_index)) => { // The free-running step overshoots the root (the solver state // sits at the natural step end, past `t_root`). The caller's // event handling assumes the solver is *at* the event instant, @@ -829,7 +870,12 @@ where solver .state_mut_back(event_t) .map_err(|err| SimError::SolverError(format!("state_mut_back: {err}")))?; - return Ok(StepAdvance { hit_root: true }); + return Ok(StepAdvance { + root: Some(SessionPendingRoot { + t: event_t, + indices: vec![root_index], + }), + }); } // Root lies beyond the requested interval: land on `target` and // defer the crossing to the next step, which resumes from @@ -1071,6 +1117,11 @@ mod tests { ); assert_eq!(session.get("force").unwrap(), Some(0.0)); + session + .advance_to(0.1) + .expect("repeating the same deadline must be idempotent"); + assert_eq!(session.get("force").unwrap(), Some(0.0)); + session .advance_to(0.101) .expect("next advance should process the event right-limit"); @@ -1142,6 +1193,9 @@ mod tests { initialization: solve::InitializationSolveSystem { residual: ComputeBlock::from_scalar_program_block(zero.clone()), row_targets: Vec::new(), + direct_families: Vec::new(), + required_target_ranges: Vec::new(), + fixed_target_ranges: Vec::new(), projection_indices: Vec::new(), projection_plan: solve::AlgebraicProjectionPlan::default(), update_rhs: solve::ScalarProgramBlock::default(), @@ -1229,11 +1283,7 @@ mod tests { zero_row(), zero_row(), ])), - row_targets: Vec::new(), - projection_indices: Vec::new(), - projection_plan: solve::AlgebraicProjectionPlan::default(), - update_rhs: solve::ScalarProgramBlock::default(), - update_targets: Vec::new(), + ..Default::default() }, discrete: solve::DiscreteSolveSystem::default(), events: solve::SolveEventPartition { @@ -1241,6 +1291,7 @@ mod tests { LinearOp::LoadY { dst: 0, index: 0 }, LinearOp::StoreOutput { src: 0 }, ]]), + root_relation_memory_targets: vec![Some(solve::scalar_slot_p(0))], ..Default::default() }, clocks: solve::SolveClockPartition::default(), diff --git a/crates/rumoca-solver-diffsol/src/tests/mod.rs b/crates/rumoca-solver-diffsol/src/tests/mod.rs index 8e7c82779..1e4289671 100644 --- a/crates/rumoca-solver-diffsol/src/tests/mod.rs +++ b/crates/rumoca-solver-diffsol/src/tests/mod.rs @@ -1,5 +1,28 @@ use super::*; +#[test] +fn root_start_boundary_is_visible_during_reset_and_rolls_back_on_failure() { + let time = Rc::new(Cell::new(1.0)); + let mode = Rc::new(RefCell::new(RootStartMode::Initial)); + + let error = with_root_start_boundary( + &time, + &mode, + 2.0, + RootStartMode::Root(vec![(3, 1.0)]), + || -> Result<(), &'static str> { + assert_eq!(time.get(), 2.0); + assert_eq!(*mode.borrow(), RootStartMode::Root(vec![(3, 1.0)])); + Err("reset failed") + }, + ) + .expect_err("failed reset should surface"); + + assert_eq!(error, "reset failed"); + assert_eq!(time.get(), 1.0); + assert_eq!(*mode.borrow(), RootStartMode::Initial); +} + macro_rules! fixture_span { () => { solve::source_span_from_offsets(49, 0, 1) @@ -180,6 +203,61 @@ fn state_only_bdf_accepts_transitive_projection_dependencies() { ); } +#[test] +fn state_only_bdf_accepts_partial_causal_hints_with_complete_producer_closure() { + let mut model = projected_derivative_model(); + model + .problem + .solve_layout + .solver_maps + .names + .push("b".to_string()); + model + .problem + .solve_layout + .solver_maps + .name_to_idx + .insert("b".to_string(), 2); + model + .problem + .solve_layout + .solver_maps + .base_to_indices + .insert("b".to_string(), vec![2]); + model.problem.solve_layout.algebraic_scalar_count = 2; + model.problem.continuous.derivative_rhs = solve::ComputeBlock::from_scalar_program_block( + solve::ScalarProgramBlock::with_source_span( + vec![vec![ + solve::LinearOp::LoadY { dst: 0, index: 2 }, + solve::LinearOp::StoreOutput { src: 0 }, + ]], + fixture_span!(), + ), + ); + model.problem.continuous.implicit_rhs = solve::ComputeBlock::from_scalar_program_block( + solve::ScalarProgramBlock::with_source_span( + vec![ + state_residual_row(), + y_minus_y_row(1, 0), + y_minus_y_row(2, 1), + ], + fixture_span!(), + ), + ); + model.problem.continuous.algebraic_projection_plan = solve::AlgebraicProjectionPlan { + blocks: vec![solve::AlgebraicProjectionBlock { + rows: vec![2, 1], + y_indices: vec![1, 2], + causal_steps: vec![solve::AlgebraicProjectionStep { row: 1, y_index: 1 }], + }], + }; + + assert!( + can_use_state_only_bdf(&model).expect("valid model should check BDF eligibility"), + "partial causal hints must not hide the remaining executable producer closure" + ); +} + #[test] fn state_only_periodic_events_do_not_advance_to_right_limit_without_state_integration() { let mut model = unit_integrator_model(); @@ -320,6 +398,94 @@ fn simulate_accepts_zero_state_solve_ir_without_building_ode_problem() { assert!(result.data.is_empty()); } +#[test] +fn simulate_zero_horizon_no_state_model_solves_nonlinear_algebraic_block_at_start() { + let mut model = solve::SolveModel::default(); + model.problem.solve_layout.algebraic_scalar_count = 1; + model.problem.solve_layout.solver_maps.names = vec!["x_zero".to_string()]; + model.problem.solve_layout.solver_maps.name_to_idx = + indexmap::IndexMap::from([("x_zero".to_string(), 0)]); + model.problem.solve_layout.solver_maps.base_to_indices = + indexmap::IndexMap::from([("x_zero".to_string(), vec![0])]); + model.problem.continuous.implicit_rhs = solve::ComputeBlock::from_scalar_program_block( + solve::ScalarProgramBlock::with_source_span( + vec![vec![ + solve::LinearOp::LoadY { dst: 0, index: 0 }, + solve::LinearOp::LoadY { dst: 1, index: 0 }, + solve::LinearOp::Binary { + dst: 2, + op: solve::BinaryOp::Mul, + lhs: 0, + rhs: 1, + }, + solve::LinearOp::Const { + dst: 3, + value: 0.25, + }, + solve::LinearOp::Binary { + dst: 4, + op: solve::BinaryOp::Sub, + lhs: 2, + rhs: 3, + }, + solve::LinearOp::StoreOutput { src: 4 }, + ]], + fixture_span!(), + ), + ); + model.artifacts.continuous.implicit_jacobian_v = solve::ComputeBlock::from_scalar_program_block( + solve::ScalarProgramBlock::with_source_span( + vec![vec![ + solve::LinearOp::Const { dst: 0, value: 2.0 }, + solve::LinearOp::LoadY { dst: 1, index: 0 }, + solve::LinearOp::Binary { + dst: 2, + op: solve::BinaryOp::Mul, + lhs: 0, + rhs: 1, + }, + solve::LinearOp::LoadSeed { dst: 3, index: 0 }, + solve::LinearOp::Binary { + dst: 4, + op: solve::BinaryOp::Mul, + lhs: 2, + rhs: 3, + }, + solve::LinearOp::StoreOutput { src: 4 }, + ]], + fixture_span!(), + ), + ); + model.problem.continuous.implicit_row_targets = vec![Some(solve::scalar_slot_y(0))]; + install_dense_algebraic_projection_plan(&mut model); + model.initial_y = vec![1.0]; + model.visible_names = vec!["x_zero".to_string()]; + model.visible_value_rows = solve::ScalarProgramBlock::with_source_span( + vec![vec![ + solve::LinearOp::LoadY { dst: 0, index: 0 }, + solve::LinearOp::StoreOutput { src: 0 }, + ]], + fixture_span!(), + ); + + let result = simulate( + &model, + &SimOptions { + t_start: 0.0, + t_end: 0.0, + ..Default::default() + }, + ) + .expect("zero-horizon no-state model should solve its algebraic block at t_start"); + + assert_eq!(result.times, vec![0.0]); + assert!( + (result.data[0][0] - 0.5).abs() <= 1.0e-7, + "expected nonlinear algebraic root 0.5, got {}", + result.data[0][0] + ); +} + #[test] fn simulate_records_visible_runtime_tail_values_for_no_state_models() { let mut model = solve::SolveModel::default(); @@ -862,6 +1028,71 @@ fn simulate_no_state_solve_ir_refreshes_periodic_event_indicator_between_ticks() ); } +#[test] +fn simulate_no_state_solve_ir_clears_lowered_change_pulse_and_alias_after_event() { + let mut model = solve::SolveModel::default(); + model.problem.solve_layout.parameter_count = 0; + model.problem.solve_layout.compiled_parameter_len = 5; + model.problem.solve_layout.discrete_valued_scalar_names = vec![ + "tick".to_string(), + "source".to_string(), + "changed".to_string(), + "trigger".to_string(), + ]; + model.problem.solve_layout.pre_param_bindings = vec![solve::PreParamBinding { + dest_p_index: 4, + source: solve::PreParamSource::P { index: 1 }, + }]; + model.problem.clocks.periodic_event_schedules = vec![solve::PeriodicEventSchedule { + phase_seconds: 0.5, + period_seconds: 0.5, + }]; + model.problem.discrete.update_targets = vec![ + solve::scalar_slot_p(0), + solve::scalar_slot_p(1), + solve::scalar_slot_p(2), + solve::scalar_slot_p(3), + ]; + model.problem.discrete.rhs = solve::ScalarProgramBlock::with_source_span( + vec![ + periodic_tick_row(0.5, 0.5), + increment_pre_parameter_on_tick_row(0, 4), + change_parameter_row(1, 4), + load_parameter_row(2), + ], + fixture_span!(), + ); + model.problem.discrete.pre_modes = vec![ + solve::DiscreteEventPreMode::FollowCurrent, + solve::DiscreteEventPreMode::EventEntry, + solve::DiscreteEventPreMode::EventEntry, + solve::DiscreteEventPreMode::FollowCurrent, + ]; + model.problem.discrete.observation_refresh = vec![true, false, true, true]; + model.parameters = vec![0.0; 5]; + model.visible_names = vec![ + "source".to_string(), + "changed".to_string(), + "trigger".to_string(), + ]; + + let result = simulate( + &model, + &SimOptions { + t_start: 0.0, + t_end: 0.75, + dt: Some(0.25), + ..Default::default() + }, + ) + .expect("lowered change pulse aliases should settle through the no-state schedule"); + + assert_eq!(result.times, vec![0.0, 0.25, 0.5, 0.75]); + assert_eq!(result.data[0], vec![0.0, 0.0, 1.0, 1.0]); + assert_eq!(result.data[1], vec![0.0; 4]); + assert_eq!(result.data[2], vec![0.0; 4]); +} + #[test] fn observation_refresh_keeps_pre_values_fixed_during_settle() { let mut model = solve::SolveModel::default(); @@ -942,8 +1173,7 @@ fn observation_refresh_resets_fixed_pre_event_history_rows() { ); } -#[test] -fn simulate_no_state_solve_ir_updates_clocked_previous_feedback_at_periodic_ticks() { +fn no_state_clocked_previous_feedback_model() -> solve::SolveModel { let mut model = solve::SolveModel::default(); model.problem.solve_layout.parameter_count = 0; model.problem.solve_layout.compiled_parameter_len = 8; @@ -1017,6 +1247,12 @@ fn simulate_no_state_solve_ir_updates_clocked_previous_feedback_at_periodic_tick "unitDelay.y".to_string(), "assignClock.u".to_string(), ]; + model +} + +#[test] +fn simulate_no_state_solve_ir_updates_clocked_previous_feedback_at_periodic_ticks() { + let model = no_state_clocked_previous_feedback_model(); let result = simulate( &model, @@ -1040,6 +1276,35 @@ fn simulate_no_state_solve_ir_updates_clocked_previous_feedback_at_periodic_tick ); } +#[test] +fn diffsol_no_state_session_rearms_periodic_sample_edges() { + let model = no_state_clocked_previous_feedback_model(); + let mut session = crate::session::SimulationSession::new( + &model, + SimOptions { + t_end: 0.05, + ..Default::default() + }, + ) + .expect("no-state periodic session should build"); + + session.advance_to(0.01).expect("advance before first tick"); + assert_eq!( + session.get("assignClock.y").expect("read output"), + Some(1.0) + ); + session.advance_to(0.02).expect("advance to first tick"); + assert_eq!( + session.get("assignClock.y").expect("read output"), + Some(2.0) + ); + session.advance_to(0.04).expect("advance to second tick"); + assert_eq!( + session.get("assignClock.y").expect("read output"), + Some(3.0) + ); +} + fn scalar_block(rows: Vec>) -> solve::ScalarProgramBlock { solve::ScalarProgramBlock::with_source_span(rows, fixture_span!()) } @@ -1084,9 +1349,22 @@ fn change_parameter_row(current_index: usize, pre_index: usize) -> Vec Vec { + increment_pre_parameter_on_tick_row(0, 2) +} + +fn increment_pre_parameter_on_tick_row( + tick_index: usize, + pre_index: usize, +) -> Vec { vec![ - solve::LinearOp::LoadP { dst: 0, index: 0 }, - solve::LinearOp::LoadP { dst: 1, index: 2 }, + solve::LinearOp::LoadP { + dst: 0, + index: tick_index, + }, + solve::LinearOp::LoadP { + dst: 1, + index: pre_index, + }, solve::LinearOp::Const { dst: 2, value: 1.0 }, solve::LinearOp::Binary { dst: 3, diff --git a/crates/rumoca-solver-diffsol/src/tests/root_events.rs b/crates/rumoca-solver-diffsol/src/tests/root_events.rs index 0eb9b444c..fe24808ee 100644 --- a/crates/rumoca-solver-diffsol/src/tests/root_events.rs +++ b/crates/rumoca-solver-diffsol/src/tests/root_events.rs @@ -33,6 +33,7 @@ fn root_reinit_does_not_interpolate_from_mutated_diffsol_state() { assert!(result.data[0].last().copied().unwrap() > 2.1); model.problem.events.root_conditions = solve::ScalarProgramBlock::default(); + model.problem.events.root_relation_memory_targets.clear(); let no_event = simulate( &model, &SimOptions { @@ -122,6 +123,7 @@ fn rising_state_with_root_reinit() -> solve::SolveModel { ]], fixture_span!(), ); + model.problem.events.root_relation_memory_targets = vec![Some(solve::scalar_slot_p(0))]; model.problem.discrete.update_targets = vec![solve::scalar_slot_y(0)]; model.problem.discrete.rhs = solve::ScalarProgramBlock::with_source_span( vec![vec![ @@ -147,7 +149,9 @@ fn rising_state_with_root_reinit() -> solve::SolveModel { ]], fixture_span!(), ); + model.problem.solve_layout.compiled_parameter_len = 1; model.initial_y = vec![0.0]; + model.parameters = vec![0.0]; model.visible_names = vec!["x".to_string()]; model } @@ -187,6 +191,7 @@ fn falling_ball_with_strict_reinit_guard() -> solve::SolveModel { ]], fixture_span!(), ); + model.problem.events.root_relation_memory_targets = vec![None]; model.problem.discrete.update_targets = vec![solve::scalar_slot_y(1)]; model.problem.discrete.pre_modes = vec![solve::DiscreteEventPreMode::Fixed]; model.problem.discrete.rhs = falling_ball_strict_reinit_rhs(); diff --git a/crates/rumoca-solver-diffsol/src/tests/runtime_value_tests.rs b/crates/rumoca-solver-diffsol/src/tests/runtime_value_tests.rs index 663148151..5ca9fb67a 100644 --- a/crates/rumoca-solver-diffsol/src/tests/runtime_value_tests.rs +++ b/crates/rumoca-solver-diffsol/src/tests/runtime_value_tests.rs @@ -64,6 +64,7 @@ fn simulate_no_state_solve_ir_stops_for_root_event_updates() { ]], fixture_span!(), ); + model.problem.events.root_relation_memory_targets = vec![None]; model.problem.discrete.rhs = solve::ScalarProgramBlock::with_source_span( vec![vec![ solve::LinearOp::LoadTime { dst: 0 }, @@ -166,6 +167,7 @@ fn no_state_root_search_refreshes_algebraic_root_dependencies() { ]], fixture_span!(), ); + model.problem.events.root_relation_memory_targets = vec![None]; model.problem.discrete.update_targets = vec![solve::scalar_slot_p(0)]; model.problem.discrete.rhs = solve::ScalarProgramBlock::with_source_span( vec![vec![ @@ -289,6 +291,7 @@ fn root_event_update_model(root_time: f64) -> solve::SolveModel { ]], fixture_span!(), ); + model.problem.events.root_relation_memory_targets = vec![None]; model.problem.discrete.rhs = solve::ScalarProgramBlock::with_source_span( vec![vec![ solve::LinearOp::LoadTime { dst: 0 }, @@ -531,10 +534,13 @@ fn project_algebraics_uses_solve_ir_row_targets_before_large_pivots() { } #[test] -fn project_algebraics_preserves_state_values_for_consistency_residuals() { +fn project_algebraics_rejects_nonzero_state_targeted_consistency_residuals() { let mut model = solve::SolveModel::default(); let rhs_rows = vec![ - vec![solve::LinearOp::Const { dst: 0, value: 0.0 }], + vec![ + solve::LinearOp::Const { dst: 0, value: 0.0 }, + solve::LinearOp::StoreOutput { src: 0 }, + ], vec![ solve::LinearOp::LoadY { dst: 0, index: 0 }, solve::LinearOp::Const { dst: 1, value: 2.0 }, @@ -548,7 +554,10 @@ fn project_algebraics_preserves_state_values_for_consistency_residuals() { ], ]; let jvp_rows = vec![ - vec![solve::LinearOp::Const { dst: 0, value: 0.0 }], + vec![ + solve::LinearOp::Const { dst: 0, value: 0.0 }, + solve::LinearOp::StoreOutput { src: 0 }, + ], vec![ solve::LinearOp::LoadSeed { dst: 0, index: 0 }, solve::LinearOp::StoreOutput { src: 0 }, @@ -569,15 +578,103 @@ fn project_algebraics_preserves_state_values_for_consistency_residuals() { let ode_model = OdeModel::new(&model).expect("ODE model should build from solve-IR rows"); let mut y = model.initial_y.clone(); - project_algebraics(&ode_model, &mut y, &[], 0.0, 1, 1.0e-12) - .expect("algebraic projection should ignore state-targeted consistency residuals"); + let err = project_algebraics(&ode_model, &mut y, &[], 0.0, 1, 1.0e-12) + .expect_err("a nonzero consistency residual in the algebraic tail must not be ignored"); // State initialization belongs to the initialization problem and fixed - // starts, not to the algebraic projector used at event boundaries. + // starts, not to the algebraic projector used at event boundaries. The + // projector therefore reports the inconsistent row without moving state. + assert!( + err.to_string() + .contains("algebraic projection plan omits implicit residual row 1") + ); assert_eq!(y[0], 0.0); assert_eq!(y[1], 0.0); } +#[test] +fn settle_algebraics_uses_full_ode_projection_semantics() { + let mut model = solve::SolveModel::default(); + model.problem.solve_layout.state_scalar_count = 1; + model.problem.solve_layout.algebraic_scalar_count = 2; + model.problem.continuous.implicit_rhs = solve::ComputeBlock::from_scalar_program_block( + solve::ScalarProgramBlock::with_source_span( + vec![ + vec![ + solve::LinearOp::Const { dst: 0, value: 0.0 }, + solve::LinearOp::StoreOutput { src: 0 }, + ], + vec![ + solve::LinearOp::LoadY { dst: 0, index: 1 }, + solve::LinearOp::Const { dst: 1, value: 2.0 }, + solve::LinearOp::Binary { + dst: 2, + op: solve::BinaryOp::Sub, + lhs: 0, + rhs: 1, + }, + solve::LinearOp::StoreOutput { src: 2 }, + ], + vec![ + solve::LinearOp::LoadY { dst: 0, index: 2 }, + solve::LinearOp::Const { dst: 1, value: 3.0 }, + solve::LinearOp::Binary { + dst: 2, + op: solve::BinaryOp::Sub, + lhs: 0, + rhs: 1, + }, + solve::LinearOp::StoreOutput { src: 2 }, + ], + ], + fixture_span!(), + ), + ); + model.artifacts.continuous.implicit_jacobian_v = solve::ComputeBlock::from_scalar_program_block( + solve::ScalarProgramBlock::with_source_span( + vec![ + vec![ + solve::LinearOp::Const { dst: 0, value: 0.0 }, + solve::LinearOp::StoreOutput { src: 0 }, + ], + vec![ + solve::LinearOp::LoadSeed { dst: 0, index: 1 }, + solve::LinearOp::StoreOutput { src: 0 }, + ], + vec![ + solve::LinearOp::LoadSeed { dst: 0, index: 2 }, + solve::LinearOp::StoreOutput { src: 0 }, + ], + ], + fixture_span!(), + ), + ); + model.problem.continuous.implicit_row_targets = vec![None; 3]; + model.problem.continuous.algebraic_projection_plan = solve::AlgebraicProjectionPlan { + blocks: vec![solve::AlgebraicProjectionBlock { + rows: vec![1], + y_indices: vec![1], + causal_steps: Vec::new(), + }], + }; + model.initial_y = vec![0.0, 0.0, 0.0]; + + let runtime = SolveRuntime::new(&model).expect("runtime should prepare the incomplete plan"); + let ode_model = OdeModel::new(&model).expect("ODE model should prepare full residual rows"); + let mut y = model.initial_y.clone(); + let mut p = model.parameters.clone(); + + let err = settle_algebraics_and_relation_memory( + &runtime, &ode_model, &mut y, &mut p, 0.0, 1, 1.0e-12, + ) + .expect_err("settling must not bypass the OdeModel semantic projector"); + + assert!( + err.to_string() + .contains("algebraic projection plan omits implicit residual row 2") + ); +} + #[test] fn simulate_seeds_algebraics_from_initial_residual_before_runtime_projection() { let mut model = solve::SolveModel::default(); diff --git a/crates/rumoca-solver-rk45/src/lib.rs b/crates/rumoca-solver-rk45/src/lib.rs index 12b115985..cc8ee9d4a 100644 --- a/crates/rumoca-solver-rk45/src/lib.rs +++ b/crates/rumoca-solver-rk45/src/lib.rs @@ -15,9 +15,10 @@ use rumoca_solver::{ RuntimeEventBoundaryHandler, RuntimeEventStop, RuntimeSolveError, SimOptions, SimResult, SimSolverMode, SimTermination, SimulationBackend, SolveStopSchedule, StepUntilOutcome, TimeoutBudget, TimeoutExceeded, clear_scheduled_root_relation_memory, - commit_pre_params_after_event, convert_variable_meta, filter_scheduled_root_crossings, - process_runtime_event_boundary, root_crossings_with_relation_memory, root_value_crossed, - runtime_event_horizon, timeline, + clear_scheduled_root_relation_memory_at_time, commit_pre_params_after_event, + convert_variable_meta, filter_scheduled_root_crossings, process_runtime_event_boundary, + root_crossings_with_relation_memory, root_value_crossed, runtime_event_horizon, + scheduled_root_relation_overrides_at_time, timeline, }; mod no_state; @@ -571,6 +572,18 @@ fn root_value_at(values: &[f64], index: usize, label: &str) -> Result bool { + let mut reads_time = false; + for op in row { + match op { + solve::LinearOp::LoadTime { .. } => reads_time = true, + solve::LinearOp::LoadY { .. } => return false, + _ => {} + } + } + reads_time +} + impl<'a> Rk45Backend<'a> { fn new(model: &'a SolveRuntime, opts: &SimOptions) -> Result { let state = model.model.initial_y[..model.state_count].to_vec(); @@ -953,7 +966,7 @@ impl<'a> Rk45Backend<'a> { )?; let old = root_value_at(&lo_roots, crossing.index, "left bisection")?; let new = root_value_at(&mid_roots, crossing.index, "midpoint bisection")?; - if root_value_crossed(old, new, self.atol) { + if root_value_crossed(old, new, self.root_bisection_tol(crossing.index)) { hi_t = mid_t; hi_state = mid_state; } else { @@ -969,6 +982,25 @@ impl<'a> Rk45Backend<'a> { }) } + fn root_bisection_tol(&self, root_index: usize) -> f64 { + let Some(row) = self + .model + .model + .problem + .events + .root_conditions + .programs + .get(root_index) + else { + return self.atol; + }; + if root_condition_is_time_only(row) { + 0.0 + } else { + self.atol + } + } + fn locate_simultaneous_crossings( &self, old_t: f64, @@ -1335,8 +1367,12 @@ impl Rk45Backend<'_> { if event.observe_right_limit || !matches!(event.pre_mode, EventPreMode::EventEntry) { return Ok(()); } - let root_indices = self.scheduled_root_indices_at_time(event_time); - self.clear_scheduled_root_relation_memory(&root_indices) + clear_scheduled_root_relation_memory_at_time( + &self.model.model, + event_time, + &mut self.params, + ) + .map_err(runtime_contract_violation) } fn clear_all_scheduled_root_relation_memory(&mut self) -> Result<(), SimError> { @@ -1368,25 +1404,18 @@ impl Rk45Backend<'_> { if event.observe_right_limit || !matches!(event.pre_mode, EventPreMode::EventEntry) { return Ok(()); } - let scheduled_indices = self.scheduled_root_indices_at_time(event_time); - if scheduled_indices.is_empty() { + let overrides = scheduled_root_relation_overrides_at_time(&self.model.model, event_time); + if overrides.is_empty() { return Ok(()); } - for index in scheduled_indices { + for (index, post_relation_memory_value) in overrides { self.pending_root_crossings.push(RootCrossing { index, - post_relation_memory_value: 1.0, + post_relation_memory_value, }); } Ok(()) } - - fn scheduled_root_indices_at_time(&self, event_time: f64) -> Vec { - timeline::scheduled_root_indices_at_time( - &self.model.model.problem.events.scheduled_root_conditions, - event_time, - ) - } } fn advance_backend_to(backend: &mut Rk45Backend<'_>, target_t: f64) -> Result<(), SimError> { diff --git a/crates/rumoca-solver-rk45/src/no_state.rs b/crates/rumoca-solver-rk45/src/no_state.rs index 3e96eb306..16255d97e 100644 --- a/crates/rumoca-solver-rk45/src/no_state.rs +++ b/crates/rumoca-solver-rk45/src/no_state.rs @@ -8,9 +8,10 @@ use rumoca_solver::{ EventActionOutcome, EventPreMode, NoStateEventStep, NoStateOrchestrationBackend, NoStateScheduledStop, RuntimeEventBoundary, RuntimeEventBoundaryHandler, RuntimeEventStop, RuntimeSolveError, SimOptions, SimTermination, SolveStopSchedule, - commit_pre_params_after_event, process_runtime_event_boundary, run_no_state_output_schedule, - runtime_event_horizon, runtime_root_event_application_time, - timeline::{event_left_limit_time, sample_time_match_with_tol}, + clear_scheduled_root_relation_memory_at_time, commit_pre_params_after_event, + process_runtime_event_boundary, run_no_state_output_schedule, runtime_event_horizon, + runtime_root_event_application_time, scheduled_root_relation_overrides_at_time, + timeline::{event_left_limit_time, sample_time_match_with_tol, scheduled_root_index_is_known}, }; use crate::{SessionState, SimError}; @@ -224,7 +225,13 @@ impl NoStateOrchestrationBackend for Rk45NoStateOrchestration<'_> { &mut self.runtime.params, self.runtime.current_t, self.opts.atol.max(1.0e-10), + )?; + clear_scheduled_root_relation_memory_at_time( + self.model, + self.runtime.current_t, + &mut self.runtime.params, ) + .map_err(SimError::SolveIr) } } @@ -346,24 +353,37 @@ impl RuntimeEventBoundaryHandler for NoStateEventBoundary<'_> { type Error = SimError; fn on_event_time(&mut self, event_t: f64, event: RuntimeEventStop) -> Result<(), Self::Error> { + let root_relation_overrides = if self.root_event + || event.observe_right_limit + || !matches!(event.pre_mode, EventPreMode::EventEntry) + { + Vec::new() + } else { + scheduled_root_relation_overrides_at_time(&self.runtime.model, event_t) + }; if !self.root_event && matches!( event.pre_mode, EventPreMode::EventEntry | EventPreMode::Fixed ) { - let left_t = event_left_limit_time(event_t); - refresh_observation_rows_and_relation_memory( - self.runtime, - self.y, - self.p, - left_t, - self.tol, - )?; + if root_relation_overrides.is_empty() { + let left_t = event_left_limit_time(event_t); + refresh_observation_rows_and_relation_memory( + self.runtime, + self.y, + self.p, + left_t, + self.tol, + )?; + } else { + clear_scheduled_root_relation_memory_at_time(&self.runtime.model, event_t, self.p) + .map_err(SimError::SolveIr)?; + } } self.event_pre_y = self.y.to_vec(); self.event_pre_p = self.p.to_vec(); - self.apply_event_updates(event_t)?; + self.apply_event_updates(event_t, &root_relation_overrides)?; refresh_observation_rows_and_relation_memory( self.runtime, self.y, @@ -378,7 +398,7 @@ impl RuntimeEventBoundaryHandler for NoStateEventBoundary<'_> { right_t: f64, _event: RuntimeEventStop, ) -> Result<(), Self::Error> { - self.apply_event_updates(right_t)?; + self.apply_event_updates(right_t, &[])?; refresh_observation_rows_and_relation_memory( self.runtime, self.y, @@ -390,7 +410,11 @@ impl RuntimeEventBoundaryHandler for NoStateEventBoundary<'_> { } impl NoStateEventBoundary<'_> { - fn apply_event_updates(&mut self, t: f64) -> Result<(), SimError> { + fn apply_event_updates( + &mut self, + t: f64, + root_relation_overrides: &[(usize, f64)], + ) -> Result<(), SimError> { let outcome = self.runtime.apply_projected_event_update( ProjectedEventUpdateInput { y: self.y, @@ -401,7 +425,7 @@ impl NoStateEventBoundary<'_> { event_pre_p: &self.event_pre_p, max_iters: NO_STATE_EVENT_UPDATE_MAX_ITERS, row_filter: EventUpdateRowFilter::All, - root_relation_overrides: &[], + root_relation_overrides, }, |y, p| refresh_algebraics_and_detect_changes(self.runtime, y, p, t, self.tol), )?; @@ -551,9 +575,16 @@ fn first_root_crossing_time( eval_refreshed_roots(runtime, y, p, t_end, tol, &mut end)?; let mut crossing = None; - for (a, b) in start.iter().zip(end.iter()) { + for (root_index, (a, b)) in start.iter().zip(end.iter()).enumerate() { + if scheduled_root_index_is_known( + &runtime.model.problem.events.scheduled_root_conditions, + root_index, + ) { + continue; + } if root_surface_crossed_or_near(*a, *b, tol) { - let root = bisect_first_root(runtime, y, p, t_start, t_end, tol, root_count)?; + let root = + bisect_first_root(runtime, y, p, (t_start, t_end), tol, root_count, root_index)?; crossing = Some(crossing.map_or(root, |current: f64| current.min(root))); } } @@ -576,10 +607,10 @@ fn bisect_first_root( runtime: &SolveRuntime, y: &[f64], p: &[f64], - mut lo: f64, - mut hi: f64, + (mut lo, mut hi): (f64, f64), tol: f64, root_count: usize, + root_index: usize, ) -> Result { let mut lo_roots = vec![0.0; root_count]; eval_refreshed_roots(runtime, y, p, lo, tol, &mut lo_roots)?; @@ -587,11 +618,9 @@ fn bisect_first_root( let mid = lo + 0.5 * (hi - lo); let mut mid_roots = vec![0.0; root_count]; eval_refreshed_roots(runtime, y, p, mid, tol, &mut mid_roots)?; - if lo_roots - .iter() - .zip(mid_roots.iter()) - .any(|(a, b)| a.signum() != b.signum() || root_surface_near_zero(*b, tol)) - { + let lo_root = lo_roots.get(root_index).copied().unwrap_or(0.0); + let mid_root = mid_roots.get(root_index).copied().unwrap_or(0.0); + if lo_root.signum() != mid_root.signum() || root_surface_near_zero(mid_root, tol) { hi = mid; } else { lo = mid; diff --git a/crates/rumoca-solver-rk45/src/tests.rs b/crates/rumoca-solver-rk45/src/tests.rs index 434021265..5b9847558 100644 --- a/crates/rumoca-solver-rk45/src/tests.rs +++ b/crates/rumoca-solver-rk45/src/tests.rs @@ -409,6 +409,58 @@ fn rk45_root_event_updates_relation_memory_for_continuous_if_branch() { ); } +#[test] +fn rk45_root_bisection_does_not_stop_at_state_atol_band() { + let mut model = single_state_model(vec![vec![ + LinearOp::Const { dst: 0, value: 0.0 }, + LinearOp::StoreOutput { src: 0 }, + ]]); + model.problem.events.root_conditions = ScalarProgramBlock::with_source_span( + vec![vec![ + LinearOp::Const { dst: 0, value: 0.5 }, + LinearOp::LoadTime { dst: 1 }, + LinearOp::Binary { + dst: 2, + op: solve::BinaryOp::Sub, + lhs: 0, + rhs: 1, + }, + LinearOp::StoreOutput { src: 2 }, + ]], + fixture_span!(), + ); + let runtime = SolveRuntime::new(&model).expect("test model should build runtime"); + let backend = Rk45Backend::new( + &runtime, + &SimOptions { + solver_mode: SimSolverMode::RkLike, + t_end: 0.501, + atol: 1.0e-6, + ..Default::default() + }, + ) + .expect("rk45 backend should initialize"); + + let root = backend + .bisect_root( + 0.5, + vec![0.0], + 0.501, + RootCrossing { + index: 0, + post_relation_memory_value: 0.0, + }, + None, + ) + .expect("time root should locate"); + + assert!( + (root.time - 0.5).abs() <= 1.0e-12, + "root location should not stop at the solver state tolerance band; t={}", + root.time + ); +} + #[test] fn rk45_applies_periodic_event_update() { let mut model = single_state_model(vec![vec![ @@ -638,6 +690,88 @@ fn rk45_clears_scheduled_sample_relation_memory_between_ticks() { assert_eq!(result.data[3], vec![1.0, 2.0, 3.0, 3.0]); } +fn periodic_tick_row(phase: f64, period: f64) -> Vec { + vec![ + LinearOp::LoadTime { dst: 0 }, + LinearOp::Const { + dst: 1, + value: phase, + }, + LinearOp::Binary { + dst: 2, + op: solve::BinaryOp::Sub, + lhs: 0, + rhs: 1, + }, + LinearOp::Const { + dst: 3, + value: -1.0e-9, + }, + LinearOp::Compare { + dst: 4, + op: solve::CompareOp::Ge, + lhs: 2, + rhs: 3, + }, + LinearOp::Const { + dst: 5, + value: period, + }, + LinearOp::Binary { + dst: 6, + op: solve::BinaryOp::Div, + lhs: 2, + rhs: 5, + }, + LinearOp::Const { dst: 7, value: 0.5 }, + LinearOp::Binary { + dst: 8, + op: solve::BinaryOp::Add, + lhs: 6, + rhs: 7, + }, + LinearOp::Unary { + dst: 9, + op: solve::UnaryOp::Floor, + arg: 8, + }, + LinearOp::Binary { + dst: 10, + op: solve::BinaryOp::Mul, + lhs: 9, + rhs: 5, + }, + LinearOp::Binary { + dst: 11, + op: solve::BinaryOp::Sub, + lhs: 2, + rhs: 10, + }, + LinearOp::Unary { + dst: 12, + op: solve::UnaryOp::Abs, + arg: 11, + }, + LinearOp::Const { + dst: 13, + value: 1.0e-9, + }, + LinearOp::Compare { + dst: 14, + op: solve::CompareOp::Le, + lhs: 12, + rhs: 13, + }, + LinearOp::Binary { + dst: 15, + op: solve::BinaryOp::And, + lhs: 4, + rhs: 14, + }, + LinearOp::StoreOutput { src: 15 }, + ] +} + #[test] fn rk45_applies_dynamic_time_event_update() { let mut model = single_state_model(vec![vec![ @@ -800,6 +934,27 @@ fn rk45_session_runs_no_state_discrete_controller() { assert_eq!(session.get("y").expect("read y"), Some(5.5)); } +#[test] +fn rk45_no_state_session_rearms_periodic_sample_edges() { + let model = no_state_periodic_sample_counter_model(); + let mut session = SimulationSession::new( + &model, + SimOptions { + t_end: 0.05, + solver_mode: SimSolverMode::RkLike, + ..Default::default() + }, + ) + .expect("no-state periodic session should build"); + + session.advance_to(0.01).expect("advance before first tick"); + assert_eq!(session.get("count").expect("read count"), Some(0.0)); + session.advance_to(0.02).expect("advance to first tick"); + assert_eq!(session.get("count").expect("read count"), Some(1.0)); + session.advance_to(0.04).expect("advance to second tick"); + assert_eq!(session.get("count").expect("read count"), Some(2.0)); +} + #[test] fn rk45_session_uses_adaptive_event_integration_for_stiff_contact() { let model = stiff_contact_model(); @@ -1076,6 +1231,9 @@ fn single_state_model(rhs_rows: Vec>) -> solve::SolveModel { initialization: solve::InitializationSolveSystem { residual: ComputeBlock::from_scalar_program_block(zero.clone()), row_targets: Vec::new(), + direct_families: Vec::new(), + required_target_ranges: Vec::new(), + fixed_target_ranges: Vec::new(), projection_indices: Vec::new(), projection_plan: solve::AlgebraicProjectionPlan::default(), update_rhs: solve::ScalarProgramBlock::default(), @@ -1217,3 +1375,87 @@ fn no_state_input_accumulator_model() -> solve::SolveModel { variable_meta: Vec::new(), } } + +fn no_state_periodic_sample_counter_model() -> solve::SolveModel { + let mut model = solve::SolveModel::default(); + model.problem.solve_layout.compiled_parameter_len = 4; + model.problem.solve_layout.discrete_valued_scalar_names = vec![ + "__pre__.sample".to_string(), + "sample".to_string(), + "count".to_string(), + "__pre__.count".to_string(), + ]; + model.problem.solve_layout.pre_param_bindings = vec![ + solve::PreParamBinding { + dest_p_index: 0, + source: solve::PreParamSource::P { index: 1 }, + }, + solve::PreParamBinding { + dest_p_index: 3, + source: solve::PreParamSource::P { index: 2 }, + }, + ]; + model.problem.clocks.periodic_event_schedules = vec![solve::PeriodicEventSchedule { + period_seconds: 0.02, + phase_seconds: 0.02, + }]; + model.problem.events.root_conditions = const_scalar_program_block(1.0); + model.problem.events.root_relation_memory_targets = vec![Some(solve::scalar_slot_p(1))]; + model.problem.events.scheduled_root_conditions = vec![solve::ScheduledRootCondition { + root_index: 0, + period_seconds: 0.02, + phase_seconds: 0.02, + }]; + model.problem.discrete.update_targets = vec![solve::scalar_slot_p(1), solve::scalar_slot_p(2)]; + model.problem.discrete.pre_modes = vec![ + solve::DiscreteEventPreMode::Fixed, + solve::DiscreteEventPreMode::Fixed, + ]; + model.problem.discrete.observation_refresh = vec![true, false]; + model.problem.discrete.rhs = ScalarProgramBlock::with_source_span( + vec![ + periodic_tick_row(0.02, 0.02), + vec![ + LinearOp::LoadP { dst: 0, index: 1 }, + LinearOp::LoadP { dst: 1, index: 0 }, + LinearOp::Unary { + dst: 2, + op: solve::UnaryOp::Not, + arg: 1, + }, + LinearOp::Binary { + dst: 3, + op: solve::BinaryOp::And, + lhs: 0, + rhs: 2, + }, + LinearOp::LoadP { dst: 4, index: 3 }, + LinearOp::Const { dst: 5, value: 1.0 }, + LinearOp::Binary { + dst: 6, + op: solve::BinaryOp::Add, + lhs: 4, + rhs: 5, + }, + LinearOp::Select { + dst: 7, + cond: 3, + if_true: 6, + if_false: 4, + }, + LinearOp::StoreOutput { src: 7 }, + ], + ], + fixture_span!(), + ); + model.parameters = vec![0.0; 4]; + model.visible_names = vec!["count".to_string()]; + model.visible_value_rows = ScalarProgramBlock::with_source_span( + vec![vec![ + LinearOp::LoadP { dst: 0, index: 2 }, + LinearOp::StoreOutput { src: 0 }, + ]], + fixture_span!(), + ); + model +} diff --git a/crates/rumoca-solver/src/lib.rs b/crates/rumoca-solver/src/lib.rs index 15fe979ca..77bee7b83 100644 --- a/crates/rumoca-solver/src/lib.rs +++ b/crates/rumoca-solver/src/lib.rs @@ -22,7 +22,8 @@ pub use runtime::no_state::{ }; pub use runtime::orchestration::{LoopStats, run_with_runtime_schedule}; pub use runtime::pre_params::{ - clear_scheduled_root_relation_memory, commit_pre_params_after_event, update_slot, + clear_scheduled_root_relation_memory, clear_scheduled_root_relation_memory_at_time, + commit_pre_params_after_event, scheduled_root_relation_overrides_at_time, update_slot, write_pre_params_from_sources, }; pub use runtime::projection::{ diff --git a/crates/rumoca-solver/src/runtime/event.rs b/crates/rumoca-solver/src/runtime/event.rs index 8733fb34e..1257d0aaf 100644 --- a/crates/rumoca-solver/src/runtime/event.rs +++ b/crates/rumoca-solver/src/runtime/event.rs @@ -84,10 +84,7 @@ fn runtime_event_final_time(event: RuntimeEventStop, event_t: f64, right_t: f64) } fn should_process_right_limit(event: RuntimeEventStop, event_t: f64, right_t: f64) -> bool { - event.observe_right_limit - && event.pre_mode == EventPreMode::FollowCurrent - && right_t > event_t - && !sample_time_match_with_tol(right_t, event_t) + event.observe_right_limit && event.pre_mode == EventPreMode::FollowCurrent && right_t > event_t } #[cfg(test)] diff --git a/crates/rumoca-solver/src/runtime/pre_params.rs b/crates/rumoca-solver/src/runtime/pre_params.rs index 8dfcf83f7..24782e97b 100644 --- a/crates/rumoca-solver/src/runtime/pre_params.rs +++ b/crates/rumoca-solver/src/runtime/pre_params.rs @@ -1,5 +1,7 @@ use rumoca_ir_solve as solve; +use crate::timeline::scheduled_root_indices_at_time; + pub fn write_pre_params_from_sources( model: &solve::SolveModel, source_y: &[f64], @@ -64,6 +66,26 @@ pub fn clear_scheduled_root_relation_memory( Ok(()) } +pub fn scheduled_root_relation_overrides_at_time( + model: &solve::SolveModel, + event_t: f64, +) -> Vec<(usize, f64)> { + scheduled_root_indices_at_time(&model.problem.events.scheduled_root_conditions, event_t) + .into_iter() + .map(|root_index| (root_index, 1.0)) + .collect() +} + +pub fn clear_scheduled_root_relation_memory_at_time( + model: &solve::SolveModel, + event_t: f64, + params: &mut [f64], +) -> Result<(), String> { + let root_indices = + scheduled_root_indices_at_time(&model.problem.events.scheduled_root_conditions, event_t); + clear_scheduled_root_relation_memory(model, &root_indices, params) +} + fn clear_pre_params_from_source_p( model: &solve::SolveModel, params: &mut [f64], diff --git a/crates/rumoca-solver/src/runtime/projection.rs b/crates/rumoca-solver/src/runtime/projection.rs index 985805acb..84b88be17 100644 --- a/crates/rumoca-solver/src/runtime/projection.rs +++ b/crates/rumoca-solver/src/runtime/projection.rs @@ -137,7 +137,7 @@ fn project_algebraics_with_plan( for _ in 0..ALGEBRAIC_PROJECTION_MAX_ITERS { seed_nonfinite_algebraics(y, state_count); model.eval_residual(y, p, t, &mut rhs)?; - let residual = projection_residual_tail(&rhs, plan, state_count)?; + let residual = projection_residual_tail(&rhs, plan, state_count, tol)?; if residual_converged(&residual, tol) { return Ok(()); } @@ -168,25 +168,73 @@ fn projection_residual_tail( rhs: &[f64], plan: &solve::AlgebraicProjectionPlan, state_count: usize, + tol: f64, ) -> Result, RuntimeSolveError> { - let mut residual = - vec![0.0; algebraic_tail_len(rhs.len(), state_count, "projection residual tail")?]; - for row in plan.blocks.iter().flat_map(|block| { - block - .rows - .iter() - .copied() - .chain(block.causal_steps.iter().map(|step| step.row)) - }) { - if row < state_count { - continue; + let covered_rows = validate_algebraic_projection_plan(plan, rhs.len(), state_count)?; + for (offset, value) in rhs[state_count..].iter().copied().enumerate() { + if !covered_rows[offset] && (!value.is_finite() || value.abs() > tol) { + return Err(RuntimeSolveError::solve_ir(format!( + "algebraic projection plan omits implicit residual row {}", + state_count + offset + ))); + } + } + Ok(rhs[state_count..].to_vec()) +} + +fn validate_algebraic_projection_plan( + plan: &solve::AlgebraicProjectionPlan, + residual_len: usize, + state_count: usize, +) -> Result, RuntimeSolveError> { + let algebraic_count = algebraic_tail_len( + residual_len, + state_count, + "validate algebraic projection plan", + )?; + let mut covered_rows = vec![false; algebraic_count]; + for (block_index, block) in plan.blocks.iter().enumerate() { + if block.rows.is_empty() && (!block.y_indices.is_empty() || !block.causal_steps.is_empty()) + { + return Err(RuntimeSolveError::solve_ir(format!( + "algebraic projection block {block_index} has targets but no residual rows" + ))); + } + if !block.rows.is_empty() && block.y_indices.is_empty() { + return Err(RuntimeSolveError::solve_ir(format!( + "algebraic projection block {block_index} has residual rows but no algebraic targets" + ))); + } + for &y_index in &block.y_indices { + if y_index < state_count || y_index >= residual_len { + return Err(RuntimeSolveError::solve_ir(format!( + "algebraic projection block {block_index} references non-algebraic y index {y_index} outside {state_count}..{residual_len}" + ))); + } } - let value = residual_at(rhs, row, "projection residual tail")?; - if let Some(slot) = residual.get_mut(row - state_count) { - *slot = value; + for &row in &block.rows { + let Some(offset) = row.checked_sub(state_count) else { + return Err(RuntimeSolveError::solve_ir(format!( + "algebraic projection block {block_index} references state residual row {row}" + ))); + }; + let Some(covered) = covered_rows.get_mut(offset) else { + return Err(RuntimeSolveError::solve_ir(format!( + "algebraic projection block {block_index} references residual row {row}, but the model evaluated only {residual_len} rows" + ))); + }; + *covered = true; + } + for step in &block.causal_steps { + if !block.rows.contains(&step.row) || !block.y_indices.contains(&step.y_index) { + return Err(RuntimeSolveError::solve_ir(format!( + "algebraic projection causal step ({}, y[{}]) is not contained in block {block_index}", + step.row, step.y_index + ))); + } } } - Ok(residual) + Ok(covered_rows) } fn algebraic_tail_len( @@ -327,6 +375,13 @@ fn project_causal_step( y.len() ))); } + if let Some(target_value) = model.eval_implicit_target_value(step.row, step.y_index, y, p, t)? { + if !target_value.is_finite() || y[step.y_index] == target_value { + return Ok(false); + } + y[step.y_index] = target_value; + return Ok(true); + } let mut rhs = vec![0.0; y.len()]; let mut seed = vec![0.0; y.len()]; let mut jv = vec![0.0; y.len()]; @@ -490,11 +545,12 @@ fn projection_error( .max_by(|(_, lhs), (_, rhs)| residual_sort_key(*lhs).total_cmp(&residual_sort_key(*rhs))); match worst { Some((row, value)) => { + let absolute_row = state_count + row; let target = model - .target_name_for_row(state_count + row) + .target_name_for_row(absolute_row) .map_or(String::new(), |name| format!(" target={name}")); RuntimeSolveError::solve_ir(format!( - "{message}: max residual row={row}{target} value={value:.6e} norm={:.6e}", + "{message}: max residual row={absolute_row}{target} value={value:.6e} norm={:.6e}", residual_norm(residual) )) } @@ -1157,6 +1213,81 @@ mod tests { } } + struct DirectAssignmentProjectionModel { + plan: solve::AlgebraicProjectionPlan, + assignment_value: f64, + } + + impl AlgebraicProjectionModel for DirectAssignmentProjectionModel { + fn eval_residual( + &self, + y: &[f64], + _p: &[f64], + _t: f64, + out: &mut [f64], + ) -> Result<(), RuntimeSolveError> { + out[0] = y[0] - 2.0; + Ok(()) + } + + fn eval_initial_residual( + &self, + y: &[f64], + p: &[f64], + t: f64, + out: &mut [f64], + ) -> Result<(), RuntimeSolveError> { + self.eval_residual(y, p, t, out) + } + + fn eval_jacobian_v( + &self, + _y: &[f64], + _p: &[f64], + _t: f64, + v: &[f64], + out: &mut [f64], + ) -> Result<(), RuntimeSolveError> { + out[0] = v[0]; + Ok(()) + } + + fn initial_residual_len(&self) -> usize { + 0 + } + + fn implicit_target(&self, row_idx: usize) -> Option { + (row_idx == 0).then(|| solve::scalar_slot_y(0)) + } + + fn initial_target(&self, _row_idx: usize) -> Option { + None + } + + fn algebraic_projection_plan(&self) -> &solve::AlgebraicProjectionPlan { + &self.plan + } + + fn has_explicit_initial_targets(&self) -> bool { + false + } + + fn target_name_for_row(&self, _row_idx: usize) -> Option<&str> { + None + } + + fn eval_implicit_target_value( + &self, + row_idx: usize, + target_y_index: usize, + _y: &[f64], + _p: &[f64], + t: f64, + ) -> Result, RuntimeSolveError> { + Ok((row_idx == 0 && target_y_index == 0).then_some(self.assignment_value + t)) + } + } + struct RectInitialProjectionModel; impl AlgebraicProjectionModel for RectInitialProjectionModel { @@ -1383,6 +1514,55 @@ mod tests { assert_eq!(y, vec![2.0, 3.0]); } + #[test] + fn project_algebraics_solves_targeted_loop_simultaneously_without_causal_steps() { + let model = BlockProjectionModel { + plan: solve::AlgebraicProjectionPlan { + blocks: vec![solve::AlgebraicProjectionBlock { + rows: vec![0, 1], + y_indices: vec![0, 1], + causal_steps: Vec::new(), + }], + }, + initial_residual_len: 0, + }; + let mut y = vec![0.0, 0.0]; + + project_algebraics(&model, &mut y, &[], 0.0, 0, 1.0e-12) + .expect("targeted algebraic loop should converge through the simultaneous solve"); + + let mut residual = vec![f64::NAN; 2]; + model + .eval_residual(&y, &[], 0.0, &mut residual) + .expect("residual evaluation should succeed"); + assert_eq!(y, vec![2.0, 3.0]); + assert_eq!(residual, vec![0.0, 0.0]); + } + + #[test] + fn project_algebraics_rejects_nonzero_residual_row_omitted_from_plan() { + let model = BlockProjectionModel { + plan: solve::AlgebraicProjectionPlan { + blocks: vec![solve::AlgebraicProjectionBlock { + rows: vec![0], + y_indices: vec![0], + causal_steps: Vec::new(), + }], + }, + initial_residual_len: 0, + }; + let mut y = vec![0.0, 0.0]; + + let err = project_algebraics(&model, &mut y, &[], 0.0, 0, 1.0e-12) + .expect_err("an omitted nonzero algebraic residual row must reject projection"); + + assert!(matches!(err, RuntimeSolveError::SolveIr { .. })); + assert!( + err.to_string() + .contains("algebraic projection plan omits implicit residual row 1") + ); + } + #[test] fn project_algebraics_rejects_state_count_past_y_length() { let model = BlockProjectionModel { @@ -1402,8 +1582,13 @@ mod tests { #[test] fn projection_residual_tail_rejects_state_count_past_rhs_length() { - let err = projection_residual_tail(&[0.0], &solve::AlgebraicProjectionPlan::default(), 2) - .expect_err("state count beyond residual length should fail"); + let err = projection_residual_tail( + &[0.0], + &solve::AlgebraicProjectionPlan::default(), + 2, + 1.0e-12, + ) + .expect_err("state count beyond residual length should fail"); assert!( err.to_string() @@ -1411,6 +1596,18 @@ mod tests { ); } + #[test] + fn projection_error_reports_absolute_implicit_row_index() { + let model = BlockProjectionModel { + plan: solve::AlgebraicProjectionPlan::default(), + initial_residual_len: 0, + }; + + let error = projection_error(&model, 3, "projection failed", &[0.25, 0.5]); + + assert!(error.to_string().contains("max residual row=4")); + } + #[test] fn project_algebraic_block_uses_svd_for_rectangular_jacobian() { let model = BlockProjectionModel { @@ -1472,6 +1669,104 @@ mod tests { ); } + #[test] + fn project_causal_step_applies_direct_assignment_without_target_incidence() { + let model = DirectAssignmentProjectionModel { + plan: solve::AlgebraicProjectionPlan::default(), + assignment_value: 2.0, + }; + let mut y = vec![0.0]; + let step = solve::AlgebraicProjectionStep { row: 0, y_index: 0 }; + + let changed = project_causal_step(&model, &mut y, &[], 0.0, &step, 1.0e-12) + .expect("direct assignment projection should succeed"); + + assert!(changed); + assert_eq!(y, vec![2.0]); + } + + #[test] + fn project_causal_step_rejects_mismatched_direct_assignment_target() { + let model = DirectAssignmentProjectionModel { + plan: solve::AlgebraicProjectionPlan::default(), + assignment_value: 2.0, + }; + let mut y = vec![0.0, 0.0]; + let step = solve::AlgebraicProjectionStep { row: 0, y_index: 1 }; + + let changed = project_causal_step(&model, &mut y, &[], 0.0, &step, 1.0e-12) + .expect("mismatched target should fall back to residual projection"); + + assert!(!changed); + assert_eq!(y, vec![0.0, 0.0]); + } + + #[test] + fn project_causal_step_preserves_sub_tolerance_time_assignment_change() { + let model = DirectAssignmentProjectionModel { + plan: solve::AlgebraicProjectionPlan::default(), + assignment_value: 2.0, + }; + let t = f64::from_bits(2.0f64.to_bits() + 1) - 2.0; + let target = 2.0 + t; + let mut y = vec![2.0]; + let step = solve::AlgebraicProjectionStep { row: 0, y_index: 0 }; + + let changed = project_causal_step(&model, &mut y, &[], t, &step, 1.0e-6) + .expect("direct assignment must preserve a representable right-limit change"); + + assert!(changed); + assert_eq!(y, vec![target]); + } + + #[test] + fn project_algebraics_validates_direct_assignment_against_full_residual() { + let model = DirectAssignmentProjectionModel { + plan: solve::AlgebraicProjectionPlan { + blocks: vec![solve::AlgebraicProjectionBlock { + rows: vec![0], + y_indices: vec![0], + causal_steps: vec![solve::AlgebraicProjectionStep { row: 0, y_index: 0 }], + }], + }, + assignment_value: 2.0, + }; + let mut y = vec![0.0]; + + project_algebraics(&model, &mut y, &[], 0.0, 0, 1.0e-12) + .expect("consistent direct assignment should satisfy the full residual"); + + let mut residual = vec![f64::NAN]; + model + .eval_residual(&y, &[], 0.0, &mut residual) + .expect("residual evaluation should succeed"); + assert_eq!(y, vec![2.0]); + assert_eq!(residual, vec![0.0]); + } + + #[test] + fn project_algebraics_rejects_direct_assignment_that_disagrees_with_residual() { + let model = DirectAssignmentProjectionModel { + plan: solve::AlgebraicProjectionPlan { + blocks: vec![solve::AlgebraicProjectionBlock { + rows: vec![0], + y_indices: vec![0], + causal_steps: vec![solve::AlgebraicProjectionStep { row: 0, y_index: 0 }], + }], + }, + assignment_value: 3.0, + }; + let mut y = vec![0.0]; + + project_algebraics(&model, &mut y, &[], 0.0, 0, 1.0e-12) + .expect_err("inconsistent direct assignment must not be masked as converged"); + let mut residual = vec![f64::NAN]; + model + .eval_residual(&y, &[], 0.0, &mut residual) + .expect("residual evaluation should succeed"); + assert_ne!(residual, vec![0.0]); + } + #[test] fn project_algebraics_accepts_scaled_residual_with_sub_tolerance_correction() { let model = ScaledResidualProjectionModel; diff --git a/crates/rumoca-solver/src/sparsity.rs b/crates/rumoca-solver/src/sparsity.rs index e2b0924b6..31edb1c9a 100644 --- a/crates/rumoca-solver/src/sparsity.rs +++ b/crates/rumoca-solver/src/sparsity.rs @@ -134,7 +134,8 @@ pub fn seeds_affecting_regs(ops: &[LinearOp]) -> HashMap> { | LinearOp::RandomState { dst, .. } | LinearOp::ImpureRandomInit { dst, .. } | LinearOp::ImpureRandom { dst, .. } - | LinearOp::ImpureRandomInteger { dst, .. } => { + | LinearOp::ImpureRandomInteger { dst, .. } + | LinearOp::ExternalCall { dst, .. } => { // Conservative: we don't track through these; leave dst absent from // the map so callers treat it as "unknown / all seeds". let _ = dst; @@ -190,6 +191,41 @@ fn checked_sparsity_row_end( .ok_or_else(|| RuntimeSolveError::solve_ir(format!("{context} overflows host index range"))) } +fn linsolve_sparsity_output_rows( + row_offset: usize, + n: usize, + output_indices: &[usize], +) -> Result<(Vec, usize), RuntimeSolveError> { + if output_indices.is_empty() { + let end = checked_sparsity_row_end(row_offset, n, "LinSolve sparsity row count")?; + return Ok(((row_offset..end).collect(), end)); + } + if output_indices.len() != n { + return Err(RuntimeSolveError::solve_ir(format!( + "LinSolve has {n} components but {} output indices while computing sparsity", + output_indices.len() + ))); + } + + let mut rows = Vec::with_capacity(n); + let mut seen = HashSet::with_capacity(n); + let mut next_row_offset = row_offset; + for &output_index in output_indices { + if !seen.insert(output_index) { + return Err(RuntimeSolveError::solve_ir(format!( + "LinSolve has duplicate output index {output_index} while computing sparsity" + ))); + } + next_row_offset = next_row_offset.max(checked_sparsity_row_end( + output_index, + 1, + "LinSolve sparsity output index", + )?); + rows.push(output_index); + } + Ok((rows, next_row_offset)) +} + fn push_matmul_jac_sparsity( column_rows: &mut [Vec], rhs_ops: &[LinearOp], @@ -302,16 +338,21 @@ impl SolveVisitor for JacSparsityVisitor { )?; } - ComputeNode::LinSolve { setup_ops, n, .. } => { + ComputeNode::LinSolve { + setup_ops, + n, + output_indices, + .. + } => { // Conservative: any seed in setup_ops may affect any output row. let reg_seeds = seeds_affecting_regs(setup_ops); let all_seeds = sorted_unique_seeds_from_regs(®_seeds); - let end = - checked_sparsity_row_end(self.row_offset, *n, "LinSolve sparsity row count")?; - for out_row in self.row_offset..end { + let (output_rows, next_row_offset) = + linsolve_sparsity_output_rows(self.row_offset, *n, output_indices)?; + for out_row in output_rows { push_row_for_seeds(&mut self.column_rows, all_seeds.iter().copied(), out_row); } - self.row_offset = end; + self.row_offset = next_row_offset; } ComputeNode::Map { @@ -525,4 +566,57 @@ mod tests { "dense: seed 1 affects both output rows" ); } + + #[test] + fn test_linsolve_sparsity_preserves_noncontiguous_output_rows_and_cursor() { + let block = ComputeBlock { + nodes: vec![ + ComputeNode::LinSolve { + setup_ops: vec![LinearOp::LoadSeed { dst: 0, index: 0 }], + matrix_start: 0, + rhs_start: 0, + n: 2, + next_reg: 1, + output_indices: vec![0, 2], + metadata: rumoca_ir_solve::TensorNodeMetadata::default(), + span: fixture_span(), + }, + ComputeNode::MatMul { + lhs_ops: vec![LinearOp::Const { dst: 0, value: 1.0 }], + lhs_start: 0, + rhs_ops: vec![LinearOp::LoadSeed { dst: 1, index: 1 }], + rhs_start: 1, + m: 1, + k: 1, + n: 1, + lhs_sparsity: SparsityPattern::Dense, + rhs_sparsity: SparsityPattern::Dense, + metadata: rumoca_ir_solve::TensorNodeMetadata::default(), + span: fixture_span(), + }, + ], + }; + + let column_rows = + compute_block_jac_sparsity(&block, 2).expect("noncontiguous LinSolve sparsity"); + + assert_eq!(column_rows[0], vec![0, 2]); + assert_eq!(column_rows[1], vec![3]); + assert!(column_rows.iter().all(|rows| !rows.contains(&1))); + } + + #[test] + fn test_linsolve_sparsity_rejects_malformed_explicit_output_indices() { + let wrong_length = linsolve_sparsity_output_rows(0, 2, &[0]) + .expect_err("explicit LinSolve map must have one row per component"); + assert!( + wrong_length + .to_string() + .contains("2 components but 1 output indices") + ); + + let duplicate = linsolve_sparsity_output_rows(0, 2, &[0, 0]) + .expect_err("explicit LinSolve map must not duplicate rows"); + assert!(duplicate.to_string().contains("duplicate output index 0")); + } } diff --git a/crates/rumoca-solver/src/timeline.rs b/crates/rumoca-solver/src/timeline.rs index 900b2da63..5d92de1a9 100644 --- a/crates/rumoca-solver/src/timeline.rs +++ b/crates/rumoca-solver/src/timeline.rs @@ -221,10 +221,7 @@ pub fn event_right_limit_time(t_event: f64) -> f64 { if !t_event.is_finite() { return t_event; } - let next = next_representable_time(t_event); - let step = (1.0e-6 * (1.0 + t_event.abs())).max(f64::EPSILON * (1.0 + t_event.abs())); - let stepped = t_event + step; - next.max(stepped) + next_representable_time(t_event) } pub fn merge_output_times_with_event_observations( @@ -289,7 +286,7 @@ mod tests { let t_event = 0.001_f64; let t_right = event_right_limit_time(t_event); assert!(t_right > t_event); - assert!(t_right - t_event >= 1.0e-6); + assert_eq!(t_right, f64::from_bits(t_event.to_bits() + 1)); } #[test] diff --git a/crates/rumoca-test-msl/src/lib.rs b/crates/rumoca-test-msl/src/lib.rs index ed04ca9e1..5c75b2f8b 100644 --- a/crates/rumoca-test-msl/src/lib.rs +++ b/crates/rumoca-test-msl/src/lib.rs @@ -6,6 +6,7 @@ pub mod msl_flamegraph; pub mod msl_tools; pub mod proc; +pub mod runtime_measurement; pub mod web_assets; use std::path::PathBuf; diff --git a/crates/rumoca-test-msl/src/msl_tools/common.rs b/crates/rumoca-test-msl/src/msl_tools/common.rs index f5d01adb4..bf4f33248 100644 --- a/crates/rumoca-test-msl/src/msl_tools/common.rs +++ b/crates/rumoca-test-msl/src/msl_tools/common.rs @@ -254,8 +254,7 @@ pub fn msl_top_packages(paths: &MslPaths) -> Vec<&'static str> { } pub fn get_omc_version() -> String { - let mut command = Command::new("omc"); - command.arg("--version"); + let mut command = omc_version_command(); let timeout = Duration::from_secs(10); let output = run_command_with_timeout(&mut command, timeout); match output { @@ -264,6 +263,24 @@ pub fn get_omc_version() -> String { } } +pub fn omc_version_command() -> Command { + if let Ok(image) = std::env::var("RUMOCA_OMC_DOCKER_IMAGE") + && !image.trim().is_empty() + { + let mut command = Command::new("docker"); + command + .arg("run") + .arg("--rm") + .arg(image.trim()) + .arg("omc") + .arg("--version"); + return command; + } + let mut command = Command::new("omc"); + command.arg("--version"); + command +} + pub fn get_git_commit(repo_root: &Path) -> String { let mut command = Command::new("git"); command.arg("rev-parse").arg("HEAD").current_dir(repo_root); diff --git a/crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference.rs b/crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference.rs index a123f8a3e..7cb974cdd 100644 --- a/crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference.rs +++ b/crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference.rs @@ -5,6 +5,7 @@ use super::common::{ has_fatal_omc_error, load_target_models, msl_load_lines, round3, summarize_batch_timings, summarize_omc_error, unix_timestamp_seconds, write_pretty_json, }; +use crate::runtime_measurement::{HostLoadSnapshot, sample_host_load}; use anyhow::{Context, Result, bail}; use clap::Args as ClapArgs; use rumoca_sim::sim_trace_compare::{ @@ -16,7 +17,7 @@ use std::cmp::Ordering; use std::collections::{BTreeMap, BTreeSet, HashMap}; use std::path::{Path, PathBuf}; use std::sync::Arc; -use std::sync::atomic::{AtomicUsize, Ordering as AtomicOrdering}; +use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering as AtomicOrdering}; use std::sync::mpsc; use std::thread; use std::time::{Duration, Instant}; @@ -126,6 +127,8 @@ struct SimModelResult { #[derive(Debug, Clone)] struct SimRunState { all_results: BTreeMap, + cached_omc_models: BTreeSet, + host_load_after: Option, // One per-model timing record. The shared `BatchTimingDetail` type is reused // (the compile reference genuinely batches); here each entry is one model. batch_timings: Vec, @@ -356,9 +359,26 @@ pub fn run(args: Args) -> Result<()> { ); let mut state = prepare_run_state(&omc_model_names); if cache_valid { - merge_cached_results_for_resume(&omc_ref_json, &omc_model_names, &mut state.all_results)?; - } - run_session_pending(&args, workers, &paths, &mut state, cache_valid)?; + merge_cached_results_for_resume(&omc_ref_json, &omc_model_names, &mut state)?; + } + let checkpoint = CheckpointContext { + args: &args, + paths: &paths, + selection: &selection, + context: FinalizeContext { + omc_version: omc_version.clone(), + git_commit: git_commit.clone(), + workers, + total, + n_batches, + effective_batch_size, + elapsed_seconds: 0.0, + cache_key: cache_key.clone(), + }, + started_at: overall_start, + }; + run_session_pending(&args, workers, &paths, &mut state, cache_valid, &checkpoint)?; + state.host_load_after = sample_host_load(); ensure_omc_trace_artifacts(&paths, &mut state.all_results); attach_rumoca_runtime(&rumoca_runtimes, &mut state.all_results); ensure_target_placeholders(&model_names, &rumoca_runtimes, &mut state.all_results); @@ -651,11 +671,19 @@ fn cached_reference_cache_key(omc_ref_json: &Path) -> Option { fn prepare_run_state(model_names: &[String]) -> SimRunState { SimRunState { all_results: BTreeMap::new(), + cached_omc_models: BTreeSet::new(), + host_load_after: None, batch_timings: Vec::new(), pending_models: model_names.to_vec(), } } +/// Message emitted by a persistent session worker. +enum SessionWorkerMessage { + Model(Box), + Fatal(SessionWorkerFatal), +} + /// Outcome of simulating one model inside a persistent session worker. struct SessionModelOutcome { idx: usize, @@ -665,11 +693,21 @@ struct SessionModelOutcome { timed_out: bool, } +/// Infrastructure failure that prevents a worker from producing physical model +/// results. These must not be serialized as per-model OMC failures. +struct SessionWorkerFatal { + worker_idx: usize, + model: String, + error: String, +} + /// Shared, immutable context handed to each session worker thread. struct SessionWorkerCtx<'a> { + worker_idx: usize, models: &'a [String], next: &'a AtomicUsize, - tx: &'a mpsc::Sender, + cancel: &'a AtomicBool, + tx: &'a mpsc::Sender, msl_exprs: &'a [String], work_dir: &'a Path, stop_time: f64, @@ -684,6 +722,14 @@ struct SessionWorkerCtx<'a> { cpu_core_id: Option, } +struct CheckpointContext<'a> { + args: &'a Args, + paths: &'a MslPaths, + selection: &'a ModelSelection, + context: FinalizeContext, + started_at: Instant, +} + /// Persistent-session execution path: the OMC analogue of the rumoca warm /// worker queue. Each worker thread owns one [`OmcSession`] that loads the MSL /// once, then pulls models from a shared atomic index and simulates them. On a @@ -695,18 +741,15 @@ fn run_session_pending( paths: &MslPaths, state: &mut SimRunState, reuse_cached: bool, + checkpoint: &CheckpointContext<'_>, ) -> Result<()> { // When the cache is valid (OMC + MSL unchanged, no --force) the cached OMC // reference has already been merged into `all_results`, so skip any model // that already has a real cached result instead of re-running OMC. - let reuse_skip: BTreeSet = if reuse_cached { - state - .all_results - .iter() - .filter(|(model_name, result)| cached_omc_result_is_reusable(paths, model_name, result)) - .map(|(name, _)| name.clone()) - .collect() + let reuse_skip = if reuse_cached { + retain_reusable_cached_models(paths, state) } else { + state.cached_omc_models.clear(); BTreeSet::new() }; let models: Vec = std::mem::take(&mut state.pending_models) @@ -735,7 +778,7 @@ fn run_session_pending( let core_plan = rumoca_worker::cpu_core_plan(worker_count); let pinned = core_plan.iter().filter(|core| core.is_some()).count(); let sim_timeout = Duration::from_secs(args.batch_timeout_seconds.max(1)); - let startup_timeout = Duration::from_secs(60); + let startup_timeout = Duration::from_secs(270); let load_timeout = Duration::from_secs(120); let total = models.len(); println!( @@ -744,11 +787,13 @@ fn run_session_pending( let models = Arc::new(models); let next = Arc::new(AtomicUsize::new(0)); - let (tx, rx) = mpsc::channel::(); + let cancel = Arc::new(AtomicBool::new(false)); + let (tx, rx) = mpsc::channel::(); let mut handles = Vec::with_capacity(worker_count); for worker_idx in 0..worker_count { let models = Arc::clone(&models); let next = Arc::clone(&next); + let cancel = Arc::clone(&cancel); let tx = tx.clone(); let msl_exprs = msl_exprs.clone(); let work_dir = paths.sim_work_dir.clone(); @@ -758,8 +803,10 @@ fn run_session_pending( let cpu_core_id = core_plan.get(worker_idx).copied().flatten(); handles.push(thread::spawn(move || { run_one_session_worker(SessionWorkerCtx { + worker_idx, models: &models, next: &next, + cancel: &cancel, tx: &tx, msl_exprs: &msl_exprs, work_dir: &work_dir, @@ -775,8 +822,51 @@ fn run_session_pending( } drop(tx); + let result = collect_session_outcomes(rx, &cancel, paths, state, checkpoint, total); + for handle in handles { + let _ = handle.join(); + } + result +} + +fn retain_reusable_cached_models(paths: &MslPaths, state: &mut SimRunState) -> BTreeSet { + let reusable = state + .all_results + .iter() + .filter(|(model_name, result)| cached_omc_result_is_reusable(paths, model_name, result)) + .map(|(name, _)| name.clone()) + .collect::>(); + state + .cached_omc_models + .retain(|model| reusable.contains(model)); + reusable +} + +fn collect_session_outcomes( + rx: mpsc::Receiver, + cancel: &AtomicBool, + paths: &MslPaths, + state: &mut SimRunState, + checkpoint: &CheckpointContext<'_>, + total: usize, +) -> Result<()> { let mut completed = 0usize; - for outcome in rx { + let mut fatal_error: Option = None; + for message in rx { + if fatal_error.is_some() { + continue; + } + let outcome = match message { + SessionWorkerMessage::Model(outcome) => outcome, + SessionWorkerMessage::Fatal(fatal) => { + cancel.store(true, AtomicOrdering::Relaxed); + fatal_error = Some(format!( + "OMC session worker {} could not start for '{}': {}", + fatal.worker_idx, fatal.model, fatal.error + )); + continue; + } + }; completed += 1; state.batch_timings.push(BatchTimingDetail { batch_idx: outcome.idx, @@ -787,16 +877,87 @@ fn run_session_pending( skipped: false, }); state.all_results.insert(outcome.model, outcome.result); + ensure_omc_trace_artifacts(paths, &mut state.all_results); + write_omc_reference_checkpoint(checkpoint, state)?; if completed.is_multiple_of(25) || completed == total { println!(" OMC session progress: {completed}/{total} models"); } } - for handle in handles { - let _ = handle.join(); + if let Some(error) = fatal_error { + bail!("{error}"); } Ok(()) } +fn write_omc_reference_checkpoint( + checkpoint: &CheckpointContext<'_>, + state: &SimRunState, +) -> Result<()> { + let mut context = checkpoint.context.clone(); + context.elapsed_seconds = checkpoint.started_at.elapsed().as_secs_f64(); + let metrics = compute_run_metrics(context.total, state); + let payload = json!({ + "msl_version": MSL_VERSION, + "omc_version": context.omc_version, + "git_commit": context.git_commit, + "cache_key": context.cache_key, + "checkpoint": true, + "target_selection": { + "source_file": checkpoint.selection.source_file.display().to_string(), + "rule": checkpoint.selection.rule, + }, + "stop_time": checkpoint.args.stop_time, + "use_experiment_stop_time": checkpoint.args.use_experiment_stop_time, + "total_models": context.total, + "processed": state.all_results.len(), + "sim_successful": metrics.sim_successful, + "sim_failed": metrics.sim_failed, + "sim_timed_out": metrics.sim_timed_out, + "simulation_success_rate_percent": round3(metrics.success_rate), + "elapsed_seconds": round3(context.elapsed_seconds), + "timing": { + "selection_seconds": round3(checkpoint.selection.selection_seconds), + "batch_size_requested": checkpoint.args.batch_size, + "batch_size_effective": context.effective_batch_size, + "batch_timeout_seconds": checkpoint.args.batch_timeout_seconds, + "workers_requested": checkpoint.args.workers, + "workers_used": context.workers, + "omc_threads": checkpoint.args.omc_threads, + "batches_total": context.n_batches, + "batches_ran": metrics.ran_batches, + "batches_skipped": metrics.skipped_batches, + "batch_elapsed_stats": metrics.batch_stats, + "batch_details": state.batch_timings, + }, + "models": state.all_results, + }); + write_checkpoint_json_atomically( + &checkpoint + .paths + .results_dir + .join("omc_simulation_reference.json"), + &payload, + ) +} + +fn write_checkpoint_json_atomically(path: &Path, payload: &Value) -> Result<()> { + if let Some(parent) = path.parent() { + std::fs::create_dir_all(parent) + .with_context(|| format!("failed to create '{}'", parent.display()))?; + } + let tmp_path = path.with_extension(format!("json.tmp.{}", std::process::id())); + let serialized = serde_json::to_string_pretty(payload).context("failed to serialize JSON")?; + std::fs::write(&tmp_path, serialized) + .with_context(|| format!("failed to write checkpoint '{}'", tmp_path.display()))?; + std::fs::rename(&tmp_path, path).with_context(|| { + format!( + "failed to replace checkpoint '{}' with '{}'", + path.display(), + tmp_path.display() + ) + }) +} + fn run_one_session_worker(ctx: SessionWorkerCtx<'_>) { // Pin before spawning OMC: on Linux the omc child inherits this thread's // CPU affinity, so the whole session stays on one core. @@ -812,31 +973,27 @@ fn run_one_session_worker(ctx: SessionWorkerCtx<'_>) { // batch (one reload per recycle, not per model). let mut models_on_session = 0usize; loop { + if ctx.cancel.load(AtomicOrdering::Relaxed) { + break; + } let idx = ctx.next.fetch_add(1, AtomicOrdering::Relaxed); let Some(model) = ctx.models.get(idx) else { break; }; if session.is_none() { - match OmcSession::spawn( - ctx.work_dir, - ctx.msl_exprs, - ctx.omc_threads, - ctx.startup_timeout, - ctx.load_timeout, - ) { + match spawn_omc_session_with_retries(&ctx) { Ok(spawned) => { session = Some(spawned); models_on_session = 0; } Err(error) => { - let _ = ctx.tx.send(SessionModelOutcome { - idx, + ctx.cancel.store(true, AtomicOrdering::Relaxed); + let _ = ctx.tx.send(SessionWorkerMessage::Fatal(SessionWorkerFatal { + worker_idx: ctx.worker_idx, model: model.clone(), - result: session_error_result(format!("omc session spawn failed: {error}")), - elapsed_seconds: 0.0, - timed_out: false, - }); - continue; + error, + })); + break; } } } @@ -863,13 +1020,15 @@ fn run_one_session_worker(ctx: SessionWorkerCtx<'_>) { ) } }; - let _ = ctx.tx.send(SessionModelOutcome { - idx, - model: model.clone(), - result, - elapsed_seconds: elapsed, - timed_out, - }); + let _ = ctx + .tx + .send(SessionWorkerMessage::Model(Box::new(SessionModelOutcome { + idx, + model: model.clone(), + result, + elapsed_seconds: elapsed, + timed_out, + }))); // `session` is still alive only on the success path (timeout/io already // took and killed it). Recycle it once it has handled enough models. if session.is_some() { @@ -883,8 +1042,38 @@ fn run_one_session_worker(ctx: SessionWorkerCtx<'_>) { } } +fn spawn_omc_session_with_retries(ctx: &SessionWorkerCtx<'_>) -> Result { + let mut last_error = None; + for attempt in 1..=SESSION_SPAWN_MAX_ATTEMPTS { + match OmcSession::spawn( + ctx.work_dir, + ctx.msl_exprs, + ctx.omc_threads, + ctx.startup_timeout, + ctx.load_timeout, + ) { + Ok(spawned) => return Ok(spawned), + Err(error) => { + last_error = Some(error); + if attempt < SESSION_SPAWN_MAX_ATTEMPTS { + eprintln!( + " OMC session worker {}: spawn attempt {attempt}/{} failed; retrying", + ctx.worker_idx, SESSION_SPAWN_MAX_ATTEMPTS + ); + thread::sleep(SESSION_SPAWN_RETRY_DELAY); + } + } + } + } + Err(last_error + .map(|error| error.to_string()) + .unwrap_or_else(|| "OMC session spawn was not attempted".to_string())) +} + /// Recycle an omc session after this many models to bound its memory growth. const SESSION_RECYCLE_MODELS: usize = 25; +const SESSION_SPAWN_MAX_ATTEMPTS: usize = 2; +const SESSION_SPAWN_RETRY_DELAY: Duration = Duration::from_secs(15); /// Build a [`SimModelResult`] from a session simulate outcome, mirroring the /// success/error classification of the former script-parsing path. @@ -1032,7 +1221,6 @@ fn cached_omc_result_is_reusable( result: &SimModelResult, ) -> bool { match result.status.as_str() { - "error" | "timeout" => true, "success" => cached_omc_success_has_trace_source(paths, model_name, result), _ => false, } @@ -1155,6 +1343,7 @@ fn load_omc_csv_trace(model_name: &str, csv_path: &Path) -> Result { } Ok(SimTrace { model_name: Some(model_name.to_string()), + n_states: None, times, names, data, @@ -1191,7 +1380,7 @@ fn parse_csv_row(row: &str) -> Vec { fn merge_cached_results_for_resume( cached_reference_path: &Path, target_models: &[String], - all_results: &mut BTreeMap, + state: &mut SimRunState, ) -> Result<()> { if !cached_reference_path.is_file() { return Ok(()); @@ -1220,12 +1409,13 @@ fn merge_cached_results_for_resume( let Ok(cached_result) = serde_json::from_value::(cached.clone()) else { continue; }; - match all_results.get_mut(model_name) { + match state.all_results.get_mut(model_name) { Some(current) => hydrate_omc_fields_from_cached(current, &cached_result), None => { - all_results.insert(model_name.clone(), cached_result); + state.all_results.insert(model_name.clone(), cached_result); } } + state.cached_omc_models.insert(model_name.clone()); } Ok(()) } @@ -1304,12 +1494,6 @@ fn quantify_trace_differences( .insert(model_name.clone(), "missing rumoca trace path".to_string()); continue; }; - let Some(omc_trace_path) = resolve_omc_trace_path(paths, model_name, omc_model) else { - report - .missing_trace - .insert(model_name.clone(), "missing omc trace path".to_string()); - continue; - }; let rumoca_trace = match load_trace_json(&rumoca_trace_path) { Ok(trace) => trace, Err(error) => { @@ -1320,6 +1504,14 @@ fn quantify_trace_differences( continue; } }; + state_selection::validate_rumoca_state_metadata(&rumoca_trace) + .with_context(|| format!("invalid state metadata contract for `{model_name}`"))?; + let Some(omc_trace_path) = resolve_omc_trace_path(paths, model_name, omc_model) else { + report + .missing_trace + .insert(model_name.clone(), "missing omc trace path".to_string()); + continue; + }; let omc_trace = match load_trace_json(&omc_trace_path) { Ok(trace) => trace, Err(error) => { @@ -1340,7 +1532,8 @@ fn quantify_trace_differences( } }; let state_selection = - state_selection::compare_model_state_selection(paths, model_name, &rumoca_trace); + state_selection::compare_model_state_selection(paths, model_name, &rumoca_trace) + .with_context(|| format!("invalid state metadata contract for `{model_name}`"))?; report.models.insert( model_name.clone(), TraceModelMetric { @@ -1417,7 +1610,6 @@ fn finalize_and_write_output( ) -> Result<()> { state.batch_timings.sort_by_key(|batch| batch.batch_idx); let metrics = compute_run_metrics(context.total, &state); - ensure_runtime_ratio_stats_present(&metrics, both_success_model_count(&state))?; let trace_summary = compute_trace_output_summary(&trace_report); let output = build_sim_output_payload( args, @@ -1471,41 +1663,6 @@ fn agreeing_model_names(trace_report: &TraceQuantification) -> BTreeSet .collect() } -fn both_success_model_count(state: &SimRunState) -> usize { - state - .all_results - .values() - .filter(|result| { - result.status == "success" && result.rumoca_status.as_deref() == Some("sim_ok") - }) - .count() -} - -fn ensure_runtime_ratio_stats_present( - metrics: &RunMetrics, - both_success_models: usize, -) -> Result<()> { - if both_success_models == 0 { - return Ok(()); - } - let mut missing = Vec::new(); - if metrics.system_ratio_both_success.is_none() { - missing.push("system_ratio_both_success"); - } - if metrics.wall_ratio_both_success.is_none() { - missing.push("wall_ratio_both_success"); - } - if missing.is_empty() { - return Ok(()); - } - bail!( - "missing runtime ratio stats for {} both-success model(s): {}. \ - Ensure OMC solver + external wall timings are captured before writing parity output.", - both_success_models, - missing.join(", ") - ); -} - fn compute_run_metrics(total: usize, state: &SimRunState) -> RunMetrics { let sim_successful = state .all_results diff --git a/crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference/omc_session.rs b/crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference/omc_session.rs index e841df604..fa80335a4 100644 --- a/crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference/omc_session.rs +++ b/crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference/omc_session.rs @@ -14,14 +14,18 @@ //! same win the rumoca warm worker provides. use anyhow::{Context, Result, anyhow}; +use std::env; +use std::ffi::OsString; +use std::fs::File; use std::path::{Path, PathBuf}; use std::process::{Child, Command, Stdio}; use std::time::{Duration, Instant}; -use super::super::common::apply_omc_thread_env; - /// Poll interval while waiting for the OMC ZeroMQ port file to appear. const PORT_FILE_POLL: Duration = Duration::from_millis(20); +/// Poll interval while waiting for a reply so a dead OMC child is noticed promptly. +const REPLY_POLL: Duration = Duration::from_millis(20); +const FALLBACK_DOCKER_OMC_IMAGE: &str = "openmodelica/openmodelica:v1.26.3-minimal"; /// Per-model timing self-reported by OMC's `SimulationResult` record. These are /// OMC's own internal phase timers, so they are independent of scheduling jitter @@ -68,6 +72,7 @@ pub(super) struct OmcSession { // The context must outlive the socket; keep it owned by the session. _ctx: zmq::Context, port_file: PathBuf, + suffix: String, } impl OmcSession { @@ -88,39 +93,48 @@ impl OmcSession { ) })?; - let mut command = Command::new("omc"); - command - .arg("--interactive=zmq") - .arg(format!("-z={suffix}")) - .arg("--locale=C"); - apply_omc_thread_env(&mut command, omc_threads); + let (mut command, command_kind) = + build_omc_interactive_command(work_dir, &suffix, omc_threads); + let spawn_log_file = work_dir.join(format!("omc_session_{suffix}.log")); + let spawn_log = File::create(&spawn_log_file).with_context(|| { + format!( + "failed to create omc session log '{}'", + spawn_log_file.display() + ) + })?; // Keep the port file and any scratch output inside our work dir so // concurrent worker sessions never collide on the default $TMPDIR file. command .env("TMPDIR", work_dir) .current_dir(work_dir) .stdin(Stdio::null()) - .stdout(Stdio::null()) - .stderr(Stdio::null()); + .stdout(Stdio::from( + spawn_log + .try_clone() + .context("failed to clone omc session log handle")?, + )) + .stderr(Stdio::from(spawn_log)); // Put omc in its own process group so that on kill we can also reap the // separate simulation executables `simulate(...)` spawns — otherwise a // model whose integration hangs leaves a grandchild pegging a CPU core // after we kill the omc parent. #[cfg(unix)] std::os::unix::process::CommandExt::process_group(&mut command, 0); - let child = command + if command_kind == OmcCommandKind::Native { + configure_native_omc_server_identity(&mut command, work_dir)?; + } + let mut child = command .spawn() .context("failed to spawn omc interactive session")?; - let port_file = match wait_for_port_file(work_dir, &suffix, startup_timeout) { - Some(path) => path, - None => { - let mut child = child; + let port_file = match wait_for_port_file(work_dir, &suffix, startup_timeout, &mut child) { + Ok(path) => path, + Err(error) => { let _ = child.kill(); let _ = child.wait(); return Err(anyhow!( - "omc session port file for suffix '{suffix}' did not appear within {:.1}s", - startup_timeout.as_secs_f64() + "{error}; {}", + omc_session_log_tail(&spawn_log_file) )); } }; @@ -144,6 +158,7 @@ impl OmcSession { socket, _ctx: ctx, port_file, + suffix, }; for expr in msl_load_exprs { @@ -160,18 +175,45 @@ impl OmcSession { /// Evaluate a single OMC expression, waiting at most `timeout` for the reply. pub(super) fn eval(&mut self, expr: &str, timeout: Duration) -> Result { - let millis = i32::try_from(timeout.as_millis().max(1)).unwrap_or(i32::MAX); - self.socket - .set_rcvtimeo(millis) - .map_err(|error| OmcEvalError::Io(anyhow!("set_rcvtimeo failed: {error}")))?; self.socket .send(expr, 0) .map_err(|error| OmcEvalError::Io(anyhow!("send failed: {error}")))?; - match self.socket.recv_string(0) { - Ok(Ok(reply)) => Ok(reply), - Ok(Err(_)) => Err(OmcEvalError::Io(anyhow!("omc reply was not valid utf-8"))), - Err(zmq::Error::EAGAIN) => Err(OmcEvalError::Timeout), - Err(error) => Err(OmcEvalError::Io(anyhow!("recv failed: {error}"))), + + let started = Instant::now(); + loop { + match self.child.try_wait() { + Ok(Some(status)) => { + return Err(OmcEvalError::Io(anyhow!( + "omc process exited while waiting for reply (status={status})" + ))); + } + Ok(None) => {} + Err(error) => { + return Err(OmcEvalError::Io(anyhow!( + "failed to poll omc process status while waiting for reply: {error}" + ))); + } + } + + let remaining = timeout.saturating_sub(started.elapsed()); + if remaining.is_zero() { + return Err(OmcEvalError::Timeout); + } + let poll = remaining.min(REPLY_POLL); + let millis = i32::try_from(poll.as_millis().max(1)).unwrap_or(i32::MAX); + self.socket + .set_rcvtimeo(millis) + .map_err(|error| OmcEvalError::Io(anyhow!("set_rcvtimeo failed: {error}")))?; + match self.socket.recv_string(0) { + Ok(Ok(reply)) => return Ok(reply), + Ok(Err(_)) => { + return Err(OmcEvalError::Io(anyhow!("omc reply was not valid utf-8"))); + } + Err(zmq::Error::EAGAIN) => {} + Err(error) => { + return Err(OmcEvalError::Io(anyhow!("recv failed: {error}"))); + } + } } } @@ -205,6 +247,7 @@ impl OmcSession { self.kill_process_group(); let _ = self.child.kill(); let _ = self.child.wait(); + self.kill_docker_container_by_suffix(); let _ = std::fs::remove_file(&self.port_file); } @@ -220,6 +263,39 @@ impl OmcSession { ); } } + + fn kill_docker_container_by_suffix(&self) { + let Ok(output) = Command::new("docker") + .args([ + "ps", + "--no-trunc", + "--filter", + &format!("ancestor={}", docker_omc_image()), + "--format", + "{{.ID}} {{.Command}}", + ]) + .output() + else { + return; + }; + if !output.status.success() { + return; + } + let stdout = String::from_utf8_lossy(&output.stdout); + for line in stdout.lines() { + if !line.contains(&self.suffix) { + continue; + } + if let Some(container_id) = line.split_whitespace().next() { + let _ = Command::new("docker") + .arg("kill") + .arg(container_id) + .stdout(Stdio::null()) + .stderr(Stdio::null()) + .status(); + } + } + } } impl Drop for OmcSession { @@ -230,10 +306,187 @@ impl Drop for OmcSession { self.kill_process_group(); let _ = self.child.kill(); let _ = self.child.wait(); + self.kill_docker_container_by_suffix(); let _ = std::fs::remove_file(&self.port_file); } } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum OmcCommandKind { + Native, + DockerWrapper, +} + +fn build_omc_interactive_command( + work_dir: &Path, + suffix: &str, + omc_threads: usize, +) -> (Command, OmcCommandKind) { + if omc_command_is_docker_wrapper() { + return ( + build_docker_omc_interactive_command(work_dir, suffix, omc_threads), + OmcCommandKind::DockerWrapper, + ); + } + let mut command = Command::new("omc"); + command + .arg("--interactive=zmq") + .arg(format!("-z={suffix}")) + .arg("--locale=C"); + apply_omc_thread_env_to_native_command(&mut command, omc_threads); + (command, OmcCommandKind::Native) +} + +fn build_docker_omc_interactive_command( + work_dir: &Path, + suffix: &str, + omc_threads: usize, +) -> Command { + let threads = omc_threads.max(1).to_string(); + let home = env::var_os("HOME").unwrap_or_else(|| OsString::from("/tmp")); + let user = docker_user_arg(); + let image = docker_omc_image(); + let mut command = Command::new("docker"); + command + .arg("run") + .arg("--rm") + .arg("-i") + .arg("--network") + .arg("host") + .arg("-v") + .arg(format!( + "{}:{}", + home.to_string_lossy(), + home.to_string_lossy() + )) + .arg("-e") + .arg(format!("HOME={}", home.to_string_lossy())) + .arg("-e") + .arg(format!("TMPDIR={}", work_dir.display())) + .arg("-e") + .arg(format!("OMP_NUM_THREADS={threads}")) + .arg("-e") + .arg(format!("OPENBLAS_NUM_THREADS={threads}")) + .arg("-e") + .arg(format!("MKL_NUM_THREADS={threads}")) + .arg("-e") + .arg(format!("NUMEXPR_NUM_THREADS={threads}")) + .arg("-w") + .arg(work_dir) + .arg("--user") + .arg(user) + .arg(image) + .arg("omc") + .arg("--interactive=zmq") + .arg(format!("-z={suffix}")) + .arg("--locale=C") + .arg(format!("--numProcs={threads}")); + command +} + +fn apply_omc_thread_env_to_native_command(command: &mut Command, omc_threads: usize) { + let threads = omc_threads.max(1).to_string(); + command.arg(format!("--numProcs={threads}")); + command.env("OMP_NUM_THREADS", &threads); + command.env("OPENBLAS_NUM_THREADS", &threads); + command.env("MKL_NUM_THREADS", &threads); + command.env("NUMEXPR_NUM_THREADS", &threads); +} + +#[cfg(unix)] +fn configure_native_omc_server_identity(command: &mut Command, work_dir: &Path) -> Result<()> { + if !native_omc_server_must_drop_root(&command_stdout("id", &["-u"])) { + return Ok(()); + } + use std::os::unix::fs::PermissionsExt; + use std::os::unix::process::CommandExt; + + // OpenModelica refuses to expose the unauthenticated server interface as + // root. The CI container itself runs as root, so only the OMC server child + // is dropped to nobody/nogroup while the Rust gate remains unchanged. + std::fs::set_permissions(work_dir, std::fs::Permissions::from_mode(0o777)).with_context( + || { + format!( + "failed to make OMC session work dir writable by unprivileged user '{}'", + work_dir.display() + ) + }, + )?; + command.uid(OMC_SERVER_UID); + command.gid(OMC_SERVER_GID); + command.env("HOME", work_dir); + command.env("USER", "nobody"); + command.env("LOGNAME", "nobody"); + Ok(()) +} + +#[cfg(not(unix))] +fn configure_native_omc_server_identity(_command: &mut Command, _work_dir: &Path) -> Result<()> { + Ok(()) +} + +#[cfg(unix)] +fn native_omc_server_must_drop_root(uid_text: &str) -> bool { + uid_text.trim() == "0" +} + +#[cfg(unix)] +const OMC_SERVER_UID: u32 = 65_534; +#[cfg(unix)] +const OMC_SERVER_GID: u32 = 65_534; + +fn docker_user_arg() -> String { + format!( + "{}:{}", + command_stdout("id", &["-u"]), + command_stdout("id", &["-g"]) + ) +} + +fn command_stdout(program: &str, args: &[&str]) -> String { + Command::new(program) + .args(args) + .output() + .ok() + .filter(|output| output.status.success()) + .map(|output| String::from_utf8_lossy(&output.stdout).trim().to_string()) + .filter(|value| !value.is_empty()) + .unwrap_or_else(|| "0".to_string()) +} + +fn omc_command_is_docker_wrapper() -> bool { + let Some(path) = find_executable_on_path("omc") else { + return false; + }; + let Ok(text) = std::fs::read_to_string(path) else { + return false; + }; + text.contains("docker run") && text.contains("openmodelica/openmodelica:") +} + +fn docker_omc_image() -> String { + if let Ok(image) = env::var("RUMOCA_OMC_DOCKER_IMAGE") + && !image.trim().is_empty() + { + return image.trim().to_string(); + } + find_executable_on_path("omc") + .and_then(|path| std::fs::read_to_string(path).ok()) + .and_then(|text| { + text.split_whitespace() + .find(|word| word.starts_with("openmodelica/openmodelica:")) + .map(|word| word.trim_end_matches('\\').to_string()) + }) + .unwrap_or_else(|| FALLBACK_DOCKER_OMC_IMAGE.to_string()) +} + +fn find_executable_on_path(name: &str) -> Option { + let path_var = env::var_os("PATH")?; + env::split_paths(&path_var) + .map(|dir| dir.join(name)) + .find(|path| path.is_file()) +} + fn unique_session_suffix() -> String { use std::sync::atomic::{AtomicU64, Ordering}; static COUNTER: AtomicU64 = AtomicU64::new(0); @@ -245,23 +498,68 @@ fn unique_session_suffix() -> String { format!("rumoca_{}_{seq}_{nanos}", std::process::id()) } -/// OMC writes its port file as `openmodelica..port.` in `$TMPDIR`. -/// We point `$TMPDIR` at `work_dir`, so look there and match on the suffix to -/// avoid depending on the resolved user name. -fn wait_for_port_file(work_dir: &Path, suffix: &str, timeout: Duration) -> Option { +/// OMC writes its port file as `openmodelica..port.` in a temp +/// directory. We point `TMPDIR` at `work_dir`, but some OpenModelica builds use +/// the process default temp dir instead, so search both and match on the unique +/// suffix rather than depending on the resolved user name. +fn wait_for_port_file( + work_dir: &Path, + suffix: &str, + timeout: Duration, + child: &mut Child, +) -> Result { let deadline = Instant::now() + timeout; let needle = format!("port.{suffix}"); + let search_dirs = omc_port_file_search_dirs(work_dir); loop { - if let Some(path) = find_port_file(work_dir, &needle) { - return Some(path); + if let Some(path) = find_port_file_in_dirs(&search_dirs, &needle) { + return Ok(path); + } + match child.try_wait() { + Ok(Some(status)) => { + return Err(format!( + "omc interactive session exited before writing port file for suffix '{suffix}' (status={status})" + )); + } + Ok(None) => {} + Err(error) => { + return Err(format!( + "failed to poll omc interactive session status: {error}" + )); + } } if Instant::now() >= deadline { - return None; + return Err(format!( + "omc session port file for suffix '{suffix}' did not appear within {:.1}s (searched {})", + timeout.as_secs_f64(), + search_dirs + .iter() + .map(|path| path.display().to_string()) + .collect::>() + .join(", ") + )); } std::thread::sleep(PORT_FILE_POLL); } } +fn omc_port_file_search_dirs(work_dir: &Path) -> Vec { + let mut dirs = vec![ + work_dir.to_path_buf(), + env::temp_dir(), + PathBuf::from("/tmp"), + ]; + dirs.sort(); + dirs.dedup(); + dirs +} + +fn find_port_file_in_dirs(search_dirs: &[PathBuf], needle: &str) -> Option { + search_dirs + .iter() + .find_map(|dir| find_port_file(dir, needle)) +} + fn find_port_file(work_dir: &Path, needle: &str) -> Option { let entries = std::fs::read_dir(work_dir).ok()?; entries.flatten().find_map(|entry| { @@ -271,6 +569,21 @@ fn find_port_file(work_dir: &Path, needle: &str) -> Option { }) } +fn omc_session_log_tail(path: &Path) -> String { + let text = std::fs::read_to_string(path).unwrap_or_default(); + let trimmed = text.trim(); + if trimmed.is_empty() { + return format!("omc session log '{}' is empty", path.display()); + } + const MAX_CHARS: usize = 4000; + let tail = if trimmed.len() > MAX_CHARS { + &trimmed[trimmed.len() - MAX_CHARS..] + } else { + trimmed + }; + format!("omc session log '{}': {tail}", path.display()) +} + /// Strip the surrounding quotes OMC puts around string replies and unescape the /// common `\n`/`\"` sequences. fn unquote_omc_string(text: &str) -> String { @@ -333,6 +646,85 @@ fn extract_record_f64(record: &str, field: &str) -> Option { mod tests { use super::*; + #[cfg(unix)] + #[test] + fn eval_reports_child_exit_promptly_instead_of_timeout() { + let ctx = zmq::Context::new(); + let endpoint = format!("inproc://omc-dead-child-{}", unique_session_suffix()); + let server = ctx.socket(zmq::REP).expect("create test REP socket"); + server.bind(&endpoint).expect("bind test REP socket"); + let socket = ctx.socket(zmq::REQ).expect("create test REQ socket"); + socket.connect(&endpoint).expect("connect test REQ socket"); + let child = Command::new("sh") + .args(["-c", "sleep 0.05; exit 23"]) + .spawn() + .expect("spawn short-lived child"); + let temp = tempfile::tempdir().expect("test tempdir"); + let mut session = OmcSession { + child, + socket, + _ctx: ctx, + port_file: temp.path().join("unused.port"), + suffix: unique_session_suffix(), + }; + + let request_budget = Duration::from_secs(2); + let started = Instant::now(); + let error = session + .eval("getVersion()", request_budget) + .expect_err("dead child must fail the request"); + let elapsed = started.elapsed(); + + assert!( + matches!(error, OmcEvalError::Io(_)), + "dead child was misclassified: {error}" + ); + assert!( + error.to_string().contains("status=exit status: 23"), + "child exit status should be actionable: {error}" + ); + assert!( + elapsed < Duration::from_secs(1), + "dead child should fail promptly, elapsed={elapsed:?}" + ); + + drop(server); + } + + #[cfg(unix)] + #[test] + fn eval_preserves_timeout_for_live_child() { + let ctx = zmq::Context::new(); + let endpoint = format!("inproc://omc-live-child-{}", unique_session_suffix()); + let server = ctx.socket(zmq::REP).expect("create test REP socket"); + server.bind(&endpoint).expect("bind test REP socket"); + let socket = ctx.socket(zmq::REQ).expect("create test REQ socket"); + socket.connect(&endpoint).expect("connect test REQ socket"); + let child = Command::new("sh") + .args(["-c", "sleep 5"]) + .spawn() + .expect("spawn live child"); + let temp = tempfile::tempdir().expect("test tempdir"); + let mut session = OmcSession { + child, + socket, + _ctx: ctx, + port_file: temp.path().join("unused.port"), + suffix: unique_session_suffix(), + }; + + let error = session + .eval("getVersion()", Duration::from_millis(100)) + .expect_err("missing reply must exhaust the request budget"); + + assert!( + matches!(error, OmcEvalError::Timeout), + "live child timeout was misclassified: {error}" + ); + + drop(server); + } + #[test] fn parse_sim_record_extracts_fields() { let record = r#"record SimulationResult @@ -376,4 +768,31 @@ end SimulationResult;"#; assert_eq!(unquote_omc_string("\"\""), ""); assert_eq!(unquote_omc_string("bare"), "bare"); } + + #[test] + fn find_port_file_in_dirs_searches_fallback_temp_dir() { + let work = tempfile::tempdir().expect("work tempdir"); + let fallback = tempfile::tempdir().expect("fallback tempdir"); + let suffix = "rumoca_test_suffix"; + let needle = format!("port.{suffix}"); + let port_file = fallback + .path() + .join(format!("openmodelica.test-user.port.{suffix}")); + std::fs::write(&port_file, "tcp://127.0.0.1:12345\n").expect("write port file"); + + let found = find_port_file_in_dirs( + &[work.path().to_path_buf(), fallback.path().to_path_buf()], + &needle, + ) + .expect("fallback port file should be found"); + assert_eq!(found, port_file); + } + + #[cfg(unix)] + #[test] + fn native_omc_server_identity_drop_is_root_only() { + assert!(native_omc_server_must_drop_root("0\n")); + assert!(!native_omc_server_must_drop_root("1001")); + assert!(!native_omc_server_must_drop_root("")); + } } diff --git a/crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference/output.rs b/crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference/output.rs index 945e0d650..bdd9be71a 100644 --- a/crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference/output.rs +++ b/crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference/output.rs @@ -392,13 +392,74 @@ fn build_timing_payload( }) } -fn build_runtime_comparison_payload(metrics: &RunMetrics) -> Value { +fn load_wall_time_provenance( + args: &Args, + paths: &MslPaths, + context: &FinalizeContext, + state: &SimRunState, +) -> crate::runtime_measurement::WallTimeMeasurementProvenance { + let rumoca_payload = std::fs::read_to_string(path_for_rumoca_results(paths)) + .ok() + .and_then(|text| serde_json::from_str::(&text).ok()); + let scheduler = rumoca_payload + .as_ref() + .and_then(|payload| payload.pointer("/timings/scheduler")); + let load_before = rumoca_payload + .as_ref() + .and_then(|payload| payload.pointer("/timings/host_load_before")) + .and_then(|value| { + serde_json::from_value::(value.clone()) + .ok() + }); + let comparable_models = state.all_results.iter().filter(|(_, result)| { + result.status == "success" + && result.rumoca_status.as_deref() == Some("sim_ok") + && runtime_pair(result.rumoca_sim_wall_seconds, result.omc_wall_seconds).is_some() + }); + let (omc_cached_sample_count, omc_fresh_sample_count) = + comparable_models.fold((0, 0), |(cached, fresh), (name, _)| { + if state.cached_omc_models.contains(name) { + (cached + 1, fresh) + } else { + (cached, fresh + 1) + } + }); + let scheduler_count = |field: &str| { + scheduler + .and_then(|value| value.get(field)) + .and_then(Value::as_u64) + .and_then(|value| usize::try_from(value).ok()) + .unwrap_or(0) + }; + crate::runtime_measurement::WallTimeMeasurementProvenance { + omc_fresh_sample_count, + omc_cached_sample_count, + affinity_requested_worker_count: scheduler_count("affinity_requested_worker_count"), + affinity_applied_worker_count: scheduler_count("affinity_applied_worker_count"), + affinity_failed_worker_count: scheduler_count("affinity_failed_worker_count"), + normalized_load_before: load_before.map(|snapshot| snapshot.normalized()), + normalized_load_after: state.host_load_after.map(|snapshot| snapshot.normalized()), + rumoca_workers_used: scheduler_count("worker_count"), + workers_used: context.workers, + omc_threads: args.omc_threads, + } +} + +fn build_runtime_comparison_payload( + args: &Args, + paths: &MslPaths, + context: &FinalizeContext, + metrics: &RunMetrics, + state: &SimRunState, +) -> Value { let runtime_ratio_stats = json!({ "system_ratio_all_positive": metrics.system_ratio_all_positive, "system_ratio_both_success": metrics.system_ratio_both_success, "wall_ratio_all_positive": metrics.wall_ratio_all_positive, "wall_ratio_both_success": metrics.wall_ratio_both_success, }); + let diagnostics = build_runtime_comparison_diagnostics(metrics); + let wall_time_provenance = load_wall_time_provenance(args, paths, context, state); json!({ "ratio_definition": "omc_over_rumoca_higher_is_better (simulation/wall runtime; this is the SIM-time comparison, distinct from the compile-speed `speedup` in msl_speed_comparison.json)", "ratio_metric_system": "omc_timeSimulation_over_rumoca_sim_seconds", @@ -412,6 +473,42 @@ fn build_runtime_comparison_payload(metrics: &RunMetrics) -> Value { "total_rumoca_sim_run_seconds": round3(metrics.total_rumoca_sim_run_seconds), "total_rumoca_sim_wall_seconds": round3(metrics.total_rumoca_sim_wall_seconds), "ratio_stats": runtime_ratio_stats, + "diagnostics": diagnostics, + "wall_time_provenance": wall_time_provenance, + }) +} + +fn build_runtime_comparison_diagnostics(metrics: &RunMetrics) -> Value { + let both_success_samples = metrics + .system_ratio_both_success + .as_ref() + .map(|stats| stats.sample_count) + .or_else(|| { + metrics + .wall_ratio_both_success + .as_ref() + .map(|stats| stats.sample_count) + }) + .unwrap_or(0); + let mut missing_ratio_stats = Vec::new(); + if metrics.system_ratio_both_success.is_none() { + missing_ratio_stats.push("system_ratio_both_success"); + } + if metrics.wall_ratio_both_success.is_none() { + missing_ratio_stats.push("wall_ratio_both_success"); + } + let reason = if missing_ratio_stats.is_empty() { + None + } else if both_success_samples == 0 { + Some("no_omc_rumoca_both_success_models") + } else { + Some("missing_runtime_ratio_stats") + }; + json!({ + "both_success_model_count": both_success_samples, + "runtime_ratio_available": missing_ratio_stats.is_empty(), + "missing_ratio_stats": missing_ratio_stats, + "unavailable_reason": reason, }) } @@ -494,7 +591,7 @@ pub(super) fn build_sim_output_payload( }); let mut timing = build_timing_payload(args, context, metrics, state); timing["selection_seconds"] = json!(round3(selection.selection_seconds)); - let runtime_comparison = build_runtime_comparison_payload(metrics); + let runtime_comparison = build_runtime_comparison_payload(args, paths, context, metrics, state); let trace_comparison = build_trace_comparison_payload(paths, trace_summary); let pipeline_progress = build_pipeline_progress_payload(context, metrics, trace_summary, state); let omc_assertion_failures = build_omc_assertion_failure_payload(state); diff --git a/crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference/state_selection.rs b/crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference/state_selection.rs index d9cf6398a..9238974e6 100644 --- a/crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference/state_selection.rs +++ b/crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference/state_selection.rs @@ -1,4 +1,5 @@ use super::{MslPaths, TraceQuantification}; +use anyhow::{Result, bail}; use rumoca_sim::sim_trace_compare::SimTrace; use serde::Serialize; use std::collections::BTreeSet; @@ -37,10 +38,16 @@ pub(super) fn compare_model_state_selection( paths: &MslPaths, model_name: &str, rumoca_trace: &SimTrace, -) -> Option { +) -> Result> { let rumoca_states = rumoca_state_names(rumoca_trace)?; - let omc_states = load_omc_state_names(paths, model_name)?; - Some(compare_state_sets(&rumoca_states, &omc_states)) + let Some(omc_states) = load_omc_state_names(paths, model_name) else { + return Ok(None); + }; + Ok(Some(compare_state_sets(&rumoca_states, &omc_states))) +} + +pub(super) fn validate_rumoca_state_metadata(trace: &SimTrace) -> Result<()> { + rumoca_state_names(trace).map(|_| ()) } pub(super) fn state_selection_summary(report: &TraceQuantification) -> StateSelectionSummary { @@ -100,15 +107,25 @@ fn compare_state_sets( } } -fn rumoca_state_names(trace: &SimTrace) -> Option> { - let states = trace - .variable_meta - .as_ref()? +fn rumoca_state_names(trace: &SimTrace) -> Result> { + let Some(expected_state_count) = trace.n_states else { + bail!("rumoca trace is missing the n_states metadata contract"); + }; + let Some(variable_meta) = trace.variable_meta.as_ref() else { + bail!("rumoca trace is missing variable_meta state metadata"); + }; + let states = variable_meta .iter() .filter(|meta| meta.role.as_deref() == Some("state")) .map(|meta| meta.name.clone()) .collect::>(); - Some(states) + if expected_state_count != states.len() { + bail!( + "rumoca trace declares {expected_state_count} states but metadata reports {}", + states.len() + ); + } + Ok(states) } fn load_omc_state_names(paths: &MslPaths, model_name: &str) -> Option> { @@ -194,6 +211,64 @@ mod tests { assert_eq!(states, BTreeSet::from(["a&b".to_string(), "x".to_string()])); } + #[test] + fn rejects_state_metadata_count_that_disagrees_with_trace_contract() { + let trace = serde_json::from_value::(serde_json::json!({ + "model_name": "BrokenMetadata", + "n_states": 2, + "times": [0.0], + "names": ["x"], + "data": [[0.0]], + "variable_meta": [{"name": "x", "role": "state"}] + })) + .expect("trace fixture should deserialize"); + + assert!(rumoca_state_names(&trace).is_err()); + } + + #[test] + fn rejects_state_metadata_without_trace_state_count_contract() { + let trace = serde_json::from_value::(serde_json::json!({ + "model_name": "MissingContract", + "times": [0.0], + "names": ["x"], + "data": [[0.0]], + "variable_meta": [{"name": "x", "role": "state"}] + })) + .expect("trace fixture should deserialize"); + + assert!(rumoca_state_names(&trace).is_err()); + } + + #[test] + fn invalid_state_metadata_contract_is_not_dropped_as_missing_comparison() { + let trace = serde_json::from_value::(serde_json::json!({ + "model_name": "BrokenMetadata", + "n_states": 2, + "times": [0.0], + "names": ["x"], + "data": [[0.0]], + "variable_meta": [{"name": "x", "role": "state"}] + })) + .expect("trace fixture should deserialize"); + let root = std::path::PathBuf::from("/nonexistent-state-contract-test"); + let paths = MslPaths { + repo_root: root.clone(), + msl_dir: root.clone(), + results_dir: root.clone(), + flat_dir: root.clone(), + work_dir: root.clone(), + sim_work_dir: root.clone(), + omc_trace_dir: root.clone(), + rumoca_trace_dir: root, + }; + + let error = compare_model_state_selection(&paths, "BrokenMetadata", &trace) + .expect_err("invalid producer metadata must fail the parity comparison"); + + assert!(error.to_string().contains("declares 2 states")); + } + #[test] fn summarizes_state_selection_agreement() { let exact = StateSelectionMetric { diff --git a/crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference/tests.rs b/crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference/tests.rs index 70db87c84..050e0007a 100644 --- a/crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference/tests.rs +++ b/crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference/tests.rs @@ -114,8 +114,8 @@ fn merge_cached_results_for_resume_hydrates_missing_omc_timing() { ) .expect("write payload"); - let mut all_results = BTreeMap::new(); - all_results.insert( + let mut state = prepare_run_state(&[model_name.to_string()]); + state.all_results.insert( model_name.to_string(), SimModelResult { status: "success".to_string(), @@ -139,9 +139,12 @@ fn merge_cached_results_for_resume_hydrates_missing_omc_timing() { }, ); - merge_cached_results_for_resume(&path, &[model_name.to_string()], &mut all_results) + merge_cached_results_for_resume(&path, &[model_name.to_string()], &mut state) .expect("merge cached results"); - let hydrated = all_results.get(model_name).expect("missing hydrated model"); + let hydrated = state + .all_results + .get(model_name) + .expect("missing hydrated model"); assert_eq!(hydrated.sim_system_seconds, Some(0.25)); assert_eq!(hydrated.total_system_seconds, Some(0.5)); assert_eq!(hydrated.omc_wall_seconds, Some(0.75)); @@ -153,6 +156,7 @@ fn merge_cached_results_for_resume_hydrates_missing_omc_timing() { hydrated.trace_file.as_deref(), Some("sim_traces/omc/Modelica.Blocks.Examples.PID_Controller.json") ); + assert!(state.cached_omc_models.contains(model_name)); } #[test] @@ -283,6 +287,30 @@ fn ensure_omc_trace_artifacts_regenerates_missing_json_from_error_result_with_cs assert!(trace_path.is_file(), "missing regenerated trace json"); } +#[test] +fn cached_error_and_timeout_results_are_not_reusable() { + let temp = tempfile::tempdir().expect("tempdir"); + let paths = MslPaths { + repo_root: temp.path().to_path_buf(), + msl_dir: temp.path().join("msl"), + results_dir: temp.path().join("results"), + flat_dir: temp.path().join("omc_flat"), + work_dir: temp.path().join("omc_work"), + sim_work_dir: temp.path().join("omc_sim_work"), + omc_trace_dir: temp.path().join("sim_traces").join("omc"), + rumoca_trace_dir: temp.path().join("sim_traces").join("rumoca"), + }; + let model_name = "Modelica.Blocks.Examples.PID_Controller"; + for status in ["error", "timeout"] { + let cached = SimModelResult { + status: status.to_string(), + error: Some("omc session port file did not appear".to_string()), + ..empty_omc_result() + }; + assert!(!cached_omc_result_is_reusable(&paths, model_name, &cached)); + } +} + #[test] fn cached_success_without_materialized_trace_source_is_not_reusable() { let temp = tempfile::tempdir().expect("tempdir"); @@ -345,6 +373,7 @@ fn cached_success_without_materialized_trace_source_is_not_reusable() { &omc_trace_dir.join(format!("{model_name}.json")), &SimTrace { model_name: Some(model_name.to_string()), + n_states: None, times: vec![0.0], names: vec!["y".to_string()], data: vec![vec![Some(1.0)]], @@ -393,6 +422,274 @@ fn compute_runtime_ratio_stats_reports_distribution() { assert!((filtered.max_ratio - 2.0).abs() < 1.0e-12); } +fn successful_runtime_pair() -> SimModelResult { + SimModelResult { + status: "success".to_string(), + error: None, + sim_system_seconds: Some(0.25), + total_system_seconds: Some(0.5), + omc_wall_seconds: Some(0.75), + result_file: None, + trace_file: None, + trace_error: None, + rumoca_status: Some("sim_ok".to_string()), + rumoca_ic_status: Some("ic_ok".to_string()), + rumoca_ic_error: None, + rumoca_ic_seconds: Some(0.1), + rumoca_sim_seconds: Some(0.4), + rumoca_sim_build_seconds: Some(0.1), + rumoca_sim_run_seconds: Some(0.3), + rumoca_sim_wall_seconds: Some(0.5), + rumoca_trace_file: None, + rumoca_trace_error: None, + } +} + +fn host_load(one_minute: f64, logical_cpus: usize) -> crate::runtime_measurement::HostLoadSnapshot { + crate::runtime_measurement::HostLoadSnapshot { + one_minute, + logical_cpus, + } +} + +fn scheduler_provenance() -> Value { + json!({ + "worker_count": 14, + "affinity_requested_worker_count": 14, + "affinity_applied_worker_count": 14, + "affinity_failed_worker_count": 0 + }) +} + +fn assert_complete_wall_time_provenance(payload: &Value) { + let provenance = &payload["runtime_comparison"]["wall_time_provenance"]; + assert_eq!(provenance["omc_cached_sample_count"], 1); + assert_eq!(provenance["omc_fresh_sample_count"], 1); + assert_eq!(provenance["affinity_requested_worker_count"], 14); + assert_eq!(provenance["affinity_applied_worker_count"], 14); + assert_eq!(provenance["affinity_failed_worker_count"], 0); + assert_eq!(provenance["normalized_load_before"], 0.5); + assert_eq!(provenance["normalized_load_after"], 0.75); + assert_eq!(provenance["rumoca_workers_used"], 14); + assert_eq!(provenance["workers_used"], 2); + assert_eq!(provenance["omc_threads"], 1); +} + +fn assert_missing_wall_time_provenance_defaults(payload: &Value) { + let provenance = &payload["runtime_comparison"]["wall_time_provenance"]; + assert_eq!(provenance["omc_cached_sample_count"], 0); + assert_eq!(provenance["omc_fresh_sample_count"], 0); + assert_eq!(provenance["affinity_requested_worker_count"], 0); + assert_eq!(provenance["affinity_applied_worker_count"], 0); + assert_eq!(provenance["affinity_failed_worker_count"], 0); + assert!(provenance["normalized_load_before"].is_null()); + assert!(provenance["normalized_load_after"].is_null()); +} + +#[test] +fn output_payload_records_fresh_cached_affinity_and_load_provenance() { + let temp = tempfile::tempdir().expect("tempdir"); + let results_dir = temp.path().join("results"); + std::fs::create_dir_all(&results_dir).expect("results dir"); + write_pretty_json( + &results_dir.join("msl_results.json"), + &json!({ + "timings": { + "host_load_before": {"one_minute": 2.0, "logical_cpus": 4}, + "scheduler": scheduler_provenance() + } + }), + ) + .expect("write rumoca results"); + let paths = MslPaths { + repo_root: temp.path().to_path_buf(), + msl_dir: temp.path().join("msl"), + results_dir: results_dir.clone(), + flat_dir: results_dir.join("omc_flat"), + work_dir: results_dir.join("omc_work"), + sim_work_dir: results_dir.join("omc_sim_work"), + omc_trace_dir: results_dir.join("sim_traces/omc"), + rumoca_trace_dir: results_dir.join("sim_traces/rumoca"), + }; + let args = Args { + dry_run: false, + batch_size: 1, + force: false, + workers: 2, + omc_threads: 1, + batch_timeout_seconds: 30, + stop_time: 1.0, + use_experiment_stop_time: false, + max_models: 0, + model_regex: None, + balance_results_file: None, + results_dir: None, + target_models_file: None, + trace_exclusions_file: None, + rumoca_sim_ok_only: false, + }; + let selection = ModelSelection { + names: vec!["cached".to_string(), "fresh".to_string()], + source_file: temp.path().join("targets.json"), + rule: "test".to_string(), + selection_seconds: 0.0, + }; + let context = FinalizeContext { + omc_version: "test".to_string(), + git_commit: "test".to_string(), + workers: 2, + total: 2, + n_batches: 2, + effective_batch_size: 1, + elapsed_seconds: 1.0, + cache_key: "test-key".to_string(), + }; + std::fs::create_dir_all(&paths.sim_work_dir).expect("sim work dir"); + std::fs::write(paths.sim_work_dir.join("cached_res.csv"), "time,y\n0,1\n") + .expect("cached result file"); + let mut reusable_cached = successful_runtime_pair(); + reusable_cached.result_file = Some("cached_res.csv".to_string()); + let cached_reference = results_dir.join("cached_reference.json"); + write_pretty_json( + &cached_reference, + &json!({ + "models": { + "cached": reusable_cached, + "fresh": successful_runtime_pair() + } + }), + ) + .expect("write cached reference"); + let model_names = ["cached".to_string(), "fresh".to_string()]; + let mut state = prepare_run_state(&model_names); + merge_cached_results_for_resume(&cached_reference, &model_names, &mut state) + .expect("merge cached results"); + assert_eq!(state.cached_omc_models, BTreeSet::from(model_names.clone())); + let reusable = retain_reusable_cached_models(&paths, &mut state); + assert_eq!(reusable, BTreeSet::from(["cached".to_string()])); + assert_eq!(state.cached_omc_models, reusable); + + state + .all_results + .insert("fresh".to_string(), successful_runtime_pair()); + state.host_load_after = Some(host_load(3.0, 4)); + let metrics = compute_run_metrics(context.total, &state); + let trace_summary = compute_trace_output_summary(&TraceQuantification::default()); + let payload = output::build_sim_output_payload( + &args, + &paths, + &selection, + &context, + &metrics, + &trace_summary, + &state, + ); + assert_complete_wall_time_provenance(&payload); +} + +#[test] +fn output_payload_keeps_empty_parity_diagnostics_when_omc_has_no_successes() { + let temp = tempfile::tempdir().expect("tempdir"); + let results_dir = temp.path().join("results"); + let paths = MslPaths { + repo_root: temp.path().to_path_buf(), + msl_dir: temp.path().join("msl"), + results_dir: results_dir.clone(), + flat_dir: results_dir.join("omc_flat"), + work_dir: results_dir.join("omc_work"), + sim_work_dir: results_dir.join("omc_sim_work"), + omc_trace_dir: results_dir.join("sim_traces").join("omc"), + rumoca_trace_dir: results_dir.join("sim_traces").join("rumoca"), + }; + let args = Args { + dry_run: false, + batch_size: 1, + force: false, + workers: 1, + omc_threads: 1, + batch_timeout_seconds: 120, + stop_time: 1.0, + use_experiment_stop_time: true, + max_models: 0, + model_regex: None, + balance_results_file: None, + results_dir: None, + target_models_file: None, + trace_exclusions_file: None, + rumoca_sim_ok_only: true, + }; + let selection = ModelSelection { + names: vec!["Modelica.Blocks.Examples.BooleanNetwork1".to_string()], + source_file: temp.path().join("targets.json"), + rule: "test".to_string(), + selection_seconds: 0.0, + }; + let context = FinalizeContext { + omc_version: "OpenModelica 1.27.0~dev".to_string(), + git_commit: "test".to_string(), + workers: 1, + total: 1, + n_batches: 1, + effective_batch_size: 1, + elapsed_seconds: 120.0, + cache_key: "test-key".to_string(), + }; + let mut all_results = BTreeMap::new(); + all_results.insert( + "Modelica.Blocks.Examples.BooleanNetwork1".to_string(), + SimModelResult { + status: "timeout".to_string(), + error: Some("omc simulate exceeded 120s budget".to_string()), + sim_system_seconds: None, + total_system_seconds: None, + omc_wall_seconds: Some(120.0), + result_file: None, + trace_file: None, + trace_error: None, + rumoca_status: Some("sim_ok".to_string()), + rumoca_ic_status: Some("ic_ok".to_string()), + rumoca_ic_error: None, + rumoca_ic_seconds: Some(0.1), + rumoca_sim_seconds: Some(0.2), + rumoca_sim_build_seconds: Some(0.05), + rumoca_sim_run_seconds: Some(0.15), + rumoca_sim_wall_seconds: Some(0.3), + rumoca_trace_file: Some("sim_traces/rumoca/A.json".to_string()), + rumoca_trace_error: None, + }, + ); + let mut state = prepare_run_state(&[]); + state.all_results = all_results; + let metrics = compute_run_metrics(context.total, &state); + let trace_summary = compute_trace_output_summary(&TraceQuantification::default()); + let payload = output::build_sim_output_payload( + &args, + &paths, + &selection, + &context, + &metrics, + &trace_summary, + &state, + ); + + assert_eq!(payload["sim_successful"], 0); + assert!(payload["runtime_comparison"]["ratio_stats"]["system_ratio_both_success"].is_null()); + assert!(payload["runtime_comparison"]["ratio_stats"]["wall_ratio_both_success"].is_null()); + assert_eq!( + payload["runtime_comparison"]["diagnostics"]["both_success_model_count"], + 0 + ); + assert_eq!( + payload["runtime_comparison"]["diagnostics"]["unavailable_reason"], + "no_omc_rumoca_both_success_models" + ); + assert_eq!( + payload["runtime_comparison"]["diagnostics"]["runtime_ratio_available"], + false + ); + assert_missing_wall_time_provenance_defaults(&payload); +} + #[test] fn quantify_trace_differences_skips_excluded_model_before_trace_loading() { let temp = tempfile::tempdir().expect("tempdir"); @@ -446,6 +743,120 @@ fn quantify_trace_differences_skips_excluded_model_before_trace_loading() { ); } +#[test] +fn quantify_rejects_invalid_state_contract_before_no_common_variable_skip() { + let (temp, paths, model_name) = trace_quantification_fixture("NoCommonVariables"); + write_pretty_json( + &paths.rumoca_trace_dir.join(format!("{model_name}.json")), + &invalid_state_contract_trace(&model_name, "x"), + ) + .expect("write rumoca trace"); + write_pretty_json( + &paths.omc_trace_dir.join(format!("{model_name}.json")), + &json!({ + "model_name": model_name, + "times": [0.0], + "names": ["y"], + "data": [[0.0]] + }), + ) + .expect("write omc trace"); + let results = trace_candidate_results(&model_name); + + let error = quantify_trace_differences(&paths, &results, &BTreeMap::new()) + .expect_err("invalid producer metadata must fail before numeric comparison skips"); + + assert!( + error + .to_string() + .contains("invalid state metadata contract") + ); + drop(temp); +} + +#[test] +fn quantify_rejects_invalid_state_contract_before_omc_load_skip() { + let (temp, paths, model_name) = trace_quantification_fixture("MalformedOmcTrace"); + write_pretty_json( + &paths.rumoca_trace_dir.join(format!("{model_name}.json")), + &invalid_state_contract_trace(&model_name, "x"), + ) + .expect("write rumoca trace"); + std::fs::write( + paths.omc_trace_dir.join(format!("{model_name}.json")), + "not json", + ) + .expect("write malformed omc trace"); + let results = trace_candidate_results(&model_name); + + let error = quantify_trace_differences(&paths, &results, &BTreeMap::new()) + .expect_err("invalid producer metadata must fail before OMC trace loading skips"); + + assert!( + error + .to_string() + .contains("invalid state metadata contract") + ); + drop(temp); +} + +fn trace_quantification_fixture(name: &str) -> (tempfile::TempDir, MslPaths, String) { + let temp = tempfile::tempdir().expect("tempdir"); + let results_dir = temp.path().join("results"); + let omc_trace_dir = results_dir.join("sim_traces/omc"); + let rumoca_trace_dir = results_dir.join("sim_traces/rumoca"); + std::fs::create_dir_all(&omc_trace_dir).expect("omc trace dir"); + std::fs::create_dir_all(&rumoca_trace_dir).expect("rumoca trace dir"); + let paths = MslPaths { + repo_root: temp.path().to_path_buf(), + msl_dir: temp.path().join("msl"), + results_dir: results_dir.clone(), + flat_dir: results_dir.join("omc_flat"), + work_dir: results_dir.join("omc_work"), + sim_work_dir: results_dir.join("omc_sim_work"), + omc_trace_dir, + rumoca_trace_dir, + }; + (temp, paths, format!("Modelica.Tests.{name}")) +} + +fn invalid_state_contract_trace(model_name: &str, variable: &str) -> Value { + json!({ + "model_name": model_name, + "n_states": 2, + "times": [0.0], + "names": [variable], + "data": [[0.0]], + "variable_meta": [{"name": variable, "role": "state"}] + }) +} + +fn trace_candidate_results(model_name: &str) -> BTreeMap { + BTreeMap::from([( + model_name.to_string(), + SimModelResult { + status: "success".to_string(), + error: None, + sim_system_seconds: None, + total_system_seconds: None, + omc_wall_seconds: None, + result_file: None, + trace_file: Some(format!("sim_traces/omc/{model_name}.json")), + trace_error: None, + rumoca_status: Some("sim_ok".to_string()), + rumoca_ic_status: Some("ic_ok".to_string()), + rumoca_ic_error: None, + rumoca_ic_seconds: None, + rumoca_sim_seconds: None, + rumoca_sim_build_seconds: None, + rumoca_sim_run_seconds: None, + rumoca_sim_wall_seconds: None, + rumoca_trace_file: Some(format!("sim_traces/rumoca/{model_name}.json")), + rumoca_trace_error: None, + }, + )]) +} + #[test] fn quantify_trace_differences_includes_error_status_model_with_existing_traces() { let temp = tempfile::tempdir().expect("tempdir"); @@ -468,10 +879,11 @@ fn quantify_trace_differences_includes_error_status_model_with_existing_traces() let model_name = "Modelica.Clocked.Examples.Elementary.RealSignals.TickBasedSine".to_string(); let trace = SimTrace { model_name: Some(model_name.clone()), + n_states: Some(0), times: vec![0.0, 0.5, 1.0], names: vec!["y".to_string()], data: vec![vec![Some(0.0), Some(1.0), Some(0.0)]], - variable_meta: None, + variable_meta: Some(Vec::new()), }; write_pretty_json(&omc_trace_dir.join(format!("{model_name}.json")), &trace) .expect("write omc trace"); @@ -515,6 +927,7 @@ fn quantify_trace_differences_includes_error_status_model_with_existing_traces() fn trace_output_summary_rolls_up_initial_condition_stats() { let rumoca = SimTrace { model_name: Some("M".to_string()), + n_states: None, times: vec![0.0, 0.5, 1.0], names: vec!["x".to_string(), "y".to_string()], data: vec![ @@ -525,6 +938,7 @@ fn trace_output_summary_rolls_up_initial_condition_stats() { }; let omc = SimTrace { model_name: Some("M".to_string()), + n_states: None, times: vec![0.0, 0.5, 1.0], names: vec!["x".to_string(), "y".to_string()], data: vec![ diff --git a/crates/rumoca-test-msl/src/msl_tools/plot_compare.rs b/crates/rumoca-test-msl/src/msl_tools/plot_compare.rs index c23208f06..9c53b55d8 100644 --- a/crates/rumoca-test-msl/src/msl_tools/plot_compare.rs +++ b/crates/rumoca-test-msl/src/msl_tools/plot_compare.rs @@ -258,6 +258,7 @@ fn trace_from_sim_result(model_name: &str, sim: &SimResult) -> SimTrace { }; SimTrace { model_name: Some(model_name.to_string()), + n_states: Some(sim.n_states), times: sim.times.clone(), names: sim.names.clone(), data, @@ -400,6 +401,7 @@ fn load_omc_csv_as_trace(model_name: &str, csv_path: &Path) -> Result Ok(SimTrace { model_name: Some(model_name.to_string()), + n_states: None, times, names, data, @@ -852,6 +854,7 @@ mod tests { fn trace(model: &str, times: Vec, names: Vec<&str>, data: Vec>) -> SimTrace { SimTrace { model_name: Some(model.to_string()), + n_states: None, times, names: names.into_iter().map(ToOwned::to_owned).collect(), data: data diff --git a/crates/rumoca-test-msl/src/runtime_measurement.rs b/crates/rumoca-test-msl/src/runtime_measurement.rs new file mode 100644 index 000000000..ec8ed1218 --- /dev/null +++ b/crates/rumoca-test-msl/src/runtime_measurement.rs @@ -0,0 +1,110 @@ +use serde::{Deserialize, Serialize}; + +/// Host load captured at one point in a wall-time measurement interval. +#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)] +pub struct HostLoadSnapshot { + pub one_minute: f64, + pub logical_cpus: usize, +} + +impl HostLoadSnapshot { + #[must_use] + pub fn normalized(self) -> f64 { + self.one_minute / self.logical_cpus as f64 + } +} + +/// Metadata needed to judge whether a wall-time comparison is trustworthy. +#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)] +pub struct WallTimeMeasurementProvenance { + pub omc_fresh_sample_count: usize, + pub omc_cached_sample_count: usize, + pub affinity_requested_worker_count: usize, + pub affinity_applied_worker_count: usize, + pub affinity_failed_worker_count: usize, + pub normalized_load_before: Option, + pub normalized_load_after: Option, + pub rumoca_workers_used: usize, + pub workers_used: usize, + pub omc_threads: usize, +} + +#[cfg(any(test, target_os = "linux", target_os = "macos"))] +fn parse_load_text(text: &str, logical_cpus: usize) -> Option { + if logical_cpus == 0 { + return None; + } + let one_minute = text + .trim() + .trim_start_matches('{') + .split_whitespace() + .next()? + .parse::() + .ok()?; + if !one_minute.is_finite() { + return None; + } + Some(HostLoadSnapshot { + one_minute, + logical_cpus, + }) +} + +/// Capture the platform one-minute load average without unsafe system calls. +#[must_use] +pub fn sample_host_load() -> Option { + #[cfg(any(target_os = "linux", target_os = "macos"))] + { + let logical_cpus = std::thread::available_parallelism().ok()?.get(); + #[cfg(target_os = "linux")] + let text = std::fs::read_to_string("/proc/loadavg").ok()?; + #[cfg(target_os = "macos")] + let text = { + let output = std::process::Command::new("sysctl") + .args(["-n", "vm.loadavg"]) + .output() + .ok()?; + if !output.status.success() { + return None; + } + String::from_utf8(output.stdout).ok()? + }; + parse_load_text(&text, logical_cpus) + } + #[cfg(not(any(target_os = "linux", target_os = "macos")))] + { + None + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn parses_linux_loadavg() { + let snapshot = parse_load_text("15.0 8.0 4.0 2/100 123", 10).unwrap(); + assert_eq!(snapshot.one_minute, 15.0); + assert_eq!(snapshot.normalized(), 1.5); + } + + #[test] + fn parses_macos_vm_loadavg() { + let snapshot = parse_load_text("{ 2.50 3.00 4.00 }", 10).unwrap(); + assert_eq!(snapshot.one_minute, 2.5); + assert_eq!(snapshot.normalized(), 0.25); + } + + #[test] + fn rejects_missing_or_non_finite_load() { + assert!(parse_load_text("unavailable", 10).is_none()); + assert!(parse_load_text("NaN 1 1", 10).is_none()); + assert!(parse_load_text("1 1 1", 0).is_none()); + } + + #[cfg(not(any(target_os = "linux", target_os = "macos")))] + #[test] + fn unsupported_platform_returns_no_host_load() { + assert!(sample_host_load().is_none()); + } +} diff --git a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_config.rs b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_config.rs index 32d554e24..72b99ec1b 100644 --- a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_config.rs +++ b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_config.rs @@ -101,6 +101,12 @@ pub(crate) struct MslParityConfig { /// Opt into the generated simulation-targets file when no explicit file or /// committed file applies. pub generated_sim_targets_file: Option, + /// Force regeneration of the OMC simulation reference cache. + pub force_omc_parity_refresh: Option, + /// OMC reference-generation worker count. + pub omc_parity_workers: Option, + /// Whole-stage OMC reference-generation timeout in seconds. + pub omc_sim_reference_batch_timeout_secs: Option, /// 1-based shard index for a sharded parity run (`--shard m/n` → `m`). The /// model set (already ordered slowest-first) is striped round-robin so this /// shard keeps every `shard_count`-th model starting at `shard_index - 1`. @@ -178,3 +184,28 @@ pub(crate) fn merge_shards_dir() -> Option { } Some(msl_workspace_root().join(path)) } + +#[test] +fn msl_parity_config_accepts_omc_and_shard_fields() { + let config: MslParityConfig = serde_json::from_str( + r#"{ + "force_omc_parity_refresh": true, + "omc_parity_workers": 6, + "omc_sim_reference_batch_timeout_secs": 900, + "shard_index": 2, + "shard_count": 4, + "merge_shards_dir": "target/msl/shards" + }"#, + ) + .expect("OMC and shard parity fields should deserialize together"); + + assert_eq!(config.force_omc_parity_refresh, Some(true)); + assert_eq!(config.omc_parity_workers, Some(6)); + assert_eq!(config.omc_sim_reference_batch_timeout_secs, Some(900)); + assert_eq!(config.shard_index, Some(2)); + assert_eq!(config.shard_count, Some(4)); + assert_eq!( + config.merge_shards_dir, + Some(PathBuf::from("target/msl/shards")) + ); +} diff --git a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_core.rs b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_core.rs index 22556dea2..9f00dfd43 100644 --- a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_core.rs +++ b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_core.rs @@ -915,6 +915,7 @@ struct InProcessWorkerRequest<'a> { fn run_compile_model_in_process_worker( worker: &mut Option, + scheduler_stats: &SchedulerStatsCollector, plan: InProcessWorkerRequest<'_>, ) -> (MslModelResult, bool) { let phase_timeout_secs = plan.budget_secs; @@ -929,7 +930,10 @@ fn run_compile_model_in_process_worker( plan.startup_timeout_secs, plan.cpu_core_id, ) { - Ok(spawned) => *worker = Some(spawned), + Ok(spawned) => { + scheduler_stats.record_worker_affinity(spawned.cpu_affinity_applied()); + *worker = Some(spawned); + } Err(error) => { return ( model_worker_failure_result( @@ -1234,16 +1238,12 @@ where requested_worker_count.min(available_core_count) }; let core_plan = cpu_core_plan(worker_count); - let pinned_workers = core_plan.iter().filter(|core| core.is_some()).count(); let scheduler_stats = SchedulerStatsCollector::new(); if worker_count < requested_worker_count { println!( " Model worker count capped at {worker_count}/{requested_worker_count} available logical CPU cores" ); } - if pinned_workers > 0 { - println!(" Model workers pinned to {pinned_workers}/{worker_count} logical CPU cores"); - } let startup_barrier = std::sync::Arc::new(std::sync::Barrier::new(worker_count)); for cpu_core_id in core_plan.iter().copied().take(worker_count) { let result_tx = result_tx.clone(); @@ -1279,7 +1279,6 @@ where requested_worker_threads: compile_threads, effective_worker_threads: effective_compile_threads, worker_count, - pinned_worker_count: pinned_workers, compile_memory_token_capacity_mb: memory_tokens .as_ref() .map(|tokens| tokens.capacity()), @@ -1335,7 +1334,10 @@ fn prepare_compiled_source_root( pub(super) fn run_msl_test(run_simulation: bool) -> MslSummary { let core_start = Instant::now(); - let mut timings = MslPhaseTimings::default(); + let mut timings = MslPhaseTimings { + host_load_before: rumoca_test_msl::runtime_measurement::sample_host_load(), + ..MslPhaseTimings::default() + }; let frontend_compile_start = Instant::now(); reset_compile_phase_timing_stats(); reset_flatten_phase_timing_stats(); diff --git a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_core/streaming_workers.rs b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_core/streaming_workers.rs index 0151703e2..fbcc39a23 100644 --- a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_core/streaming_workers.rs +++ b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_core/streaming_workers.rs @@ -24,6 +24,9 @@ struct SchedulerStatsInner { models_started: AtomicUsize, active_workers: AtomicUsize, max_active_workers: AtomicUsize, + affinity_requested_workers: AtomicUsize, + affinity_applied_workers: AtomicUsize, + affinity_failed_workers: AtomicUsize, memory_token_wait_nanos: std::sync::atomic::AtomicU64, active_model_wall_nanos: std::sync::atomic::AtomicU64, } @@ -33,7 +36,6 @@ pub(super) struct SchedulerTimingInputs { pub(super) requested_worker_threads: usize, pub(super) effective_worker_threads: usize, pub(super) worker_count: usize, - pub(super) pinned_worker_count: usize, pub(super) compile_memory_token_capacity_mb: Option, pub(super) compile_memory_model_cost_mb: Option, pub(super) elapsed_seconds: f64, @@ -44,6 +46,36 @@ struct ActiveModelGuard<'a> { started: Instant, } +pub(super) struct AffinityCounts { + pub(super) requested: usize, + pub(super) applied: usize, + pub(super) failed: usize, +} + +pub(super) fn affinity_counts(results: impl IntoIterator>) -> AffinityCounts { + results.into_iter().fold( + AffinityCounts { + requested: 0, + applied: 0, + failed: 0, + }, + |mut counts, result| { + match result { + Some(true) => { + counts.requested += 1; + counts.applied += 1; + } + Some(false) => { + counts.requested += 1; + counts.failed += 1; + } + None => {} + } + counts + }, + ) +} + impl SchedulerStatsCollector { pub(super) fn new() -> Self { Self { @@ -51,6 +83,9 @@ impl SchedulerStatsCollector { models_started: AtomicUsize::new(0), active_workers: AtomicUsize::new(0), max_active_workers: AtomicUsize::new(0), + affinity_requested_workers: AtomicUsize::new(0), + affinity_applied_workers: AtomicUsize::new(0), + affinity_failed_workers: AtomicUsize::new(0), memory_token_wait_nanos: std::sync::atomic::AtomicU64::new(0), active_model_wall_nanos: std::sync::atomic::AtomicU64::new(0), }), @@ -68,6 +103,19 @@ impl SchedulerStatsCollector { ); } + pub(super) fn record_worker_affinity(&self, applied: Option) { + let counts = affinity_counts([applied]); + self.inner + .affinity_requested_workers + .fetch_add(counts.requested, Ordering::Relaxed); + self.inner + .affinity_applied_workers + .fetch_add(counts.applied, Ordering::Relaxed); + self.inner + .affinity_failed_workers + .fetch_add(counts.failed, Ordering::Relaxed); + } + fn enter_active_model(&self) -> ActiveModelGuard<'_> { let active = self.inner.active_workers.fetch_add(1, Ordering::Relaxed) + 1; update_atomic_max(&self.inner.max_active_workers, active); @@ -81,12 +129,23 @@ impl SchedulerStatsCollector { let worker_slot_wall_seconds = inputs.elapsed_seconds * inputs.worker_count as f64; let active_model_wall_seconds = nanos_to_seconds(self.inner.active_model_wall_nanos.load(Ordering::Relaxed)); + let affinity_applied_worker_count = + self.inner.affinity_applied_workers.load(Ordering::Relaxed); MslSchedulerTimings { selected_model_count: inputs.selected_model_count, requested_worker_threads: inputs.requested_worker_threads, effective_worker_threads: inputs.effective_worker_threads, worker_count: inputs.worker_count, - pinned_worker_count: inputs.pinned_worker_count, + affinity_requested_worker_count: self + .inner + .affinity_requested_workers + .load(Ordering::Relaxed), + affinity_applied_worker_count, + affinity_failed_worker_count: self + .inner + .affinity_failed_workers + .load(Ordering::Relaxed), + pinned_worker_count: affinity_applied_worker_count, cpu_token_capacity: 0, compile_memory_token_capacity_mb: inputs.compile_memory_token_capacity_mb, compile_memory_model_cost_mb: inputs.compile_memory_model_cost_mb, @@ -165,7 +224,12 @@ pub(super) fn run_model_worker_queue(queue: ModelWorkerQueue<'_>) { startup_timeout_secs, queue.cpu_core_id, ) { - Ok(spawned) => worker = Some(spawned), + Ok(spawned) => { + queue + .scheduler_stats + .record_worker_affinity(spawned.cpu_affinity_applied()); + worker = Some(spawned); + } Err(error) => { let entry = model_worker_failure_result( name, @@ -183,6 +247,7 @@ pub(super) fn run_model_worker_queue(queue: ModelWorkerQueue<'_>) { .is_some_and(|names| names.contains(name)); let (entry, _keep_worker) = run_compile_model_in_process_worker( &mut worker, + &queue.scheduler_stats, InProcessWorkerRequest { source_root_path: queue.source_root_path, cpu_core_id: queue.cpu_core_id, diff --git a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_core/tests.rs b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_core/tests.rs index 321239c35..62021cc3b 100644 --- a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_core/tests.rs +++ b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_core/tests.rs @@ -52,6 +52,14 @@ fn empty_compilation_result() -> CompilationResult { } } +#[test] +fn affinity_counts_distinguish_requested_success_and_failure() { + let counts = affinity_counts([Some(true), Some(false), None]); + assert_eq!(counts.requested, 2); + assert_eq!(counts.applied, 1); + assert_eq!(counts.failed, 1); +} + #[test] fn compile_chunk_progress_loop_exits_promptly_after_flag_clears() { let compile_in_flight = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(true)); diff --git a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_merge.rs b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_merge.rs index 590ef19c2..3b7b537d0 100644 --- a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_merge.rs +++ b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_merge.rs @@ -408,8 +408,8 @@ fn sum_rumoca_runtime_seconds(models: &Map) -> f64 { .sum() } -fn build_runtime_comparison(models: &Map) -> Value { - json!({ +fn build_runtime_comparison(models: &Map, omc_payloads: &[Value]) -> Value { + let mut comparison = json!({ "ratio_definition": "omc_over_rumoca_higher_is_better (simulation/wall runtime; this is the SIM-time comparison, distinct from the compile-speed `speedup` in msl_speed_comparison.json)", "ratio_metric_system": "omc_timeSimulation_over_rumoca_sim_seconds", "ratio_metric_wall": "omc_external_wall_over_rumoca_external_wall", @@ -427,7 +427,59 @@ fn build_runtime_comparison(models: &Map) -> Value { "wall_ratio_all_positive": runtime_ratio_bucket(models, true, false), "wall_ratio_both_success": runtime_ratio_bucket(models, true, true), }, - }) + }); + if let Some(provenance) = merge_wall_time_provenance(omc_payloads) { + comparison["wall_time_provenance"] = provenance; + } + comparison +} + +fn merge_wall_time_provenance(omc_payloads: &[Value]) -> Option { + let provenances = omc_payloads + .iter() + .map(|payload| payload.pointer("/runtime_comparison/wall_time_provenance")) + .collect::>>()?; + let checked_sum = |key: &str| { + provenances.iter().try_fold(0_usize, |total, value| { + total.checked_add(json_usize(value, &[key])?) + }) + }; + let same = |key: &str| { + let values = provenances + .iter() + .map(|value| json_usize(value, &[key])) + .collect::>>()?; + values + .first() + .copied() + .filter(|first| values.iter().all(|value| value == first)) + }; + let max_finite = |key: &str| { + provenances + .iter() + .try_fold(f64::NEG_INFINITY, |maximum, value| { + let current = json_f64(value, &[key])?; + current.is_finite().then_some(maximum.max(current)) + }) + .filter(|value| value.is_finite()) + }; + let mut merged = json!({ + "omc_fresh_sample_count": checked_sum("omc_fresh_sample_count")?, + "omc_cached_sample_count": checked_sum("omc_cached_sample_count")?, + "affinity_requested_worker_count": checked_sum("affinity_requested_worker_count")?, + "affinity_applied_worker_count": checked_sum("affinity_applied_worker_count")?, + "affinity_failed_worker_count": checked_sum("affinity_failed_worker_count")?, + "rumoca_workers_used": checked_sum("rumoca_workers_used")?, + "normalized_load_before": max_finite("normalized_load_before"), + "normalized_load_after": max_finite("normalized_load_after"), + }); + if let Some(value) = same("workers_used") { + merged["workers_used"] = json!(value); + } + if let Some(value) = same("omc_threads") { + merged["omc_threads"] = json!(value); + } + Some(merged) } fn merge_initial_condition_summary(trace_values: &[&Value]) -> Result { @@ -980,11 +1032,7 @@ fn merge_timing_payload(omc_payloads: &[Value]) -> Value { .sum::(); root.insert(key.to_string(), json!(values)); } - let workers = omc_payloads - .iter() - .filter_map(|payload| json_usize(payload, &["timing", "workers_used"])) - .sum::(); - if workers > 0 { + if let Some(workers) = optional_same_usize(omc_payloads, &["timing", "workers_used"]) { root.insert("workers_used".to_string(), json!(workers)); } if let Some(omc_threads) = optional_same_usize(omc_payloads, &["timing", "omc_threads"]) { @@ -1060,7 +1108,7 @@ fn merge_omc_reference_payloads( root.insert("timing".to_string(), merge_timing_payload(omc_payloads)); root.insert( "runtime_comparison".to_string(), - build_runtime_comparison(&omc_models), + build_runtime_comparison(&omc_models, omc_payloads), ); root.insert("trace_comparison".to_string(), trace_summary.clone()); root.insert( @@ -1302,7 +1350,20 @@ fn shard_omc_reference_fixture(model: &str) -> Value { "omc_version": "OpenModelica 1.26.1", "total_models": 1, "timing": shard_timing_fixture(), - "runtime_comparison": { "ratio_stats": { + "runtime_comparison": { + "wall_time_provenance": { + "omc_fresh_sample_count": 1, + "omc_cached_sample_count": 0, + "affinity_requested_worker_count": 7, + "affinity_applied_worker_count": 7, + "affinity_failed_worker_count": 0, + "rumoca_workers_used": 7, + "normalized_load_before": 0.2, + "normalized_load_after": 0.3, + "workers_used": 3, + "omc_threads": 1 + }, + "ratio_stats": { "system_ratio_both_success": shard_ratio_stats_fixture(), "wall_ratio_both_success": shard_ratio_stats_fixture() }}, @@ -1368,10 +1429,12 @@ fn write_shard_fixture(dir: &Path, shard: usize, model: &str) { &shard_dir.join("msl_results.json"), &serde_json::to_value(shard_summary(model)).expect("serialize summary"), ); - write_json_fixture( - &shard_dir.join(SHARD_OMC_REFERENCE_FILE), - &shard_omc_reference_fixture(model), - ); + let mut omc = shard_omc_reference_fixture(model); + if shard == 2 { + omc["runtime_comparison"]["wall_time_provenance"]["normalized_load_before"] = json!(0.4); + omc["runtime_comparison"]["wall_time_provenance"]["normalized_load_after"] = json!(0.5); + } + write_json_fixture(&shard_dir.join(SHARD_OMC_REFERENCE_FILE), &omc); write_json_fixture( &shard_dir.join(SHARD_TRACE_COMPARISON_FILE), &shard_trace_comparison_fixture(model), @@ -1413,8 +1476,56 @@ fn merge_shard_parity_artifacts_writes_full_omc_and_trace_inputs() { Some(2) ); assert_eq!(json_usize(&trace, &["models_compared"]), Some(2)); + let provenance = &omc["runtime_comparison"]["wall_time_provenance"]; + assert_eq!(provenance["omc_fresh_sample_count"], 2); + assert_eq!(provenance["rumoca_workers_used"], 14); + assert_eq!(provenance["affinity_requested_worker_count"], 14); + assert_eq!(provenance["workers_used"], 3); + assert_eq!(provenance["normalized_load_before"], 0.4); + assert_eq!(provenance["normalized_load_after"], 0.5); assert_eq!( trace.get("models").and_then(Value::as_object).map(Map::len), Some(2) ); } + +#[test] +fn merge_shard_provenance_omits_untrustworthy_fields() { + let mut first = shard_omc_reference_fixture("A"); + let mut second = shard_omc_reference_fixture("B"); + second["runtime_comparison"]["wall_time_provenance"]["workers_used"] = json!(4); + second["runtime_comparison"]["wall_time_provenance"]["omc_threads"] = json!(2); + second["runtime_comparison"]["wall_time_provenance"]["normalized_load_before"] = Value::Null; + let merged = merge_wall_time_provenance(&[first.clone(), second]).expect("aggregate exists"); + assert!(merged.get("workers_used").is_none()); + assert!(merged.get("omc_threads").is_none()); + assert!(merged["normalized_load_before"].is_null()); + + first["runtime_comparison"] + .as_object_mut() + .unwrap() + .remove("wall_time_provenance"); + assert!(merge_wall_time_provenance(&[first]).is_none()); +} + +#[test] +fn merge_shard_provenance_rejects_missing_cached_sample_count() { + let first = shard_omc_reference_fixture("A"); + let mut second = shard_omc_reference_fixture("B"); + second["runtime_comparison"]["wall_time_provenance"] + .as_object_mut() + .unwrap() + .remove("omc_cached_sample_count"); + assert!(merge_wall_time_provenance(&[first, second]).is_none()); +} + +#[test] +fn merge_shard_provenance_rejects_missing_affinity_failed_count() { + let first = shard_omc_reference_fixture("A"); + let mut second = shard_omc_reference_fixture("B"); + second["runtime_comparison"]["wall_time_provenance"] + .as_object_mut() + .unwrap() + .remove("affinity_failed_worker_count"); + assert!(merge_wall_time_provenance(&[first, second]).is_none()); +} diff --git a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate.rs b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate.rs index 47529b038..25c245c2c 100644 --- a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate.rs +++ b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate.rs @@ -5,8 +5,10 @@ mod cache; mod status; #[cfg(test)] mod tests; +mod wall_time; use cache::*; use status::*; +use wall_time::*; // ============================================================================= // MSL quality gate (compile/balance strict + simulation tolerant gate) @@ -38,6 +40,8 @@ pub(super) const MSL_STAGE_COUNT_ALLOWED_DROP_MIN_DENOMINATOR: usize = 100; pub(super) const TRACE_MODELS_COMPARED_ALLOWED_DROP: usize = 2; /// Allowed relative drop in runtime speedup median (omc/rumoca) before failing. pub(super) const RUNTIME_RATIO_MEDIAN_REL_TOLERANCE: f64 = 0.35; +/// Maximum normalized one-minute host load for a blocking wall-time comparison. +pub(super) const WALL_TIME_NORMALIZED_LOAD_MAX: f64 = 1.5; /// OMC per-model timeout budget for simulation reference generation. OMC's /// `simulate()` generates C code and invokes gcc per model, which on shared CI /// runners far exceeds rumoca's warm per-model budget; giving OMC the same tiny @@ -45,11 +49,26 @@ pub(super) const RUNTIME_RATIO_MEDIAN_REL_TOLERANCE: f64 = 0.35; /// rumoca, so a larger budget only lets more OMC models complete (the measured /// `timeSimulation` it reports is unchanged, so the timing comparison stays fair). pub(super) const OMC_SIM_REFERENCE_BATCH_TIMEOUT_SECONDS: u64 = 120; +/// Whole-stage watchdog for generating the OMC simulation reference. +/// +/// This covers the complete batch of Rumoca-sim-ok models, not a single OMC +/// `simulate()` call. Keep it aligned with the outer parity-stage budget so +/// slower local or shared runners can finish a progressing reference batch +/// without weakening any model-level timeout or parity quality gate. +pub(super) const OMC_SIM_REFERENCE_STAGE_TIMEOUT_SECONDS: u64 = 20_400; /// Force low-impact OpenMP/BLAS threading in OMC child processes. pub(super) const OMC_PARITY_THREADS_DEFAULT: usize = 1; pub(super) const MSL_QUALITY_GATE_VERSION: u32 = 1; pub(super) const MSL_QUALITY_RUN_SCOPE_FULL: &str = "full"; pub(super) const MSL_QUALITY_RUN_SCOPE_PARTIAL: &str = "partial"; +/// Default OMC worker cap for parity reference generation. +/// +/// OMC is often accessed through a Docker-backed wrapper on macOS. Running one +/// OMC process per local CPU can make otherwise quick Clocked examples hit the +/// per-model timeout and collapse trace coverage. Keep this conservative by +/// default; `cargo xtask verify msl-parity --omc-parity-workers` can raise or +/// lower it for a documented one-off run. +pub(super) const OMC_PARITY_WORKERS_DEFAULT_MAX: usize = 2; pub(super) const MSL_QUALITY_BASELINE_FILE_REL: &str = "tests/msl_tests/msl_quality_baseline.json"; pub(super) const MSL_QUALITY_CURRENT_FILE_REL: &str = "msl_quality_current.json"; pub(super) const MSL_SIM_TARGETS_FILE_REL: &str = "msl_simulation_targets.json"; @@ -233,6 +252,7 @@ pub(super) struct MslParityGateInput { omc_version: Option, runtime_context: Option, runtime_ratio_stats: Option, + wall_time_provenance: Option, trace_accuracy_stats: Option, omc_assertion_failure_models: usize, omc_assertion_failure_examples: Vec, @@ -434,6 +454,15 @@ fn parse_runtime_ratio_stats(payload: &serde_json::Value) -> Option Option { + serde_json::from_value( + payload + .pointer("/runtime_comparison/wall_time_provenance")? + .clone(), + ) + .ok() +} + fn parse_trace_bounded_normalized_l1( trace: &serde_json::Value, models_compared: usize, @@ -606,6 +635,7 @@ pub(super) fn load_msl_parity_gate_input(path: &Path) -> io::Result Vec io::Result { +) -> io::Result> { let path = omc_simulation_reference_path(); + load_msl_parity_gate_input_optional_from_path(&path, expected_sim_target_models) +} + +fn load_required_msl_parity_gate_input_from_path( + path: &Path, + expected_sim_target_models: usize, +) -> io::Result { if !path.is_file() { return Err(io::Error::other(format!( - "missing required OMC parity file '{}'", + "required OMC parity reference '{}' is missing; generate the current full-run OMC runtime and trace reference before accepting the MSL quality gate", path.display() ))); } - let parity = load_msl_parity_gate_input(&path)?; - validate_parity_total_models(&path, &parity, expected_sim_target_models)?; + let parity = load_msl_parity_gate_input(path)?; + validate_parity_total_models(path, &parity, expected_sim_target_models)?; + validate_required_msl_parity_gate_input(path, parity) +} + +fn load_current_required_msl_parity_gate_input( + expected_sim_target_models: usize, +) -> io::Result { + load_required_msl_parity_gate_input_from_path( + &omc_simulation_reference_path(), + expected_sim_target_models, + ) +} + +fn load_msl_parity_gate_input_optional_from_path( + path: &Path, + expected_sim_target_models: usize, +) -> io::Result> { + if !path.is_file() { + return Ok(None); + } + let parity = load_msl_parity_gate_input(path)?; + let Some(parity_total_models) = parity.total_models else { + return Err(io::Error::other(format!( + "OMC parity file '{}' is missing total_models/models metadata", + path.display() + ))); + }; + if parity_total_models != expected_sim_target_models { + return Ok(None); + } + validate_msl_parity_metadata(path, &parity)?; + if !parity_has_required_runtime_and_trace_metrics(&parity) { + return Ok(None); + } + validate_required_msl_parity_gate_input(path, parity).map(Some) +} + +fn parity_has_required_runtime_and_trace_metrics(parity: &MslParityGateInput) -> bool { + let Some(runtime_stats) = parity.runtime_ratio_stats.as_ref() else { + return false; + }; + if runtime_stats.system_ratio_both_success.sample_count == 0 + || runtime_stats.wall_ratio_both_success.sample_count == 0 + { + return false; + } + let Some(trace_stats) = parity.trace_accuracy_stats.as_ref() else { + return false; + }; + trace_stats.models_compared > 0 && trace_model_bucket_percentages(trace_stats).is_some() +} + +fn validate_msl_parity_metadata(path: &Path, parity: &MslParityGateInput) -> io::Result<()> { if parity.omc_version.is_none() { return Err(io::Error::other(format!( "OMC parity file '{}' is missing omc_version metadata; regenerate OMC simulation reference", @@ -688,6 +777,36 @@ pub(super) fn load_current_msl_parity_gate_input_required( examples ))); } + Ok(()) +} + +fn validate_parity_total_models( + path: &Path, + parity: &MslParityGateInput, + expected_sim_target_models: usize, +) -> io::Result<()> { + let parity_total_models = parity.total_models.ok_or_else(|| { + io::Error::other(format!( + "OMC parity file '{}' is missing total_models/models metadata", + path.display() + )) + })?; + if parity_total_models != expected_sim_target_models { + return Err(io::Error::other(format!( + "OMC parity file '{}' is stale: total_models={} but current sim_target_models={}; regenerate OMC simulation reference for the active target set", + path.display(), + parity_total_models, + expected_sim_target_models + ))); + } + Ok(()) +} + +fn validate_required_msl_parity_gate_input( + path: &Path, + parity: MslParityGateInput, +) -> io::Result { + validate_msl_parity_metadata(path, &parity)?; let runtime_stats = parity.runtime_ratio_stats.as_ref().ok_or_else(|| { io::Error::other(format!( "OMC parity file '{}' is missing runtime_ratio_stats", @@ -712,7 +831,7 @@ pub(super) fn load_current_msl_parity_gate_input_required( })?; if trace_stats.models_compared == 0 { return Err(io::Error::other(format!( - "OMC parity file '{}' has models_compared=0 (no OMC/Rumoca traces were compared)", + "OMC parity file '{}' is missing comparable trace metrics: models_compared=0 (no OMC/Rumoca traces were compared)", path.display() ))); } @@ -725,38 +844,6 @@ pub(super) fn load_current_msl_parity_gate_input_required( Ok(parity) } -pub(super) fn load_current_msl_parity_gate_input_optional( - expected_sim_target_models: usize, -) -> io::Result> { - let path = omc_simulation_reference_path(); - if !path.is_file() { - return Ok(None); - } - load_current_msl_parity_gate_input_required(expected_sim_target_models).map(Some) -} - -fn validate_parity_total_models( - path: &Path, - parity: &MslParityGateInput, - expected_sim_target_models: usize, -) -> io::Result<()> { - let parity_total_models = parity.total_models.ok_or_else(|| { - io::Error::other(format!( - "OMC parity file '{}' is missing total_models/models metadata", - path.display() - )) - })?; - if parity_total_models != expected_sim_target_models { - return Err(io::Error::other(format!( - "OMC parity file '{}' is stale: total_models={} but current sim_target_models={}; regenerate OMC simulation reference for the active target set", - path.display(), - parity_total_models, - expected_sim_target_models - ))); - } - Ok(()) -} - pub(super) fn resolve_msl_tools_exe_inner() -> Result { for env_key in [ "CARGO_BIN_EXE_rumoca-msl-tools", @@ -852,6 +939,7 @@ struct ParityStepContext { omc_version: String, workers: usize, omc_threads: usize, + sim_batch_timeout_seconds: u64, } fn run_simulation_parity_reference_command( @@ -872,7 +960,7 @@ fn run_simulation_parity_reference_command( "--rumoca-sim-ok-only".to_string(), "--use-experiment-stop-time".to_string(), "--model-timeout-seconds".to_string(), - OMC_SIM_REFERENCE_BATCH_TIMEOUT_SECONDS.to_string(), + context.sim_batch_timeout_seconds.to_string(), "--workers".to_string(), context.workers.to_string(), "--omc-threads".to_string(), @@ -894,8 +982,15 @@ fn ensure_simulation_parity_reference( sim_targets_path: &Path, sim_targets: &[String], ) -> io::Result<()> { - let _sim_ref_watchdog = StageAbortWatchdog::new("parity_simulation_reference", 3600); - let sim_policy = current_simulation_parity_cache_policy(); + let _sim_ref_watchdog = StageAbortWatchdog::new( + "parity_simulation_reference", + OMC_SIM_REFERENCE_STAGE_TIMEOUT_SECONDS, + ); + let sim_policy = current_simulation_parity_cache_policy( + context.workers, + context.omc_threads, + context.sim_batch_timeout_seconds, + ); let omc_simulation_reference = omc_simulation_reference_path(); let sim_cache_key = simulation_parity_cache_key( sim_targets, @@ -932,7 +1027,14 @@ fn ensure_simulation_parity_reference( &context.omc_version, sim_policy, )? && simulation_parity_cache_has_required_metrics(&omc_simulation_reference)?; - if force_refresh || !canonical_cache_matches { + let canonical_cache_can_resume = simulation_parity_cache_can_resume( + &omc_simulation_reference, + sim_targets, + &summary.msl_version, + &context.omc_version, + sim_policy, + )?; + if force_refresh || (!canonical_cache_matches && !canonical_cache_can_resume) { println!( "MSL parity cache miss/incomplete for simulation reference; regenerating {}", omc_simulation_reference.display() @@ -940,7 +1042,7 @@ fn ensure_simulation_parity_reference( run_simulation_parity_reference_command(context, sim_targets_path, false)?; } else { println!( - "MSL parity cache hit: reusing {} (refreshing Rumoca trace comparison via --resume)", + "MSL parity cache hit/resume: reusing {} (refreshing Rumoca trace comparison via --resume)", omc_simulation_reference.display() ); run_simulation_parity_reference_command(context, sim_targets_path, true)?; @@ -950,32 +1052,32 @@ fn ensure_simulation_parity_reference( } pub(super) fn ensure_required_msl_parity_references(summary: &MslSummary) -> io::Result<()> { - if summary.sim_attempted == 0 { + let focused_or_partial = !requires_msl_parity_artifacts(); + if !full_parity_is_required(focused_or_partial, summary.sim_attempted) { + if focused_or_partial { + println!( + "MSL parity stage: skipped because this is a focused/partial run; no full quality-gate parity claim is made." + ); + } return Ok(()); } let stage_start = Instant::now(); let force_refresh = force_omc_parity_refresh_enabled(); let (sim_targets_path, sim_targets) = load_sim_parity_targets()?; - let omc_version = match current_omc_version() { - Ok(version) => version, - Err(error) => { - println!( - "MSL parity stage: OMC unavailable; skipping parity reference generation ({error})" - ); - return Ok(()); - } - }; + let omc_version = require_omc_version_for_full_parity(current_omc_version())?; let context = ParityStepContext { tools_exe: resolve_msl_tools_exe()?, omc_version, workers: omc_parity_workers(), omc_threads: omc_parity_threads(), + sim_batch_timeout_seconds: omc_sim_reference_batch_timeout_seconds(), }; println!( - "MSL parity targets: simulation={} (workers={})", + "MSL parity targets: simulation={} (workers={}, sim_timeout={}s)", sim_targets.len(), - context.workers + context.workers, + context.sim_batch_timeout_seconds ); // The OMC reference comes solely from the persistent-zmq simulation pass, @@ -995,7 +1097,8 @@ pub(super) fn ensure_required_msl_parity_references(summary: &MslSummary) -> io: sim_ref_start.elapsed().as_secs_f64() ); - let _ = load_current_msl_parity_gate_input_required(sim_targets.len())?; + load_current_required_msl_parity_gate_input(sim_targets.len())?; + println!("MSL parity reference includes required runtime and trace comparison metrics."); println!( "MSL parity total step time: {:.2}s", stage_start.elapsed().as_secs_f64() @@ -1003,6 +1106,18 @@ pub(super) fn ensure_required_msl_parity_references(summary: &MslSummary) -> io: Ok(()) } +fn full_parity_is_required(focused_or_partial: bool, sim_attempted: usize) -> bool { + !focused_or_partial && sim_attempted > 0 +} + +fn require_omc_version_for_full_parity(result: io::Result) -> io::Result { + result.map_err(|error| { + io::Error::other(format!( + "required OMC prerequisite is unavailable for the full MSL parity run: {error}" + )) + }) +} + pub(super) fn current_omc_parity_workers() -> usize { omc_parity_workers() } @@ -1074,6 +1189,7 @@ pub(super) fn current_msl_quality_baseline( fn current_msl_quality_snapshot_json( summary: &MslSummary, parity_input: Option<&MslParityGateInput>, + promoted_baseline: Option<&MslQualityBaseline>, partial: bool, ) -> io::Result { let baseline = current_msl_quality_baseline(summary, parity_input); @@ -1138,6 +1254,31 @@ fn current_msl_quality_snapshot_json( )) })?, ); + root.insert( + "wall_time_provenance".to_string(), + parity_input + .and_then(|parity| parity.wall_time_provenance.as_ref()) + .map_or(serde_json::Value::Null, |provenance| { + serde_json::json!(provenance) + }), + ); + let wall_decision = promoted_baseline.map_or_else( + || WallTimeStatusContent { + status: "ADVISORY", + trusted: false, + reasons: vec!["runtime baseline missing".to_string()], + observed_median: parity_input + .and_then(|parity| parity.runtime_ratio_stats.as_ref()) + .map(|stats| stats.wall_ratio_both_success.median), + baseline_median: None, + floor: None, + }, + |baseline| wall_time_status_content(baseline, parity_input), + ); + root.insert( + "runtime_wall_decision".to_string(), + serde_json::json!(wall_decision), + ); if partial { root.insert("partial".to_string(), serde_json::Value::Bool(true)); } @@ -1159,9 +1300,11 @@ pub(super) fn write_current_msl_quality_snapshot(summary: &MslSummary) -> io::Re } let parity_input = load_current_msl_parity_gate_input_optional(summary.sim_target_models.len())?; + let promoted_baseline = load_msl_quality_baseline(&msl_quality_baseline_path()).ok(); let snapshot = current_msl_quality_snapshot_json( summary, parity_input.as_ref(), + promoted_baseline.as_ref(), should_skip_msl_quality_gate(), )?; let baseline_path = msl_quality_current_path(); @@ -1536,7 +1679,9 @@ pub(super) fn push_runtime_ratio_regression_reasons( let allowed_wall_median = baseline_runtime.wall_ratio_both_success.median * (1.0 - RUNTIME_RATIO_MEDIAN_REL_TOLERANCE); - if current_runtime.wall_ratio_both_success.median + SIM_RATE_GATE_EPSILON < allowed_wall_median + if wall_time_trust_decision(baseline, parity_input).trusted + && current_runtime.wall_ratio_both_success.median + SIM_RATE_GATE_EPSILON + < allowed_wall_median { reasons.push(format!( "runtime wall speedup median regressed: current={:.6e} < floor={:.6e} (baseline={:.6e}, tolerance={:.1}%)", @@ -1682,11 +1827,12 @@ pub(super) fn msl_quality_gate_failure_message( } pub(super) fn enforce_msl_quality_gate(summary: &MslSummary) -> io::Result<()> { - if require_selected_targets_success() { + let focused_or_partial = should_skip_msl_quality_gate(); + if require_selected_targets_success() && focused_or_partial { return enforce_all_selected_targets_succeeded(summary); } if summary.sim_attempted == 0 { - if should_skip_msl_quality_gate() { + if focused_or_partial { println!("MSL quality gate: skipped for compile/balance-only run."); return Ok(()); } @@ -1695,7 +1841,7 @@ pub(super) fn enforce_msl_quality_gate(summary: &MslSummary) -> io::Result<()> { summary.sim_target_models.len() ))); } - if should_skip_msl_quality_gate() { + if focused_or_partial { println!( "MSL quality gate: skipped for focused/non-baseline run (committed target scope, explicit target file, subset, or partial sim set)." ); @@ -1708,17 +1854,16 @@ pub(super) fn enforce_msl_quality_gate(summary: &MslSummary) -> io::Result<()> { let baseline_path = msl_quality_baseline_path(); let baseline = load_msl_quality_baseline(&baseline_path)?; let parity_input = - load_current_msl_parity_gate_input_optional(summary.sim_target_models.len())?; - let gate_failure = - msl_quality_gate_failure_message(gate_input, &baseline, parity_input.as_ref()); + load_current_required_msl_parity_gate_input(summary.sim_target_models.len())?; + let gate_failure = msl_quality_gate_failure_message(gate_input, &baseline, Some(&parity_input)); if let Some(message) = gate_failure { panic!("MSL quality gate: {message}."); } print_compile_and_sim_gate_pass(gate_input, &baseline); - print_trace_gate_status(&baseline, parity_input.as_ref()); - print_runtime_ratio_status(&baseline, parity_input.as_ref()); + print_trace_gate_status(&baseline, Some(&parity_input)); + print_runtime_ratio_status(&baseline, Some(&parity_input)); println!("MSL quality baseline source: {}", baseline_path.display()); Ok(()) @@ -1757,10 +1902,10 @@ pub(super) fn selected_target_failures(summary: &MslSummary) -> Vec { .filter(|result| target_set.contains(result.model_name.as_str())) .filter_map(|result| { seen_targets.insert(result.model_name.as_str()); - if result.sim_status.as_deref() == Some("sim_ok") { + let status = selected_target_result_status(summary, result); + if status == "ok" { return None; } - let status = result.sim_status.as_deref().unwrap_or("not-simulated"); Some(format!("{} ({status})", result.model_name)) }) .collect(); @@ -1773,12 +1918,47 @@ pub(super) fn selected_target_failures(summary: &MslSummary) -> Vec { failures } +fn selected_target_result_status(summary: &MslSummary, result: &MslModelResult) -> String { + if summary.sim_attempted > 0 { + return match result.sim_status.as_deref() { + Some("sim_ok") => "ok".to_string(), + Some(status) => status.to_string(), + None => "not-simulated".to_string(), + }; + } + if result.phase_reached != "Success" { + return match result.phase_reached.as_str() { + "" => "missing-phase", + "Resolve" => "resolve-failed", + "Instantiate" => "instantiate-failed", + "Typecheck" => "typecheck-failed", + "Flatten" => "flatten-failed", + "ToDae" => "todae-failed", + "NeedsInner" => "needs-inner", + "NonSim" => "non-sim", + _ => "phase-failed", + } + .to_string(); + } + if result.is_balanced == Some(false) { + return "unbalanced".to_string(); + } + if result.initial_balance_ok == Some(false) { + return "initial-unbalanced".to_string(); + } + "ok".to_string() +} + +pub(super) fn requires_msl_parity_artifacts() -> bool { + msl_target_scope() == MslTargetScope::RootExamples + && sim_targets_file_override().is_none() + && sim_subset_patterns().is_empty() + && sim_subset_limit().is_none() + && sim_set_mode() == SimSetMode::Full +} + pub(super) fn should_skip_msl_quality_gate() -> bool { - msl_target_scope() != MslTargetScope::RootExamples - || sim_targets_file_override().is_some() - || !sim_subset_patterns().is_empty() - || sim_subset_limit().is_some() - || sim_set_mode() != SimSetMode::Full + !requires_msl_parity_artifacts() // A shard sees only its stripe of the model set, so the aggregate // baseline ratchet + sim-ok floor are meaningless here; the fan-in job // runs the gate once on the merged results. Also stamps the snapshot diff --git a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/cache.rs b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/cache.rs index f2df0001a..ee55f11e5 100644 --- a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/cache.rs +++ b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/cache.rs @@ -59,6 +59,8 @@ pub(super) fn parity_target_set_cache_key( #[derive(Debug, Clone, Copy, PartialEq)] pub(super) struct SimulationParityCachePolicy { pub(super) batch_timeout_seconds: u64, + pub(super) workers: usize, + pub(super) omc_threads: usize, pub(super) use_experiment_stop_time: bool, pub(super) stop_time_override: Option, } @@ -68,10 +70,23 @@ pub(super) fn simulation_stop_time_override() -> Option { None } -pub(super) fn current_simulation_parity_cache_policy() -> SimulationParityCachePolicy { +pub(super) fn omc_sim_reference_batch_timeout_seconds() -> u64 { + parity_config() + .omc_sim_reference_batch_timeout_secs + .filter(|value| *value > 0) + .unwrap_or(OMC_SIM_REFERENCE_BATCH_TIMEOUT_SECONDS) +} + +pub(super) fn current_simulation_parity_cache_policy( + workers: usize, + omc_threads: usize, + batch_timeout_seconds: u64, +) -> SimulationParityCachePolicy { let stop_time_override = simulation_stop_time_override(); SimulationParityCachePolicy { - batch_timeout_seconds: OMC_SIM_REFERENCE_BATCH_TIMEOUT_SECONDS, + batch_timeout_seconds, + workers, + omc_threads, use_experiment_stop_time: stop_time_override.is_none(), stop_time_override, } @@ -98,6 +113,10 @@ pub(super) fn simulation_parity_cache_key( hash = fnv1a64_update(hash, &[0xfc]); hash = fnv1a64_update(hash, policy.batch_timeout_seconds.to_string().as_bytes()); hash = fnv1a64_update(hash, &[0xfb]); + hash = fnv1a64_update(hash, policy.workers.to_string().as_bytes()); + hash = fnv1a64_update(hash, &[0xf9]); + hash = fnv1a64_update(hash, policy.omc_threads.to_string().as_bytes()); + hash = fnv1a64_update(hash, &[0xf8]); hash = fnv1a64_update(hash, &[u8::from(policy.use_experiment_stop_time)]); hash = fnv1a64_update(hash, &[0xfa]); if let Some(stop_time_override) = policy.stop_time_override { @@ -222,9 +241,7 @@ pub(super) fn persist_simulation_parity_cache_entry( } pub(super) fn current_omc_version() -> io::Result { - let output = std::process::Command::new("omc") - .arg("--version") - .output()?; + let output = omc_version_command().output()?; if !output.status.success() { return Err(io::Error::other(format!( "failed to query OMC version (status={})", @@ -244,6 +261,24 @@ pub(super) fn current_omc_version() -> io::Result { Ok(version) } +fn omc_version_command() -> std::process::Command { + if let Ok(image) = std::env::var("RUMOCA_OMC_DOCKER_IMAGE") + && !image.trim().is_empty() + { + let mut command = std::process::Command::new("docker"); + command + .arg("run") + .arg("--rm") + .arg(image.trim()) + .arg("omc") + .arg("--version"); + return command; + } + let mut command = std::process::Command::new("omc"); + command.arg("--version"); + command +} + pub(super) fn parity_cache_matches_targets_and_msl( path: &Path, target_models: &[String], @@ -306,6 +341,22 @@ pub(super) fn simulation_parity_cache_matches( if batch_timeout_seconds != Some(policy.batch_timeout_seconds) { return Ok(false); } + let workers_used = payload + .get("timing") + .and_then(serde_json::Value::as_object) + .and_then(|timing| timing.get("workers_used")) + .and_then(serde_json::Value::as_u64); + if workers_used != Some(policy.workers as u64) { + return Ok(false); + } + let omc_threads = payload + .get("timing") + .and_then(serde_json::Value::as_object) + .and_then(|timing| timing.get("omc_threads")) + .and_then(serde_json::Value::as_u64); + if omc_threads != Some(policy.omc_threads as u64) { + return Ok(false); + } let use_experiment_stop_time = payload .get("use_experiment_stop_time") .and_then(serde_json::Value::as_bool); @@ -321,6 +372,106 @@ pub(super) fn simulation_parity_cache_matches( })) } +pub(super) fn simulation_parity_cache_can_resume( + path: &Path, + target_models: &[String], + msl_version: &str, + omc_version: &str, + policy: SimulationParityCachePolicy, +) -> io::Result { + if !path.is_file() { + return Ok(false); + } + let payload: serde_json::Value = + serde_json::from_reader(File::open(path)?).map_err(|error| { + io::Error::other(format!( + "invalid simulation parity JSON ({}): {error}", + path.display() + )) + })?; + if payload + .get("msl_version") + .and_then(serde_json::Value::as_str) + .is_none_or(|cached| canonical_msl_version(cached) != canonical_msl_version(msl_version)) + { + return Ok(false); + } + if payload + .get("omc_version") + .and_then(serde_json::Value::as_str) + .is_none_or(|cached| canonical_omc_version(cached) != canonical_omc_version(omc_version)) + { + return Ok(false); + } + let Some(cached_models) = model_names_from_omc_models_map(&payload) else { + return Ok(false); + }; + if cached_models.is_empty() { + return Ok(false); + } + if !omc_models_map_has_success(&payload) { + return Ok(false); + } + let target_models = normalize_model_names(target_models.to_vec()); + let target_set = target_models + .into_iter() + .collect::>(); + if cached_models + .iter() + .any(|model| !target_set.contains(model)) + { + return Ok(false); + } + let batch_timeout_seconds = payload + .get("timing") + .and_then(serde_json::Value::as_object) + .and_then(|timing| timing.get("batch_timeout_seconds")) + .and_then(serde_json::Value::as_u64); + if batch_timeout_seconds != Some(policy.batch_timeout_seconds) { + return Ok(false); + } + let workers_used = payload + .get("timing") + .and_then(serde_json::Value::as_object) + .and_then(|timing| timing.get("workers_used")) + .and_then(serde_json::Value::as_u64); + if workers_used != Some(policy.workers as u64) { + return Ok(false); + } + let omc_threads = payload + .get("timing") + .and_then(serde_json::Value::as_object) + .and_then(|timing| timing.get("omc_threads")) + .and_then(serde_json::Value::as_u64); + if omc_threads != Some(policy.omc_threads as u64) { + return Ok(false); + } + let use_experiment_stop_time = payload + .get("use_experiment_stop_time") + .and_then(serde_json::Value::as_bool); + if use_experiment_stop_time != Some(policy.use_experiment_stop_time) { + return Ok(false); + } + let Some(stop_time_override) = policy.stop_time_override else { + return Ok(true); + }; + let stop_time = payload.get("stop_time").and_then(serde_json::Value::as_f64); + Ok(stop_time.is_some_and(|value| { + (value - stop_time_override).abs() <= f64::EPSILON.max(stop_time_override.abs() * 1e-12) + })) +} + +fn omc_models_map_has_success(payload: &serde_json::Value) -> bool { + payload + .get("models") + .and_then(serde_json::Value::as_object) + .is_some_and(|models| { + models.values().any(|model| { + model.get("status").and_then(serde_json::Value::as_str) == Some("success") + }) + }) +} + pub(super) fn run_msl_tool_command(exe: &Path, args: I) -> io::Result<()> where I: IntoIterator, @@ -356,7 +507,10 @@ where } pub(super) fn omc_parity_workers() -> usize { - msl_stage_parallelism() + parity_config() + .omc_parity_workers + .filter(|value| *value > 0) + .unwrap_or_else(|| msl_stage_parallelism().clamp(1, OMC_PARITY_WORKERS_DEFAULT_MAX)) } pub(super) fn omc_parity_threads() -> usize { @@ -364,5 +518,5 @@ pub(super) fn omc_parity_threads() -> usize { } pub(super) fn force_omc_parity_refresh_enabled() -> bool { - false + parity_config().force_omc_parity_refresh.unwrap_or(false) } diff --git a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/status.rs b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/status.rs index e9213b0c0..c7d3ef0b3 100644 --- a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/status.rs +++ b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/status.rs @@ -180,15 +180,30 @@ pub(super) fn print_runtime_ratio_status( .and_then(|context| context.omc_threads); if let Some(baseline_runtime) = baseline.runtime_ratio_stats.as_ref() { + let system_floor = baseline_runtime.system_ratio_both_success.median + * (1.0 - RUNTIME_RATIO_MEDIAN_REL_TOLERANCE); + let system_status = if current_runtime.system_ratio_both_success.median + + SIM_RATE_GATE_EPSILON + < system_floor + { + "FAIL" + } else { + "PASS" + }; println!( - "MSL speed gate: PASS system_median={:.3e} (baseline={:.3e}), wall_median={:.3e} (baseline={:.3e}), workers={}, omc_threads={}.", + "MSL system speed gate: {system_status} median={:.3e}, baseline={:.3e}, floor={:.3e} (tolerance={:.1}%), workers={}, omc_threads={}.", current_runtime.system_ratio_both_success.median, baseline_runtime.system_ratio_both_success.median, - current_runtime.wall_ratio_both_success.median, - baseline_runtime.wall_ratio_both_success.median, + system_floor, + RUNTIME_RATIO_MEDIAN_REL_TOLERANCE * 100.0, fmt_opt_usize(current_workers), fmt_opt_usize(current_omc_threads) ); + + println!( + "{}", + format_wall_time_status(&wall_time_status_content(baseline, parity_input)) + ); return; } diff --git a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/tests.rs b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/tests.rs index 44a3d528f..d49d79db5 100644 --- a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/tests.rs +++ b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/tests.rs @@ -7,6 +7,21 @@ use std::path::Path; use std::path::PathBuf; use tempfile::tempdir; +mod cache_resume; +mod wall_time; +const ISOLATED_QUALITY_GATE_CHILD_MARKER: &str = "target/msl/isolated-quality-gate-child.marker"; +fn is_isolated_quality_gate_child() -> bool { + Path::new(ISOLATED_QUALITY_GATE_CHILD_MARKER).is_file() +} + +fn write_isolated_quality_gate_child_marker(workspace: &Path) { + let marker = workspace.join(ISOLATED_QUALITY_GATE_CHILD_MARKER); + fs::create_dir_all(marker.parent().expect("marker parent")) + .expect("create isolated child marker directory"); + fs::write(marker, "isolated quality-gate child\n") + .expect("write isolated quality-gate child marker"); +} + fn assert_distribution_parsed(input: Value, expected: MslDistributionStats) { let stats = parse_distribution_stats(&input).expect("expected distribution stats"); assert_eq!(stats.sample_count, expected.sample_count); @@ -96,6 +111,39 @@ fn panic_message(payload: &Box) -> String { "".to_string() } +fn run_in_isolated_quality_gate_workspace(test_name: &str, config: Option) { + let workspace = tempdir().expect("temporary workspace"); + fs::write(workspace.path().join("Cargo.toml"), "[workspace]\n") + .expect("write workspace manifest"); + let crate_dir = workspace.path().join("crates/rumoca-test-msl"); + fs::create_dir_all(&crate_dir).expect("create temporary crate directory"); + fs::write( + crate_dir.join("Cargo.toml"), + "[package]\nname = \"rumoca-test-msl\"\nversion = \"0.0.0\"\n", + ) + .expect("write temporary crate manifest"); + if let Some(config) = config { + let config_path = workspace.path().join("target/msl/parity-config.json"); + fs::create_dir_all(config_path.parent().expect("config parent")) + .expect("create parity config directory"); + fs::write( + config_path, + serde_json::to_vec_pretty(&config).expect("serialize parity config"), + ) + .expect("write parity config"); + } + write_isolated_quality_gate_child_marker(workspace.path()); + + let status = std::process::Command::new(std::env::current_exe().expect("current test binary")) + .current_dir(workspace.path()) + .arg("--exact") + .arg(test_name) + .arg("--nocapture") + .status() + .expect("run isolated default-parity regression"); + assert!(status.success(), "isolated regression failed: {test_name}"); +} + fn baseline_quality_template() -> MslQualityBaseline { MslQualityBaseline { quality_gate_version: MSL_QUALITY_GATE_VERSION, @@ -163,6 +211,7 @@ fn valid_summary_template() -> MslSummary { fn selected_target_failures_report_non_sim_ok_models() { let mut summary = valid_summary_template(); summary.sim_target_models = vec!["A".to_string(), "B".to_string()]; + summary.sim_attempted = summary.sim_target_models.len(); let mut ok = phase_error_result("A".to_string(), "Success", None, None); ok.sim_status = Some("sim_ok".to_string()); let mut fail = phase_error_result("B".to_string(), "Success", None, None); @@ -196,6 +245,7 @@ fn selected_target_failures_report_missing_results_in_target_order() { fn selected_target_gate_returns_error_instead_of_asserting() { let mut summary = valid_summary_template(); summary.sim_target_models = vec!["A".to_string()]; + summary.sim_attempted = 1; let mut fail = phase_error_result("A".to_string(), "Success", None, None); fail.sim_status = Some("sim_solver_fail".to_string()); summary.model_results = vec![fail]; @@ -209,6 +259,14 @@ fn selected_target_gate_returns_error_instead_of_asserting() { #[test] fn full_quality_gate_rejects_zero_simulation_attempts() { + if !is_isolated_quality_gate_child() { + run_in_isolated_quality_gate_workspace( + "balance_pipeline::balance_pipeline_quality_gate::tests::full_quality_gate_rejects_zero_simulation_attempts", + None, + ); + return; + } + let mut summary = valid_summary_template(); summary.sim_target_models = vec!["A".to_string(), "B".to_string()]; @@ -219,10 +277,145 @@ fn full_quality_gate_rejects_zero_simulation_attempts() { assert!(message.contains("0 simulations attempted for 2 selected simulation target(s)")); } +#[test] +fn root_examples_full_shard_requires_parity_artifacts_but_skips_aggregate_gate() { + if is_isolated_quality_gate_child() { + assert!( + requires_msl_parity_artifacts(), + "a root-examples/full shard must generate its own parity artifacts for fan-in" + ); + assert!( + should_skip_msl_quality_gate(), + "a shard must still skip the aggregate baseline ratchet" + ); + return; + } + + run_in_isolated_quality_gate_workspace( + "balance_pipeline::balance_pipeline_quality_gate::tests::root_examples_full_shard_requires_parity_artifacts_but_skips_aggregate_gate", + Some(json!({ + "target_scope": "root-examples", + "sim_set": "full", + "shard_index": 1, + "shard_count": 4 + })), + ); +} + +#[test] +fn required_full_parity_rejects_unavailable_omc() { + let error = require_omc_version_for_full_parity(Err(io::Error::new( + io::ErrorKind::NotFound, + "omc executable not found", + ))) + .expect_err("a full MSL parity run must require OMC"); + + let message = error.to_string(); + assert!(message.contains("required OMC prerequisite is unavailable")); + assert!(message.contains("omc executable not found")); +} + +#[test] +fn required_full_parity_rejects_reference_without_comparable_metrics() { + let dir = tempdir().expect("tempdir"); + let path = dir.path().join("omc_simulation_reference.json"); + let mut payload = valid_simulation_parity_payload(); + payload["runtime_comparison"]["ratio_stats"]["system_ratio_both_success"] = + serde_json::Value::Null; + payload["runtime_comparison"]["ratio_stats"]["wall_ratio_both_success"] = + serde_json::Value::Null; + payload["trace_comparison"]["models_compared"] = json!(0); + fs::write( + &path, + serde_json::to_vec_pretty(&payload).expect("serialize payload"), + ) + .expect("write payload"); + + let error = load_required_msl_parity_gate_input_from_path(&path, 7) + .expect_err("full parity preparation must reject zero comparable metrics"); + let message = error.to_string(); + assert!(message.contains("missing runtime_ratio_stats")); + assert!(message.contains("omc_simulation_reference.json")); +} + +#[test] +fn full_quality_gate_rejects_missing_current_parity_input() { + let dir = tempdir().expect("tempdir"); + let path = dir.path().join("omc_simulation_reference.json"); + + let error = load_required_msl_parity_gate_input_from_path(&path, 7) + .expect_err("a full quality gate must require current parity input"); + let message = error.to_string(); + assert!(message.contains("required OMC parity reference")); + assert!(message.contains("is missing")); +} + +#[test] +fn full_root_selected_target_switch_still_requires_current_parity() { + if is_isolated_quality_gate_child() { + let mut summary = valid_summary_template(); + summary.sim_target_models = vec!["A".to_string()]; + summary.sim_attempted = 1; + summary.sim_ok = 1; + let mut result = phase_error_result("A".to_string(), "Success", None, None); + result.sim_status = Some("sim_ok".to_string()); + summary.model_results = vec![result]; + + let error = enforce_msl_quality_gate(&summary) + .expect_err("full root gate must not bypass required parity via selected-target mode"); + assert!(error.to_string().contains("required OMC parity reference")); + return; + } + + let workspace = tempdir().expect("temporary workspace"); + fs::write(workspace.path().join("Cargo.toml"), "[workspace]\n") + .expect("write workspace manifest"); + let crate_dir = workspace.path().join("crates/rumoca-test-msl"); + fs::create_dir_all(&crate_dir).expect("create temporary crate directory"); + fs::write( + crate_dir.join("Cargo.toml"), + "[package]\nname = \"rumoca-test-msl\"\nversion = \"0.0.0\"\n", + ) + .expect("write temporary crate manifest"); + let config_path = workspace.path().join("target/msl/parity-config.json"); + fs::create_dir_all(config_path.parent().expect("config parent")) + .expect("create parity config directory"); + let results_dir = workspace.path().join("results"); + fs::write( + &config_path, + serde_json::to_vec_pretty(&json!({ + "results_dir": results_dir, + "require_selected_targets_success": true + })) + .expect("serialize parity config"), + ) + .expect("write parity config"); + write_isolated_quality_gate_child_marker(workspace.path()); + + let status = std::process::Command::new(std::env::current_exe().expect("current test binary")) + .current_dir(workspace.path()) + .arg("--exact") + .arg("balance_pipeline::balance_pipeline_quality_gate::tests::full_root_selected_target_switch_still_requires_current_parity") + .arg("--nocapture") + .status() + .expect("run isolated full-root gate regression"); + assert!( + status.success(), + "isolated full-root gate regression failed" + ); +} + +#[test] +fn focused_or_partial_runs_do_not_require_full_parity() { + assert!(!full_parity_is_required(true, 1)); + assert!(!full_parity_is_required(false, 0)); + assert!(full_parity_is_required(false, 1)); +} + #[test] fn current_quality_snapshot_marks_only_partial_runs() { let summary = valid_summary_template(); - let full = current_msl_quality_snapshot_json(&summary, None, false) + let full = current_msl_quality_snapshot_json(&summary, None, None, false) .expect("full snapshot should serialize"); assert_eq!( full.get("quality_gate_version").and_then(Value::as_u64), @@ -237,7 +430,7 @@ fn current_quality_snapshot_marks_only_partial_runs() { "full baseline snapshots should omit the partial marker" ); - let partial = current_msl_quality_snapshot_json(&summary, None, true) + let partial = current_msl_quality_snapshot_json(&summary, None, None, true) .expect("partial snapshot should serialize"); assert_eq!( partial.get("run_scope").and_then(Value::as_str), @@ -254,12 +447,13 @@ fn current_quality_snapshot_records_parity_omc_version() { omc_version: Some("OpenModelica 1.26.1".to_string()), runtime_context: None, runtime_ratio_stats: None, + wall_time_provenance: None, trace_accuracy_stats: None, omc_assertion_failure_models: 0, omc_assertion_failure_examples: Vec::new(), }; - let snapshot = current_msl_quality_snapshot_json(&summary, Some(&parity), false) + let snapshot = current_msl_quality_snapshot_json(&summary, Some(&parity), None, false) .expect("snapshot should serialize"); assert_eq!( snapshot.get("omc_version").and_then(Value::as_str), @@ -278,12 +472,13 @@ fn current_quality_snapshot_records_runtime_ratio_stats() { omc_threads: Some(1), }), runtime_ratio_stats: Some(runtime_ratio_stats(5.0, 4.0)), + wall_time_provenance: None, trace_accuracy_stats: None, omc_assertion_failure_models: 0, omc_assertion_failure_examples: Vec::new(), }; - let snapshot = current_msl_quality_snapshot_json(&summary, Some(&parity), false) + let snapshot = current_msl_quality_snapshot_json(&summary, Some(&parity), None, false) .expect("snapshot should serialize"); assert_eq!( snapshot @@ -314,6 +509,7 @@ fn quality_context_reports_omc_version_mismatch_for_pinned_baseline() { omc_version: Some("OpenModelica 1.27.0".to_string()), runtime_context: None, runtime_ratio_stats: None, + wall_time_provenance: None, trace_accuracy_stats: None, omc_assertion_failure_models: 0, omc_assertion_failure_examples: Vec::new(), @@ -338,6 +534,7 @@ fn quality_context_accepts_omc_package_rebuild_suffix_drift() { omc_version: Some("OpenModelica 1.26.7~2-ge74480f".to_string()), runtime_context: None, runtime_ratio_stats: None, + wall_time_provenance: None, trace_accuracy_stats: None, omc_assertion_failure_models: 0, omc_assertion_failure_examples: Vec::new(), @@ -385,7 +582,7 @@ fn current_quality_snapshot_includes_pipeline_progression() { vec!["Modelica.NotAStandaloneRoot".to_string()], ); - let snapshot = current_msl_quality_snapshot_json(&summary, None, false) + let snapshot = current_msl_quality_snapshot_json(&summary, None, None, false) .expect("snapshot should serialize"); let pipeline = snapshot .get("pipeline_progress") @@ -471,7 +668,7 @@ fn current_quality_snapshot_includes_mls_contract_category_coverage() { ); summary.model_results = vec![array_result, connector_result]; - let snapshot = current_msl_quality_snapshot_json(&summary, None, false) + let snapshot = current_msl_quality_snapshot_json(&summary, None, None, false) .expect("snapshot should serialize"); let coverage = snapshot .get("mls_contract_coverage") @@ -806,6 +1003,14 @@ fn valid_msl_summary_rejects_resolve_errors() { #[test] fn valid_msl_summary_rejects_baseline_sim_run_below_hard_floor() { + if !is_isolated_quality_gate_child() { + run_in_isolated_quality_gate_workspace( + "balance_pipeline::balance_pipeline_quality_gate::tests::valid_msl_summary_rejects_baseline_sim_run_below_hard_floor", + None, + ); + return; + } + let mut summary = valid_summary_template(); summary.total_models = SIM_SET_LIMIT_DEFAULT; summary.sim_attempted = SIM_SET_LIMIT_DEFAULT; @@ -950,37 +1155,6 @@ fn trace_accuracy_deviation_migrated_to_near() -> MslTraceAccuracyStatsBaseline } } -#[test] -fn runtime_ratio_regression_reason_triggers_on_large_drop() { - let baseline = MslQualityBaseline { - runtime_ratio_stats: Some(runtime_ratio_stats(2.0, 1.5)), - ..baseline_quality_template() - }; - let parity = MslParityGateInput { - total_models: Some(10), - omc_version: Some("OpenModelica 1.26.1".to_string()), - runtime_context: None, - runtime_ratio_stats: Some(runtime_ratio_stats(1.0, 0.5)), - trace_accuracy_stats: None, - omc_assertion_failure_models: 0, - omc_assertion_failure_examples: Vec::new(), - }; - - let mut reasons = Vec::new(); - push_runtime_ratio_regression_reasons(&mut reasons, &baseline, Some(&parity)); - assert_eq!(reasons.len(), 2); - assert!( - reasons - .iter() - .any(|reason| reason.contains("runtime system speedup median")) - ); - assert!( - reasons - .iter() - .any(|reason| reason.contains("runtime wall speedup median")) - ); -} - #[test] fn msl_quality_regression_reasons_include_runtime_ratio_drop() { let mut baseline = baseline_quality_template(); @@ -991,6 +1165,7 @@ fn msl_quality_regression_reasons_include_runtime_ratio_drop() { omc_version: Some("OpenModelica 1.26.1".to_string()), runtime_context: None, runtime_ratio_stats: Some(runtime_ratio_stats(1.0, 1.5)), + wall_time_provenance: None, trace_accuracy_stats: Some(trace_accuracy_baseline()), omc_assertion_failure_models: 0, omc_assertion_failure_examples: Vec::new(), @@ -1017,6 +1192,7 @@ fn runtime_ratio_gate_allows_observed_ci_runner_delta() { omc_version: Some("OpenModelica 1.26.8".to_string()), runtime_context: None, runtime_ratio_stats: Some(runtime_ratio_stats(0.887_313_1, 0.887_313_1)), + wall_time_provenance: None, trace_accuracy_stats: None, omc_assertion_failure_models: 0, omc_assertion_failure_examples: Vec::new(), @@ -1041,6 +1217,7 @@ fn trace_bucket_and_channel_regression_reasons_trigger_when_thresholds_are_excee omc_version: Some("OpenModelica 1.26.1".to_string()), runtime_context: None, runtime_ratio_stats: None, + wall_time_provenance: None, trace_accuracy_stats: Some(trace_accuracy_regressed()), omc_assertion_failure_models: 0, omc_assertion_failure_examples: Vec::new(), @@ -1073,6 +1250,7 @@ fn trace_channel_share_tolerances_allow_small_runner_drift() { omc_version: Some("OpenModelica 1.26.1".to_string()), runtime_context: None, runtime_ratio_stats: None, + wall_time_provenance: None, trace_accuracy_stats: Some(trace_accuracy_small_channel_drift()), omc_assertion_failure_models: 0, omc_assertion_failure_examples: Vec::new(), @@ -1105,6 +1283,7 @@ fn trace_near_to_high_promotion_does_not_trigger_regression() { omc_version: Some("OpenModelica 1.26.1".to_string()), runtime_context: None, runtime_ratio_stats: None, + wall_time_provenance: None, trace_accuracy_stats: Some(trace_accuracy_near_promoted_to_high()), omc_assertion_failure_models: 0, omc_assertion_failure_examples: Vec::new(), @@ -1131,6 +1310,7 @@ fn trace_deviation_to_near_migration_does_not_trigger_regression() { omc_version: Some("OpenModelica 1.26.1".to_string()), runtime_context: None, runtime_ratio_stats: None, + wall_time_provenance: None, trace_accuracy_stats: Some(trace_accuracy_deviation_migrated_to_near()), omc_assertion_failure_models: 0, omc_assertion_failure_examples: Vec::new(), @@ -1155,6 +1335,7 @@ fn trace_acceptable_band_regression_reason_triggers_on_real_drop() { omc_version: Some("OpenModelica 1.26.1".to_string()), runtime_context: None, runtime_ratio_stats: None, + wall_time_provenance: None, trace_accuracy_stats: Some(trace_accuracy_acceptable_band_regressed()), omc_assertion_failure_models: 0, omc_assertion_failure_examples: Vec::new(), @@ -1202,6 +1383,7 @@ fn trace_fixed_denominator_gate_accepts_current_ci_delta() { omc_version: Some("OpenModelica 1.26.1".to_string()), runtime_context: None, runtime_ratio_stats: None, + wall_time_provenance: None, trace_accuracy_stats: Some(current_trace), omc_assertion_failure_models: 0, omc_assertion_failure_examples: Vec::new(), @@ -1218,6 +1400,7 @@ fn trace_fixed_denominator_gate_accepts_current_ci_delta() { fn valid_simulation_parity_payload() -> Value { json!({ "total_models": 7, + "omc_version": "OpenModelica 1.26.1", "runtime_comparison": { "ratio_stats": { "system_ratio_both_success": { "sample_count": 5, @@ -1262,6 +1445,88 @@ fn valid_simulation_parity_payload() -> Value { }) } +#[test] +fn optional_parity_input_treats_target_count_mismatch_as_absent() { + let dir = tempdir().expect("tempdir"); + let path = dir.path().join("omc_simulation_reference.json"); + fs::write( + &path, + serde_json::to_vec_pretty(&valid_simulation_parity_payload()).expect("serialize payload"), + ) + .expect("write payload"); + + let optional = load_msl_parity_gate_input_optional_from_path(&path, 1) + .expect("stale optional parity input should not error"); + assert!( + optional.is_none(), + "stale optional parity input should be ignored for focused snapshots" + ); + + let stale = validate_parity_total_models( + &path, + &load_msl_parity_gate_input(&path).expect("load parity input"), + 1, + ) + .expect_err("required parity input must reject a stale target count"); + assert!( + stale.to_string().contains("is stale"), + "required parity input should still explain the stale reference, got {stale}" + ); +} + +#[test] +fn optional_parity_input_treats_missing_comparison_metrics_as_absent() { + let dir = tempdir().expect("tempdir"); + let path = dir.path().join("omc_simulation_reference.json"); + let mut payload = valid_simulation_parity_payload(); + payload["runtime_comparison"]["ratio_stats"]["system_ratio_both_success"] = + serde_json::Value::Null; + payload["runtime_comparison"]["ratio_stats"]["wall_ratio_both_success"] = + serde_json::Value::Null; + payload["trace_comparison"]["models_compared"] = json!(0); + fs::write( + &path, + serde_json::to_vec_pretty(&payload).expect("serialize payload"), + ) + .expect("write payload"); + + let optional = load_msl_parity_gate_input_optional_from_path(&path, 7) + .expect("missing comparison metrics should make optional parity absent"); + assert!( + optional.is_none(), + "OMC runs with no comparable OMC/Rumoca samples should not hard-fail optional parity" + ); + + let required = validate_required_msl_parity_gate_input( + &path, + load_msl_parity_gate_input(&path).expect("load parity input"), + ) + .expect_err("required parity input must still reject missing comparison metrics"); + assert!( + required.to_string().contains("missing runtime_ratio_stats"), + "required parity should explain the missing metrics, got {required}" + ); +} + +#[test] +fn required_full_parity_rejects_trace_only_missing_comparable_metrics() { + let dir = tempdir().expect("tempdir"); + let path = dir.path().join("omc_simulation_reference.json"); + let mut payload = valid_simulation_parity_payload(); + payload["trace_comparison"]["models_compared"] = json!(0); + fs::write( + &path, + serde_json::to_vec_pretty(&payload).expect("serialize payload"), + ) + .expect("write payload"); + + let error = load_required_msl_parity_gate_input_from_path(&path, 7) + .expect_err("full parity must reject a reference with no comparable traces"); + let message = error.to_string(); + assert!(message.contains("missing comparable trace metrics")); + assert!(message.contains("models_compared=0")); +} + #[test] fn simulation_parity_cache_requires_runtime_and_trace_metrics() { fn write_payload(path: &Path, payload: &Value) { @@ -1449,6 +1714,7 @@ fn parity_total_models_guard_checks_stale_and_matching_counts() { omc_version: Some("OpenModelica 1.26.1".to_string()), runtime_context: None, runtime_ratio_stats: None, + wall_time_provenance: None, trace_accuracy_stats: None, omc_assertion_failure_models: 0, omc_assertion_failure_examples: Vec::new(), @@ -1463,6 +1729,7 @@ fn parity_total_models_guard_checks_stale_and_matching_counts() { omc_version: Some("OpenModelica 1.26.1".to_string()), runtime_context: None, runtime_ratio_stats: None, + wall_time_provenance: None, trace_accuracy_stats: None, omc_assertion_failure_models: 0, omc_assertion_failure_examples: Vec::new(), @@ -1517,6 +1784,8 @@ fn simulation_parity_cache_key_changes_with_policy() { "OpenModelica 1.26.1", SimulationParityCachePolicy { batch_timeout_seconds: 600, + workers: 2, + omc_threads: 1, use_experiment_stop_time: true, stop_time_override: None, }, @@ -1527,6 +1796,8 @@ fn simulation_parity_cache_key_changes_with_policy() { "OpenModelica 1.26.1", SimulationParityCachePolicy { batch_timeout_seconds: 900, + workers: 2, + omc_threads: 1, use_experiment_stop_time: true, stop_time_override: None, }, @@ -1537,6 +1808,8 @@ fn simulation_parity_cache_key_changes_with_policy() { "OpenModelica 1.26.1", SimulationParityCachePolicy { batch_timeout_seconds: 600, + workers: 2, + omc_threads: 1, use_experiment_stop_time: false, stop_time_override: Some(30.0), }, @@ -1557,7 +1830,9 @@ fn simulation_parity_cache_matches_rejects_mismatched_policy() { "stop_time": 10.0, "use_experiment_stop_time": true, "timing": { - "batch_timeout_seconds": 600 + "batch_timeout_seconds": 600, + "workers_used": 2, + "omc_threads": 1 }, "models": { "A": { "status": "success" }, @@ -1570,6 +1845,8 @@ fn simulation_parity_cache_matches_rejects_mismatched_policy() { let matching = SimulationParityCachePolicy { batch_timeout_seconds: 600, + workers: 2, + omc_threads: 1, use_experiment_stop_time: true, stop_time_override: None, }; @@ -1577,8 +1854,14 @@ fn simulation_parity_cache_matches_rejects_mismatched_policy() { batch_timeout_seconds: 900, ..matching }; + let mismatched_workers = SimulationParityCachePolicy { + workers: 3, + ..matching + }; let mismatched_override = SimulationParityCachePolicy { batch_timeout_seconds: 600, + workers: 2, + omc_threads: 1, use_experiment_stop_time: false, stop_time_override: Some(30.0), }; @@ -1604,6 +1887,17 @@ fn simulation_parity_cache_matches_rejects_mismatched_policy() { .expect("mismatched timeout should parse"), "batch-timeout drift should invalidate cache entry" ); + assert!( + !simulation_parity_cache_matches( + &path, + &["A".to_string(), "B".to_string()], + "4.1.0", + "OpenModelica 1.26.1", + mismatched_workers, + ) + .expect("mismatched workers should parse"), + "OMC worker drift should invalidate cache entry" + ); assert!( !simulation_parity_cache_matches( &path, @@ -1616,3 +1910,90 @@ fn simulation_parity_cache_matches_rejects_mismatched_policy() { "stop-time policy drift should invalidate cache entry" ); } + +#[test] +fn simulation_parity_cache_can_resume_partial_target_subset() { + let temp = tempdir().expect("tempdir"); + let path = temp.path().join("omc_simulation_reference.json"); + fs::write( + &path, + serde_json::to_vec_pretty(&json!({ + "msl_version": "4.1.0", + "omc_version": "OpenModelica 1.26.1", + "stop_time": 10.0, + "use_experiment_stop_time": true, + "timing": { + "batch_timeout_seconds": 600, + "workers_used": 2, + "omc_threads": 1 + }, + "models": { + "A": { "status": "success" } + } + })) + .expect("serialize cache payload"), + ) + .expect("write cache payload"); + + assert!( + simulation_parity_cache_can_resume( + &path, + &["A".to_string(), "B".to_string()], + "4.1.0", + "OpenModelica 1.26.1", + SimulationParityCachePolicy { + batch_timeout_seconds: 600, + workers: 2, + omc_threads: 1, + use_experiment_stop_time: true, + stop_time_override: None, + }, + ) + .expect("partial checkpoint should parse"), + "a policy-compatible partial checkpoint should be resumable" + ); +} + +#[test] +fn simulation_parity_cache_can_resume_rejects_models_outside_target_set() { + let temp = tempdir().expect("tempdir"); + let path = temp.path().join("omc_simulation_reference.json"); + fs::write( + &path, + serde_json::to_vec_pretty(&json!({ + "msl_version": "4.1.0", + "omc_version": "OpenModelica 1.26.1", + "stop_time": 10.0, + "use_experiment_stop_time": true, + "timing": { + "batch_timeout_seconds": 600, + "workers_used": 2, + "omc_threads": 1 + }, + "models": { + "A": { "status": "success" }, + "Outside": { "status": "success" } + } + })) + .expect("serialize cache payload"), + ) + .expect("write cache payload"); + + assert!( + !simulation_parity_cache_can_resume( + &path, + &["A".to_string(), "B".to_string()], + "4.1.0", + "OpenModelica 1.26.1", + SimulationParityCachePolicy { + batch_timeout_seconds: 600, + workers: 2, + omc_threads: 1, + use_experiment_stop_time: true, + stop_time_override: None, + }, + ) + .expect("partial checkpoint should parse"), + "a checkpoint containing non-target models must not be resumed" + ); +} diff --git a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/tests/cache_resume.rs b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/tests/cache_resume.rs new file mode 100644 index 000000000..b9e22b108 --- /dev/null +++ b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/tests/cache_resume.rs @@ -0,0 +1,47 @@ +use super::*; + +#[test] +fn simulation_parity_cache_can_resume_rejects_all_error_checkpoint() { + let temp = tempdir().expect("tempdir"); + let path = temp.path().join("omc_simulation_reference.json"); + fs::write( + &path, + serde_json::to_vec_pretty(&json!({ + "msl_version": "4.1.0", + "omc_version": "OpenModelica 1.26.1", + "stop_time": 10.0, + "use_experiment_stop_time": true, + "timing": { + "batch_timeout_seconds": 600, + "workers_used": 2, + "omc_threads": 1 + }, + "models": { + "A": { + "status": "error", + "error": "omc session spawn failed: port file did not appear" + } + } + })) + .expect("serialize cache payload"), + ) + .expect("write cache payload"); + + assert!( + !simulation_parity_cache_can_resume( + &path, + &["A".to_string(), "B".to_string()], + "4.1.0", + "OpenModelica 1.26.1", + SimulationParityCachePolicy { + batch_timeout_seconds: 600, + workers: 2, + omc_threads: 1, + use_experiment_stop_time: true, + stop_time_override: None, + }, + ) + .expect("all-error checkpoint should parse"), + "a checkpoint without any reusable successful OMC result must not be resumed" + ); +} diff --git a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/tests/wall_time.rs b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/tests/wall_time.rs new file mode 100644 index 000000000..93b144693 --- /dev/null +++ b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/tests/wall_time.rs @@ -0,0 +1,379 @@ +use super::*; + +fn provenance( + fresh: usize, + cached: usize, + affinity_requested: usize, + affinity_applied: usize, + affinity_failed: usize, + load_before: Option, + load_after: Option, +) -> MslWallTimeProvenance { + MslWallTimeProvenance { + omc_fresh_sample_count: fresh, + omc_cached_sample_count: cached, + affinity_requested_worker_count: affinity_requested, + affinity_applied_worker_count: affinity_applied, + affinity_failed_worker_count: affinity_failed, + normalized_load_before: load_before, + normalized_load_after: load_after, + rumoca_workers_used: affinity_requested, + workers_used: affinity_requested, + omc_threads: 1, + } +} + +#[test] +fn rumoca_and_omc_worker_topologies_are_independent() { + let baseline = baseline_with_runtime(2.0, 1.5, 2, 1); + let mut wall = provenance(8, 0, 14, 14, 0, Some(0.2), Some(0.3)); + wall.rumoca_workers_used = 14; + wall.workers_used = 2; + let parity = parity_with_provenance(runtime_ratio_stats(2.0, 1.5), wall); + assert!(wall_time_trust_decision(&baseline, Some(&parity)).trusted); +} + +#[test] +fn wall_status_content_covers_pass_fail_and_advisory() { + let baseline = baseline_with_runtime(2.0, 2.0, 2, 1); + let trusted = provenance(8, 0, 2, 2, 0, Some(0.2), Some(0.3)); + let pass = wall_time_status_content( + &baseline, + Some(&parity_with_provenance( + runtime_ratio_stats(2.0, 1.5), + trusted.clone(), + )), + ); + assert_eq!(pass.status, "PASS"); + assert!(format_wall_time_status(&pass).contains("MSL wall speed gate: PASS")); + let fail = wall_time_status_content( + &baseline, + Some(&parity_with_provenance( + runtime_ratio_stats(2.0, 1.0), + trusted, + )), + ); + assert_eq!(fail.status, "FAIL"); + assert!(format_wall_time_status(&fail).contains("MSL wall speed gate: FAIL")); + let advisory = wall_time_status_content(&baseline, None); + assert_eq!(advisory.status, "ADVISORY"); + assert!(!advisory.trusted); + assert!(format_wall_time_status(&advisory).contains("missing parity input")); +} + +fn parity_with_provenance( + runtime_ratio_stats: MslRuntimeRatioStatsBaseline, + wall_time_provenance: MslWallTimeProvenance, +) -> MslParityGateInput { + MslParityGateInput { + total_models: Some(10), + omc_version: Some("OpenModelica 1.26.1".to_string()), + runtime_context: Some(MslParityRuntimeContext { + workers_used: Some(wall_time_provenance.workers_used), + omc_threads: Some(wall_time_provenance.omc_threads), + }), + runtime_ratio_stats: Some(runtime_ratio_stats), + wall_time_provenance: Some(wall_time_provenance), + trace_accuracy_stats: None, + omc_assertion_failure_models: 0, + omc_assertion_failure_examples: Vec::new(), + } +} + +fn baseline_with_runtime( + system_median: f64, + wall_median: f64, + workers_used: usize, + omc_threads: usize, +) -> MslQualityBaseline { + MslQualityBaseline { + runtime_context: Some(MslParityRuntimeContext { + workers_used: Some(workers_used), + omc_threads: Some(omc_threads), + }), + runtime_ratio_stats: Some(runtime_ratio_stats(system_median, wall_median)), + ..baseline_quality_template() + } +} + +#[test] +fn cached_omc_wall_time_regression_is_advisory_but_system_regression_blocks() { + let baseline = baseline_with_runtime(2.0, 1.5, 2, 1); + let parity = parity_with_provenance( + runtime_ratio_stats(1.0, 0.5), + provenance(0, 8, 2, 2, 0, Some(0.2), Some(0.3)), + ); + let mut reasons = Vec::new(); + push_runtime_ratio_regression_reasons(&mut reasons, &baseline, Some(&parity)); + assert!( + reasons + .iter() + .any(|reason| reason.contains("runtime system speedup median")) + ); + assert!( + !reasons + .iter() + .any(|reason| reason.contains("runtime wall speedup median")) + ); + assert!( + wall_time_trust_decision(&baseline, Some(&parity)) + .reasons + .iter() + .any(|reason| reason.contains("cached")) + ); +} + +#[test] +fn trusted_wall_time_regression_remains_blocking() { + let baseline = baseline_with_runtime(2.0, 1.5, 2, 1); + let parity = parity_with_provenance( + runtime_ratio_stats(2.0, 0.5), + provenance(8, 0, 2, 2, 0, Some(0.5), Some(0.6)), + ); + let mut reasons = Vec::new(); + push_runtime_ratio_regression_reasons(&mut reasons, &baseline, Some(&parity)); + assert!( + reasons + .iter() + .any(|reason| reason.contains("runtime wall speedup median")) + ); +} + +#[test] +fn affinity_failure_makes_wall_time_advisory() { + let baseline = baseline_with_runtime(2.0, 1.5, 2, 1); + let parity = parity_with_provenance( + runtime_ratio_stats(2.0, 1.5), + provenance(8, 0, 2, 1, 1, Some(0.5), Some(0.6)), + ); + let decision = wall_time_trust_decision(&baseline, Some(&parity)); + assert!(!decision.trusted); + assert!( + decision + .reasons + .iter() + .any(|reason| reason.contains("affinity")) + ); +} + +#[test] +fn excessive_host_load_makes_wall_time_advisory() { + let baseline = baseline_with_runtime(2.0, 1.5, 2, 1); + let parity = parity_with_provenance( + runtime_ratio_stats(2.0, 1.5), + provenance(8, 0, 2, 2, 0, Some(1.51), Some(0.6)), + ); + let decision = wall_time_trust_decision(&baseline, Some(&parity)); + assert!(!decision.trusted); + assert!( + decision + .reasons + .iter() + .any(|reason| reason.contains("load")) + ); +} + +#[test] +fn missing_host_load_makes_wall_time_advisory() { + let baseline = baseline_with_runtime(2.0, 1.5, 2, 1); + let parity = parity_with_provenance( + runtime_ratio_stats(2.0, 1.5), + provenance(8, 0, 2, 2, 0, None, Some(0.6)), + ); + let decision = wall_time_trust_decision(&baseline, Some(&parity)); + assert!(!decision.trusted); + assert!( + decision + .reasons + .iter() + .any(|reason| reason.contains("missing load")) + ); +} + +#[test] +fn missing_provenance_makes_wall_time_advisory() { + let baseline = baseline_with_runtime(2.0, 1.5, 2, 1); + let parity = MslParityGateInput { + total_models: Some(10), + omc_version: Some("OpenModelica 1.26.1".to_string()), + runtime_context: Some(MslParityRuntimeContext { + workers_used: Some(2), + omc_threads: Some(1), + }), + runtime_ratio_stats: Some(runtime_ratio_stats(2.0, 1.5)), + wall_time_provenance: None, + trace_accuracy_stats: None, + omc_assertion_failure_models: 0, + omc_assertion_failure_examples: Vec::new(), + }; + let decision = wall_time_trust_decision(&baseline, Some(&parity)); + assert!(!decision.trusted); + assert!( + decision + .reasons + .iter() + .any(|reason| reason.contains("missing provenance")) + ); +} + +#[test] +fn malformed_provenance_makes_wall_time_advisory() { + let baseline = baseline_with_runtime(2.0, 1.5, 2, 1); + let dir = tempdir().expect("tempdir"); + let path = dir.path().join("omc_simulation_reference.json"); + let mut payload = valid_simulation_parity_payload(); + payload["runtime_comparison"]["wall_time_provenance"] = json!({ + "omc_fresh_sample_count": "not-a-count", + "omc_cached_sample_count": 0 + }); + fs::write( + &path, + serde_json::to_vec_pretty(&payload).expect("serialize payload"), + ) + .expect("write payload"); + let parity = load_msl_parity_gate_input(&path).expect("load parity payload"); + let decision = wall_time_trust_decision(&baseline, Some(&parity)); + assert!(!decision.trusted); + assert!( + decision + .reasons + .iter() + .any(|reason| reason.contains("missing provenance")) + ); +} + +#[test] +fn mismatched_runtime_context_makes_wall_time_advisory() { + let baseline = baseline_with_runtime(2.0, 1.5, 2, 1); + let mut parity = parity_with_provenance( + runtime_ratio_stats(2.0, 1.5), + provenance(8, 0, 2, 2, 0, Some(0.5), Some(0.6)), + ); + parity.runtime_context = Some(MslParityRuntimeContext { + workers_used: Some(4), + omc_threads: Some(2), + }); + let decision = wall_time_trust_decision(&baseline, Some(&parity)); + assert!(!decision.trusted); + assert!( + decision + .reasons + .iter() + .any(|reason| reason.contains("runtime context")) + ); +} + +#[test] +fn self_consistent_wall_time_context_mismatching_baseline_is_advisory() { + let baseline = baseline_with_runtime(2.0, 1.5, 2, 1); + let mut wall_provenance = provenance(8, 0, 4, 4, 0, Some(0.5), Some(0.6)); + wall_provenance.omc_threads = 2; + let parity = parity_with_provenance(runtime_ratio_stats(2.0, 1.5), wall_provenance); + + let decision = wall_time_trust_decision(&baseline, Some(&parity)); + assert!(!decision.trusted); + assert!( + decision + .reasons + .iter() + .any(|reason| reason.contains("baseline runtime context")) + ); +} + +#[test] +fn wall_time_fresh_count_not_covering_compared_samples_is_advisory() { + let baseline = baseline_with_runtime(2.0, 1.5, 2, 1); + let mut runtime = runtime_ratio_stats(2.0, 1.5); + runtime.wall_ratio_both_success.sample_count = 10; + let parity = parity_with_provenance(runtime, provenance(1, 0, 2, 2, 0, Some(0.5), Some(0.6))); + + let decision = wall_time_trust_decision(&baseline, Some(&parity)); + assert!(!decision.trusted); + assert!( + decision + .reasons + .iter() + .any(|reason| reason.contains("sample count mismatch")) + ); +} + +#[test] +fn missing_baseline_wall_time_runtime_context_is_advisory() { + let baseline = MslQualityBaseline { + runtime_ratio_stats: Some(runtime_ratio_stats(2.0, 1.5)), + ..baseline_quality_template() + }; + let parity = parity_with_provenance( + runtime_ratio_stats(2.0, 1.5), + provenance(8, 0, 2, 2, 0, Some(0.5), Some(0.6)), + ); + + let decision = wall_time_trust_decision(&baseline, Some(&parity)); + assert!(!decision.trusted); + assert!( + decision + .reasons + .iter() + .any(|reason| reason.contains("baseline runtime context missing")) + ); +} + +#[test] +fn current_quality_snapshot_serializes_wall_decisions_without_polluting_baseline() { + let summary = valid_summary_template(); + let baseline = baseline_with_runtime(2.0, 2.0, 2, 1); + let mut trusted = provenance(8, 0, 14, 14, 0, Some(0.2), Some(0.3)); + trusted.workers_used = 2; + for (median, provenance_value, expected) in [ + (1.5, Some(trusted.clone()), "PASS"), + (1.0, Some(trusted.clone()), "FAIL"), + (1.0, None, "ADVISORY"), + ] { + let mut parity = parity_with_provenance(runtime_ratio_stats(2.0, median), trusted.clone()); + parity.wall_time_provenance = provenance_value; + let snapshot = + current_msl_quality_snapshot_json(&summary, Some(&parity), Some(&baseline), false) + .expect("snapshot should serialize"); + let decision = &snapshot["runtime_wall_decision"]; + assert_eq!(decision["status"], expected); + assert_eq!(decision["observed_median"], median); + assert_eq!(decision["baseline_median"], 2.0); + assert_eq!(decision["floor"], 1.3); + assert!(snapshot.get("wall_time_provenance").is_some()); + } + let promoted = serde_json::to_value(&baseline).expect("serialize promoted baseline"); + assert!(promoted.get("wall_time_provenance").is_none()); + assert!(promoted.get("runtime_wall_decision").is_none()); +} + +#[test] +fn missing_runtime_baseline_is_untrusted_in_status_and_snapshot() { + let summary = valid_summary_template(); + let baseline = MslQualityBaseline { + runtime_context: Some(MslParityRuntimeContext { + workers_used: Some(2), + omc_threads: Some(1), + }), + runtime_ratio_stats: None, + ..baseline_quality_template() + }; + let mut trusted = provenance(8, 0, 14, 14, 0, Some(0.2), Some(0.3)); + trusted.workers_used = 2; + let parity = parity_with_provenance(runtime_ratio_stats(2.0, 1.5), trusted); + + let content = wall_time_status_content(&baseline, Some(&parity)); + assert_eq!(content.status, "ADVISORY"); + assert!(!content.trusted); + assert_eq!(content.reasons, ["runtime baseline missing"]); + assert!(format_wall_time_status(&content).contains("runtime baseline missing")); + + let snapshot = + current_msl_quality_snapshot_json(&summary, Some(&parity), Some(&baseline), false) + .expect("snapshot should serialize"); + let decision = &snapshot["runtime_wall_decision"]; + assert_eq!(decision["status"], "ADVISORY"); + assert_eq!(decision["trusted"], false); + assert_eq!(decision["reasons"], json!(["runtime baseline missing"])); + assert!(decision["baseline_median"].is_null()); + assert!(decision["floor"].is_null()); +} diff --git a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/wall_time.rs b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/wall_time.rs new file mode 100644 index 000000000..f17094e99 --- /dev/null +++ b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/wall_time.rs @@ -0,0 +1,219 @@ +use super::*; + +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +pub(super) struct MslWallTimeProvenance { + pub(super) omc_fresh_sample_count: usize, + pub(super) omc_cached_sample_count: usize, + pub(super) affinity_requested_worker_count: usize, + pub(super) affinity_applied_worker_count: usize, + pub(super) affinity_failed_worker_count: usize, + pub(super) normalized_load_before: Option, + pub(super) normalized_load_after: Option, + pub(super) rumoca_workers_used: usize, + pub(super) workers_used: usize, + pub(super) omc_threads: usize, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(super) struct WallTimeTrustDecision { + pub(super) trusted: bool, + pub(super) reasons: Vec, +} + +pub(super) fn wall_time_trust_decision( + baseline: &MslQualityBaseline, + parity_input: Option<&MslParityGateInput>, +) -> WallTimeTrustDecision { + let Some(expected_context) = complete_runtime_context(baseline.runtime_context.as_ref()) else { + return advisory("baseline runtime context missing; no wall-time policy comparator"); + }; + let Some(parity) = parity_input else { + return advisory("missing parity input"); + }; + let Some(provenance) = parity.wall_time_provenance.as_ref() else { + return advisory("missing provenance for wall time"); + }; + + let mut reasons = Vec::new(); + push_sample_count_reasons(&mut reasons, parity, provenance); + push_affinity_reasons(&mut reasons, provenance); + push_load_reason(&mut reasons, "before", provenance.normalized_load_before); + push_load_reason(&mut reasons, "after", provenance.normalized_load_after); + if !runtime_context_matches(expected_context, parity, provenance) { + reasons.push(format!( + "baseline runtime context mismatch: baseline={:?}, current={:?}, provenance=workers:{},omc_threads:{}", + baseline.runtime_context, + parity.runtime_context, + provenance.workers_used, + provenance.omc_threads + )); + } + + WallTimeTrustDecision { + trusted: reasons.is_empty(), + reasons, + } +} + +fn push_sample_count_reasons( + reasons: &mut Vec, + parity: &MslParityGateInput, + provenance: &MslWallTimeProvenance, +) { + let Some(wall_sample_count) = parity + .runtime_ratio_stats + .as_ref() + .map(|stats| stats.wall_ratio_both_success.sample_count) + else { + reasons.push("missing wall-time runtime samples".to_string()); + return; + }; + if provenance.omc_fresh_sample_count == 0 { + reasons.push("no fresh OMC wall-time samples".to_string()); + } + if provenance.omc_cached_sample_count > 0 { + reasons.push(format!( + "cached OMC wall-time samples present: {}", + provenance.omc_cached_sample_count + )); + } + let classified_samples = provenance + .omc_fresh_sample_count + .checked_add(provenance.omc_cached_sample_count); + if classified_samples != Some(wall_sample_count) { + reasons.push(format!( + "wall-time sample count mismatch: fresh={} + cached={} != compared={wall_sample_count}", + provenance.omc_fresh_sample_count, provenance.omc_cached_sample_count + )); + } +} + +fn advisory(reason: &str) -> WallTimeTrustDecision { + WallTimeTrustDecision { + trusted: false, + reasons: vec![reason.to_string()], + } +} + +fn push_affinity_reasons(reasons: &mut Vec, provenance: &MslWallTimeProvenance) { + if provenance.affinity_requested_worker_count != provenance.rumoca_workers_used { + reasons.push(format!( + "affinity coverage mismatch: requested={} rumoca_workers_used={}", + provenance.affinity_requested_worker_count, provenance.rumoca_workers_used + )); + } + if provenance.affinity_requested_worker_count == 0 + || provenance.affinity_applied_worker_count != provenance.affinity_requested_worker_count + || provenance.affinity_failed_worker_count > 0 + { + reasons.push(format!( + "affinity not fully applied: requested={} applied={} failed={}", + provenance.affinity_requested_worker_count, + provenance.affinity_applied_worker_count, + provenance.affinity_failed_worker_count + )); + } +} + +#[derive(Debug, Clone, Serialize, PartialEq)] +pub(super) struct WallTimeStatusContent { + pub(super) status: &'static str, + pub(super) trusted: bool, + pub(super) reasons: Vec, + pub(super) observed_median: Option, + pub(super) baseline_median: Option, + pub(super) floor: Option, +} + +pub(super) fn wall_time_status_content( + baseline: &MslQualityBaseline, + parity_input: Option<&MslParityGateInput>, +) -> WallTimeStatusContent { + let observed = parity_input + .and_then(|parity| parity.runtime_ratio_stats.as_ref()) + .map(|stats| stats.wall_ratio_both_success.median); + let baseline_median = baseline + .runtime_ratio_stats + .as_ref() + .map(|stats| stats.wall_ratio_both_success.median); + let Some(baseline_median) = baseline_median else { + return WallTimeStatusContent { + status: "ADVISORY", + trusted: false, + reasons: vec!["runtime baseline missing".to_string()], + observed_median: observed, + baseline_median: None, + floor: None, + }; + }; + let floor = Some(baseline_median * (1.0 - RUNTIME_RATIO_MEDIAN_REL_TOLERANCE)); + let trust = wall_time_trust_decision(baseline, parity_input); + let status = match (trust.trusted, observed, floor) { + (true, Some(observed), Some(floor)) if observed + SIM_RATE_GATE_EPSILON < floor => "FAIL", + (true, Some(_), Some(_)) => "PASS", + _ => "ADVISORY", + }; + WallTimeStatusContent { + status, + trusted: trust.trusted, + reasons: trust.reasons, + observed_median: observed, + baseline_median: Some(baseline_median), + floor, + } +} + +pub(super) fn format_wall_time_status(content: &WallTimeStatusContent) -> String { + let reason_text = if content.reasons.is_empty() { + "trusted provenance".to_string() + } else { + content.reasons.join("; ") + }; + match ( + content.observed_median, + content.baseline_median, + content.floor, + ) { + (Some(observed), Some(baseline), Some(floor)) => format!( + "MSL wall speed gate: {} median={observed:.3e}, baseline={baseline:.3e}, floor={floor:.3e} (tolerance={:.1}%); provenance: {reason_text}.", + content.status, + RUNTIME_RATIO_MEDIAN_REL_TOLERANCE * 100.0, + ), + _ => format!( + "MSL wall speed gate: {}; provenance: {reason_text}.", + content.status + ), + } +} + +fn push_load_reason(reasons: &mut Vec, phase: &str, load: Option) { + match load { + None => reasons.push(format!("missing load {phase}")), + Some(load) if !load.is_finite() || load < 0.0 => { + reasons.push(format!("invalid load {phase}: {load}")); + } + Some(load) if load > WALL_TIME_NORMALIZED_LOAD_MAX => reasons.push(format!( + "excessive load {phase}: {load:.3} > {WALL_TIME_NORMALIZED_LOAD_MAX:.3}" + )), + Some(_) => {} + } +} + +fn complete_runtime_context(context: Option<&MslParityRuntimeContext>) -> Option<(usize, usize)> { + let context = context?; + Some((context.workers_used?, context.omc_threads?)) +} + +fn runtime_context_matches( + expected: (usize, usize), + parity: &MslParityGateInput, + provenance: &MslWallTimeProvenance, +) -> bool { + expected.0 > 0 + && expected.1 > 0 + && parity.runtime_context.as_ref().is_some_and(|context| { + context.workers_used == Some(expected.0) && context.omc_threads == Some(expected.1) + }) + && provenance.workers_used == expected.0 + && provenance.omc_threads == expected.1 +} diff --git a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_render_sim.rs b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_render_sim.rs index a190935be..1a7e85872 100644 --- a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_render_sim.rs +++ b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_render_sim.rs @@ -93,15 +93,15 @@ fn simulation_settings_from_parts( .filter(|seconds| seconds.is_finite()) .unwrap_or(0.0); let mut t_end = experiment_stop_time - .filter(|seconds| seconds.is_finite() && *seconds > t_start) + .filter(|seconds| seconds.is_finite() && *seconds >= t_start) .unwrap_or(t_start + DEFAULT_SIM_END_TIME_SECS); - if !t_start.is_finite() || !t_end.is_finite() || t_end <= t_start { + if !t_start.is_finite() || !t_end.is_finite() || t_end < t_start { t_start = 0.0; t_end = DEFAULT_SIM_END_TIME_SECS; } if let Some(stop_time) = simulation_stop_time_override() - && stop_time > t_start + && stop_time >= t_start { t_end = stop_time; } @@ -467,3 +467,16 @@ pub(super) fn begin_chunked_render_sim_setup( ) -> RenderSimSetup { RenderSimSetup::new_from_compile_scope(compile_scope_names, run_simulation) } + +#[cfg(test)] +mod simulation_settings_tests { + use super::simulation_settings_from_parts; + + #[test] + fn preserves_zero_duration_experiment_horizon() { + let settings = simulation_settings_from_parts(Some(0.0), Some(0.0), None, None, None); + + assert_eq!(settings.t_start, 0.0); + assert_eq!(settings.t_end, 0.0); + } +} diff --git a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_selection.rs b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_selection.rs index 68c379b90..248de4874 100644 --- a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_selection.rs +++ b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_selection.rs @@ -1,4 +1,5 @@ use super::*; +#[cfg(target_os = "linux")] use indexmap::IndexSet; // ============================================================================= diff --git a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_stats_report.rs b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_stats_report.rs index e253ccff2..514188397 100644 --- a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_stats_report.rs +++ b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_stats_report.rs @@ -486,7 +486,8 @@ fn last_stage_seconds(report: &MslParityTimingReport) -> f64 { fn run_simulation_parity_stages(summary: &MslSummary, report: &mut MslParityTimingReport) { let parity_start = Instant::now(); - let _parity_watchdog = StageAbortWatchdog::new("parity_stage", 7200); + let _parity_watchdog = + StageAbortWatchdog::new("parity_stage", OMC_SIM_REFERENCE_STAGE_TIMEOUT_SECONDS); println!("MSL parity stage: ensuring OMC references + trace comparison..."); run_timed_parity_stage_or_panic( report, @@ -876,6 +877,9 @@ mod tests { requested_worker_threads: 16, effective_worker_threads: 14, worker_count: 14, + affinity_requested_worker_count: 14, + affinity_applied_worker_count: 14, + affinity_failed_worker_count: 0, pinned_worker_count: 14, cpu_token_capacity: 0, compile_memory_token_capacity_mb: Some(24_000), diff --git a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_summary.rs b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_summary.rs index 6403773df..e893c27e5 100644 --- a/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_summary.rs +++ b/crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_summary.rs @@ -35,18 +35,28 @@ pub(super) fn summarize_msl_results(results: &[MslModelResult]) -> ResultCounter _ => {} } } - if let Some(ref status) = result.ic_status { - counters.ic_attempted += 1; - match status.as_str() { - "ic_ok" => counters.ic_ok += 1, - "ic_solver_fail" => counters.ic_solver_fail += 1, - _ => {} - } - } + process_initial_condition_status(result, &mut counters); } counters } +fn process_initial_condition_status(result: &MslModelResult, counters: &mut ResultCounters) { + if let Some(ref status) = result.ic_status { + counters.ic_attempted += 1; + match status.as_str() { + "ic_ok" => counters.ic_ok += 1, + "ic_solver_fail" => counters.ic_solver_fail += 1, + _ => {} + } + return; + } + + if result.sim_status.as_deref() == Some("sim_ok") { + counters.ic_attempted += 1; + counters.ic_ok += 1; + } +} + pub(super) fn finalize_msl_summary_from_results( model_results: Vec, sim_target_models: Vec, @@ -303,4 +313,24 @@ VmHWM:\t 524288 kB Some(&1) ); } + + #[test] + fn summarize_msl_results_counts_sim_ok_as_initial_condition_success() { + let mut result = phase_error_result( + "Modelica.Test.DiscreteController".to_string(), + "Success", + None, + None, + ); + result.sim_status = Some("sim_ok".to_string()); + result.ic_status = None; + + let counters = summarize_msl_results(&[result]); + + assert_eq!(counters.sim_attempted, 1); + assert_eq!(counters.sim_ok, 1); + assert_eq!(counters.ic_attempted, 1); + assert_eq!(counters.ic_ok, 1); + assert_eq!(counters.ic_solver_fail, 0); + } } diff --git a/crates/rumoca-test-msl/tests/casadi_msl_test.rs b/crates/rumoca-test-msl/tests/casadi_msl_test.rs index 2cebc5cb3..fac3fd33c 100644 --- a/crates/rumoca-test-msl/tests/casadi_msl_test.rs +++ b/crates/rumoca-test-msl/tests/casadi_msl_test.rs @@ -262,6 +262,7 @@ fn casadi_simulate( Ok(SimTrace { model_name: Some(model_name.to_string()), + n_states: None, times, names, data, @@ -272,6 +273,7 @@ fn casadi_simulate( fn sim_result_to_trace(sim: &SimResult, model_name: &str) -> SimTrace { SimTrace { model_name: Some(model_name.to_string()), + n_states: Some(sim.n_states), times: sim.times.clone(), names: sim.names.clone(), data: sim diff --git a/crates/rumoca-test-msl/tests/msl_tests.rs b/crates/rumoca-test-msl/tests/msl_tests.rs index 6f40e835d..6bf22bb40 100644 --- a/crates/rumoca-test-msl/tests/msl_tests.rs +++ b/crates/rumoca-test-msl/tests/msl_tests.rs @@ -562,6 +562,12 @@ struct MslSchedulerTimings { #[serde(default)] worker_count: usize, #[serde(default)] + affinity_requested_worker_count: usize, + #[serde(default)] + affinity_applied_worker_count: usize, + #[serde(default)] + affinity_failed_worker_count: usize, + #[serde(default)] pinned_worker_count: usize, #[serde(default)] cpu_token_capacity: usize, @@ -587,6 +593,8 @@ struct MslSchedulerTimings { #[derive(Debug, Clone, Default, Serialize, Deserialize)] struct MslPhaseTimings { + #[serde(default, skip_serializing_if = "Option::is_none")] + host_load_before: Option, parse_seconds: f64, session_build_seconds: f64, frontend_compile_seconds: f64, diff --git a/crates/rumoca-test-msl/tests/msl_tests/README.md b/crates/rumoca-test-msl/tests/msl_tests/README.md index ee5557432..5955a7b3d 100644 --- a/crates/rumoca-test-msl/tests/msl_tests/README.md +++ b/crates/rumoca-test-msl/tests/msl_tests/README.md @@ -51,6 +51,15 @@ This directory contains helper includes for `tests/msl_tests.rs`. source tag before running this separate gate. - Focused subset controls (`RUMOCA_MSL_SIM_MATCH`, `RUMOCA_MSL_SIM_LIMIT`) are for iterative simulation work and must not be treated as baseline runs. +- Full root-example runs require OMC parity and fail closed when OMC is + unavailable, the current OMC reference is missing or stale, or it contains no + comparable runtime and trace metrics. A passing full gate therefore requires + non-empty system and wall runtime samples plus at least one compared trace + with model bucket percentages; compile/simulation counts do not substitute + for parity. +- Explicit focused/partial runs keep their existing skip behavior. Their output + states that OMC parity is skipped because the run is focused/partial and does + not imply that the full MSL quality gate passed. - By default, the pipeline uses the root-example baseline scope when no target environment override is provided. The committed explicit target file can still be selected with `RUMOCA_MSL_TARGET_SCOPE=committed-targets` for local triage, @@ -92,6 +101,12 @@ This directory contains helper includes for `tests/msl_tests.rs`. - Baseline JSON also captures OMC parity distributions for this set (runtime speedup ratio + trace-accuracy min/median/mean/max), populated from `omc_simulation_reference.json`. +- OMC simulation reference generation checkpoints after each model into + `target/msl/results/omc_simulation_reference.json`. CI preserves that JSON, + `omc_sim_work`, `sim_traces/omc`, and `omc_parity_cache` so a progressing + full-library reference run can resume across GitHub-hosted job attempts + without dropping models, loosening per-model timeouts, or skipping trace + parity. - Trace parity excludes known stochastic random-input examples listed in: - `tests/msl_tests/msl_trace_compare_exclusions.json` - these models remain in compile/balance/sim stats, but are skipped from diff --git a/crates/rumoca-test-msl/tests/msl_tests/msl_quality_baseline.json b/crates/rumoca-test-msl/tests/msl_tests/msl_quality_baseline.json index 370f0933f..bb17d0f90 100644 --- a/crates/rumoca-test-msl/tests/msl_tests/msl_quality_baseline.json +++ b/crates/rumoca-test-msl/tests/msl_tests/msl_quality_baseline.json @@ -171,14 +171,14 @@ "error_code_counts": {} }, "CLK_SM": { - "models": 77, - "compiled": 77, + "models": 74, + "compiled": 74, "solve_ir": 70, "balanced": 70, "sim_ok": 56, "phase_counts": { "Success": 70, - "ToDae": 7 + "ToDae": 4 }, "error_code_counts": {} }, @@ -206,13 +206,14 @@ "error_code_counts": {} }, "FUNC": { - "models": 1, - "compiled": 1, + "models": 4, + "compiled": 4, "solve_ir": 1, "balanced": 1, "sim_ok": 0, "phase_counts": { - "Success": 1 + "Success": 1, + "ToDae": 3 }, "error_code_counts": {} }, diff --git a/crates/rumoca-worker/src/bin/rumoca-worker.rs b/crates/rumoca-worker/src/bin/rumoca-worker.rs index 8e028e122..4abc70e72 100644 --- a/crates/rumoca-worker/src/bin/rumoca-worker.rs +++ b/crates/rumoca-worker/src/bin/rumoca-worker.rs @@ -17,8 +17,10 @@ use rumoca_compile::compile::{ }; use rumoca_sim::{ BuildSimulationTimings, PreparedSimulation, SimError, SimOptions, SimResult, SimSolverMode, + boundary_reduced_dae_for_simulation_artifact, build_simulation_with_stage_timing_and_solve_model, check_prepared_initialization, run_prepared_simulation, structurally_lowered_dae_for_simulation_artifact, + structurally_prepared_dae_for_simulation_artifact, }; use rumoca_worker::{ MODEL_WORKER_PARTIAL_RESULT_FILE, MODEL_WORKER_PROTOCOL_VERSION, MODEL_WORKER_RESULT_FILE, @@ -579,30 +581,6 @@ fn strict_dae_failure_phase(failure_summary: &str) -> &'static str { .unwrap_or("ToDae") } -fn initialization_balance_check( - dae: &rumoca_compile::compile::Dae, - scalar_unknowns: i64, - scalar_equations: i64, -) -> (i64, i64, i64, i64, i64) { - let deficit_before = (scalar_unknowns - scalar_equations).max(0); - let initial_equation_scalars = dae - .initialization - .equations - .iter() - .map(|eq| eq.scalar_count as i64) - .sum::(); - let initial_algorithm_scalars = 0; - let closure_used = (initial_equation_scalars + initial_algorithm_scalars).min(deficit_before); - let deficit_after = deficit_before - closure_used; - ( - deficit_before, - initial_equation_scalars, - initial_algorithm_scalars, - closure_used, - deficit_after, - ) -} - fn summarize_dae_success( model_name: &str, result: &rumoca_compile::compile::DaeCompilationResult, @@ -610,16 +588,12 @@ fn summarize_dae_success( ) -> WorkerModelResult { let detail = &result.balance_detail; let (scalar_equations, scalar_unknowns) = detail.equations_unknowns(); + let closure_detail = + rumoca_compile::analysis::initial_closure_balance_detail(result.dae.as_ref()) + .expect("successful DAE compilation has valid balance metadata"); let scalar_equations = scalar_equations as i64; let scalar_unknowns = scalar_unknowns as i64; - let ( - deficit_before, - initial_equation_scalars, - initial_algorithm_scalars, - closure_used, - deficit_after, - ) = initialization_balance_check(result.dae.as_ref(), scalar_unknowns, scalar_equations); - let scalar_equations_with_init = scalar_equations + closure_used; + let scalar_equations_with_init = scalar_equations + closure_detail.closure_used; let input_scalars = result .dae .variables @@ -647,12 +621,12 @@ fn summarize_dae_success( row.class_type = Some(result.dae.metadata.class_type.as_str().to_string()); row.scalar_equations = usize::try_from(scalar_equations_for_report).ok(); row.scalar_unknowns = usize::try_from(scalar_unknowns_for_report).ok(); - row.initial_equation_scalars = usize::try_from(initial_equation_scalars).ok(); - row.initial_algorithm_scalars = usize::try_from(initial_algorithm_scalars).ok(); - row.initial_balance_deficit_before = Some(deficit_before); - row.initial_closure_used = usize::try_from(closure_used).ok(); - row.initial_balance_deficit_after = Some(deficit_after); - row.initial_balance_ok = Some(deficit_after == 0); + row.initial_equation_scalars = usize::try_from(closure_detail.initial_equation_scalars).ok(); + row.initial_algorithm_scalars = usize::try_from(closure_detail.initial_algorithm_scalars).ok(); + row.initial_balance_deficit_before = Some(closure_detail.deficit_before); + row.initial_closure_used = usize::try_from(closure_detail.closure_used).ok(); + row.initial_balance_deficit_after = Some(closure_detail.deficit_after); + row.initial_balance_ok = Some(closure_detail.deficit_after == 0); row.compile_seconds = Some(compile_seconds); row } @@ -662,18 +636,8 @@ fn sim_timeout_secs() -> f64 { } fn simulation_settings(result: &rumoca_compile::compile::DaeCompilationResult) -> SimSettings { - let mut t_start = result - .experiment_start_time - .filter(|seconds| seconds.is_finite()) - .unwrap_or(0.0); - let mut t_end = result - .experiment_stop_time - .filter(|seconds| seconds.is_finite() && *seconds > t_start) - .unwrap_or(t_start + DEFAULT_SIM_END_TIME_SECS); - if !t_start.is_finite() || !t_end.is_finite() || t_end <= t_start { - t_start = 0.0; - t_end = DEFAULT_SIM_END_TIME_SECS; - } + let (t_start, t_end) = + simulation_horizon(result.experiment_start_time, result.experiment_stop_time); let tolerance = result .experiment_tolerance .filter(|value| value.is_finite() && *value > 0.0); @@ -693,6 +657,23 @@ fn simulation_settings(result: &rumoca_compile::compile::DaeCompilationResult) - } } +fn simulation_horizon( + experiment_start_time: Option, + experiment_stop_time: Option, +) -> (f64, f64) { + let mut t_start = experiment_start_time + .filter(|seconds| seconds.is_finite()) + .unwrap_or(0.0); + let mut t_end = experiment_stop_time + .filter(|seconds| seconds.is_finite() && *seconds >= t_start) + .unwrap_or(t_start + DEFAULT_SIM_END_TIME_SECS); + if !t_start.is_finite() || !t_end.is_finite() || t_end < t_start { + t_start = 0.0; + t_end = DEFAULT_SIM_END_TIME_SECS; + } + (t_start, t_end) +} + fn root_standalone_example_name(model_name: &str) -> bool { if !model_name.starts_with("Modelica.") || !model_name.contains(".Examples.") { return false; @@ -991,9 +972,18 @@ fn initial_structural_dae_artifact_error( if !request.emit_json && !request.emit_modelica { return None; } + let mut error = match structurally_prepared_dae_for_simulation_artifact(dae, opts) { + Ok(prepared_dae) => write_prepared_structural_dae_artifacts(request, &prepared_dae), + Err(error) => Some(error.to_string()), + }; + error = error.or( + match boundary_reduced_dae_for_simulation_artifact(dae, opts) { + Ok(reduced_dae) => write_boundary_reduced_dae_artifacts(request, &reduced_dae), + Err(error) => Some(error.to_string()), + }, + ); match structurally_lowered_dae_for_simulation_artifact(dae, opts) { Ok(structural_dae) => { - let mut error = None; if request.emit_modelica { error = error.or(write_modelica_dae_artifact( request, @@ -1016,6 +1006,48 @@ fn initial_structural_dae_artifact_error( } } +fn write_boundary_reduced_dae_artifacts( + request: &ModelWorkerRequest, + reduced_dae: &rumoca_compile::compile::Dae, +) -> Option { + let mut error = None; + if request.emit_modelica { + error = error.or(write_modelica_dae_artifact( + request, + "ir-boundary-reduced-dae.mo", + reduced_dae, + ) + .err()); + } + if request.emit_json { + error = error + .or(write_artifact_json(request, "ir-boundary-reduced-dae.json", reduced_dae).err()); + } + error +} + +fn write_prepared_structural_dae_artifacts( + request: &ModelWorkerRequest, + prepared_dae: &rumoca_compile::compile::Dae, +) -> Option { + let mut error = None; + if request.emit_modelica { + error = error.or(write_modelica_dae_artifact( + request, + "ir-prepared-structural-dae.mo", + prepared_dae, + ) + .err()); + } + if request.emit_json { + error = + error.or( + write_artifact_json(request, "ir-prepared-structural-dae.json", prepared_dae).err(), + ); + } + error +} + fn observe_simulation_build_stage( stage: &str, progress: &ProgressLog, @@ -1386,10 +1418,14 @@ fn write_control_message(message: &ModelWorkerControlMessage) -> Result<(), Stri .map_err(|error| format!("failed to flush model worker control message: {error}")) } -fn run_worker_daemon(source_root_path: &Path) -> Result<(), String> { +fn run_worker_daemon( + source_root_path: &Path, + cpu_affinity_applied: Option, +) -> Result<(), String> { let mut session = load_source_root(source_root_path)?; write_control_message(&ModelWorkerControlMessage::Ready { protocol_version: MODEL_WORKER_PROTOCOL_VERSION, + cpu_affinity_applied, })?; for line in std::io::stdin().lock().lines() { let line = line.map_err(|error| format!("failed to read model worker command: {error}"))?; @@ -1419,12 +1455,14 @@ fn run_worker_daemon(source_root_path: &Path) -> Result<(), String> { } fn run_worker_entry(args: Args) -> Result<(), String> { - if let Some(cpu_core_id) = args.cpu_core_id { - pin_current_thread_to_cpu_core(cpu_core_id)?; - } + let cpu_affinity_applied = args.cpu_core_id.map(|cpu_core_id| { + pin_current_thread_to_cpu_core(cpu_core_id) + .inspect_err(|error| eprintln!("warning: {error}; continuing without CPU pinning")) + .is_ok() + }); match (args.request_json.as_ref(), args.source_root_path.as_ref()) { (Some(_), None) => run_worker(args), - (None, Some(source_root_path)) => run_worker_daemon(source_root_path), + (None, Some(source_root_path)) => run_worker_daemon(source_root_path, cpu_affinity_applied), (Some(_), Some(_)) => Err("use either --request-json or --source-root-path".to_string()), (None, None) => Err("missing --request-json or --source-root-path".to_string()), } @@ -1449,3 +1487,13 @@ fn main() { std::process::exit(1); } } + +#[cfg(test)] +mod tests { + use super::simulation_horizon; + + #[test] + fn simulation_horizon_preserves_zero_duration_experiment() { + assert_eq!(simulation_horizon(Some(0.0), Some(0.0)), (0.0, 0.0)); + } +} diff --git a/crates/rumoca-worker/src/lib.rs b/crates/rumoca-worker/src/lib.rs index e7e2935b2..77053cd09 100644 --- a/crates/rumoca-worker/src/lib.rs +++ b/crates/rumoca-worker/src/lib.rs @@ -55,6 +55,8 @@ pub enum ModelWorkerCommand { pub enum ModelWorkerControlMessage { Ready { protocol_version: u32, + #[serde(default)] + cpu_affinity_applied: Option, }, Result { response: Box, @@ -206,6 +208,7 @@ pub struct ModelWorkerDaemon { child: Child, stdin: ChildStdin, messages: mpsc::Receiver>, + cpu_affinity_applied: Option, } impl ModelWorkerDaemon { @@ -243,6 +246,7 @@ impl ModelWorkerDaemon { child, stdin, messages, + cpu_affinity_applied: None, }; match worker.wait_for_ready(timeout_secs) { Ok(()) => Ok(worker), @@ -257,6 +261,10 @@ impl ModelWorkerDaemon { self.child.id() } + pub fn cpu_affinity_applied(&self) -> Option { + self.cpu_affinity_applied + } + pub fn run_request( &mut self, request: &ModelWorkerRequest, @@ -345,12 +353,16 @@ impl ModelWorkerDaemon { let start = Instant::now(); loop { match self.messages.try_recv() { - Ok(Ok(ModelWorkerControlMessage::Ready { protocol_version })) - if protocol_version == MODEL_WORKER_PROTOCOL_VERSION => - { + Ok(Ok(ModelWorkerControlMessage::Ready { + protocol_version, + cpu_affinity_applied, + })) if protocol_version == MODEL_WORKER_PROTOCOL_VERSION => { + self.cpu_affinity_applied = cpu_affinity_applied; return Ok(()); } - Ok(Ok(ModelWorkerControlMessage::Ready { protocol_version })) => { + Ok(Ok(ModelWorkerControlMessage::Ready { + protocol_version, .. + })) => { return Err(format!( "model worker ready protocol {}, expected {}", protocol_version, MODEL_WORKER_PROTOCOL_VERSION @@ -828,6 +840,50 @@ fn write_json_file(path: &std::path::Path, value: &T) -> Result<() mod tests { use super::*; + #[test] + fn ready_message_preserves_affinity_result() { + let ready = ModelWorkerControlMessage::Ready { + protocol_version: MODEL_WORKER_PROTOCOL_VERSION, + cpu_affinity_applied: Some(false), + }; + let encoded = serde_json::to_string(&ready).unwrap(); + let decoded: ModelWorkerControlMessage = serde_json::from_str(&encoded).unwrap(); + assert!(matches!( + decoded, + ModelWorkerControlMessage::Ready { + cpu_affinity_applied: Some(false), + .. + } + )); + } + + #[test] + fn ready_message_allows_unrequested_affinity() { + let ready = ModelWorkerControlMessage::Ready { + protocol_version: MODEL_WORKER_PROTOCOL_VERSION, + cpu_affinity_applied: None, + }; + assert!( + serde_json::to_string(&ready) + .unwrap() + .contains("cpu_affinity_applied") + ); + } + + #[test] + fn ready_message_deserializes_old_payload_without_affinity() { + let decoded: ModelWorkerControlMessage = + serde_json::from_str(r#"{"event":"ready","protocol_version":1}"#) + .expect("old ready payload should deserialize"); + assert!(matches!( + decoded, + ModelWorkerControlMessage::Ready { + cpu_affinity_applied: None, + .. + } + )); + } + #[test] fn cpu_core_plan_has_one_entry_per_worker() { assert_eq!(cpu_core_plan(0), Vec::>::new()); diff --git a/crates/rumoca/src/cli.rs b/crates/rumoca/src/cli.rs index 181fba20b..99c5dc090 100644 --- a/crates/rumoca/src/cli.rs +++ b/crates/rumoca/src/cli.rs @@ -893,10 +893,11 @@ fn run_configured_simulation(args: SimCommandArgs) -> Result<()> { return run_simulation(SimulationRun { dae: result.dae.as_ref(), model: &compiled_model, + t_start: SimOptions::default().t_start, t_end: configured_sim_t_end(args.t_end, config.sim.t_end), dt: Some(configured_sim_dt(args.dt, config.sim.dt)), - atol: configured_sim_option(args.atol, config.sim.atol), - rtol: configured_sim_option(args.rtol, config.sim.rtol), + atol: configured_sim_tolerance(args.atol, config.sim.atol, SimOptions::default().atol), + rtol: configured_sim_tolerance(args.rtol, config.sim.rtol, SimOptions::default().rtol), solver_mode, solver_label: &solver_label, output: args.output.as_deref().or(config.sim.output.as_deref()), @@ -957,6 +958,34 @@ fn configured_sim_option(cli_value: Option, config_value: Option) -> O cli_value.or(config_value) } +fn configured_sim_tolerance( + cli_value: Option, + config_value: Option, + default: f64, +) -> f64 { + cli_value.or(config_value).unwrap_or(default) +} + +#[cfg(test)] +mod configured_sim_tolerance_tests { + use super::configured_sim_tolerance; + + #[test] + fn configured_sim_tolerance_prefers_cli_value() { + assert_eq!(configured_sim_tolerance(Some(1.0), Some(2.0), 3.0), 1.0); + } + + #[test] + fn configured_sim_tolerance_uses_configured_value_without_cli_value() { + assert_eq!(configured_sim_tolerance(None, Some(2.0), 3.0), 2.0); + } + + #[test] + fn configured_sim_tolerance_uses_backend_default_without_overrides() { + assert_eq!(configured_sim_tolerance(None, None, 3.0), 3.0); + } +} + #[cfg(feature = "scheduled-sim")] fn run_config_check(args: SimCheckArgs) -> Result<()> { let config_path = args.config_path()?; @@ -1246,13 +1275,15 @@ fn run_direct_simulation(args: SimCommandArgs) -> Result<()> { } let workspace_root = discover_workspace_root_for_model_file(&input.model_file); let solver = simulate_solver_or_auto(args.solver, result.experiment_solver.as_deref()); + let sim_defaults = direct_sim_defaults(args.t_end, args.dt, args.atol, args.rtol, &result); run_simulation(SimulationRun { dae: result.dae.as_ref(), model: &model, - t_end: direct_sim_t_end(args.t_end), - dt: args.dt, - atol: args.atol, - rtol: args.rtol, + t_start: sim_defaults.t_start, + t_end: sim_defaults.t_end, + dt: sim_defaults.dt, + atol: sim_defaults.atol, + rtol: sim_defaults.rtol, solver_mode: solver.into(), solver_label: solver.as_label(), output: args.output.as_deref(), @@ -1287,8 +1318,42 @@ fn simulate_solver_or_auto( } } -fn direct_sim_t_end(t_end: Option) -> f64 { - t_end.unwrap_or(1.0) +#[derive(Debug, Clone, Copy)] +pub(crate) struct DirectSimDefaults { + pub t_start: f64, + pub t_end: f64, + pub dt: Option, + pub atol: f64, + pub rtol: f64, +} + +pub(crate) fn direct_sim_defaults( + t_end: Option, + dt: Option, + atol: Option, + rtol: Option, + result: &DaeCompilationResult, +) -> DirectSimDefaults { + let default_opts = SimOptions::default(); + let t_start = result.experiment_start_time.unwrap_or(default_opts.t_start); + DirectSimDefaults { + t_start, + t_end: t_end + .or(result.experiment_stop_time) + .unwrap_or(default_opts.t_end) + .max(t_start), + dt: dt.or_else(|| { + result + .experiment_interval + .filter(|value| value.is_finite() && *value > 0.0) + }), + atol: atol + .or(result.experiment_tolerance) + .unwrap_or(default_opts.atol), + rtol: rtol + .or(result.experiment_tolerance) + .unwrap_or(default_opts.rtol), + } } fn run_lint(args: LintArgs) -> Result<()> { @@ -1651,10 +1716,11 @@ fn print_summary(model: &str, result: &CompilationResult) { struct SimulationRun<'a> { dae: &'a Dae, model: &'a str, + t_start: f64, t_end: f64, dt: Option, - atol: Option, - rtol: Option, + atol: f64, + rtol: f64, solver_mode: SimSolverMode, solver_label: &'a str, output: Option<&'a str>, @@ -1674,23 +1740,18 @@ fn run_simulation(run: SimulationRun<'_>) -> Result<()> { ); } - let mut opts = SimOptions { + let opts = SimOptions { + t_start: run.t_start, t_end: run.t_end, dt: run.dt, + atol: run.atol, + rtol: run.rtol, solver_mode: run.solver_mode, // `--solver esdirk34` / `trbdf2` selects an implicit SDIRK tableau on // the diffsol path; other names leave the BDF default. diffsol_method: diffsol_method_for_solver_label(run.solver_label), ..SimOptions::default() }; - // Explicit --atol/--rtol override the backend default so a host's tolerance - // policy can be reproduced exactly from the CLI. - if let Some(atol) = run.atol { - opts.atol = atol; - } - if let Some(rtol) = run.rtol { - opts.rtol = rtol; - } eprintln!("Simulating {} to t={}...", run.model, run.t_end); // On a non-finite-suggestive failure (e.g. a model divide-by-zero showing up diff --git a/crates/rumoca/src/cli/value.rs b/crates/rumoca/src/cli/value.rs index e63daa2cc..b005ed5a6 100644 --- a/crates/rumoca/src/cli/value.rs +++ b/crates/rumoca/src/cli/value.rs @@ -22,7 +22,7 @@ use super::{ CompilationResult, CompileArgs, CompilePhase, EarlyIrArtifact, SimCommandArgs, SimOptions, SimulationRequestSummary, SimulationRunMetrics, TemplateIr, compile_str_dae_with_inferred_model, compile_str_early_ir_with_inferred_model, - compile_str_with_inferred_model, diffsol_method_for_solver_label, direct_sim_t_end, + compile_str_with_inferred_model, diffsol_method_for_solver_label, direct_sim_defaults, render_early_ir_as_modelica_ast, render_early_ir_as_modelica_flat, render_ir_as_modelica, simulate_solver_or_auto, target_manifest, }; @@ -154,20 +154,18 @@ pub fn simulate_to_value(args: &SimCommandArgs, source: &str) -> Result { )?; let compile_seconds = compile_started.elapsed().as_secs_f64(); let solver = simulate_solver_or_auto(args.solver, result.experiment_solver.as_deref()); - - let mut opts = SimOptions { - t_end: direct_sim_t_end(args.t_end), - dt: args.dt, + let sim_defaults = direct_sim_defaults(args.t_end, args.dt, args.atol, args.rtol, &result); + + let opts = SimOptions { + t_start: sim_defaults.t_start, + t_end: sim_defaults.t_end, + dt: sim_defaults.dt, + atol: sim_defaults.atol, + rtol: sim_defaults.rtol, solver_mode: solver.into(), diffsol_method: diffsol_method_for_solver_label(solver.as_label()), ..SimOptions::default() }; - if let Some(atol) = args.atol { - opts.atol = atol; - } - if let Some(rtol) = args.rtol { - opts.rtol = rtol; - } let sim_started = Instant::now(); // Dispatch on `opts.solver_mode` (auto / bdf / rk-like) exactly like the diff --git a/crates/rumoca/src/compiler.rs b/crates/rumoca/src/compiler.rs index 8e41f93a9..036556f13 100644 --- a/crates/rumoca/src/compiler.rs +++ b/crates/rumoca/src/compiler.rs @@ -1,5 +1,9 @@ //! High-level API for compiling Modelica models to DAE representations. //! +//! SPEC_0021 file-size exception: compiler facade tests still live beside the +//! public API while session facade migration is stabilizing. split plan: move +//! FMI/codegen and session regression tests into compiler submodules. +//! //! This module provides a clean, ergonomic interface for using rumoca as a library. //! The main entry point is the [`Compiler`] struct, which uses a builder pattern //! for configuration. @@ -55,7 +59,7 @@ use rumoca_compile::source_roots::{ referenced_unloaded_source_root_paths, render_source_root_status_message, resolve_source_root_cache_dir, source_root_source_set_key, }; -use rumoca_sim::{lower_solve_artifacts, lower_solve_problem}; +use rumoca_sim::lower_solve_problem; use serde_json::{Map, Value}; use crate::error::CompilerError; @@ -96,13 +100,17 @@ pub enum TemplateIr { } fn build_solve_template_renderer(dae_model: &Dae) -> Result { - let problem = lower_solve_problem(dae_model) - .map_err(|err| CompilerError::TemplateError(CodegenError::template(err.to_string())))?; - let artifacts = lower_solve_artifacts(&problem) - .map_err(|err| CompilerError::TemplateError(CodegenError::template(err.to_string())))?; let template_dae = dae_for_solve_template_context(dae_model)?; - SolveTemplateRenderer::new_with_dae(&problem, &artifacts, template_dae) - .map_err(CompilerError::TemplateError) + let solve_model = rumoca_sim::lower_dae_to_solve_model_owned(template_dae.clone()) + .map_err(|err| CompilerError::TemplateError(CodegenError::template(err.to_string())))?; + SolveTemplateRenderer::new_with_dae_and_visible_outputs( + &solve_model.problem, + &solve_model.artifacts, + template_dae, + solve_model.visible_names, + solve_model.visible_value_rows, + ) + .map_err(CompilerError::TemplateError) } fn build_solve_template_renderer_without_dae( @@ -439,6 +447,10 @@ impl CompilationResult { } } + pub fn scalarized_template_dae(&self) -> Dae { + self.dae.clone() + } + /// Equation balance (equations - unknowns). pub fn balance(&self) -> i64 { self.balance_detail.balance() @@ -504,6 +516,8 @@ pub struct Compiler { source_root_paths: Vec, /// Enable verbose output. verbose: bool, + /// Explicit compatibility mode for non-standard library annotations. + allow_non_param_evaluate_annotation: bool, } impl Compiler { @@ -546,6 +560,23 @@ impl Compiler { self } + /// Allow `annotation(Evaluate=true)` on non-parameter/non-constant components. + /// + /// Strict MLS behavior remains the default. This opt-in exists for external + /// libraries that carry this non-standard annotation on otherwise ordinary + /// connector/helper components. + pub fn allow_non_param_evaluate_annotation(mut self, allow: bool) -> Self { + self.allow_non_param_evaluate_annotation = allow; + self + } + + fn session_config(&self) -> SessionConfig { + SessionConfig { + evaluate_scope_is_error: !self.allow_non_param_evaluate_annotation, + ..SessionConfig::default() + } + } + /// Load a source-root path into the session. /// /// Handles both single files and directories recursively. @@ -721,6 +752,18 @@ impl Compiler { self.compile_str_dae(&source, path) } + /// Compile a Modelica file through DAE for diagnostics while retaining an + /// unbalanced DAE result. + pub fn compile_file_dae_allow_unbalanced_for_diagnostics( + &self, + path: &str, + ) -> Result { + let source = + fs::read_to_string(path).map_err(|e| CompilerError::io_error(path, e.to_string()))?; + + self.compile_str_dae_allow_unbalanced_for_diagnostics(&source, path) + } + /// Compile a Modelica file from a Path through DAE only. pub fn compile_path_dae(&self, path: &Path) -> Result { let path_str = path.to_string_lossy().to_string(); @@ -744,7 +787,7 @@ impl Compiler { } // Create a session and add the document - let mut session = Session::new(SessionConfig::default()); + let mut session = Session::new(self.session_config()); self.load_required_source_roots(&mut session, source)?; if self.verbose { @@ -827,7 +870,7 @@ impl Compiler { eprintln!("[rumoca] Phase 1-2: Parsing and resolving..."); } - let mut session = Session::new(SessionConfig::default()); + let mut session = Session::new(self.session_config()); self.load_required_source_roots(&mut session, source)?; self.load_local_compile_unit(&mut session, source, file_name)?; session @@ -852,7 +895,7 @@ impl Compiler { eprintln!("[rumoca] Phase 1-5: Parsing, resolving, and flattening..."); } - let mut session = Session::new(SessionConfig::default()); + let mut session = Session::new(self.session_config()); self.load_required_source_roots(&mut session, source)?; self.load_local_compile_unit(&mut session, source, file_name)?; session @@ -876,7 +919,7 @@ impl Compiler { eprintln!("[rumoca] Source file: {}", file_name); } - let mut session = Session::new(SessionConfig::default()); + let mut session = Session::new(self.session_config()); self.load_required_source_roots(&mut session, source)?; if self.verbose { @@ -927,6 +970,45 @@ impl Compiler { Ok(*result) } + + /// Compile Modelica source code through DAE for diagnostics while retaining + /// an unbalanced DAE result. + pub fn compile_str_dae_allow_unbalanced_for_diagnostics( + &self, + source: &str, + file_name: &str, + ) -> Result { + let model_name = self + .model_name + .as_ref() + .ok_or(CompilerError::NoModelSpecified)?; + + if self.verbose { + eprintln!("[rumoca] Compiling model through diagnostic DAE: {model_name}"); + eprintln!("[rumoca] Source file: {file_name}"); + } + + let mut session = Session::new(self.session_config()); + self.load_required_source_roots(&mut session, source)?; + + if self.verbose { + eprintln!("[rumoca] Phase 1-2: Parsing and resolving..."); + } + self.load_local_compile_unit(&mut session, source, file_name)?; + + if self.verbose { + eprintln!("[rumoca] Phase 3-6: Diagnostic DAE compile..."); + } + + session + .compile_model_dae_allow_unbalanced_for_diagnostics(model_name) + .map(|result| *result) + .map_err(|summary| CompilerError::CompileDiagnosticsError { + summary, + failures: Vec::new(), + source_map: None, + }) + } } #[cfg(test)] @@ -936,6 +1018,11 @@ mod tests { use quick_xml::events::{BytesStart, Event}; use tempfile::tempdir; + fn builtin_fmi2_template(path: &str) -> &'static str { + rumoca_phase_codegen::templates::builtin_template_source("fmi2", path) + .expect("FMI2 template should be registered") + } + #[test] fn test_simple_model() { let source = r#" @@ -1705,6 +1792,183 @@ mod tests { ); } + #[test] + fn test_render_fmi2_model_restores_output_observable_signs() { + let source = r#" + model OutputAlias + Real x(start = 2, fixed = true); + output Real y; + output Real negY; + equation + der(x) = -x; + y = x; + negY = -y; + end OutputAlias; + "#; + + let result = Compiler::new() + .model("OutputAlias") + .compile_str(source, "output_alias.mo") + .expect("compilation should succeed"); + let rendered = result + .render_template_str_with_name_and_ir( + builtin_fmi2_template("model.c.jinja"), + "OutputAlias", + TemplateIr::Solve, + ) + .expect("prepared named FMI2 model render should succeed"); + + assert!( + rendered.contains("y = x;"), + "expected output alias y to preserve positive x sign; got:\n{rendered}" + ); + assert!( + rendered.contains("negY = (-y);"), + "expected negY to remain the negative alias; got:\n{rendered}" + ); + } + + #[test] + fn test_render_fmi2_model_propagates_tunable_parameter_bindings_on_initialization() { + let source = r#" + model Child + parameter Real p = 1; + Real x(start = p, fixed = true); + initial equation + x = p; + equation + der(x) = 0; + end Child; + + model BindingProbe + parameter Real root = 5; + Child child(p = root); + end BindingProbe; + "#; + + let result = Compiler::new() + .model("BindingProbe") + .compile_str(source, "binding_probe.mo") + .expect("compilation should succeed"); + let rendered = result + .render_template_str_with_name_and_ir( + builtin_fmi2_template("model.c.jinja"), + "BindingProbe", + TemplateIr::Solve, + ) + .expect("prepared FMI2 C render should succeed"); + let rendered = rendered.replace("\r\n", "\n"); + + assert!( + rendered.contains("p = root;\n m->p[1] = p; /* binding child.p */"), + "FMI2 C must re-evaluate modifier-derived child.p from root after setReal; got:\n{rendered}" + ); + assert!( + rendered.contains("m->x[0] = p; /* initial equation: child.x */"), + "FMI2 C must apply direct initial equation x = p after parameter bindings; got:\n{rendered}" + ); + assert!( + rendered.contains("apply_parameter_bindings(m);\n apply_initial_equations(m);"), + "fmi2ExitInitializationMode must apply parameter bindings before initial equations; got:\n{rendered}" + ); + assert!( + rendered.contains( + "apply_parameter_bindings(m);\n m->dirty_values = 1;\n return fmi2OK;" + ), + "fmi2SetReal must re-apply dependent parameter bindings after batched writes; got:\n{rendered}" + ); + } + + #[test] + fn test_render_fmi2_model_description_resolves_symbolic_starts_to_literals() { + let source = r#" + model Child + parameter Real p = 1; + Real x(start = p, fixed = true); + initial equation + x = p; + equation + der(x) = 0; + end Child; + + model BindingProbe + parameter Real root = 5; + Child child(p = root); + end BindingProbe; + "#; + + let result = Compiler::new() + .model("BindingProbe") + .compile_str(source, "binding_probe.mo") + .expect("compilation should succeed"); + let rendered = result + .render_fmi_model_description_template_str_with_name( + builtin_fmi2_template("modelDescription.xml.jinja"), + "BindingProbe", + ) + .expect("prepared FMI2 modelDescription render should succeed"); + let rendered = rendered.replace("\r\n", "\n"); + + assert!( + rendered.contains(r#"name="root" valueReference="2" causality="parameter" variability="fixed" initial="exact"> + "#), + "numeric parameter starts should still be emitted; got:\n{rendered}" + ); + assert!( + rendered.contains(r#"name="child.p" valueReference="3" causality="parameter" variability="fixed" initial="exact"> + "#), + "resolvable symbolic parameter starts must be emitted as typed literals; got:\n{rendered}" + ); + assert!( + rendered.contains(r#"name="child.x" valueReference="0" causality="local" variability="continuous" initial="exact"> + "#), + "resolvable symbolic state starts must be emitted as typed literals; got:\n{rendered}" + ); + } + + #[test] + fn test_render_fmi2_model_preserves_array_slice_modifier_bindings() { + let source = r#" + model Child + parameter Real p = 1; + Real x(start = p, fixed = true); + initial equation + x = p; + equation + der(x) = 0; + end Child; + + model ArrayModifierProbe + parameter Real root[5] = {11, 12, 13, 14, 15}; + Child child[5](p = root[1:5]); + end ArrayModifierProbe; + "#; + + let result = Compiler::new() + .model("ArrayModifierProbe") + .compile_str(source, "array_modifier_probe.mo") + .expect("compilation should succeed"); + let rendered = result + .render_template_str_with_name_and_ir( + builtin_fmi2_template("model.c.jinja"), + "ArrayModifierProbe", + TemplateIr::Solve, + ) + .expect("prepared FMI2 C render should succeed"); + let rendered = rendered.replace("\r\n", "\n"); + + assert!( + rendered.contains( + "child_1_p = root_1;\n m->p[5] = child_1_p; /* binding child[1].p */" + ), + "array component modifier bindings must preserve element-wise source dependencies; got:\n{rendered}" + ); + assert!( + rendered.contains("m->x[0] = child_1_p; /* initial equation: child[1].x */"), + "initial equations must use the dependent array modifier parameter; got:\n{rendered}" + ); + } + #[test] fn test_strict_reachable_requested_success_ignores_unreachable_failures() { let source = r#" @@ -1733,6 +1997,35 @@ mod tests { ); } + #[test] + fn test_non_param_evaluate_annotation_requires_explicit_compatibility_opt_in() { + let source = r#" + model NonParamEvaluate + Real x(start = 1) annotation(Evaluate = true); + equation + der(x) = 0; + end NonParamEvaluate; + "#; + + let strict_err = Compiler::new() + .model("NonParamEvaluate") + .compile_str(source, "NonParamEvaluate.mo") + .expect_err("strict MLS mode must reject Evaluate=true on a non-parameter"); + let strict_message = strict_err.to_string(); + assert!( + strict_message.contains( + "annotation Evaluate is only allowed on parameter or constant components" + ), + "strict failure should preserve the Evaluate-scope diagnostic: {strict_message}" + ); + + Compiler::new() + .model("NonParamEvaluate") + .allow_non_param_evaluate_annotation(true) + .compile_str(source, "NonParamEvaluate.mo") + .expect("explicit compatibility opt-in should downgrade ER070 to a warning"); + } + #[test] fn test_strict_reachable_requested_failure_excludes_unreachable_context() { let source = r#" diff --git a/crates/rumoca/src/fmu.rs b/crates/rumoca/src/fmu.rs index 88791dc0b..4532f4c60 100644 --- a/crates/rumoca/src/fmu.rs +++ b/crates/rumoca/src/fmu.rs @@ -14,6 +14,22 @@ pub(crate) fn build_fmu( ) -> Result<()> { use std::process::Command; + let build_script = out_dir.join("build.sh"); + if build_script.is_file() { + eprintln!(" running {}", build_script.display()); + let status = Command::new("sh") + .arg("build.sh") + .current_dir(out_dir) + .status()?; + if !status.success() { + bail!( + "FMU build script failed with exit code {}", + status.code().unwrap_or(-1) + ); + } + return Ok(()); + } + let (platform, lib_ext) = fmu_binary_platform(target_name)?; // Compile shared library @@ -116,3 +132,25 @@ fn create_fmu_zip(out_dir: &Path, fmu_path: &Path) -> Result<()> { zip.finish()?; Ok(()) } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn build_fmu_runs_target_build_script_when_present() { + let dir = tempfile::tempdir().expect("temp fmu dir"); + let script = dir.path().join("build.sh"); + std::fs::write(&script, "printf ran > build-script-marker\n").expect("write build script"); + + build_fmu(dir.path(), "Demo", Some("fmi2")).expect("build script should run"); + + let marker = + std::fs::read_to_string(dir.path().join("build-script-marker")).expect("read marker"); + assert_eq!(marker, "ran"); + assert!( + !dir.path().join("binaries").exists(), + "target-provided build script should own FMU build outputs" + ); + } +} diff --git a/crates/rumoca/src/target_manifest.rs b/crates/rumoca/src/target_manifest.rs index 681149880..2e938eab9 100644 --- a/crates/rumoca/src/target_manifest.rs +++ b/crates/rumoca/src/target_manifest.rs @@ -284,9 +284,12 @@ fn render_manifest_files( .render(result, &file.path, model_identifier) .with_context(|| format!("Render target output path '{}'", file.path))?; let template = bundle.template_source(&file.template)?; + let file_renderer = file + .ir + .map(|ir| ManifestRenderer::Ir(template_ir_to_cli(ir))); let content = render_manifest_template( result, - renderer, + file_renderer.as_ref().unwrap_or(renderer), file.render_context, template.as_ref(), model_identifier, @@ -679,9 +682,12 @@ fn write_manifest_file( } let template = bundle.template_source(&file.template)?; + let file_renderer = file + .ir + .map(|ir| ManifestRenderer::Ir(template_ir_to_cli(ir))); let rendered = render_manifest_template( result, - renderer, + file_renderer.as_ref().unwrap_or(renderer), file.render_context, template.as_ref(), model_identifier, @@ -982,6 +988,120 @@ end ScalarCudaSmoke; Command::new(command).arg("--version").output().is_ok() } + #[test] + fn solve_target_can_render_dae_metadata_file_before_solve_files() { + let result = compile_scalar_cuda_smoke_demo(); + let target_dir = tempfile::tempdir().expect("temp target dir"); + std::fs::write( + target_dir.path().join("target.toml"), + r#" +version = 1 +ir = "solve" +name = "mixed-ir-target" + +[[files]] +path = "resources/ir-kind.txt" +template = "ir-kind.txt.jinja" +ir = "dae" + +[[files]] +path = "solve-kind.txt" +template = "solve-kind.txt.jinja" +"#, + ) + .expect("write target manifest"); + std::fs::write( + target_dir.path().join("ir-kind.txt.jinja"), + "metadata={{ ir_kind }} model={{ model_name }}", + ) + .expect("write DAE template"); + std::fs::write( + target_dir.path().join("solve-kind.txt.jinja"), + "runtime={{ ir_kind }} model={{ model_name }}", + ) + .expect("write solve template"); + + let bundle = TargetBundle::load( + target_dir + .path() + .to_str() + .expect("target dir path should be utf-8"), + ) + .expect("load temp target"); + let manifest = bundle.parse_manifest().expect("parse temp manifest"); + let out_dir = tempfile::tempdir().expect("temp output dir"); + + compile_manifest_target( + &result, + "ScalarCudaSmoke", + &bundle, + &manifest, + Some(out_dir.path().to_path_buf()), + ) + .expect("mixed IR target should render"); + + let metadata = std::fs::read_to_string(out_dir.path().join("resources/ir-kind.txt")) + .expect("read DAE metadata file"); + let runtime = std::fs::read_to_string(out_dir.path().join("solve-kind.txt")) + .expect("read solve file"); + assert_eq!(metadata, "metadata=dae model=ScalarCudaSmoke"); + assert_eq!(runtime, "runtime=solve model=ScalarCudaSmoke"); + } + + #[test] + fn solve_target_context_exposes_visible_value_rows() { + let result = compile_scalar_cuda_smoke_demo(); + let target_dir = tempfile::tempdir().expect("temp target dir"); + std::fs::write( + target_dir.path().join("target.toml"), + r#" +version = 1 +ir = "solve" +name = "visible-row-target" + +[[files]] +path = "visible.txt" +template = "visible.txt.jinja" +"#, + ) + .expect("write target manifest"); + std::fs::write( + target_dir.path().join("visible.txt.jinja"), + "names={{ solve.visible_names | length }} rows={{ solve.visible_value_rows | length }}", + ) + .expect("write solve template"); + + let bundle = TargetBundle::load( + target_dir + .path() + .to_str() + .expect("target dir path should be utf-8"), + ) + .expect("load temp target"); + let manifest = bundle.parse_manifest().expect("parse temp manifest"); + let out_dir = tempfile::tempdir().expect("temp output dir"); + + compile_manifest_target( + &result, + "ScalarCudaSmoke", + &bundle, + &manifest, + Some(out_dir.path().to_path_buf()), + ) + .expect("solve target should render visible-row context"); + + let rendered = + std::fs::read_to_string(out_dir.path().join("visible.txt")).expect("read visible file"); + assert!( + rendered.starts_with("names=") && rendered.contains(" rows="), + "visible row metadata should render from solve context: {rendered}" + ); + assert!( + !rendered.contains("names=0 rows=0"), + "visible row metadata should not be empty for scalar model: {rendered}" + ); + } + #[test] fn solve_target_capabilities_allow_scalar_tensor_fallback() { let result = compile_tensor_target_demo(); diff --git a/crates/rumoca/tests/architecture_hardening/env_var_registry.rs b/crates/rumoca/tests/architecture_hardening/env_var_registry.rs index f8f25bbab..465006990 100644 --- a/crates/rumoca/tests/architecture_hardening/env_var_registry.rs +++ b/crates/rumoca/tests/architecture_hardening/env_var_registry.rs @@ -28,10 +28,18 @@ use super::*; use std::collections::BTreeMap; use std::path::PathBuf; -/// Registered `RUMOCA_*` environment variables. Intentionally empty: the policy -/// is literal zero (see module docs). Adding an entry here is a deliberate, -/// reviewable policy exception — not the default escape hatch. -const REGISTERED_ENV_VARS: &[&str] = &[]; +/// Registered `RUMOCA_*` environment variables. Intentionally tiny: the policy +/// is literal zero by default (see module docs). Adding an entry here is a +/// deliberate, reviewable policy exception — not the default escape hatch. +/// +/// `RUMOCA_OMC_DOCKER_IMAGE` is a host/CI compatibility override for the +/// Docker-backed `omc` wrapper used by MSL parity verification. The parity +/// harness cannot pass argv into the shell wrapper, and the committed MSL +/// baseline is tied to a specific OpenModelica image version. +/// +/// `RUMOCA_CI_HEAD_SHA` carries the workflow-selected commit into checkout +/// actions and independent shell jobs, which cannot share a CLI argument. +const REGISTERED_ENV_VARS: &[&str] = &["RUMOCA_CI_HEAD_SHA", "RUMOCA_OMC_DOCKER_IMAGE"]; /// Extract every `RUMOCA_` token on one source line that is used as an /// environment variable, across both Rust and JS/TS source. A token qualifies diff --git a/crates/rumoca/tests/architecture_hardening_test.rs b/crates/rumoca/tests/architecture_hardening_test.rs index 76f246b34..b01da8cea 100644 --- a/crates/rumoca/tests/architecture_hardening_test.rs +++ b/crates/rumoca/tests/architecture_hardening_test.rs @@ -63,6 +63,35 @@ fn test_no_direct_dot_tokenization_for_model_paths() { assert_no_direct_dot_tokenization_for_model_paths(); } +#[test] +fn test_builtin_codegen_has_no_kelvin_or_boptest_adapters() { + let root = workspace_root(); + let codegen_root = root.join("crates/rumoca-phase-codegen/src"); + let mut rs_files = Vec::new(); + collect_rs_files(&codegen_root, &mut rs_files); + + let banned_terms = ["boptest", "top_down", "top-down", "kelvin"]; + let mut offenders = Vec::new(); + + for path in rs_files { + let Ok(content) = fs::read_to_string(&path) else { + continue; + }; + // `kelvin` is also the standard SI base-dimension JSON key. + let lower = content.to_ascii_lowercase().replace("\"kelvin\":", ""); + for term in banned_terms { + if lower.contains(term) { + offenders.push(format!("{} contains {term}", path.display())); + } + } + } + + assert!( + offenders.is_empty(), + "Rumoca built-in codegen must stay project-neutral; Kelvin/BOPTEST adapters belong in Kelvin: {offenders:?}" + ); +} + #[test] fn test_semantic_code_does_not_add_textual_model_path_recovery() { assert_semantic_code_does_not_add_textual_model_path_recovery(); @@ -1363,9 +1392,9 @@ fn test_sim_facade_cross_crate_exports_are_curated() { ); assert!( root_exports.iter().any(|export| { - export == "pub use rumoca_phase_solve::{lower_solve_artifacts, lower_solve_problem};" + export == "pub use rumoca_phase_solve::{ lower_dae_to_solve_model_owned, lower_solve_artifacts, lower_solve_problem, };" }), - "rumoca-sim may expose solve lowering/artifact preparation as its simulation-preparation facade" + "rumoca-sim may expose direct SolveModel lowering plus solve preparation helpers as its simulation-preparation facade" ); assert!( root_exports diff --git a/crates/rumoca/tests/array_component_binding_regression.rs b/crates/rumoca/tests/array_component_binding_regression.rs new file mode 100644 index 000000000..dee6f6f81 --- /dev/null +++ b/crates/rumoca/tests/array_component_binding_regression.rs @@ -0,0 +1,188 @@ +use rumoca::Compiler; +use rumoca_core::ExpressionVisitor; +use std::collections::HashSet; + +const TWO_DIMENSIONAL_COMPONENT_ARRAY: &str = r#" +block Source + parameter Real p; + output Real y; +equation + y = p; +end Source; + +partial model Icon + parameter Real display; +end Icon; + +partial model BaseCell + extends Icon(final display=y); + parameter Real p; + Source source(p=p); + output Real y = source.y; +end BaseCell; + +model ConcreteCell + extends BaseCell; +end ConcreteCell; + +partial model BaseStack + parameter Integer n = 2; + parameter Integer m = 2; + parameter Real q[n,m] = [1, 2; 3, 4]; + replaceable BaseCell cell[n,m](p=q, y(start=q)); +end BaseStack; + +model TwoDimensionalComponentArray + extends BaseStack(redeclare ConcreteCell cell(p=q, y(start=q))); +end TwoDimensionalComponentArray; +"#; + +#[test] +fn two_dimensional_component_array_retains_nested_declaration_bindings() { + let compiled = Compiler::new() + .model("TwoDimensionalComponentArray") + .compile_str( + TWO_DIMENSIONAL_COMPONENT_ARRAY, + "array_component_binding.mo", + ) + .expect("two-dimensional component-array declaration bindings must balance"); + + for name in ["cell[1,1].y", "cell[1,2].y", "cell[2,1].y", "cell[2,2].y"] { + let binding = compiled + .flat + .variables + .get(&rumoca_core::VarName::new(name)) + .and_then(|variable| variable.binding.as_ref()); + assert!(binding.is_some(), "missing declaration binding for {name}"); + } +} + +const TWO_DIMENSIONAL_COMPONENT_ARRAY_NESTED_EQUATIONS: &str = r#" +partial model BaseDynamicLeaf + Real x(start=0, fixed=true); + Real y; +equation + der(x) = 1; + y = x; +end BaseDynamicLeaf; + +model DynamicLeaf + extends BaseDynamicLeaf; +end DynamicLeaf; + +partial model BaseCellWithNestedEquations + replaceable BaseDynamicLeaf cell; +end BaseCellWithNestedEquations; + +model CellWithNestedEquations + extends BaseCellWithNestedEquations(redeclare DynamicLeaf cell); +end CellWithNestedEquations; + +partial model BaseTwoDimensionalComponentArrayNestedEquations + replaceable BaseCellWithNestedEquations cell[2,2]; +end BaseTwoDimensionalComponentArrayNestedEquations; + +model StackWithNestedEquations + extends BaseTwoDimensionalComponentArrayNestedEquations( + redeclare CellWithNestedEquations cell); +end StackWithNestedEquations; + +model TwoDimensionalComponentArrayNestedEquations + StackWithNestedEquations stack; +end TwoDimensionalComponentArrayNestedEquations; +"#; + +#[test] +fn two_dimensional_component_array_qualifies_nested_equations_per_element() { + let compiled = Compiler::new() + .model("TwoDimensionalComponentArrayNestedEquations") + .compile_str( + TWO_DIMENSIONAL_COMPONENT_ARRAY_NESTED_EQUATIONS, + "array_component_nested_equations.mo", + ) + .expect("nested equations in a component array must bind to each scalar instance"); + + for (x, y) in [ + ("stack.cell[1,1].cell.x", "stack.cell[1,1].cell.y"), + ("stack.cell[1,2].cell.x", "stack.cell[1,2].cell.y"), + ("stack.cell[2,1].cell.x", "stack.cell[2,1].cell.y"), + ("stack.cell[2,2].cell.x", "stack.cell[2,2].cell.y"), + ] { + assert!( + compiled + .dae + .variables + .states + .contains_key(&rumoca_core::VarName::new(x)), + "nested derivative must classify {x} as a state" + ); + assert!( + compiled + .dae + .variables + .algebraics + .contains_key(&rumoca_core::VarName::new(y)), + "nested algebraic equation must retain {y}" + ); + + let equation_census = compiled + .dae + .continuous + .equations + .iter() + .map(|equation| EquationCensus::from_expression(&equation.rhs)) + .collect::>(); + assert_eq!( + equation_census + .iter() + .filter(|equation| equation.has_der && equation.refs.contains(x)) + .count(), + 1, + "expected one indexed der({x}) equation" + ); + assert_eq!( + equation_census + .iter() + .filter(|equation| { + !equation.has_der && equation.refs.contains(x) && equation.refs.contains(y) + }) + .count(), + 1, + "expected one indexed {y} = {x} equation" + ); + } +} + +#[derive(Default)] +struct EquationCensus { + refs: HashSet, + has_der: bool, +} + +impl EquationCensus { + fn from_expression(expression: &rumoca_core::Expression) -> Self { + let mut census = Self::default(); + census.visit_expression(expression); + census + } +} + +impl ExpressionVisitor for EquationCensus { + fn visit_var_ref( + &mut self, + name: &rumoca_core::Reference, + subscripts: &[rumoca_core::Subscript], + ) { + self.refs.insert(name.as_str().to_string()); + self.walk_var_ref(name, subscripts); + } + + fn visit_builtin_call( + &mut self, + function: &rumoca_core::BuiltinFunction, + args: &[rumoca_core::Expression], + ) { + self.has_der |= *function == rumoca_core::BuiltinFunction::Der; + self.walk_builtin_call(function, args); + } +} diff --git a/crates/rumoca/tests/array_dependent_parameter_propagation.rs b/crates/rumoca/tests/array_dependent_parameter_propagation.rs new file mode 100644 index 000000000..ea6cc1f83 --- /dev/null +++ b/crates/rumoca/tests/array_dependent_parameter_propagation.rs @@ -0,0 +1,254 @@ +use rumoca::Compiler; +use rumoca_core::VarName; +use rumoca_ir_solve::{LinearOp, ScalarSlot}; +use rumoca_sim::{ + SimOptions, SimSolverMode, lower_for_simulation_with_overrides, simulate_dae_with_diagnostics, +}; + +const SOURCE: &str = r#" +model Arr + parameter Real a = 1; + parameter Real arr[3] = {a, 2*a, 3*a}; + Real x[3](each start = 0); +equation + for i in 1:3 loop + der(x[i]) = arr[i]; + end for; +end Arr; +"#; + +const SINGLETON_DIM_SOURCE: &str = r#" +model SingletonDimArr + parameter Real a = 1; + parameter Real arr[1,3] = [a, 2*a, 3*a]; + Real x[1,3](each start = 0); +equation + for i in 1:1 loop + for j in 1:3 loop + der(x[i,j]) = arr[i,j]; + end for; + end for; +end SingletonDimArr; +"#; + +fn compile_arr() -> rumoca::CompilationResult { + Compiler::new() + .model("Arr") + .compile_str(SOURCE, "array_dependent_parameter.mo") + .expect("Arr should compile") +} + +fn compile_singleton_dim_arr() -> rumoca::CompilationResult { + Compiler::new() + .model("SingletonDimArr") + .compile_str( + SINGLETON_DIM_SOURCE, + "array_dependent_parameter_singleton_dim.mo", + ) + .expect("SingletonDimArr should compile") +} + +fn p_index(model: &rumoca_ir_solve::SolveModel, name: &str) -> usize { + match model.problem.layout.binding(name) { + Some(ScalarSlot::P { index, .. }) => index, + other => panic!("{name} must have a P slot, got {other:?}"), + } +} + +fn parameter_values(model: &rumoca_ir_solve::SolveModel) -> [f64; 4] { + ["a", "arr[1]", "arr[2]", "arr[3]"].map(|name| model.parameters[p_index(model, name)]) +} + +#[test] +fn array_dependent_parameters_preserve_dae_slots_and_derivative_lanes() { + let compiled = compile_arr(); + let a = compiled + .dae + .variables + .parameters + .get(&VarName::new("a")) + .expect("DAE parameter a"); + let arr = compiled + .dae + .variables + .parameters + .get(&VarName::new("arr")) + .expect("DAE parameter arr"); + assert!(a.start.is_some(), "DAE parameter a must retain its binding"); + assert_eq!(arr.dims, vec![3], "DAE arr must retain its declared shape"); + let arr_binding = arr.start.as_ref().expect("DAE arr binding"); + let mut binding_refs = Vec::new(); + arr_binding.collect_var_refs(&mut binding_refs); + assert_eq!(binding_refs, vec![VarName::new("a")]); + + let artifact_opts = SimOptions { + solver_mode: SimSolverMode::Bdf, + ..SimOptions::default() + }; + let prepared = rumoca_sim::structurally_prepared_dae_for_simulation_artifact( + &compiled.dae, + &artifact_opts, + ) + .expect("prepared DAE"); + let boundary = + rumoca_sim::boundary_reduced_dae_for_simulation_artifact(&compiled.dae, &artifact_opts) + .expect("boundary-reduced DAE"); + assert_eq!(prepared.continuous.structured_equations.len(), 1); + assert_eq!(boundary.continuous.structured_equations.len(), 1); + assert!( + !boundary.continuous.structured_equations[0].interiors_materialized, + "boundary elimination must retain the compact family that owns interior lanes" + ); + + let base = lower_for_simulation_with_overrides(&compiled.dae, &SimOptions::default()) + .expect("base solve model"); + assert_eq!(parameter_values(&base), [1.0, 1.0, 2.0, 3.0]); + + let override_opts = SimOptions { + param_overrides: vec![("a".to_string(), 10.0)], + ..SimOptions::default() + }; + let overridden = lower_for_simulation_with_overrides(&compiled.dae, &override_opts) + .expect("override solve model"); + assert_eq!(parameter_values(&overridden), [10.0, 10.0, 20.0, 30.0]); + + let scalar_rows = + rumoca_eval_solve::to_scalar_program_block(&base.problem.continuous.derivative_rhs) + .expect("derivative tensor nodes should have a scalar view"); + assert_eq!(scalar_rows.programs.len(), 3); + for (lane, program) in scalar_rows.programs.iter().enumerate() { + let expected = p_index(&base, &format!("arr[{}]", lane + 1)); + let loaded = program + .iter() + .filter_map(|op| match op { + LinearOp::LoadP { index, .. } => Some(*index), + _ => None, + }) + .collect::>(); + assert_eq!( + loaded, + vec![expected], + "derivative lane {} must read arr[{}], nodes={:?}, rows={:?}, outputs={:?}", + lane + 1, + lane + 1, + base.problem.continuous.derivative_rhs.nodes, + scalar_rows.programs, + scalar_rows.output_indices, + ); + } + for name in ["x[1]", "x[2]", "x[3]"] { + assert!( + base.visible_names.iter().any(|candidate| candidate == name), + "runtime-visible names must contain {name}: {:?}", + base.visible_names + ); + } +} + +fn final_values(dae: &rumoca_ir_dae::Dae, opts: SimOptions) -> [f64; 3] { + final_values_for_names(dae, opts, ["x[1]", "x[2]", "x[3]"]) +} + +fn final_values_for_names( + dae: &rumoca_ir_dae::Dae, + opts: SimOptions, + names: [&str; 3], +) -> [f64; 3] { + let sim = simulate_dae_with_diagnostics(dae, &opts).expect("Arr should simulate"); + names.map(|name| { + let index = sim + .names + .iter() + .position(|candidate| candidate == name) + .unwrap_or_else(|| panic!("simulation names must contain {name}: {:?}", sim.names)); + sim.data[index] + .last() + .copied() + .expect("simulation output row") + }) +} + +#[test] +fn array_dependent_parameter_trajectories_follow_base_and_override() { + let compiled = compile_arr(); + let base = final_values( + &compiled.dae, + SimOptions { + t_end: 0.5, + dt: Some(0.01), + solver_mode: SimSolverMode::Bdf, + ..SimOptions::default() + }, + ); + for (actual, expected) in base.into_iter().zip([0.5, 1.0, 1.5]) { + assert!((actual - expected).abs() < 1.0e-8, "{actual} != {expected}"); + } + + let overridden = final_values( + &compiled.dae, + SimOptions { + t_end: 0.5, + dt: Some(0.01), + solver_mode: SimSolverMode::Bdf, + param_overrides: vec![("a".to_string(), 10.0)], + ..SimOptions::default() + }, + ); + for (actual, expected) in overridden.into_iter().zip([5.0, 10.0, 15.0]) { + assert!((actual - expected).abs() < 1.0e-8, "{actual} != {expected}"); + } +} + +#[test] +fn singleton_dimension_array_parameters_reach_solve_and_bdf() { + let compiled = compile_singleton_dim_arr(); + let base = lower_for_simulation_with_overrides(&compiled.dae, &SimOptions::default()) + .expect("base singleton-dimension solve model"); + assert_eq!( + ["a", "arr[1,1]", "arr[1,2]", "arr[1,3]"].map(|name| base.parameters[p_index(&base, name)]), + [1.0, 1.0, 2.0, 3.0] + ); + let scalar_rows = + rumoca_eval_solve::to_scalar_program_block(&base.problem.continuous.derivative_rhs) + .expect("singleton-dimension derivative tensor should scalarize"); + for (lane, program) in scalar_rows.programs.iter().enumerate() { + let expected = p_index(&base, &format!("arr[1,{}]", lane + 1)); + let loaded = program + .iter() + .filter_map(|op| match op { + LinearOp::LoadP { index, .. } => Some(*index), + _ => None, + }) + .collect::>(); + assert_eq!(loaded, vec![expected], "singleton-dimension lane {lane}"); + } + + let base_final = final_values_for_names( + &compiled.dae, + SimOptions { + t_end: 0.5, + dt: Some(0.01), + solver_mode: SimSolverMode::Bdf, + ..SimOptions::default() + }, + ["x[1,1]", "x[1,2]", "x[1,3]"], + ); + for (actual, expected) in base_final.into_iter().zip([0.5, 1.0, 1.5]) { + assert!((actual - expected).abs() < 1.0e-8, "{actual} != {expected}"); + } + + let overridden_final = final_values_for_names( + &compiled.dae, + SimOptions { + t_end: 0.5, + dt: Some(0.01), + solver_mode: SimSolverMode::Bdf, + param_overrides: vec![("a".to_string(), 10.0)], + ..SimOptions::default() + }, + ["x[1,1]", "x[1,2]", "x[1,3]"], + ); + for (actual, expected) in overridden_final.into_iter().zip([5.0, 10.0, 15.0]) { + assert!((actual - expected).abs() < 1.0e-8, "{actual} != {expected}"); + } +} diff --git a/crates/rumoca/tests/backend_template_runtime_regression.rs b/crates/rumoca/tests/backend_template_runtime_regression.rs index e4590b0e3..110b4a18a 100644 --- a/crates/rumoca/tests/backend_template_runtime_regression.rs +++ b/crates/rumoca/tests/backend_template_runtime_regression.rs @@ -18,6 +18,12 @@ use rumoca_phase_codegen::templates; use rumoca_sim::{SimOptions, SimResult, SimSolverMode, simulate_dae_with_diagnostics}; use tempfile::Builder; +#[cfg(feature = "template-runtime-tests")] +#[path = "backend_template_runtime_regression/python_runtime.rs"] +mod python_runtime; +#[cfg(feature = "template-runtime-tests")] +use python_runtime::run_python; + // ============================================================================ // Tolerance — max bounded relative error: |a-b| / max(|a|, |b|, 1.0) // ============================================================================ @@ -30,20 +36,6 @@ const CASADI_TOLERANCE: f64 = 0.01; #[cfg(feature = "template-runtime-tests")] const C_TOLERANCE: f64 = 0.05; -// ============================================================================ -// Runtime detection -// ============================================================================ - -#[cfg(feature = "template-runtime-tests")] -fn python_command() -> &'static str { - for candidate in ["python3", "python"] { - if Command::new(candidate).arg("--version").output().is_ok() { - return candidate; - } - } - panic!("expected python3 or python to be available"); -} - #[cfg(feature = "template-runtime-tests")] fn strict_runtime_dependencies() -> bool { std::path::PathBuf::from(env!("CARGO_MANIFEST_DIR")) @@ -292,36 +284,6 @@ fn assert_traces_match( } } -// ============================================================================ -// Python execution helper (returns stdout) -// ============================================================================ - -#[cfg(feature = "template-runtime-tests")] -fn run_python(rendered: &str, driver: &str) -> String { - let dir = Builder::new() - .prefix("rumoca_runtime_test_") - .tempdir() - .expect("create temp dir"); - let model_path = dir.path().join("model.py"); - let driver_path = dir.path().join("driver.py"); - fs::write(&model_path, rendered).expect("write model.py"); - fs::write(&driver_path, driver).expect("write driver.py"); - - let output = Command::new(python_command()) - .arg(driver_path.to_str().unwrap()) - .output() - .expect("run Python driver"); - - assert!( - output.status.success(), - "Python execution failed\nstdout:\n{}\nstderr:\n{}", - String::from_utf8_lossy(&output.stdout), - String::from_utf8_lossy(&output.stderr) - ); - - String::from_utf8(output.stdout).expect("stdout is utf8") -} - #[cfg(feature = "template-runtime-tests")] fn run_julia(rendered: &str, driver: &str) -> String { let dir = Builder::new() @@ -467,6 +429,16 @@ equation end Oscillator; "#; +const ARRAY_ACCESS_SOURCE: &str = r#" +model ArrayAccess + parameter Real nominal_air_flow[6] = {1, 2, 3, 4, 5, 6}; + parameter Real floor_internal_gain[2, 2] = [1, 2; 3, 4]; + Real x(start = 0); +equation + der(x) = -x + sum(nominal_air_flow) + nominal_air_flow[5 + 1] + floor_internal_gain[1, 1]; +end ArrayAccess; +"#; + const MATRIX_DER_PRODUCT_SOURCE: &str = r#" model MatrixDerProduct Real R[3,3](start={{1,0,0},{0,1,0},{0,0,1}}, fixed=true); @@ -494,6 +466,25 @@ fn native_simulates_matrix_derivative_product() { ); } +const INDEXED_COMPONENT_FIELD_SOURCE: &str = r#" +package IndexedComponentFieldProbe + model Zone + Real T(start = 290); + equation + der(T) = -0.01 * (T - 290); + end Zone; + + model Main + Zone zone[3]; + Real y[2]; + equation + for i in 1:2 loop + y[i] = zone[i + 1].T; + end for; + end Main; +end IndexedComponentFieldProbe; +"#; + // ============================================================================ // CasADi driver — outputs CSV: time,state1,state2,... // ============================================================================ @@ -537,7 +528,7 @@ for i, t in enumerate(tgrid): #[cfg(feature = "template-runtime-tests")] fn casadi_trace_test(source: &str, model_name: &str, template: &str) { let rendered = render_template(source, model_name, template); - let csv = run_python(&rendered, CASADI_CSV_DRIVER); + let csv = run_python(&rendered, CASADI_CSV_DRIVER, "CasADi"); let backend_traces = parse_csv_traces(&csv); let (dae, sim) = reference_trace(source, model_name, 1.0); assert_traces_match(&backend_traces, &dae.dae, &sim, CASADI_TOLERANCE, "CasADi"); @@ -791,6 +782,30 @@ fn fmi2_event_reinit_runtime() { assert_event_reinit_trace(&csv, "FMI2"); } +#[test] +fn fmi2_array_access_component_compiles() { + let compiled = compile_model(ARRAY_ACCESS_SOURCE, "ArrayAccess"); + let model_c = render_fmi_solve_template(&compiled, "fmi2", "model.c.jinja", "ArrayAccess"); + let driver_c = + render_fmi_solve_template(&compiled, "fmi2", "test_driver.c.jinja", "ArrayAccess"); + compile_and_run_c( + &[("model.c", &model_c), ("driver.c", &driver_c)], + &["--t-end", "0.01", "--dt", "0.001"], + ); +} + +#[test] +fn fmi2_indexed_component_field_compiles() { + let model = "IndexedComponentFieldProbe.Main"; + let compiled = compile_model(INDEXED_COMPONENT_FIELD_SOURCE, model); + let model_c = render_fmi_solve_template(&compiled, "fmi2", "model.c.jinja", model); + let driver_c = render_fmi_solve_template(&compiled, "fmi2", "test_driver.c.jinja", model); + compile_and_run_c( + &[("model.c", &model_c), ("driver.c", &driver_c)], + &["--t-end", "0.01", "--dt", "0.001"], + ); +} + // ============================================================================ // FMI 3.0 runtime tests // ============================================================================ @@ -850,6 +865,18 @@ fn fmi3_event_reinit_runtime() { assert_event_reinit_trace(&csv, "FMI3"); } +#[test] +fn fmi3_array_access_component_compiles() { + let compiled = compile_model(ARRAY_ACCESS_SOURCE, "ArrayAccess"); + let model_c = render_fmi_solve_template(&compiled, "fmi3", "model.c.jinja", "ArrayAccess"); + let driver_c = + render_fmi_solve_template(&compiled, "fmi3", "test_driver.c.jinja", "ArrayAccess"); + compile_and_run_c( + &[("model.c", &model_c), ("driver.c", &driver_c)], + &["--t-end", "0.01", "--dt", "0.001"], + ); +} + #[test] #[cfg(feature = "template-runtime-tests")] fn fmi3_matrix_derivative_product_runtime() { @@ -1586,7 +1613,7 @@ fn sympy_trace_test(source: &str, model_name: &str) { ) .expect("render template"); - let stdout = run_python(&rendered, SYMPY_EVAL_DRIVER); + let stdout = run_python(&rendered, SYMPY_EVAL_DRIVER, "SymPy"); let result: serde_json::Value = serde_json::from_str(stdout.trim()).expect("parse JSON output"); // Get reference derivatives at t=0 from rumoca simulator @@ -1650,26 +1677,14 @@ spec.loader.exec_module(mod) print(mod.simulate()) "#; -#[cfg(feature = "template-runtime-tests")] -fn python_has_onnx() -> bool { - Command::new(python_command()) - .args(["-c", "import onnx; import onnxruntime; import numpy"]) - .output() - .map(|o| o.status.success()) - .unwrap_or(false) -} - #[cfg(feature = "template-runtime-tests")] fn onnx_trace_test(source: &str, model_name: &str) { - if !runtime_dependency_available(python_has_onnx(), "onnx/onnxruntime") { - return; - } let rendered = render_template( source, model_name, templates::builtin_template_source("onnx", "onnx.py.jinja").unwrap(), ); - let csv = run_python(&rendered, ONNX_CSV_DRIVER); + let csv = run_python(&rendered, ONNX_CSV_DRIVER, "ONNX"); let backend_traces = parse_csv_traces(&csv); let (dae, sim) = reference_trace(source, model_name, 1.0); assert_traces_match(&backend_traces, &dae.dae, &sim, C_TOLERANCE, "ONNX"); @@ -1714,26 +1729,14 @@ spec.loader.exec_module(mod) print(mod.simulate_csv()) "#; -#[cfg(feature = "template-runtime-tests")] -fn python_has_jax() -> bool { - Command::new(python_command()) - .args(["-c", "import jax; import diffrax; import numpy"]) - .output() - .map(|o| o.status.success()) - .unwrap_or(false) -} - #[cfg(feature = "template-runtime-tests")] fn jax_trace_test(source: &str, model_name: &str) { - if !runtime_dependency_available(python_has_jax(), "jax/diffrax") { - return; - } let rendered = render_template( source, model_name, templates::builtin_template_source("jax", "jax.py.jinja").unwrap(), ); - let csv = run_python(&rendered, JAX_CSV_DRIVER); + let csv = run_python(&rendered, JAX_CSV_DRIVER, "JAX"); let backend_traces = parse_csv_traces(&csv); let (dae, sim) = reference_trace(source, model_name, 1.0); assert_traces_match(&backend_traces, &dae.dae, &sim, JAX_TOLERANCE, "JAX"); diff --git a/crates/rumoca/tests/backend_template_runtime_regression/python_runtime.rs b/crates/rumoca/tests/backend_template_runtime_regression/python_runtime.rs new file mode 100644 index 000000000..a1cb8c552 --- /dev/null +++ b/crates/rumoca/tests/backend_template_runtime_regression/python_runtime.rs @@ -0,0 +1,178 @@ +use std::{fs, process::Command}; + +use tempfile::Builder; + +fn python_modules(backend: &str) -> &'static [&'static str] { + match backend { + "CasADi" => &["casadi", "numpy"], + "SymPy" => &["sympy"], + "ONNX" => &["onnx", "onnxruntime", "numpy"], + "JAX" => &["jax", "diffrax", "numpy"], + _ => panic!("unknown Python template backend: {backend}"), + } +} + +fn probe_python_modules(interpreter: &str, backend: &str, modules: &[&str]) -> Result<(), String> { + let output = Command::new(interpreter) + .args(["-c", &format!("import {}", modules.join(", "))]) + .output() + .map_err(|error| { + format!("{backend} dependency probe could not run {interpreter}: {error}") + })?; + if output.status.success() { + return Ok(()); + } + Err(format!( + "{backend} Python runtime dependency probe failed\ninterpreter: {interpreter}\nrequired modules: {}\nstdout:\n{}\nstderr:\n{}", + modules.join(", "), + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr), + )) +} + +fn resolve_python<'a>( + candidates: impl IntoIterator, + backend: &str, + modules: &[&str], +) -> Result<&'a str, String> { + let mut failures = Vec::new(); + for candidate in candidates { + match Command::new(candidate).arg("--version").output() { + Ok(output) if output.status.success() => { + match probe_python_modules(candidate, backend, modules) { + Ok(()) => return Ok(candidate), + Err(error) => failures.push(error), + } + } + Ok(output) => failures.push(format!( + "interpreter {candidate} --version failed\nstdout:\n{}\nstderr:\n{}", + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr), + )), + Err(error) => failures.push(format!("interpreter {candidate} unavailable: {error}")), + } + } + Err(format!( + "{backend} Python runtime resolution failed; required modules: {}\n{}", + modules.join(", "), + failures.join("\n---\n") + )) +} + +fn python_command(backend: &str) -> &'static str { + resolve_python(["python3", "python"], backend, python_modules(backend)) + .unwrap_or_else(|error| panic!("{error}")) +} + +pub(super) fn run_python(rendered: &str, driver: &str, backend: &str) -> String { + let python = python_command(backend); + let dir = Builder::new() + .prefix("rumoca_runtime_test_") + .tempdir() + .expect("create temp dir"); + let model_path = dir.path().join("model.py"); + let driver_path = dir.path().join("driver.py"); + fs::write(&model_path, rendered).expect("write model.py"); + fs::write(&driver_path, driver).expect("write driver.py"); + + let output = Command::new(python) + .arg(driver_path.to_str().unwrap()) + .output() + .expect("run Python driver"); + + assert!( + output.status.success(), + "Python execution failed\nstdout:\n{}\nstderr:\n{}", + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr) + ); + + String::from_utf8(output.stdout).expect("stdout is utf8") +} + +#[cfg(unix)] +fn probe_executable(dir: &std::path::Path, name: &str, body: &str) -> std::path::PathBuf { + use std::os::unix::fs::PermissionsExt; + + let path = dir.join(name); + fs::write(&path, body).expect("write fake Python executable"); + let mut permissions = fs::metadata(&path) + .expect("fake Python metadata") + .permissions(); + permissions.set_mode(0o755); + fs::set_permissions(&path, permissions).expect("make fake Python executable"); + path +} + +#[cfg(unix)] +const CANDIDATE_PROBE: &str = r#"#!/bin/sh +dir=${0%/*} +name=${0##*/} +printf '%s\n' "$@" > "$dir/$name.args" +if [ "$1" = --version ]; then exit 0; fi +if [ "$name" = python-capable ]; then exit 0; fi +echo "ModuleNotFoundError: $name lacks ONNX" >&2 +exit 7 +"#; + +#[test] +fn python_runtime_probe_contract_modules_are_backend_specific() { + assert_eq!(python_modules("CasADi"), &["casadi", "numpy"]); + assert_eq!(python_modules("SymPy"), &["sympy"]); + assert_eq!(python_modules("ONNX"), &["onnx", "onnxruntime", "numpy"]); + assert_eq!(python_modules("JAX"), &["jax", "diffrax", "numpy"]); +} + +#[test] +#[cfg(unix)] +fn python_runtime_probe_contract_selects_backend_capable_candidate() { + let dir = Builder::new() + .prefix("rumoca_python_candidates_") + .tempdir() + .expect("create Python candidate dir"); + let first = probe_executable(dir.path(), "python-first", CANDIDATE_PROBE); + let second = probe_executable(dir.path(), "python-capable", CANDIDATE_PROBE); + let candidates = [ + first.to_str().expect("UTF-8 first candidate"), + second.to_str().expect("UTF-8 second candidate"), + ]; + + assert_eq!( + resolve_python(candidates, "ONNX", python_modules("ONNX")).expect("resolve Python"), + candidates[1] + ); + for candidate in ["python-first", "python-capable"] { + assert_eq!( + fs::read_to_string(dir.path().join(format!("{candidate}.args"))) + .expect("read candidate argv"), + "-c\nimport onnx, onnxruntime, numpy\n" + ); + } +} + +#[test] +#[cfg(unix)] +fn python_runtime_probe_contract_preserves_backend_import_errors() { + let dir = Builder::new() + .prefix("rumoca_python_probe_") + .tempdir() + .expect("create Python probe dir"); + let first = probe_executable(dir.path(), "python-first", CANDIDATE_PROBE); + let second = probe_executable(dir.path(), "python-second", CANDIDATE_PROBE); + let candidates = [ + first.to_str().expect("UTF-8 first candidate"), + second.to_str().expect("UTF-8 second candidate"), + ]; + let error = resolve_python(candidates, "ONNX", python_modules("ONNX")) + .expect_err("missing ONNX dependencies must fail closed"); + for expected in [ + "ONNX", + "onnx, onnxruntime, numpy", + candidates[0], + candidates[1], + "ModuleNotFoundError: python-first lacks ONNX", + "ModuleNotFoundError: python-second lacks ONNX", + ] { + assert!(error.contains(expected), "missing {expected:?}: {error}"); + } +} diff --git a/crates/rumoca/tests/casadi_vs_rumoca_test.rs b/crates/rumoca/tests/casadi_vs_rumoca_test.rs index 44a6fb50d..30b9d43f2 100644 --- a/crates/rumoca/tests/casadi_vs_rumoca_test.rs +++ b/crates/rumoca/tests/casadi_vs_rumoca_test.rs @@ -168,6 +168,7 @@ fn casadi_simulate( fn sim_result_to_trace(sim: &SimResult, model_name: &str) -> SimTrace { SimTrace { model_name: Some(model_name.to_string()), + n_states: Some(sim.n_states), times: sim.times.clone(), names: sim.names.clone(), data: sim @@ -194,6 +195,7 @@ fn casadi_trace_to_sim_trace(raw: &CasadiRawTrace, model_name: &str) -> SimTrace SimTrace { model_name: Some(model_name.to_string()), + n_states: None, times: raw.times.clone(), names: raw.names.clone(), data, diff --git a/crates/rumoca/tests/clocked_sample_regression.rs b/crates/rumoca/tests/clocked_sample_regression.rs index e46967a8f..60e41f88d 100644 --- a/crates/rumoca/tests/clocked_sample_regression.rs +++ b/crates/rumoca/tests/clocked_sample_regression.rs @@ -20,8 +20,11 @@ model SampleTime block PeriodicClock parameter Real period = 0.1; ClockOutput y; + protected + Clock c; equation - y = Clock(period); + c = Clock(period); + y = c; end PeriodicClock; block AssignClock @@ -163,6 +166,35 @@ fn native_simulation_updates_condition_memory_after_clocked_sample_time() { ); } +#[test] +fn native_simulation_rearms_clock_alias_edge_at_every_periodic_tick() { + let compiled = rumoca::Compiler::new() + .model("SampleTime") + .compile_str(SAMPLE_TIME_SOURCE, "sample_time.mo") + .expect("clocked alias model should compile"); + let sim = simulate_dae( + &compiled.dae, + &SimOptions { + t_end: 0.3, + dt: Some(0.1), + ..SimOptions::default() + }, + ) + .expect("clocked alias model should simulate"); + + let y = trace_values(&sim, "assignClock.y"); + assert!( + (y[1] - 0.1).abs() <= 1.0e-12, + "clock alias should activate the assignment at the first periodic tick; got {}", + y[1] + ); + assert!( + (y[2] - 0.2).abs() <= 1.0e-12, + "clock alias should rearm the assignment at the second periodic tick; got {}", + y[2] + ); +} + #[test] fn rk_like_simulation_updates_condition_memory_after_clocked_sample_time() { let compiled = rumoca::Compiler::new() diff --git a/crates/rumoca/tests/discrete_boolean_array_regression.rs b/crates/rumoca/tests/discrete_boolean_array_regression.rs new file mode 100644 index 000000000..e05a1e305 --- /dev/null +++ b/crates/rumoca/tests/discrete_boolean_array_regression.rs @@ -0,0 +1,55 @@ +use rumoca_compile::{Session, SessionConfig}; +use rumoca_sim::{SimOptions, SimSolverMode, simulate_dae_with_diagnostics}; + +const BOOLEAN_ARRAY_FANOUT: &str = r#" +model BooleanArrayFanout + Boolean trigger; + Boolean fanout[2]; + discrete Integer hits[2](each start = 0, each fixed = true); + Real x(start = 0, fixed = true); +equation + der(x) = 0; + trigger = time >= 0.5; + fanout = fill(trigger, 2); + when edge(fanout[1]) then + hits[1] = pre(hits[1]) + 1; + end when; + when edge(fanout[2]) then + hits[2] = pre(hits[2]) + 1; + end when; +end BooleanArrayFanout; +"#; + +#[test] +fn scalarized_boolean_array_fanout_reaches_both_event_consumers() { + let mut session = Session::new(SessionConfig::default()); + session + .add_document("boolean_array_fanout.mo", BOOLEAN_ARRAY_FANOUT) + .expect("Boolean array fanout fixture should parse"); + let compiled = session + .compile_model("BooleanArrayFanout") + .expect("Boolean array fanout fixture should compile"); + + let sim = simulate_dae_with_diagnostics( + &compiled.dae, + &SimOptions { + t_end: 1.0, + solver_mode: SimSolverMode::RkLike, + ..SimOptions::default() + }, + ) + .expect("Boolean array fanout fixture should simulate"); + + for name in ["hits[1]", "hits[2]"] { + let index = sim + .names + .iter() + .position(|candidate| candidate == name) + .unwrap_or_else(|| panic!("missing trace for {name}: {:?}", sim.names)); + assert_eq!( + sim.data[index].last().copied(), + Some(1.0), + "both scalarized Boolean lanes must settle in the same event" + ); + } +} diff --git a/crates/rumoca/tests/discrete_enum_array_start_test.rs b/crates/rumoca/tests/discrete_enum_array_start_test.rs new file mode 100644 index 000000000..495f476a1 --- /dev/null +++ b/crates/rumoca/tests/discrete_enum_array_start_test.rs @@ -0,0 +1,118 @@ +//! Regression for array-valued enumeration starts crossing the DAE -> Solve boundary. + +use rumoca::Compiler; +use rumoca_core::{BuiltinFunction, Expression}; +use rumoca_phase_solve::lower_dae_to_solve_model; + +const ENUM_ARRAY_START_MODEL: &str = r#" +within; +type Logic = enumeration(U, X); + +model Register + parameter Integer n = 1; + Logic nextstate[n](start = fill(Logic.U, n)); +equation + when time >= 1 then + nextstate = pre(nextstate); + end when; +end Register; + +model EnumArrayStart + Register register(n = 2); +end EnumArrayStart; +"#; + +#[test] +fn dae_to_solve_preserves_filled_enum_array_start_for_generated_pre_variable() { + let compiled = Compiler::new() + .model("EnumArrayStart") + .compile_str(ENUM_ARRAY_START_MODEL, "EnumArrayStart.mo") + .expect("enum array model should compile to DAE"); + + let n = compiled + .flat + .variables + .get(&rumoca_core::VarName::new("register.n")) + .expect("flattening should retain the modified instance parameter"); + assert!( + matches!( + n.binding, + Some(Expression::Literal { + value: rumoca_core::Literal::Integer(2), + .. + }) + ), + "the flattened instance parameter must keep its modification: {:?}", + n.binding + ); + + let nextstate = compiled + .dae + .variables + .discrete_valued + .get(&rumoca_core::VarName::new("register.nextstate")) + .expect("DAE should retain the declared enum array"); + assert_eq!(nextstate.dims, [2]); + assert!( + matches!( + nextstate.start, + Some(Expression::BuiltinCall { + function: BuiltinFunction::Fill, + ref args, + .. + }) if matches!( + args.as_slice(), + [_, Expression::Literal { + value: rumoca_core::Literal::Integer(2), + .. + }] + ) + ), + "the instance parameter modification must reach the start expression: {:?}", + nextstate.start + ); + + let pre_nextstate = compiled + .dae + .variables + .parameters + .get(&rumoca_core::VarName::new("__pre__.register.nextstate")) + .expect("pre(nextstate) should generate an explicit DAE parameter"); + assert_eq!(pre_nextstate.dims, [2]); + assert!( + matches!( + pre_nextstate.start, + Some(Expression::BuiltinCall { + function: BuiltinFunction::Fill, + ref args, + .. + }) if matches!( + args.as_slice(), + [_, Expression::Literal { + value: rumoca_core::Literal::Integer(2), + .. + }] + ) + ), + "the generated pre parameter must preserve the shaped start" + ); + + let prepared = lower_dae_to_solve_model(&compiled.dae) + .expect("filled enum array starts should preserve both values in Solve lowering"); + for name in [ + "register.nextstate[1]", + "register.nextstate[2]", + "__pre__.register.nextstate[1]", + "__pre__.register.nextstate[2]", + ] { + let rumoca_ir_solve::ScalarSlot::P { index, .. } = prepared + .problem + .layout + .binding(name) + .unwrap_or_else(|| panic!("{name} should have a Solve parameter slot")) + else { + panic!("{name} should lower to a Solve parameter slot"); + }; + assert_eq!(prepared.parameters[index], 1.0, "wrong start for {name}"); + } +} diff --git a/crates/rumoca/tests/examples_smoke.rs b/crates/rumoca/tests/examples_smoke.rs index a1d4b766f..3e8811376 100644 --- a/crates/rumoca/tests/examples_smoke.rs +++ b/crates/rumoca/tests/examples_smoke.rs @@ -588,10 +588,10 @@ fn quadrotor_acro_roll_command_generates_body_rate_when_cmm_available() { return; }; - for (axis_input, gyro_output) in [ - ("stick_roll", "gyro[1]"), - ("stick_pitch", "gyro[2]"), - ("stick_yaw", "gyro[3]"), + for (axis_input, gyro_output, minimum_rate) in [ + ("stick_roll", "gyro[1]", 0.05), + ("stick_pitch", "gyro[2]", 0.05), + ("stick_yaw", "gyro[3]", 0.015), ] { let mut session = SimulationSession::new_with_diagnostics( &result.dae, @@ -632,9 +632,9 @@ fn quadrotor_acro_roll_command_generates_body_rate_when_cmm_available() { .copied() .unwrap_or_else(|| panic!("quadrotor state should contain {gyro_output}")); assert!( - rate.abs() > 0.05, + rate > minimum_rate, "{axis_input} should generate visible body rate at hover throttle; \ - {gyro_output}={rate}" + {gyro_output}={rate}, expected > {minimum_rate}" ); } } diff --git a/crates/rumoca/tests/flowmodel_modifier_scope_test.rs b/crates/rumoca/tests/flowmodel_modifier_scope_test.rs new file mode 100644 index 000000000..539860969 --- /dev/null +++ b/crates/rumoca/tests/flowmodel_modifier_scope_test.rs @@ -0,0 +1,54 @@ +//! Regression tests for FlowModel modifier scope resolution. + +#[test] +fn test_flowmodel_modifier_keeps_enclosing_port_scope() { + let source = r#" + package Medium + function setState_p + input Real p; + output Real s; + algorithm + s := p; + end setState_p; + end Medium; + + connector FluidPort + Real p; + flow Real m_flow; + end FluidPort; + + model FlowModel + parameter Real states[2]; + Real m_flows[1]; + end FlowModel; + + model StaticPipe + FluidPort port_a; + FluidPort port_b; + FlowModel flowModel(states={ + Medium.setState_p(port_a.p), + Medium.setState_p(port_b.p)}); + equation + port_a.m_flow = flowModel.m_flows[1]; + end StaticPipe; + + model Top + StaticPipe pipe; + end Top; + "#; + + let compiled = rumoca::Compiler::new() + .model("Top") + .compile_str(source, "test.mo") + .expect("Top should compile"); + + let flat_dump = format!("{:#?}", compiled.flat); + assert!( + !flat_dump.contains("pipe.flowModel.port_a.p"), + "flowModel modifier should resolve port_a.p in enclosing scope, got over-qualified ref" + ); + assert!( + flat_dump.contains("pipe.port_a.p"), + "expected canonical enclosing connector pressure path in flattened model" + ); +} diff --git a/crates/rumoca/tests/indexed_cross_add_projection.rs b/crates/rumoca/tests/indexed_cross_add_projection.rs new file mode 100644 index 000000000..4c1721224 --- /dev/null +++ b/crates/rumoca/tests/indexed_cross_add_projection.rs @@ -0,0 +1,65 @@ +use rumoca_compile::{Session, SessionConfig}; +use rumoca_eval_solve::{eval_scalar_program_block, to_scalar_program_block}; +use rumoca_phase_solve::lower_solve_problem; +use rumoca_sim::{SimOptions, structurally_prepared_dae_for_simulation_artifact}; + +#[test] +fn modelica_cross_sum_preserves_each_indexed_component_through_solve_ir() { + let source = r#" +model IndexedCrossSum + Real r[3, 4]; + Real f[3, 4]; + Real m[3]; +equation + r = [1, 0, 0, 2; 0, 1, 0, 1; 0, 0, 1, 0]; + f = [0, 0, 1, 1; 1, 0, 0, 0; 0, 1, 0, 3]; + m = cross(r[:, 1], f[:, 1]) + + cross(r[:, 2], f[:, 2]) + + cross(r[:, 3], f[:, 3]) + + cross(r[:, 4], f[:, 4]); +end IndexedCrossSum; +"#; + let mut session = Session::new(SessionConfig::default()); + session + .add_document("indexed_cross_sum.mo", source) + .expect("fixture should parse and resolve"); + let dae = session + .compile_model("IndexedCrossSum") + .expect("fixture should compile through Modelica to DAE") + .dae; + let prepared = structurally_prepared_dae_for_simulation_artifact(&dae, &SimOptions::default()) + .expect("fixture DAE should pass structural preparation"); + let solve = lower_solve_problem(&prepared).expect("prepared DAE should lower to Solve IR"); + let residuals = to_scalar_program_block(&solve.continuous.residual) + .expect("residual ComputeBlock should have a scalar fallback"); + let mut y = vec![0.0; solve.solve_layout.solver_scalar_count()]; + for (name, value) in [ + ("r[1,1]", 1.0), + ("r[2,2]", 1.0), + ("r[3,3]", 1.0), + ("r[1,4]", 2.0), + ("r[2,4]", 1.0), + ("f[2,1]", 1.0), + ("f[3,2]", 1.0), + ("f[1,3]", 1.0), + ("f[1,4]", 1.0), + ("f[3,4]", 3.0), + ] { + let index = solve + .solve_layout + .solver_maps + .names + .iter() + .position(|candidate| candidate == name) + .unwrap_or_else(|| panic!("missing solver slot {name}")); + y[index] = value; + } + let mut outputs = vec![0.0; residuals.output_count()]; + eval_scalar_program_block(&residuals, &y, &[], 0.0, None, &mut outputs) + .expect("residual scalar programs should evaluate"); + + assert!( + outputs.windows(3).any(|window| window == [-4.0, 5.0, 0.0]), + "expected the three contact-moment residuals to stay distinct, got {outputs:?}" + ); +} diff --git a/crates/rumoca/tests/mod_propagation_nested_component_record.rs b/crates/rumoca/tests/mod_propagation_nested_component_record.rs new file mode 100644 index 000000000..2e36ba7b1 --- /dev/null +++ b/crates/rumoca/tests/mod_propagation_nested_component_record.rs @@ -0,0 +1,172 @@ +fn flat_expr_is_numeric_value(expr: &rumoca_core::Expression, expected: i64) -> bool { + match expr { + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Integer(value), + .. + } => *value == expected, + rumoca_core::Expression::Literal { + value: rumoca_core::Literal::Real(value), + .. + } => (*value - expected as f64).abs() <= f64::EPSILON, + _ => false, + } +} + +fn flat_expr_mentions_name(expr: &rumoca_core::Expression, needle: &str) -> bool { + match expr { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } => { + name.as_str().contains(needle) + || subscripts.iter().any(|subscript| match subscript { + rumoca_core::Subscript::Expr { expr, .. } => { + flat_expr_mentions_name(expr, needle) + } + _ => false, + }) + } + rumoca_core::Expression::FieldAccess { base, field, .. } => { + field.contains(needle) || flat_expr_mentions_name(base, needle) + } + rumoca_core::Expression::FunctionCall { name, args, .. } => { + name.as_str().contains(needle) + || args.iter().any(|arg| flat_expr_mentions_name(arg, needle)) + } + rumoca_core::Expression::Array { elements, .. } => elements + .iter() + .any(|element| flat_expr_mentions_name(element, needle)), + _ => false, + } +} + +#[test] +fn test_nested_component_record_modifier_resolves_sibling_alias_scope() { + let source = r#" + record CoreParameters + parameter Real p = 1; + end CoreParameters; + + record Data + parameter CoreParameters core; + end Data; + + model Core + parameter CoreParameters coreParameters; + parameter Real use = coreParameters.p; + end Core; + + partial model PartialMachine + parameter CoreParameters coreParameters; + Core core(final coreParameters = coreParameters); + end PartialMachine; + + model Machine + extends PartialMachine; + end Machine; + + model Motor + parameter Data data; + Machine machine(coreParameters = data.core); + end Motor; + + model Top + parameter Data data(core(p = 3)); + Motor motor(data = data); + end Top; + "#; + + let compiled = rumoca::Compiler::new() + .model("Top") + .compile_str(source, "test.mo") + .expect("Top should compile"); + + let binding = compiled + .flat + .variables + .iter() + .find(|(name, _)| name.as_str() == "motor.machine.core.coreParameters.p") + .and_then(|(_, var)| var.binding.as_ref()) + .expect("motor.machine.core.coreParameters.p should have binding"); + + assert!( + !flat_expr_mentions_name(binding, "motor.machine.data.core"), + "inherited nested record modifier must not be scoped under the machine component; binding={binding:?}" + ); + assert!( + !flat_expr_mentions_name(binding, "motor.machine.coreParameters"), + "inherited nested record modifier must not leave an intermediate record alias for DAE lowering; binding={binding:?}" + ); + assert!( + flat_expr_mentions_name(binding, "motor.data.core") + || flat_expr_mentions_name(binding, "data.core") + || flat_expr_is_numeric_value(binding, 3), + "inherited nested record modifier should resolve through the outer sibling alias scope; binding={binding:?}" + ); + + rumoca_phase_dae::to_dae(&compiled.flat).expect("Top should lower to DAE"); +} + +#[test] +fn test_nested_modifier_on_array_component_selects_element_row() { + let source = r#" + record Curve + parameter Real eta[:]; + end Curve; + + record Performance + parameter Curve motorEfficiency(eta={1.0}); + end Performance; + + model Pump + parameter Performance per; + Real y = per.motorEfficiency.eta[1]; + end Pump; + + model Top + parameter Real motorEta[2, 2] = {{0.87, 0.88}, {0.77, 0.78}}; + Pump pumps[2](per(motorEfficiency(eta=motorEta))); + Pump shared[2](per(motorEfficiency(each eta={5, 6}))); + end Top; + "#; + + let compiled = rumoca::Compiler::new() + .model("Top") + .compile_str(source, "test.mo") + .expect("nested modifier should select one row for each array element"); + + for index in 1..=2 { + let name = format!("pumps[{index}].per.motorEfficiency.eta"); + let eta = compiled + .flat + .variables + .get(&rumoca_core::VarName::new(&name)) + .unwrap_or_else(|| panic!("{name} should be present")); + assert_eq!(eta.dims, vec![2]); + match eta.binding.as_ref().expect("binding should be preserved") { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } => { + assert_eq!(name.as_str(), "motorEta"); + assert!( + matches!( + subscripts.as_slice(), + [rumoca_core::Subscript::Index { value, .. }] if *value == index + ), + "expected row {index}, got {subscripts:?}" + ); + } + other => panic!("expected selected source row, got {other:?}"), + } + + let shared = format!("shared[{index}].per.motorEfficiency.eta"); + let binding = compiled.flat.variables[&rumoca_core::VarName::new(&shared)] + .binding + .as_ref() + .expect("each binding should be preserved"); + let rumoca_core::Expression::Array { elements, .. } = binding else { + panic!("each modifier should keep the full array, got {binding:?}"); + }; + assert!(flat_expr_is_numeric_value(&elements[0], 5)); + assert!(flat_expr_is_numeric_value(&elements[1], 6)); + } +} diff --git a/crates/rumoca/tests/mod_propagation_test.rs b/crates/rumoca/tests/mod_propagation_test.rs index ba4a74406..15aa217ec 100644 --- a/crates/rumoca/tests/mod_propagation_test.rs +++ b/crates/rumoca/tests/mod_propagation_test.rs @@ -86,6 +86,53 @@ fn assert_bool_binding(overlay: &ast::InstanceOverlay, comp_name: &str, expected } } +fn assert_real_array_binding(data: &ast::InstanceData, expected: &[&str]) { + let binding = data.binding.as_ref().unwrap_or_else(|| { + panic!( + "{} should have binding", + data.qualified_name.to_flat_string() + ) + }); + let ast::Expression::Array { elements, .. } = binding else { + panic!( + "{} binding should be an array row, got {:?}", + data.qualified_name.to_flat_string(), + binding + ); + }; + let actual = elements + .iter() + .map(|expr| match expr { + ast::Expression::Terminal { + terminal_type: ast::TerminalType::UnsignedReal, + token, + .. + } => token.text.as_ref(), + other => panic!("expected real literal in row binding, got {other:?}"), + }) + .collect::>(); + assert_eq!(actual, expected); +} + +fn assert_indexed_single_ref(expr: &ast::Expression, name: &str, index: &str) { + let ast::Expression::ComponentReference(reference) = expr else { + panic!("expected indexed component reference source, got {expr:?}"); + }; + assert_eq!(reference.parts.len(), 1); + assert_eq!(reference.parts[0].ident.text.as_ref(), name); + let Some(subscripts) = reference.parts[0].subs.as_ref() else { + panic!("expected {name} to carry an element subscript"); + }; + assert_eq!(subscripts.len(), 1); + let ast::Subscript::Expression(ast::Expression::Terminal { token, .. }) = &subscripts[0] else { + panic!( + "expected literal integer subscript, got {:?}", + subscripts[0] + ); + }; + assert_eq!(token.text.as_ref(), index); +} + /// Helper: Assert that a component has the expected dimensions. fn assert_dims(overlay: &ast::InstanceOverlay, comp_name: &str, expected_dims: &[i64]) { let data = @@ -355,6 +402,756 @@ fn test_array_modifier_distribution_preserves_binding_source_scope() { ); } +#[test] +fn test_array_component_modifier_reference_selects_element_row() { + let source = r#" + model Tower + parameter Real v_flow_rate[:]; + end Tower; + + model TowerGroup + parameter Integer n = 2; + parameter Real v_flow_rate[n, 3]; + Tower ct[n](v_flow_rate = v_flow_rate); + end TowerGroup; + + model Top + TowerGroup group(v_flow_rate = {{1.0, 2.0, 3.0}, {4.0, 5.0, 6.0}}); + end Top; + "#; + + let (_tree, overlay) = instantiate_test_model(source, "Top"); + let first = find_component(&overlay, "group.ct[1].v_flow_rate") + .expect("first tower v_flow_rate should exist"); + let second = find_component(&overlay, "group.ct[2].v_flow_rate") + .expect("second tower v_flow_rate should exist"); + + assert!( + first.binding_from_modification, + "first tower binding should come from the array component modifier" + ); + assert!( + second.binding_from_modification, + "second tower binding should come from the array component modifier" + ); + assert_real_array_binding(first, &["1.0", "2.0", "3.0"]); + assert_real_array_binding(second, &["4.0", "5.0", "6.0"]); + + let first_source = first + .binding_source + .as_ref() + .expect("first tower should keep symbolic binding source"); + assert_indexed_single_ref(first_source, "v_flow_rate", "1"); +} + +#[test] +fn test_forwarded_colon_parameter_drives_record_size_dimension() { + let source = r#" + record Curve + parameter Real V_flow[:]; + parameter Real dp[size(V_flow, 1)]; + end Curve; + + model Mover + parameter Real V_flow[:]; + parameter Real dp[:]; + Curve pressure(V_flow = V_flow, dp = dp); + end Mover; + + model Top + parameter Real flows[2] = {1.0, 2.0}; + Mover m(V_flow = flows, dp = flows); + Real y; + equation + y = m.pressure.dp[1]; + end Top; + "#; + + let compiled = rumoca::Compiler::new() + .model("Top") + .compile_str(source, "test.mo") + .expect("forwarded record size dimensions should compile"); + + let dims = compiled + .flat + .variables + .get(&rumoca_core::VarName::new("m.pressure.dp")) + .map(|var| var.dims.clone()) + .expect("m.pressure.dp should be in flat variables"); + assert_eq!(dims, vec![2]); +} + +#[test] +fn test_nested_record_parameter_size_dimension_resolves_through_alias() { + let source = r#" + record FlowParameters + parameter Real V_flow[:]; + parameter Real dp[size(V_flow, 1)]; + end FlowParameters; + + record EfficiencyParameters + parameter Real V_flow[:]; + parameter Real eta[size(V_flow, 1)]; + end EfficiencyParameters; + + record Generic + parameter FlowParameters pressure; + parameter EfficiencyParameters motorEfficiency; + parameter EfficiencyParameters hydraulicEfficiency; + end Generic; + + model Interface + parameter Generic per; + FlowParameters pressure = per.pressure; + EfficiencyParameters motorEfficiency = per.motorEfficiency; + final parameter Real motDer[size(per.motorEfficiency.V_flow, 1)]; + final parameter Real hydDer[size(per.hydraulicEfficiency.V_flow, 1)]; + end Interface; + + model Top + parameter Real flows[3] = {1.0, 2.0, 3.0}; + parameter Real heads[3] = {10.0, 20.0, 30.0}; + parameter Real effs[3] = {0.1, 0.2, 0.3}; + Interface mover( + per( + pressure(V_flow = flows, dp = heads), + motorEfficiency(V_flow = flows, eta = effs), + hydraulicEfficiency(V_flow = flows, eta = effs))); + Real y; + equation + y = mover.pressure.dp[1] + mover.motDer[1] + mover.hydDer[1]; + end Top; + "#; + + let compiled = rumoca::Compiler::new() + .model("Top") + .compile_str(source, "test.mo") + .expect("nested record size dimensions should compile"); + + for name in [ + "mover.per.pressure.dp", + "mover.per.motorEfficiency.eta", + "mover.pressure.dp", + "mover.motorEfficiency.eta", + "mover.motDer", + "mover.hydDer", + ] { + let dims = compiled + .flat + .variables + .get(&rumoca_core::VarName::new(name)) + .map(|var| var.dims.clone()) + .unwrap_or_else(|| panic!("{name} should be in flat variables")); + assert_eq!(dims, vec![3], "{name} should keep the V_flow row count"); + } +} + +#[test] +fn test_array_component_nested_record_defaults_drive_inner_size_dimensions() { + let source = r#" + record FlowParameters + parameter Real V_flow[:]; + parameter Real dp[size(V_flow, 1)]; + end FlowParameters; + + record EfficiencyParameters + parameter Real V_flow[:]; + parameter Real eta[size(V_flow, 1)]; + end EfficiencyParameters; + + record PowerParameters + parameter Real V_flow[:]; + parameter Real P[size(V_flow, 1)]; + end PowerParameters; + + record Generic + parameter FlowParameters pressure(V_flow = {0.0, 0.0}, dp = {0.0, 0.0}); + parameter EfficiencyParameters hydraulicEfficiency(V_flow = {0.0}, eta = {0.7}); + parameter EfficiencyParameters motorEfficiency(V_flow = {0.0}, eta = {0.7}); + parameter PowerParameters power(V_flow = {0.0}, P = {0.0}); + parameter Boolean motorCooledByFluid = true; + end Generic; + + model Interface + parameter Generic per; + parameter Integer nOri; + final parameter Real motDer[size(per.motorEfficiency.V_flow, 1)]; + final parameter Real hydDer[size(per.hydraulicEfficiency.V_flow, 1)]; + end Interface; + + partial model PartialMover + parameter Generic per; + final parameter Integer nOri = size(per.pressure.V_flow, 1); + Interface eff( + per( + final hydraulicEfficiency = per.hydraulicEfficiency, + final motorEfficiency = per.motorEfficiency, + final motorCooledByFluid = per.motorCooledByFluid, + final power = per.power), + final nOri = nOri); + end PartialMover; + + model SpeedMover + extends PartialMover( + eff(per(final pressure = per.pressure))); + end SpeedMover; + + model WithoutMotor + parameter Real HydEff[:]; + parameter Real MotEff[:]; + parameter Real VolFloCur[:]; + parameter Real PreCur[:]; + SpeedMover varSpeFloMov( + per( + pressure(V_flow = VolFloCur, dp = PreCur), + hydraulicEfficiency(eta = HydEff, V_flow = VolFloCur), + motorEfficiency(eta = MotEff, V_flow = VolFloCur))); + end WithoutMotor; + + model PumpSystem + parameter Integer n = 2; + parameter Real HydEff[n, 3]; + parameter Real MotEff[n, 3]; + parameter Real VolFloCur[n, 3]; + parameter Real PreCur[n, 3]; + WithoutMotor pum[n]( + HydEff = HydEff, + MotEff = MotEff, + VolFloCur = VolFloCur, + PreCur = PreCur); + end PumpSystem; + + model Top + PumpSystem sys( + HydEff = {{1.0, 1.0, 1.0}, {1.0, 1.0, 1.0}}, + MotEff = {{0.6, 0.7, 0.8}, {0.6, 0.7, 0.8}}, + VolFloCur = {{0.0, 1.0, 2.0}, {0.0, 1.0, 2.0}}, + PreCur = {{20.0, 10.0, 0.0}, {20.0, 10.0, 0.0}}); + Real y; + equation + y = sys.pum[1].varSpeFloMov.eff.motDer[1] + + sys.pum[1].varSpeFloMov.eff.hydDer[1]; + end Top; + "#; + + let compiled = rumoca::Compiler::new() + .model("Top") + .compile_str(source, "test.mo") + .expect("array component record defaults should drive inner size dimensions"); + + for name in [ + "sys.pum[1].varSpeFloMov.per.pressure.dp", + "sys.pum[1].varSpeFloMov.per.motorEfficiency.eta", + "sys.pum[1].varSpeFloMov.eff.per.pressure.dp", + "sys.pum[1].varSpeFloMov.eff.per.motorEfficiency.eta", + "sys.pum[1].varSpeFloMov.eff.motDer", + "sys.pum[1].varSpeFloMov.eff.hydDer", + ] { + let dims = compiled + .flat + .variables + .get(&rumoca_core::VarName::new(name)) + .map(|var| var.dims.clone()) + .unwrap_or_else(|| panic!("{name} should be in flat variables")); + assert_eq!( + dims, + vec![3], + "{name} should keep the selected pump curve row count" + ); + } +} + +#[test] +fn test_array_component_symbolic_row_modifier_keeps_source_shape_for_record_size() { + let source = r#" + record Curve + parameter Real x[:]; + parameter Real y[size(x, 1)]; + end Curve; + + model Child + parameter Real row[:]; + Curve curve(x = row, y = row); + end Child; + + model Group + parameter Integer n = 2; + parameter Real rows[n, 3]; + Child child[n](row = rows); + end Group; + + model Top + parameter Real upstream[2, 3] = {{1.0, 2.0, 3.0}, {4.0, 5.0, 6.0}}; + Group group(rows = upstream); + Real y; + equation + y = group.child[1].curve.y[1]; + end Top; + "#; + + let compiled = rumoca::Compiler::new() + .model("Top") + .compile_str(source, "test.mo") + .expect("symbolic row modifiers should preserve source array shape"); + + for name in [ + "group.child[1].row", + "group.child[1].curve.x", + "group.child[1].curve.y", + ] { + let dims = compiled + .flat + .variables + .get(&rumoca_core::VarName::new(name)) + .map(|var| var.dims.clone()) + .unwrap_or_else(|| panic!("{name} should be in flat variables")); + assert_eq!(dims, vec![3], "{name} should keep the selected row width"); + } + + let row = compiled + .flat + .variables + .get(&rumoca_core::VarName::new("group.child[1].row")) + .expect("group.child[1].row should be in flat variables"); + let row_binding = row + .binding + .as_ref() + .expect("group.child[1].row should keep a flat binding"); + match row_binding { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } => { + assert_eq!(name.as_str(), "group.rows"); + assert!( + matches!( + subscripts.as_slice(), + [rumoca_core::Subscript::Index { value: 1, .. }] + ), + "group.child[1].row should select the first source row, got {subscripts:?}" + ); + } + other => panic!("group.child[1].row should bind to selected source row, got {other:?}"), + } +} + +#[test] +fn test_nested_array_component_modifier_selects_element_row() { + let source = r#" + model Tower + parameter Real v_flow_rate[:]; + end Tower; + + model TowerGroup + parameter Integer n = 3; + Tower ct[n]; + end TowerGroup; + + model Wrapper + parameter Real rows[3, 3]; + TowerGroup group(ct(v_flow_rate = rows)); + end Wrapper; + + model Top + parameter Real upstream[3, 3] = { + {0.0, 0.5, 1.0}, + {0.0, 0.6, 1.0}, + {0.0, 0.7, 1.0}}; + Wrapper wrapper(rows = upstream); + Real y; + equation + y = wrapper.group.ct[1].v_flow_rate[1]; + end Top; + "#; + + let compiled = rumoca::Compiler::new() + .model("Top") + .compile_str(source, "test.mo") + .expect("nested array component modifier should select row per element"); + + let first = compiled + .flat + .variables + .get(&rumoca_core::VarName::new( + "wrapper.group.ct[1].v_flow_rate", + )) + .expect("first tower v_flow_rate should be in flat variables"); + assert_eq!( + first.dims, + vec![3], + "nested modifier should distribute the selected row shape" + ); + match first + .binding + .as_ref() + .expect("first tower should keep a flat binding") + { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } => { + assert_eq!(name.as_str(), "wrapper.rows"); + assert!( + matches!( + subscripts.as_slice(), + [rumoca_core::Subscript::Index { value: 1, .. }] + ), + "nested modifier should select the first source row, got {subscripts:?}" + ); + } + other => panic!("expected first tower binding to select source row, got {other:?}"), + } +} + +#[test] +fn test_record_field_modifier_drives_sibling_size_dimension_in_nested_array_component() { + let source = r#" + record Fan + parameter Real r_V[:]; + parameter Real r_P[size(r_V, 1)]; + end Fan; + + model YorkCalc + parameter Fan fanRelPow(r_V = {0.0, 0.5, 1.0}, r_P = {0.0, 0.2, 1.0}); + final parameter Real fanRelPowDer[size(fanRelPow.r_V, 1)]; + end YorkCalc; + + model Tower + parameter Real v_flow_rate[:]; + parameter Real eta[:]; + YorkCalc yorkCalc(fanRelPow(r_V = v_flow_rate, r_P = eta)); + end Tower; + + model TowerGroup + parameter Integer n = 3; + Tower ct[n]; + end TowerGroup; + + model Wrapper + parameter Real rows[3, 3]; + parameter Real powers[3, 3]; + TowerGroup group(ct(v_flow_rate = rows, eta = powers)); + end Wrapper; + + model Top + parameter Real upstream[3, 3] = { + {0.0, 0.5, 1.0}, + {0.0, 0.6, 1.0}, + {0.0, 0.7, 1.0}}; + parameter Real fanPower[3, 3] = { + {0.0, 0.2, 1.0}, + {0.0, 0.3, 1.0}, + {0.0, 0.4, 1.0}}; + Wrapper wrapper(rows = upstream, powers = fanPower); + Real y; + equation + y = wrapper.group.ct[1].yorkCalc.fanRelPow.r_P[1] + + wrapper.group.ct[1].yorkCalc.fanRelPowDer[1]; + end Top; + "#; + + let compiled = rumoca::Compiler::new() + .model("Top") + .compile_str(source, "test.mo") + .expect("record field modifier should size sibling record fields"); + + for name in [ + "wrapper.group.ct[1].yorkCalc.fanRelPow.r_V", + "wrapper.group.ct[1].yorkCalc.fanRelPow.r_P", + "wrapper.group.ct[1].yorkCalc.fanRelPowDer", + ] { + let dims = compiled + .flat + .variables + .get(&rumoca_core::VarName::new(name)) + .map(|var| var.dims.clone()) + .unwrap_or_else(|| panic!("{name} should be in flat variables")); + assert_eq!( + dims, + vec![3], + "{name} should keep the selected tower fan curve row count" + ); + } +} + +#[test] +fn test_nested_attribute_modifier_on_array_component_selects_element_row() { + let source = r#" + record Curve + parameter Real eta[:]; + end Curve; + + record Performance + parameter Curve motorEfficiency(eta={1.0}); + end Performance; + + model Pump + parameter Performance per; + Real y; + equation + y = per.motorEfficiency.eta[1]; + end Pump; + + model PumpSystem + parameter Integer n = 3; + parameter Real Motor_eta[n, 2] = { + {0.87, 0.88}, + {0.77, 0.78}, + {0.67, 0.68}}; + Pump pumConSpe[n](per(motorEfficiency(eta=Motor_eta))); + end PumpSystem; + + model Top + PumpSystem system; + Real y; + equation + y = system.pumConSpe[1].y; + end Top; + "#; + + let compiled = rumoca::Compiler::new() + .model("Top") + .compile_str(source, "test.mo") + .expect("nested attribute modifier on array component should select one row"); + + let eta = compiled + .flat + .variables + .get(&rumoca_core::VarName::new( + "system.pumConSpe[1].per.motorEfficiency.eta", + )) + .expect("first pump motor efficiency eta should be in flat variables"); + assert_eq!( + eta.dims, + vec![2], + "first pump motor efficiency eta should keep one selected row" + ); + match eta + .binding + .as_ref() + .expect("first pump motor efficiency eta should keep a flat binding") + { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } => { + assert_eq!(name.as_str(), "system.Motor_eta"); + assert!( + matches!( + subscripts.as_slice(), + [rumoca_core::Subscript::Index { value: 1, .. }] + ), + "nested attribute modifier should select first source row, got {subscripts:?}" + ); + } + other => panic!("expected selected source row binding, got {other:?}"), + } +} + +#[test] +fn test_forwarded_parent_parameter_remains_visible_to_child_modifier_rhs() { + let source = r#" + model Damper + parameter Real dpValve_nominal; + Real y; + equation + y = dpValve_nominal; + end Damper; + + model Terminal + parameter Real PreDroAir; + Damper dam(dpValve_nominal=PreDroAir); + Real y; + equation + y = dam.y; + end Terminal; + + model FiveZone + parameter Real PreDroAir1; + Terminal vAV1(PreDroAir=PreDroAir1); + Real y; + equation + y = vAV1.y; + end FiveZone; + + model Floor + parameter Real PreDroAir1; + FiveZone fivZonVAV(PreDroAir1=PreDroAir1); + Real y; + equation + y = fivZonVAV.y; + end Floor; + + model Wrapper + parameter Real PreDroAir[5] = {200, 124, 124, 124, 124}; + Floor floor(PreDroAir1=PreDroAir[1]); + Real y; + equation + y = floor.y; + end Wrapper; + "#; + + let compiled = rumoca::Compiler::new() + .model("Wrapper") + .compile_str(source, "test.mo") + .expect("forwarded parent parameter should remain visible to child modifier RHS"); + + let pre_dro_air = compiled + .flat + .variables + .get(&rumoca_core::VarName::new("floor.fivZonVAV.vAV1.PreDroAir")) + .expect("terminal PreDroAir should be in flat variables"); + assert!( + pre_dro_air.binding.is_some(), + "terminal PreDroAir should keep the forwarded binding" + ); + + let dp = compiled + .flat + .variables + .get(&rumoca_core::VarName::new( + "floor.fivZonVAV.vAV1.dam.dpValve_nominal", + )) + .expect("damper dpValve_nominal should be in flat variables"); + assert!( + dp.binding.is_some(), + "damper dpValve_nominal should bind through terminal PreDroAir" + ); + match dp + .binding + .as_ref() + .expect("damper dpValve_nominal should keep a flat binding") + { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } => { + assert_eq!(name.as_str(), "floor.fivZonVAV.vAV1.PreDroAir"); + assert!( + subscripts.is_empty(), + "scalar forwarded parameter must not receive an array-element subscript, got {subscripts:?}" + ); + } + other => { + panic!("expected damper dpValve_nominal to reference terminal PreDroAir, got {other:?}") + } + } +} + +#[test] +fn test_array_element_local_vector_modifier_rhs_is_not_indexed_again() { + let source = r#" + record PressureCurve + parameter Real V_flow[:]; + end PressureCurve; + + record Performance + parameter PressureCurve pressure; + end Performance; + + model FlowMover + parameter Performance per; + Real y; + equation + y = per.pressure.V_flow[1]; + end FlowMover; + + model Pump + parameter Real VolFloCur[:]; + FlowMover varSpeFloMov(per(pressure(V_flow=VolFloCur))); + Real y; + equation + y = varSpeFloMov.y; + end Pump; + + model PumpSystem + parameter Integer n = 2; + parameter Real VolFloCur[n, 3] = { + {0.1, 0.2, 0.3}, + {1.1, 1.2, 1.3}}; + Pump pum[n](VolFloCur=VolFloCur); + end PumpSystem; + + model Top + PumpSystem system; + Real y; + equation + y = system.pum[1].y; + end Top; + "#; + + let compiled = rumoca::Compiler::new() + .model("Top") + .compile_str(source, "test.mo") + .expect("array element local vector modifier RHS should keep vector shape"); + + let v_flow = compiled + .flat + .variables + .get(&rumoca_core::VarName::new( + "system.pum[1].varSpeFloMov.per.pressure.V_flow", + )) + .expect("pressure V_flow should be in flat variables"); + assert_eq!(v_flow.dims, vec![3]); + match v_flow + .binding + .as_ref() + .expect("pressure V_flow should keep a binding") + { + rumoca_core::Expression::VarRef { + name, subscripts, .. + } => { + assert_eq!(name.as_str(), "system.pum[1].VolFloCur"); + assert!( + subscripts.is_empty(), + "local vector parameter must not receive outer array element subscript, got {subscripts:?}" + ); + } + other => panic!("expected pressure V_flow to reference local VolFloCur, got {other:?}"), + } +} + +#[test] +fn test_array_element_colon_parameter_uses_local_binding_shape() { + let source = r#" + model Base + parameter Real stageInputs[:]; + end Base; + + model Flow + extends Base(stageInputs=massFlowRates); + parameter Real m_flow_nominal; + parameter Real per_speeds[1] = {1}; + parameter Real massFlowRates[:] = + m_flow_nominal * {per_speeds[i] / per_speeds[end] for i in 1:size(per_speeds, 1)}; + end Flow; + + model System + parameter Integer n = 3; + parameter Real m_flow_nominal[n] = {2, 3, 4}; + Flow pump[n](m_flow_nominal=m_flow_nominal); + end System; + + model Top + System system; + end Top; + "#; + + let compiled = rumoca::Compiler::new() + .model("Top") + .compile_str(source, "test.mo") + .expect("array element colon parameter should infer shape from local binding"); + + for name in ["system.pump[1].stageInputs", "system.pump[1].massFlowRates"] { + let var = compiled + .flat + .variables + .get(&rumoca_core::VarName::new(name)) + .unwrap_or_else(|| panic!("{name} should be in flat variables")); + assert_eq!( + var.dims, + vec![1], + "{name} should use the local per_speeds binding shape, not the parent pump array length" + ); + } +} + /// Test that typecheck_instanced evaluates dimensions correctly. #[test] fn test_dimension_evaluation_after_typecheck() { @@ -1176,56 +1973,3 @@ where } } } - -#[test] -fn test_flowmodel_modifier_keeps_enclosing_port_scope() { - let source = r#" - package Medium - function setState_p - input Real p; - output Real s; - algorithm - s := p; - end setState_p; - end Medium; - - connector FluidPort - Real p; - flow Real m_flow; - end FluidPort; - - model FlowModel - parameter Real states[2]; - Real m_flows[1]; - end FlowModel; - - model StaticPipe - FluidPort port_a; - FluidPort port_b; - FlowModel flowModel(states={ - Medium.setState_p(port_a.p), - Medium.setState_p(port_b.p)}); - equation - port_a.m_flow = flowModel.m_flows[1]; - end StaticPipe; - - model Top - StaticPipe pipe; - end Top; - "#; - - let compiled = rumoca::Compiler::new() - .model("Top") - .compile_str(source, "test.mo") - .expect("Top should compile"); - - let flat_dump = format!("{:#?}", compiled.flat); - assert!( - !flat_dump.contains("pipe.flowModel.port_a.p"), - "flowModel modifier should resolve port_a.p in enclosing scope, got over-qualified ref" - ); - assert!( - flat_dump.contains("pipe.port_a.p"), - "expected canonical enclosing connector pressure path in flattened model" - ); -} diff --git a/crates/rumoca/tests/msl_sim_regression.rs b/crates/rumoca/tests/msl_sim_regression.rs index 9f3a07a1f..82253552d 100644 --- a/crates/rumoca/tests/msl_sim_regression.rs +++ b/crates/rumoca/tests/msl_sim_regression.rs @@ -80,6 +80,15 @@ fn max_abs_series_delta(left: &[f64], right: &[f64]) -> f64 { .fold(0.0, f64::max) } +fn series_value_at(result: &SimResult, name: &str, time: f64) -> f64 { + let sample_idx = result + .times + .iter() + .position(|candidate| (*candidate - time).abs() <= 1.0e-12) + .unwrap_or_else(|| panic!("simulation result missing t={time}: {:?}", result.times)); + result_series(result, &[name])[sample_idx] +} + fn variable_is_state(result: &SimResult, name: &str) -> bool { result .variable_meta @@ -184,3 +193,43 @@ fn pid_msl_responds_to_step_error() { "expected PIDMSL controller output to become nonzero, max |pid.y|={pid_y_max}" ); } + +#[test] +fn exactly_clocked_drive_refreshes_inferred_sample_and_controller_on_same_tick() { + let msl_compiler = require_msl_compiler(); + let model_path = cached_msl_root() + .expect("cached MSL root is required for this concrete example") + .join("Modelica 4.1.0/Clocked/Examples/SimpleControlledDrive/ExactlyClockedWithDiscreteController.mo"); + let compiled = msl_compiler + .model( + "Modelica.Clocked.Examples.SimpleControlledDrive.ExactlyClockedWithDiscreteController", + ) + .compile_file(model_path.to_string_lossy().as_ref()) + .expect("ExactlyClockedWithDiscreteController should compile"); + let result = simulate_dae_with_diagnostics( + &compiled.dae, + &SimOptions { + // Run the model's full experiment horizon. A short run through the + // second clock tick misses later continuous/event instability. + t_end: 5.0, + max_wall_seconds: Some(10.0), + ..SimOptions::default() + }, + ) + .expect("ExactlyClockedWithDiscreteController should simulate"); + + let sampled_speed = series_value_at(&result, "sample1.y", 0.2); + let speed = series_value_at(&result, "speed.w", 0.2); + assert!( + (sampled_speed - speed).abs() <= 1.0e-9 && sampled_speed > 0.1, + "the inferred-clock sample must read the projected speed source on the second tick: \ + sample1.y={sampled_speed}, speed.w={speed}" + ); + let pi_y = series_value_at(&result, "PI.y", 0.2); + let held = series_value_at(&result, "hold1.y", 0.2); + assert!( + (pi_y - 3.3).abs() <= 1.0e-8 && (held - pi_y).abs() <= 1.0e-9, + "the PI and hold chain must consume the refreshed sample on the same tick: \ + PI.y={pi_y}, hold1.y={held}" + ); +} diff --git a/crates/rumoca/tests/msl_table_regression.rs b/crates/rumoca/tests/msl_table_regression.rs index d95431d93..4b2c3d9e1 100644 --- a/crates/rumoca/tests/msl_table_regression.rs +++ b/crates/rumoca/tests/msl_table_regression.rs @@ -1,7 +1,9 @@ //! Regression tests for MSL table-driven no-state simulation behavior. +use rumoca::Compiler; use rumoca_core::{SourceId, Span}; use rumoca_ir_dae as dae; +use rumoca_ir_solve::{LinearOp, LinearOpSliceKind, SolveVisitor}; use rumoca_sim::{SimOptions, simulate_dae}; fn fixture_span() -> Span { @@ -452,7 +454,13 @@ fn msl_buf3s_no_state_model() -> dae::Dae { insert_buf3s_parameters(&mut dae_model); insert_buf3s_discrete_vars(&mut dae_model); populate_buf3s_equations(&mut dae_model); - dae_model.events.scheduled_time_events.push(0.0); + dae_model + .events + .scheduled_time_events + .push(dae::DaeScheduledTimeEvent { + time: 0.0, + source_span: Some(fixture_span()), + }); dae_model } @@ -577,3 +585,294 @@ fn native_no_state_simulation_refreshes_table_driven_boolean_discrete_chain() { .expect("trace should contain c"); assert_eq!(sim.data[c_idx], vec![1.0, 1.0, 0.0]); } + +#[test] +fn source_table_algorithm_keeps_last_write_at_each_threshold() { + let source = r#" +block TwoPointTable + parameter Integer x[:] = {4}; + parameter Real t[size(x, 1)] = {1}; + parameter Integer y0 = 1; + final parameter Integer n = size(x, 1); + output Integer y; +algorithm + y := y0; + for i in 1:n loop + if time >= t[i] then + y := x[i]; + end if; + end for; +end TwoPointTable; + +model TwoThresholdTable + TwoPointTable table(y0 = 3, x = {4, 3}, t = {1, 3}); +end TwoThresholdTable; +"#; + let compiled = Compiler::new() + .model("TwoThresholdTable") + .compile_str(source, "TwoThresholdTable.mo") + .expect("two-threshold table fixture should compile"); + let sim = simulate_dae( + &compiled.dae, + &SimOptions { + t_end: 4.0, + dt: Some(1.0), + ..Default::default() + }, + ) + .expect("two-threshold table fixture should simulate"); + let y_idx = sim + .names + .iter() + .position(|name| name == "table.y") + .expect("trace should contain table.y"); + let value_at = |time: f64| { + sim.times + .iter() + .zip(&sim.data[y_idx]) + .rev() + .find(|(sample_time, _)| (**sample_time - time).abs() <= 1.0e-12) + .map(|(_, value)| *value) + .unwrap_or_else(|| panic!("trace should contain t={time}")) + }; + + // MLS §11.1/§11.2 and Appendix B: statements and loop iterations execute + // sequentially, so the last active assignment owns the algorithm output. + assert_eq!(value_at(2.0), 4.0, "first threshold should select x[1]"); + assert_eq!(value_at(3.0), 3.0, "second threshold should select x[2]"); +} + +const MODIFIED_TABLES_SOURCE: &str = r#" +package Modelica +package Blocks +package Types +function ExternalCombiTimeTable + input String tableName; + input String fileName; + input Real table[:, :]; + input Real startTime; + input Integer columns[:]; + input Integer smoothness; + input Integer extrapolation; + input Real shiftTime = 0.0; + input Integer timeEvents = 1; + input Boolean verboseRead = false; + input String delimiter = ","; + input Integer nHeaderLines = 0; + output Real tableID; +external "C" tableID = ModelicaStandardTables_CombiTimeTable_init3( + fileName, tableName, table, size(table, 1), size(table, 2), startTime, + columns, size(columns, 1), smoothness, extrapolation, shiftTime, timeEvents, + verboseRead, delimiter, nHeaderLines); +end ExternalCombiTimeTable; +end Types; + +package Tables +package Internal +pure function getTimeTableValueNoDer + input Real tableID; + input Integer icol; + input Real timeIn; + input Real nextTimeEvent; + input Real pre_nextTimeEvent; + output Real y; +external "C" y = ModelicaStandardTables_CombiTimeTable_getValue( + tableID, icol, timeIn, nextTimeEvent, pre_nextTimeEvent); +end getTimeTableValueNoDer; +end Internal; +end Tables; +end Blocks; +end Modelica; + +block CombiTimeTable + parameter Real table[:, :]; + parameter Integer columns[:] = {2}; + parameter Integer smoothness = 3; + parameter Integer extrapolation = 1; + output Real y; +protected + parameter Real tableID = Modelica.Blocks.Types.ExternalCombiTimeTable( + "NoName", "NoName", table, 0.0, columns, smoothness, extrapolation, + 0.0, 1, false, ",", 0); +equation + y = Modelica.Blocks.Tables.Internal.getTimeTableValueNoDer( + tableID, 1, time, 0.0, 0.0); +end CombiTimeTable; + +block BooleanTable + parameter Real table[:] = {0, 1}; + parameter Boolean startValue = false; + final parameter Integer n = size(table, 1); + output Boolean y; + CombiTimeTable combiTimeTable( + final table = if n > 0 then + if startValue then + [table[1], 1.0; table, {mod(i + 1, 2.0) for i in 1:n}] + else + [table[1], 0.0; table, {mod(i, 2.0) for i in 1:n}] + else [0.0, 0.0]); +equation + y = combiTimeTable.y >= 0.5; +end BooleanTable; + +block IntegerTable + parameter Real table[:, 2] = [0, 0]; + output Integer y; + CombiTimeTable combiTimeTable(final table = table); +equation + y = integer(combiTimeTable.y); +end IntegerTable; + +model ModifiedTables + BooleanTable booleanTable(table = {.05, .15}); + IntegerTable integerTable(table = [0, 1; .025, 2; .05, 0; .075, -1]); +end ModifiedTables; +"#; + +struct LookupCallCounter(usize); + +impl rumoca_core::ExpressionVisitor for LookupCallCounter { + fn visit_function_call( + &mut self, + name: &rumoca_core::Reference, + args: &[rumoca_core::Expression], + is_constructor: bool, + ) { + let segments = name.segments(); + if name.last_segment() == "getTimeTableValueNoDer" + || (segments.last() == Some(&"y") + && segments.get(segments.len().saturating_sub(2)) + == Some(&"getTimeTableValueNoDer")) + { + self.0 += 1; + } + self.walk_function_call(name, args, is_constructor); + } +} + +#[derive(Default)] +struct TableLookupCounter(usize); + +impl SolveVisitor for TableLookupCounter { + type Error = std::convert::Infallible; + + fn visit_linear_op( + &mut self, + _kind: LinearOpSliceKind, + _op_index: usize, + op: &LinearOp, + ) -> Result<(), Self::Error> { + if matches!(op, LinearOp::TableLookup { .. }) { + self.0 += 1; + } + Ok(()) + } +} + +fn prepared_lookup_call_count(prepared_dae: &dae::Dae) -> usize { + let mut counter = LookupCallCounter(0); + for equation in &prepared_dae.continuous.equations { + rumoca_core::ExpressionVisitor::visit_expression(&mut counter, &equation.rhs); + } + counter.0 +} + +fn solve_lookup_count(solve_model: &rumoca_ir_solve::SolveModel) -> usize { + let mut counter = TableLookupCounter::default(); + counter + .visit_solve_model(solve_model) + .expect("infallible SolveModel traversal"); + counter.0 +} + +fn lookup_table(tables: &[rumoca_core::ExternalTableData], table_id: u64, time: f64) -> f64 { + let row = [ + LinearOp::Const { + dst: 0, + value: table_id as f64, + }, + LinearOp::Const { dst: 1, value: 1.0 }, + LinearOp::Const { + dst: 2, + value: time, + }, + LinearOp::TableLookup { + dst: 3, + table_id: 0, + column: 1, + input: 2, + }, + LinearOp::StoreOutput { src: 3 }, + ]; + rumoca_eval_solve::eval_row_with_context( + &row, + &[], + &[], + time, + rumoca_eval_solve::RowEvalContext { + external_tables: Some(tables), + ..Default::default() + }, + ) + .expect("materialized SolveModel TableLookup op should evaluate") +} + +#[test] +fn source_instance_modifiers_materialize_boolean_and_integer_time_tables() { + let compiled = Compiler::new() + .model("ModifiedTables") + .compile_str(MODIFIED_TABLES_SOURCE, "ModifiedTables.mo") + .expect("real Modelica table fixture should compile"); + let options = SimOptions { + t_end: 0.15, + dt: Some(0.001), + ..Default::default() + }; + let prepared_dae = + rumoca_sim::structurally_prepared_dae_for_simulation_artifact(&compiled.dae, &options) + .expect("real Modelica table fixture should survive structural preparation"); + assert_eq!( + prepared_lookup_call_count(&prepared_dae), + 2, + "numeric table equations must not be pruned as String metadata" + ); + let solve_model = rumoca_sim::lower_dae_for_simulation(&compiled.dae, &options) + .expect("real Modelica table fixture should lower to simulation SolveModel"); + let tables = solve_model.external_tables.as_slice(); + + assert!( + tables + .iter() + .any(|table| table.data == vec![vec![0.05, 0.0], vec![0.05, 1.0], vec![0.15, 0.0]]), + "BooleanTable modifier must materialize its duplicate-knot matrix; tables={tables:?}" + ); + assert!( + tables.iter().any(|table| table.data + == vec![ + vec![0.0, 1.0], + vec![0.025, 2.0], + vec![0.05, 0.0], + vec![0.075, -1.0], + ]), + "IntegerTable modifier must materialize its matrix; tables={tables:?}" + ); + + assert!( + solve_lookup_count(&solve_model) >= 2, + "both qualified getTimeTableValueNoDer equations must survive structural lowering" + ); + + let boolean_table = tables + .iter() + .find(|table| table.data.len() == 3) + .expect("BooleanTable data"); + let integer_table = tables + .iter() + .find(|table| table.data.len() == 4) + .expect("IntegerTable data"); + assert_eq!(lookup_table(tables, boolean_table.id, 0.049), 0.0); + assert_eq!(lookup_table(tables, boolean_table.id, 0.05), 1.0); + assert_eq!(lookup_table(tables, boolean_table.id, 0.149), 1.0); + assert_eq!(lookup_table(tables, boolean_table.id, 0.15), 0.0); + assert_eq!(lookup_table(tables, integer_table.id, 0.0), 1.0); +} diff --git a/crates/rumoca/tests/omc_differential_semantics.rs b/crates/rumoca/tests/omc_differential_semantics.rs index 22da12788..3b42fa8f7 100644 --- a/crates/rumoca/tests/omc_differential_semantics.rs +++ b/crates/rumoca/tests/omc_differential_semantics.rs @@ -1,7 +1,7 @@ use std::fs; use std::process::Command; -use tempfile::tempdir; +use tempfile::tempdir_in; const ENCAPSULATED_SCOPE_SOURCE: &str = r#" package P @@ -23,24 +23,25 @@ fn encapsulated_scope_rejection_matches_omc() { return; } - let dir = tempdir().expect("tempdir"); + let cwd = std::env::current_dir().expect("current dir"); + let dir = tempdir_in(cwd).expect("tempdir in mounted workspace"); + let model_path = dir.path().join("EncapsulatedScope.mo"); + let script_path = dir.path().join("check.mos"); + fs::write(&model_path, ENCAPSULATED_SCOPE_SOURCE).expect("write model"); fs::write( - dir.path().join("EncapsulatedScope.mo"), - ENCAPSULATED_SCOPE_SOURCE, - ) - .expect("write model"); - fs::write( - dir.path().join("check.mos"), - r#"loadFile("EncapsulatedScope.mo"); + &script_path, + format!( + r#"loadFile("{}"); checkModel(P.M); getErrorString(); "#, + model_path.display() + ), ) .expect("write OMC script"); let omc = Command::new("omc") - .arg("check.mos") - .current_dir(dir.path()) + .arg(&script_path) .output() .expect("run omc"); let omc_output = format!( diff --git a/crates/rumoca/tests/pipeline_cases/alias_scope_and_dims.rs b/crates/rumoca/tests/pipeline_cases/alias_scope_and_dims.rs index 1e483ea92..4b37149e8 100644 --- a/crates/rumoca/tests/pipeline_cases/alias_scope_and_dims.rs +++ b/crates/rumoca/tests/pipeline_cases/alias_scope_and_dims.rs @@ -1,6 +1,171 @@ +// SPEC_0021 file-size exception: alias/scope/dimension pipeline cases share +// package fixtures across resolve, typecheck, flatten, and DAE assertions. split plan: +// move package override, record alias, and dimension cases by fixture. use super::*; mod package_override_dimension_tests; + +/// MLS §4.6 + §4.7: record component type checks must classify the resolved +/// component type class, not a later short-name lookup. Unit aliases such as +/// `Temperature = Real(...)` remain specialized class `type` declarations and +/// are legal record fields. +#[test] +fn test_record_component_type_alias_uses_resolved_type_def_id_for_er023() { + let source = r#" +model Temperature +end Temperature; + +package Modelica + package Units + package SI + type Temperature = Real(final quantity="ThermodynamicTemperature", final unit="K"); + type Density = Real(final quantity="Density", final unit="kg/m3"); + end SI; + end Units; + + package Media + import Modelica.Units.SI; + + package Interfaces + package Types + type Temperature = SI.Temperature( + min=1, + max=1.e4, + nominal=300, + start=288.15); + type Density = SI.Density( + min=0, + max=1.e5, + nominal=1, + start=1); + + package IdealGas + record FluidConstants + Temperature criticalTemperature; + Density criticalDensity; + end FluidConstants; + end IdealGas; + end Types; + end Interfaces; + end Media; +end Modelica; + +model UsesFluidConstants + Modelica.Media.Interfaces.Types.IdealGas.FluidConstants constants; +equation + constants.criticalTemperature = 300; + constants.criticalDensity = 1; +end UsesFluidConstants; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("record type-alias fixture should parse"); + + session + .compile_model("UsesFluidConstants") + .expect("record fields using resolved type aliases must not raise ER023"); +} + +/// MLS §5.3 + §7.2: the right-hand side of a component modifier is resolved in +/// the lexical scope where the modifier occurs, not against the modified field. +#[test] +fn test_same_name_component_modifier_binding_uses_enclosing_parameter() { + let source = r#" +model TunedComponent + parameter Real p_start = p_start; + Real y; +equation + y = p_start; +end TunedComponent; + +model UsesOuterStart + parameter Real p_start = 3; + TunedComponent c(p_start = p_start); + Real y; +equation + y = c.y; +end UsesOuterStart; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("same-name modifier binding should resolve without ER007"); + + let result = session + .compile_model("UsesOuterStart") + .expect("compile should succeed"); + + let flat_code = render_flat_template_with_name( + &result.flat, + templates::builtin_template_source("flat-modelica", "flat_modelica.mo.jinja").unwrap(), + "UsesOuterStart", + ) + .expect("flat rendering should succeed"); + + assert!( + flat_code.contains("parameter Real c.p_start") && flat_code.contains("= p_start;"), + "component modifier should preserve enclosing parameter reference, got:\n{flat_code}" + ); +} + +/// MLS §7.2.4: a same-name modifier forwarding a runtime Boolean input is +/// still resolved in the enclosing lexical scope. It must not be folded from +/// the modified child's Boolean default before the flow equation is built. +#[test] +fn test_same_name_runtime_boolean_modifier_drives_nested_flow_connector() { + let source = r#" +connector CountPort + Real dummy; + flow Real count; +end CountPort; + +model FlowAdapter + input Boolean localActive; + CountPort port; +equation + port.count = if localActive then 1.0 else 0.0; +end FlowAdapter; + +model RuntimeFlowSource + output Boolean localActive; + FlowAdapter adapter(localActive = localActive); + CountPort port; +equation + localActive = time >= 0; + connect(adapter.port, port); +end RuntimeFlowSource; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("runtime Boolean forwarding fixture should parse"); + + let result = session + .compile_model("RuntimeFlowSource") + .expect("runtime Boolean forwarding fixture should compile"); + + let flat_code = render_flat_template_with_name( + &result.flat, + templates::builtin_template_source("flat-modelica", "flat_modelica.mo.jinja").unwrap(), + "RuntimeFlowSource", + ) + .expect("flat rendering should succeed"); + + assert!( + flat_code.contains("adapter.localActive(start = false) = localActive"), + "nested runtime Boolean modifier must retain its enclosing source, got:\n{flat_code}" + ); + assert!( + flat_code.contains("if adapter.localActive then 1") + || flat_code.contains("if adapter.localActive then 1.0"), + "nested flow source must remain runtime-dependent, got:\n{flat_code}" + ); +} + /// MLS §7.3: extends-clause package redeclarations that forward through a local /// alias (`redeclare package Medium = Medium`) must resolve using the active /// modification environment when instantiated through a component modifier. @@ -19,6 +184,7 @@ end PartialMedium; package RealMedium extends PartialMedium; + constant String extraPropertiesNames[:] = fill("", 0); redeclare model extends BaseProperties Real R_s; @@ -1019,3 +1185,2033 @@ end FluidLike; "pipe2.flowModel.x should be instantiated from inherited FlowModel local class" ); } + +/// MLS §7.3: a component with a single package redeclare override must expose +/// that package's constants in component scope, even when the component type is +/// fully qualified (e.g. `Modelica.Fluid.Sources.Boundary_pT` style). +#[test] +fn test_single_package_override_applies_to_fully_qualified_component_type_scope() { + let source = r#" +package PartialMedium + constant Integer nX = 1; + constant Integer nXi = 0; + + replaceable partial model BaseProperties + Real Xi[nXi]; + Real X[nX]; + equation + for i in 1:nX loop + X[i] = 1; + end for; + end BaseProperties; +end PartialMedium; + +package MixMedium + extends PartialMedium( + nX=2, + nXi=1); + + redeclare model extends BaseProperties + end BaseProperties; +end MixMedium; + +package Sources + model Boundary_pT + replaceable package Medium = PartialMedium; + Medium.BaseProperties medium; + equation + medium.Xi[1] = 0.5; + end Boundary_pT; +end Sources; + +model Top + Sources.Boundary_pT src(redeclare package Medium = MixMedium); +end Top; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("parse failed"); + + let result = session.compile_model("Top").expect("compile failed"); + + let xi_dims = result + .flat + .variables + .iter() + .find(|(name, _)| name.as_str() == "src.medium.Xi") + .map(|(_, var)| var.dims.clone()) + .expect("src.medium.Xi should exist"); + assert_eq!( + xi_dims, + vec![1], + "src.medium.Xi should use nXi from src's redeclared Medium package" + ); + + let x_dims = result + .flat + .variables + .iter() + .find(|(name, _)| name.as_str() == "src.medium.X") + .map(|(_, var)| var.dims.clone()) + .expect("src.medium.X should exist"); + assert_eq!( + x_dims, + vec![2], + "src.medium.X should use nX from src's redeclared Medium package" + ); +} + +/// MLS §7.3: local package aliases with class modifications (e.g. +/// `package Medium = PureMedium(AbsolutePressure(max=...))`) must preserve the +/// aliased package constants for member model dimensions (`Medium.nX/nXi`). +#[test] +fn test_local_package_alias_with_class_modification_preserves_member_model_dims() { + let source = r#" +package PartialMedium + type AbsolutePressure = Real; + constant Integer nS = 2; + final constant Integer nX = nS; + final constant Integer nXi = 0; + + replaceable model BaseProperties + AbsolutePressure p; + Real h; + Real d; + Real X[nX]; + input Real Xi[nXi]; + equation + d = p + h; + X = fill(1.0, nX); + Xi = fill(0.0, nXi); + end BaseProperties; +end PartialMedium; + +package PureMedium + extends PartialMedium(nS = 1); +end PureMedium; + +model UsesAliasWithModification + package Medium = PureMedium(AbsolutePressure(max = 1e6)); + Medium.BaseProperties medium; + Medium.BaseProperties medium2; +equation + medium.p = 1; + medium.h = 2; + medium2.p = 3; + medium2.h = 4; +end UsesAliasWithModification; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("parse failed"); + + let result = session + .compile_model("UsesAliasWithModification") + .expect("compile failed"); + + let medium_x_dims = result + .flat + .variables + .iter() + .find(|(name, _)| name.as_str() == "medium.X") + .map(|(_, var)| var.dims.clone()) + .expect("medium.X should exist"); + assert_eq!( + medium_x_dims, + vec![1], + "medium.X should use Medium.nX=1 from the aliased PureMedium package" + ); + + let medium2_x_dims = result + .flat + .variables + .iter() + .find(|(name, _)| name.as_str() == "medium2.X") + .map(|(_, var)| var.dims.clone()) + .expect("medium2.X should exist"); + assert_eq!( + medium2_x_dims, + vec![1], + "medium2.X should use Medium.nX=1 from the aliased PureMedium package" + ); + + let medium_xi_dims = result + .flat + .variables + .iter() + .find(|(name, _)| name.as_str() == "medium.Xi") + .map(|(_, var)| var.dims.clone()) + .expect("medium.Xi should exist"); + assert_eq!( + medium_xi_dims, + vec![0], + "medium.Xi should use Medium.nXi=0 from the aliased PureMedium package" + ); +} + +/// MLS §7.3 + §10.1: forwarding redeclares (`redeclare package Medium = Medium`) +/// inside nested components must evaluate dimensions against the enclosing +/// effective package override, not the local default package. +#[test] +fn test_forwarding_package_redeclare_applies_to_nested_stream_dimension() { + let source = r#" +package P + package MediumBase + constant Integer nC = 0; + end MediumBase; + + package MediumCO2 + extends MediumBase(nC = 1); + end MediumCO2; + + connector Port + replaceable package Medium = MediumBase; + Real p; + flow Real m_flow; + stream Real C_outflow[Medium.nC]; + end Port; + + model Source + replaceable package Medium = MediumBase; + Port port(redeclare package Medium = Medium); + equation + port.p = 0; + port.C_outflow = fill(0.0, Medium.nC); + end Source; + + model M + Source s(redeclare package Medium = MediumCO2); + end M; +end P; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("parse failed"); + + let result = session.compile_model("P.M").expect("compile failed"); + + let dims = result + .flat + .variables + .iter() + .find(|(name, _)| name.as_str() == "s.port.C_outflow") + .map(|(_, var)| var.dims.clone()) + .expect("s.port.C_outflow should exist"); + assert_eq!( + dims, + vec![1], + "s.port.C_outflow should use MediumCO2.nC through the forwarding redeclare" + ); +} + +/// MLS §7.3: chained forwarding redeclares inside component instances must keep +/// the active package override when a nested component instantiates +/// `Medium.BaseProperties`. +#[test] +fn test_component_redeclare_chain_instantiates_concrete_baseproperties() { + let source = r#" +package PartialMedium + replaceable partial model BaseProperties + Real p; + Real h; + Real d; + equation + d = p + h; + end BaseProperties; +end PartialMedium; + +package RealMedium + extends PartialMedium; + + redeclare replaceable model BaseProperties + Real p; + Real h; + Real d; + Real marker; + equation + d = p + h; + marker = d - p; + end BaseProperties; +end RealMedium; + +model Volume + replaceable package Medium = PartialMedium; + + model Balance + replaceable package Medium = PartialMedium; + Medium.BaseProperties medium; + equation + medium.p = 3; + medium.h = 2; + end Balance; + + Balance dynBal(redeclare package Medium = Medium); +end Volume; + +partial model MediumCarrier + replaceable package Medium = PartialMedium; +end MediumCarrier; + +partial model PortCarrier + replaceable package Medium = PartialMedium; +end PortCarrier; + +model FanBase + extends MediumCarrier; + extends PortCarrier; + Volume vol(redeclare package Medium = Medium); +end FanBase; + +model Top + package MediumAir = RealMedium(extraPropertiesNames = {"CO2"}); + + model AirHandler + replaceable package MediumAir = PartialMedium; + FanBase fan(redeclare package Medium = MediumAir); + end AirHandler; + + AirHandler ahu(redeclare package MediumAir = MediumAir); + Real y; +equation + y = ahu.fan.vol.dynBal.medium.marker; +end Top; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("parse failed"); + + let result = session.compile_model("Top").expect("compile failed"); + + let flat_var_names: Vec<_> = result + .flat + .variables + .keys() + .map(|k| k.to_string()) + .collect(); + assert!( + flat_var_names + .iter() + .any(|name| name == "ahu.fan.vol.dynBal.medium.marker"), + "nested Medium.BaseProperties should use RealMedium, vars={flat_var_names:?}" + ); +} + +const INHERITED_VOLUME_MEDIUM_BASEPROPERTIES_SOURCE: &str = r#" +package PartialMedium + replaceable partial model BaseProperties + Real p; + Real h; + end BaseProperties; +end PartialMedium; + +package CondensingMedium + extends PartialMedium; +end CondensingMedium; + +package RealMedium + extends PartialMedium; + + redeclare replaceable model BaseProperties + Real p; + Real h; + Real marker; + equation + marker = p + h; + end BaseProperties; +end RealMedium; + +package Interfaces + partial model LumpedVolumeDeclarations + replaceable package Medium = PartialMedium; + end LumpedVolumeDeclarations; + + model ConservationEquation + extends LumpedVolumeDeclarations; + Medium.BaseProperties medium; + equation + medium.p = 1; + medium.h = 2; + end ConservationEquation; +end Interfaces; + +partial model PartialMixingVolume + extends Interfaces.LumpedVolumeDeclarations; + Interfaces.ConservationEquation dynBal(redeclare final package Medium = Medium); +end PartialMixingVolume; + +model MixingVolume + extends PartialMixingVolume; +end MixingVolume; + +model MixingVolumeHeatPort + extends PartialMixingVolume; +end MixingVolumeHeatPort; + +model MixingVolumeHeatMoisturePort + extends PartialMixingVolume; +end MixingVolumeHeatMoisturePort; + +model FourPortHexBase + replaceable package Medium1 = PartialMedium; + replaceable package Medium2 = PartialMedium; + replaceable MixingVolumeHeatPort vol1 constrainedby + MixingVolumeHeatPort(redeclare final package Medium = Medium1); + replaceable MixingVolume vol2 constrainedby + MixingVolumeHeatPort(redeclare final package Medium = Medium2); +end FourPortHexBase; + +model BaseHex + extends FourPortHexBase; +end BaseHex; + +model UsesDefaultConstrainedbyVolume + extends FourPortHexBase( + redeclare package Medium1 = RealMedium, + redeclare package Medium2 = RealMedium); + Real y; +equation + y = vol1.dynBal.medium.marker; +end UsesDefaultConstrainedbyVolume; + +model LatentHex + extends BaseHex( + redeclare final MixingVolumeHeatPort vol1, + redeclare final MixingVolumeHeatMoisturePort vol2); +end LatentHex; + +partial model PartialFourPort + replaceable package Medium1 = PartialMedium; + replaceable package Medium2 = PartialMedium; +end PartialFourPort; + +model DryCoil + extends PartialFourPort; + replaceable model HexElement = BaseHex; + HexElement ele[1]( + redeclare each package Medium1 = Medium1, + redeclare each package Medium2 = Medium2); +end DryCoil; + +model WetCoil + extends DryCoil( + redeclare replaceable package Medium2 = CondensingMedium, + redeclare model HexElement = LatentHex); +end WetCoil; + +model CoilWrapper + replaceable package MediumAir = PartialMedium; + replaceable package MediumWat = PartialMedium; + WetCoil cooCoi( + redeclare package Medium1 = MediumWat, + redeclare package Medium2 = MediumAir); +end CoilWrapper; + +partial model WatCoil + replaceable package MediumAir = PartialMedium; + replaceable package MediumWat = PartialMedium; +end WatCoil; + +model CoolingCoil + extends WatCoil; + CoilWrapper coi( + redeclare package MediumAir = MediumAir, + redeclare package MediumWat = MediumWat); +end CoolingCoil; + +model Top + package MediumWater = RealMedium; + package MediumAir = RealMedium; + CoolingCoil cooCoi( + redeclare package MediumAir = MediumAir, + redeclare package MediumWat = MediumWater); + Real y; +equation + y = cooCoi.coi.cooCoi.ele[1].vol1.dynBal.medium.marker; +end Top; +"#; + +/// MLS §7.3: inherited replaceable model arrays must keep active package +/// redeclarations when nested volumes instantiate `Medium.BaseProperties`. +#[test] +fn test_inherited_replaceable_model_array_keeps_active_medium_for_baseproperties() { + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", INHERITED_VOLUME_MEDIUM_BASEPROPERTIES_SOURCE) + .expect("parse failed"); + + let direct_result = session + .compile_model("UsesDefaultConstrainedbyVolume") + .expect("default constrainedby volume compile failed"); + + let direct_var_names: Vec<_> = direct_result + .flat + .variables + .keys() + .map(|k| k.to_string()) + .collect(); + assert!( + direct_var_names + .iter() + .any(|name| name == "vol1.dynBal.medium.marker"), + "default replaceable volume should use constrainedby RealMedium; vars={direct_var_names:?}" + ); + + let result = session.compile_model("Top").expect("compile failed"); + + let flat_var_names: Vec<_> = result + .flat + .variables + .keys() + .map(|k| k.to_string()) + .collect(); + assert!( + flat_var_names + .iter() + .any(|name| name == "cooCoi.coi.cooCoi.ele[1].vol1.dynBal.medium.marker"), + "array element Medium.BaseProperties should use active RealMedium; vars={flat_var_names:?}" + ); +} + +/// MLS §7.3: package aliases in sibling models must not leak into the active +/// model's alias resolution. Compiling `Examples.A` should use `A.Medium`. +#[test] +fn test_sibling_model_package_alias_does_not_pollute_active_model_dims() { + let source = r#" +package PartialMedium + constant Integer nX = 1; + replaceable model BaseProperties + Real p; + Real h; + Real d; + Real X[nX]; + equation + d = p + h; + X = fill(1.0, nX); + end BaseProperties; +end PartialMedium; + +package MediumOne + extends PartialMedium(nX = 1); +end MediumOne; + +package MediumTwo + extends PartialMedium(nX = 2); +end MediumTwo; + +package Examples + model A + package Medium = MediumOne; + Medium.BaseProperties medium; + equation + medium.p = 1; + medium.h = 2; + end A; + + model B + package Medium = MediumTwo; + Medium.BaseProperties medium; + equation + medium.p = 3; + medium.h = 4; + end B; +end Examples; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("parse failed"); + + let result = session.compile_model("Examples.A").expect("compile failed"); + + let x_dims = result + .flat + .variables + .iter() + .find(|(name, _)| name.as_str() == "medium.X") + .map(|(_, var)| var.dims.clone()) + .expect("medium.X should exist"); + assert_eq!( + x_dims, + vec![1], + "Examples.A.medium.X should use A.Medium (MediumOne.nX=1), not sibling alias values" + ); +} + +/// MLS §7.3: inherited package-constant chains (`PartialMedium -> +/// PartialPureSubstance -> PartialSimpleMedium`) must preserve `nS/nX/nXi` +/// when used through a local `package Medium = ...` alias. +#[test] +fn test_local_medium_alias_preserves_partial_pure_substance_constants() { + let source = r#" +package Interfaces + partial package PartialMedium + type SpecificEnthalpy = Real; + constant Boolean reducedX = true; + constant Boolean fixedX = false; + constant String substanceNames[:] = {"single"}; + final constant Integer nS = size(substanceNames, 1); + constant Integer nX = nS; + constant Integer nXi = if fixedX then 0 else if reducedX then nS - 1 else nS; + constant Real reference_X[nX] = fill(1.0 / nX, nX); + + replaceable partial model BaseProperties + Real p; + SpecificEnthalpy h; + Real d; + Real X[nX]; + input Real Xi[nXi]; + equation + X = reference_X; + Xi = X[1:nXi]; + d = p + h; + end BaseProperties; + end PartialMedium; + + partial package PartialPureSubstance + extends PartialMedium(final reducedX = true, final fixedX = true); + end PartialPureSubstance; + + partial package PartialSimpleMedium + extends PartialPureSubstance; + end PartialSimpleMedium; +end Interfaces; + +package WaterLike + extends Interfaces.PartialSimpleMedium; + redeclare model extends BaseProperties + end BaseProperties; +end WaterLike; + +model UsesLocalMediumAlias + package Medium = WaterLike(SpecificEnthalpy(max = 1e6)); + Medium.BaseProperties medium; +equation + medium.p = 1; + medium.h = 2; +end UsesLocalMediumAlias; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("parse failed"); + + let result = session + .compile_model("UsesLocalMediumAlias") + .expect("compile failed"); + + let x_dims = result + .flat + .variables + .iter() + .find(|(name, _)| name.as_str() == "medium.X") + .map(|(_, var)| var.dims.clone()) + .expect("medium.X should exist"); + assert_eq!( + x_dims, + vec![1], + "medium.X should use nX=1 for PartialPureSubstance aliases" + ); + + let xi_dims = result + .flat + .variables + .iter() + .find(|(name, _)| name.as_str() == "medium.Xi") + .map(|(_, var)| var.dims.clone()) + .expect("medium.Xi should exist"); + assert_eq!( + xi_dims, + vec![0], + "medium.Xi should use nXi=0 for PartialPureSubstance aliases" + ); +} + +/// MLS §7.3: extends-modification redeclarations inside a package (e.g. +/// `extends PartialMedium(redeclare record ThermodynamicState=...)`) must be +/// visible when resolving dotted member types (`Medium.ThermodynamicState`). +#[test] +fn test_package_extends_redeclare_record_alias_applies_to_dotted_member_type() { + let source = r#" +package Common + record BaseProps_Tpoly + Real T; + Real p; + end BaseProps_Tpoly; +end Common; + +package Interfaces + partial package PartialMedium + replaceable record ThermodynamicState + Real x; + end ThermodynamicState; + + replaceable function setState_pTX + input Real p; + input Real T; + output ThermodynamicState state; + algorithm + state := ThermodynamicState(); + end setState_pTX; + end PartialMedium; +end Interfaces; + +package TableBased + extends Interfaces.PartialMedium( + redeclare record ThermodynamicState = Common.BaseProps_Tpoly + ); + + redeclare function setState_pTX + input Real p; + input Real T; + output ThermodynamicState state; + algorithm + state := Common.BaseProps_Tpoly(T=T, p=p); + end setState_pTX; +end TableBased; + +model UsesTableBasedState + package Medium = TableBased; + Medium.ThermodynamicState state = Medium.setState_pTX(1, 2); +end UsesTableBasedState; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("parse failed"); + + let result = session + .compile_model("UsesTableBasedState") + .expect("compile failed"); + + let flat_var_names: Vec<_> = result + .flat + .variables + .keys() + .map(|k| k.to_string()) + .collect(); + + assert!( + flat_var_names.iter().any(|n| n == "state.T"), + "state.T should come from redeclared ThermodynamicState; vars={flat_var_names:?}" + ); + assert!( + flat_var_names.iter().any(|n| n == "state.p"), + "state.p should come from redeclared ThermodynamicState; vars={flat_var_names:?}" + ); + assert!( + !flat_var_names.iter().any(|n| n == "state.x"), + "base ThermodynamicState field should be replaced by redeclare; vars={flat_var_names:?}" + ); +} + +/// MLS §5.3 + §7.3: A `redeclare record extends T` inside a package whose base +/// package is reached by a fully-qualified `extends` clause must find the +/// inherited replaceable type. This is the compact form of the Buildings +/// `Media.Air` pattern used by Kelvin's BOPTEST parity artifact export. +#[test] +fn test_fully_qualified_package_extends_redeclare_record_extends_inherited_type() { + let source = r#" +package Modelica + package Icons + package Package + end Package; + end Icons; + + package Media + package Interfaces + partial package PartialMedium + replaceable record ThermodynamicState + end ThermodynamicState; + + replaceable model BaseProperties + input Real p; + input Real T; + ThermodynamicState state; + end BaseProperties; + end PartialMedium; + + partial package PartialCondensingGases + extends Modelica.Media.Interfaces.PartialMedium; + end PartialCondensingGases; + end Interfaces; + end Media; +end Modelica; + +package Buildings + package Media + package Air + extends Modelica.Media.Interfaces.PartialCondensingGases; + extends Modelica.Icons.Package; + + redeclare record extends ThermodynamicState + Real p; + Real T; + end ThermodynamicState; + + redeclare model extends BaseProperties + equation + state.p = p; + state.T = T; + end BaseProperties; + end Air; + end Media; +end Buildings; + +model UsesAir + package Medium = Buildings.Media.Air; + Medium.BaseProperties medium; +equation + medium.p = 101325; + medium.T = 295.15; +end UsesAir; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("fully-qualified redeclare fixture should parse"); + + let result = session + .compile_model("UsesAir") + .expect("fully-qualified package extends must expose inherited redeclare target"); + + let flat_var_names: Vec<_> = result + .flat + .variables + .keys() + .map(|k| k.to_string()) + .collect(); + assert!( + flat_var_names.iter().any(|n| n == "medium.state.p"), + "medium.state.p should come from the Air ThermodynamicState redeclare; vars={flat_var_names:?}" + ); + assert!( + flat_var_names.iter().any(|n| n == "medium.state.T"), + "medium.state.T should come from the Air ThermodynamicState redeclare; vars={flat_var_names:?}" + ); +} + +/// MLS §5.3 + §7.3: A redeclared class body in the derived package uses that +/// package's active inherited/redeclared type members for short type names. +/// Buildings `Media.Air` declares `redeclare replaceable model BaseProperties` +/// with `ThermodynamicState state` in this form. +#[test] +fn test_redeclared_model_body_uses_sibling_redeclared_record_type() { + let source = r#" +package BaseMedium + replaceable record ThermodynamicState + end ThermodynamicState; + + replaceable model BaseProperties + input Real p; + input Real T; + ThermodynamicState state; + end BaseProperties; +end BaseMedium; + +package Air + extends BaseMedium; + + redeclare record extends ThermodynamicState + Real p; + Real T; + end ThermodynamicState; + + redeclare replaceable model BaseProperties + input Real p; + input Real T; + ThermodynamicState state; + equation + state.p = p; + state.T = T; + end BaseProperties; +end Air; + +model UsesAir + Air.BaseProperties medium; +equation + medium.p = 101325; + medium.T = 295.15; +end UsesAir; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("sibling redeclare fixture should parse"); + + let result = session + .compile_model("UsesAir") + .expect("redeclared model body must use sibling redeclared record type"); + + let flat_var_names: Vec<_> = result + .flat + .variables + .keys() + .map(|k| k.to_string()) + .collect(); + assert!( + flat_var_names.iter().any(|n| n == "medium.state.p"), + "medium.state.p should come from sibling redeclared ThermodynamicState; vars={flat_var_names:?}" + ); + assert!( + flat_var_names.iter().any(|n| n == "medium.state.T"), + "medium.state.T should come from sibling redeclared ThermodynamicState; vars={flat_var_names:?}" + ); +} + +/// MLS §5.3 + §7.3: inherited replaceable record/function members remain +/// visible through intermediate package extends. Buildings media packages use +/// this for `redeclare record extends ThermodynamicState` and +/// `redeclare function extends dynamicViscosity` through +/// `PartialCondensingGases -> PartialMixtureMedium`. +#[test] +fn test_redeclare_extends_finds_replaceable_members_through_intermediate_package_extends() { + let source = r#" +package Modelica + package Media + package Interfaces + partial package PartialMedium + replaceable record ThermodynamicState + end ThermodynamicState; + + replaceable partial function dynamicViscosity + input ThermodynamicState state; + output Real eta; + end dynamicViscosity; + + replaceable partial function temperature + input ThermodynamicState state; + output Real T; + end temperature; + + replaceable partial function specificEnthalpy + input ThermodynamicState state; + output Real h; + end specificEnthalpy; + end PartialMedium; + + partial package PartialMixtureMedium + extends Modelica.Media.Interfaces.PartialMedium; + end PartialMixtureMedium; + + partial package PartialCondensingGases + extends Modelica.Media.Interfaces.PartialMixtureMedium; + end PartialCondensingGases; + end Interfaces; + end Media; +end Modelica; + +package Buildings + package Media + package Air + extends Modelica.Media.Interfaces.PartialCondensingGases; + + redeclare record extends ThermodynamicState + Real p; + Real T; + end ThermodynamicState; + + redeclare function extends dynamicViscosity + algorithm + eta := state.p + state.T; + end dynamicViscosity; + + redeclare function extends temperature + algorithm + T := state.T; + end temperature; + + redeclare replaceable function extends specificEnthalpy + algorithm + h := state.p + state.T; + end specificEnthalpy; + end Air; + end Media; +end Buildings; + +model UsesAirFunctions + Buildings.Media.Air.ThermodynamicState state; + Real eta; + Real T; + Real h; +equation + state.p = 101325; + state.T = 295.15; + eta = Buildings.Media.Air.dynamicViscosity(state); + T = Buildings.Media.Air.temperature(state); + h = Buildings.Media.Air.specificEnthalpy(state); +end UsesAirFunctions; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("intermediate package extends fixture should parse"); + + let result = session + .compile_model("UsesAirFunctions") + .expect("redeclare extends must find inherited members through intermediate packages"); + + let flat_var_names: Vec<_> = result + .flat + .variables + .keys() + .map(|k| k.to_string()) + .collect(); + assert!( + flat_var_names.iter().any(|n| n == "state.p"), + "state.p should come from the Air ThermodynamicState redeclare; vars={flat_var_names:?}" + ); + assert!( + flat_var_names.iter().any(|n| n == "state.T"), + "state.T should come from the Air ThermodynamicState redeclare; vars={flat_var_names:?}" + ); +} + +/// Buildings.Media.Air redeclares `ThermodynamicState` with an empty +/// `record extends` body. The inherited record fields and inherited functions +/// must remain visible for later `redeclare function extends ...` clauses. +#[test] +fn test_empty_redeclare_record_extends_preserves_inherited_function_members() { + let source = r#" +package Modelica + package Media + package Interfaces + partial package PartialMedium + replaceable record ThermodynamicState + Real p; + Real T; + end ThermodynamicState; + + replaceable partial function temperature + input ThermodynamicState state; + output Real T; + end temperature; + end PartialMedium; + + partial package PartialMixtureMedium + extends Modelica.Media.Interfaces.PartialMedium; + end PartialMixtureMedium; + + partial package PartialCondensingGases + extends Modelica.Media.Interfaces.PartialMixtureMedium; + end PartialCondensingGases; + end Interfaces; + end Media; +end Modelica; + +package Buildings + package Media + package Air + extends Modelica.Media.Interfaces.PartialCondensingGases; + + redeclare record extends ThermodynamicState + end ThermodynamicState; + + redeclare function extends temperature + algorithm + T := state.T; + end temperature; + end Air; + end Media; +end Buildings; + +model UsesAirTemperature + Buildings.Media.Air.ThermodynamicState state; + Real T; +equation + state.p = 101325; + state.T = 295.15; + T = Buildings.Media.Air.temperature(state); +end UsesAirTemperature; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("empty record extends fixture should parse"); + + let result = session + .compile_model("UsesAirTemperature") + .expect("empty redeclare record extends must preserve inherited function members"); + + let flat_var_names: Vec<_> = result + .flat + .variables + .keys() + .map(|k| k.to_string()) + .collect(); + assert!( + flat_var_names.iter().any(|n| n == "state.T"), + "state.T should remain inherited through empty record extends; vars={flat_var_names:?}" + ); +} + +/// A class modification on the package extends must not hide inherited +/// replaceable function members from later `redeclare function extends ...` +/// lookup. Buildings.Media.Air uses a heavily modified +/// `extends PartialCondensingGases(...)` clause before its function redeclares. +#[test] +fn test_class_modified_package_extends_preserves_inherited_function_members() { + let source = r#" +package BaseMedium + constant Boolean flag = false; + + replaceable record ThermodynamicState + Real T; + end ThermodynamicState; + + replaceable partial function temperature + input ThermodynamicState state; + output Real T; + end temperature; +end BaseMedium; + +package MixtureMedium + extends BaseMedium; +end MixtureMedium; + +package CondensingGases + extends MixtureMedium; +end CondensingGases; + +package Air + extends CondensingGases(flag = true); + + redeclare record extends ThermodynamicState + end ThermodynamicState; + + redeclare function extends temperature + algorithm + T := state.T; + end temperature; +end Air; + +model UsesAirTemperature + Air.ThermodynamicState state; + Real T; +equation + state.T = 295.15; + T = Air.temperature(state); +end UsesAirTemperature; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("class-modified function redeclare fixture should parse"); + + session + .compile_model("UsesAirTemperature") + .expect("class-modified package extends must preserve inherited function members"); +} + +/// Buildings library layout declares the concrete medium in a separate file as +/// `within Buildings.Media; package Air ...`. Inherited `redeclare function +/// extends ...` lookup must use the fully-qualified enclosing package name, +/// not the local file-root name. +#[test] +fn test_within_package_redeclare_extends_finds_inherited_function_members() { + let modelica_source = r#" +package Modelica + package Media + package Interfaces + partial package PartialMedium + replaceable record ThermodynamicState + Real T; + end ThermodynamicState; + + replaceable partial function temperature + input ThermodynamicState state; + output Real T; + end temperature; + end PartialMedium; + + partial package PartialMixtureMedium + extends Modelica.Media.Interfaces.PartialMedium; + end PartialMixtureMedium; + + partial package PartialCondensingGases + extends Modelica.Media.Interfaces.PartialMixtureMedium; + end PartialCondensingGases; + end Interfaces; + end Media; +end Modelica; + +package Buildings + package Media + end Media; +end Buildings; +"#; + + let air_source = r#" +within Buildings.Media; +package Air + extends Modelica.Media.Interfaces.PartialCondensingGases; + + redeclare record extends ThermodynamicState + end ThermodynamicState; + + redeclare function extends temperature + algorithm + T := state.T; + end temperature; +end Air; +"#; + + let probe_source = r#" +model UsesWithinAirTemperature + Buildings.Media.Air.ThermodynamicState state; + Real T; +equation + state.T = 295.15; + T = Buildings.Media.Air.temperature(state); +end UsesWithinAirTemperature; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("modelica.mo", modelica_source) + .expect("modelica fixture should parse"); + session + .add_document("Air.mo", air_source) + .expect("within Air fixture should parse"); + session + .add_document("probe.mo", probe_source) + .expect("probe fixture should parse"); + + session + .compile_model("UsesWithinAirTemperature") + .expect("within package redeclare extends must find inherited function members"); +} + +/// Plant performance records use a colon-sized abscissa and sibling arrays +/// sized by `size(V_flow, 1)`, with the concrete abscissa supplied through a +/// record modification. Buildings mover and cooling-tower curves use this +/// pattern for pump pressure/efficiency and fan-power data. +#[test] +fn test_record_modification_array_dims_feed_sibling_size_dimensions() { + let source = r#" +record FlowParameters + parameter Real V_flow[:] = {0.0, 1.0}; + parameter Real dp[size(V_flow, 1)] = {1.0, 0.0}; +end FlowParameters; + +record PerformanceData + parameter FlowParameters pressure( + V_flow = {0.0, 1.0, 2.0}, + dp = {2.0, 1.0, 0.0}); +end PerformanceData; + +model Pump + parameter PerformanceData per; + parameter Real curveDer[size(per.pressure.V_flow, 1)] = {0.0, 0.0, 0.0}; +end Pump; + +model UsesPumpRecordCurves + Pump pump; +end UsesPumpRecordCurves; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("record curve fixture should parse"); + + let result = session + .compile_model("UsesPumpRecordCurves") + .expect("record modifications should drive sibling size() dimensions"); + + let pump_der_dims = result + .flat + .variables + .iter() + .find_map(|(name, var)| (name.to_string() == "pump.curveDer").then_some(var.dims.clone())) + .expect("pump.curveDer should be present"); + assert_eq!( + pump_der_dims, + vec![3], + "pump.der should use modified per.pressure.V_flow length" + ); + + let pressure_dp_dims = result + .flat + .variables + .iter() + .find_map(|(name, var)| { + (name.to_string() == "pump.per.pressure.dp").then_some(var.dims.clone()) + }) + .expect("pump.per.pressure.dp should be present"); + assert_eq!( + pressure_dp_dims, + vec![3], + "record field dp should use sibling V_flow length from record modification" + ); +} + +/// Array component modifications distribute over component elements. For a +/// component array `pump[n](VolFloCur=VolFloCur)` where the modifier expression +/// is `VolFloCur[n,:]`, each `pump[i].VolFloCur[:]` receives row `i`; downstream +/// record curves sized by `size(V_flow, 1)` must see that row length. +#[test] +fn test_array_component_row_modification_feeds_record_size_dimensions() { + let source = r#" +record FlowParameters + parameter Real V_flow[:] = {0.0, 1.0}; + parameter Real dp[size(V_flow, 1)] = {1.0, 0.0}; +end FlowParameters; + +record PerformanceData + parameter FlowParameters pressure( + V_flow = {0.0, 1.0}, + dp = {1.0, 0.0}); +end PerformanceData; + +model WithoutMotor + parameter Real VolFloCur[:]; + parameter Real PreCur[:]; + parameter PerformanceData per( + pressure(V_flow = VolFloCur, dp = PreCur)); +end WithoutMotor; + +model PumpSystem + parameter Integer n = 2; + parameter Real VolFloCur[n, :] = {{0.0, 1.0, 2.0} for i in linspace(1, n, n)}; + parameter Real PreCur[n, :] = {{2.0, 1.0, 0.0} for i in linspace(1, n, n)}; + WithoutMotor pump[n]( + VolFloCur = VolFloCur, + PreCur = PreCur); +end PumpSystem; + +model UsesPumpSystemRecordCurves + PumpSystem sys; +end UsesPumpSystemRecordCurves; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("array component row fixture should parse"); + + let result = session + .compile_model("UsesPumpSystemRecordCurves") + .expect("array component row modifiers should drive record curve dimensions"); + + let dims = result + .flat + .variables + .iter() + .find_map(|(name, var)| { + (name.to_string() == "sys.pump[1].per.pressure.dp").then_some(var.dims.clone()) + }) + .expect("sys.pump[1].per.pressure.dp should be present"); + assert_eq!( + dims, + vec![3], + "array component row modifier should provide one VolFloCur row" + ); +} + +/// Buildings air-side heat exchangers redeclare `Medium2 = MediumAir`, where +/// `MediumAir` is a package alias with class modifications. The active medium's +/// inherited `ThermodynamicState.X[nX]` must evaluate `nX` in the concrete +/// medium package scope, not in the component instance scope. +#[test] +fn test_redeclared_modified_medium_alias_drives_thermodynamic_state_nx() { + let source = r#" +package PartialMedium + constant String substanceNames[:] = {"base"}; + final constant Integer nS = size(substanceNames, 1); + final constant Integer nX = nS; + constant String extraPropertiesNames[:] = fill("", 0); + + replaceable record ThermodynamicState + Real X[nX]; + end ThermodynamicState; +end PartialMedium; + +package MixtureMedium + extends PartialMedium(substanceNames = {"water", "air"}); + + redeclare replaceable record extends ThermodynamicState + Real p; + Real T; + end ThermodynamicState; +end MixtureMedium; + +package Air + extends MixtureMedium(extraPropertiesNames = {"CO2"}); + + redeclare record extends ThermodynamicState + end ThermodynamicState; +end Air; + +model Coil + replaceable package Medium2 = PartialMedium; + Medium2.ThermodynamicState state_a2_inflow; +end Coil; + +model UsesModifiedAirMedium + package MediumAir = Air(extraPropertiesNames = {"CO2"}); + Coil coil(redeclare package Medium2 = MediumAir); +equation + coil.state_a2_inflow.X[1] = 0.7; + coil.state_a2_inflow.X[2] = 0.3; + coil.state_a2_inflow.p = 101325; + coil.state_a2_inflow.T = 295.15; +end UsesModifiedAirMedium; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("modified medium alias fixture should parse"); + + let result = session + .compile_model("UsesModifiedAirMedium") + .expect("modified medium alias should provide ThermodynamicState.nX"); + + let dims = result + .flat + .variables + .iter() + .find_map(|(name, var)| { + (name.to_string() == "coil.state_a2_inflow.X").then_some(var.dims.clone()) + }) + .expect("coil.state_a2_inflow.X should be present"); + assert_eq!(dims, vec![2], "state X should use Air.nX=2"); +} + +/// MLS §7.2 + §7.3: a package redeclare with class modifications must still +/// replace the package used by inherited component types. Buildings media +/// examples use `extends PartialProperties(redeclare package Medium = +/// Buildings.Media.Steam(p_default=...))` before accessing +/// `basPro.state.p/T`. +#[test] +fn test_inherited_component_type_uses_redeclared_package_with_class_modification() { + let source = r#" +package BaseMedium + constant Real p_default = 1; + + replaceable record ThermodynamicState + Real x; + end ThermodynamicState; + + replaceable model BaseProperties + input Real p; + input Real T; + ThermodynamicState state; + end BaseProperties; +end BaseMedium; + +package Steam + extends BaseMedium; + constant Real p_default = 2; + + redeclare record ThermodynamicState + Real p; + Real T; + end ThermodynamicState; + + redeclare replaceable model extends BaseProperties + equation + state.p = p; + state.T = T; + end BaseProperties; +end Steam; + +partial model PartialProperties + replaceable package Medium = BaseMedium; + parameter Real p = Medium.p_default; + Medium.BaseProperties basPro; +end PartialProperties; + +model SteamProperties + extends PartialProperties( + redeclare package Medium = Steam(p_default = 200000), + p = 200000); +equation + basPro.p = p; + basPro.T = 295.15; +end SteamProperties; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("class-modified package redeclare fixture should parse"); + + let result = session + .compile_model("SteamProperties") + .expect("class-modified package redeclare must control inherited component type"); + + let flat_var_names: Vec<_> = result + .flat + .variables + .keys() + .map(|k| k.to_string()) + .collect(); + assert!( + flat_var_names.iter().any(|n| n == "basPro.state.p"), + "basPro.state.p should come from redeclared package Medium; vars={flat_var_names:?}" + ); + assert!( + flat_var_names.iter().any(|n| n == "basPro.state.T"), + "basPro.state.T should come from redeclared package Medium; vars={flat_var_names:?}" + ); +} + +/// MLS §5.3 + §7.3: unqualified constants from the enclosing package of a +/// redeclared member model remain visible after instantiation. Buildings +/// `Media.Steam.BaseProperties` uses this as `MM = steam.MM` where `steam` is a +/// package-level constant record. +#[test] +fn test_redeclared_member_model_reads_enclosing_package_record_constant_field() { + let source = r#" +package BaseMedium + replaceable model BaseProperties + Real MM; + end BaseProperties; +end BaseMedium; + +package Data + record Species + Real MM; + Real R_s; + end Species; + + constant Species H2O(MM=18.01528, R_s=461.5); +end Data; + +package Steam + extends BaseMedium; + +protected + record GasProperties + Real MM; + Real R; + end GasProperties; + + constant GasProperties steam(MM=Data.H2O.MM, R=Data.H2O.R_s); + +public + redeclare replaceable model extends BaseProperties + equation + MM = steam.MM; + end BaseProperties; +end Steam; + +partial model PartialProperties + replaceable package Medium = BaseMedium; + Medium.BaseProperties basPro; +end PartialProperties; + +model SteamProperties + extends PartialProperties(redeclare package Medium = Steam); +end SteamProperties; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("record constant field fixture should parse"); + + let phase_result = session + .compile_model_phases("SteamProperties") + .expect("redeclared member model should compile through DAE lowering"); + assert!( + matches!(phase_result, PhaseResult::Success(_)), + "record constant field should not remain as unresolved basPro.steam.MM; got {phase_result:?}" + ); +} + +/// MLS §5.3: sibling functions can read constant fields from a lexical nested +/// record class namespace. MSL IF97 uses `data.RH2O` inside functions nested +/// under `BaseIF97.Basic`, where `data` is a record class in `BaseIF97`. +#[test] +fn test_function_body_reads_lexical_nested_record_class_constant_field() { + let source = r#" +package IF97 + record data + constant Real RH2O = 461.526; + end data; + + package Basic + function g2 + input Real p; + output Real r; + algorithm + r := data.RH2O * p; + end g2; + end Basic; +end IF97; + +model UsesIF97 + Real y; +equation + y = IF97.Basic.g2(2.0); +end UsesIF97; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("lexical record class constant fixture should parse"); + + let phase_result = session + .compile_model_phases("UsesIF97") + .expect("lexical record class constant fixture should compile"); + assert!( + matches!(phase_result, PhaseResult::Success(_)), + "function body should qualify and substitute data.RH2O; got {phase_result:?}" + ); +} + +/// MLS §5.3: lexical nested record class constants must also be substituted +/// when they appear in function-body branch conditions. MSL IF97 uses +/// `if T < data.TLIMIT1 then ...` in nested utility functions. +#[test] +fn test_function_body_if_condition_reads_lexical_nested_record_class_constant_field() { + let source = r#" +package IF97 + record data + constant Real TLIMIT1 = 623.15; + end data; + + package Basic + function region + input Real T; + output Real r; + algorithm + if T < data.TLIMIT1 then + r := 1.0; + else + r := 2.0; + end if; + end region; + end Basic; +end IF97; + +model UsesIF97 + Real y; +equation + y = IF97.Basic.region(300.0); +end UsesIF97; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("lexical record class branch constant fixture should parse"); + + let phase_result = session + .compile_model_phases("UsesIF97") + .expect("lexical record class branch constant fixture should compile"); + assert!( + matches!(phase_result, PhaseResult::Success(_)), + "function branch condition should substitute data.TLIMIT1; got {phase_result:?}" + ); +} + +/// MLS §5.3: lexical class aliases must climb from a nested package function +/// to an ancestor package sibling record. This mirrors +/// `IF97_Utilities.BaseIF97.Basic.region_*` reading `BaseIF97.data.TLIMIT1`. +#[test] +fn test_nested_package_function_reads_ancestor_sibling_record_constant_field() { + let source = r#" +package IF97_Utilities + package BaseIF97 + record data + constant Real TLIMIT1 = 623.15; + end data; + + package Basic + function region + input Real T; + output Real r; + algorithm + if T < data.TLIMIT1 then + r := 1.0; + else + r := 2.0; + end if; + end region; + end Basic; + end BaseIF97; +end IF97_Utilities; + +model UsesIF97 + Real y; +equation + y = IF97_Utilities.BaseIF97.Basic.region(300.0); +end UsesIF97; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("ancestor sibling record constant fixture should parse"); + + let phase_result = session + .compile_model_phases("UsesIF97") + .expect("ancestor sibling record constant fixture should compile"); + assert!( + matches!(phase_result, PhaseResult::Success(_)), + "nested package function should substitute ancestor sibling data.TLIMIT1; got {phase_result:?}" + ); +} + +/// MLS §5.3: constants in assertion messages and builtin String arguments must +/// use the same lexical class alias substitution as numeric branch conditions. +/// MSL IF97 emits messages like `String(data.TLIMIT1)`. +#[test] +fn test_assert_message_string_reads_ancestor_sibling_record_constant_field() { + let source = r#" +package IF97_Utilities + package BaseIF97 + record data + constant Real TLIMIT1 = 623.15; + end data; + + package Basic + function bounded + input Real T; + output Real r; + algorithm + assert(T >= data.TLIMIT1, "T < " + String(data.TLIMIT1)); + r := T; + end bounded; + end Basic; + end BaseIF97; +end IF97_Utilities; + +model UsesIF97 + Real y; +equation + y = IF97_Utilities.BaseIF97.Basic.bounded(700.0); +end UsesIF97; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("assert message record constant fixture should parse"); + + let phase_result = session + .compile_model_phases("UsesIF97") + .expect("assert message record constant fixture should compile"); + assert!( + matches!(phase_result, PhaseResult::Success(_)), + "assert message should substitute data.TLIMIT1 in String(); got {phase_result:?}" + ); +} + +/// MLS §5.3: dependency functions collected from another function call keep +/// the lexical class aliases of their own package scope. MSL SteamProperties +/// reaches IF97 helpers through a chain of function dependencies. +#[test] +fn test_dependency_function_reads_ancestor_sibling_record_constant_field() { + let source = r#" +package IF97_Utilities + package BaseIF97 + record data + constant Real TLIMIT1 = 623.15; + end data; + + package Inverses + function bounded + input Real T; + output Real r; + algorithm + assert(T >= data.TLIMIT1, "T < " + String(data.TLIMIT1)); + r := T - data.TLIMIT1; + end bounded; + end Inverses; + + package Regions + function region + input Real T; + output Real r; + algorithm + r := Inverses.bounded(T); + end region; + end Regions; + end BaseIF97; +end IF97_Utilities; + +model UsesIF97 + Real y; +equation + y = IF97_Utilities.BaseIF97.Regions.region(700.0); +end UsesIF97; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("dependency function record constant fixture should parse"); + + let phase_result = session + .compile_model_phases("UsesIF97") + .expect("dependency function record constant fixture should compile"); + assert!( + matches!(phase_result, PhaseResult::Success(_)), + "dependency function should substitute ancestor sibling data.TLIMIT1; got {phase_result:?}" + ); +} + +/// MLS §5.3 + §7.3: function bodies can read fields from a fully-qualified +/// package-level constant record. Buildings.Media.Steam uses +/// `Buildings.Media.Steam.steam.MM` in a function body; the referenced class +/// scope is `Buildings.Media.Steam`, not the constant record instance +/// `Buildings.Media.Steam.steam`. +#[test] +fn test_function_body_reads_fully_qualified_package_record_constant_field() { + let source = r#" +package Buildings + package Media + package Steam + record GasProperties + Real MM; + Real R; + end GasProperties; + + constant GasProperties steam(MM=18.01528, R=461.5); + + function molarMass + output Real MM; + algorithm + MM := Buildings.Media.Steam.steam.MM; + end molarMass; + end Steam; + end Media; +end Buildings; + +model UsesSteam + Real y; +equation + y = Buildings.Media.Steam.molarMass(); +end UsesSteam; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("fully-qualified package record constant fixture should parse"); + + let phase_result = session + .compile_model_phases("UsesSteam") + .expect("fully-qualified package record constant fixture should compile"); + assert!( + matches!(phase_result, PhaseResult::Success(_)), + "function body should substitute Buildings.Media.Steam.steam.MM; got {phase_result:?}" + ); +} + +/// MLS §5.3 + §7.3: model-level `extends(... redeclare package Medium=...)` +/// must override unrelated import aliases, and short member types inside +/// `Medium.BaseProperties` (e.g. `ThermodynamicState state`) must resolve to +/// the redeclared package record. +#[test] +fn test_model_redeclare_package_controls_member_dims_and_short_record_type() { + let source = r#" +package Common + record BaseProps_Tpoly + Real T; + Real p; + end BaseProps_Tpoly; +end Common; + +package Interfaces + partial package PartialMedium + constant Boolean reducedX = true; + constant Boolean fixedX = false; + constant String substanceNames[:] = {"single"}; + final constant Integer nS = size(substanceNames, 1); + constant Integer nX = nS; + constant Integer nXi = if fixedX then 0 else if reducedX then nS - 1 else nS; + constant Real reference_X[nX] = fill(1.0 / nX, nX); + + replaceable record ThermodynamicState + Real x; + end ThermodynamicState; + + replaceable model BaseProperties + input Real Xi[nXi]; + Real X[nX]; + ThermodynamicState state; + equation + X = reference_X; + Xi = X[1:nXi]; + end BaseProperties; + end PartialMedium; +end Interfaces; + +package TableBased + extends Interfaces.PartialMedium( + final reducedX = true, + final fixedX = true, + redeclare record ThermodynamicState = Common.BaseProps_Tpoly + ); + + redeclare model extends BaseProperties + equation + state.T = 1; + state.p = 2; + end BaseProperties; +end TableBased; + +model Base + replaceable package Medium = Interfaces.PartialMedium; + Medium.BaseProperties medium; +equation + medium.Xi = Medium.reference_X[1:Medium.nXi]; +end Base; + +model Probe + extends Base(redeclare package Medium = TableBased); +end Probe; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("parse failed"); + + let result = session.compile_model("Probe").expect("compile failed"); + + let medium_x_dims = result + .flat + .variables + .iter() + .find(|(name, _)| name.as_str() == "medium.X") + .map(|(_, var)| var.dims.clone()) + .expect("medium.X should exist"); + assert_eq!( + medium_x_dims, + vec![1], + "medium.X should use redeclared Medium.nX=1" + ); + + let medium_xi_dims = result + .flat + .variables + .iter() + .find(|(name, _)| name.as_str() == "medium.Xi") + .map(|(_, var)| var.dims.clone()) + .expect("medium.Xi should exist"); + assert_eq!( + medium_xi_dims, + vec![0], + "medium.Xi should use redeclared Medium.nXi=0" + ); + + let flat_var_names: Vec<_> = result + .flat + .variables + .keys() + .map(|k| k.to_string()) + .collect(); + assert!( + flat_var_names.iter().any(|n| n == "medium.state.T"), + "medium.state.T should come from redeclared ThermodynamicState; vars={flat_var_names:?}" + ); + assert!( + flat_var_names.iter().any(|n| n == "medium.state.p"), + "medium.state.p should come from redeclared ThermodynamicState; vars={flat_var_names:?}" + ); + assert!( + !flat_var_names.iter().any(|n| n == "medium.state.x"), + "base ThermodynamicState field should be replaced by redeclare; vars={flat_var_names:?}" + ); +} + +/// MLS §5.3: a local nested package declaration must shadow import aliases +/// with the same name when evaluating constants used in dimensions. +#[test] +fn test_local_package_shadows_import_alias_for_dimension_constants() { + let source = r#" +package Interfaces + partial package PartialMedium + constant String mediumName = "unset"; + constant String substanceNames[:] = {mediumName}; + final constant Integer nS = size(substanceNames, 1); + constant Integer nX = nS; + constant Integer nXi = 0; + + replaceable model BaseProperties + Real p; + Real h; + Real X[nX]; + input Real Xi[nXi]; + equation + X = fill(1.0, nX); + Xi = fill(0.0, nXi); + end BaseProperties; + end PartialMedium; + + partial package PartialPureSubstance + extends PartialMedium; + end PartialPureSubstance; + + partial package PartialSimpleMedium + extends PartialPureSubstance; + end PartialSimpleMedium; +end Interfaces; + +package MediumTwo + extends Interfaces.PartialSimpleMedium( + mediumName = "two", + substanceNames = {"A", "B"} + ); +end MediumTwo; + +package MediumOne + extends Interfaces.PartialSimpleMedium( + mediumName = "one", + substanceNames = {"A"} + ); +end MediumOne; + +model Target + import Medium = MediumTwo; + package Medium = MediumOne; + Medium.BaseProperties medium; +equation + medium.p = 1; + medium.h = 2; +end Target; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("parse failed"); + + let result = session.compile_model("Target").expect("compile failed"); + + let x_dims = result + .flat + .variables + .iter() + .find(|(name, _)| name.as_str() == "medium.X") + .map(|(_, var)| var.dims.clone()) + .expect("medium.X should exist"); + assert_eq!( + x_dims, + vec![1], + "local package Medium=MediumOne must shadow import Medium=MediumTwo for nX" + ); +} diff --git a/crates/rumoca/tests/pipeline_cases/alias_scope_and_dims/package_override_dimension_tests.rs b/crates/rumoca/tests/pipeline_cases/alias_scope_and_dims/package_override_dimension_tests.rs index cca8d6665..99d7bfec9 100644 --- a/crates/rumoca/tests/pipeline_cases/alias_scope_and_dims/package_override_dimension_tests.rs +++ b/crates/rumoca/tests/pipeline_cases/alias_scope_and_dims/package_override_dimension_tests.rs @@ -588,6 +588,60 @@ end P; ); } +/// MLS section 7.3: forwarding package redeclares must carry non-dimensional +/// package constants as well as size constants. Trace-substance fluid models +/// use `Medium.C_nominal` in nested balance equations. +#[test] +fn test_forwarding_package_redeclare_applies_to_nested_real_array_constant() { + let source = r#" +package P + package MediumBase + constant Integer nC = 0; + constant Real C_nominal[nC] = fill(0.0, nC); + end MediumBase; + + package MediumCO2 + extends MediumBase(nC = 1, C_nominal = {0.001519}); + end MediumCO2; + + model Duct + replaceable package Medium = MediumBase; + Real mCs_scaled[2, Medium.nC]; + parameter Real mbC_flows[2, Medium.nC] = fill(1.0, 2, Medium.nC); + equation + for i in 1:2 loop + der(mCs_scaled[i, :]) = mbC_flows[i, :] ./ Medium.C_nominal; + end for; + end Duct; + + model Source + replaceable package Medium = MediumBase; + Duct duct(redeclare package Medium = Medium); + end Source; + + model M + Source s(redeclare package Medium = MediumCO2); + end M; +end P; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("parse failed"); + + let result = session.compile_model("P.M").expect("compile failed"); + let dae_debug = format!("{:?}", result.dae.continuous.equations); + assert!( + !dae_debug.contains("Medium.C_nominal"), + "forwarded Medium.C_nominal should be substituted before DAE conversion: {dae_debug}" + ); + assert!( + dae_debug.contains("0.001519"), + "forwarded Medium.C_nominal value should come from MediumCO2: {dae_debug}" + ); +} + /// MLS §7.3: package aliases in sibling models must not leak into the active /// model's alias resolution. Compiling `Examples.A` should use `A.Medium`. #[test] diff --git a/crates/rumoca/tests/pipeline_cases/enum_if_branch_selection.rs b/crates/rumoca/tests/pipeline_cases/enum_if_branch_selection.rs new file mode 100644 index 000000000..f6ae0423d --- /dev/null +++ b/crates/rumoca/tests/pipeline_cases/enum_if_branch_selection.rs @@ -0,0 +1,97 @@ +//! Regression for enum-parameter branch selection in mismatched if-equations. + +use super::*; + +const SOURCE: &str = r#" +type Init = enumeration(NoInit, InitialState, InitialOutput); + +model FilterLike + parameter Init initType = Init.NoInit; + Real x; + Real y; +equation + if initType == Init.InitialState then + x = 1; + elseif initType == Init.InitialOutput then + x = 2; + y = 3; + end if; +end FilterLike; + +model UsesNestedEnumIfBranchSelection + FilterLike filter(initType = Init.InitialOutput); +end UsesNestedEnumIfBranchSelection; + +model UsesArrayEnumIfBranchSelection + FilterLike filter[2](each initType = Init.InitialOutput); +end UsesArrayEnumIfBranchSelection; + +record GenericChillerData + final parameter Integer nCapFunT = 6; + parameter Real capFunT[nCapFunT]; +end GenericChillerData; + +record ConcreteChillerData = GenericChillerData( + capFunT = {1, 2, 3, 4, 5, 6}); + +model UsesArrayRecordFieldDimension + parameter ConcreteChillerData per[2]; + Real y; +equation + y = per[2].capFunT[6]; +end UsesArrayRecordFieldDimension; +"#; + +#[test] +fn test_mismatched_if_uses_scoped_enum_parameter_value() { + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", SOURCE) + .expect("enum if fixture should parse"); + + let result = session + .compile_model("UsesNestedEnumIfBranchSelection") + .expect("scoped enum parameter should select the matching if branch"); + + assert!( + rumoca_phase_dae::balance::is_balanced(&result.dae).expect("valid DAE balance fixture"), + "selected branch should keep the model balanced: {}", + rumoca_phase_dae::balance::balance_detail(&result.dae).expect("valid DAE balance fixture") + ); +} + +#[test] +fn test_array_record_field_dimension_uses_canonical_parameter_value() { + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", SOURCE) + .expect("array record dimension fixture should parse"); + + let result = session + .compile_model("UsesArrayRecordFieldDimension") + .expect("array record field dimension should resolve for second element"); + + assert!( + rumoca_phase_dae::balance::is_balanced(&result.dae).expect("valid DAE balance fixture"), + "array record field dimension model should remain balanced: {}", + rumoca_phase_dae::balance::balance_detail(&result.dae).expect("valid DAE balance fixture") + ); +} + +#[test] +fn test_mismatched_if_uses_canonical_array_enum_parameter_value() { + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", SOURCE) + .expect("enum array if fixture should parse"); + + let result = session + .compile_model("UsesArrayEnumIfBranchSelection") + .expect("array enum parameter should select the matching if branch"); + + assert!( + rumoca_phase_dae::balance::is_balanced(&result.dae).expect("valid DAE balance fixture"), + "selected array branches should keep the model balanced: {}", + rumoca_phase_dae::balance::balance_detail(&result.dae).expect("valid DAE balance fixture") + ); +} diff --git a/crates/rumoca/tests/pipeline_cases/flat_output_regressions.rs b/crates/rumoca/tests/pipeline_cases/flat_output_regressions.rs index aff32a256..6429cb4b1 100644 --- a/crates/rumoca/tests/pipeline_cases/flat_output_regressions.rs +++ b/crates/rumoca/tests/pipeline_cases/flat_output_regressions.rs @@ -262,12 +262,23 @@ end AssertEmission; ) .expect("flat rendering should succeed"); + let flat_code_normalized = flat_code.replace("\r\n", "\n"); + let equation_idx = flat_code_normalized + .find("\nequation\n") + .expect("flat output should contain equation section"); + let initial_idx = flat_code_normalized + .find("\ninitial equation\n") + .expect("flat output should contain initial equation section"); + let equation_section = &flat_code_normalized[equation_idx..initial_idx]; + let initial_section = &flat_code_normalized[initial_idx..]; assert!( - flat_code.contains(r#"assert((k > 0), "k must stay positive");"#), + equation_section.contains(r#""k must stay positive""#) + && equation_section.contains("assert("), "expected equation-section assert in flat output, got:\n{flat_code}" ); assert!( - flat_code.contains(r#"assert((k >= 0), "k must stay nonnegative");"#), + initial_section.contains(r#""k must stay nonnegative""#) + && initial_section.contains("assert("), "expected initial-equation assert in flat output, got:\n{flat_code}" ); } diff --git a/crates/rumoca/tests/pipeline_cases/medium_base_properties_forwarding.rs b/crates/rumoca/tests/pipeline_cases/medium_base_properties_forwarding.rs new file mode 100644 index 000000000..0a3944014 --- /dev/null +++ b/crates/rumoca/tests/pipeline_cases/medium_base_properties_forwarding.rs @@ -0,0 +1,691 @@ +//! Regression for Buildings-style `Medium.BaseProperties` forwarding through +//! nested `redeclare package Medium = Medium` / `redeclare package Medium = MediumAir`. + +use super::*; + +const SOURCE: &str = r#" +package PartialMedium + replaceable partial model BaseProperties + Real p; + end BaseProperties; +end PartialMedium; + +package BuildingsMediaAir + extends PartialMedium; + + redeclare model extends BaseProperties + Real h; + equation + h = p + 1; + end BaseProperties; +end BuildingsMediaAir; + +block LumpedVolumeDeclarations + replaceable package Medium = PartialMedium; +end LumpedVolumeDeclarations; + +model ConservationEquation + extends LumpedVolumeDeclarations; + + Medium.BaseProperties medium; +equation + medium.p = 101325; +end ConservationEquation; + +model PartialMixingVolume + extends LumpedVolumeDeclarations; + + ConservationEquation dynBal(redeclare final package Medium = Medium); +end PartialMixingVolume; + +model PartialFlowMachine + extends LumpedVolumeDeclarations; + + PartialMixingVolume vol(redeclare package Medium = Medium); +end PartialFlowMachine; + +model WithoutMotorLike + replaceable package Medium = PartialMedium; + + PartialFlowMachine varSpeFloMov(redeclare package Medium = Medium); +end WithoutMotorLike; + +model Floor + replaceable package MediumAir = BuildingsMediaAir; + + WithoutMotorLike mover(redeclare package Medium = MediumAir); +end Floor; + +model UsesForwardedMediumBaseProperties + Floor floor1; +end UsesForwardedMediumBaseProperties; +"#; + +#[test] +fn test_nested_medium_base_properties_forwarding_compiles() { + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", SOURCE) + .expect("nested Medium forwarding fixture should parse"); + + session + .compile_model("UsesForwardedMediumBaseProperties") + .expect("nested package Medium forwarding must instantiate concrete BaseProperties"); +} + +const WATER_FINAL_BINDING_SOURCE: &str = r#" +package PartialMedium + type MolarMass = Real(unit = "kg/mol"); + + replaceable partial model BaseProperties + Real p; + end BaseProperties; +end PartialMedium; + +package BuildingsMediaWater + extends PartialMedium; + + constant MolarMass MM_const = 0.01801528; + + redeclare replaceable model BaseProperties + Real p; + final MolarMass MM = MM_const; + equation + p = 300000; + end BaseProperties; +end BuildingsMediaWater; + +model ConservationEquation + replaceable package Medium = PartialMedium; + + Medium.BaseProperties medium; +end ConservationEquation; + +model UsesWaterFinalBasePropertiesBinding + ConservationEquation dynBal(redeclare final package Medium = BuildingsMediaWater); +end UsesWaterFinalBasePropertiesBinding; +"#; + +#[test] +fn test_forwarded_water_base_properties_keeps_final_binding_equation() { + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", WATER_FINAL_BINDING_SOURCE) + .expect("water final binding fixture should parse"); + + let result = session + .compile_model("UsesWaterFinalBasePropertiesBinding") + .expect("water final binding fixture should compile"); + + assert!( + result + .dae + .continuous + .equations + .iter() + .any(|eq| eq.origin.contains("binding equation for dynBal.medium.MM")), + "final medium field binding should define dynBal.medium.MM; balance={}", + rumoca_phase_dae::balance::balance_detail(&result.dae).expect("valid DAE balance fixture") + ); + assert!( + rumoca_phase_dae::balance::is_balanced(&result.dae).expect("valid DAE balance fixture"), + "model should remain balanced: {}", + rumoca_phase_dae::balance::balance_detail(&result.dae).expect("valid DAE balance fixture") + ); +} + +const HEX_SOURCE: &str = r#" +package PartialMedium + replaceable partial model BaseProperties + Real p; + end BaseProperties; +end PartialMedium; + +package BuildingsMediaAir + extends PartialMedium; + + redeclare model extends BaseProperties + Real h; + equation + h = p + 1; + end BaseProperties; +end BuildingsMediaAir; + +package BuildingsMediaWater + extends PartialMedium; + + redeclare model extends BaseProperties + Real h; + equation + h = p + 2; + end BaseProperties; +end BuildingsMediaWater; + +block LumpedVolumeDeclarations + replaceable package Medium = PartialMedium; +end LumpedVolumeDeclarations; + +model ConservationEquation + extends LumpedVolumeDeclarations; + + Medium.BaseProperties medium; +equation + medium.p = 101325; +end ConservationEquation; + +model PartialMixingVolume + extends LumpedVolumeDeclarations; + + ConservationEquation dynBal(redeclare final package Medium = Medium); +end PartialMixingVolume; + +model PartialHexElementLike + replaceable package Medium1 = PartialMedium; + replaceable package Medium2 = PartialMedium; + + PartialMixingVolume vol1(redeclare final package Medium = Medium1); + PartialMixingVolume vol2(redeclare final package Medium = Medium2); +end PartialHexElementLike; + +model WetCoilLike + replaceable package MediumWat = BuildingsMediaWater; + replaceable package MediumAir = BuildingsMediaAir; + + PartialHexElementLike ele[1]( + redeclare each package Medium1 = MediumWat, + redeclare each package Medium2 = MediumAir); +end WetCoilLike; + +model UsesHexMediumBaseProperties + WetCoilLike coil; +end UsesHexMediumBaseProperties; +"#; + +const CONSTRAINEDBY_HEX_SOURCE: &str = r#" +package PartialMedium + replaceable partial model BaseProperties + Real p; + end BaseProperties; +end PartialMedium; + +package BuildingsMediaWater + extends PartialMedium; + + redeclare model extends BaseProperties + Real h; + equation + h = p + 2; + end BaseProperties; +end BuildingsMediaWater; + +block LumpedVolumeDeclarations + replaceable package Medium = PartialMedium; +end LumpedVolumeDeclarations; + +model ConservationEquation + extends LumpedVolumeDeclarations; + + Medium.BaseProperties medium; +equation + medium.p = 101325; +end ConservationEquation; + +model PartialMixingVolume + extends LumpedVolumeDeclarations; + + ConservationEquation dynBal(redeclare final package Medium = Medium); +end PartialMixingVolume; + +model PartialHexElementLike + replaceable package Medium1 = PartialMedium; + + replaceable PartialMixingVolume vol1 constrainedby PartialMixingVolume( + redeclare final package Medium = Medium1); +end PartialHexElementLike; + +model WetCoilLike + replaceable package MediumWat = BuildingsMediaWater; + + PartialHexElementLike ele[1](redeclare each package Medium1 = MediumWat); +end WetCoilLike; + +model UsesConstrainedbyHexMediumBaseProperties + WetCoilLike coil; +end UsesConstrainedbyHexMediumBaseProperties; +"#; + +const HEX_ELEMENT_LATENT_SOURCE: &str = r#" +package PartialMedium + replaceable partial model BaseProperties + Real p; + end BaseProperties; +end PartialMedium; + +package BuildingsMediaWater + extends PartialMedium; + + redeclare model extends BaseProperties + Real h; + equation + h = p + 2; + end BaseProperties; +end BuildingsMediaWater; + +block LumpedVolumeDeclarations + replaceable package Medium = PartialMedium; +end LumpedVolumeDeclarations; + +model ConservationEquation + extends LumpedVolumeDeclarations; + + Medium.BaseProperties medium; +equation + medium.p = 101325; +end ConservationEquation; + +model PartialMixingVolume + extends LumpedVolumeDeclarations; + + parameter Boolean initialize_p = true; + ConservationEquation dynBal( + redeclare final package Medium = Medium, + final initialize_p = initialize_p); +end PartialMixingVolume; + +model PartialHexElementLike + replaceable package Medium1 = PartialMedium; + parameter Boolean initialize_p1 = true; + + replaceable PartialMixingVolume vol1 constrainedby PartialMixingVolume( + redeclare final package Medium = Medium1); +end PartialHexElementLike; + +model HexElementLatentLike + extends PartialHexElementLike( + redeclare final PartialMixingVolume vol1(final initialize_p = initialize_p1)); +end HexElementLatentLike; + +model WetCoilLike + replaceable package MediumWat = BuildingsMediaWater; + + HexElementLatentLike ele[1](redeclare each package Medium1 = MediumWat); +end WetCoilLike; + +model UsesHexElementLatentMediumBaseProperties + WetCoilLike coil; +end UsesHexElementLatentMediumBaseProperties; +"#; + +#[test] +fn test_hex_element_latent_constrainedby_medium_forwarding_compiles() { + let mut session = Session::new(SessionConfig::default()); + session + .add_document("hex_element_latent_test.mo", HEX_ELEMENT_LATENT_SOURCE) + .expect("hex element latent constrainedby fixture should parse"); + + session + .compile_model("UsesHexElementLatentMediumBaseProperties") + .expect("hex element latent constrainedby Medium forwarding must compile"); +} + +const HEX_ELEMENT_LATENT_VOLUME_REDECLARE_SOURCE: &str = r#" +connector RealInput = input Real; + +model ConservationEquation + parameter Boolean use_mWat_flow = false; + RealInput mWat_flow if use_mWat_flow; + Real mWat_flow_internal; +equation + if use_mWat_flow then + mWat_flow_internal = mWat_flow; + else + mWat_flow_internal = 0; + end if; +end ConservationEquation; + +model PartialMixingVolume + parameter Real V = 1; + ConservationEquation dynBal(final use_mWat_flow = false); +end PartialMixingVolume; + +model MixingVolume + extends PartialMixingVolume; +end MixingVolume; + +model MixingVolumeHeatPort + extends PartialMixingVolume; +end MixingVolumeHeatPort; + +model MixingVolumeHeatMoisturePort + extends PartialMixingVolume(dynBal(final use_mWat_flow = true)); + + RealInput mWat_flow; +equation + connect(mWat_flow, dynBal.mWat_flow); +end MixingVolumeHeatMoisturePort; + +model FourPortHeatMassExchangerLike + replaceable MixingVolume vol2 constrainedby MixingVolumeHeatPort; +end FourPortHeatMassExchangerLike; + +model PartialHexElementLike + extends FourPortHeatMassExchangerLike; +end PartialHexElementLike; + +model HexElementLatentLike + extends PartialHexElementLike( + redeclare final MixingVolumeHeatMoisturePort vol2); +end HexElementLatentLike; + +model UsesHexElementLatentVolumeRedeclare + HexElementLatentLike ele[1]; + Real z; +equation + ele[1].vol2.mWat_flow = 1; + z = ele[1].vol2.dynBal.mWat_flow_internal; +end UsesHexElementLatentVolumeRedeclare; +"#; + +#[test] +fn test_hex_element_latent_component_redeclare_uses_redeclared_volume_modifiers() { + let mut session = Session::new(SessionConfig::default()); + session + .add_document( + "hex_element_latent_volume_redeclare_test.mo", + HEX_ELEMENT_LATENT_VOLUME_REDECLARE_SOURCE, + ) + .expect("hex element latent volume redeclare fixture should parse"); + + let result = session + .compile_model("UsesHexElementLatentVolumeRedeclare") + .expect("component redeclare must instantiate the redeclared volume type"); + let json = serde_json::to_value(&result.dae).expect("DAE should serialize"); + + assert!( + json["u"].get("ele[1].vol2.mWat_flow").is_some() + || json["w"].get("ele[1].vol2.mWat_flow").is_some() + || json["y"].get("ele[1].vol2.mWat_flow").is_some(), + "redeclared volume should expose mWat_flow connector; got DAE:\n{json:#}" + ); + assert_eq!( + json["p"]["ele[1].vol2.dynBal.use_mWat_flow"]["start"]["Literal"]["value"]["Boolean"], true, + "redeclared volume's inherited dynBal modifier must force use_mWat_flow=true" + ); +} + +const COMPONENT_REDECLARE_OVERRIDE_SCOPE_SOURCE: &str = r#" +connector RealInput = input Real; + +model ConservationEquation + parameter Boolean use_mWat_flow = false; + RealInput mWat_flow if use_mWat_flow; + Real mWat_flow_internal; +equation + if use_mWat_flow then + mWat_flow_internal = mWat_flow; + else + mWat_flow_internal = 0; + end if; +end ConservationEquation; + +model PartialMixingVolume + ConservationEquation dynBal(final use_mWat_flow = false); +end PartialMixingVolume; + +model MixingVolume + extends PartialMixingVolume; +end MixingVolume; + +model MixingVolumeHeatPort + extends PartialMixingVolume; +end MixingVolumeHeatPort; + +model MixingVolumeHeatMoisturePort + extends PartialMixingVolume(dynBal(final use_mWat_flow = true)); + RealInput mWat_flow; +equation + connect(mWat_flow, dynBal.mWat_flow); +end MixingVolumeHeatMoisturePort; + +model FourPortLike + replaceable MixingVolume vol2 constrainedby MixingVolumeHeatPort; +end FourPortLike; + +model LatentElementLike + extends FourPortLike(redeclare final MixingVolumeHeatMoisturePort vol2); +end LatentElementLike; + +model EightPortLike + MixingVolume vol2; +end EightPortLike; + +model InternalHexLike + extends EightPortLike(vol2(final V = 2)); +end InternalHexLike; + +model UsesLatentAndInternalHex + LatentElementLike latent; + InternalHexLike internalHex; + Real z1; + Real z2; +equation + latent.vol2.mWat_flow = 1; + z1 = latent.vol2.dynBal.mWat_flow_internal; + z2 = internalHex.vol2.dynBal.mWat_flow_internal; +end UsesLatentAndInternalHex; +"#; + +#[test] +fn test_component_redeclare_override_does_not_leak_to_unrelated_same_named_component() { + let mut session = Session::new(SessionConfig::default()); + session + .add_document( + "component_redeclare_override_scope_test.mo", + COMPONENT_REDECLARE_OVERRIDE_SCOPE_SOURCE, + ) + .expect("component redeclare override scope fixture should parse"); + + session + .compile_model("UsesLatentAndInternalHex") + .expect("component redeclare override must not leak to an unrelated same-named component"); +} + +#[test] +fn test_hex_medium1_medium2_base_properties_forwarding_compiles() { + let mut session = Session::new(SessionConfig::default()); + session + .add_document("hex_test.mo", HEX_SOURCE) + .expect("hex Medium1/Medium2 forwarding fixture should parse"); + + session + .compile_model("UsesHexMediumBaseProperties") + .expect("hex element Medium1/Medium2 forwarding must instantiate concrete BaseProperties"); +} + +const MEDIUM_NXI_FOR_SOURCE: &str = r#" +package PartialMedium + constant Integer nXi = 1; +end PartialMedium; + +package MediumAir + extends PartialMedium(nXi = 2); +end MediumAir; + +block LumpedVolumeDeclarations + replaceable package Medium = PartialMedium; +end LumpedVolumeDeclarations; + +model ConservationEquation + extends LumpedVolumeDeclarations; + + Real s[Medium.nXi]; +equation + for i in 1:Medium.nXi loop + s[i] = i; + end for; +end ConservationEquation; + +model PartialMixingVolume + extends LumpedVolumeDeclarations; + + ConservationEquation dynBal(redeclare final package Medium = Medium); +end PartialMixingVolume; + +model PartialFlowMachine + extends LumpedVolumeDeclarations; + + PartialMixingVolume vol(redeclare package Medium = Medium); +end PartialFlowMachine; + +model WithoutMotorLike + replaceable package Medium = PartialMedium; + + PartialFlowMachine varSpeFloMov(redeclare package Medium = MediumAir); +end WithoutMotorLike; + +model UsesNestedMediumNxiForRange + WithoutMotorLike mover; +end UsesNestedMediumNxiForRange; +"#; + +#[test] +fn test_nested_medium_nxi_for_range_compiles() { + let mut session = Session::new(SessionConfig::default()); + session + .add_document("medium_nxi_for_test.mo", MEDIUM_NXI_FOR_SOURCE) + .expect("nested Medium.nXi for-range fixture should parse"); + + let result = session + .compile_model("UsesNestedMediumNxiForRange") + .expect("nested Medium forwarding must resolve Medium.nXi in dynBal for-equations"); + + assert!( + rumoca_phase_dae::balance::is_balanced(&result.dae).expect("valid DAE balance fixture"), + "nested dynBal for-range should expand with MediumAir.nXi=2: {}", + rumoca_phase_dae::balance::balance_detail(&result.dae).expect("valid DAE balance fixture") + ); +} + +#[test] +fn test_constrainedby_medium_base_properties_forwarding_compiles() { + let mut session = Session::new(SessionConfig::default()); + session + .add_document("constrainedby_hex_test.mo", CONSTRAINEDBY_HEX_SOURCE) + .expect("constrainedby Medium forwarding fixture should parse"); + + session + .compile_model("UsesConstrainedbyHexMediumBaseProperties") + .expect("constrainedby Medium binding must forward to concrete BaseProperties"); +} + +const MEDIUM_FUNCTION_FORWARDING_SOURCE: &str = r#" +package Modelica + package Media + package Interfaces + partial package PartialMedium + replaceable record ThermodynamicState + Real p; + Real T; + end ThermodynamicState; + + replaceable partial function setState_pTX + input Real p; + input Real T; + output ThermodynamicState state; + end setState_pTX; + + replaceable function specificEnthalpy + input ThermodynamicState state; + output Real h; + algorithm + h := state.T + state.p; + end specificEnthalpy; + + replaceable function specificEnthalpy_pTX + input Real p; + input Real T; + output Real h; + algorithm + h := specificEnthalpy(setState_pTX(p, T)); + end specificEnthalpy_pTX; + end PartialMedium; + end Interfaces; + end Media; +end Modelica; + +package BuildingsMediaAir + extends Modelica.Media.Interfaces.PartialMedium; + + redeclare record extends ThermodynamicState + end ThermodynamicState; + + redeclare function extends setState_pTX + algorithm + state := ThermodynamicState(p=p, T=T); + end setState_pTX; +end BuildingsMediaAir; + +block LumpedVolumeDeclarations + replaceable package Medium = Modelica.Media.Interfaces.PartialMedium; +end LumpedVolumeDeclarations; + +model ConservationEquation + extends LumpedVolumeDeclarations; + + Real h; +equation + h = Medium.specificEnthalpy(Medium.setState_pTX(101325, 293.15)); +end ConservationEquation; + +model InheritedFunctionCaller + extends LumpedVolumeDeclarations; + + Real h; +equation + h = Medium.specificEnthalpy_pTX(101325, 293.15); +end InheritedFunctionCaller; + +model PartialMixingVolume + extends LumpedVolumeDeclarations; + + ConservationEquation dynBal(redeclare final package Medium = Medium); + InheritedFunctionCaller funBal(redeclare final package Medium = Medium); +end PartialMixingVolume; + +model PartialFlowMachine + extends LumpedVolumeDeclarations; + + PartialMixingVolume vol(redeclare package Medium = Medium); +end PartialFlowMachine; + +model WithoutMotorLike + replaceable package Medium = Modelica.Media.Interfaces.PartialMedium; + + PartialFlowMachine varSpeFloMov(redeclare package Medium = Medium); +end WithoutMotorLike; + +model Floor + replaceable package MediumAir = BuildingsMediaAir; + + WithoutMotorLike mover(redeclare package Medium = MediumAir); +end Floor; + +model UsesNestedMediumFunctionForwarding + Floor floor1; +end UsesNestedMediumFunctionForwarding; +"#; + +#[test] +fn test_nested_medium_function_forwarding_compiles() { + let mut session = Session::new(SessionConfig::default()); + session + .add_document( + "medium_function_forwarding_test.mo", + MEDIUM_FUNCTION_FORWARDING_SOURCE, + ) + .expect("nested Medium function forwarding fixture should parse"); + + session + .compile_model("UsesNestedMediumFunctionForwarding") + .expect("nested Medium forwarding must rewrite setState_pTX to the concrete medium"); +} diff --git a/crates/rumoca/tests/pipeline_cases/mod.rs b/crates/rumoca/tests/pipeline_cases/mod.rs index ad4925952..d4149c13b 100644 --- a/crates/rumoca/tests/pipeline_cases/mod.rs +++ b/crates/rumoca/tests/pipeline_cases/mod.rs @@ -1,6 +1,8 @@ use super::*; mod alias_scope_and_dims; +mod enum_if_branch_selection; mod flat_output_regressions; +mod medium_base_properties_forwarding; mod package_alias_regressions; mod record_constant_arrays; diff --git a/crates/rumoca/tests/pipeline_cases/package_alias_regressions.rs b/crates/rumoca/tests/pipeline_cases/package_alias_regressions.rs index 44985d472..9c11ce10a 100644 --- a/crates/rumoca/tests/pipeline_cases/package_alias_regressions.rs +++ b/crates/rumoca/tests/pipeline_cases/package_alias_regressions.rs @@ -1,5 +1,72 @@ use super::*; +/// MLS §7.3.2: if a replaceable package declaration has no explicit +/// `constrainedby`, the declaration type is its implicit constraining type. +/// +/// A later redeclare must be checked against that original constraining type, +/// not against the local alias class introduced by the replaceable declaration. +#[test] +fn test_replaceable_package_implicit_constraint_uses_declared_type() { + let source = r#" +package Modelica + package Media + package Interfaces + partial package PartialMedium + end PartialMedium; + + partial package PartialMixtureMedium + extends PartialMedium; + end PartialMixtureMedium; + + partial package PartialCondensingGases + extends PartialMixtureMedium; + end PartialCondensingGases; + end Interfaces; + end Media; +end Modelica; + +package Buildings + package Fluid + package Interfaces + partial model PartialFourPort + replaceable package Medium2 = + Modelica.Media.Interfaces.PartialMedium; + end PartialFourPort; + + partial model PartialFourPortInterface + extends PartialFourPort; + end PartialFourPortInterface; + end Interfaces; + + package HeatExchangers + model DryCoilCounterFlow + extends Interfaces.PartialFourPortInterface; + end DryCoilCounterFlow; + + model WetCoilCounterFlow + extends DryCoilCounterFlow( + redeclare replaceable package Medium2 = + Modelica.Media.Interfaces.PartialCondensingGases); + end WetCoilCounterFlow; + end HeatExchangers; + end Fluid; +end Buildings; + +model Probe + Buildings.Fluid.HeatExchangers.WetCoilCounterFlow coil; +end Probe; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("parse failed"); + + session + .compile_model("Probe") + .expect("implicit constraining type should accept PartialCondensingGases"); +} + /// MLS §7.3: constant evaluation in nested package aliases must use alias-local scope. /// /// Without scope-aware lookup, `size(substanceNames, 1)` for `Medium.nS` can resolve to diff --git a/crates/rumoca/tests/pipeline_test.rs b/crates/rumoca/tests/pipeline_test.rs index caa5c6af2..608e47e9e 100644 --- a/crates/rumoca/tests/pipeline_test.rs +++ b/crates/rumoca/tests/pipeline_test.rs @@ -95,6 +95,82 @@ fn expr_if_branch_mentions_var(expr: &rumoca_core::Expression, var_name: &str) - } } +fn find_function_call_args<'a>( + expr: &'a rumoca_core::Expression, + target_name: &str, +) -> Option<&'a [rumoca_core::Expression]> { + match expr { + rumoca_core::Expression::FunctionCall { name, args, .. } => { + if name.as_str() == target_name || name.as_str().ends_with(&format!(".{target_name}")) { + Some(args) + } else { + args.iter() + .find_map(|arg| find_function_call_args(arg, target_name)) + } + } + rumoca_core::Expression::Binary { lhs, rhs, .. } => { + find_function_call_args(lhs, target_name) + .or_else(|| find_function_call_args(rhs, target_name)) + } + rumoca_core::Expression::Unary { rhs, .. } => find_function_call_args(rhs, target_name), + rumoca_core::Expression::BuiltinCall { args, .. } => args + .iter() + .find_map(|arg| find_function_call_args(arg, target_name)), + rumoca_core::Expression::If { + branches, + else_branch, + .. + } => branches + .iter() + .find_map(|(_, value)| find_function_call_args(value, target_name)) + .or_else(|| find_function_call_args(else_branch, target_name)), + rumoca_core::Expression::Array { elements, .. } + | rumoca_core::Expression::Tuple { elements, .. } => elements + .iter() + .find_map(|element| find_function_call_args(element, target_name)), + rumoca_core::Expression::Range { + start, step, end, .. + } => find_function_call_args(start, target_name) + .or_else(|| { + step.as_ref() + .and_then(|step| find_function_call_args(step, target_name)) + }) + .or_else(|| find_function_call_args(end, target_name)), + rumoca_core::Expression::ArrayComprehension { + expr, + indices, + filter, + .. + } => find_function_call_args(expr, target_name) + .or_else(|| { + indices + .iter() + .find_map(|range_idx| find_function_call_args(&range_idx.range, target_name)) + }) + .or_else(|| { + filter + .as_ref() + .and_then(|filter| find_function_call_args(filter, target_name)) + }), + rumoca_core::Expression::Index { + base, subscripts, .. + } => find_function_call_args(base, target_name).or_else(|| { + subscripts.iter().find_map(|subscript| match subscript { + rumoca_core::Subscript::Expr { expr, .. } => { + find_function_call_args(expr, target_name) + } + _ => None, + }) + }), + rumoca_core::Expression::FieldAccess { base, .. } => { + find_function_call_args(base, target_name) + } + rumoca_core::Expression::VarRef { .. } + | rumoca_core::Expression::Literal { .. } + | rumoca_core::Expression::Empty { .. } => None, + } +} + // ============================================================================= // Basic pipeline tests // ============================================================================= @@ -615,6 +691,173 @@ end RealRangeDims; ); } +#[test] +fn test_nested_modifier_colon_dimension_feeds_sibling_size() { + let source = r#" +package ET004Fixture + model Curve + parameter Real V_flow[:]; + parameter Real dp[size(V_flow, 1)]; + end Curve; + + model Performance + parameter Real q_flow[:] = {1.0, 2.0, 3.0}; + Curve pressure(V_flow=q_flow, dp={10.0, 20.0, 30.0}); + end Performance; + + model Mover + parameter Real VolFloCur[:] = {1.0, 2.0, 3.0}; + Performance per(q_flow=VolFloCur); + Real y; + equation + y = per.pressure.dp[1]; + end Mover; +end ET004Fixture; +"#; + + let result = compile_model(source, "ET004Fixture.Mover") + .expect("nested modifier colon dimensions should feed sibling size()"); + + let v_flow_dims = result + .variables + .parameters + .get(&DaeVarName::new("per.pressure.V_flow")) + .map(|v| v.dims.clone()) + .expect("per.pressure.V_flow should exist in DAE parameters"); + assert_eq!( + v_flow_dims, + vec![3], + "V_flow[:] should infer dimensions from the active modifier binding" + ); + + let dp_dims = result + .variables + .parameters + .get(&DaeVarName::new("per.pressure.dp")) + .map(|v| v.dims.clone()) + .expect("per.pressure.dp should exist in DAE parameters"); + assert_eq!( + dp_dims, + vec![3], + "dp[size(V_flow, 1)] should resolve through the sibling V_flow dimension" + ); +} + +#[test] +fn test_array_component_modifier_colon_dimension_feeds_sibling_size() { + let source = r#" +package ET004ArrayFixture + model Curve + parameter Real V_flow[:]; + parameter Real dp[size(V_flow, 1)]; + end Curve; + + model Mover + parameter Real VolFloCur[:] = {1.0, 2.0, 3.0}; + Curve pressure(V_flow=VolFloCur, dp={10.0, 20.0, 30.0}); + end Mover; + + model Plant + parameter Real curves[2, 3] = [1.0, 2.0, 3.0; 4.0, 5.0, 6.0]; + Mover pum[2](VolFloCur=curves); + Real y; + equation + y = pum[2].pressure.dp[1]; + end Plant; +end ET004ArrayFixture; +"#; + + let result = compile_model(source, "ET004ArrayFixture.Plant") + .expect("array component modifiers should feed sibling size()"); + + let v_flow_dims = result + .variables + .parameters + .get(&DaeVarName::new("pum[2].pressure.V_flow")) + .map(|v| v.dims.clone()) + .expect("pum[2].pressure.V_flow should exist in DAE parameters"); + assert_eq!( + v_flow_dims, + vec![3], + "V_flow[:] should infer dimensions after array-component modifier distribution" + ); + + let dp_dims = result + .variables + .parameters + .get(&DaeVarName::new("pum[2].pressure.dp")) + .map(|v| v.dims.clone()) + .expect("pum[2].pressure.dp should exist in DAE parameters"); + assert_eq!( + dp_dims, + vec![3], + "dp[size(V_flow, 1)] should resolve for the distributed array element" + ); +} + +#[test] +fn test_extends_forwarded_modifier_dotted_size_feeds_sibling_dimension() { + let source = r#" +package ET004DottedFixture + model Curve + parameter Real V_flow[:]; + parameter Real eta[size(V_flow, 1)]; + end Curve; + + model PerformanceBase + parameter Real flo[:]; + Curve motorEfficiency(V_flow=flo, eta=fill(1.0, size(flo, 1))); + end PerformanceBase; + + model Performance + extends PerformanceBase; + end Performance; + + model MoverBase + parameter Real VolFloCur[:]; + Performance per(flo=VolFloCur); + parameter Integer n = size(per.motorEfficiency.V_flow, 1); + parameter Real eff[n] = per.motorEfficiency.eta; + end MoverBase; + + model Mover + extends MoverBase(VolFloCur={1.0, 2.0, 3.0}); + Real y; + equation + y = eff[1]; + end Mover; +end ET004DottedFixture; +"#; + + let result = compile_model(source, "ET004DottedFixture.Mover").expect( + "extends-forwarded nested modifier dimensions should feed dotted size() in active scope", + ); + + let v_flow_dims = result + .variables + .parameters + .get(&DaeVarName::new("per.motorEfficiency.V_flow")) + .map(|v| v.dims.clone()) + .expect("per.motorEfficiency.V_flow should exist in DAE parameters"); + assert_eq!( + v_flow_dims, + vec![3], + "nested V_flow[:] should infer dimensions from the forwarded modifier" + ); + + let eff_dims = result + .variables + .parameters + .get(&DaeVarName::new("eff")) + .map(|v| v.dims.clone()) + .expect("eff should exist in DAE parameters"); + assert_eq!( + eff_dims, + vec![3], + "eff[size(per.motorEfficiency.V_flow, 1)] should resolve in MoverBase scope" + ); +} + #[test] fn test_array_comprehension_function_call_equation_preserves_dependencies() { use std::collections::HashSet; @@ -771,6 +1014,75 @@ end P; ); } +#[test] +fn test_named_record_actual_is_decomposed_by_callee_parameter_name() { + let source = r#" +package P + record R + Real a; + Real b; + end R; + + function f + input R r; + output Real y; + algorithm + y := r.a; + end f; + + model NamedRecordActual + R rec(a = 1, b = 2); + Real y; + equation + y = P.f(r = P.R(rec.a, rec.b)); + end NamedRecordActual; +end P; +"#; + + let mut session = Session::new(SessionConfig::default()); + session + .add_document("test.mo", source) + .expect("parse/resolve/typecheck failed"); + + let phase_result = session + .compile_model_phases("P.NamedRecordActual") + .expect("phase compilation should succeed"); + let result = match phase_result { + PhaseResult::Success(result) => result, + other => panic!( + "expected successful phase result, got {:?}", + std::mem::discriminant(&other) + ), + }; + + let function = result + .flat + .functions + .get(&rumoca_core::VarName::new("P.f")) + .expect("function P.f should remain in flat IR"); + let input_names = function + .inputs + .iter() + .map(|input| input.name.as_str()) + .collect::>(); + assert_eq!(input_names, vec!["r_a", "r_b"]); + + let call_args = result + .flat + .equations + .iter() + .find_map(|eq| find_function_call_args(&eq.residual, "P.f")) + .expect("flat equation should contain rewritten P.f call"); + let actual_names = call_args + .iter() + .map(|arg| match arg { + rumoca_core::Expression::VarRef { name, .. } => name.as_str().to_string(), + other => format!("{other:?}"), + }) + .collect::>(); + assert_eq!(actual_names, vec!["rec.a", "rec.b"]); +} + #[test] fn test_binding_equation_kept_when_explicit_rhs_refs_subscripted_unknowns() { let source = r#" diff --git a/crates/rumoca/tests/spec_budget_test.rs b/crates/rumoca/tests/spec_budget_test.rs index 654df6aa3..b0e5c802d 100644 --- a/crates/rumoca/tests/spec_budget_test.rs +++ b/crates/rumoca/tests/spec_budget_test.rs @@ -195,6 +195,154 @@ fn test_spec_0025_aligns_with_pr_template() { ); } +fn collect_missing_spec_recovery_contract(spec_recovery: &str, missing: &mut Vec) { + let required_spec_contract = [ + "Explicitly authorized ClimaMind Rumoca broken-main recovery batch", + "normal reviewer gate remains unchanged", + "`authorization_ref` MUST identify a durable record in the validation integration PR body or a maintainer-controlled GitHub artifact", + "`authorized_by` MUST identify a ClimaMind Rumoca repository maintainer", + "this task's explicit maintainer authorization is sufficient; no additional maintainer or approval is required", + "required fields: `authorization_ref`, `authorized_by`, `batch_id`, authorized ordered `owner_prs`, `target_branch`, and RFC 3339 UTC `expires_at`", + "automatically becomes inactive and MUST fail closed as soon as any one of these conditions is true", + "RFC 3339 `expires_at` has passed", + "every authorized owner PR has landed", + "all required CI checks on the target `main` are green", + "Before each owner PR merge, the authoritative record MUST exist, match the recorded batch, PR, head, and target values, and remain unexpired", + "independent technical review", + "owner mechanism test", + "evidence to that owner PR's final `head_sha`", + "exact-head integration hosted CI is green", + "all required hosted CI checks are green", + "then merge in sequence", + "owner PR `head_sha` values", + "Every listed final owner `head_sha`, including the recovery-rule PR `head_sha`, MUST be a Git ancestor of the integration `head_sha`", + "recorded target baseline `head_sha` MUST be a Git ancestor of the integration `head_sha`", + "Cherry-pick, patch-id, squash, or content equivalence is not exact provenance", + "target baseline, the listed exact owner histories, and signed merge commits only", + "Every such merge commit MUST carry exactly one `Signed-off-by` trailer and no `Co-Authored-By` trailer", + "MUST NOT contain any integration-only production, test, spec, workflow, baseline, validator, tolerance, fixture, or content commit", + "hosted CI workflow `head_sha` MUST equal the recorded integration PR `head_sha`", + "Any owner PR `head_sha`, target baseline `head_sha`, or integration PR `head_sha` change MUST invalidate affected evidence and fail closed", + "reconstruct the integration PR, refresh affected review or mechanism-test evidence, and rerun all required hosted CI", + "No GitHub approving review is required only for owner PRs in that active batch", + "Draft", + "validation-only", + "MUST NEVER merge", + "MUST NOT contain unique fixes", + "MUST NOT weaken or bypass any existing gate", + "MUST NOT apply to third-party contributors or an unauthorized batch", + ]; + for required in required_spec_contract { + if !spec_recovery.contains(required) { + missing.push(format!("SPEC_0025 missing recovery contract: `{required}`")); + } + } +} + +fn collect_missing_template_recovery_contract(template_recovery: &str, missing: &mut Vec) { + let required_template_contract = [ + "## Authorized Broken-Main Recovery (optional)", + "Leave blank for normal PRs", + "Explicitly authorized ClimaMind Rumoca broken-main recovery batch", + "`authorization_ref` (durable validation integration PR body or maintainer-controlled GitHub artifact):", + "`authorized_by` (ClimaMind Rumoca repository maintainer):", + "`batch_id`:", + "Authorized ordered `owner_prs`:", + "`target_branch` / baseline `head_sha`:", + "RFC 3339 UTC `expires_at`:", + "Owner PR / final `head_sha`:", + "Independent technical review / reviewed `head_sha`:", + "Owner mechanism test / tested `head_sha`:", + "Recovery-rule PR / final `head_sha`:", + "Integration PR / `head_sha`:", + "Hosted CI workflow / `head_sha`:", + "Authorization exists, matches this merge, and is unexpired.", + "`authorized_by` is a ClimaMind Rumoca maintainer; this task's explicit authorization is sufficient, with no additional maintainer or approval.", + "Recovery is inactive and fails closed if `expires_at` passed, every authorized owner PR landed, or target `main` has all required CI green.", + "Evidence is bound to the owner final head and recorded in order; merge only after all required hosted CI is green on the integration head.", + "Every listed final owner head, including the recovery-rule PR head, is a Git ancestor of the integration head; no cherry-pick, patch-id, squash, or content-equivalent substitute.", + "Recorded target baseline `head_sha` is a Git ancestor of the integration head.", + "Integration history = target baseline + listed exact owner histories + signed merge commits only; no integration-only production, test, spec, workflow, baseline, validator, tolerance, fixture, or content commit.", + "Every integration merge commit has exactly one `Signed-off-by` trailer and no `Co-Authored-By` trailer.", + "CI workflow head = integration head.", + "Any owner, baseline, or integration head change fails closed; rebuild and rerun affected evidence and CI.", + "Draft, validation-only, never merge", + ]; + for required in required_template_contract { + if !template_recovery.contains(required) { + missing.push(format!( + "PR template missing recovery linkage: `{required}`" + )); + } + } +} + +fn collect_forbidden_recovery_contract( + spec_recovery: &str, + template_recovery: &str, + missing: &mut Vec, +) { + let forbidden_contract = [ + "`authorization_url`", + "independent maintainer", + "another maintainer", + "not the author of an owner PR", + "not self-attested by an owner-PR author", + ]; + for forbidden in forbidden_contract { + if spec_recovery.contains(forbidden) { + missing.push(format!( + "SPEC_0025 recovery contract retains forbidden requirement: `{forbidden}`" + )); + } + if template_recovery.contains(forbidden) { + missing.push(format!( + "PR template recovery linkage retains forbidden requirement: `{forbidden}`" + )); + } + } +} + +#[test] +fn test_spec_0025_preserves_authorized_broken_main_recovery_contract() { + // The narrow recovery path is documentation-enforced policy. Keep its + // activation boundary, ordered evidence, expiry, and integration-only + // restrictions mechanically visible so a later edit cannot broaden it. + let root = workspace_root(); + let spec = fs::read_to_string(root.join("spec/SPEC_0025_PR_REVIEW_PROCESS.md")) + .expect("read SPEC_0025"); + let template = fs::read_to_string(root.join(".github/pull_request_template.md")) + .expect("read PR template"); + + let spec_heading = "### 6a. Authorized Broken-Main Recovery (optional)"; + let spec_tail = &spec[spec.find(spec_heading).expect("SPEC_0025 recovery section")..]; + let spec_recovery = &spec_tail[..spec_tail + .find("\n### 7.") + .expect("SPEC_0025 recovery section end")]; + let template_heading = "## Authorized Broken-Main Recovery (optional)"; + let template_tail = &template[template + .find(template_heading) + .expect("PR template recovery section")..]; + let template_after_heading = &template_tail[template_heading.len()..]; + let template_end = template_after_heading + .find("\n## ") + .map_or(template_tail.len(), |offset| { + template_heading.len() + offset + }); + let template_recovery = &template_tail[..template_end]; + + let mut missing = Vec::new(); + collect_missing_spec_recovery_contract(spec_recovery, &mut missing); + collect_missing_template_recovery_contract(template_recovery, &mut missing); + collect_forbidden_recovery_contract(spec_recovery, template_recovery, &mut missing); + + assert!( + missing.is_empty(), + "authorized broken-main recovery contract is incomplete:\n {}", + missing.join("\n "), + ); +} + #[test] fn test_specs_have_required_status_marker() { // SPEC_0000 §"Required Sections": every spec must declare a parseable diff --git a/crates/rumoca/tests/sympy_template_regression.rs b/crates/rumoca/tests/sympy_template_regression.rs index 6c9172e9d..af702fc50 100644 --- a/crates/rumoca/tests/sympy_template_regression.rs +++ b/crates/rumoca/tests/sympy_template_regression.rs @@ -9,11 +9,15 @@ use tempfile::Builder; #[cfg(feature = "template-runtime-tests")] fn python_command() -> &'static str { for candidate in ["python3", "python"] { - if Command::new(candidate).arg("--version").output().is_ok() { + if Command::new(candidate) + .args(["-c", "import sympy"]) + .output() + .is_ok_and(|output| output.status.success()) + { return candidate; } } - panic!("expected python3 or python to be available for SymPy template regression tests"); + panic!("expected python3 or python with sympy installed for SymPy template regression tests"); } fn render_template(source: &str, model_name: &str, file_name: &str) -> String { diff --git a/crates/rumoca/tests/template_target_ci.rs b/crates/rumoca/tests/template_target_ci.rs index 294ced246..c96e90203 100644 --- a/crates/rumoca/tests/template_target_ci.rs +++ b/crates/rumoca/tests/template_target_ci.rs @@ -32,6 +32,24 @@ equation end Smoke; "#; +const FMI_EXTERNAL_DEPENDENCIES_MODEL: &str = "FmiExternalDependencies"; +const FMI_EXTERNAL_DEPENDENCIES_SOURCE: &str = r#" +function NativeProbe + input Real u; + output Real y; + external "C" y = native_probe(u) + annotation( + IncludeDirectory="modelica://NativeProbe/Resources/Include", + Library="native_probe"); +end NativeProbe; + +model FmiExternalDependencies + Real x(start = 1); +equation + der(x) = NativeProbe(x); +end FmiExternalDependencies; +"#; + const FMI_START_EXPRESSIONS_MODEL: &str = "FmiStartExpressions"; const FMI_START_EXPRESSIONS_SOURCE: &str = r#" model FmiStartExpressions @@ -253,6 +271,103 @@ fn galec_target_rejects_continuous_fixture_via_capability_gate() { ); } +#[test] +fn builtin_target_allow_empty_entries_are_exactly_fmi_dependency_lists() { + let mut actual = BTreeSet::new(); + for target in templates::builtin_targets() { + let manifest = parse_target_manifest(target.manifest) + .unwrap_or_else(|err| panic!("target {} manifest should parse: {err}", target.name)); + for file in manifest.files { + if file.allow_empty { + actual.insert((target.name.to_string(), file.path)); + } + } + } + let expected = [ + ("fmi2", "resources/externalLibraries.txt"), + ("fmi2", "resources/externalIncludeDirectories.txt"), + ("fmi3", "resources/externalLibraries.txt"), + ("fmi3", "resources/externalIncludeDirectories.txt"), + ] + .into_iter() + .map(|(target, path)| (target.to_string(), path.to_string())) + .collect(); + + assert_eq!(actual, expected); +} + +#[test] +fn fmi_no_external_function_allows_only_declared_dependency_lists_to_be_empty() { + let fixture = compile_fixture(SMOKE_MODEL, SMOKE_SOURCE); + let optional_paths = [ + "resources/externalLibraries.txt", + "resources/externalIncludeDirectories.txt", + ]; + for target_name in ["fmi2", "fmi3"] { + let target = templates::builtin_target(target_name).expect("built-in FMI target"); + let manifest = parse_target_manifest(target.manifest).expect("parse FMI target manifest"); + let files = render_target_files(&fixture.compiled, fixture.model_name, target_name, None) + .expect("render FMI target without external functions"); + assert_eq!(files.len(), manifest.files.len()); + + for path in optional_paths { + let declared = manifest + .files + .iter() + .find(|file| file.path == path) + .unwrap_or_else(|| panic!("{target_name} must declare {path}")); + assert!(declared.allow_empty, "{target_name}:{path} must opt in"); + let rendered = find_rendered_file(&files, path); + assert_eq!( + rendered.content.as_bytes(), + b"", + "{target_name}:{path} must be present and zero bytes" + ); + } + + for (declared, rendered) in manifest + .files + .iter() + .zip(&files) + .filter(|(file, _)| !optional_paths.contains(&file.path.as_str())) + { + assert!( + !declared.allow_empty, + "{target_name}:{} must remain required", + declared.path + ); + assert!( + !rendered.content.trim().is_empty(), + "{target_name}:{} must remain non-empty", + declared.path + ); + } + } +} + +#[test] +fn fmi_external_function_populates_optional_dependency_lists() { + let fixture = compile_fixture( + FMI_EXTERNAL_DEPENDENCIES_MODEL, + FMI_EXTERNAL_DEPENDENCIES_SOURCE, + ); + for target_name in ["fmi2", "fmi3"] { + let render = |template_name| { + let template = templates::builtin_template_source(target_name, template_name) + .expect("built-in FMI dependency template"); + fixture + .compiled + .render_template_str_with_name_and_ir(template, fixture.model_name, TemplateIr::Dae) + .expect("render FMI dependency metadata from DAE") + }; + assert_eq!(render("externalLibraries.txt.jinja"), "native_probe\n"); + assert_eq!( + render("externalIncludeDirectories.txt.jinja"), + "modelica://NativeProbe/Resources/Include\n" + ); + } +} + #[test] fn fmi2_target_model_description_serializes_start_expressions_as_literals_issue_289() { let xml = render_fmi_model_description_xml("fmi2"); @@ -528,7 +643,8 @@ fn assert_manifest_only_target(target: &templates::BuiltinTarget, manifest: &Tar /// Render every `[[files]]` entry through the real CLI path (capability /// validation, path templates, name-dispatched renderers) and assert each -/// rendered file is non-empty. +/// rendered file is non-empty unless its manifest entry explicitly allows +/// empty output. fn render_manifest_target_files( fixture: &Fixture, target: &'static templates::BuiltinTarget, @@ -547,14 +663,14 @@ fn render_manifest_target_files( "target {} rendered a different file count than its manifest declares", target.name ); - for file in &files { + for (declared, file) in manifest.files.iter().zip(&files) { assert!( !file.path.is_empty(), "target {} rendered an empty output path", target.name ); assert!( - !file.content.trim().is_empty(), + declared.allow_empty || !file.content.trim().is_empty(), "target {} rendered empty content for {}", target.name, file.path diff --git a/crates/rumoca/tests/tier_cases/tiers_a.rs b/crates/rumoca/tests/tier_cases/tiers_a.rs index 4ef298d22..06da91f2e 100644 --- a/crates/rumoca/tests/tier_cases/tiers_a.rs +++ b/crates/rumoca/tests/tier_cases/tiers_a.rs @@ -784,6 +784,46 @@ end ForRangeStep; ); } + #[test] + fn t4_05b_for_range_uses_derived_enum_parameter_bound() { + let source = r#" +type AnalogFilter = enumeration(CriticalDamping, Bessel); + +block DerivedEnumRangeFilter + parameter AnalogFilter analogFilter = AnalogFilter.CriticalDamping; + parameter Integer order = 2; + parameter Integer na = if analogFilter == AnalogFilter.CriticalDamping then 0 else integer(order / 2); + parameter Integer nr = if analogFilter == AnalogFilter.CriticalDamping then order else mod(order, 2); + Real x[order]; + Real uu[na + nr + 1]; + input Real u; + output Real y; +equation + uu[1] = u; + for i in 1:nr loop + der(x[i]) = x[i] - uu[i]; + end for; + for i in 1:nr loop + uu[i + 1] = x[i]; + end for; + y = uu[nr + na + 1]; +end DerivedEnumRangeFilter; + +model DerivedEnumRangeBound + parameter Integer order = 3; + DerivedEnumRangeFilter filter(order = order); +equation + filter.u = 1; +end DerivedEnumRangeBound; +"#; + let r = assert_compiles(source, "DerivedEnumRangeBound"); + assert_eq!( + r.f_x_count, 9, + "outer-parameter-modified derived enum parameter nr=3 should expand both for-equations through i=3" + ); + assert_eq!(r.balance, 0); + } + #[test] fn t4_06_nested_for_equation() { let source = r#" diff --git a/crates/rumoca/tests/tier_cases/tiers_b.rs b/crates/rumoca/tests/tier_cases/tiers_b.rs index d5b2adfc9..01b923b78 100644 --- a/crates/rumoca/tests/tier_cases/tiers_b.rs +++ b/crates/rumoca/tests/tier_cases/tiers_b.rs @@ -989,10 +989,10 @@ end Wrapper; r.dae.metadata.oc_break_edge_scalar_count, 1, "Should detect 1 break edge when root class and sub-component both have VCG branches" ); - // balance=2 because the wrapper has open external connectors that generate - // redundant flow equations. The VCG correctly detects 1 break edge. - // Real MSL models that use this pattern are balanced by interface flow counting. - assert_eq!(r.balance, 2); + assert_eq!( + r.balance, 0, + "open external connector flows are balanced by interface flow counting" + ); } /// Array overconstrained connectors should derive optional VCG edges per element. diff --git a/crates/rumoca/tests/tier_cases/tiers_c.rs b/crates/rumoca/tests/tier_cases/tiers_c.rs index bb1ef32a5..9b27e971e 100644 --- a/crates/rumoca/tests/tier_cases/tiers_c.rs +++ b/crates/rumoca/tests/tier_cases/tiers_c.rs @@ -826,16 +826,23 @@ end PlugToPinsNLike; r.balance, 0, "aggregate connector-field equations should balance" ); - let total_scalar_eq: usize = r + let explicit_scalar_eq: usize = r .dae .continuous .equations .iter() .map(|eq| eq.scalar_count) .sum(); + let interface_flow_needed = r.algebraics.saturating_sub(explicit_scalar_eq); + let effective_interface_flow = r + .dae + .metadata + .interface_flow_count + .min(interface_flow_needed); assert_eq!( - total_scalar_eq, 12, - "m=3 should yield 12 scalar equations (including unconnected flows)" + explicit_scalar_eq + effective_interface_flow, + 12, + "m=3 should yield 12 effective scalar equations including clamped interface-flow metadata" ); } @@ -919,6 +926,125 @@ end MiniBusTranscriptionArr; "GALEC projection should keep the redundant output-to-known bus connection skipped; origins={projection_origins:?}" ); } + + fn assert_nested_projection_flat_identity(flat: &rumoca_ir_flat::Model) { + let lane_refs = [ + "transcription.outBus.cellBus[1,1].x", + "transcription.outBus.cellBus[2,1].x", + ] + .map(|name| { + flat.variables[&rumoca_core::VarName::new(name)] + .component_ref + .as_ref() + .expect("real flattened lane must retain structured identity") + }); + assert!(lane_refs.iter().all(|reference| reference.def_id.is_some())); + assert_ne!( + lane_refs[0].def_id, lane_refs[1].def_id, + "flatten assigns a distinct resolved identity to each scalar lane instance" + ); + assert_eq!( + lane_refs[0].span, lane_refs[1].span, + "scalar lane instances must retain their common source declaration span" + ); + assert!(!lane_refs[0].span.is_dummy()); + let aggregate_ref = flat.variables + [&rumoca_core::VarName::new("transcription.outBus.cellBus.x")] + .component_ref + .as_ref() + .expect("synthetic aggregate projection must retain a structured path"); + assert_eq!(aggregate_ref.def_id, None); + } + + /// Regression pattern from a stack bus containing an array of nested + /// connectors. The aggregate field projection (`inBus.cellBus.x`) is only + /// a view over the connected scalar lanes (`inBus.cellBus[1,1].x`, ...), + /// not an unused discrete expandable member. + #[test] + fn t10k_08_nested_expandable_bus_array_projection_stays_continuous() { + let source = r#" +connector RealInput = input Real; +connector RealOutput = output Real; + +connector CellBus + Real x; +end CellBus; + +expandable connector StackBus + CellBus cellBus[2,1]; +end StackBus; + +block GainR + RealInput u; + RealOutput y; +equation + y = u; +end GainR; + +model NestedBusTranscription + StackBus inBus; + StackBus outBus; +protected + GainR g[2,1]; +equation + connect(g.u, inBus.cellBus.x); + connect(g.y, outBus.cellBus.x); +end NestedBusTranscription; + +block BusSource + StackBus bus; +equation + for i in 1:2 loop + for j in 1:1 loop + bus.cellBus[i,j].x = 1.0; + end for; + end for; +end BusSource; + +model NestedBusTranscriptionSystem + BusSource source; + NestedBusTranscription transcription; +equation + connect(source.bus, transcription.inBus); +end NestedBusTranscriptionSystem; +"#; + + let r = assert_compiles(source, "NestedBusTranscriptionSystem"); + assert_eq!(r.balance, 0); + assert_nested_projection_flat_identity(&r.flat); + assert!( + r.dae + .variables + .discrete_reals + .keys() + .all(|name| name.as_str() != "transcription.inBus.cellBus.x" + && name.as_str() != "transcription.outBus.cellBus.x"), + "connected aggregate projections must not become discrete variables" + ); + let origins = r + .dae + .continuous + .equations + .iter() + .map(|equation| equation.origin.as_str()) + .collect::>(); + assert!( + origins.iter().any(|origin| { + origin.contains("transcription.inBus.cellBus.x") + && origin.contains("transcription.g[1,1].u") + }), + "input-side scalar lane connection must remain in f_x; origins={origins:?}" + ); + assert!( + origins.iter().all(|origin| { + !(origin.contains("connection equation") + && origin.contains("transcription.outBus.cellBus") + && origin.contains("transcription.g[") + && origin.contains(".y")) + }), + "nested output lanes already defined by gain equations must not add repeated connection constraints; origins={origins:?}" + ); + } } // ============================================================================= diff --git a/crates/rumoca/tests/tiered_models.rs b/crates/rumoca/tests/tiered_models.rs index 2445a62e0..519842bbb 100644 --- a/crates/rumoca/tests/tiered_models.rs +++ b/crates/rumoca/tests/tiered_models.rs @@ -98,6 +98,7 @@ fn contains_pre_param_ref(expr: &rumoca_core::Expression) -> bool { /// Result of compiling a model with diagnostic information. #[derive(Debug)] struct CompileResult { + flat: rumoca_ir_flat::Model, dae: Dae, states: usize, algebraics: usize, @@ -128,8 +129,10 @@ fn compile(source: &str, model_name: &str) -> Result { .compile_model(model_name) .map_err(|e| format!("Instantiate/Flatten/ToDae: {:?}", e))?; + let flat = result.flat; let dae = result.dae; Ok(CompileResult { + flat, states: dae.variables.states.len(), algebraics: dae.variables.algebraics.len(), parameters: dae.variables.parameters.len(), diff --git a/crates/xtask/src/review_scan_cmd.rs b/crates/xtask/src/review_scan_cmd.rs index b15ecd52c..7db7dc7c3 100644 --- a/crates/xtask/src/review_scan_cmd.rs +++ b/crates/xtask/src/review_scan_cmd.rs @@ -401,7 +401,11 @@ fn scan_file_size(path: &str, content: &str) -> Vec { } let line_count = content.lines().count(); let (severity, rule) = if line_count > 2000 { - ("high", "file-size-hard-limit") + if has_spec_0021_file_size_exception(content) { + ("low", "file-size-exception-audit") + } else { + ("high", "file-size-hard-limit") + } } else if line_count >= 1800 { ("low", "file-size-near-limit") } else { @@ -412,10 +416,18 @@ fn scan_file_size(path: &str, content: &str) -> Vec { rule, path: path.to_string(), line: 1, - excerpt: format!("{line_count} lines"), + excerpt: if rule == "file-size-exception-audit" { + format!("{line_count} lines with explicit SPEC_0021 file-size exception and split plan") + } else { + format!("{line_count} lines") + }, }] } +fn has_spec_0021_file_size_exception(content: &str) -> bool { + content.contains("SPEC_0021") && content.contains("file-size") && content.contains("split plan") +} + fn is_line_count_checked_rust_source(path: &str) -> bool { !path.contains("/generated/") } @@ -479,6 +491,10 @@ fn line_trips_ir_behavior_boundary(path: &str, line: &str) -> bool { mod tests { use super::*; + fn over_limit_rust_source(header: &str) -> String { + format!("{header}\n{}", "fn fixture() {}\n".repeat(2001)) + } + #[test] fn scan_content_reports_review_gate_needles() { let findings = scan_content( @@ -568,6 +584,64 @@ rumoca-eval-dae = { workspace = true } assert!(generated_file.is_empty()); } + #[test] + fn scan_file_size_accepts_complete_spec_0021_exception_as_non_forbidden_audit() { + let content = over_limit_rust_source( + "// SPEC_0021 file-size exception: cohesive fixture. split plan: move fixtures by owner.", + ); + + let findings = scan_file_size("crates/example/src/lib.rs", &content); + + assert!( + findings + .iter() + .all(|finding| finding.rule != "file-size-hard-limit"), + "a complete SPEC_0021 exception must not produce a hard-limit finding: {findings:#?}" + ); + assert_eq!(forbidden_finding_count(&findings), 0); + assert!( + findings + .iter() + .any(|finding| finding.rule == "file-size-exception-audit"), + "the explicit exception should remain visible to reviewers" + ); + } + + #[test] + fn scan_file_size_requires_every_spec_0021_exception_marker() { + for header in [ + "// file-size exception: cohesive fixture. split plan: move fixtures by owner.", + "// SPEC_0021 exception: cohesive fixture. split plan: move fixtures by owner.", + "// SPEC_0021 file-size exception: cohesive fixture.", + ] { + let findings = + scan_file_size("crates/example/src/lib.rs", &over_limit_rust_source(header)); + + assert!( + findings + .iter() + .any(|finding| finding.rule == "file-size-hard-limit"), + "incomplete exception marker set must remain a hard failure: {header}" + ); + assert_eq!(forbidden_finding_count(&findings), 1); + } + } + + #[test] + fn scan_file_size_keeps_new_unexcepted_threshold_crossing_forbidden() { + let findings = scan_file_size( + "crates/example/src/new_module.rs", + &over_limit_rust_source("// New module without an exception"), + ); + + assert!( + findings + .iter() + .any(|finding| finding.rule == "file-size-hard-limit") + ); + assert_eq!(forbidden_finding_count(&findings), 1); + } + #[test] fn scan_content_reports_float_epsilon_audit_needles() { let findings = scan_content( diff --git a/crates/xtask/src/verify_cmd.rs b/crates/xtask/src/verify_cmd.rs index 8f338e77b..940a21a7a 100644 --- a/crates/xtask/src/verify_cmd.rs +++ b/crates/xtask/src/verify_cmd.rs @@ -1,3 +1,6 @@ +// SPEC_0021 file-size exception: verify command wiring still owns CLI parsing, +// gate orchestration, and harness config serialization. split plan: move MSL +// parity config serialization and GitHub baseline retrieval into submodules. use anyhow::{Context, Result, ensure}; use clap::{Args, Subcommand, ValueEnum}; use serde::Serialize; @@ -129,6 +132,15 @@ pub(crate) struct VerifyMslParityArgs { /// Total simulation memory budget in MB (caps the sim worker count) #[arg(long)] sim_total_memory_mb: Option, + /// Force regeneration of the OMC simulation reference cache + #[arg(long)] + force_omc_parity_refresh: bool, + /// OMC reference-generation worker count + #[arg(long)] + omc_parity_workers: Option, + /// Whole-stage OMC reference-generation timeout in seconds + #[arg(long)] + omc_sim_reference_batch_timeout_secs: Option, /// Explicit MSL quality baseline JSON for baseline-relative gates #[arg(long)] quality_baseline: Option, @@ -219,6 +231,15 @@ impl VerifyMslParityArgs { if let Some(value) = self.sim_total_memory_mb { config.insert("sim_total_memory_mb".into(), value.into()); } + if self.force_omc_parity_refresh { + config.insert("force_omc_parity_refresh".into(), true.into()); + } + if let Some(value) = self.omc_parity_workers { + config.insert("omc_parity_workers".into(), value.into()); + } + if let Some(value) = self.omc_sim_reference_batch_timeout_secs { + config.insert("omc_sim_reference_batch_timeout_secs".into(), value.into()); + } if let Some(value) = &self.quality_baseline { config.insert( "quality_baseline_file".into(), diff --git a/crates/xtask/src/verify_cmd/msl_quality_baseline.rs b/crates/xtask/src/verify_cmd/msl_quality_baseline.rs index f9d54220e..5e76e31ac 100644 --- a/crates/xtask/src/verify_cmd/msl_quality_baseline.rs +++ b/crates/xtask/src/verify_cmd/msl_quality_baseline.rs @@ -1,13 +1,16 @@ use anyhow::{Context, Result, ensure}; +use std::env; use std::fs; use std::io::Read; use std::path::{Path, PathBuf}; use super::VerifyMslParityArgs; -const MSL_QUALITY_BASELINE_ASSET_URL: &str = "https://github.com/CogniPilot/rumoca/releases/download/msl-quality-baseline/msl_quality_baseline.json"; +const MSL_QUALITY_BASELINE_ASSET_URL_FALLBACK: &str = "https://github.com/CogniPilot/rumoca/releases/download/msl-quality-baseline/msl_quality_baseline.json"; const MSL_QUALITY_BASELINE_FALLBACK_REL: &str = "crates/rumoca-test-msl/tests/msl_tests/msl_quality_baseline.json"; +const MSL_QUALITY_BASELINE_RELEASE_TAG: &str = "msl-quality-baseline"; +const MSL_QUALITY_BASELINE_ASSET_NAME: &str = "msl_quality_baseline.json"; pub(super) fn resolve_msl_quality_baseline( root: &Path, @@ -71,11 +74,12 @@ fn resolve_workspace_path(root: &Path, path: &Path) -> PathBuf { fn download_msl_quality_baseline_asset(root: &Path) -> Result> { let output_path = downloaded_msl_quality_baseline_path(root); + let asset_url = msl_quality_baseline_asset_url(); println!( "MSL quality baseline: downloading latest promoted asset from {}", - MSL_QUALITY_BASELINE_ASSET_URL + asset_url ); - let response = match ureq::get(MSL_QUALITY_BASELINE_ASSET_URL).call() { + let response = match ureq::get(&asset_url).call() { Ok(response) => response, Err(error) => { eprintln!( @@ -117,9 +121,46 @@ fn download_msl_quality_baseline_asset(root: &Path) -> Result> { Ok(Some(output_path)) } +fn msl_quality_baseline_asset_url() -> String { + match current_github_repo_url() { + Some(repo_url) => format!( + "{repo_url}/releases/download/{MSL_QUALITY_BASELINE_RELEASE_TAG}/{MSL_QUALITY_BASELINE_ASSET_NAME}" + ), + None => MSL_QUALITY_BASELINE_ASSET_URL_FALLBACK.to_string(), + } +} + +fn current_github_repo_url() -> Option { + current_github_repo_url_from( + env::var("GITHUB_SERVER_URL").ok().as_deref(), + env::var("GITHUB_REPOSITORY").ok().as_deref(), + ) +} + +fn current_github_repo_url_from( + server_url: Option<&str>, + repository: Option<&str>, +) -> Option { + let repository = repository?.trim(); + if repository.is_empty() { + return None; + } + + let server_url = server_url + .map(str::trim) + .filter(|url| !url.is_empty()) + .unwrap_or("https://github.com"); + Some(format!( + "{}/{}", + server_url.trim_end_matches('/'), + repository.trim_matches('/') + )) +} + #[cfg(test)] mod tests { use super::super::VerifyMslParityArgs; + use super::*; use std::path::PathBuf; #[test] @@ -149,4 +190,32 @@ mod tests { }; assert!(!short_run.uses_baseline_relative_quality_gate()); } + + #[test] + fn current_github_repo_url_uses_actions_repository() { + assert_eq!( + current_github_repo_url_from(Some("https://github.com/"), Some("climamind/rumoca")), + Some("https://github.com/climamind/rumoca".to_string()) + ); + } + + #[test] + fn current_github_repo_url_defaults_to_github_server_url() { + assert_eq!( + current_github_repo_url_from(None, Some("climamind/rumoca")), + Some("https://github.com/climamind/rumoca".to_string()) + ); + } + + #[test] + fn current_github_repo_url_ignores_missing_repository() { + assert_eq!( + current_github_repo_url_from(Some("https://github.com"), None), + None + ); + assert_eq!( + current_github_repo_url_from(Some("https://github.com"), Some("")), + None + ); + } } diff --git a/crates/xtask/src/vscode_cmd.rs b/crates/xtask/src/vscode_cmd.rs index 3a773fa7a..703190ea3 100644 --- a/crates/xtask/src/vscode_cmd.rs +++ b/crates/xtask/src/vscode_cmd.rs @@ -1496,12 +1496,11 @@ fn cargo_target_cc_env_suffix(target: &str) -> String { mod tests { use super::{ VscodeMslSmokeSummary, VscodeNpmDependencyMode, VscodeNpmInstallPlan, VscodePackageTarget, - VscodeSmokeEnvironment, VscodeSmokeLaunchMode, VscodeSmokeOptions, - cargo_target_cc_env_suffix, cargo_target_linker_env_suffix, - mirror_cached_vscode_smoke_install, prepare_install_check_workspace, replace_staged_binary, - resolve_install_check_document, resolve_install_check_profile_root, - resolve_vscode_npm_install_plan, resolve_workspace_dir, select_vscode_smoke_launch_mode, - should_copy_vscode_smoke_root_entry, should_install_vscode_smoke_prereqs, + VscodeSmokeEnvironment, VscodeSmokeLaunchMode, cargo_target_cc_env_suffix, + cargo_target_linker_env_suffix, mirror_cached_vscode_smoke_install, + prepare_install_check_workspace, replace_staged_binary, resolve_install_check_document, + resolve_install_check_profile_root, resolve_vscode_npm_install_plan, resolve_workspace_dir, + select_vscode_smoke_launch_mode, should_copy_vscode_smoke_root_entry, should_retry_vscode_npm_ci_after_clean, stage_vscode_smoke_workspace, }; use anyhow::anyhow; @@ -1847,6 +1846,8 @@ mod tests { #[cfg(target_os = "linux")] #[test] fn install_prereqs_flag_only_applies_to_linux_missing_tools() { + use super::{VscodeSmokeOptions, should_install_vscode_smoke_prereqs}; + assert!(should_install_vscode_smoke_prereqs( smoke_environment(true, true, true, false), VscodeSmokeOptions { diff --git a/docs/OPEN_SOURCE.md b/docs/OPEN_SOURCE.md new file mode 100644 index 000000000..003db6f4d --- /dev/null +++ b/docs/OPEN_SOURCE.md @@ -0,0 +1,70 @@ +# Open Source Readiness + +This checklist defines what must be true before making the ClimaMind-maintained +Rumoca repository public or publishing artifacts from it. + +## Repository Boundary + +Publish only the Rumoca compiler, tooling, editor, binding, and packaging +sources tracked by git. Keep the ClimaMind product layer outside this +repository: + +- no customer, building, site, or calibration data; +- no Kelvin runtime artifacts, training outputs, or top-down experiment results; +- no cloud account identifiers, private buckets, tokens, API keys, or local + machine credentials; +- no generated `target/`, downloaded MSL trees, local caches, release staging + directories, or one-off probe outputs. + +## Required Preflight + +Run these checks from the repository root before flipping repository visibility: + +```bash +git status --short --branch +git ls-files | rg -i '(^|/)(\.env|.*\.pem|.*\.key|.*\.p12|.*\.sqlite|.*\.db|.*\.zip|.*\.tar|.*\.gz|.*\.parquet|.*\.csv|.*\.jsonl)$|secret|credential|token' +git log --all --name-only --pretty=format: | sort -u | rg -i '(^|/)(\.env|.*\.pem|.*\.key|.*\.p12|.*\.sqlite|.*\.db|.*\.zip|.*\.tar|.*\.gz|.*\.parquet|.*\.csv|.*\.jsonl)$|secret|credential|token' +cargo metadata --format-version 1 > target/open-source-cargo-metadata.json +jq -r '.packages[] | select(.source != null) | [.name, .version, (.license // "NO_LICENSE")] | @tsv' target/open-source-cargo-metadata.json | sort -u +``` + +Treat any secret, credential, private artifact, or unexplained non-source file +as a blocker. If a blocker appears only in git history, rewrite or recreate the +public repository instead of publishing that history. + +## License Position + +Rumoca is Apache-2.0. Preserve: + +- `LICENSE`; +- `NOTICE`; +- license fields in Cargo, Python, and VS Code package metadata; +- third-party notices required by bundled binary/editor artifacts. + +The observed Cargo graph is mostly permissive. Known licenses that require +release attention are MPL-2.0 for `option-ext`, Unicode-3.0 packages, and +dual-licensed packages where the permissive option should be selected. npm +packages must be checked separately before VS Code or web-editor distribution. + +## Release Surfaces + +Source publication is lower risk than binary publication. Before publishing +GitHub Releases, PyPI wheels, VS Code `.vsix` files, Docker images, or GitHub +Pages assets, verify each artifact independently: + +- the artifact contains `LICENSE` and required notices; +- package metadata points at the public ClimaMind repository; +- generated files do not embed local absolute paths except in tests or fixtures; +- release workflows publish under the intended GitHub organization and package + namespace; +- the `ghcr.io/climamind/rumoca-dev:main` image exists before requiring CI jobs + that use it as a container; +- MSL data is downloaded from the official release during CI and not vendored + into the repository. + +## Upstream Attribution + +This repository should state that Rumoca originated in the CogniPilot community +and that ClimaMind maintains this public line. Keep attribution factual; do not +imply that ClimaMind owns upstream trademarks, third-party packages, or Modelica +Association materials. diff --git a/docs/dev-guide/book.toml b/docs/dev-guide/book.toml index 3df321773..5b2626911 100644 --- a/docs/dev-guide/book.toml +++ b/docs/dev-guide/book.toml @@ -8,8 +8,8 @@ src = "src" build-dir = "book" [output.html] -git-repository-url = "https://github.com/CogniPilot/rumoca" -edit-url-template = "https://github.com/CogniPilot/rumoca/edit/main/docs/dev-guide/{path}" +git-repository-url = "https://github.com/climamind/rumoca" +edit-url-template = "https://github.com/climamind/rumoca/edit/main/docs/dev-guide/{path}" default-theme = "rust" preferred-dark-theme = "ayu" # Live Modelica example runner, shared with the user guide (single source). diff --git a/docs/dev-guide/src/contributing/specs-process.md b/docs/dev-guide/src/contributing/specs-process.md index 5a77897f3..d7eb39506 100644 --- a/docs/dev-guide/src/contributing/specs-process.md +++ b/docs/dev-guide/src/contributing/specs-process.md @@ -17,7 +17,7 @@ it contains no rules itself, and neither does this book. | Diagnostics, spans, error codes, tracing | [SPEC_0008](https://github.com/CogniPilot/rumoca/blob/main/spec/SPEC_0008_PHASE_ERRORS.md) | | Tool config (`rumoca-tool-*`, env-var policy) | [SPEC_0018](https://github.com/CogniPilot/rumoca/blob/main/spec/SPEC_0018_TOOL_CONFIG.md) | | Function length, nesting, file size, determinism | [SPEC_0021](https://github.com/CogniPilot/rumoca/blob/main/spec/SPEC_0021_CODE_COMPLEXITY.md) | -| Development workflow, bug triage, root-cause proof | [SPEC_0032](https://github.com/CogniPilot/rumoca/blob/main/spec/SPEC_0032_DEVELOPMENT_PROCESS.md) | +| Development workflow, bug triage, root-cause proof | [SPEC_0033](https://github.com/CogniPilot/rumoca/blob/main/spec/SPEC_0033_DEVELOPMENT_PROCESS.md) | | Opening a PR | [SPEC_0025](https://github.com/CogniPilot/rumoca/blob/main/spec/SPEC_0025_PR_REVIEW_PROCESS.md) | | Scope/philosophy questions ("should this live in the compiler?") | [SPEC_0031](https://github.com/CogniPilot/rumoca/blob/main/spec/SPEC_0031_COMPILER_PHILOSOPHY.md) | diff --git a/docs/dev-guide/src/tooling/msl-quality-gate.md b/docs/dev-guide/src/tooling/msl-quality-gate.md index d22d93373..663c15d6d 100644 --- a/docs/dev-guide/src/tooling/msl-quality-gate.md +++ b/docs/dev-guide/src/tooling/msl-quality-gate.md @@ -41,9 +41,57 @@ The stage checks are cumulative over the fixed root-example denominator: parse/IR-AST, flatten/IR-flat, DAE/IR-DAE, solve/IR-Solve, initial-condition solve, and simulation. Increasing an early-stage pass count is always treated as an improvement; the gate fails when any cumulative stage -count drops below the committed baseline for the same target set. - -`msl_quality_current.json` also records release review metadata: +count drops below the resolved baseline for the same target set. + +The root-example/full quality gate is fail-closed for OMC parity. OMC must be +available during parity preparation, and the current +`omc_simulation_reference.json` must match the active target count and contain +non-empty system and wall runtime samples plus at least one comparable trace +with model bucket percentages. Missing OMC, a missing or stale reference, or a +reference with no comparable runtime/trace metrics fails the run with an +actionable error; compile and simulation counts alone are not a passing full +quality gate. + +Runtime system-time speedup uses the committed baseline median and is always a +blocking gate when it regresses by more than 35%. Wall-time uses the same 35% +floor only when its measurement provenance is trustworthy. The +`runtime_comparison.wall_time_provenance` object records: + +- fresh and cached OMC sample counts; +- requested, successfully applied, and failed CPU-affinity worker counts, plus + the independently observed Rumoca scheduler worker count; +- normalized one-minute host load sampled before and after parity work; and +- the worker count and OMC thread count paired with `timing` runtime context. + +A wall-time comparison is trustworthy only when it contains fresh samples and +no cached samples, the fresh plus cached counts exactly cover the compared wall +samples, affinity was requested and applied successfully for every Rumoca +scheduler worker, and +both normalized load samples are present and at most `1.5`. The current runtime +context and provenance worker/thread policy must also both exactly match the +promoted baseline runtime context; a baseline without a complete runtime +context has no policy comparator and is advisory. Missing or malformed +provenance makes wall-time advisory; it does not make missing or malformed +parity, runtime-ratio, or trace data acceptable. + +The official shard fan-in sums sample, Rumoca-worker, and affinity counts. It +retains OMC worker/thread values only when every shard reports the same value, +and retains each normalized load value only when every shard reports a finite +sample, using the maximum across shards. Missing or mismatched shard provenance +therefore remains advisory rather than being reconstructed as trusted. + +The console reports `MSL wall speed gate: PASS` when trusted wall-time is above +the floor, `FAIL` when trusted wall-time regresses past it, and `ADVISORY` when +the measurement is not trusted. Advisory output still shows the observed +median, baseline, 35% floor, and every provenance reason. Correctness and +system-time failures remain blocking regardless of wall-time status. + +`msl_quality_current.json` records `wall_time_provenance` and a +`runtime_wall_decision` object containing `status`, `trusted`, `reasons`, and, +when a runtime baseline exists, the observed median, baseline median, and 35% +floor. These audit fields belong only to the current snapshot and are not part +of the promoted baseline comparison schema. It also records release review +metadata: - `omc_version` records the OpenModelica build used for OMC trace parity; the quality gate compares the upstream release version and tolerates distro @@ -73,7 +121,9 @@ Promotion requires a full-run snapshot with non-empty `omc_version` metadata. Focused debugging runs can use `RUMOCA_MSL_SIM_MATCH`, `RUMOCA_MSL_SIM_LIMIT`, `RUMOCA_MSL_SIM_TARGETS_FILE`, or `RUMOCA_MSL_TARGET_SCOPE=committed-targets`, but those runs are not baseline -updates. +updates. These explicit focused/partial modes skip required OMC parity and say +so in their output; they do not report or imply that the full quality gate +passed. For commit-to-commit regression diffs, run both worktrees with the same focused target JSON, then generate machine-readable buckets with: diff --git a/docs/superpowers/completed/2026-07-12-fix-ci-lockfile-merge.md b/docs/superpowers/completed/2026-07-12-fix-ci-lockfile-merge.md new file mode 100644 index 000000000..587293244 --- /dev/null +++ b/docs/superpowers/completed/2026-07-12-fix-ci-lockfile-merge.md @@ -0,0 +1,30 @@ +# CI Lockfile Merge Repair — Completed + +**Status:** completed on 2026-07-12. This is a provenance record, not an active implementation plan. + +## Outcome + +The upstream merge repair was closed without weakening CI or MSL quality policy: + +- regenerated the merged Cargo lock graph and restored locked parsing; +- reconciled the structural, DAE, Solve, compiler, CLI, codegen, simulation-session, and MSL parity merge artifacts; +- removed orphan DAE and diffsol helpers left by the merge; +- restored the repository lint gate and the task-scoped tests. + +The old plan's unchecked Task 9 Step 4 was stale bookkeeping: the following final step and the final gate record were already marked complete in `a4ea7868`. + +## Commit provenance + +- `72585949` — regenerate merged Cargo lockfile +- `c0d6157a` — reconcile upstream merge artifacts +- `ae06da6d` — reconcile DAE merge artifacts +- `e4912bd3` — reconcile Solve merge artifacts +- `ea3f1ba1` — register the codegen JSON filter +- `f2d2026b` — remove orphan DAE test helper +- `96fdc1da` — remove orphan diffsol helper +- `25869d1d` — restore discrete simulation sessions +- `1ba5faab` — reconcile compiler and CLI merge APIs +- `21a560bd` — reconcile MSL parity configuration +- `a4ea7868` — record the completed final CI gates + +All listed implementation commits carry a DCO `Signed-off-by` trailer. The detailed step-by-step shell transcript was intentionally removed from the active plan area; Git history remains the authoritative audit trail. diff --git a/docs/superpowers/completed/2026-07-12-fix-msl-gate-regressions.md b/docs/superpowers/completed/2026-07-12-fix-msl-gate-regressions.md new file mode 100644 index 000000000..6ade2e2bc --- /dev/null +++ b/docs/superpowers/completed/2026-07-12-fix-msl-gate-regressions.md @@ -0,0 +1,74 @@ +# MSL Gate Repairs — Completed Work + +**Status:** the items below are complete. This is a concise provenance record, not an active problem list. + +## Original semantic repairs + +- connector instance identity: `94f6033a`; +- near-future root tolerance: `bfd070fe`; +- periodic Clock alias ticks with builtin identity proof: `8fbc5035`, `6ff3c7cf`; +- algebraic refresh alternatives and complete projection validation: `8d402084`, `66ba724d`; +- record-constructor projections: `30a19e77`; +- scheduled-root invalidation: `d9894cb5`; +- projected result cardinality and call dependencies: `8e709eed`, `176e5fba`. + +## Architecture and follow-up repairs + +- SPEC_0021 runtime-test split: `8fb3e5ae`; +- review-scan parity with file-size exceptions: `6cde10b7`; +- sequential instance algorithm bindings: `ebb75ad4`; +- shaped discrete enum-array starts: `678a15a2`; +- reference-aware indexed aggregate identity: `7c0355dd`. + +## Gate finalization repairs + +- fail fast when a warm OpenModelica process exits: `974ae96d`; +- install before verification and pin exact OpenModelica `1.27.0-1`: `b11f0615`; +- require full-run runtime/trace parity and close the selected-target bypass: `82fc87ae`, `438a3947`; +- align selected-target fixtures and isolate full-gate defaults: `427b204e`, `9cd395d9`; +- settle symbolic dimensions without redundant evaluation while preserving every dependency-driven retry: `77ad16a0`, `c30ffec5`, `8e2c6f56`, `e72e724b`, `cf1c0af6`. +- ModelicaTest semantic gate at `9e8f2e09`: selected target `ModelicaTest.Blocks.MuxDemux` completed successfully. + +The symbolic-dimension implementation passed independent review with no findings. At clean commit `cf1c0af67d00108e59329b16def8fa8044bb1e1a`, 540 `rumoca-phase-flatten` library tests, strict crate clippy, formatting, and diff checks passed. The clean focused artifact `/tmp/rumoca-task15-dimension-cf1/` records 4/4 Flat success with no timeout or budget increase. + +## Stack-family structural repair + +The former Stack/StackRC deficit was closed at its producing mechanisms rather than compensated downstream: + +- Flat preserves multidimensional member bindings and indexed nested-component boundaries/provenance: `aec66b13`, `bd7b9541`, `22e16326`, `eec5fbdb`, `abed9d04`; +- inherited references resolve through explicit record aliases and fail closed when instance ownership is not proven; the rejected path-prefix heuristic is not present: `5fb3a38e`, `ddff7f88`; +- ToDae canonicalizes and validates expandable-record projection identity before admitting it: `1071d894`, `a31e3be4`, `3711b4c1`; +- structural elimination preserves runtime references, retires invalid boundary targets, rewrites proven indexed substitutions, and rejects ambiguous or symbolic target identity: `c558d2a6`, `dc8fcb4f`, `d5bf4868`, `bdc53572`, `4b92974d`, `9e8f2e09`. + +The complete mechanism chain passed independent review. At clean semantic commit `9e8f2e09a31e4e8f8834f121db3825ec883318c3`, focused structural admission is balanced for all controls: `CCCV_Stack` `134/134`, `CCCV_StackRC` `176/176`, `CCCV_Cell` `19/19`, and `CCCV_CellRC` `20/20`. + +## Finalization and handoff + +- `f84f3cb3` split oversized regression modules without changing behavior; +- `3ba4b26d` gave the Clock fixture compiler-owned builtin `DefId` identity; +- `21ccab5d` extracted the structural runtime-partition test helper; +- `6a5be4b9` replaced hidden fixture sentinels with deterministic spans and an explicit fixed marker. + +The complete documentation gate, formatting, strict workspace all-target/all-feature clippy, `cargo test --workspace`, and review scan all passed; review scan reported zero forbidden findings. DCO and clean-worktree checks passed. The final independent review's only finding was stale completed work in the active plan, which this closeout removes. + +## Full promoted-ratchet closure + +The remaining full-gate regressions were closed at their owning mechanisms: + +- `046a2691` stops cache-only symbolic-dimension repairs from restarting the Flat fixpoint; `HeatingSystem` now reaches ToDae in 2.74 seconds instead of exhausting the unchanged 45-second Flat budget; +- `b35191ce` preserves a surviving physical producer when derivative-dependent aggregate leaves are structurally eliminated; +- `63c1d7ee` assigns residual ownership from complete producer incidence and keeps coupled rows simultaneous; +- `e7d8de30` projects coupled algebraics from the accepted local branch; +- `58be4144` refreshes direct and coupled algebraic observations, consumes root relations atomically, and performs branch-preserving consistency polish before publishing reconstructed connection flows; +- `6c546777` preserves accepted derivative/root/observation seeds across BDF restarts, clamps scheduled left limits, retains the triggering root index, and defers deadline roots consistently in session mode; +- `ebf3518a` closes the final review findings: converged Newton polish is trust-bounded to the accepted branch, and coincident relation memories update atomically before confirmed root overrides win; +- `b173267c` splits runtime values and branch-projection regressions into owned modules so every affected Rust source remains below the SPEC_0021 hard limit; +- `5bfcebc3` makes the no-causal-target projection path explicit instead of relying on a default empty vector. + +The unchanged `cargo xtask verify msl-parity` gate passed at `b173267c` over all 566 MSL 4.1.0 root examples. The final snapshot records Parse `566`, Flat `565`, DAE `556`, Solve `407`, balanced `543`, initialization success `222`, simulation success `180`, compared traces `172`, high-agreement traces `124`, minor-deviation traces `27`, no-severe traces `149`, severe traces `23`, severe channels `162`, and exact state-set matches `156`. The resolved promoted baseline was not changed; no timeout, selected target, trace exclusion, tolerance, or parity requirement was relaxed. + +After the full gate, the explicit empty-projection refactor passed all `943` `rumoca-phase-solve` tests, the strict workspace lint gate, formatting, diff checks, DCO checks, and a high/forbidden-failing review scan of that final commit. The whole-stack review audit's high-severity matches were limited to existing test-fixture dummy spans/defaults plus two new `#[cfg(test)]` fixture spans; the only new production-code match was the implicit default removed by `5bfcebc3`. + +No branch push was performed because the user did not authorize one. This is handoff state, not an active problem. + +Detailed implementation scripts, stale worktree paths, obsolete pre-fix counts, and already completed checklists were removed. Git history and the cited clean artifacts are the authoritative audit trail. diff --git a/docs/superpowers/plans/2026-07-14-cli-52-upstream-equivalent-cleanup.md b/docs/superpowers/plans/2026-07-14-cli-52-upstream-equivalent-cleanup.md new file mode 100644 index 000000000..2c0cde8e7 --- /dev/null +++ b/docs/superpowers/plans/2026-07-14-cli-52-upstream-equivalent-cleanup.md @@ -0,0 +1,135 @@ +# CLI-52 Upstream-Equivalent Cleanup Implementation Plan + +> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking. + +**Goal:** Remove the remaining ClimaMind-owned pure-discrete session implementation and make the unified simulation facade use the upstream solver backends' no-state sessions. + +**Architecture:** Keep `SimulationSession` as the upstream-owned facade and delegate every session to the selected Diffsol or RK45 backend. Delete the facade-level `discrete_stepper` implementation, which duplicates and bypasses the canonical backend no-state event orchestration. Preserve all later ClimaMind compiler/runtime deltas that are unrelated to the 28 upstream-equivalent source commits. + +**Tech Stack:** Rust workspace, Cargo tests, Rumoca Solve IR/session APIs. + +## Global Constraints + +- Latest verified refs are `origin/main=17b13432ff6c86c1df6089478a74891f049283d3` and `upstream/main=e6884d035f700955e4e2cccaaa6230ca79f21e20` after a fresh fetch on 2026-07-14. +- Preserve the public Git DAG and the Linear provenance document; do not rebase/rewrite the 411 non-merge commits and do not revert the 28 historical source commits. +- Do not merge the five unrelated commits after merge-base `24209c803442ce0402a2fe47c3533280758c88d3` as part of CLI-52; `git merge-tree --write-tree origin/main upstream/main` reports broad unrelated conflicts. +- Upstream implementations are authoritative for all 28 `upstream-equivalent` items; no old/new dual path may remain. +- Preserve non-equivalent later ClimaMind deltas; do not replace whole current files from `upstream/main`. +- Do not modify CogniPilot upstream and do not create an external PR. +- Follow TDD: the deadline-sensitive no-state session regression must fail against the old facade stepper before production code changes. +- Final commits must include `Signed-off-by` and must not include AI `Co-Authored-By` trailers. + +--- + +### Task 1: Delegate zero-state sessions to upstream solver backends + +**Files:** +- Modify: `crates/rumoca-bind-wasm/src/tests.rs` +- Modify: `crates/rumoca-sim/src/simulation_session.rs` +- Modify: `crates/rumoca-sim/src/lib.rs` +- Delete: `crates/rumoca-sim/src/discrete_stepper.rs` +- Modify: `crates/rumoca-solver/src/runtime/pre_params.rs` +- Modify: `crates/rumoca-solver/src/lib.rs` +- Modify: `crates/rumoca-solver-rk45/src/no_state.rs` +- Modify: `crates/rumoca-solver-rk45/src/lib.rs` +- Modify: `crates/rumoca-solver-rk45/src/tests.rs` +- Modify: `crates/rumoca-solver-diffsol/src/tests/mod.rs` + +**Interfaces:** +- Consumes: `rumoca_solver_diffsol::SimulationSession::from_solve_model` and `rumoca_solver_rk45::SimulationSession::from_solve_model`, both of which own zero-state behavior. +- Produces: unchanged public `rumoca_sim::SimulationSession` methods (`new_with_diagnostics`, `set_input`, `reset`, `advance_to`, `step`, `time`, `get`, `state`, `input_names`, `variable_names`). + +- [x] **Step 1: Write the failing deadline-sensitive regression** + +Extend `test_interactive_session_runs_pure_discrete_model_with_guarded_dynamic_subscript` so the controller updates inside `when sample(0.02, 0.02) then ... end when`. Advance first to `0.01` and assert `y` and `k` have not changed; then advance to `0.02` and assert the first sample update; advance to `0.04` with a changed input and assert the second update uses `pre(k)` and the runtime-indexed table. + +The behavioral shape must be: + +```rust +session.set_input("u", 1.5).expect("set input u"); +session.advance_to(0.01).expect("advance before first sample"); +assert_eq!(session.get("y").expect("read y"), Some(0.0)); + +session.advance_to(0.02).expect("advance to first sample"); +assert_eq!(session.get("y").expect("read y"), Some(1.5)); + +session.set_input("u", 2.0).expect("set input u"); +session.advance_to(0.04).expect("advance to second sample"); +assert_eq!(session.get("y").expect("read y"), Some(4.0)); +``` + +- [x] **Step 2: Run the regression and verify RED** + +Run: + +```bash +cargo test -p rumoca-bind-wasm --features sim-rk45 test_interactive_session_runs_pure_discrete_model_with_guarded_dynamic_subscript -- --nocapture +``` + +Expected: FAIL at the first scheduled deadline (`0.02`): the facade-owned `discrete_stepper` has no scheduled-event orchestration. Without `--features sim-rk45`, Cargo silently filters out this feature-gated test and reports `running 0 tests`. + +Observed RED: the assertions at `0.01` passed and the first scheduled update at `0.02` failed (`y = 0.0`, expected `1.5`). After removing the facade stepper, the same regression exposed a backend-owned RK45 no-state bug: the first event updated correctly, but the second event at `0.04` remained at `1.5` instead of `4.0`. + +- [x] **Step 3: Remove the duplicate implementation** + +In `simulation_session.rs`, remove: + +```rust +SimulationSessionInner::Discrete(...) +``` + +and every corresponding match arm, remove `discrete_session_from_solve_model`, and remove both `state_scalar_count() == 0` early returns. Keep the existing Diffsol and RK45 construction paths unchanged so their canonical no-state sessions own zero-state behavior. + +In `lib.rs`, remove: + +```rust +mod discrete_stepper; +``` + +Delete `crates/rumoca-sim/src/discrete_stepper.rs` entirely. + +The delegated RK45 path also needs the same scheduled-root lifecycle already used by its stateful backend. Its generic root search must ignore roots owned by `scheduled_root_conditions`; at an exact scheduled `EventEntry`, it must clear the relation memory before the event snapshot and seed the projected update with the scheduled root override. Shared at-time helpers in `rumoca-solver` keep the stateful and no-state paths on one mechanism. The ordinary left-limit refresh remains unchanged for non-scheduled caller deadlines. + +- [x] **Step 4: Verify GREEN and focused backend behavior** + +Run: + +```bash +cargo test -p rumoca-bind-wasm --features sim-rk45 test_interactive_session_runs_pure_discrete_model_with_guarded_dynamic_subscript -- --nocapture +cargo test -p rumoca-solver-rk45 rk45_session_runs_no_state_discrete_controller -- --nocapture +cargo test -p rumoca-solver-rk45 rk45_no_state_session_rearms_periodic_sample_edges -- --nocapture +cargo test -p rumoca-bind-wasm --features sim-diffsol test_interactive_session_runs_pure_discrete_model_with_guarded_dynamic_subscript -- --nocapture +cargo test -p rumoca-solver-diffsol no_state -- --nocapture +cargo test -p rumoca-sim --all-features +``` + +Expected: PASS; no zero-state facade variant remains. + +- [x] **Step 5: Run the five upstream-equivalent focused groups** + +Run: + +```bash +cargo test -p rumoca-phase-solve discrete_history_update_from_input_keeps_state_as_target +cargo test -p rumoca-phase-solve lower_discrete_rhs_keeps_pre_selector_branch_as_runtime_select +cargo test -p rumoca-phase-solve lower_expression_dynamic_param_subscript_emits_indexed_load +cargo test -p rumoca-phase-codegen test_embedded_c_templates_render_solve_ir +cargo test -p rumoca --features template-runtime-tests --test backend_template_runtime_regression embedded_c_ -- --nocapture +cargo test -p rumoca-opt +cargo test -p rumoca --test neural_ode_tensor_solve_ir +cargo test -p rumoca-sim --all-features parses_multirate_lockstep_config +cargo test -p rumoca-sim --all-features rejects_invalid_multirate_lockstep_rates +``` + +Expected: PASS with no duplicate-path restoration. + +The two `rumoca-sim` name-filter commands include `--all-features` so they select their feature-gated tests and each run one test. The full `rumoca-sim --all-features` gate runs all 85 tests. + +- [x] **Step 6: Commit the implementation** + +```bash +git add docs/superpowers/plans/2026-07-14-cli-52-upstream-equivalent-cleanup.md crates/rumoca-bind-wasm/src/tests.rs crates/rumoca-sim/src/simulation_session.rs crates/rumoca-sim/src/lib.rs crates/rumoca-sim/src/discrete_stepper.rs crates/rumoca-solver/src/runtime/pre_params.rs crates/rumoca-solver/src/lib.rs crates/rumoca-solver-rk45/src/no_state.rs crates/rumoca-solver-rk45/src/lib.rs crates/rumoca-solver-rk45/src/tests.rs crates/rumoca-solver-diffsol/src/tests/mod.rs +git commit -s -m "refactor(sim): adopt upstream no-state sessions" +``` + +Expected: one signed-off implementation commit, net-negative production code. diff --git a/docs/superpowers/plans/2026-07-14-msl-wall-time-trust-gate.md b/docs/superpowers/plans/2026-07-14-msl-wall-time-trust-gate.md new file mode 100644 index 000000000..750a9e319 --- /dev/null +++ b/docs/superpowers/plans/2026-07-14-msl-wall-time-trust-gate.md @@ -0,0 +1,445 @@ +# MSL Wall-Time Trust Gate Implementation Plan + +> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking. + +**Goal:** Keep MSL wall-time regression blocking only for fresh, affinity-correct, healthy-host measurements while preserving every existing correctness and system-time gate. + +**Architecture:** Extend the worker handshake so scheduler metrics describe actual affinity, add a small MSL measurement-provenance module that records cache and host-health evidence, then make the quality gate classify wall-time as `PASS`, `FAIL`, or `ADVISORY`. The runtime tolerance remains unchanged; only the trust qualification changes whether a wall regression is blocking. + +**Tech Stack:** Rust 2024, serde/serde_json, Cargo integration tests, Rumoca model-worker protocol, MSL parity harness. + +## Global Constraints + +- Keep the existing runtime regression tolerance exactly `0.35`. +- System-time runtime regression remains blocking on every complete parity run. +- Stage-count, balance, simulation, trace, OMC availability, runtime-sample presence, and target-set checks remain fail-closed. +- Wall time is blocking only when every compared OMC runtime sample is fresh, requested Rumoca worker affinity succeeded, normalized one-minute load is present and `<= 1.5` before and after measured work, and runtime worker/thread context matches policy. +- Missing or malformed wall-time trust provenance makes only the wall regression advisory; it must never make missing parity data valid. +- Do not change or promote `msl_quality_baseline.json`. +- Do not add retry loops, bypass environment variables, fitted constants, or generated-artifact patches. +- Update `SPEC_0025_PR_REVIEW_PROCESS.md` and `docs/dev-guide/src/tooling/msl-quality-gate.md` with the final policy. +- Commit and push only to the ClimaMind `origin` feature branch `cli-52-upstream-equivalent-cleanup`; never push or open a PR against `upstream`. + +--- + +### Task 1: Report Actual Worker Affinity + +**Files:** +- Modify: `crates/rumoca-worker/src/lib.rs` +- Modify: `crates/rumoca-worker/src/bin/rumoca-worker.rs` +- Modify: `crates/rumoca-test-msl/tests/balance_pipeline/mod.rs` +- Modify: `crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_core.rs` +- Modify: `crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_core/streaming_workers.rs` +- Test: `crates/rumoca-worker/src/lib.rs` +- Test: `crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_core/tests.rs` + +**Interfaces:** +- Produces: `ModelWorkerControlMessage::Ready { protocol_version, cpu_affinity_applied: Option }`. +- Produces: `ModelWorkerDaemon::cpu_affinity_applied(&self) -> Option`. +- Produces: scheduler counts `affinity_requested_worker_count`, `affinity_applied_worker_count`, and `affinity_failed_worker_count`; `pinned_worker_count` remains serialized for compatibility but equals actual affinity successes. +- Consumes: existing `pin_current_thread_to_cpu_core` and `cpu_core_plan`. + +- [ ] **Step 1: Write failing worker-protocol tests** + +Add tests that serialize and deserialize both requested-affinity success and no-request cases: + +```rust +#[test] +fn ready_message_preserves_affinity_result() { + let ready = ModelWorkerControlMessage::Ready { + protocol_version: MODEL_WORKER_PROTOCOL_VERSION, + cpu_affinity_applied: Some(false), + }; + let encoded = serde_json::to_string(&ready).unwrap(); + let decoded: ModelWorkerControlMessage = serde_json::from_str(&encoded).unwrap(); + assert!(matches!( + decoded, + ModelWorkerControlMessage::Ready { + cpu_affinity_applied: Some(false), + .. + } + )); +} + +#[test] +fn ready_message_allows_unrequested_affinity() { + let ready = ModelWorkerControlMessage::Ready { + protocol_version: MODEL_WORKER_PROTOCOL_VERSION, + cpu_affinity_applied: None, + }; + assert!(serde_json::to_string(&ready).unwrap().contains("cpu_affinity_applied")); +} +``` + +- [ ] **Step 2: Run RED test** + +Run: + +```bash +CARGO_TARGET_DIR=/Users/hechuan/workspace/climamind/rumoca/target \ + cargo test -p rumoca-worker ready_message_ -- --nocapture +``` + +Expected: compilation fails because `Ready` has no `cpu_affinity_applied` field. + +- [ ] **Step 3: Implement the handshake and daemon accessor** + +Compute affinity once in `run_worker_entry`: + +```rust +let cpu_affinity_applied = args.cpu_core_id.map(|cpu_core_id| { + pin_current_thread_to_cpu_core(cpu_core_id) + .inspect_err(|error| eprintln!("warning: {error}; continuing without CPU pinning")) + .is_ok() +}); +``` + +Pass the value to `run_worker_daemon`, include it in `Ready`, retain it in +`ModelWorkerDaemon::wait_for_ready`, and expose it through the accessor. Keep +`None` for workers for which affinity was not requested. + +- [ ] **Step 4: Write failing scheduler aggregation tests** + +Add a pure aggregation helper and tests: + +```rust +#[test] +fn affinity_counts_distinguish_requested_success_and_failure() { + let counts = affinity_counts([Some(true), Some(false), None]); + assert_eq!(counts.requested, 2); + assert_eq!(counts.applied, 1); + assert_eq!(counts.failed, 1); +} +``` + +Run: + +```bash +CARGO_TARGET_DIR=/Users/hechuan/workspace/climamind/rumoca/target \ + cargo test -p rumoca-test-msl --features msl-full-test --test msl_tests \ + affinity_counts_ -- --nocapture +``` + +Expected: FAIL because the aggregation helper and fields do not exist. + +- [ ] **Step 5: Aggregate actual results** + +Record the daemon accessor value once per spawned worker. Extend +`SchedulerStatsCollector` with atomics for requested, applied, and failed +affinity. Populate `MslSchedulerTimings` from those atomics, and set +`pinned_worker_count = affinity_applied_worker_count`. Do not count planned +core IDs as successful pinning. + +- [ ] **Step 6: Run GREEN tests and affected package tests** + +Run: + +```bash +CARGO_TARGET_DIR=/Users/hechuan/workspace/climamind/rumoca/target \ + cargo test -p rumoca-worker -- --nocapture +CARGO_TARGET_DIR=/Users/hechuan/workspace/climamind/rumoca/target \ + cargo test -p rumoca-test-msl --features msl-full-test --test msl_tests \ + affinity_ -- --nocapture +``` + +Expected: all selected tests pass with no warnings from test code. + +- [ ] **Step 7: Commit** + +```bash +git add crates/rumoca-worker/src/lib.rs \ + crates/rumoca-worker/src/bin/rumoca-worker.rs \ + crates/rumoca-test-msl/tests/balance_pipeline/mod.rs \ + crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_core.rs \ + crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_core/streaming_workers.rs \ + crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_core/tests.rs +git commit -s -m "fix(msl): report actual worker affinity" +``` + +--- + +### Task 2: Capture Wall-Time Measurement Provenance + +**Files:** +- Create: `crates/rumoca-test-msl/src/runtime_measurement.rs` +- Modify: `crates/rumoca-test-msl/src/lib.rs` +- Modify: `crates/rumoca-test-msl/tests/balance_pipeline/mod.rs` +- Modify: `crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_core.rs` +- Modify: `crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_summary.rs` +- Modify: `crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference.rs` +- Modify: `crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference/output.rs` +- Test: `crates/rumoca-test-msl/src/runtime_measurement.rs` +- Test: `crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference/tests.rs` + +**Interfaces:** +- Consumes: Task 1 scheduler affinity counts serialized in `msl_results.json`. +- Produces: `HostLoadSnapshot { one_minute: f64, logical_cpus: usize }` with `normalized() -> f64`. +- Produces: `WallTimeMeasurementProvenance` serialized under `runtime_comparison.wall_time_provenance`. +- Produces fields: `omc_fresh_sample_count`, `omc_cached_sample_count`, `affinity_requested_worker_count`, `affinity_applied_worker_count`, `affinity_failed_worker_count`, `normalized_load_before`, `normalized_load_after`, `workers_used`, and `omc_threads`. + +- [ ] **Step 1: Write failing host-load parsing and threshold tests** + +Create the module with tests first: + +```rust +#[test] +fn parses_linux_loadavg() { + let snapshot = parse_load_text("15.0 8.0 4.0 2/100 123", 10).unwrap(); + assert_eq!(snapshot.one_minute, 15.0); + assert_eq!(snapshot.normalized(), 1.5); +} + +#[test] +fn parses_macos_vm_loadavg() { + let snapshot = parse_load_text("{ 2.50 3.00 4.00 }", 10).unwrap(); + assert_eq!(snapshot.one_minute, 2.5); + assert_eq!(snapshot.normalized(), 0.25); +} + +#[test] +fn rejects_missing_or_non_finite_load() { + assert!(parse_load_text("unavailable", 10).is_none()); + assert!(parse_load_text("NaN 1 1", 10).is_none()); + assert!(parse_load_text("1 1 1", 0).is_none()); +} +``` + +- [ ] **Step 2: Run RED load tests** + +Run: + +```bash +CARGO_TARGET_DIR=/Users/hechuan/workspace/climamind/rumoca/target \ + cargo test -p rumoca-test-msl runtime_measurement -- --nocapture +``` + +Expected: FAIL because the module and parser do not exist. + +- [ ] **Step 3: Implement safe platform load sampling** + +Implement `parse_load_text` without `unsafe`. On Linux read `/proc/loadavg`; +on macOS execute `sysctl -n vm.loadavg`; on other platforms return `None`. +Use `std::thread::available_parallelism()` for the denominator. Record the +before snapshot at the beginning of `run_msl_test` and serialize it in +`MslPhaseTimings` with `#[serde(default)]` compatibility. Record the after +snapshot in `omc_simulation_reference` after its fresh/resume session work, so +the trust interval spans both Rumoca and OMC measurement phases. + +- [ ] **Step 4: Write failing OMC cache-provenance tests** + +Extend the existing resume fixtures so one cached result and one newly run +result produce exact counts: + +```rust +assert_eq!(payload["runtime_comparison"]["wall_time_provenance"]["omc_cached_sample_count"], 1); +assert_eq!(payload["runtime_comparison"]["wall_time_provenance"]["omc_fresh_sample_count"], 1); +``` + +Also assert that affinity and load values from the Rumoca result summary are +copied into the same provenance object. + +- [ ] **Step 5: Run RED provenance tests** + +Run: + +```bash +CARGO_TARGET_DIR=/Users/hechuan/workspace/climamind/rumoca/target \ + cargo test -p rumoca-test-msl omc_simulation_reference -- --nocapture +``` + +Expected: FAIL because `wall_time_provenance` is absent. + +- [ ] **Step 6: Track cache origin and emit provenance** + +Extend `SimRunState` with a set of cached OMC model names populated by +`merge_cached_results_for_resume`. At output time, count only models that +participate in `wall_ratio_both_success`; classify each as fresh or cached. +Read scheduler affinity and host-load snapshots from `msl_results.json` and +write the structured provenance alongside the runtime ratio stats. Do not +infer freshness from `batches_skipped` alone. + +- [ ] **Step 7: Run GREEN tests** + +Run both commands from Steps 2 and 5. Expected: all selected tests pass and +the provenance counts exactly match their fixtures. + +- [ ] **Step 8: Commit** + +```bash +git add crates/rumoca-test-msl/src/runtime_measurement.rs \ + crates/rumoca-test-msl/src/lib.rs \ + crates/rumoca-test-msl/tests/balance_pipeline/mod.rs \ + crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_core.rs \ + crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_summary.rs \ + crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference.rs \ + crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference/output.rs \ + crates/rumoca-test-msl/src/msl_tools/omc_simulation_reference/tests.rs +git commit -s -m "feat(msl): record wall-time measurement provenance" +``` + +--- + +### Task 3: Enforce Trust-Qualified Wall-Time Policy + +**Files:** +- Modify: `crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate.rs` +- Modify: `crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/status.rs` +- Modify: `crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/tests.rs` +- Modify: `spec/SPEC_0025_PR_REVIEW_PROCESS.md` +- Modify: `docs/dev-guide/src/tooling/msl-quality-gate.md` + +**Interfaces:** +- Consumes: Task 2 `runtime_comparison.wall_time_provenance`. +- Produces: `WallTimeTrustDecision { trusted: bool, reasons: Vec }`. +- Produces: blocking reasons containing wall regression only when `trusted` is true. +- Produces: console status `MSL wall speed gate: PASS|FAIL|ADVISORY` with provenance reasons. + +- [ ] **Step 1: Write failing trust-decision tests** + +Add fixtures and tests for all policy branches: + +```rust +#[test] +fn cached_omc_wall_regression_is_advisory_but_system_regression_blocks() { + let parity = parity_with_provenance( + runtime_ratio_stats(1.0, 0.5), + provenance(/* fresh */ 0, /* cached */ 10, 2, 2, 0, 0.2, 0.3), + ); + let mut reasons = Vec::new(); + push_runtime_ratio_regression_reasons(&mut reasons, &baseline_with_runtime(2.0, 1.5), Some(&parity)); + assert!(reasons.iter().any(|reason| reason.contains("runtime system speedup median"))); + assert!(!reasons.iter().any(|reason| reason.contains("runtime wall speedup median"))); + assert!(wall_time_trust_decision(Some(&parity)).reasons.iter().any(|reason| reason.contains("cached"))); +} + +#[test] +fn trusted_wall_regression_remains_blocking() { + let parity = parity_with_provenance( + runtime_ratio_stats(2.0, 0.5), + provenance(10, 0, 2, 2, 0, 0.5, 0.6), + ); + let mut reasons = Vec::new(); + push_runtime_ratio_regression_reasons(&mut reasons, &baseline_with_runtime(2.0, 1.5), Some(&parity)); + assert!(reasons.iter().any(|reason| reason.contains("runtime wall speedup median"))); +} +``` + +Add separate tests proving affinity failure, load `> 1.5`, missing load, missing +provenance, and mismatched worker/thread policy each produce `trusted == false`. + +- [ ] **Step 2: Run RED gate tests** + +Run: + +```bash +CARGO_TARGET_DIR=/Users/hechuan/workspace/climamind/rumoca/target \ + cargo test -p rumoca-test-msl --features msl-full-test --test msl_tests \ + wall_time_ -- --nocapture +``` + +Expected: FAIL because the trust model and advisory behavior do not exist. + +- [ ] **Step 3: Implement parsing and pure trust decision** + +Add serde-compatible provenance types and parse the Task 2 JSON object. The +pure decision must append stable reason strings for cached samples, affinity +failure, missing/excessive load, and runtime-context mismatch. `trusted` is +true only when the reason list is empty. + +- [ ] **Step 4: Apply trust decision to blocking reasons and status** + +Keep the existing unconditional system comparison. Wrap only the wall +regression insertion with `wall_time_trust_decision(...).trusted`. Update +status output so advisory measurements print their observed median, baseline, +35% floor, and reason list without saying `PASS`. + +- [ ] **Step 5: Update policy documentation** + +Change the runtime row in `SPEC_0025_PR_REVIEW_PROCESS.md` to state: + +```text +Runtime system-time speedup median MUST NOT regress by > 35%. Wall-time uses +the same 35% limit only for fresh, affinity-correct, healthy-host paired +measurements; otherwise it remains visible as ADVISORY and does not mask any +correctness or system-time failure. +``` + +Add the same behavior, provenance fields, and `PASS|FAIL|ADVISORY` meanings to +the MSL quality-gate developer guide. + +- [ ] **Step 6: Run focused GREEN tests** + +Run: + +```bash +CARGO_TARGET_DIR=/Users/hechuan/workspace/climamind/rumoca/target \ + cargo test -p rumoca-test-msl --features msl-full-test --test msl_tests \ + runtime_ratio_ -- --nocapture +CARGO_TARGET_DIR=/Users/hechuan/workspace/climamind/rumoca/target \ + cargo test -p rumoca-test-msl --features msl-full-test --test msl_tests \ + wall_time_ -- --nocapture +``` + +Expected: every selected runtime and trust test passes. + +- [ ] **Step 7: Run affected verification** + +Run: + +```bash +cargo fmt --check +CARGO_TARGET_DIR=/Users/hechuan/workspace/climamind/rumoca/target \ + cargo clippy -p rumoca-worker -p rumoca-test-msl --all-targets --all-features -- -D warnings +CARGO_TARGET_DIR=/Users/hechuan/workspace/climamind/rumoca/target \ + cargo test -p rumoca-worker +CARGO_TARGET_DIR=/Users/hechuan/workspace/climamind/rumoca/target \ + cargo test -p rumoca-test-msl --features msl-full-test --test msl_tests \ + balance_pipeline::balance_pipeline_quality_gate::tests -- --nocapture +CARGO_TARGET_DIR=/Users/hechuan/workspace/climamind/rumoca/target \ + cargo test -p rumoca --test architecture_hardening_test +cargo xtask verify docs +``` + +Expected: all commands exit 0. + +- [ ] **Step 8: Run the real MSL parity acceptance gate** + +Run: + +```bash +CARGO_TARGET_DIR=/Users/hechuan/workspace/climamind/rumoca/target \ + cargo xtask verify msl-parity +``` + +Expected on the current cached/high-load path: correctness and system-time +checks pass; a wall regression prints `ADVISORY` with cache, affinity, or load +reasons and does not fail the run. Do not claim a trusted wall-time pass unless +the emitted provenance actually satisfies every trust condition. + +- [ ] **Step 9: Commit** + +```bash +git add crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate.rs \ + crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/status.rs \ + crates/rumoca-test-msl/tests/balance_pipeline/balance_pipeline_quality_gate/tests.rs \ + spec/SPEC_0025_PR_REVIEW_PROCESS.md \ + docs/dev-guide/src/tooling/msl-quality-gate.md +git commit -s -m "fix(msl): trust-qualify wall-time regressions" +``` + +--- + +## Final Branch Verification and Delivery + +- [ ] Generate a whole-branch review package from merge base `6936a2043843ac701cc0cfac57bef61ba9de9a68` to final HEAD and dispatch the final code reviewer. +- [ ] Resolve every Critical or Important finding through one fix subagent and re-review. +- [ ] Re-run the affected verification commands after the final fix commit. +- [ ] Confirm `git status --short` is empty and every commit has `Signed-off-by`. +- [ ] Push only: + +```bash +git push origin cli-52-upstream-equivalent-cleanup +``` + +- [ ] Fetch and prove local branch, `origin/cli-52-upstream-equivalent-cleanup`, and the intended final SHA are equal. +- [ ] Confirm no push or pull request was made to the `upstream` remote. diff --git a/docs/superpowers/specs/2026-07-14-msl-wall-time-trust-gate-design.md b/docs/superpowers/specs/2026-07-14-msl-wall-time-trust-gate-design.md new file mode 100644 index 000000000..a5c8343ff --- /dev/null +++ b/docs/superpowers/specs/2026-07-14-msl-wall-time-trust-gate-design.md @@ -0,0 +1,156 @@ +# MSL Wall-Time Trust Gate Design + +## Context + +The full MSL quality gate compares Rumoca runtime speedup against a promoted +baseline. System-time and wall-time medians currently use the same blocking +policy: either may regress by at most 35%. + +The wall comparison is not always paired. The parity harness can reuse cached +OMC timings while refreshing Rumoca timings on the current host. In the +observed failure, the source commit did not change between runs, system-time +speedup stayed stable, and wall-time speedup moved materially while the host +was heavily loaded and worker CPU pinning failed. Treating that result as a +code regression makes the gate sensitive to measurement provenance and host +scheduling rather than only to Rumoca performance. + +This is a ClimaMind-fork change. It will be committed and pushed only to the +ClimaMind `origin` feature branch; no pull request or push to the CogniPilot +`upstream` remote is part of this work. + +## Goals + +- Keep correctness, coverage, trace, and system-time regression checks + fail-closed. +- Keep the existing 35% runtime tolerance. +- Make wall-time blocking conditional on a trustworthy, paired measurement. +- Preserve wall-time metrics and warnings when the measurement is not trusted. +- Record enough provenance in generated artifacts to explain every blocking or + advisory decision. + +## Non-Goals + +- Do not relax stage-count, balance, simulation, trace, or system-time gates. +- Do not change the promoted baseline to make the current result pass. +- Do not silently ignore missing runtime data. +- Do not add retry loops or environment-specific bypass variables. +- Do not submit this downstream policy change to CogniPilot upstream. + +## Considered Approaches + +### 1. Increase the 35% tolerance + +Rejected. A wider threshold hides real regressions without fixing the +comparison between cached and fresh measurements. + +### 2. Make wall time permanently advisory + +Rejected. Wall time remains valuable on a controlled, paired benchmark and can +catch scheduler, allocation, or I/O regressions that system time misses. + +### 3. Trust-qualified wall-time gate + +Selected. System time remains blocking on every complete parity run. Wall time +is blocking only when its provenance and host-health evidence show that the OMC +and Rumoca measurements are comparable. Otherwise the same regression is +reported as an advisory warning. + +## Design + +### Measurement provenance + +The parity artifact and quality snapshot will record a structured wall-time +trust context. It will include: + +- whether all OMC runtime samples used by the comparison were generated fresh + in the current invocation or whether any were resumed from cache; +- requested and effective worker counts and OMC thread count; +- actual Rumoca worker CPU-affinity success and failure counts; +- normalized one-minute host load sampled before and after the measured work, + when the platform exposes it safely; +- a reason list explaining why wall time is trusted or advisory. + +Missing provenance is conservative: it makes wall time advisory, not trusted. +It does not make the full parity artifact valid when required runtime or trace +samples are missing. + +### Trust decision + +Wall time is blocking only when all of the following are true: + +1. Every compared OMC runtime sample is fresh in the current invocation. +2. The Rumoca workers requested for pinning report successful affinity. +3. Host load is available and remains within the documented normalized limit + before and after the measured work. +4. Runtime worker/thread context matches the comparison policy. + +If any condition fails, the wall median and its baseline delta remain visible, +but `push_runtime_ratio_regression_reasons` does not add the wall regression to +the blocking reason list. It adds an advisory status with the failed trust +conditions instead. System-time regression remains blocking regardless of the +wall-time trust decision. + +Normalized load is the one-minute load average divided by available logical +CPUs. The initial limit will be `1.5`. This allows the benchmark itself to +occupy the host while rejecting the observed multi-fold oversubscription. +Platforms without a safe load source report load as unavailable and therefore +cannot produce a blocking wall-time verdict. + +### Affinity reporting + +The model-worker ready handshake will report whether a requested CPU affinity +was applied. The MSL scheduler will aggregate actual successes and failures; +the existing `pinned_worker_count` field will represent successful pinning, +not merely the number of planned core assignments. This keeps the trust +decision based on worker evidence rather than parsing warning text. + +### Output and documentation + +The console status, `omc_simulation_reference.json`, and +`msl_quality_current.json` will distinguish: + +- `PASS`: trusted wall measurement and no regression; +- `FAIL`: trusted wall measurement exceeds the existing tolerance; +- `ADVISORY`: wall measurement is present but not trusted, with reasons. + +`SPEC_0025_PR_REVIEW_PROCESS.md` and the MSL quality-gate developer guide will +state that the 35% wall-time rule applies only to trust-qualified paired +measurements. The system-time rule remains unconditional. + +## Error Handling + +- Missing OMC parity, empty runtime samples, stale target sets, or missing + trace data continue to fail closed under the existing rules. +- Missing trust metadata downgrades only the wall regression verdict to + advisory. +- Malformed trust metadata is treated as untrusted and surfaced in the status + reason list. +- No fallback promotes a cached or unhealthy wall measurement to trusted. + +## Test Strategy + +Development follows red-green-refactor: + +1. Add a failing unit test proving a cached OMC reference cannot create a + blocking wall regression while system regression still blocks. +2. Add failing trust-decision tests for affinity failure, excessive normalized + load, missing provenance, and a fully trusted measurement. +3. Add protocol tests proving worker affinity status survives the ready + handshake and is aggregated as actual success/failure counts. +4. Add snapshot/parsing tests for the new provenance and advisory status. +5. Run the focused runtime-ratio and cache-resume test filters. +6. Run formatting, clippy for affected crates, architecture/spec tests, and the + repository-required verification appropriate to the final diff. + +The heavyweight MSL gate will be used to confirm that an unhealthy or cached +measurement reports `ADVISORY` rather than a false performance regression. A +clean-host trusted-wall pass is reported only if the run actually satisfies +all trust conditions. + +## Delivery + +Implementation will remain on `cli-52-upstream-equivalent-cleanup` in the +ClimaMind fork. The branch will be pushed only to +`https://github.com/climamind/rumoca.git`. The `upstream` remote will be used +only as a read-only reference and will not receive commits, branches, or pull +requests. diff --git a/docs/user-guide/book.toml b/docs/user-guide/book.toml index d97569c55..3e4a8918c 100644 --- a/docs/user-guide/book.toml +++ b/docs/user-guide/book.toml @@ -8,8 +8,8 @@ src = "src" build-dir = "book" [output.html] -git-repository-url = "https://github.com/CogniPilot/rumoca" -edit-url-template = "https://github.com/CogniPilot/rumoca/edit/main/docs/user-guide/{path}" +git-repository-url = "https://github.com/climamind/rumoca" +edit-url-template = "https://github.com/climamind/rumoca/edit/main/docs/user-guide/{path}" default-theme = "rust" preferred-dark-theme = "ayu" additional-js = ["live/rumoca-live.js"] diff --git a/docs/user-guide/src/introduction.md b/docs/user-guide/src/introduction.md index fa971a976..e7c2bc852 100644 --- a/docs/user-guide/src/introduction.md +++ b/docs/user-guide/src/introduction.md @@ -362,5 +362,5 @@ today. - **Code Generation** covers built-in and custom targets. Developers who want to understand or modify the compiler itself should read -the companion [Rumoca Dev Guide](https://cognipilot.github.io/rumoca/dev-guide/) +the companion [Rumoca Dev Guide](https://climamind.github.io/rumoca/dev-guide/) book. diff --git a/flake.nix b/flake.nix index 72e8b04bb..19af5f876 100644 --- a/flake.nix +++ b/flake.nix @@ -89,7 +89,8 @@ # Release-mode artifacts for the MSL parity gate, built as one Cargo # graph so the shard / merge / ModelicaTest consumers restore them via - # Cachix instead of recompiling + re-LTO'ing the workspace. A single + # a GitHub Actions closure artifact instead of recompiling + re-LTO'ing + # the workspace. A single # derivation keeps rumoca-worker, rumoca-sim-worker, rumoca-msl-tools, # and the libtest harness in one target directory; separate derivations # rebuild the same workspace crates and made rumoca-worker a serial diff --git a/infra/install/install.ps1 b/infra/install/install.ps1 index 655f906cb..0fffd8563 100644 --- a/infra/install/install.ps1 +++ b/infra/install/install.ps1 @@ -1,7 +1,7 @@ [CmdletBinding()] param( [string]$Version = $env:RUMOCA_INSTALL_VERSION, - [string]$Repo = $(if ($env:RUMOCA_INSTALL_REPO) { $env:RUMOCA_INSTALL_REPO } else { "cognipilot/rumoca" }), + [string]$Repo = $(if ($env:RUMOCA_INSTALL_REPO) { $env:RUMOCA_INSTALL_REPO } else { "climamind/rumoca" }), [string]$BinDir = $(if ($env:RUMOCA_INSTALL_BIN_DIR) { $env:RUMOCA_INSTALL_BIN_DIR } else { Join-Path $env:LOCALAPPDATA "rumoca\bin" }), [switch]$WithLsp ) diff --git a/infra/install/install.sh b/infra/install/install.sh index a3a057612..14db7aa6f 100755 --- a/infra/install/install.sh +++ b/infra/install/install.sh @@ -1,7 +1,7 @@ #!/usr/bin/env bash set -euo pipefail -REPO="${RUMOCA_INSTALL_REPO:-cognipilot/rumoca}" +REPO="${RUMOCA_INSTALL_REPO:-climamind/rumoca}" BIN_DIR="${RUMOCA_INSTALL_BIN_DIR:-$HOME/.local/bin}" VERSION="${RUMOCA_INSTALL_VERSION:-latest}" WITH_LSP="${RUMOCA_INSTALL_WITH_LSP:-0}" diff --git a/packages/vscode/README.md b/packages/vscode/README.md index a4d41447c..0ab041d11 100644 --- a/packages/vscode/README.md +++ b/packages/vscode/README.md @@ -1,8 +1,8 @@ # Rumoca Modelica -A VS Code extension providing language support for [Modelica](https://modelica.org/) using the [rumoca](https://github.com/cognipilot/rumoca) compiler. +A VS Code extension providing language support for [Modelica](https://modelica.org/) using the [rumoca](https://github.com/climamind/rumoca) compiler. -📖 **Documentation:** [Rumoca User Guide](https://cognipilot.github.io/rumoca/user-guide/) · [Rumoca Dev Guide](https://cognipilot.github.io/rumoca/dev-guide/) · [Web Playground](https://cognipilot.github.io/rumoca/) — also available from the command palette: **Rumoca: Open Rumoca User Guide**. +📖 **Documentation:** [Rumoca User Guide](https://climamind.github.io/rumoca/user-guide/) · [Rumoca Dev Guide](https://climamind.github.io/rumoca/dev-guide/) · [Web Playground](https://climamind.github.io/rumoca/) — also available from the command palette: **Rumoca: Open Rumoca User Guide**. ## Features @@ -35,7 +35,7 @@ version mismatches. **From VSIX file:** -1. Download the `.vsix` file for your platform from [GitHub Releases](https://github.com/cognipilot/rumoca/releases) +1. Download the `.vsix` file for your platform from [GitHub Releases](https://github.com/climamind/rumoca/releases) 2. In VS Code, open the Command Palette (`Ctrl+Shift+P` / `Cmd+Shift+P`) 3. Run "Extensions: Install from VSIX..." 4. Select the downloaded `.vsix` file @@ -70,10 +70,10 @@ If you need to install `rumoca-lsp` manually: ```bash # From GitHub Releases installer -curl --proto '=https' --tlsv1.2 -LsSf https://raw.githubusercontent.com/cognipilot/rumoca/main/infra/install/install.sh | bash -s -- --with-lsp +curl --proto '=https' --tlsv1.2 -LsSf https://raw.githubusercontent.com/climamind/rumoca/main/infra/install/install.sh | bash -s -- --with-lsp # Or from source -git clone https://github.com/cognipilot/rumoca.git +git clone https://github.com/climamind/rumoca.git cd rumoca cargo install --path crates/rumoca-tool-lsp ``` diff --git a/packages/vscode/package.json b/packages/vscode/package.json index 112b074a9..1bbc34257 100644 --- a/packages/vscode/package.json +++ b/packages/vscode/package.json @@ -6,11 +6,11 @@ "publisher": "JamesGoppert", "repository": { "type": "git", - "url": "https://github.com/cognipilot/rumoca" + "url": "https://github.com/climamind/rumoca" }, - "homepage": "https://github.com/cognipilot/rumoca#readme", + "homepage": "https://github.com/climamind/rumoca#readme", "bugs": { - "url": "https://github.com/cognipilot/rumoca/issues" + "url": "https://github.com/climamind/rumoca/issues" }, "license": "Apache-2.0", "icon": "icons/rumoca.png", diff --git a/scripts/ci/apt-install.sh b/scripts/ci/apt-install.sh new file mode 100755 index 000000000..85b188ec4 --- /dev/null +++ b/scripts/ci/apt-install.sh @@ -0,0 +1,49 @@ +#!/usr/bin/env bash +set -euo pipefail + +if [ "$#" -eq 0 ]; then + echo "usage: $0 [apt-get-install-flags...] ..." >&2 + exit 64 +fi + +if [ "${EUID:-$(id -u)}" -eq 0 ]; then + SUDO=() +else + SUDO=(sudo) +fi + +disable_unstable_runner_sources() { + local source_path + + # GitHub's ubuntu-24.04 runner can ship Microsoft apt sources that return + # transient invalid or unauthorized metadata, which makes apt-get update fail + # before Ubuntu repository packages are used by this script. + while IFS= read -r source_path; do + if [ -n "$source_path" ] && [ -e "$source_path" ]; then + echo "Disabling unstable apt source: $source_path" + "${SUDO[@]}" mv "$source_path" "$source_path.disabled" + fi + done < <( + grep -Erl "packages\.microsoft\.com/(repos/azure-cli|ubuntu/24\.04/prod)" \ + /etc/apt/sources.list /etc/apt/sources.list.d 2>/dev/null || true + ) +} + +apt_update_with_retry() { + local attempt + + for attempt in 1 2 3; do + if "${SUDO[@]}" apt-get update; then + return 0 + fi + if [ "$attempt" -lt 3 ]; then + sleep "$((attempt * 5))" + fi + done + + return 1 +} + +disable_unstable_runner_sources +apt_update_with_retry +"${SUDO[@]}" env DEBIAN_FRONTEND=noninteractive apt-get install -y "$@" diff --git a/scripts/ci/install-openmodelica.sh b/scripts/ci/install-openmodelica.sh new file mode 100755 index 000000000..6fa9fd21e --- /dev/null +++ b/scripts/ci/install-openmodelica.sh @@ -0,0 +1,68 @@ +#!/usr/bin/env bash +set -euo pipefail + +repo_root="$(cd "$(dirname "${BASH_SOURCE[0]}")/../.." && pwd)" +expected_package_version="$(tr -d '[:space:]' < "$repo_root/toolchains/openmodelica-version")" +package_manifest="$repo_root/toolchains/openmodelica-packages.txt" +package_verifier="$repo_root/scripts/ci/verify-openmodelica-packages.sh" +if [[ ! "$expected_package_version" =~ ^([0-9]+\.[0-9]+\.[0-9]+)~1-g[0-9a-f]+-[0-9]+$ ]]; then + echo "Invalid OpenModelica package pin: $expected_package_version" >&2 + exit 1 +fi +expected_version="${expected_package_version%-*}" +if [[ "${EUID:-$(id -u)}" -eq 0 ]]; then + SUDO=() +else + SUDO=(sudo) +fi + +installed_version() { + local version_output="$1" + [[ "$version_output" =~ ^OpenModelica\ (.+)$ ]] || return 1 + printf '%s\n' "${BASH_REMATCH[1]}" +} + +check_version() { + local version_output="$1" + local actual_version + actual_version="$(installed_version "$version_output" || true)" + if [[ "$actual_version" != "$expected_version" ]]; then + echo "OpenModelica version mismatch: expected $expected_version, got ${actual_version:-unknown}" >&2 + return 1 + fi + printf '%s\n' "$version_output" +} + +if [[ "${1:-}" == "--check-output" ]]; then + [[ "$#" -eq 2 ]] || { echo "usage: $0 --check-output ''" >&2; exit 64; } + check_version "$2" + exit +fi + +package_list="$("$package_verifier" "$package_manifest" "$expected_package_version")" +package_directory="$(mktemp -d)" +trap 'rm -rf -- "$package_directory"' EXIT + +"${SUDO[@]}" apt-get update +"${SUDO[@]}" apt-get install -y --no-install-recommends \ + build-essential ca-certificates clang cmake curl \ + libexpat1-dev liblapack-dev unzip zip +while read -r filename url; do + curl --proto '=https' --tlsv1.2 --fail --location --show-error \ + --output "$package_directory/$filename" "$url" +done <<< "$package_list" +"$package_verifier" "$package_manifest" "$expected_package_version" "$package_directory" + +"${SUDO[@]}" apt-get install -y --no-install-recommends \ + "$package_directory/omc_${expected_package_version}_amd64.deb" \ + "$package_directory/omc-common_${expected_package_version}_all.deb" \ + "$package_directory/libomc_${expected_package_version}_amd64.deb" \ + "$package_directory/libomcsimulation_${expected_package_version}_amd64.deb" +for package in omc omc-common libomc libomcsimulation; do + installed_package_version="$(dpkg-query -W -f='${Version}' "$package")" + if [[ "$installed_package_version" != "$expected_package_version" ]]; then + echo "Installed $package version mismatch: expected $expected_package_version, got $installed_package_version" >&2 + exit 1 + fi +done +check_version "$(omc --version)" diff --git a/scripts/ci/verify-openmodelica-packages.sh b/scripts/ci/verify-openmodelica-packages.sh new file mode 100755 index 000000000..502e0c7dd --- /dev/null +++ b/scripts/ci/verify-openmodelica-packages.sh @@ -0,0 +1,72 @@ +#!/usr/bin/env bash +set -euo pipefail + +if [[ "$#" -lt 2 || "$#" -gt 3 ]]; then + echo "usage: $0 MANIFEST EXPECTED_VERSION [PACKAGE_DIRECTORY]" >&2 + exit 64 +fi + +manifest="$1" +expected_version="$2" +package_directory="${3:-}" +pool="https://build.openmodelica.org/apt/pool/contrib-noble" +declare -a names architectures filenames hashes urls +count=0 +seen=" " + +while read -r name version architecture filename hash url extra; do + [[ -z "$name" || "$name" == \#* ]] && continue + [[ -z "${extra:-}" ]] || { echo "Unexpected manifest fields for $name" >&2; exit 1; } + case "$name" in + omc|libomc|libomcsimulation) required_architecture="amd64" ;; + omc-common) required_architecture="all" ;; + *) echo "Unexpected OpenModelica package: $name" >&2; exit 1 ;; + esac + [[ "$seen" != *" $name "* ]] || { echo "Duplicate OpenModelica package: $name" >&2; exit 1; } + [[ "$version" == "$expected_version" ]] || { echo "Manifest version mismatch for $name" >&2; exit 1; } + [[ "$architecture" == "$required_architecture" ]] || { echo "Manifest architecture mismatch for $name" >&2; exit 1; } + expected_filename="${name}_${expected_version}_${architecture}.deb" + [[ "$filename" == "$expected_filename" ]] || { echo "Manifest filename mismatch for $name" >&2; exit 1; } + [[ "$hash" =~ ^[0-9a-f]{64}$ ]] || { echo "Invalid SHA-256 for $name" >&2; exit 1; } + [[ "$url" == "$pool/$filename" ]] || { echo "Manifest URL mismatch for $name" >&2; exit 1; } + names[count]="$name" + architectures[count]="$architecture" + filenames[count]="$filename" + hashes[count]="$hash" + urls[count]="$url" + count=$((count + 1)) + seen+="$name " +done < "$manifest" + +for required in omc omc-common libomc libomcsimulation; do + [[ "$seen" == *" $required "* ]] || { echo "Missing OpenModelica package: $required" >&2; exit 1; } +done +[[ "$count" -eq 4 ]] || { echo "Expected four OpenModelica packages, got $count" >&2; exit 1; } + +if [[ -z "$package_directory" ]]; then + for ((index = 0; index < count; index++)); do + printf '%s %s\n' "${filenames[index]}" "${urls[index]}" + done + exit +fi + +deb_count="$(find "$package_directory" -maxdepth 1 -type f -name '*.deb' | wc -l | tr -d '[:space:]')" +[[ "$deb_count" == "4" ]] || { echo "Package directory must contain exactly four .deb files" >&2; exit 1; } + +for ((index = 0; index < count; index++)); do + artifact="$package_directory/${filenames[index]}" + [[ -f "$artifact" ]] || { echo "Missing OpenModelica artifact: ${filenames[index]}" >&2; exit 1; } + if command -v sha256sum >/dev/null 2>&1; then + actual_hash="$(sha256sum "$artifact" | awk '{print $1}')" + else + actual_hash="$(shasum -a 256 "$artifact" | awk '{print $1}')" + fi + [[ "$actual_hash" == "${hashes[index]}" ]] || { echo "Checksum mismatch for ${filenames[index]}" >&2; exit 1; } + + actual_package="$(dpkg-deb -f "$artifact" Package)" + actual_version="$(dpkg-deb -f "$artifact" Version)" + actual_architecture="$(dpkg-deb -f "$artifact" Architecture)" + [[ "$actual_package" == "${names[index]}" ]] || { echo "Package mismatch for ${filenames[index]}" >&2; exit 1; } + [[ "$actual_version" == "$expected_version" ]] || { echo "Version mismatch for ${filenames[index]}" >&2; exit 1; } + [[ "$actual_architecture" == "${architectures[index]}" ]] || { echo "Architecture mismatch for ${filenames[index]}" >&2; exit 1; } +done diff --git a/spec/README.md b/spec/README.md index 5d07d7ed9..39f3f16a0 100644 --- a/spec/README.md +++ b/spec/README.md @@ -25,6 +25,7 @@ For setup and day-to-day usage, see [CONTRIBUTING.md](../CONTRIBUTING.md). | [SPEC_0029](SPEC_0029_CRATE_BOUNDARIES.md) | Crate Boundaries as Collaboration Guardrails | architecture | ~340 | ACCEPTED | | [SPEC_0031](SPEC_0031_COMPILER_PHILOSOPHY.md) | Compiler Scope and Philosophy | architecture | ~150 | ACCEPTED | | [SPEC_0032](SPEC_0032_RANGE_PRESERVING_TENSORS.md) | Range-Preserving Tensor IR | IR | ~85 | ACCEPTED | +| [SPEC_0033](SPEC_0033_DEVELOPMENT_PROCESS.md) | Development Process | process | ~105 | ACCEPTED | | [SPEC_0034](SPEC_0034_GALEC_EFMI_EXPORT.md) | eFMI/GALEC Algorithm Code Export | target/codegen | ~180 | DRAFT | ## Deferred Specifications diff --git a/spec/SPEC_0001_DEFID.md b/spec/SPEC_0001_DEFID.md index c4292b6b8..4a84270c9 100644 --- a/spec/SPEC_0001_DEFID.md +++ b/spec/SPEC_0001_DEFID.md @@ -43,6 +43,7 @@ component and MUST have unique post-instantiation identity. | Flat/DAE runtime components use instance-unique identity | Flatten and later | Reused classes create distinct unknowns | | `DefId(0)` is reserved for root/global scope | Core ids | Stable sentinel | | DefIds are local to a compilation unit | Whole compiler | No cross-unit global registry | +| Builtin type DefIds come from the compiler-owned typed builtin catalog | Resolve / Core | Downstream phases compare builtin identity without inspecting rendered names | Instance-unique identity may be represented by a dedicated instance `DefId`, or by a small structured key whose identity fields are all `DefId` / diff --git a/spec/SPEC_0007_IR_PIPELINE.md b/spec/SPEC_0007_IR_PIPELINE.md index 9701bdd30..c67f4d690 100644 --- a/spec/SPEC_0007_IR_PIPELINE.md +++ b/spec/SPEC_0007_IR_PIPELINE.md @@ -6,7 +6,7 @@ ACCEPTED ## Summary Rumoca transforms Modelica through AST → Flat → DAE → Solve IRs. Each -stage has a contract: contents, ownership, boundary leaks. +stage owns contract. ## The Four IR Stages @@ -48,6 +48,8 @@ Modelica source (.mo) adapters wrap toolchains, packaging, runtime calls, or JIT APIs, not semantics, DAE lowering, structural rewrites, or template policy. +**Neutral codegen:** IR; no shims. + --- ### Stage 1 — AST (`rumoca-ir-ast`) @@ -141,6 +143,7 @@ has the same meaning as the default. Incompatible schema changes bump | Rule | Where | Why | |---|---|---| +| Each proven matrix-product result lane expands to a complete inner-dimension dot sum; unknown or mismatched shapes fail closed | DAE lowering | Preserves exact matrix-product semantics without guessing shape facts | | No source temporal operators (`pre`, `edge`, `change`, `sample`, `previous`) survive in f_x, f_z, f_m, f_c, relations, or initialization equations | DAE lowering rewrites them into Appendix B constructs: explicit `__pre__.*` inputs, relation/c variables, scheduled events, clock metadata, and ordinary equations over `v` | MLS Appendix B states the DAE as functions over `v` and `relation(v)`; source temporal operators are not computable DAE/Solve graph nodes | | No `der()` on RHS | derivatives flow via `dae.states` + equation structure | Inline `der()` would hide state identity | | No `initial()` in f_x/f_z/f_m/f_c | initial phase is handled separately | Avoids mixing initialization into runtime equations | @@ -149,6 +152,7 @@ has the same meaning as the default. Incompatible schema changes bump | `sample(...)` and clocked `previous(...)` are represented by DAE event/clock metadata plus ordinary equations over current/pre slots | Runtime scheduling data is explicit DAE metadata; sampled values are `__pre__.*` reads where needed | Keeps clock semantics at DAE level while keeping compute functions ordinary | | `reinit(x, expr)` is lowered into guarded discrete state-update equations before DAE validation | DAE lowering converts state resets into ordinary Appendix B update equations over current/pre slots | Keeps state reset semantics in the numeric update system instead of exposing a source operator to runtimes | | `assert(...)` and `terminate(...)` are represented as `events.event_actions`, not as residual/value expressions | DAE lowering converts integration-flow statements into guarded event actions with source spans | Keeps Appendix B compute graphs pure while preserving solver-visible runtime actions | +| DAE row projection preserves proven array-product semantics | When a DAE consumer requires scalar rows, `Mul` projects matrix/vector result lanes as full inner-dimension dot products, scalar-array products project only the array side, and `MulElem` stays same-lane elementwise; unknown, incompatible, zero-inner, or unsupported-rank shapes fail with the source span instead of broadcasting or selecting one lane | Scalar views must remain mathematically equivalent to the symbolic DAE expression | | `appendix_b_validation` rejects any surviving source temporal operator | `phase-dae/src/appendix_b_validation.rs::validate_no_source_temporal_operator_survives` | Positive enforcement gate, not defensive code | **What to do here:** DAE-level passes such as pre-lowering and alias @@ -186,14 +190,16 @@ Canonical terminology: | `TensorProgramNode` | `ComputeNode::{MatMul, LinSolve, AffineStencil, ...}` | A tensor-level kernel with explicit shape/layout metadata and scalar fallback | | `ComputeBlock` | `ComputeBlock` | Ordered mix of scalar program blocks and tensor program nodes | -`ScalarProgramBlock` and `ComputeNode::ScalarPrograms` are the public source-code -names. New Solve-IR APIs must use `ScalarProgram` / `ScalarProgramBlock` -terminology and must not reintroduce `RowBlock` / `ScalarRows` naming. +Public APIs use `ScalarProgram`/`ScalarProgramBlock`; `RowBlock`/`ScalarRows` +must not return. + +GPU initialization requires exact, nonoverlapping, source-spanned Y coverage; +adjacency may merge, unsupported semantics never fall back, and settlement +shares one runtime/table context. -`ComputeNode::AffineStencil` is source-proven: it comes from preserved DAE -structured-family domains plus affine operand proofs. It carries the compact -iteration domain and strides; Solve lowering must not recover stencils by -scanning unstructured scalar rows after structured-family metadata is discarded. +`ComputeNode::AffineStencil` is source-proven from preserved DAE family domains +and affine operand proofs; Solve lowering must not recover stencils from +unstructured scalar rows. The root `schema_version` field is mandatory on serialized Solve payloads. Deserializers reject unsupported versions and the Solve wire format does not @@ -278,10 +284,7 @@ rendering (those live in DAE-IR/upstream lowering, `rumoca-exec-*`, or ## Structural Lowering Scope -Rumoca performs OpenModelica-class structural lowering between DAE and Solve. -Structural lowering is DAE-to-DAE: it rewrites or annotates mathematical -structure for downstream lowering without changing IR stage. The supported -transformations are listed here to keep scope and ownership clear. +Rumoca performs these DAE-to-DAE transformations before Solve. **In scope:** diff --git a/spec/SPEC_0022_MLS_COMPILER_COMPLIANCE.md b/spec/SPEC_0022_MLS_COMPILER_COMPLIANCE.md index 943c394de..a1d9ae1e1 100644 --- a/spec/SPEC_0022_MLS_COMPILER_COMPLIANCE.md +++ b/spec/SPEC_0022_MLS_COMPILER_COMPLIANCE.md @@ -607,7 +607,7 @@ Defines state-to-state transitions with priority and timing control. | FUNC-015 | Component types | §12.2 | "Function must not contain model, block, operator, or connector components" | | FUNC-016 | Not in connections | §12.2 | "Functions shall not be used in connections" | | FUNC-017 | Return in algorithm only | §12.1.2 | "Return statement can only be used in an algorithm section of a function" | -| FUNC-018 | Input ordering significant | §12.1.1 | "Relative ordering between input formal parameter declarations is significant" | +| FUNC-018 | Input ordering significant | §12.1.1 | "Relative ordering between input formal parameter declarations is significant"; flattening preserves positional actual order when actual and formal counts match | | FUNC-019 | Named arg slot error | §12.4.1 | "Error if named argument slot is already filled" | | FUNC-020 | Unfilled slots error | §12.4.1 | "Error if any unfilled slots remain after argument processing" | | FUNC-021 | Impure inheritance | §12.3 | "If function declared impure, any extending function shall be declared impure" | diff --git a/spec/SPEC_0025_PR_REVIEW_PROCESS.md b/spec/SPEC_0025_PR_REVIEW_PROCESS.md index 73305a994..ef72ba945 100644 --- a/spec/SPEC_0025_PR_REVIEW_PROCESS.md +++ b/spec/SPEC_0025_PR_REVIEW_PROCESS.md @@ -132,7 +132,7 @@ Rust developer workflow MUST remain Cargo-native. | Balanced / OMC-agreement counts MUST NOT decrease | These are headline correctness and numerical-quality numbers | | Focused or limited MSL runs MUST mark quality snapshots as partial and partial snapshots MUST NOT be promoted | Prevents local-debug subsets from becoming the committed release baseline | | Trace-quality metrics MUST be gated against the resolved promoted baseline when OMC parity data is available | Prevents balanced-but-numerically-worse simulations from passing unnoticed | -| Runtime speedup medians (system & wall) MUST NOT regress by > 35% | Tolerates 4-core hosted-runner noise without hiding material regressions | +| Runtime system-time speedup median MUST NOT regress by > 35%. Wall-time uses the same 35% limit only for all-fresh, affinity-correct, healthy-host paired measurements whose current and provenance worker/thread contexts exactly match the promoted baseline context; otherwise it remains visible as ADVISORY and does not mask any correctness or system-time failure. | Keeps the system-time regression gate unconditional while preventing cached, incomparable, or noisy wall measurements from blocking a correct run | | Promoted baseline release-asset updates require a successful full main CI run and a non-regressing ratchet decision; checked-in fallback updates remain explicit via `cargo xtask repo msl promote-quality-baseline` | Prevents silent baseline drift | | Coverage trim/gate updates follow `cargo xtask coverage {run,report,gate}` workflow | Coverage promotion is explicit only | @@ -173,6 +173,30 @@ net_added_lines: | No new trait without ≥ 2 concrete impls | Single-impl traits are noise | | No old/new code paths left side-by-side without explicit migration plan | Dead-but-alive code accretes | +### 6a. Authorized Broken-Main Recovery (optional) + +The normal reviewer gate remains unchanged. The following is the sole +exception: an Explicitly authorized ClimaMind Rumoca broken-main recovery batch +may waive the GitHub approving review for its owner PRs only. + +| Rule | Owner / Where | Brief justification | +|---|---|---| +| The authoritative record MUST live outside every owner PR branch. `authorization_ref` MUST identify a durable record in the validation integration PR body or a maintainer-controlled GitHub artifact. `authorized_by` MUST identify a ClimaMind Rumoca repository maintainer; this task's explicit maintainer authorization is sufficient; no additional maintainer or approval is required. Its required fields: `authorization_ref`, `authorized_by`, `batch_id`, authorized ordered `owner_prs`, `target_branch`, and RFC 3339 UTC `expires_at`. | Recovery authorization | Makes activation durable and explicit | +| The batch automatically becomes inactive and MUST fail closed as soon as any one of these conditions is true: its RFC 3339 `expires_at` has passed; every authorized owner PR has landed; or all required CI checks on the target `main` are green. | Maintainer | Limits the exception to broken-main recovery | +| Before each owner PR merge, the authoritative record MUST exist, match the recorded batch, PR, head, and target values, and remain unexpired; it MUST also describe an active batch under the preceding rule. A maintainer MUST verify those values against the final owner head; missing, mismatched, expired, or inactive authorization MUST fail closed. | Maintainer | Makes activation fail closed | +| Each owner PR records an independent technical review and owner mechanism test, binding both pieces of evidence to that owner PR's final `head_sha`; then exact-head integration hosted CI is green and all required hosted CI checks are green; then merge in sequence. Evidence order is mandatory: verify authorization; record the independent technical review; record the passing owner mechanism test; construct exact-head integration; record hosted CI; then merge. No later step may occur before its successful recorded predecessor. | Owner PR | Preserves ordered evidence | +| Exact-head provenance records the owner PR `head_sha` values, including the recovery-rule PR. Every listed final owner `head_sha`, including the recovery-rule PR `head_sha`, MUST be a Git ancestor of the integration `head_sha`; the recorded target baseline `head_sha` MUST be a Git ancestor of the integration `head_sha` too. Cherry-pick, patch-id, squash, or content equivalence is not exact provenance. | Integration PR | Proves the tested commits are the owner commits | +| The integration history MUST contain the target baseline, the listed exact owner histories, and signed merge commits only. Every such merge commit MUST carry exactly one `Signed-off-by` trailer and no `Co-Authored-By` trailer. It MUST NOT contain any integration-only production, test, spec, workflow, baseline, validator, tolerance, fixture, or content commit. | Integration PR | Prevents validation-only changes from manufacturing green status | +| All required hosted CI checks MUST run on the recorded integration head; the hosted CI workflow `head_sha` MUST equal the recorded integration PR `head_sha`. | Integration PR | Prevents stale CI reuse | +| Any owner PR `head_sha`, target baseline `head_sha`, or integration PR `head_sha` change MUST invalidate affected evidence and fail closed; reconstruct the integration PR, refresh affected review or mechanism-test evidence, and rerun all required hosted CI. | Maintainer | Rejects stale evidence | +| No GitHub approving review is required only for owner PRs in that active batch. Every other §6 rule remains required. | Maintainer | Keeps waiver narrow | +| Integration PR is Draft and validation-only; it MUST NEVER merge and MUST NOT contain unique fixes. | Integration PR | Keeps validation disposable | + +**PROHIBITED:** +- MUST NOT weaken or bypass any existing gate, including CI, sign-off, or sequencing. +- MUST NOT apply to third-party contributors or an unauthorized batch. +- MUST NOT merge the integration PR or place a repair only on its branch. + ### 7. Maintainability Quick Reference See SPEC_0021 for the authoritative function-length, nesting, and arg-count diff --git a/spec/SPEC_0029_CRATE_BOUNDARIES.md b/spec/SPEC_0029_CRATE_BOUNDARIES.md index 9efaa7330..c0720f996 100644 --- a/spec/SPEC_0029_CRATE_BOUNDARIES.md +++ b/spec/SPEC_0029_CRATE_BOUNDARIES.md @@ -119,7 +119,9 @@ structural-parameter values are available only after instantiation. ### 5. Evaluation Decoupled from Representation -Evaluation crates are aligned to IR ownership: `rumoca-eval-ast`, `rumoca-eval-flat`, and `rumoca-eval-dae`. `rumoca-eval-solve` builds on DAE evaluation primitives for solver-facing row evaluation. This keeps evaluation entry points explicit per representation and avoids cross-layer helper crates that hide where behavior lives. +Evaluation follows IR ownership: `rumoca-eval-ast`, `rumoca-eval-flat`, +`rumoca-eval-dae`, and solver-facing `rumoca-eval-solve`. Entry points stay +representation-specific; cross-layer helpers must not hide ownership. Phase crates MAY depend on the evaluation crate for the IR they are actively processing when the phase needs compile-time evaluation of that representation. For example, @@ -241,8 +243,10 @@ compiler/session → DAE structural → solve-IR lowering → runtime contracts | DAE structural analysis (Pantelides, BLT, tearing, demotion) | `rumoca-phase-structural` | SPEC_0007 §Structural Transformation Scope | | Solver-facing prepared data + row ops | `rumoca-ir-solve` | Backend-neutral execution IR | | DAE → solve-IR lowering | `rumoca-phase-solve` | Lowering only, not structural mutation | +| Compact Map evaluation | `rumoca-eval-solve` | Backend-neutral; no scalar reconstruction | +| GPU initialization settlement | `rumoca-sim` | Orchestration only; no solver dependency | | Optimization/training orchestration | `rumoca-opt` | Consumes Solve/eval APIs; no Modelica semantics | -| Textual generated artifacts and templates | `rumoca-phase-codegen` | Jinja/minijinja rendering owns generated C, Rust, CUDA C, MLIR, FMI/eFMI and FMU/eFMU packaging text | +| Generated text/templates | `rumoca-phase-codegen` | Renders C, Rust, CUDA, MLIR, FMI/eFMI, and packaging text | | GALEC `.alg` text (recorded exception) | `rumoca-ir-galec` | Typed AST printing per eFMI conformance; routed via template context (SPEC_0034 GAL-009) | | eFMI packaging XML (`__content.xml`, manifests) | `rumoca-phase-codegen` | Rendered like FMI `modelDescription`; validators + generic checksum/container build step, not typed serializers (SPEC_0034 D3 amended) | | Compiled/JIT execution adapter crates | `rumoca-exec-*` | Invoke tools, load artifacts, wrap Cranelift/LLVM/CUDA/NVRTC APIs, expose ergonomic runtime calls; no compiler semantics | @@ -264,19 +268,21 @@ Inkwell, LLVM ORC bindings, CUDA Driver APIs, or NVRTC; backend bytecode, native/JIT execution, and device launch policy belong in `rumoca-exec-*`, above the IR-lowering phase. -Target-language and target-format policy belongs in manifests/templates, not -Rust control flow. Rust MAY provide generic manifest parsing, template -rendering, safe path handling, schema validation, and language-neutral feature -probes over IR data. Rust MUST NOT hard-code target-language capabilities, file -layouts, emitted language names, or backend feature tables for textual targets +Target-language/format policy belongs in manifests/templates, not Rust control +flow. Rust MAY provide generic parsing, rendering, safe paths, schema validation, +and language-neutral IR probes. Rust MUST NOT hard-code target capabilities, +layouts, emitted language names, or feature tables for textual targets (C, Rust, CUDA C, MLIR, FMI/eFMI, or future custom targets). A textual/codegen -target should be addable with `target.toml` plus Jinja templates; required -capability declarations or unsupported-feature contracts must live in that -manifest schema and be enforced by generic validation. Unsupported manifest -capability failures MUST report stable `unsupported-feature:` from +target should be addable with `target.toml` plus Jinja; required capabilities +and unsupported-feature contracts must live in its schema and use generic +validation. Unsupported capability failures MUST report stable +`unsupported-feature:` from the manifest feature ID so CI, MSL reports, and release summaries can aggregate gaps without knowing the target language. +Textual `[[files]]` default non-empty. Format-required empty artifacts MUST set +`allow_empty = true` per entry; render coverage MUST enforce it. + JIT targets follow the same layering rule as execution adapters, not textual template targets. Cranelift, LLVM ORC/Inkwell, CUDA NVRTC/Driver, and browser WebAssembly compilation are allowed only in backend-facing execution crates or diff --git a/spec/SPEC_0032_RANGE_PRESERVING_TENSORS.md b/spec/SPEC_0032_RANGE_PRESERVING_TENSORS.md index 9ceb0df15..8cfd78f9f 100644 --- a/spec/SPEC_0032_RANGE_PRESERVING_TENSORS.md +++ b/spec/SPEC_0032_RANGE_PRESERVING_TENSORS.md @@ -35,12 +35,29 @@ explicitly when identity crosses phase boundaries. | View ordering is deterministic | Domain enumeration | Backend agreement | | Views carry provenance | Scalar-view metadata | Diagnostics and fallback | | No scalar-row reassembly | Solve lowering | Prevents fragile recovery | +| Unmaterialized interior rows are non-semantic placeholders | Flat/DAE structural metadata | Corner proof remains authoritative | Domains enumerate in binder declaration order, lexicographic with the innermost binder varying fastest, respecting explicit step direction. For each index tuple, body equations emit in source/body order. Scalar views must preserve parent structured/tensor id, index tuple, scalar row id, and instantiated lhs/rhs or output expression. +Function projection derives slice shape from selector kind: `:` preserves the +axis, a confirmed scalar selector removes it without evaluating its value, and +compatible elementwise binary array operands retain that shape. Function- +projection shape inference declines unknown or array-valued selectors and +ranges with unknown compile-time length. + +For a regular family whose interiors are not materialized, only the base and +per-binder neighbor rows carry the reconstruction proof. Structural rewrites of +an interior placeholder do not invalidate that proof; rewrites of a corner row +must discard the family metadata unless a new proof is produced. + +An extent-1 binder has no distinct neighbor row: the base row is sufficient and +that binder contributes zero stride to every affine access and output map. Corner +selection skips such binders while retaining base/neighbor proof for all other +dimensions. An all-singleton non-empty domain therefore uses only its base row; +an empty domain produces zero rows and no corner proof. ### 3. DAE Canonical Form @@ -49,6 +66,15 @@ lhs/rhs or output expression. | Structured DAE contains no source `der(...)` | DAE lowering | MLS Appendix B form | | Derivative families map to canonical slots | DAE structured family | Explicit state identity | | No parallel scalarized owner | DAE IR | Avoids drift | +| Orphan pruning counts exact scalar references on both equation sides | Structural phases | An explicit scalar lhs is a live owner use; a shaped slice/base lhs owns only the exact scalar projection proven by DAE dimensions and `scalar_count`; an aggregate base alias alone does not keep unrelated scalar leaves | + +DAE lowers colon-slice multiplication to a scalar dot product only when both +operands are proven rank-one vectors of equal width. Proven scalar operands, +including scalar compound expressions, retain elementwise vector scaling. A +conditional is proven scalar only when every condition, branch value, and the +else value are proven scalar; unresolved function/builtin calls, unknown widths, +and higher-rank shapes remain unprojected rather than acquiring broadcast +semantics. A source family such as `der(u[i, j]) = w[i, j]` is represented as residuals over canonical derivative slots/state metadata. The structured node owns the @@ -63,6 +89,7 @@ slot. | `ComputeNode::AffineStencil` is neighborhood access | Solve IR | Affine offset semantics | | Solve grouping is semantic | `rumoca-phase-solve` | Backends do not redefine IR | | Scalar fallback uses shared scalarization | `rumoca-eval-solve` | One ordering implementation | +| Unprojectable direct single-array-output calls stay whole | Solve projection / array runtime boundary | Never duplicate one array call as scalar lanes | `Map` represents canonical DAE residual families that are elementwise over a compact domain, including `der(u) = w` after DAE canonicalization. `AffineStencil` @@ -71,6 +98,26 @@ must not rediscover stencils by scanning anonymous scalar rows. Backends may fuse or split generated kernels as target-local codegen, but the reported kernel inventory must match the generated work. +For direct structured initialization, the same domain can pair a residual `Map` +with a compact target `TensorOutputMap`. The target map is the sole scalar-view +mapping for that family; creating parallel `row_targets`, `StructuredProgram`, +or `Vec>` ownership is forbidden on the compact path. The +`ComputeBlock` remains the sole owner of the Map; initialization metadata refers +to it by node index. `rumoca-eval-solve` executes the base program and affine +strides natively over the domain, without per-cell `LinearOp` construction. +Direct and fixed-start target ranges form an exact, non-overlapping affine +partition. Fixed-start array coverage is derived from the resolved contiguous +layout base and shape without scalar row-target materialization. Descending +source binders are normalized to an ascending execution domain by selecting the +corresponding source base and corners; target maps therefore remain canonical +positive-stride maps without changing source-index semantics. +Corner-derived load, constant, and target strides are admissible only after +Solve lowering proves the reconstructed program against every materialized +family cell. A family whose interiors are unavailable, whose values are not +affine, or whose initializer contains random/impure operations fails closed at +the first source row that breaks the proof; executing a self-consistent but +unproven affine reconstruction is forbidden. + ### 5. Ownership Boundaries | Thing | Owner/Where | Brief Justification | diff --git a/spec/SPEC_0032_DEVELOPMENT_PROCESS.md b/spec/SPEC_0033_DEVELOPMENT_PROCESS.md similarity index 99% rename from spec/SPEC_0032_DEVELOPMENT_PROCESS.md rename to spec/SPEC_0033_DEVELOPMENT_PROCESS.md index 25b9e2a1d..db62a9f6d 100644 --- a/spec/SPEC_0032_DEVELOPMENT_PROCESS.md +++ b/spec/SPEC_0033_DEVELOPMENT_PROCESS.md @@ -1,4 +1,4 @@ -# SPEC_0032: Development Process +# SPEC_0033: Development Process ## Status ACCEPTED diff --git a/toolchains/openmodelica-packages.txt b/toolchains/openmodelica-packages.txt new file mode 100644 index 000000000..241fe01c4 --- /dev/null +++ b/toolchains/openmodelica-packages.txt @@ -0,0 +1,5 @@ +# source_commit d7e2907f419d8061ce0a461ac9709dcb84fa70a1 +omc 1.27.0~1-gd7e2907-1 amd64 omc_1.27.0~1-gd7e2907-1_amd64.deb a4511acb19f7377275347f9fc27af92c307db0903e7f46e44085410a060d86ab https://build.openmodelica.org/apt/pool/contrib-noble/omc_1.27.0~1-gd7e2907-1_amd64.deb +omc-common 1.27.0~1-gd7e2907-1 all omc-common_1.27.0~1-gd7e2907-1_all.deb 9781f4efa44274ecff4538227ef8054b10d3b9abe6cf3d36bb9b23237c0b1119 https://build.openmodelica.org/apt/pool/contrib-noble/omc-common_1.27.0~1-gd7e2907-1_all.deb +libomc 1.27.0~1-gd7e2907-1 amd64 libomc_1.27.0~1-gd7e2907-1_amd64.deb 90b650ecf9da174cd477115a8f61d54f8bf22b6c87ba1707cb5482d933ae3f5f https://build.openmodelica.org/apt/pool/contrib-noble/libomc_1.27.0~1-gd7e2907-1_amd64.deb +libomcsimulation 1.27.0~1-gd7e2907-1 amd64 libomcsimulation_1.27.0~1-gd7e2907-1_amd64.deb bdc2c5c307aacf516017f8b64fc4878e6a75d7b0587366d2b93ec8ae3590632d https://build.openmodelica.org/apt/pool/contrib-noble/libomcsimulation_1.27.0~1-gd7e2907-1_amd64.deb diff --git a/toolchains/openmodelica-version b/toolchains/openmodelica-version new file mode 100644 index 000000000..ee5338121 --- /dev/null +++ b/toolchains/openmodelica-version @@ -0,0 +1 @@ +1.27.0~1-gd7e2907-1