From 357c26bd866f34e3a9d6345df8851bf018c7f680 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8E=8B=E9=B9=8F?= Date: Wed, 8 Jul 2026 17:20:32 +0800 Subject: [PATCH] feat: add grid trading calculations --- src/grid_trading/domain/__init__.py | 1 + src/grid_trading/domain/calculations.py | 217 ++++++++++++++++++++++++ src/grid_trading/domain/models.py | 146 ++++++++++++++++ tests/test_calculations.py | 156 +++++++++++++++++ 4 files changed, 520 insertions(+) create mode 100644 src/grid_trading/domain/__init__.py create mode 100644 src/grid_trading/domain/calculations.py create mode 100644 src/grid_trading/domain/models.py create mode 100644 tests/test_calculations.py diff --git a/src/grid_trading/domain/__init__.py b/src/grid_trading/domain/__init__.py new file mode 100644 index 0000000..9646dfc --- /dev/null +++ b/src/grid_trading/domain/__init__.py @@ -0,0 +1 @@ +"""Domain models and calculations for grid trading.""" diff --git a/src/grid_trading/domain/calculations.py b/src/grid_trading/domain/calculations.py new file mode 100644 index 0000000..7e8e60d --- /dev/null +++ b/src/grid_trading/domain/calculations.py @@ -0,0 +1,217 @@ +from __future__ import annotations + +from collections import defaultdict +from dataclasses import dataclass +from datetime import date +from decimal import Decimal, ROUND_HALF_UP +from typing import Iterable, Mapping, Sequence + +from grid_trading.domain.models import ( + Account, + AccountSummary, + CashLedgerEntry, + FeeEstimate, + FeeRules, + Instrument, + PositionSummary, + Trade, + TradeGroup, + TradeSide, +) + +MONEY_PLACES = Decimal("0.01") +RATE_PLACES = Decimal("0.0001") + + +class CalculationError(ValueError): + pass + + +@dataclass +class _GroupState: + quantity: int = 0 + cost: Decimal = Decimal("0") + realized_pnl: Decimal = Decimal("0") + + +def money(value: Decimal) -> Decimal: + return value.quantize(MONEY_PLACES, rounding=ROUND_HALF_UP) + + +def price(value: Decimal) -> Decimal: + return value.quantize(MONEY_PLACES, rounding=ROUND_HALF_UP) + + +def estimate_fees( + side: TradeSide, + trade_price: Decimal, + quantity: int, + rules: FeeRules, +) -> FeeEstimate: + gross = trade_price * Decimal(quantity) + commission = max(gross * rules.commission_rate, rules.min_commission) + stamp_tax = gross * rules.stamp_tax_rate if side is TradeSide.SELL else Decimal("0") + transfer_fee = gross * rules.transfer_fee_rate + return FeeEstimate( + commission=money(commission), + stamp_tax=money(stamp_tax), + transfer_fee=money(transfer_fee), + ) + + +def calculate_positions( + instruments: Sequence[Instrument], + trades: Sequence[Trade], + *, + as_of: date, + instrument_cashflows: Mapping[int, Decimal] | None = None, +) -> list[PositionSummary]: + cashflows = instrument_cashflows or {} + instrument_by_id = {instrument.id: instrument for instrument in instruments if instrument.id is not None} + states: dict[int, dict[TradeGroup, _GroupState]] = defaultdict( + lambda: {group: _GroupState() for group in TradeGroup} + ) + net_invested: dict[int, Decimal] = defaultdict(lambda: Decimal("0")) + last_trade_price: dict[int, Decimal] = {} + today_buys: dict[int, int] = defaultdict(int) + + for trade in sorted(trades, key=lambda item: (item.trade_date, item.id or 0)): + if trade.trade_date > as_of: + continue + _validate_trade(trade) + instrument = instrument_by_id.get(trade.instrument_id) + if instrument is None: + raise CalculationError(f"Trade references unknown instrument {trade.instrument_id}") + + group_state = states[trade.instrument_id][trade.trade_group] + gross = money(trade.gross_amount) + fees = money(trade.total_fee) + last_trade_price[trade.instrument_id] = trade.price + + if trade.side is TradeSide.BUY: + group_state.quantity += trade.quantity + group_state.cost = money(group_state.cost + gross + fees) + net_invested[trade.instrument_id] = money(net_invested[trade.instrument_id] + gross + fees) + if trade.trade_date == as_of: + today_buys[trade.instrument_id] += trade.quantity + else: + if trade.quantity > group_state.quantity: + raise CalculationError( + f"Insufficient {trade.trade_group.value} position for {instrument.code}" + ) + average_cost = group_state.cost / Decimal(group_state.quantity) + released_cost = money(average_cost * Decimal(trade.quantity)) + net_income = money(gross - fees) + realized = money(net_income - released_cost) + group_state.quantity -= trade.quantity + group_state.cost = money(group_state.cost - released_cost) + group_state.realized_pnl = money(group_state.realized_pnl + realized) + net_invested[trade.instrument_id] = money(net_invested[trade.instrument_id] - net_income) + + summaries: list[PositionSummary] = [] + for instrument in instruments: + if instrument.id is None: + continue + group_states = states[instrument.id] + base_quantity = group_states[TradeGroup.BASE].quantity + grid_quantity = group_states[TradeGroup.GRID].quantity + other_quantity = group_states[TradeGroup.OTHER].quantity + total_quantity = base_quantity + grid_quantity + other_quantity + remaining_cost = money(sum((state.cost for state in group_states.values()), Decimal("0"))) + realized_pnl = money(sum((state.realized_pnl for state in group_states.values()), Decimal("0"))) + grid_profit = money(group_states[TradeGroup.GRID].realized_pnl) + current_price, price_source = _resolve_current_price(instrument, last_trade_price.get(instrument.id)) + market_value = money(current_price * Decimal(total_quantity)) if current_price is not None else None + floating_pnl = money(market_value - remaining_cost) if market_value is not None else None + available_quantity = max(0, total_quantity - today_buys[instrument.id]) + adjusted_net_invested = money(net_invested[instrument.id] - cashflows.get(instrument.id, Decimal("0"))) + + summaries.append( + PositionSummary( + instrument_id=instrument.id, + code=instrument.code, + name=instrument.name, + market=instrument.market, + current_price=current_price, + price_source=price_source, + total_quantity=total_quantity, + available_quantity=available_quantity, + base_quantity=base_quantity, + grid_quantity=grid_quantity, + other_quantity=other_quantity, + remaining_cost=remaining_cost, + cost_price=_per_share(remaining_cost, total_quantity), + position_breakeven_price=_per_share(remaining_cost - grid_profit, total_quantity), + account_breakeven_price=_per_share(adjusted_net_invested, total_quantity), + realized_pnl=realized_pnl, + grid_profit=grid_profit, + floating_pnl=floating_pnl, + market_value=market_value, + ) + ) + + return summaries + + +def calculate_account_summary( + account: Account, + positions: Iterable[PositionSummary], + cash_ledger_entries: Sequence[CashLedgerEntry], + trades: Sequence[Trade], +) -> AccountSummary: + trade_cash = Decimal("0") + for trade in trades: + gross = money(trade.gross_amount) + fees = money(trade.total_fee) + if trade.side is TradeSide.BUY: + trade_cash -= gross + fees + else: + trade_cash += gross - fees + + ledger_cash = sum((entry.amount for entry in cash_ledger_entries), Decimal("0")) + cash = money(account.initial_cash + trade_cash + ledger_cash) + market_value = money( + sum((position.market_value or Decimal("0") for position in positions), Decimal("0")) + ) + floating_pnl = money( + sum((position.floating_pnl or Decimal("0") for position in positions), Decimal("0")) + ) + total_assets = money(cash + market_value) + capital_usage_rate = ( + (market_value / total_assets).quantize(RATE_PLACES, rounding=ROUND_HALF_UP) + if total_assets > 0 + else Decimal("0") + ) + return AccountSummary( + total_assets=total_assets, + cash=cash, + market_value=market_value, + floating_pnl=floating_pnl, + capital_usage_rate=capital_usage_rate, + ) + + +def _validate_trade(trade: Trade) -> None: + if trade.price < 0: + raise CalculationError("Trade price cannot be negative") + if trade.quantity <= 0: + raise CalculationError("Trade quantity must be positive") + if trade.total_fee < 0: + raise CalculationError("Trade fees cannot be negative") + + +def _resolve_current_price( + instrument: Instrument, + fallback_trade_price: Decimal | None, +) -> tuple[Decimal | None, str]: + if instrument.manual_price is not None: + return instrument.manual_price, "manual" + if fallback_trade_price is not None: + return fallback_trade_price, "last_trade" + return None, "missing" + + +def _per_share(total: Decimal, quantity: int) -> Decimal | None: + if quantity <= 0: + return None + return price(total / Decimal(quantity)) diff --git a/src/grid_trading/domain/models.py b/src/grid_trading/domain/models.py new file mode 100644 index 0000000..f23cbba --- /dev/null +++ b/src/grid_trading/domain/models.py @@ -0,0 +1,146 @@ +from __future__ import annotations + +from dataclasses import dataclass +from datetime import date +from decimal import Decimal +from enum import StrEnum + + +class TradeSide(StrEnum): + BUY = "buy" + SELL = "sell" + + +class TradeGroup(StrEnum): + BASE = "base" + GRID = "grid" + OTHER = "other" + + +@dataclass(frozen=True) +class Account: + id: int | None + name: str + initial_cash: Decimal + notes: str = "" + + +@dataclass(frozen=True) +class Instrument: + id: int | None + code: str + name: str + market: str = "A" + lot_size: int = 100 + manual_price: Decimal | None = None + allow_odd_lot: bool = False + active: bool = True + + +@dataclass(frozen=True) +class StrategyTemplate: + id: int | None + name: str + grid_spacing_pct: Decimal + amount_per_grid: Decimal + base_target_amount: Decimal + max_position_amount: Decimal + min_lot: int = 100 + is_default: bool = False + + +@dataclass(frozen=True) +class StrategyOverride: + instrument_id: int + template_id: int | None = None + grid_spacing_pct: Decimal | None = None + amount_per_grid: Decimal | None = None + base_target_amount: Decimal | None = None + max_position_amount: Decimal | None = None + min_lot: int | None = None + + +@dataclass(frozen=True) +class Trade: + id: int | None + account_id: int + instrument_id: int + trade_date: date + side: TradeSide + price: Decimal + quantity: int + commission: Decimal = Decimal("0") + stamp_tax: Decimal = Decimal("0") + transfer_fee: Decimal = Decimal("0") + trade_group: TradeGroup = TradeGroup.GRID + notes: str = "" + + @property + def gross_amount(self) -> Decimal: + return self.price * Decimal(self.quantity) + + @property + def total_fee(self) -> Decimal: + return self.commission + self.stamp_tax + self.transfer_fee + + +@dataclass(frozen=True) +class CashLedgerEntry: + id: int | None + account_id: int + entry_date: date + amount: Decimal + category: str + instrument_id: int | None = None + notes: str = "" + + +@dataclass(frozen=True) +class FeeRules: + commission_rate: Decimal + min_commission: Decimal + stamp_tax_rate: Decimal + transfer_fee_rate: Decimal + + +@dataclass(frozen=True) +class FeeEstimate: + commission: Decimal + stamp_tax: Decimal + transfer_fee: Decimal + + @property + def total(self) -> Decimal: + return self.commission + self.stamp_tax + self.transfer_fee + + +@dataclass(frozen=True) +class PositionSummary: + instrument_id: int + code: str + name: str + market: str + current_price: Decimal | None + price_source: str + total_quantity: int + available_quantity: int + base_quantity: int + grid_quantity: int + other_quantity: int + remaining_cost: Decimal + cost_price: Decimal | None + position_breakeven_price: Decimal | None + account_breakeven_price: Decimal | None + realized_pnl: Decimal + grid_profit: Decimal + floating_pnl: Decimal | None + market_value: Decimal | None + + +@dataclass(frozen=True) +class AccountSummary: + total_assets: Decimal + cash: Decimal + market_value: Decimal + floating_pnl: Decimal + capital_usage_rate: Decimal diff --git a/tests/test_calculations.py b/tests/test_calculations.py new file mode 100644 index 0000000..aab7930 --- /dev/null +++ b/tests/test_calculations.py @@ -0,0 +1,156 @@ +from datetime import date, timedelta +from decimal import Decimal + +import pytest + +from grid_trading.domain.calculations import CalculationError, calculate_positions, estimate_fees +from grid_trading.domain.models import FeeRules, Instrument, Trade, TradeGroup, TradeSide + + +def make_trade( + *, + trade_id: int, + instrument_id: int = 1, + trade_date: date, + side: TradeSide, + price: str, + quantity: int, + commission: str = "0", + stamp_tax: str = "0", + transfer_fee: str = "0", + trade_group: TradeGroup = TradeGroup.GRID, +) -> Trade: + return Trade( + id=trade_id, + account_id=1, + instrument_id=instrument_id, + trade_date=trade_date, + side=side, + price=Decimal(price), + quantity=quantity, + commission=Decimal(commission), + stamp_tax=Decimal(stamp_tax), + transfer_fee=Decimal(transfer_fee), + trade_group=trade_group, + ) + + +def test_buy_sell_grid_profit_and_breakeven(): + today = date(2026, 7, 8) + instrument = Instrument(id=1, code="510300", name="沪深300ETF", manual_price=Decimal("9.50")) + trades = [ + make_trade( + trade_id=1, + trade_date=today - timedelta(days=2), + side=TradeSide.BUY, + price="10", + quantity=100, + commission="1", + trade_group=TradeGroup.GRID, + ), + make_trade( + trade_id=2, + trade_date=today - timedelta(days=1), + side=TradeSide.SELL, + price="11", + quantity=100, + commission="1", + trade_group=TradeGroup.GRID, + ), + make_trade( + trade_id=3, + trade_date=today, + side=TradeSide.BUY, + price="9", + quantity=100, + commission="1", + trade_group=TradeGroup.BASE, + ), + ] + + [summary] = calculate_positions([instrument], trades, as_of=today) + + assert summary.total_quantity == 100 + assert summary.base_quantity == 100 + assert summary.grid_quantity == 0 + assert summary.remaining_cost == Decimal("901.00") + assert summary.realized_pnl == Decimal("98.00") + assert summary.grid_profit == Decimal("98.00") + assert summary.cost_price == Decimal("9.01") + assert summary.position_breakeven_price == Decimal("8.03") + assert summary.account_breakeven_price == Decimal("8.03") + assert summary.market_value == Decimal("950.00") + assert summary.floating_pnl == Decimal("49.00") + + +def test_t_plus_one_available_quantity_excludes_today_buys(): + today = date(2026, 7, 8) + instrument = Instrument(id=1, code="600000", name="浦发银行", manual_price=Decimal("10")) + trades = [ + make_trade( + trade_id=1, + trade_date=today - timedelta(days=1), + side=TradeSide.BUY, + price="10", + quantity=200, + trade_group=TradeGroup.BASE, + ), + make_trade( + trade_id=2, + trade_date=today, + side=TradeSide.BUY, + price="9.8", + quantity=100, + trade_group=TradeGroup.BASE, + ), + ] + + [summary] = calculate_positions([instrument], trades, as_of=today) + + assert summary.total_quantity == 300 + assert summary.available_quantity == 200 + + +def test_sell_more_than_group_position_raises(): + today = date(2026, 7, 8) + instrument = Instrument(id=1, code="159915", name="创业板ETF") + trades = [ + make_trade( + trade_id=1, + trade_date=today - timedelta(days=1), + side=TradeSide.BUY, + price="2", + quantity=100, + trade_group=TradeGroup.GRID, + ), + make_trade( + trade_id=2, + trade_date=today, + side=TradeSide.SELL, + price="2.1", + quantity=200, + trade_group=TradeGroup.GRID, + ), + ] + + with pytest.raises(CalculationError, match="Insufficient grid position"): + calculate_positions([instrument], trades, as_of=today) + + +def test_estimate_fees_uses_min_commission_and_sell_tax(): + rules = FeeRules( + commission_rate=Decimal("0.00025"), + min_commission=Decimal("5"), + stamp_tax_rate=Decimal("0.0005"), + transfer_fee_rate=Decimal("0.00001"), + ) + + buy_fees = estimate_fees(TradeSide.BUY, Decimal("10"), 100, rules) + sell_fees = estimate_fees(TradeSide.SELL, Decimal("10"), 100, rules) + + assert buy_fees.commission == Decimal("5.00") + assert buy_fees.stamp_tax == Decimal("0.00") + assert buy_fees.transfer_fee == Decimal("0.01") + assert sell_fees.commission == Decimal("5.00") + assert sell_fees.stamp_tax == Decimal("0.50") + assert sell_fees.transfer_fee == Decimal("0.01")