Skip to content
This repository was archived by the owner on Apr 21, 2026. It is now read-only.

Commit fbe10f9

Browse files
committed
lint fixes
1 parent 9980444 commit fbe10f9

14 files changed

Lines changed: 86 additions & 174 deletions

pytradebacktest/backtest.py

Lines changed: 6 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -13,23 +13,23 @@
1313

1414

1515
class Backtest:
16-
1716
def __init__(
1817
self,
1918
data: MarketData,
2019
kstrategy: Type[FxStrategy],
2120
cash: float,
22-
comission: float = 0.0,
21+
commission: float = 0.0,
2322
margin: float = 1.0,
23+
spread: float = 0.0,
2424
):
2525
self.data = data
2626
self._strategy = kstrategy
2727
self.cash = cash
28-
self.comission = comission
28+
self.commission = commission
2929
self.margin = margin
30+
self.spread = spread
3031

3132
async def run(self, **kwargs):
32-
3333
# Monkey patch indicators so their update does not
3434
# need to recalculate after each increment, instead
3535
# store their initial values from the full data context
@@ -41,14 +41,13 @@ def increment_indicator(self):
4141

4242
Indicator._update = increment_indicator
4343

44-
self.broker = BacktestBroker(self.data, self.cash, self.comission, self.margin)
44+
self.broker = BacktestBroker(self.data, self.cash, self.commission, self.margin, spread=self.spread)
4545

4646
strategy = self._strategy(self.broker, self.data, **kwargs)
4747
strategy.init()
4848

4949
with ProgressBar(max_value=len(self.data), redirect_stdout=True) as bar:
5050
while self.data.next():
51-
5251
self.broker.next()
5352
strategy.next()
5453

@@ -68,9 +67,7 @@ def plot(self):
6867
for instrument in instruments:
6968
goog_data = self.data.get(instrument, Granularity.M1)
7069
trades = [
71-
trade
72-
for trade in self.broker.closed_trades
73-
if trade.instrument == instrument
70+
trade for trade in self.broker.closed_trades if trade.instrument == instrument
7471
]
7572
equity_df = pd.DataFrame(
7673
self.broker._equity, index=self.data._market_index, columns=["Equity"]

pytradebacktest/broker.py

Lines changed: 10 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,6 @@
1414

1515

1616
class BacktestBroker(IBroker):
17-
1817
def __init__(
1918
self,
2019
data: MarketData,
@@ -24,6 +23,7 @@ def __init__(
2423
trade_on_close=False,
2524
hedging=False,
2625
exclusive_orders=False,
26+
spread=0.0,
2727
):
2828
self._data = data
2929
self._cash = cash
@@ -32,6 +32,7 @@ def __init__(
3232
self._trade_on_close = trade_on_close
3333
self._hedging = hedging
3434
self._exclusive_orders = exclusive_orders
35+
self._spread = spread
3536

3637
self._equity = np.tile(np.nan, len(self._data))
3738
self.orders: list[Order] = []
@@ -53,7 +54,6 @@ def leverage(self) -> float:
5354
return self._leverage
5455

5556
def order(self, order: Order):
56-
5757
# If exclusive orders (each new order auto-closes previous orders/position),
5858
# cancel all non-contingent orders and close all open trades beforehand
5959
if self._exclusive_orders:
@@ -65,14 +65,10 @@ def order(self, order: Order):
6565

6666
self.orders.append(order)
6767

68-
def load_instrument_candles(
69-
self, instrument: Instrument, granularity: Granularity, count: int
70-
):
68+
def load_instrument_candles(self, instrument: Instrument, granularity: Granularity, count: int):
7169
self._data.load_instrument_candles(instrument, granularity, count)
7270

73-
def subscribe(
74-
self, instrument: Instrument, granularity: Granularity
75-
) -> IInstrumentData:
71+
def subscribe(self, instrument: Instrument, granularity: Granularity) -> IInstrumentData:
7672
return self._data.get(instrument, granularity)
7773

7874
def get_position(self, instrument: Instrument) -> Position:
@@ -81,9 +77,7 @@ def get_position(self, instrument: Instrument) -> Position:
8177
def close_position(self, instrument: Instrument):
8278
position = self.get_position(instrument)
8379
for trade in position.trades:
84-
self.orders.insert(
85-
0, Order(trade.instrument, -trade.size, parent_trade=trade)
86-
)
80+
self.orders.insert(0, Order(trade.instrument, -trade.size, parent_trade=trade))
8781

8882
def next(self):
8983
self._process_orders()
@@ -107,7 +101,6 @@ def _update_equity(self):
107101
raise OutOfMoneyError
108102

109103
def _process_orders(self):
110-
111104
reprocess_orders = False
112105
for order in list(self.orders):
113106
_data = self._get_instrument_data(order.instrument)
@@ -119,6 +112,7 @@ def _process_orders(self):
119112
self._commission,
120113
self._leverage,
121114
self.margin_available,
115+
self._spread,
122116
)
123117
# Related SL/TP order already removed
124118
if order not in self.orders:
@@ -172,9 +166,7 @@ def _process_contingent_order(self, ctx: OrderContext):
172166
_order_size = ctx.adjusted_size
173167

174168
if trade in self.trades:
175-
closed = self._reduce_trade(
176-
trade, _order_size, ctx.entry_price, ctx.entry_time
177-
)
169+
closed = self._reduce_trade(trade, _order_size, ctx.entry_price, ctx.entry_time)
178170
# If this is a SL/TP closing the trade already removed it.
179171
if closed and order in self.orders:
180172
self.orders.remove(order)
@@ -230,9 +222,7 @@ def _update_position(self, ctx: OrderContext):
230222
def _open_trade(self, ctx: OrderContext, tag: Optional[str] = None):
231223
order = ctx.order
232224
size = ctx.adjusted_size
233-
trade = Trade(
234-
order.instrument, size, ctx.entry_price, ctx.entry_time, ctx.data, tag
235-
)
225+
trade = Trade(order.instrument, size, ctx.entry_price, ctx.entry_time, ctx.data, tag)
236226
self.trades.append(trade)
237227

238228
if order.take_profit_on_fill:
@@ -255,9 +245,7 @@ def _open_trade(self, ctx: OrderContext, tag: Optional[str] = None):
255245
trade.sl = sl_order
256246
self.orders.insert(0, sl_order)
257247

258-
def _reduce_trade(
259-
self, trade: Trade, size: int, price: float, timestamp: Timestamp
260-
):
248+
def _reduce_trade(self, trade: Trade, size: int, price: float, timestamp: Timestamp):
261249
size_left = trade.size + size
262250
closed = False
263251

@@ -281,9 +269,7 @@ def _reduce_trade(
281269

282270
def close_trades(self):
283271
for trade in self.trades:
284-
self.orders.insert(
285-
0, Order(trade.instrument, -trade.size, parent_trade=trade)
286-
)
272+
self.orders.insert(0, Order(trade.instrument, -trade.size, parent_trade=trade))
287273

288274
def _close_trade(self, trade: Trade, price: float, timestamp: Timestamp):
289275
self.trades.remove(trade)

pytradebacktest/data.py

Lines changed: 22 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -6,17 +6,14 @@
66
import pandas as pd
77
from pandas import DatetimeIndex, Timestamp
88
from pytrade.events.event import Event
9-
from pytrade.instruments import Granularity, Instrument
9+
from pytrade.instruments import Granularity, Instrument, UPDATE_MAP
1010
from pytrade.interfaces.data import IDataContext, IInstrumentData
1111

1212
from pytradebacktest.utils import load_csv
1313

1414

1515
class InstrumentData(IInstrumentData):
16-
17-
def __init__(
18-
self, instrument: Instrument, granularity: Granularity, df: pd.DataFrame
19-
):
16+
def __init__(self, instrument: Instrument, granularity: Granularity, df: pd.DataFrame):
2017
self.__df = df
2118
self.__i = len(df) - 1
2219
self.__pip: Optional[float] = None
@@ -37,9 +34,7 @@ def granularity(self) -> Granularity:
3734

3835
@property
3936
def df(self) -> pd.DataFrame:
40-
return (
41-
self.__df.iloc[: self.__i + 1] if self.__i < len(self.__df) else self.__df
42-
)
37+
return self.__df.iloc[: self.__i + 1] if self.__i < len(self.__df) else self.__df
4338

4439
@property
4540
def on_update(self) -> Event:
@@ -74,28 +69,24 @@ def index(self, value: Timestamp):
7469

7570

7671
class DataSource:
77-
7872
def __init__(self, instrument: Instrument, granularity: Granularity):
7973
self.instrument = instrument
8074
self.granularity = granularity
8175

8276

8377
class CsvDataSource(DataSource):
84-
8578
def __init__(self, path: str, instrument: Instrument, granularity: Granularity):
8679
super().__init__(instrument, granularity)
8780
self.path = path
8881

8982

9083
class MarketDataLoader:
91-
9284
@abstractmethod
9385
def load(self) -> list[InstrumentData]:
9486
raise NotImplementedError
9587

9688

9789
class CsvMarketDataLoader(MarketDataLoader):
98-
9990
def __init__(self, sources: list[CsvDataSource]):
10091
self.sources = sources
10192

@@ -116,7 +107,6 @@ def load(self) -> list[InstrumentData]:
116107

117108

118109
class MarketData(IDataContext):
119-
120110
_index: pd.Timestamp
121111

122112
def __init__(self, loader: MarketDataLoader):
@@ -137,9 +127,7 @@ def universe(self):
137127
def index(self) -> pd.Timestamp:
138128
return self._index
139129

140-
def load_instrument_candles(
141-
self, instrument: Instrument, granularity: Granularity, count: int
142-
):
130+
def load_instrument_candles(self, instrument: Instrument, granularity: Granularity, count: int):
143131
_data = self.get(instrument, granularity)
144132
_instrument_timestamp: pd.Timestamp = _data.df.index[count]
145133
if _instrument_timestamp > self._index:
@@ -172,7 +160,6 @@ def next(self):
172160
return result
173161

174162
def __next(self):
175-
176163
# Slice index incase some candles were loaded prior to starting
177164
# the test run
178165
_index = self._market_index[self.i :]
@@ -185,8 +172,24 @@ def __next(self):
185172
yield False
186173

187174
def get(self, instrument: Instrument, granularity: Granularity) -> IInstrumentData:
188-
return next(
175+
_data: IInstrumentData = next(
189176
src
190177
for src in self._sources
191-
if src.instrument == instrument and src.granularity == granularity
178+
if src.instrument == instrument and src.granularity == Granularity.M1
192179
)
180+
181+
if not _data:
182+
raise RuntimeError(f"Data not loaded for {instrument}")
183+
184+
if granularity != Granularity.M1:
185+
_df = _data.df.resample(UPDATE_MAP[granularity]).agg(
186+
{
187+
"open": "first", # First value in each 15-min period
188+
"high": "max", # Maximum value in each 15-min period
189+
"low": "min", # Minimum value in each 15-min period
190+
"close": "last", # Last value in each 15-min period
191+
}
192+
)
193+
_data = InstrumentData(instrument=instrument, granularity=granularity, df=_df)
194+
195+
return _data

pytradebacktest/order.py

Lines changed: 8 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,6 @@
77

88

99
class OrderContext:
10-
1110
def __init__(
1211
self,
1312
order: Order,
@@ -16,13 +15,15 @@ def __init__(
1615
commission: float,
1716
leverage: float,
1817
margin_available: float,
18+
spread: float = 0.0,
1919
):
2020
self._order = order
2121
self._data = data
2222
self._trade_on_close = trade_on_close
2323
self._commission = commission
2424
self._leverage = leverage
2525
self._margin_availalbe = margin_available
26+
self._spread = spread
2627

2728
@property
2829
def order(self):
@@ -82,9 +83,12 @@ def entry_price(self):
8283

8384
@property
8485
def adjusted_entry_price(self):
85-
# Need to update to account for currency pairs and base vs counter currency
86-
# as well as spread
87-
return self.entry_price * (1 + copysign(self._commission, self.order.size))
86+
# Apply commission and spread
87+
# For longs (buy): pay ask price (mid + spread/2)
88+
# For shorts (sell): receive bid price (mid - spread/2)
89+
spread_adjustment = copysign(self._spread / 2, self.order.size)
90+
commission_adjustment = copysign(self._commission, self.order.size)
91+
return self.entry_price * (1 + commission_adjustment) + spread_adjustment
8892

8993
@property
9094
def adjusted_size(self):

pytradebacktest/plot.py

Lines changed: 1 addition & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -8,10 +8,7 @@
88

99

1010
def plot(data: IInstrumentData, equity: pd.DataFrame, trades: list[Trade]):
11-
12-
fig = make_subplots(
13-
rows=2, cols=1, row_heights=[0.2, 0.8], subplot_titles=("Equity", "Trades")
14-
)
11+
fig = make_subplots(rows=2, cols=1, row_heights=[0.2, 0.8], subplot_titles=("Equity", "Trades"))
1512

1613
_plot_equity(fig, equity)
1714
_plot_ohlc(fig, data)

pytradebacktest/position.py

Lines changed: 2 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -24,9 +24,7 @@ def __bool__(self):
2424

2525
@property
2626
def trades(self):
27-
return [
28-
trade for trade in self.__trades if trade.instrument == self.__instrument
29-
]
27+
return [trade for trade in self.__trades if trade.instrument == self.__instrument]
3028

3129
@property
3230
def size(self) -> float:
@@ -57,6 +55,4 @@ def is_short(self) -> bool:
5755
return self.size < 0
5856

5957
def __repr__(self):
60-
return (
61-
f"<Position[{self.__instrument}]: {self.size} ({len(self.trades)} trades)>"
62-
)
58+
return f"<Position[{self.__instrument}]: {self.size} ({len(self.trades)} trades)>"

0 commit comments

Comments
 (0)