diff --git a/vecxt/src-js/doublematrix.scala b/vecxt/src-js/doublematrix.scala index 7dcc0057..691f8730 100644 --- a/vecxt/src-js/doublematrix.scala +++ b/vecxt/src-js/doublematrix.scala @@ -303,6 +303,10 @@ object JsDoubleMatrix: out end * + /** Per-column fused multiply-add: `out(i, j) = m(i, j) * multiply(j) + add(j)`, as a fresh dense column-major + * matrix. See `JvmDoubleMatrix.fmaCols`; this platform uses the shared scalar loop. + */ + def fmaCols(multiply: Array[Double], add: Array[Double]): Matrix[Double] = FmaCols.loop(m, multiply, add) end extension end JsDoubleMatrix diff --git a/vecxt/src-jvm/doublematrix.scala b/vecxt/src-jvm/doublematrix.scala index 38963707..c746a622 100644 --- a/vecxt/src-jvm/doublematrix.scala +++ b/vecxt/src-jvm/doublematrix.scala @@ -384,6 +384,54 @@ object JvmDoubleMatrix: end += + /** Per-column fused multiply-add: `out(i, j) = m(i, j) * multiply(j) + add(j)`, as a fresh dense column-major + * matrix. + * + * A column scale and a column broadcast-add in one pass: `m` is read once and the result written once, where + * `copy`, `+= arr` and a column scale would each make a full pass. Reference BLAS has no equivalent - `dscal` per + * column scales but cannot add, and `A * diag(x)` routines (`dgmm`) likewise stop at the multiply. + * + * When `rowStride == 1` each column is a contiguous run and is processed with `DoubleVector.fma` against two + * broadcast lanes; any other layout (row-major, transposed or doubly-strided views) falls back to the shared + * scalar loop in [[FmaCols.loop]]. The input is never modified. + * + * @param multiply + * one multiplier per column; must have length `m.cols` + * @param add + * one addend per column; must have length `m.cols` + */ + def fmaCols(multiply: Array[Double], add: Array[Double]): Matrix[Double] = + if m.rowStride != 1 then FmaCols.loop(m, multiply, add) + else + FmaCols.check(m, multiply, add) + val spd = doublearrays.spd + val spdl = doublearrays.spdl + val rows = m.rows + val bound = spd.loopBound(rows) + val out = new Array[Double](m.numel) + var j = 0 + while j < m.cols do + val src = m.offset + j * m.colStride + val dst = j * rows + val mulS = multiply(j) + val addS = add(j) + val mul = DoubleVector.broadcast(spd, mulS) + val ad = DoubleVector.broadcast(spd, addS) + var i = 0 + while i < bound do + DoubleVector.fromArray(spd, m.raw, src + i).fma(mul, ad).intoArray(out, dst + i) + i += spdl + end while + while i < rows do + out(dst + i) = Math.fma(m.raw(src + i), mulS, addS) + i += 1 + end while + j += 1 + end while + Matrix[Double](out, m.rows, m.cols) + end if + end fmaCols + def +=(n: Double): Unit = if m.hasSimpleContiguousMemoryLayout then vecxt.doublearrays.+=(m.raw)(n) diff --git a/vecxt/src-native/doublematrix_native.scala b/vecxt/src-native/doublematrix_native.scala index b4fdcd80..6883922b 100644 --- a/vecxt/src-native/doublematrix_native.scala +++ b/vecxt/src-native/doublematrix_native.scala @@ -308,6 +308,11 @@ object NativeDoubleMatrix: m.*=(vec, out, alpha, 0.0) out end * + + /** Per-column fused multiply-add: `out(i, j) = m(i, j) * multiply(j) + add(j)`, as a fresh dense column-major + * matrix. See `JvmDoubleMatrix.fmaCols`; this platform uses the shared scalar loop. + */ + def fmaCols(multiply: Array[Double], add: Array[Double]): Matrix[Double] = FmaCols.loop(m, multiply, add) end extension end NativeDoubleMatrix diff --git a/vecxt/src/fmaCols.scala b/vecxt/src/fmaCols.scala new file mode 100644 index 00000000..3ad060a9 --- /dev/null +++ b/vecxt/src/fmaCols.scala @@ -0,0 +1,47 @@ +package vecxt + +import vecxt.matrix.* + +/** Shared pieces of `fmaCols`: `out(i, j) = m(i, j) * multiply(j) + add(j)`. + * + * Each platform exposes the public `fmaCols` extension (`JvmDoubleMatrix`, `JsDoubleMatrix`, `NativeDoubleMatrix`). JS + * and Native delegate straight to [[loop]]; the JVM uses a SIMD path when rows are contiguous and falls back to + * [[loop]] otherwise. Deliberately not exported from `vecxt.all` - it is the kernel, not the API. + */ +private[vecxt] object FmaCols: + + /** Throws unless both vectors have one entry per column of `m`. Message built only on the failing path. */ + inline def check(m: Matrix[Double], multiply: Array[Double], add: Array[Double]): Unit = + if multiply.length != m.cols then + throw new IllegalArgumentException(s"multiply length ${multiply.length} != expected ${m.cols}") + end if + if add.length != m.cols then throw new IllegalArgumentException(s"add length ${add.length} != expected ${m.cols}") + end if + end check + + /** Layout-agnostic kernel: reads through the strides, writes a fresh dense column-major matrix. + * + * Columns outer, rows inner, so the write is always sequential. Uses a plain `x * mul + ad` rather than `Math.fma`, + * which is emulated (slowly) on Scala.js; results may therefore differ from the JVM SIMD path in the last bit. + */ + def loop(m: Matrix[Double], multiply: Array[Double], add: Array[Double]): Matrix[Double] = + check(m, multiply, add) + val rows = m.rows + val out = new Array[Double](m.numel) + var j = 0 + while j < m.cols do + val src = m.offset + j * m.colStride + val dst = j * rows + val mul = multiply(j) + val ad = add(j) + var i = 0 + while i < rows do + out(dst + i) = m.raw(src + i * m.rowStride) * mul + ad + i += 1 + end while + j += 1 + end while + Matrix[Double](out, m.rows, m.cols) + end loop + +end FmaCols diff --git a/vecxt/test/src/fmaCols.test.scala b/vecxt/test/src/fmaCols.test.scala new file mode 100644 index 00000000..3ab12c0b --- /dev/null +++ b/vecxt/test/src/fmaCols.test.scala @@ -0,0 +1,80 @@ +package vecxt + +import all.* + +/** `fmaCols` is `out(i, j) = m(i, j) * multiply(j) + add(j)`. The JVM has a SIMD path for `rowStride == 1` and falls + * back to the shared scalar loop otherwise; JS and Native always use the loop. Every case is checked against a naive + * reference read through `m(i, j)`, so the layouts that pick different branches are held to the same answer. + */ +class FmaColsSuite extends munit.FunSuite: + + private def reference(m: Matrix[Double], multiply: Array[Double], add: Array[Double]): Matrix[Double] = + val out = Array.ofDim[Double](m.rows * m.cols) + for j <- 0 until m.cols; i <- 0 until m.rows do out(i + j * m.rows) = m(i, j) * multiply(j) + add(j) + end for + Matrix[Double](out, m.rows, m.cols) + end reference + + private def assertDenseColMajor(m: Matrix[Double])(using munit.Location): Unit = + assertEquals((m.rowStride, m.colStride, m.offset), (1, m.rows, 0)) + + private def check(m: Matrix[Double], multiply: Array[Double], add: Array[Double])(using munit.Location): Unit = + val out = m.fmaCols(multiply, add) + assertDenseColMajor(out) + assertMatrixEquals(out, reference(m, multiply, add)) + end check + + test("small hand computed example"): + val m = Matrix.fromRows[Double](Array(1.0, 2.0), Array(3.0, 4.0)) + val out = m.fmaCols(Array(10.0, -1.0), Array(0.5, 100.0)) + assertMatrixEquals(out, Matrix.fromRows[Double](Array(10.5, 98.0), Array(30.5, 96.0))) + + test("dense column-major, row count not a multiple of common SIMD widths"): + val m = Matrix[Double](Array.tabulate(37 * 3)(i => i * 0.25 - 3.0), 37, 3) + check(m, Array(2.0, -0.5, 0.0), Array(1.0, 7.0, -3.0)) + + test("dense column-major, large enough for many SIMD iterations per column"): + val m = Matrix[Double](Array.tabulate(1000 * 4)(i => math.sin(i.toDouble)), 1000, 4) + check(m, Array(1.5, -2.0, 3.25, 1e6), Array(-1.0, 0.0, 2.0, 1e-3)) + + test("row-major input"): + val m = Matrix[Double](Array.tabulate(5 * 3)(_.toDouble), 5, 3, 3, 1, 0) + assertEquals(m.rowStride, 3) + check(m, Array(1.0, 2.0, 3.0), Array(-1.0, -2.0, -3.0)) + + test("padded column stride (rowStride 1, not contiguous)"): + // rows=2, cols=2, colStride=3: columns are raw(0..1) and raw(3..4), raw(2)/raw(5) are padding + val m = Matrix[Double](Array(1.0, 2.0, 99.0, 3.0, 4.0, 99.0), 2, 2, 1, 3, 0) + assert(!m.hasSimpleContiguousMemoryLayout) + val out = m.fmaCols(Array(2.0, 3.0), Array(1.0, 1.0)) + assertMatrixEquals(out, Matrix.fromRows[Double](Array(3.0, 10.0), Array(5.0, 13.0))) + + test("submatrix view with offset"): + val base = Matrix[Double](Array.tabulate(20 * 6)(_.toDouble), 20, 6) + // rows 3..17 of columns 2..4 of a 20 x 6 column-major matrix + val view = Matrix[Double](base.raw, 15, 3, 1, 20, 3 + 2 * 20) + assertEquals(view(0, 0), 43.0) + check(view, Array(-1.0, 0.5, 2.0), Array(10.0, 20.0, 30.0)) + + test("input is not modified"): + val raw = Array.tabulate(9 * 2)(_.toDouble) + val m = Matrix[Double](raw.clone, 9, 2) + m.fmaCols(Array(3.0, 4.0), Array(1.0, 2.0)) + assertVecEquals(m.raw, raw) + + test("empty matrices"): + val noRows = Matrix[Double](Array.empty[Double], 0, 3) + val out = noRows.fmaCols(Array(1.0, 2.0, 3.0), Array(1.0, 2.0, 3.0)) + assertEquals((out.rows, out.cols), (0, 3)) + val noCols = Matrix[Double](Array.empty[Double], 4, 0) + val out2 = noCols.fmaCols(Array.empty[Double], Array.empty[Double]) + assertEquals((out2.rows, out2.cols), (4, 0)) + + test("wrong vector lengths throw"): + val m = Matrix[Double](Array.fill(6)(1.0), 3, 2) + intercept[IllegalArgumentException](m.fmaCols(Array(1.0), Array(1.0, 2.0))) + intercept[IllegalArgumentException](m.fmaCols(Array(1.0, 2.0), Array(1.0, 2.0, 3.0))) + val rowMajor = Matrix[Double](Array.fill(6)(1.0), 3, 2, 2, 1, 0) + intercept[IllegalArgumentException](rowMajor.fmaCols(Array(1.0), Array(1.0, 2.0))) + +end FmaColsSuite diff --git a/vecxt_re/src/CurrencyValue.scala b/vecxt_re/src/CurrencyValue.scala new file mode 100644 index 00000000..7e47b652 --- /dev/null +++ b/vecxt_re/src/CurrencyValue.scala @@ -0,0 +1,75 @@ +package vecxt_re + +enum Ccy derives CanEqual: + case USD, EUR, GBP, CHF, JPY, AUD, CAD, NZD +end Ccy + +/** The scale a monetary amount is quoted in. `multiplier` is the number of major currency units in one of this unit. */ +enum AmountUnit(val multiplier: Long): + case One extends AmountUnit(1L) + case K extends AmountUnit(1_000L) + case Mn extends AmountUnit(1_000_000L) + case Bn extends AmountUnit(1_000_000_000L) +end AmountUnit + +/** This should _NOT_ be used for actual cash values. We use double here because this intended for financial simulation + * where errors of <0.01 are meaningless. An absolute monetary quantity denominated in a single currency, e.g. a + * notional, limit, premium or loss. + * + * Amounts in different currencies are never implicitly combined: arithmetic between currencies fails, and conversion + * requires an explicit FX rate. Scaling by a [[Rel]] yields a CurrencyAmount (premium = notional × coupon); dividing + * two amounts in the same currency yields a [[Rel]]. + * + * Stored as a Double in major currency units; [[AmountUnit]] only affects construction and display. + */ +final case class CurrencyAmount private (amount: Double, ccy: Ccy) derives CanEqual: + + /** The amount expressed in unit `u`, e.g. `CurrencyAmount(2.5, Mn, USD).in(K) == 2500.0` */ + def in(u: AmountUnit): Double = amount / u.multiplier + + def +(y: CurrencyAmount): CurrencyAmount = + sameCcy(y); new CurrencyAmount(amount + y.amount, ccy) + end + + + def -(y: CurrencyAmount): CurrencyAmount = + sameCcy(y); new CurrencyAmount(amount - y.amount, ccy) + end - + + def unary_- : CurrencyAmount = new CurrencyAmount(-amount, ccy) + + /** premium = notional * coupon */ + def *(r: Rel): CurrencyAmount = + new CurrencyAmount(amount * r.canonical, ccy) + + def *(k: Double): CurrencyAmount = + new CurrencyAmount(amount * k, ccy) + + /** Ratio of two same-currency amounts, e.g. a loss ratio. */ + def /(y: CurrencyAmount): Rel = + sameCcy(y); Rel(amount / y.amount, PriceUnit.One) + end / + + def approxEq(y: CurrencyAmount, tol: Double = 1e-6): Boolean = + sameCcy(y); math.abs(amount - y.amount) <= tol + end approxEq + + override def toString: String = f"$amount%.2f $ccy" + + private def sameCcy(y: CurrencyAmount): Unit = + require(ccy == y.ccy, s"Currency mismatch: $ccy vs ${y.ccy}") +end CurrencyAmount + +object CurrencyAmount: + /** e.g. `CurrencyAmount(2.5, AmountUnit.Mn, Ccy.USD)` is 2,500,000 USD */ + def apply(amount: Double, unit: AmountUnit, ccy: Ccy): CurrencyAmount = + new CurrencyAmount(amount * unit.multiplier, ccy) + + def zero(ccy: Ccy): CurrencyAmount = new CurrencyAmount(0.0, ccy) + + given Ordering[CurrencyAmount] with + def compare(a: CurrencyAmount, b: CurrencyAmount): Int = + require(a.ccy == b.ccy, s"Cannot order ${a.ccy} vs ${b.ccy}") + java.lang.Double.compare(a.amount, b.amount) + end compare + end given +end CurrencyAmount diff --git a/vecxt_re/src/NamedRecord.scala b/vecxt_re/src/NamedRecord.scala new file mode 100644 index 00000000..842e19e2 --- /dev/null +++ b/vecxt_re/src/NamedRecord.scala @@ -0,0 +1,55 @@ +package vecxt_re + +import scala.NamedTuple.{AnyNamedTuple, DropNames, NamedTuple, Names} +import scala.compiletime.{constValue, constValueTuple, erasedValue, error, summonFrom} + +/** Compile-time helpers for accepting a named tuple by field name rather than by field position. + * + * A caller may supply `(id = "1", price = p, notional = n, name = "A")` where the required shape is + * `(price: Price, notional: CurrencyAmount, name: String, id: String)`: field order is irrelevant and extra fields are + * ignored, but every required field must be present with a conforming type or compilation fails. + */ +object NamedRecord: + + /** Marker produced by [[Lookup]] when a field name is absent. */ + sealed trait Missing + + /** The value type of field `K` in the named tuple with names `N` and values `V`, or [[Missing]]. */ + type Lookup[N <: Tuple, V <: Tuple, K] = (N, V) match + case (K *: _, v *: _) => v + case (_ *: ns, _ *: vs) => Lookup[ns, vs, K] + case (EmptyTuple, EmptyTuple) => Missing + + /** Fails compilation unless the named tuple `(N, V)` has every field of `Required` with a conforming type. + * + * An unconstrained `N` (inferred as `Tuple` when passing `Nil`) is accepted, as there are no values to check. + */ + inline def check[N <: Tuple, V <: Tuple, Required <: AnyNamedTuple]: Unit = + summonFrom { + case _: (N =:= Tuple) => () + case _ => checkFields[N, V, Names[Required], DropNames[Required]] + } + + private inline def checkFields[N <: Tuple, V <: Tuple, RN <: Tuple, RV <: Tuple]: Unit = + inline erasedValue[RN] match + case _: EmptyTuple => () + case _: (k *: rns) => + inline erasedValue[RV] match + case _: (t *: rvs) => + inline erasedValue[Lookup[N, V, k]] match + case _: Missing => error("Missing field: " + constValue[k & String]) + case _: t => checkFields[N, V, rns, rvs] + case _ => error("Field has the wrong type: " + constValue[k & String]) + + /** The field names of `N` as runtime strings; empty for an unconstrained `N` (see [[check]]). */ + inline def names[N <: Tuple]: List[String] = + summonFrom { + case _: (N =:= Tuple) => Nil + case _ => constValueTuple[N].toList.asInstanceOf[List[String]] + } + + /** Zips pre-computed field `names` with the values of `nt` into a name -> value map. */ + def toMap[N <: Tuple, V <: Tuple](names: List[String], nt: NamedTuple[N, V]): Map[String, Any] = + names.iterator.zip(nt.toTuple.productIterator).toMap + +end NamedRecord diff --git a/vecxt_re/src/PortfolioCalc.scala b/vecxt_re/src/PortfolioCalc.scala new file mode 100644 index 00000000..90053ecf --- /dev/null +++ b/vecxt_re/src/PortfolioCalc.scala @@ -0,0 +1,221 @@ +package vecxt_re + +import scala.NamedTuple.NamedTuple +import java.time.LocalDate +import vecxt.all.* + +/** Portfolio level calculations over a collection of positions. + * + * Positions are record-like views over a `Map[String, Any]`, exposing typed fields via [[scala.Selectable]]. Callers + * need not construct them directly: the public entry points accept any named tuple carrying the required fields (see + * [[NamedRecord]]). + */ +object PortfolioCalc: + + /** A holding for forward-looking (e.g. P&L) calculations. + * + * @note + * Fields are currently identical to [[PositionT0]]; kept separate as the two are expected to diverge. + */ + class PositionT1(values: Map[String, Any]) extends Selectable: + /** @param price + * the market price, as a [[Price]] relative to par + * @param notional + * the face amount, carrying its own currency + * @param name + * a human readable label, used in error messages + * @param id + * a unique identifier, used in error messages + */ + type Fields = + ( + price: Price, + riskFree: RiskFreeRate, + spread: Spread, + maturity: LocalDate, + notional: CurrencyAmount, + name: String, + id: String + ) + def selectDynamic(f: String): Any = values(f) + end PositionT1 + + /** A holding as observed at the valuation date (day 0). + * + * Normally built from a named tuple by [[day0MarketValues]]; constructing one directly from a `Map` is unchecked, + * and a missing key or a wrongly typed value only surfaces when the field is accessed. + */ + class PositionT0(values: Map[String, Any]) extends Selectable: + /** @param price + * the market price, as a [[Price]] relative to par + * @param notional + * the face amount, carrying its own currency + * @param name + * a human readable label, used in error messages + * @param id + * a unique identifier, used in error messages + */ + type Fields = + (price: Price, notional: CurrencyAmount, name: String, id: String) + def selectDynamic(f: String): Any = values(f) + end PositionT0 + + /** The day 0 market value of each position, in the portfolio currency. + * + * For each position the value is `xcRate(notional.ccy) * notional * price`, where the price is taken in canonical + * units (so `Price(102, Pts)` contributes a factor of 1.02). + * + * Positions may be supplied as any named tuple holding the fields of [[PositionT0.Fields]], in any order; extra + * fields are ignored. A missing or mistyped field is a compile error. All elements of the list must share one field + * order, as the list has a single element type. + * + * {{{ + * day0MarketValues( + * List((id = "1", name = "Bond A", price = Price(102, PriceUnit.Pts), notional = CurrencyAmount(1, Mn, USD))), + * portfolioCurrency = Ccy.CHF, + * xcRates = Map(Ccy.USD -> 0.9) + * ) // Array(918000.0) + * }}} + * + * @param positions + * the holdings to value + * @param portfolioCurrency + * the currency results are expressed in. Positions already in this currency use a rate of 1.0 and need no entry in + * `xcRates`. + * @param xcRates + * units of `portfolioCurrency` per one unit of the keyed currency. Entries for currencies no position uses are + * ignored. + * @return + * one value per position, in input order + * @throws IllegalArgumentException + * if any position's currency has no FX rate. All such positions are reported together, one per line, each with its + * id and name. + */ + inline def day0MarketValues[N <: Tuple, V <: Tuple]( + positions: List[NamedTuple[N, V]], + portfolioCurrency: Ccy, + xcRates: Map[Ccy, Double] + ): Array[Double] = + NamedRecord.check[N, V, PositionT0#Fields] + val names = NamedRecord.names[N] + day0MarketValuesOf(positions.map(p => PositionT0(NamedRecord.toMap(names, p))), portfolioCurrency, xcRates) + end day0MarketValues + + /** As [[day0MarketValues]], for positions already built as [[PositionT0]]. + * + * @throws IllegalArgumentException + * if any position's currency has no FX rate; see [[day0MarketValues]] + */ + def day0MarketValuesOf( + positions: List[PositionT0], + portfolioCurrency: Ccy, + xcRates: Map[Ccy, Double] + ): Array[Double] = + val results: List[Either[String, Double]] = positions.map { p => + val ccy = p.notional.ccy + val xcRate = if ccy == portfolioCurrency then Some(1.0) else xcRates.get(ccy) + xcRate + .map(_ * (p.notional * p.price).amount) + .toRight( + s"Position id=${p.id} name=${p.name}: no FX rate from $ccy to $portfolioCurrency" + ) + } + val errors = results.collect { case Left(e) => e } + if errors.nonEmpty then + throw new IllegalArgumentException( + s"Failed to compute day 0 market values for ${errors.size} position(s):\n${errors.mkString("\n")}" + ) + end if + results.collect { case Right(v) => v }.toArray + end day0MarketValuesOf + + def marketValueForecast1Year( + holdings: IndexedSeq[PositionT1], + lossMatrix: Matrix[Double], + t0Date: LocalDate, + xcRates: Map[Ccy, Double], + portfolioCurrency: Ccy + ): Matrix[Double] = + require( + holdings.length == lossMatrix.cols, + s"Each entry in the lossMatrix should have a holdings entry. Loss matrix width ${lossMatrix.cols}, got ${holdings.length} holdings" + ) + val oneYearYield = for h <- holdings yield + val priceT1 = PositionCalculations.priceForecast1Year(h.price, t0Date, h.maturity) + val ccy = h.notional.ccy + val xcRate = + if ccy == portfolioCurrency then Right[String, Double](1.0) + else + xcRates + .get(ccy) + .toRight( + s"Position id=${h.id} name=${h.name}: no FX rate from $ccy to $portfolioCurrency" + ) + + xcRate.map { xc => + val notionalInPortfolioCurrency = xc * h.notional.amount + ((priceT1 + h.riskFree + h.spread).canonical, notionalInPortfolioCurrency) + } + + val errors = oneYearYield.collect { case Left(e) => e } + if errors.nonEmpty then + throw new IllegalArgumentException( + s"Failed to compute 1 year market value forecast for ${errors.size} position(s):\n${errors.mkString("\n")}" + ) + end if + + val (unitValueT1, notionals) = oneYearYield.collect { case Right(r) => r }.unzip + val n = notionals.toArray + + // MV(i, j) = n(j) * (unitValueT1(j) - loss(i, j)) = loss(i, j) * -n(j) + unitValueT1(j) * n(j), in one pass + lossMatrix.fmaCols(multiply = n * -1.0, add = unitValueT1.toArray * n) + end marketValueForecast1Year + + /** The one year relative return of each position in each loss scenario, in canonical units (`0.05` is 5%). + * + * The return is measured against the day 0 market value `notional * price`: + * + * `r(i, j) = (priceT1(j) + riskFree(j) + spread(j) - loss(i, j)) / price(j) - 1` + * + * i.e. [[marketValueForecast1Year]] divided by the day 0 market value, minus one. Notional and FX cancel (both are + * taken at the day 0 rate), so no FX rates are needed. + * + * @param lossMatrix + * one row per scenario, one column per holding; positive losses in canonical units of notional + * @return + * a `lossMatrix.rows x holdings.length` matrix of canonical relative returns + * @throws IllegalArgumentException + * if the column count does not match the holdings, or any holding has a non-positive price (all such holdings are + * reported together, one per line) + */ + def relativeReturn1Year( + holdings: IndexedSeq[PositionT1], + lossMatrix: Matrix[Double], + t0Date: LocalDate + ): Matrix[Double] = + require( + holdings.length == lossMatrix.cols, + s"Each entry in the lossMatrix should have a holdings entry. Loss matrix width ${lossMatrix.cols}, got ${holdings.length} holdings" + ) + val perHolding = holdings.map { h => + val price0 = h.price.canonical + if price0 > 0.0 then + val unitValueT1 = + (PositionCalculations.priceForecast1Year(h.price, t0Date, h.maturity) + h.riskFree + h.spread).canonical + Right((-1.0 / price0, unitValueT1 / price0 - 1.0)) + else Left(s"Position id=${h.id} name=${h.name}: price ${h.price} must be positive to compute a relative return") + end if + } + val errors = perHolding.collect { case Left(e) => e } + if errors.nonEmpty then + throw new IllegalArgumentException( + s"Failed to compute 1 year relative return for ${errors.size} position(s):\n${errors.mkString("\n")}" + ) + end if + + val (multiply, add) = perHolding.collect { case Right(r) => r }.unzip + // r(i, j) = (c(j) - loss(i, j)) / p0(j) - 1 = loss(i, j) * (-1 / p0(j)) + (c(j) / p0(j) - 1), in one pass + lossMatrix.fmaCols(multiply.toArray, add.toArray) + end relativeReturn1Year + +end PortfolioCalc diff --git a/vecxt_re/src/PositionCalculations.scala b/vecxt_re/src/PositionCalculations.scala index 0e52d568..b9f4d095 100644 --- a/vecxt_re/src/PositionCalculations.scala +++ b/vecxt_re/src/PositionCalculations.scala @@ -10,9 +10,7 @@ object PositionCalculations: * This allows us to ignore seasonality, which is assumed to be annual. * * @param price - * the current price, expressed in `priceUnit` - * @param priceUnit - * the unit `price` is quoted in (e.g. `Pts` for a price of 102 meaning 102% of par) + * the current price * @param priceDate * the date on which `price` was observed; the forecast is for one year after this date * @param maturity @@ -21,12 +19,11 @@ object PositionCalculations: * the forecast price one year after `priceDate`, in `priceUnit`. If the bond matures within the year, this is par. */ def priceForecast1Year( - price: Double, - priceUnit: PriceUnit, + price: Price, priceDate: LocalDate, maturity: LocalDate - ): Double = - price + pull2Parity1Year(price, priceUnit, priceDate, maturity) + ): Price = + price + pull2Parity1Year(price, priceDate, maturity) /** The change in price over one year of a bond which is being pulled linearly (by calendar days) to parity at * maturity. @@ -48,24 +45,23 @@ object PositionCalculations: * negative above par. If the bond matures within the year, this is the full distance to par. */ def pull2Parity1Year( - price: Double, - priceUnit: PriceUnit, + price: Price, priceDate: LocalDate, maturity: LocalDate - ): Double = + ): Price = import scala.math.Ordered.orderingToOrdered assert( maturity >= priceDate, s"Maturity $maturity and pricing date $priceDate suggest this bond has already matured. It should not pull to parity" ) - val one = PriceUnit.One.convert(1.0, priceUnit) + val one = Rel(1.0, PriceUnit.One) val projectionDate = priceDate.plusYears(1L) val days2Maturity = ChronoUnit.DAYS.between(priceDate, maturity) val days2Project = ChronoUnit.DAYS.between(priceDate, projectionDate) - if maturity < projectionDate then one - price - else days2Project.toDouble / days2Maturity.toDouble * (one - price) + if maturity < projectionDate then -(price - Rel.one) + else days2Project.toDouble / days2Maturity.toDouble * -(price - Rel.one) end if end pull2Parity1Year @@ -88,19 +84,13 @@ object PositionCalculations: * plus pull to parity minus that loss */ def pnlForecast1Year( - price: Double, - priceUnit: PriceUnit, + price: Price, priceDate: LocalDate, maturity: LocalDate, lossVector: Array[Double], - riskFreeRate: Double, - riskFreeRateUnit: PriceUnit, - spread: Double, - spreadUnit: PriceUnit + riskFreeRate: RiskFreeRate, + spread: Spread ): Array[Double] = - val riskFreeInterest = riskFreeRateUnit.convert(1.0, PriceUnit.One) * riskFreeRate - val spreadInOne = spreadUnit.convert(spread, PriceUnit.One) - val priceInOne = priceUnit.convert(price, PriceUnit.One) - (spreadInOne + riskFreeInterest + pull2Parity1Year(priceInOne, PriceUnit.One, priceDate, maturity)) - lossVector + (spread + riskFreeRate + pull2Parity1Year(price, priceDate, maturity)).canonical - lossVector end pnlForecast1Year end PositionCalculations diff --git a/vecxt_re/src/Rel.scala b/vecxt_re/src/Rel.scala new file mode 100644 index 00000000..75641766 --- /dev/null +++ b/vecxt_re/src/Rel.scala @@ -0,0 +1,61 @@ +package vecxt_re + +// TODO: RelVec(values: Array[Double], unit) for bulk ops, avoid Array[Rel] + +/** Relative Units that have a unit attached. + * + * @param value + * @param unit + */ +final case class Rel(relativeToPriceUnit: Double, unit: PriceUnit): + def canonical: Double = unit.convert(relativeToPriceUnit, PriceUnit.One) + def to(u: PriceUnit): Rel = Rel(unit.convert(relativeToPriceUnit, u), u) + inline def value = canonical + inline def in(u: PriceUnit) = to(u) + + def approxEq(y: Rel, tol: Double = 1e-12): Boolean = + math.abs(canonical - y.canonical) <= tol +end Rel + +extension (k: Double) def *[R <: Rel](r: R): R = r * k +end extension + +object Rel: + // Arithmetic preserves the static type of the left operand, so Price + Rel is a Price. + // Rel is final, so every R <: Rel is a Rel (or an opaque alias of it) at runtime and the cast is safe. + extension [R <: Rel](x: R) + def +(y: Rel): R = + Rel(x.relativeToPriceUnit + y.unit.convert(y.relativeToPriceUnit, x.unit), x.unit).asInstanceOf[R] + def -(y: Rel): R = + Rel(x.relativeToPriceUnit - y.unit.convert(y.relativeToPriceUnit, x.unit), x.unit).asInstanceOf[R] + def unary_- : R = x.copy(relativeToPriceUnit = -x.relativeToPriceUnit).asInstanceOf[R] + def *(k: Double): R = x.copy(relativeToPriceUnit = x.relativeToPriceUnit * k).asInstanceOf[R] + end extension + + given Ordering[Rel] = Ordering.by[Rel, Double](_.canonical) + lazy val one = Rel(1.0, PriceUnit.One) +end Rel + +opaque type Spread <: Rel = Rel +object Spread: + def apply(v: Double, u: PriceUnit): Spread = Rel(v, u) + def from(r: Rel): Spread = r + given CanEqual[Spread, Spread] = CanEqual.derived + given Ordering[Spread] = summon[Ordering[Rel]] +end Spread + +opaque type RiskFreeRate <: Rel = Rel +object RiskFreeRate: + def apply(v: Double, u: PriceUnit): RiskFreeRate = Rel(v, u) + def from(r: Rel): RiskFreeRate = r + given CanEqual[RiskFreeRate, RiskFreeRate] = CanEqual.derived + given Ordering[RiskFreeRate] = summon[Ordering[Rel]] +end RiskFreeRate + +opaque type Price <: Rel = Rel +object Price: + def apply(v: Double, u: PriceUnit): Price = Rel(v, u) + def from(r: Rel): Price = r + given CanEqual[Price, Price] = CanEqual.derived + given Ordering[Price] = summon[Ordering[Rel]] +end Price diff --git a/vecxt_re/test/src/currencyAmount.test.scala b/vecxt_re/test/src/currencyAmount.test.scala new file mode 100644 index 00000000..d664b71f --- /dev/null +++ b/vecxt_re/test/src/currencyAmount.test.scala @@ -0,0 +1,110 @@ +package vecxt_re + +import AmountUnit.* +import Ccy.* + +class CurrencyAmountSuite extends munit.FunSuite: + + test("construction scales by unit"): + assertEquals(CurrencyAmount(2.5, One, USD).amount, 2.5) + assertEquals(CurrencyAmount(2.5, K, USD).amount, 2_500.0) + assertEquals(CurrencyAmount(2.5, Mn, USD).amount, 2_500_000.0) + assertEquals(CurrencyAmount(2.5, Bn, USD).amount, 2_500_000_000.0) + + test("the same money in different units is equal"): + assertEquals(CurrencyAmount(1, Mn, EUR), CurrencyAmount(1_000, K, EUR)) + assertEquals(CurrencyAmount(1, Bn, EUR), CurrencyAmount(1_000_000_000, One, EUR)) + assertEquals(CurrencyAmount(1, Mn, EUR).hashCode, CurrencyAmount(1_000, K, EUR).hashCode) + + test("the same number in different currencies is not equal"): + assertNotEquals(CurrencyAmount(1, Mn, USD), CurrencyAmount(1, Mn, EUR)) + + test("in converts to the requested unit"): + val x = CurrencyAmount(2.5, Mn, USD) + assertEquals(x.in(One), 2_500_000.0) + assertEquals(x.in(K), 2_500.0) + assertEquals(x.in(Mn), 2.5) + assertEqualsDouble(x.in(Bn), 0.0025, 1e-15) + + test("addition and subtraction across units"): + val x = CurrencyAmount(1, Mn, GBP) + val y = CurrencyAmount(250, K, GBP) + assertEquals(x + y, CurrencyAmount(1.25, Mn, GBP)) + assertEquals(x - y, CurrencyAmount(750, K, GBP)) + assertEquals(y - x, CurrencyAmount(-750, K, GBP)) + + test("negation"): + assertEquals(-CurrencyAmount(3, K, CHF), CurrencyAmount(-3, K, CHF)) + assertEquals(-(-CurrencyAmount(3, K, CHF)), CurrencyAmount(3, K, CHF)) + + test("zero is the additive identity"): + val x = CurrencyAmount(42, Mn, JPY) + assertEquals(x + CurrencyAmount.zero(JPY), x) + assertEquals(x - x, CurrencyAmount.zero(JPY)) + + test("arithmetic between currencies fails"): + val usd = CurrencyAmount(1, Mn, USD) + val eur = CurrencyAmount(1, Mn, EUR) + intercept[IllegalArgumentException](usd + eur) + intercept[IllegalArgumentException](usd - eur) + intercept[IllegalArgumentException](usd / eur) + + test("premium = notional * coupon, in any price unit"): + val notional = CurrencyAmount(100, Mn, USD) + val expected = CurrencyAmount(500, K, USD) + assertEquals(notional * Rel(50, PriceUnit.Bps), expected) + assertEquals(notional * Rel(0.5, PriceUnit.Pts), expected) + assertEquals(notional * Rel(0.005, PriceUnit.One), expected) + assertEquals(notional * Spread(50, PriceUnit.Bps), expected) + + test("scaling by a Double keeps the currency"): + val x = CurrencyAmount(10, Mn, AUD) * 0.3 + assertEquals(x.ccy, AUD) + assert(x.approxEq(CurrencyAmount(3, Mn, AUD))) + + test("dividing two amounts gives a Rel"): + val loss = CurrencyAmount(60, Mn, USD) + val premium = CurrencyAmount(80, Mn, USD) + assertEquals(loss / premium, Rel(75, PriceUnit.Pts)) + assertEquals(premium / premium, Rel.one) + + test("dividing then multiplying round trips"): + val a = CurrencyAmount(123.456, K, NZD) + val b = CurrencyAmount(7.89, Mn, NZD) + assert((b * (a / b)).approxEq(a)) + + test("approxEq tolerates floating point noise"): + val sum = CurrencyAmount(0.1, One, USD) + CurrencyAmount(0.2, One, USD) + val exact = CurrencyAmount(0.3, One, USD) + assertNotEquals(sum, exact) + assert(sum.approxEq(exact)) + + test("approxEq respects the tolerance"): + val x = CurrencyAmount(1, Mn, USD) + val y = CurrencyAmount(1_000_000.01, One, USD) + assert(!x.approxEq(y)) + assert(x.approxEq(y, tol = 0.1)) + + test("approxEq across currencies fails"): + intercept[IllegalArgumentException](CurrencyAmount(1, Mn, USD).approxEq(CurrencyAmount(1, Mn, EUR))) + + test("ordering is by amount regardless of the unit used to construct"): + val amounts = List(CurrencyAmount(2, Mn, CAD), CurrencyAmount(500, K, CAD), CurrencyAmount(0.001, Bn, CAD)) + assertEquals( + amounts.sorted, + List(CurrencyAmount(500, K, CAD), CurrencyAmount(1, Mn, CAD), CurrencyAmount(2, Mn, CAD)) + ) + assertEquals(amounts.max, CurrencyAmount(2, Mn, CAD)) + + test("ordering across currencies fails"): + intercept[IllegalArgumentException](List(CurrencyAmount(1, Mn, USD), CurrencyAmount(1, Mn, EUR)).sorted) + + test("very large amounts do not overflow"): + val x = CurrencyAmount(1_000_000, Bn, USD) + assertEquals((x + x).in(Bn), 2_000_000.0) + + test("toString shows the amount in major units with the currency"): + assertEquals(CurrencyAmount(2.5, Mn, USD).toString, "2500000.00 USD") + assertEquals(CurrencyAmount(-1.234, K, EUR).toString, "-1234.00 EUR") + +end CurrencyAmountSuite diff --git a/vecxt_re/test/src/portfolioCalc.test.scala b/vecxt_re/test/src/portfolioCalc.test.scala new file mode 100644 index 00000000..c6fc375b --- /dev/null +++ b/vecxt_re/test/src/portfolioCalc.test.scala @@ -0,0 +1,171 @@ +package vecxt_re + +import java.time.LocalDate + +import vecxt.all.* + +import PortfolioCalc.* + +class PortfolioCalcSuite extends munit.FunSuite: + + private def amt(v: Double, ccy: Ccy): CurrencyAmount = CurrencyAmount(v, AmountUnit.One, ccy) + + private val tol = 1e-9 + + test("empty portfolio gives empty array"): + assertEquals(day0MarketValues(Nil, Ccy.CHF, Map.empty).length, 0) + + test("positions in portfolio currency need no FX rate"): + val result = day0MarketValues( + List( + (price = Price(102, PriceUnit.Pts), notional = amt(1_000_000, Ccy.CHF), name = "A", id = "1", unused = "???") + ), + Ccy.CHF, + Map.empty + ) + assertEquals(result.length, 1) + assertEqualsDouble(result(0), 1_020_000.0, tol) + + test("applies FX rate, notional and canonical price, preserving order"): + val positions = List( + (id = "1", name = "A", price = Price(102, PriceUnit.Pts), notional = amt(1_000_000, Ccy.USD)), + (id = "2", name = "B", price = Price(9500, PriceUnit.Bps), notional = amt(2_000_000, Ccy.EUR)), + (id = "3", name = "C", price = Price(1.0, PriceUnit.One), notional = amt(500_000, Ccy.CHF)) + ) + val result = day0MarketValues(positions, Ccy.CHF, Map(Ccy.USD -> 0.9, Ccy.EUR -> 0.95)) + assertEquals(result.length, 3) + assertEqualsDouble(result(0), 0.9 * 1_000_000 * 1.02, tol) + assertEqualsDouble(result(1), 0.95 * 2_000_000 * 0.95, tol) + assertEqualsDouble(result(2), 500_000.0, tol) + + test("field order does not matter and extra fields are ignored"): + val result = day0MarketValues( + List((desk = "rates", notional = amt(100, Ccy.USD), id = "1", price = Price(100, PriceUnit.Pts), name = "A")), + Ccy.CHF, + Map(Ccy.USD -> 2.0) + ) + assertEqualsDouble(result(0), 200.0, tol) + + test("unused FX rates are ignored"): + val result = day0MarketValues( + List((price = Price(100, PriceUnit.Pts), notional = amt(100, Ccy.USD), name = "A", id = "1")), + Ccy.CHF, + Map(Ccy.USD -> 2.0, Ccy.JPY -> 0.006, Ccy.GBP -> 1.1) + ) + assertEqualsDouble(result(0), 200.0, tol) + + test("missing or mistyped fields do not compile"): + assertNoDiff( + compileErrors( + """day0MarketValues(List((price = Price(100, PriceUnit.Pts), name = "A", id = "1")), Ccy.CHF, Map.empty)""" + ).linesIterator.find(_.startsWith("error:")).getOrElse(""), + "error: Missing field: notional" + ) + assertNoDiff( + compileErrors( + """day0MarketValues(List((price = Price(100, PriceUnit.Pts), notional = 100.0, name = "A", id = "1")), Ccy.CHF, Map.empty)""" + ).linesIterator.find(_.startsWith("error:")).getOrElse(""), + "error: Field has the wrong type: notional" + ) + + test("missing FX rate throws with position id and name"): + val ex = intercept[IllegalArgumentException]( + day0MarketValues( + List((id = "P1", name = "Foo Bond", price = Price(100, PriceUnit.Pts), notional = amt(100, Ccy.USD))), + Ccy.CHF, + Map.empty + ) + ) + val lines = ex.getMessage.split("\n").toList + assertEquals(lines.length, 2) + assert(lines.head.contains("1 position(s)"), lines.head) + assertEquals(lines(1), "Position id=P1 name=Foo Bond: no FX rate from USD to CHF") + + test("all failures are reported in one exception, one per line, in input order"): + val positions = List( + (id = "P1", name = "Foo", price = Price(100, PriceUnit.Pts), notional = amt(100, Ccy.USD)), + (id = "P2", name = "Ok", price = Price(100, PriceUnit.Pts), notional = amt(100, Ccy.EUR)), + (id = "P3", name = "Bar", price = Price(100, PriceUnit.Pts), notional = amt(100, Ccy.JPY)), + (id = "P4", name = "Home", price = Price(100, PriceUnit.Pts), notional = amt(100, Ccy.CHF)), + (id = "P5", name = "Baz", price = Price(100, PriceUnit.Pts), notional = amt(100, Ccy.USD)) + ) + val ex = intercept[IllegalArgumentException]( + day0MarketValues(positions, Ccy.CHF, Map(Ccy.EUR -> 0.95)) + ) + val lines = ex.getMessage.split("\n").toList + assert(lines.head.contains("3 position(s)"), lines.head) + assertEquals( + lines.tail, + List( + "Position id=P1 name=Foo: no FX rate from USD to CHF", + "Position id=P3 name=Bar: no FX rate from JPY to CHF", + "Position id=P5 name=Baz: no FX rate from USD to CHF" + ) + ) + + private val t0 = LocalDate.of(2026, 1, 1) + + private def pos(price: Price, rf: Double, spread: Double, maturity: LocalDate, notional: CurrencyAmount, id: String) = + PositionT1( + Map( + "price" -> price, + "riskFree" -> RiskFreeRate(rf, PriceUnit.Bps), + "spread" -> Spread(spread, PriceUnit.Bps), + "maturity" -> maturity, + "notional" -> notional, + "name" -> s"Bond $id", + "id" -> id + ) + ) + + private val t1Holdings = IndexedSeq( + // at par, long dated: no pull to par, carry 5%, 900k CHF + pos(Price(100, PriceUnit.Pts), 300, 200, LocalDate.of(2036, 1, 1), amt(1_000_000, Ccy.USD), "A"), + // matures within the year: pulled fully to par, no carry, already CHF + pos(Price(98, PriceUnit.Pts), 0, 0, LocalDate.of(2026, 7, 1), amt(500_000, Ccy.CHF), "B") + ) + + // two scenarios (rows) x two positions (columns); losses are positive fractions of notional + private val losses = Matrix.fromRows[Double](Array(0.0, 0.2), Array(0.1, 0.0)) + + test( + "marketValueForecast1Year: notional in portfolio ccy times (price T1 + carry - loss), per scenario and position" + ): + val mv = marketValueForecast1Year(t1Holdings, losses, t0, Map(Ccy.USD -> 0.9), Ccy.CHF) + assertEquals((mv.rows, mv.cols), (2, 2)) + assertEqualsDouble(mv(0, 0), 900_000 * 1.05, 1e-6) + assertEqualsDouble(mv(1, 0), 900_000 * (1.05 - 0.1), 1e-6) + assertEqualsDouble(mv(0, 1), 500_000 * (1.0 - 0.2), 1e-6) + assertEqualsDouble(mv(1, 1), 500_000 * 1.0, 1e-6) + + test("relativeReturn1Year: canonical return on day 0 market value, per scenario and position"): + val r = relativeReturn1Year(t1Holdings, losses, t0) + assertEquals((r.rows, r.cols), (2, 2)) + assertEqualsDouble(r(0, 0), 0.05, 1e-12) + assertEqualsDouble(r(1, 0), -0.05, 1e-12) + assertEqualsDouble(r(0, 1), 0.8 / 0.98 - 1.0, 1e-12) + assertEqualsDouble(r(1, 1), 1.0 / 0.98 - 1.0, 1e-12) + + test("relativeReturn1Year agrees with marketValueForecast1Year / day 0 market value - 1"): + val fx = Map(Ccy.USD -> 0.9) + val mv1 = marketValueForecast1Year(t1Holdings, losses, t0, fx, Ccy.CHF) + val mv0 = Array(900_000 * 1.0, 500_000 * 0.98) + val r = relativeReturn1Year(t1Holdings, losses, t0) + for i <- 0 until 2; j <- 0 until 2 do assertEqualsDouble(r(i, j), mv1(i, j) / mv0(j) - 1.0, 1e-12) + end for + + test("relativeReturn1Year rejects non-positive prices, reporting all of them"): + val bad = IndexedSeq( + pos(Price(0, PriceUnit.Pts), 0, 0, LocalDate.of(2030, 1, 1), amt(1, Ccy.CHF), "Z"), + t1Holdings(0), + pos(Price(-5, PriceUnit.Pts), 0, 0, LocalDate.of(2030, 1, 1), amt(1, Ccy.CHF), "N") + ) + val ex = intercept[IllegalArgumentException](relativeReturn1Year(bad, Matrix[Double](Array.fill(3)(0.0), 1, 3), t0)) + val lines = ex.getMessage.split("\n").toList + assert(lines.head.contains("2 position(s)"), lines.head) + assertEquals(lines.tail.map(_.takeWhile(_ != ':')), List("Position id=Z name=Bond Z", "Position id=N name=Bond N")) + + test("relativeReturn1Year rejects a loss matrix whose width does not match the holdings"): + intercept[IllegalArgumentException](relativeReturn1Year(t1Holdings, Matrix[Double](Array.fill(3)(0.0), 1, 3), t0)) + +end PortfolioCalcSuite diff --git a/vecxt_re/test/src/priceForecast.test.scala b/vecxt_re/test/src/priceForecast.test.scala index 74c76de9..b2aeb129 100644 --- a/vecxt_re/test/src/priceForecast.test.scala +++ b/vecxt_re/test/src/priceForecast.test.scala @@ -7,67 +7,63 @@ class PositionCalculationsSuite extends munit.FunSuite: test("one year maturity"): val reportDate = LocalDate.of(2026, 1, 1) - val forecastPrice = PositionCalculations.priceForecast1Year(102, PriceUnit.Pts, reportDate, reportDate.plusYears(1)) + val forecastPrice = + PositionCalculations.priceForecast1Year(Price(102, PriceUnit.Pts), reportDate, reportDate.plusYears(1)) - assertEqualsDouble(forecastPrice, 100.0, 0.0000001) + assertEquals(forecastPrice, Price(100.0, PriceUnit.Pts)) test("Maturity same date"): val reportDate = LocalDate.of(2026, 1, 1) - val forecastPrice = PositionCalculations.priceForecast1Year(102, PriceUnit.Pts, reportDate, reportDate) + val forecastPrice = PositionCalculations.priceForecast1Year(Price(102, PriceUnit.Pts), reportDate, reportDate) - assertEqualsDouble(forecastPrice, 100.0, 0.0000001) + assertEquals(forecastPrice, Price(100.0, PriceUnit.Pts)) test("Early maturity"): val reportDate = LocalDate.of(2026, 1, 1) val forecastPrice = - PositionCalculations.priceForecast1Year(102, PriceUnit.Pts, reportDate, reportDate.plusMonths(6)) - val forecastPrice2 = PositionCalculations.priceForecast1Year(102, PriceUnit.Pts, reportDate, reportDate.plusDays(1)) + PositionCalculations.priceForecast1Year(Price(102, PriceUnit.Pts), reportDate, reportDate.plusMonths(6)) + val forecastPrice2 = + PositionCalculations.priceForecast1Year(Price(102, PriceUnit.Pts), reportDate, reportDate.plusDays(1)) - assertEqualsDouble(forecastPrice, 100.0, 0.0000001) - assertEqualsDouble(forecastPrice2, 100.0, 0.0000001) + assertEquals(forecastPrice, Price(100.0, PriceUnit.Pts)) + assertEquals(forecastPrice2, Price(100.0, PriceUnit.Pts)) test("refuses already matured"): val reportDate = LocalDate.of(2026, 1, 1) intercept[java.lang.AssertionError]( - PositionCalculations.priceForecast1Year(102, PriceUnit.Pts, reportDate, reportDate.plusYears(-1)) + PositionCalculations.priceForecast1Year(Price(102, PriceUnit.Pts), reportDate, reportDate.plusYears(-1)) ) test("Interpolation to a point"): val reportDate = LocalDate.of(2026, 1, 1) val forecastPrice = - PositionCalculations.priceForecast1Year(10200, PriceUnit.Bps, reportDate, reportDate.plusYears(2)) + PositionCalculations.priceForecast1Year(Price(10200, PriceUnit.Bps), reportDate, reportDate.plusYears(2)) - assertEqualsDouble(forecastPrice, 10100.0, 0.0000001) + assertEquals(forecastPrice, Price(10100.0, PriceUnit.Bps)) + assertEquals(forecastPrice.unit, PriceUnit.Bps) test("pnlForecast"): - val reportDate = LocalDate.of(2026, 1, 1) val losses = Array[Double](0.0, 0.1, 0.0, 0.2) val forecastPrice = PositionCalculations.pnlForecast1Year( - 102, - PriceUnit.Pts, + Price(102, PriceUnit.Pts), reportDate, reportDate.plusYears(2), losses, - 325.0, - PriceUnit.Bps, - 0, - PriceUnit.Bps + RiskFreeRate(325.0, PriceUnit.Bps), + Spread(0.0, PriceUnit.Bps) ) val forecast = Array[Double](0.0225, -0.0775, 0.0225, -0.1775) assertVecEquals(forecastPrice, forecast) val forecastPriceWSpread = PositionCalculations.pnlForecast1Year( - 102, - PriceUnit.Pts, + Price(102, PriceUnit.Pts), reportDate, reportDate.plusYears(2), losses, - 325.0, - PriceUnit.Bps, - 100.0, - PriceUnit.Bps + RiskFreeRate(325.0, PriceUnit.Bps), + Spread(100.0, PriceUnit.Bps) ) assertVecEquals(forecastPriceWSpread, forecast + 0.01) end PositionCalculationsSuite diff --git a/vecxt_re/test/src/vecEquals.scala b/vecxt_re/test/src/vecEquals.scala index 924a940c..f42f098f 100644 --- a/vecxt_re/test/src/vecEquals.scala +++ b/vecxt_re/test/src/vecEquals.scala @@ -29,3 +29,10 @@ def assertVecEquals(v1: Array[Long], v2: Array[Long])(implicit loc: munit.Locati i += 1 end while end assertVecEquals + +/** Makes `assertEquals` on Rel (and Price / Spread / RiskFreeRate) compare canonical values within a tolerance, so + * `Price(100, Pts)` equals `Price(1.0, One)` and floating point noise doesn't fail tests. + */ +given relCompare[R <: Rel]: munit.Compare[R, R] with + def isEqual(obtained: R, expected: R): Boolean = obtained.approxEq(expected, 1e-9) +end relCompare