// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package integrate import ( "sourcedock.dev/petrbalvin/tensor/internal/base" "sourcedock.dev/petrbalvin/tensor/internal/core" ) import ( "math" "testing" ) // decay returns f for y' = −y, the reference every scheme must nail. func decay(t float64, y *core.Array) (*core.Array, error) { return core.MulF(y, -1), nil } // TestIntegrateODEExponential checks the adaptive solver against the // analytic exponential decay, forward and backward in time. func TestIntegrateODEExponential(t *testing.T) { y0 := mustFloats(t, []float64{1}) end, err := IntegrateODE(decay, 0, 1, y0, ODEOptions{RelTol: 1e-10, AbsTol: 1e-12}) if err != nil { t.Fatalf("IntegrateODE: %v", err) } if math.Abs(end.FloatAt(0)-math.Exp(-1)) > 1e-9 { t.Fatalf("y(1) = %.14g, want %.14g", end.FloatAt(0), math.Exp(-1)) } // Backward integration from t=1 to t=0 must return the start. start, err := IntegrateODE(decay, 1, 0, mustFloats(t, []float64{math.Exp(-1)}), ODEOptions{RelTol: 1e-10, AbsTol: 1e-12}) if err != nil { t.Fatalf("IntegrateODE backward: %v", err) } if math.Abs(start.FloatAt(0)-1) > 1e-8 { t.Fatalf("backward y(0) = %.14g, want 1", start.FloatAt(0)) } } // TestIntegrateODEOscillator checks a two-dimensional linear system // against the analytic phase rotation. func TestIntegrateODEOscillator(t *testing.T) { f := func(t float64, y *core.Array) (*core.Array, error) { return core.FromFloats([]float64{y.FloatAt(1), -y.FloatAt(0)}, 2) } y0 := mustFloats(t, []float64{1, 0}) end, err := IntegrateODE(f, 0, math.Pi/2, y0, ODEOptions{RelTol: 1e-11, AbsTol: 1e-13}) if err != nil { t.Fatalf("IntegrateODE: %v", err) } // y1 = cos t, y2 = −sin t, so a quarter period lands on (0, −1). if math.Abs(end.FloatAt(0)) > 1e-7 || math.Abs(end.FloatAt(1)+1) > 1e-7 { t.Fatalf("quarter period = (%.10g, %.10g), want (0, -1)", end.FloatAt(0), end.FloatAt(1)) } // A full period returns to the start. full, err := IntegrateODE(f, 0, 2*math.Pi, y0, ODEOptions{}) if err != nil { t.Fatalf("IntegrateODE full: %v", err) } for i := range 2 { if math.Abs(full.FloatAt(i)-y0.FloatAt(i)) > 1e-6 { t.Fatalf("full period[%d] = %v, want %v", i, full.FloatAt(i), y0.FloatAt(i)) } } } // TestIntegrateRk4 checks the fixed-step scheme's fourth-order // convergence: halving h must shrink the error by roughly sixteen. func TestIntegrateRk4(t *testing.T) { errAt := func(steps int) float64 { end, err := IntegrateRK4(decay, 0, 1, mustFloats(t, []float64{1}), steps) if err != nil { t.Fatalf("IntegrateRK4(%d): %v", steps, err) } return math.Abs(end.FloatAt(0) - math.Exp(-1)) } e10, e20 := errAt(10), errAt(20) if e10 < 1e-13 { t.Skipf("error already at round-off (%v)", e10) } ratio := e10 / e20 if ratio < 12 || ratio > 20 { t.Fatalf("error ratio over a halved step = %.2g, want ≈ 16 for a fourth-order scheme", ratio) } } // TestIntegrateBackwardEulerStiff demonstrates the reason an implicit // scheme exists: on y' = −1000y the implicit Euler stays bounded and // matches its closed-form damping at a step where explicit schemes // blow up. func TestIntegrateBackwardEulerStiff(t *testing.T) { const lambda = 1000.0 f := func(t float64, y *core.Array) (*core.Array, error) { return core.MulF(y, -lambda), nil } // Ten steps of h = 0.01: h·λ = 10, far outside RK4's stability // region but perfectly damped for the implicit scheme. y0 := mustFloats(t, []float64{1}) end, err := IntegrateBackwardEuler(f, 0, 0.1, y0, 10, ODEOptions{}) if err != nil { t.Fatalf("IntegrateBackwardEuler: %v", err) } // Closed form of one implicit Euler step: y_{n+1} = y_n/(1+hλ), // so y_10 = (1/11)^10. damped := math.Pow(1/(1+0.01*lambda), 10) if math.Abs(end.FloatAt(0)-damped) > 1e-9*math.Abs(damped) { t.Fatalf("stiff result = %.12g, want %.12g", end.FloatAt(0), damped) } if end.FloatAt(0) <= 0 { t.Fatalf("the implicit scheme must stay positive on decay, got %v", end.FloatAt(0)) } } // TestIntegrateBackwardEulerAccuracy checks that with a sane step the // implicit scheme tracks the analytic decay as well. func TestIntegrateBackwardEulerAccuracy(t *testing.T) { end, err := IntegrateBackwardEuler(decay, 0, 1, mustFloats(t, []float64{1}), 1000, ODEOptions{}) if err != nil { t.Fatalf("IntegrateBackwardEuler: %v", err) } // First order: the global error is O(h), about h/2 for decay. if math.Abs(end.FloatAt(0)-math.Exp(-1)) > 1e-3 { t.Fatalf("y(1) = %.14g, want %.14g ± 1e-3", end.FloatAt(0), math.Exp(-1)) } } // TestIntegrateBackwardEulerVectorState pins the Newton solve on a // two-dimensional stiff system: both modes decay at their own rate, // and constant steps give each component the closed form // y(t1) = y0·(1+h·λ)^{−steps} an implicit Euler step has on y' = −λy. func TestIntegrateBackwardEulerVectorState(t *testing.T) { const slow = 1.0 const fast = 1000.0 f := func(t float64, y *core.Array) (*core.Array, error) { return core.FromFloats([]float64{-fast * y.FloatAt(0), -slow * y.FloatAt(1)}, 2) } y0 := mustFloats(t, []float64{1, 1}) end, err := IntegrateBackwardEuler(f, 0, 1, y0, 10, ODEOptions{}) if err != nil { t.Fatalf("IntegrateBackwardEuler: %v", err) } wantFast := math.Pow(1/(1+0.1*fast), 10) if math.Abs(end.FloatAt(0)-wantFast) > 1e-9*math.Abs(wantFast) { t.Fatalf("fast mode = %.14g, want %.14g", end.FloatAt(0), wantFast) } wantSlow := math.Pow(1/(1+0.1*slow), 10) if math.Abs(end.FloatAt(1)-wantSlow) > 1e-9*math.Abs(wantSlow) { t.Fatalf("slow mode = %.14g, want %.14g", end.FloatAt(1), wantSlow) } } // TestODEErrors pins the error contracts shared by the solvers. func TestODEErrors(t *testing.T) { wrongShape := func(t float64, y *core.Array) (*core.Array, error) { return core.FromFloats([]float64{1, 1}, 2) } y0 := mustFloats(t, []float64{1}) if _, err := IntegrateODE(wrongShape, 0, 1, y0, ODEOptions{}); err == nil { t.Fatal("expected an error when f returns the wrong shape") } if _, err := IntegrateODE(wrongShape, 0, 1, y0, ODEOptions{}); err == nil { t.Fatal("expected an error when f returns the wrong shape in RK4 as well") } if _, err := IntegrateRK4(decay, 0, 1, y0, 0); err == nil { t.Fatal("expected an error for zero steps") } matrixState := mustFloats(t, []float64{1, 1}, 1, 2) if _, err := IntegrateODE(decay, 0, 1, matrixState, ODEOptions{}); err == nil { t.Fatal("expected an error for a rank-2 state") } empty := mustFloats(t, nil) if _, err := IntegrateODE(decay, 0, 1, empty, ODEOptions{}); err == nil { t.Fatal("expected an error for an empty state") } // A tight budget on a slow decay must report the budget, not lie. if _, err := IntegrateODE(decay, 0, 1, y0, ODEOptions{MaxSteps: 3}); err == nil { t.Fatal("expected an error for an exhausted step budget") } } // TestIntegrateODEPathExponential checks the sampled trajectory of // y' = −y against the analytic decay at every sample point. func TestIntegrateODEPathExponential(t *testing.T) { const n = 5 times, states, err := IntegrateODEPath(decay, 0, 2, mustFloats(t, []float64{1}), n, ODEOptions{RelTol: 1e-10, AbsTol: 1e-12}) if err != nil { t.Fatalf("IntegrateODEPath: %v", err) } if len(times) != n || len(states) != n { t.Fatalf("lengths (%d, %d), want (%d, %d)", len(times), len(states), n, n) } for i := range n { want := 0.5 * float64(i) if math.Abs(times[i]-want) > 1e-12 { t.Fatalf("times[%d] = %.14g, want %.14g", i, times[i], want) } got := states[i].FloatAt(0) if math.Abs(got-math.Exp(-want)) > 1e-9 { t.Fatalf("y(%.1f) = %.14g, want %.14g", want, got, math.Exp(-want)) } } } // TestIntegrateODEPathBackward samples a backwards integration: the // times descend and the analytic law still holds per sample. func TestIntegrateODEPathBackward(t *testing.T) { times, states, err := IntegrateODEPath(decay, 2, 0, mustFloats(t, []float64{math.Exp(-2)}), 3, ODEOptions{RelTol: 1e-10, AbsTol: 1e-12}) if err != nil { t.Fatalf("IntegrateODEPath backward: %v", err) } for i, want := range []float64{2, 1, 0} { if math.Abs(times[i]-want) > 1e-12 { t.Fatalf("times[%d] = %.14g, want %g", i, times[i], want) } if math.Abs(states[i].FloatAt(0)-math.Exp(-want)) > 1e-9 { t.Fatalf("y(%g) = %.14g, want %.14g", want, states[i].FloatAt(0), math.Exp(-want)) } } } // TestIntegrateODEPathStatesIndependent checks the returned states do // not alias one another or the initial vector: a later integration // must never rewrite an earlier sample. func TestIntegrateODEPathStatesIndependent(t *testing.T) { y0 := mustFloats(t, []float64{1, 0}) f := func(t float64, y *core.Array) (*core.Array, error) { return core.FromFloats([]float64{y.FloatAt(1), -y.FloatAt(0)}, 2) } _, states, err := IntegrateODEPath(f, 0, 1, y0, 4, ODEOptions{}) if err != nil { t.Fatalf("IntegrateODEPath: %v", err) } states[3].RawFloats()[0] = 99 // must not leak into y0 or the other samples if y0.FloatAt(0) != 1 { t.Fatal("mutating a sample changed the caller's initial state") } if states[2].FloatAt(0) == 99 { t.Fatal("samples alias one another") } } // TestIntegrateODEPathErrors pins the path-specific error contract. func TestIntegrateODEPathErrors(t *testing.T) { y0 := mustFloats(t, []float64{1}) if _, _, err := IntegrateODEPath(decay, 0, 1, y0, 1, ODEOptions{}); err == nil { t.Fatal("expected an error for a single sample") } // An f that fails past the midpoint must fail the whole path. boom := func(t float64, y *core.Array) (*core.Array, error) { if t > 0.6 { return nil, base.Errf("detector tripped") } return core.MulF(y, -1), nil } if _, _, err := IntegrateODEPath(boom, 0, 1, y0, 5, ODEOptions{}); err == nil { t.Fatal("expected the operator error to propagate") } if _, _, err := IntegrateODEPath(decay, 0, 1, mustFloats(t, nil), 3, ODEOptions{}); err == nil { t.Fatal("expected an error for an empty state") } } // TestIntegrateODEEventsProjectile drops a projectile and watches its // height: the crossing of zero is the flight time 2v₀/g, analytically // known, with the state's velocity the exact mirror of the launch. func TestIntegrateODEEventsProjectile(t *testing.T) { const g = 9.81 const v0 = 10 f := func(now float64, y *core.Array) (*core.Array, error) { return mustFloats(t, []float64{y.FloatAt(1), -g}), nil } height := func(now float64, y *core.Array) (float64, error) { return y.FloatAt(0), nil } hits, final, err := IntegrateODEEvents(f, 0, 5, mustFloats(t, []float64{0, v0}), []ODEWatch{{Function: height, Direction: -1}}, ODEOptions{RelTol: 1e-12, AbsTol: 1e-14}) if err != nil { t.Fatalf("IntegrateODEEvents: %v", err) } if len(hits) != 1 { t.Fatalf("hits = %d, want 1", len(hits)) } wantT := 2 * v0 / g if math.Abs(hits[0].Time-wantT) > 1e-9 { t.Fatalf("impact at %.14g, want %.14g", hits[0].Time, wantT) } if math.Abs(hits[0].State.FloatAt(1)+v0) > 1e-7 { t.Fatalf("impact speed %.10g, want %.10g", hits[0].State.FloatAt(1), float64(-v0)) } if hits[0].Rising { t.Fatal("height crossing on the way down must not be rising") } if final.Len() != 2 { t.Fatalf("final state shape %v", final.Shape()) } } // TestIntegrateODEEventsOscillator watches the oscillator's position // over five periods: cos crosses zero once per half period, and the // falling-only filter keeps every second one, at π/2 + 2πk. func TestIntegrateODEEventsOscillator(t *testing.T) { f := func(now float64, y *core.Array) (*core.Array, error) { return mustFloats(t, []float64{y.FloatAt(1), -y.FloatAt(0)}), nil } position := func(now float64, y *core.Array) (float64, error) { return y.FloatAt(0), nil } hits, _, err := IntegrateODEEvents(f, 0, 5*2*math.Pi, mustFloats(t, []float64{1, 0}), []ODEWatch{{Function: position, Direction: -1}}, ODEOptions{RelTol: 1e-12, AbsTol: 1e-14}) if err != nil { t.Fatalf("IntegrateODEEvents: %v", err) } if len(hits) != 5 { t.Fatalf("hits = %d, want 5 falling crossings", len(hits)) } for k, hit := range hits { want := math.Pi/2 + float64(2*k)*math.Pi if math.Abs(hit.Time-want) > 1e-8 { t.Fatalf("hit %d at %.12g, want %.12g", k, hit.Time, want) } } // Without the filter every half period fires: ten crossings. hits, _, err = IntegrateODEEvents(f, 0, 5*2*math.Pi, mustFloats(t, []float64{1, 0}), []ODEWatch{{Function: position}}, ODEOptions{RelTol: 1e-12, AbsTol: 1e-14}) if err != nil { t.Fatalf("IntegrateODEEvents unfiltered: %v", err) } if len(hits) != 10 { t.Fatalf("unfiltered hits = %d, want 10", len(hits)) } for i := 1; i < len(hits); i++ { if hits[i].Time <= hits[i-1].Time { t.Fatal("hits are not sorted by time") } } } // TestIntegrateODEStepsErrors pins the error contract of the step // recorder. func TestIntegrateODEStepsErrors(t *testing.T) { y0 := mustFloats(t, []float64{1}) if _, _, err := IntegrateODESteps(decay, 0, 1, y0, ODEOptions{MaxSteps: 2}); err == nil { t.Fatal("expected an error for an exhausted step budget") } if _, _, err := IntegrateODESteps(decay, 0, 1, mustFloats(t, nil), ODEOptions{}); err == nil { t.Fatal("expected an error for an empty state") } } // TestIntegrateODEStepsForward records the accepted steps of a decay: // the trace starts at the initial state, ends exactly at t1, and every // node sits on the analytic curve at solver tolerance. func TestIntegrateODEStepsForward(t *testing.T) { times, states, err := IntegrateODESteps(decay, 0, 1, mustFloats(t, []float64{1}), ODEOptions{RelTol: 1e-8, AbsTol: 1e-12}) if err != nil { t.Fatalf("IntegrateODESteps: %v", err) } if len(times) != len(states) || len(times) < 3 { t.Fatalf("trace has %d times and %d states, want matching lengths of at least 3", len(times), len(states)) } if times[0] != 0 || times[len(times)-1] != 1 { t.Fatalf("trace spans [%g, %g], want [0, 1]", times[0], times[len(times)-1]) } for i := 1; i < len(times); i++ { if times[i] <= times[i-1] { t.Fatalf("times not strictly increasing at %d: %g after %g", i, times[i], times[i-1]) } } if states[0].FloatAt(0) != 1 { t.Fatalf("first state = %v, want the initial state 1", states[0].FloatAt(0)) } for i := range times { if math.Abs(states[i].FloatAt(0)-math.Exp(-times[i])) > 1e-6 { t.Fatalf("y(%g) = %.14g, want %.14g", times[i], states[i].FloatAt(0), math.Exp(-times[i])) } } } // TestIntegrateODEStepsDegenerate covers the zero-span trace and a // backward pass with descending times. func TestIntegrateODEStepsDegenerate(t *testing.T) { times, states, err := IntegrateODESteps(decay, 1, 1, mustFloats(t, []float64{1}), ODEOptions{}) if err != nil { t.Fatalf("IntegrateODESteps zero span: %v", err) } if len(times) != 1 || times[0] != 1 || states[0].FloatAt(0) != 1 { t.Fatalf("zero-span trace = %v, want the single initial node", times) } times, states, err = IntegrateODESteps(decay, 1, 0, mustFloats(t, []float64{math.Exp(-1)}), ODEOptions{RelTol: 1e-8, AbsTol: 1e-12}) if err != nil { t.Fatalf("IntegrateODESteps backward: %v", err) } if times[0] != 1 || times[len(times)-1] != 0 { t.Fatalf("backward trace spans [%g, %g], want [1, 0]", times[0], times[len(times)-1]) } for i := 1; i < len(times); i++ { if times[i] >= times[i-1] { t.Fatalf("backward times not strictly decreasing at %d", i) } } } // TestIntegrateODEEventsErrors pins the validation and error paths. func TestIntegrateODEEventsErrors(t *testing.T) { f := decay if _, _, err := IntegrateODEEvents(f, 0, 1, mustFloats(t, []float64{1}), nil, ODEOptions{}); err == nil { t.Fatal("no watches: want an error") } if _, _, err := IntegrateODEEvents(f, 0, 1, mustFloats(t, []float64{1}), []ODEWatch{{}}, ODEOptions{}); err == nil { t.Fatal("empty watch: want an error") } boom := func(t float64, y *core.Array) (float64, error) { return 0, base.Errf("watch failed") } if _, _, err := IntegrateODEEvents(f, 0, 1, mustFloats(t, []float64{1}), []ODEWatch{{Function: boom}}, ODEOptions{}); err == nil { t.Fatal("watch error: want an error") } // An f that fails on the refinement path must surface. broken := func(t float64, y *core.Array) (*core.Array, error) { if t > 0.5 { return nil, base.Errf("integrand failed") } return core.MulF(y, -1), nil } cross := func(t float64, y *core.Array) (float64, error) { return 1 - t, nil } if _, _, err := IntegrateODEEvents(broken, 0, 2, mustFloats(t, []float64{1}), []ODEWatch{{Function: cross}}, ODEOptions{}); err == nil { t.Fatal("refinement error: want an error") } } // TestODEArrivalGuard pins the ulp-level arrival rule: an endpoint // missed by rounding terminates the loop, a genuine remaining distance // does not, and a plain integration across magnitudes still lands. func TestODEArrivalGuard(t *testing.T) { if !odeArrived(1, 1) { t.Fatal("equal times must count as arrived") } if !odeArrived(1, 1+4*2.220446049250313e-16) { t.Fatal("a four-ulp miss must count as arrived") } if odeArrived(1, 1.1) { t.Fatal("a genuine remaining distance must not count as arrived") } if odeArrived(0, -1e-20) { t.Fatal("a tiny but representable distance near zero must not count as arrived") } // An integration whose endpoints differ well beyond Sterbenz still // terminates and reports the endpoint state. f := func(tt float64, y *core.Array) (*core.Array, error) { return core.FromFloats([]float64{y.FloatAt(0) * 0}, 1) } y0, _ := core.FromFloats([]float64{3}, 1) got, err := IntegrateODE(f, 1e9, 1e9+0.5, y0, ODEOptions{}) if err != nil { t.Fatalf("IntegrateODE across magnitudes: %v", err) } if got.FloatAt(0) != 3 { t.Fatalf("constant state changed: %g", got.FloatAt(0)) } }