diff --git a/pyproject.toml b/pyproject.toml index f88a999..d9e3691 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -12,7 +12,7 @@ requires-python = ">=3.10" dependencies = [ "pyarrow>=14.0", "jsonschema>=4.20", - "quant-data-kit @ git+https://github.com/PureSaber/quant-data-kit.git@ba136c2fa2eea121bfb2ad7887b536c3952586f7", + "quant-data-kit @ git+https://github.com/PureSaber/quant-data-kit.git@fd788b2956a10490aa00c399b717ec796ae371b1", ] [project.optional-dependencies] diff --git a/requirements.lock b/requirements.lock index 1390d34..62b7f15 100644 --- a/requirements.lock +++ b/requirements.lock @@ -67,7 +67,7 @@ pytz==2026.3.post1 # via pandas pyyaml==6.0.3 # via quant-data-kit -quant-data-kit @ git+https://github.com/PureSaber/quant-data-kit.git@ba136c2fa2eea121bfb2ad7887b536c3952586f7 +quant-data-kit @ git+https://github.com/PureSaber/quant-data-kit.git@fd788b2956a10490aa00c399b717ec796ae371b1 # via quant-execution (pyproject.toml) referencing==0.37.0 # via diff --git a/src/quant_execution/_fixed.py b/src/quant_execution/_fixed.py index 66eeb09..b604694 100644 --- a/src/quant_execution/_fixed.py +++ b/src/quant_execution/_fixed.py @@ -3,34 +3,49 @@ from __future__ import annotations from decimal import ROUND_DOWN, ROUND_HALF_EVEN, Decimal -from functools import lru_cache from quant_data_kit import FixedPoint from quant_data_kit.exceptions import ValidationError -@lru_cache(maxsize=8192) def decimal(value: FixedPoint) -> Decimal: if not isinstance(value, FixedPoint): raise ValidationError("value must be a FixedPoint") return value.to_decimal() -@lru_cache(maxsize=8192) def fixed( value: Decimal | int | str, scale: int, *, rounding: str | None = ROUND_HALF_EVEN, ) -> FixedPoint: - decimal_value = value if isinstance(value, Decimal) else Decimal(str(value)) - if not decimal_value.is_finite(): - raise ValidationError("fixed-point value must be finite") - scaled = decimal_value.scaleb(scale) - integral = scaled.to_integral_value(rounding=rounding) if rounding else scaled - if rounding is None and scaled != scaled.to_integral_value(): - raise ValidationError(f"value {value!r} is not exact at scale {scale}") - return FixedPoint(units=int(integral), scale=scale) + return FixedPoint.from_decimal(value, scale, rounding=rounding) + + +def add_decimal_exact(left: Decimal, right: Decimal) -> Decimal: + """Add finite decimals exactly without consulting the ambient context.""" + + if not isinstance(left, Decimal) or not isinstance(right, Decimal): + raise ValidationError("exact decimal addition requires Decimal values") + if not left.is_finite() or not right.is_finite(): + raise ValidationError("exact decimal addition requires finite values") + + left_parts = left.as_tuple() + right_parts = right.as_tuple() + exponent = min(int(left_parts.exponent), int(right_parts.exponent)) + + def aligned(parts) -> int: + coefficient = 0 + for digit in parts.digits: + coefficient = coefficient * 10 + digit + coefficient *= 10 ** (int(parts.exponent) - exponent) + return -coefficient if parts.sign and coefficient else coefficient + + total = aligned(left_parts) + aligned(right_parts) + magnitude = abs(total) + digits = tuple(int(digit) for digit in str(magnitude)) if magnitude else (0,) + return Decimal((int(total < 0), digits, exponent)) def floor_to_scale(value: Decimal, scale: int) -> FixedPoint: diff --git a/src/quant_execution/ledger.py b/src/quant_execution/ledger.py index 47bcc38..b6a72ee 100644 --- a/src/quant_execution/ledger.py +++ b/src/quant_execution/ledger.py @@ -29,7 +29,7 @@ ) from quant_data_kit.exceptions import ValidationError -from quant_execution._fixed import decimal, fixed +from quant_execution._fixed import add_decimal_exact, decimal, fixed from quant_execution._json import fixed_token, flat_sequence_bytes, string_token, utc_token from quant_execution.artifacts import ( fee_bytes, @@ -1822,16 +1822,19 @@ def _post(self, transaction: LedgerTransaction, *, local_rollback: bool = True) if not local_rollback: for posting in transaction.postings: key = (posting.ledger_account, posting.currency, posting.instrument_id) - self._accounts[key] = self._accounts.get(key, Decimal(0)) + decimal(posting.amount) + self._accounts[key] = add_decimal_exact( + self._accounts.get(key, Decimal(0)), decimal(posting.amount) + ) if ( posting.ledger_account == "assets:position" and posting.instrument_id is not None and posting.quantity_delta is not None ): instrument_id = posting.instrument_id - self._positions[instrument_id] = self._positions.get( - instrument_id, Decimal(0) - ) + decimal(posting.quantity_delta) + self._positions[instrument_id] = add_decimal_exact( + self._positions.get(instrument_id, Decimal(0)), + decimal(posting.quantity_delta), + ) if self._artifact_sink is None: self._transactions.append(transaction) else: @@ -1849,7 +1852,9 @@ def _post(self, transaction: LedgerTransaction, *, local_rollback: bool = True) for posting in transaction.postings: key = (posting.ledger_account, posting.currency, posting.instrument_id) prior_accounts.setdefault(key, self._accounts.get(key, missing)) - self._accounts[key] = self._accounts.get(key, Decimal(0)) + decimal(posting.amount) + self._accounts[key] = add_decimal_exact( + self._accounts.get(key, Decimal(0)), decimal(posting.amount) + ) if ( posting.ledger_account == "assets:position" and posting.instrument_id is not None @@ -1859,9 +1864,10 @@ def _post(self, transaction: LedgerTransaction, *, local_rollback: bool = True) prior_positions.setdefault( instrument_id, self._positions.get(instrument_id, missing) ) - self._positions[instrument_id] = self._positions.get( - instrument_id, Decimal(0) - ) + decimal(posting.quantity_delta) + self._positions[instrument_id] = add_decimal_exact( + self._positions.get(instrument_id, Decimal(0)), + decimal(posting.quantity_delta), + ) if self._artifact_sink is None: self._transactions.append(transaction) else: diff --git a/tests/test_precision.py b/tests/test_precision.py new file mode 100644 index 0000000..f107b32 --- /dev/null +++ b/tests/test_precision.py @@ -0,0 +1,213 @@ +from __future__ import annotations + +from decimal import Decimal, Inexact, Rounded, localcontext + +import pytest +from conftest import T0, fp, spec +from quant_data_kit import AssetClass, FixedPoint +from quant_data_kit.exceptions import ValidationError + +from quant_execution._fixed import decimal, fixed +from quant_execution.artifacts import ledger_transaction_bytes +from quant_execution.contracts import LedgerEventType, LedgerTransaction, Posting +from quant_execution.ledger import ExactAccountLedger + +STOCK = "HK:TEST" + + +def ledger(*, cash: str = "100") -> ExactAccountLedger: + instrument = spec( + STOCK, + asset_class=AssetClass.EQUITY, + product_type="cash_equity", + settlement_currency="HKD", + ) + return ExactAccountLedger( + account_id="precision-account", + base_currency="HKD", + instruments={STOCK: instrument}, + initial_cash={"HKD": fp(cash)}, + ) + + +def cash_transaction(reference: str, amount: FixedPoint) -> LedgerTransaction: + return LedgerTransaction( + transaction_id=f"tx:{reference}", + idempotency_key=f"precision:{reference}", + event_time=T0, + event_type=LedgerEventType.FEE, + reference_id=reference, + postings=( + Posting(ledger_account="assets:cash", currency="HKD", amount=amount), + Posting( + ledger_account="equity:precision-counter", + currency="HKD", + amount=FixedPoint(-amount.units, amount.scale), + ), + ), + ) + + +def position_transaction(reference: str, quantity: FixedPoint) -> LedgerTransaction: + zero = FixedPoint(0, 0) + return LedgerTransaction( + transaction_id=f"tx:{reference}", + idempotency_key=f"precision:{reference}", + event_time=T0, + event_type=LedgerEventType.CORPORATE_ACTION, + reference_id=reference, + postings=( + Posting( + ledger_account="assets:position", + currency="HKD", + amount=zero, + instrument_id=STOCK, + quantity_delta=quantity, + ), + Posting( + ledger_account="memo:position_counter", + currency="HKD", + amount=zero, + instrument_id=STOCK, + quantity_delta=FixedPoint(-quantity.units, quantity.scale), + ), + ), + ) + + +@pytest.mark.parametrize("local_rollback", [False, True]) +def test_post_adds_small_cash_exactly_under_low_precision(local_rollback: bool) -> None: + account = ledger() + transaction = cash_transaction("cent", FixedPoint(1, 2)) + + with localcontext() as context: + context.prec = 2 + context.traps[Inexact] = True + context.traps[Rounded] = True + account._post(transaction, local_rollback=local_rollback) + + assert account.cash_balance("HKD") == Decimal("100.01") + assert ledger_transaction_bytes(account.transactions[-1]) == ledger_transaction_bytes( + transaction + ) + + +def test_fixed_conversion_is_exact_and_call_order_cannot_pollute_results() -> None: + value = Decimal("100.01") + results = [] + for precision in (2, 28, 80, 2): + with localcontext() as context: + context.prec = precision + context.traps[Inexact] = True + context.traps[Rounded] = True + results.append(fixed(value, 2, rounding=None)) + + assert results == [FixedPoint(10001, 2)] * 4 + assert decimal(results[0]) == value + with localcontext() as context: + context.prec = 2 + context.traps[Inexact] = True + context.traps[Rounded] = True + with pytest.raises(ValidationError, match="not exact"): + fixed(Decimal("100.001"), 2, rounding=None) + assert fixed(Decimal("1.005"), 2).units == 100 + assert fixed(Decimal("1.015"), 2).units == 102 + with pytest.raises(ValidationError, match="FixedPoint"): + decimal(Decimal(1)) # type: ignore[arg-type] + + +def test_post_accumulates_signed_mixed_scales_and_unit_quantities_exactly() -> None: + account = ledger() + operations = ( + cash_transaction("plus-cent", FixedPoint(1, 2)), + cash_transaction("minus-mills", FixedPoint(-2, 3)), + position_transaction("one-unit", FixedPoint(1, 0)), + position_transaction("one-hundredth", FixedPoint(1, 2)), + position_transaction("minus-one-thousandth", FixedPoint(-1, 3)), + ) + + with localcontext() as context: + context.prec = 2 + context.traps[Inexact] = True + context.traps[Rounded] = True + for transaction in operations: + account._post(transaction) + + assert account.cash_balance("HKD") == Decimal("100.008") + assert account._positions[STOCK] == Decimal("1.009") + + +def test_precision_and_traps_do_not_change_post_bytes_balances_or_journal() -> None: + observations = [] + for precision in (2, 6, 28, 80): + account = ledger() + transaction = cash_transaction("stable", FixedPoint(1, 2)) + with localcontext() as context: + context.prec = precision + context.traps[Inexact] = True + context.traps[Rounded] = True + account._post(transaction) + observations.append( + ( + account.cash_balance("HKD"), + ledger_transaction_bytes(account.transactions[-1]), + account.journal_sha256, + ) + ) + + assert observations == [observations[0]] * len(observations) + + +def test_failed_stream_append_restores_successful_prior_post_and_all_new_state() -> None: + class FailingSink: + def append(self, stream: str, payload: bytes) -> None: + assert stream == "ledger_transactions" + assert payload + raise RuntimeError("injected artifact append failure") + + account = ledger() + account._post(cash_transaction("committed", FixedPoint(1, 2))) + before = account.capture_state() + before_hash = account.journal_sha256 + failing = LedgerTransaction( + transaction_id="tx:failing", + idempotency_key="precision:failing", + event_time=T0, + event_type=LedgerEventType.CORPORATE_ACTION, + reference_id="failing", + postings=( + Posting( + ledger_account="assets:cash", + currency="HKD", + amount=FixedPoint(1, 3), + ), + Posting( + ledger_account="equity:precision-counter", + currency="HKD", + amount=FixedPoint(-1, 3), + ), + Posting( + ledger_account="assets:position", + currency="HKD", + amount=FixedPoint(0, 0), + instrument_id=STOCK, + quantity_delta=FixedPoint(1, 3), + ), + ), + ) + account._artifact_sink = FailingSink() + + with ( + localcontext() as context, + pytest.raises(RuntimeError, match="injected artifact append failure"), + ): + context.prec = 2 + context.traps[Inexact] = True + context.traps[Rounded] = True + account._post(failing) + + account._artifact_sink = None + assert account.capture_state() == before + assert account.journal_sha256 == before_hash + assert account.cash_balance("HKD") == Decimal("100.01") + assert STOCK not in account._positions