From 2dc0b26ad4f1255b7ba44c3fca49fa52036b1a18 Mon Sep 17 00:00:00 2001 From: Pigbibi <20649888+Pigbibi@users.noreply.github.com> Date: Thu, 1 Oct 2026 11:19:22 +0800 Subject: [PATCH] fix(schwab): retain bounded native cash and account type facts Co-Authored-By: Codex --- src/quant_platform_kit/schwab/portfolio.py | 79 +++++++++++-- tests/test_schwab_portfolio.py | 129 ++++++++++++++++++++- 2 files changed, 198 insertions(+), 10 deletions(-) diff --git a/src/quant_platform_kit/schwab/portfolio.py b/src/quant_platform_kit/schwab/portfolio.py index 1d326b17..79685494 100644 --- a/src/quant_platform_kit/schwab/portfolio.py +++ b/src/quant_platform_kit/schwab/portfolio.py @@ -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.""" @@ -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. @@ -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, ) diff --git a/tests/test_schwab_portfolio.py b/tests/test_schwab_portfolio.py index 61cb40a6..0c8626c5 100644 --- a/tests/test_schwab_portfolio.py +++ b/tests/test_schwab_portfolio.py @@ -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: @@ -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})