Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down
2 changes: 1 addition & 1 deletion requirements.lock
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
37 changes: 26 additions & 11 deletions src/quant_execution/_fixed.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
24 changes: 15 additions & 9 deletions src/quant_execution/ledger.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand All @@ -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
Expand All @@ -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:
Expand Down
213 changes: 213 additions & 0 deletions tests/test_precision.py
Original file line number Diff line number Diff line change
@@ -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
Loading