259 lines
7.6 KiB
Python
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")
|