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
79 changes: 71 additions & 8 deletions src/quant_platform_kit/schwab/portfolio.py
Original file line number Diff line number Diff line change
@@ -1,15 +1,23 @@
from __future__ import annotations

from datetime import datetime, timezone
from decimal import Decimal, InvalidOperation
from hashlib import sha256
import json
import math
import re
from typing import Any, Iterable

from quant_platform_kit.common.models import PortfolioSnapshot, Position
from .market_data import _request_with_retries, decode_response_json


_BROKER_ACCOUNT_TYPE_RE = re.compile(r"[A-Za-z_]{1,32}\Z", re.ASCII)
_MAX_CASH_BALANCE_RAW_LENGTH = 128
_MAX_CASH_BALANCE_INTEGER_DIGITS = 15
_MAX_CASH_BALANCE_FRACTION_DIGITS = 8


def _payload_digest(payload: Any) -> str:
"""Bind a normalized snapshot to the broker response without retaining it."""

Expand Down Expand Up @@ -44,6 +52,51 @@ def _optional_finite_balance(balances: dict[str, Any], key: str) -> float | None
return _finite_balance(balances, key)


def _optional_decimal_text(value: Any) -> str | None:
"""Preserve a bounded broker decimal fact without routing it through float sizing."""

if isinstance(value, bool) or not isinstance(value, (int, float, str, Decimal)):
return None
try:
raw_text = str(value)
except (TypeError, ValueError, OverflowError):
return None
if len(raw_text) > _MAX_CASH_BALANCE_RAW_LENGTH:
return None
try:
amount = Decimal(raw_text)
except (InvalidOperation, TypeError, ValueError):
return None
if not amount.is_finite():
return None
sign, digits, exponent = amount.as_tuple()
integer_digits = max(1, len(digits) + exponent)
fraction_digits = max(0, -exponent)
if (
integer_digits > _MAX_CASH_BALANCE_INTEGER_DIGITS
or fraction_digits > _MAX_CASH_BALANCE_FRACTION_DIGITS
):
return None
expected_output_length = (
sign
+ integer_digits
+ (1 + fraction_digits if fraction_digits else 0)
)
if expected_output_length > (
1 + _MAX_CASH_BALANCE_INTEGER_DIGITS + 1 + _MAX_CASH_BALANCE_FRACTION_DIGITS
):
return None
return format(amount, "f")


def _optional_account_type(value: Any) -> str | None:
"""Keep only a bounded raw provider token; do not classify account type."""

if not isinstance(value, str) or _BROKER_ACCOUNT_TYPE_RE.fullmatch(value) is None:
return None
return value


def _resolve_buying_power(balances: dict[str, Any], *, cash_available_for_trading: float) -> tuple[float, str]:
"""Map broker buying-power fields without inventing leverage.

Expand Down Expand Up @@ -142,18 +195,28 @@ def fetch_account_snapshot(
total_equity = cash_for_equity + all_position_market_value
total_equity_source = "cash_available_plus_all_position_market_values"

metadata: dict[str, Any] = {
"account_hash": account_hash,
"cash_available_for_trading": cash_for_equity,
"cash_available_for_withdrawal": raw_withdrawable,
"buying_power_source": buying_power_source,
"total_equity_source": total_equity_source,
"source_digest_sha256": _payload_digest(account_payload),
}
broker_account_type = _optional_account_type(account.get("type"))
if broker_account_type is not None:
metadata["broker_account_type"] = broker_account_type
metadata["broker_account_type_source"] = "securitiesAccount.type"
broker_cash_balance = _optional_decimal_text(balances.get("cashBalance"))
if broker_cash_balance is not None:
metadata["broker_cash_balance"] = broker_cash_balance
metadata["broker_cash_balance_source"] = "cashBalance"

return PortfolioSnapshot(
as_of=datetime.now(timezone.utc),
total_equity=total_equity,
buying_power=buying_power,
cash_balance=cash_for_equity,
positions=tuple(positions),
metadata={
"account_hash": account_hash,
"cash_available_for_trading": cash_for_equity,
"cash_available_for_withdrawal": raw_withdrawable,
"buying_power_source": buying_power_source,
"total_equity_source": total_equity_source,
"source_digest_sha256": _payload_digest(account_payload),
},
metadata=metadata,
)
129 changes: 127 additions & 2 deletions tests/test_schwab_portfolio.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,10 @@
from unittest.mock import patch

from quant_platform_kit.risk.engine import RiskEngine
from quant_platform_kit.schwab.portfolio import fetch_account_snapshot
from quant_platform_kit.schwab.portfolio import _optional_decimal_text, fetch_account_snapshot


_MISSING = object()


class FakeResponse:
Expand Down Expand Up @@ -60,14 +63,136 @@ def _install_fake_schwab_module(self):
)
return patch.dict(sys.modules, {"schwab": schwab_module, "schwab.client": client_module})

def _snapshot_with_balances(self, balances):
def _snapshot_with_balances(self, balances, *, account_type=_MISSING, cash_balance=_MISSING):
payload = FakeClient().get_account("abc123", "POSITIONS").json()
payload["securitiesAccount"]["currentBalances"] = balances
if account_type is not _MISSING:
payload["securitiesAccount"]["type"] = account_type
if cash_balance is not _MISSING:
payload["securitiesAccount"]["currentBalances"]["cashBalance"] = cash_balance
with self._install_fake_schwab_module(), patch.object(
FakeClient, "get_account", return_value=FakeResponse(payload)
):
return fetch_account_snapshot(FakeClient(), strategy_symbols=("TQQQ",))

def test_optional_native_facts_come_from_the_same_account_response(self) -> None:
class CountingClient(FakeClient):
def __init__(self):
self.account_number_calls = 0
self.account_calls = 0

def get_account_numbers(self):
self.account_number_calls += 1
return super().get_account_numbers()

def get_account(self, account_hash, fields):
self.account_calls += 1
payload = super().get_account(account_hash, fields).json()
payload["securitiesAccount"]["type"] = "MARGIN_UNKNOWN"
payload["securitiesAccount"]["currentBalances"].update(
{"cashBalance": "1234.5600", "buyingPower": 5000.0}
)
return FakeResponse(payload)

api_client = CountingClient()
with self._install_fake_schwab_module():
snapshot = fetch_account_snapshot(api_client, strategy_symbols=("TQQQ",))

self.assertEqual((1, 1), (api_client.account_number_calls, api_client.account_calls))
self.assertEqual("MARGIN_UNKNOWN", snapshot.metadata["broker_account_type"])
self.assertEqual("securitiesAccount.type", snapshot.metadata["broker_account_type_source"])
self.assertEqual("1234.5600", snapshot.metadata["broker_cash_balance"])
self.assertEqual("cashBalance", snapshot.metadata["broker_cash_balance_source"])
self.assertEqual(1000.0, snapshot.cash_balance)
self.assertEqual(5000.0, snapshot.buying_power)
self.assertEqual(1210.0, snapshot.total_equity)
self.assertEqual(("TQQQ",), tuple(position.symbol for position in snapshot.positions))

def test_raw_broker_account_type_is_preserved_only_as_a_bounded_token(self) -> None:
for raw, expected in (
("CASH", "CASH"),
("MARGIN", "MARGIN"),
("PROVIDER_UNKNOWN", "PROVIDER_UNKNOWN"),
(None, None),
("", None),
("contains space", None),
("type-1", None),
("é", None),
("X" * 33, None),
(42, None),
(True, None),
):
with self.subTest(raw=raw):
snapshot = self._snapshot_with_balances(
{"cashAvailableForTrading": 1000.0}, account_type=raw
)
if expected is None:
self.assertNotIn("broker_account_type", snapshot.metadata)
self.assertNotIn("broker_account_type_source", snapshot.metadata)
else:
self.assertEqual(expected, snapshot.metadata["broker_account_type"])
self.assertEqual(
"securitiesAccount.type",
snapshot.metadata["broker_account_type_source"],
)

def test_native_cash_balance_is_canonical_decimal_text_without_fallback(self) -> None:
for raw, expected in (
("123.4500", "123.4500"),
(123, "123"),
(12.5, "12.5"),
(0, "0"),
("-25.125", "-25.125"),
):
with self.subTest(raw=raw):
snapshot = self._snapshot_with_balances(
{"cashAvailableForTrading": 1000.0}, cash_balance=raw
)
self.assertEqual(expected, snapshot.metadata["broker_cash_balance"])
self.assertEqual("cashBalance", snapshot.metadata["broker_cash_balance_source"])
self.assertEqual(1000.0, snapshot.cash_balance)
self.assertEqual(1000.0, snapshot.buying_power)

for raw in (None, True, False, "NaN", "Infinity", "-Infinity", "bad", {}, []):
with self.subTest(invalid_raw=repr(raw)):
snapshot = self._snapshot_with_balances(
{"cashAvailableForTrading": 1000.0}, cash_balance=raw
)
self.assertNotIn("broker_cash_balance", snapshot.metadata)
self.assertNotIn("broker_cash_balance_source", snapshot.metadata)

for raw in (float("nan"), float("inf"), float("-inf")):
with self.subTest(non_finite_number=repr(raw)):
self.assertIsNone(_optional_decimal_text(raw))

missing = self._snapshot_with_balances({"cashAvailableForTrading": 1000.0})
self.assertNotIn("broker_cash_balance", missing.metadata)
self.assertNotIn("broker_cash_balance_source", missing.metadata)

def test_cash_balance_formatting_is_bounded_before_fixed_point_expansion(self) -> None:
for raw in ("1e1000", "1e-1000", "1e999999999", "1e-999999999", "9" * 129):
with self.subTest(raw_prefix=raw[:20], raw_length=len(raw)):
self.assertIsNone(_optional_decimal_text(raw))
snapshot = self._snapshot_with_balances(
{"cashAvailableForTrading": 1000.0}, cash_balance=raw
)
self.assertNotIn("broker_cash_balance", snapshot.metadata)
self.assertNotIn("broker_cash_balance_source", snapshot.metadata)
self.assertEqual(1000.0, snapshot.cash_balance)
self.assertEqual(1000.0, snapshot.buying_power)
self.assertEqual(1210.0, snapshot.total_equity)

self.assertEqual("125", _optional_decimal_text("1.25e2"))
self.assertEqual("0.0125", _optional_decimal_text("1.25e-2"))
self.assertEqual("0", _optional_decimal_text("0"))
self.assertEqual("-1.25", _optional_decimal_text("-1.25"))
self.assertEqual(
"123456789012345.12345678",
_optional_decimal_text("123456789012345.12345678"),
)
self.assertIsNone(_optional_decimal_text("1234567890123456"))
self.assertIsNone(_optional_decimal_text("0.123456789"))

def test_cash_available_for_trading_is_required(self) -> None:
with self.assertRaises(ValueError) as raised:
self._snapshot_with_balances({"liquidationValue": 2_500.0})
Expand Down
Loading