Files
Grid_Trading/tests/test_calculations.py
2026-07-14 16:21:32 +08:00

259 lines
7.6 KiB
Python

from datetime import date, timedelta
from decimal import Decimal
import pytest
from grid_trading.domain.calculations import (
CalculationError,
calculate_account_summary,
calculate_positions,
estimate_fees,
)
from grid_trading.domain.models import Account, FeeRules, Instrument, QuoteSnapshot, 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")
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,
),
]
quotes = {
1: QuoteSnapshot(
symbol="sh510300",
code="510300",
name="沪深300ETF",
price=Decimal("9.50"),
source="tencent",
)
}
[summary] = calculate_positions([instrument], trades, as_of=today, quote_snapshots=quotes)
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.current_price == Decimal("9.50")
assert summary.price_source == "tencent"
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_quote_snapshot_supplies_current_price_for_position_value():
today = date(2026, 7, 8)
instrument = Instrument(id=1, code="510300", name="沪深300ETF", manual_price=Decimal("3.90"))
trades = [
make_trade(
trade_id=1,
trade_date=today - timedelta(days=1),
side=TradeSide.BUY,
price="4.00",
quantity=1000,
trade_group=TradeGroup.BASE,
)
]
quotes = {
1: QuoteSnapshot(
symbol="sh510300",
code="510300",
name="沪深300ETF",
price=Decimal("4.12"),
source="tencent",
)
}
[summary] = calculate_positions([instrument], trades, as_of=today, quote_snapshots=quotes)
assert summary.current_price == Decimal("4.12")
assert summary.price_source == "tencent"
assert summary.market_value == Decimal("4120.00")
def test_missing_quote_does_not_use_manual_or_last_trade_price_for_current_price():
today = date(2026, 7, 8)
instrument = Instrument(id=1, code="510300", name="沪深300ETF", manual_price=Decimal("3.90"))
trades = [
make_trade(
trade_id=1,
trade_date=today - timedelta(days=1),
side=TradeSide.BUY,
price="4.00",
quantity=1000,
trade_group=TradeGroup.BASE,
)
]
[summary] = calculate_positions([instrument], trades, as_of=today)
assert summary.current_price is None
assert summary.price_source == "missing"
assert summary.market_value is None
assert summary.floating_pnl is None
def test_account_summary_caps_capital_usage_when_cash_is_negative():
today = date(2026, 7, 8)
account = Account(id=1, name="主账户", initial_cash=Decimal("100"))
instrument = Instrument(id=1, code="600588", name="用友网络")
trades = [
make_trade(
trade_id=1,
trade_date=today,
side=TradeSide.BUY,
price="1.50",
quantity=100,
trade_group=TradeGroup.BASE,
)
]
quotes = {
1: QuoteSnapshot(
symbol="sh600588",
code="600588",
name="用友网络",
price=Decimal("1.00"),
source="tencent",
)
}
positions = calculate_positions([instrument], trades, as_of=today, quote_snapshots=quotes)
summary = calculate_account_summary(account, positions, [], trades)
assert summary.cash == Decimal("-50.00")
assert summary.market_value == Decimal("100.00")
assert summary.total_assets == Decimal("50.00")
assert summary.capital_usage_rate == Decimal("1.0000")
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")