PyAlgoEngine 0.12.3__cp315-cp315-win_amd64.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (156) hide show
  1. algo_engine/__infra__.pxd +3 -0
  2. algo_engine/__init__.pxd +3 -0
  3. algo_engine/__init__.py +107 -0
  4. algo_engine/apps/__init__.py +17 -0
  5. algo_engine/apps/backtest/__init__.py +20 -0
  6. algo_engine/apps/backtest/doc_server.py +328 -0
  7. algo_engine/apps/backtest/static/styles/dash.css +48 -0
  8. algo_engine/apps/backtest/templates/dash.html +20 -0
  9. algo_engine/apps/backtest/templates/index.html +40 -0
  10. algo_engine/apps/backtest/tester.py +252 -0
  11. algo_engine/apps/backtest/web_app.py +125 -0
  12. algo_engine/apps/bokeh_server.py +245 -0
  13. algo_engine/apps/demo/__init__.py +0 -0
  14. algo_engine/apps/demo/test.py +40 -0
  15. algo_engine/apps/sim_input/__init__.py +23 -0
  16. algo_engine/apps/sim_input/client.py +412 -0
  17. algo_engine/apps/sim_input/sim_keyboard.py +88 -0
  18. algo_engine/apps/sim_input/sim_mouse.py +137 -0
  19. algo_engine/apps/sim_input/window.py +162 -0
  20. algo_engine/backtest/__init__.py +19 -0
  21. algo_engine/backtest/__main__.py +51 -0
  22. algo_engine/backtest/metrics.py +179 -0
  23. algo_engine/backtest/replay.py +635 -0
  24. algo_engine/backtest/sim_match.py +506 -0
  25. algo_engine/base/__infra__.pxd +3 -0
  26. algo_engine/base/__init__.pxd +3 -0
  27. algo_engine/base/__init__.py +60 -0
  28. algo_engine/base/c_allocator_protocol.c +11608 -0
  29. algo_engine/base/c_allocator_protocol.cp315-win_amd64.pyd +0 -0
  30. algo_engine/base/c_allocator_protocol.pxd +24 -0
  31. algo_engine/base/c_allocator_protocol.pyi +68 -0
  32. algo_engine/base/c_allocator_protocol.pyx +111 -0
  33. algo_engine/base/c_intern_string.c +5908 -0
  34. algo_engine/base/c_intern_string.cp315-win_amd64.pyd +0 -0
  35. algo_engine/base/c_intern_string.pxd +14 -0
  36. algo_engine/base/c_intern_string.pyi +22 -0
  37. algo_engine/base/c_intern_string.pyx +17 -0
  38. algo_engine/base/c_market_data/__infra__.pxd +196 -0
  39. algo_engine/base/c_market_data/__init__.pxd +196 -0
  40. algo_engine/base/c_market_data/__init__.py +24 -0
  41. algo_engine/base/c_market_data/c_candlestick.c +18996 -0
  42. algo_engine/base/c_market_data/c_candlestick.cp315-win_amd64.pyd +0 -0
  43. algo_engine/base/c_market_data/c_candlestick.pxd +27 -0
  44. algo_engine/base/c_market_data/c_candlestick.pyi +217 -0
  45. algo_engine/base/c_market_data/c_candlestick.pyx +255 -0
  46. algo_engine/base/c_market_data/c_internal.c +14059 -0
  47. algo_engine/base/c_market_data/c_internal.cp315-win_amd64.pyd +0 -0
  48. algo_engine/base/c_market_data/c_internal.pxd +14 -0
  49. algo_engine/base/c_market_data/c_internal.pyi +47 -0
  50. algo_engine/base/c_market_data/c_internal.pyx +41 -0
  51. algo_engine/base/c_market_data/c_market_data.c +30420 -0
  52. algo_engine/base/c_market_data/c_market_data.cp315-win_amd64.pyd +0 -0
  53. algo_engine/base/c_market_data/c_market_data.h +1464 -0
  54. algo_engine/base/c_market_data/c_market_data.pxd +414 -0
  55. algo_engine/base/c_market_data/c_market_data.pyi +550 -0
  56. algo_engine/base/c_market_data/c_market_data.pyx +701 -0
  57. algo_engine/base/c_market_data/c_market_data_buffer.c +29407 -0
  58. algo_engine/base/c_market_data/c_market_data_buffer.cp315-win_amd64.pyd +0 -0
  59. algo_engine/base/c_market_data/c_market_data_buffer.h +941 -0
  60. algo_engine/base/c_market_data/c_market_data_buffer.pxd +143 -0
  61. algo_engine/base/c_market_data/c_market_data_buffer.pyi +355 -0
  62. algo_engine/base/c_market_data/c_market_data_buffer.pyx +571 -0
  63. algo_engine/base/c_market_data/c_market_data_config.h +29 -0
  64. algo_engine/base/c_market_data/c_tick.c +44303 -0
  65. algo_engine/base/c_market_data/c_tick.cp315-win_amd64.pyd +0 -0
  66. algo_engine/base/c_market_data/c_tick.pxd +53 -0
  67. algo_engine/base/c_market_data/c_tick.pyi +466 -0
  68. algo_engine/base/c_market_data/c_tick.pyx +673 -0
  69. algo_engine/base/c_market_data/c_trade_utils.c +28702 -0
  70. algo_engine/base/c_market_data/c_trade_utils.cp315-win_amd64.pyd +0 -0
  71. algo_engine/base/c_market_data/c_trade_utils.pxd +53 -0
  72. algo_engine/base/c_market_data/c_trade_utils.pyi +602 -0
  73. algo_engine/base/c_market_data/c_trade_utils.pyx +609 -0
  74. algo_engine/base/c_market_data/c_transaction.c +23558 -0
  75. algo_engine/base/c_market_data/c_transaction.cp315-win_amd64.pyd +0 -0
  76. algo_engine/base/c_market_data/c_transaction.pxd +27 -0
  77. algo_engine/base/c_market_data/c_transaction.pyi +433 -0
  78. algo_engine/base/c_market_data/c_transaction.pyx +460 -0
  79. algo_engine/base/console_utils.py +1070 -0
  80. algo_engine/base/finance_decimal.py +258 -0
  81. algo_engine/base/telemetrics.py +18 -0
  82. algo_engine/engine/__infra__.pxd +10 -0
  83. algo_engine/engine/__init__.pxd +10 -0
  84. algo_engine/engine/__init__.py +40 -0
  85. algo_engine/engine/algo_engine.py +904 -0
  86. algo_engine/engine/c_event_engine.c +16325 -0
  87. algo_engine/engine/c_event_engine.cp315-win_amd64.pyd +0 -0
  88. algo_engine/engine/c_event_engine.pxd +25 -0
  89. algo_engine/engine/c_event_engine.pyi +68 -0
  90. algo_engine/engine/c_market_engine.c +24258 -0
  91. algo_engine/engine/c_market_engine.cp315-win_amd64.pyd +0 -0
  92. algo_engine/engine/c_market_engine.pxd +87 -0
  93. algo_engine/engine/c_market_engine.pyi +357 -0
  94. algo_engine/engine/event_engine.py +53 -0
  95. algo_engine/engine/trade_engine.py +2037 -0
  96. algo_engine/exchange_profile/__infra__.pxd +100 -0
  97. algo_engine/exchange_profile/__init__.pxd +100 -0
  98. algo_engine/exchange_profile/__init__.py +53 -0
  99. algo_engine/exchange_profile/c_ex_profile_base.c +87 -0
  100. algo_engine/exchange_profile/c_ex_profile_base.h +1204 -0
  101. algo_engine/exchange_profile/c_ex_profile_cn.c +968 -0
  102. algo_engine/exchange_profile/c_ex_profile_cn.h +39 -0
  103. algo_engine/exchange_profile/c_exchange_profile.c +52375 -0
  104. algo_engine/exchange_profile/c_exchange_profile.cp315-win_amd64.pyd +0 -0
  105. algo_engine/exchange_profile/c_exchange_profile.pxd +336 -0
  106. algo_engine/exchange_profile/c_exchange_profile.pyi +883 -0
  107. algo_engine/exchange_profile/c_exchange_profile.pyx +1495 -0
  108. algo_engine/exchange_profile/c_profile_cn.c +7798 -0
  109. algo_engine/exchange_profile/c_profile_cn.cp315-win_amd64.pyd +0 -0
  110. algo_engine/exchange_profile/c_profile_cn.pxd +1 -0
  111. algo_engine/exchange_profile/c_profile_cn.pyi +3 -0
  112. algo_engine/exchange_profile/c_profile_cn.pyx +1 -0
  113. algo_engine/exchange_profile/c_profile_default.c +7798 -0
  114. algo_engine/exchange_profile/c_profile_default.cp315-win_amd64.pyd +0 -0
  115. algo_engine/exchange_profile/c_profile_default.pxd +1 -0
  116. algo_engine/exchange_profile/c_profile_default.pyi +3 -0
  117. algo_engine/exchange_profile/c_profile_default.pyx +1 -0
  118. algo_engine/exchange_profile/c_profile_dispatcher.c +7798 -0
  119. algo_engine/exchange_profile/c_profile_dispatcher.cp315-win_amd64.pyd +0 -0
  120. algo_engine/exchange_profile/c_profile_dispatcher.pxd +1 -0
  121. algo_engine/exchange_profile/c_profile_dispatcher.pyi +3 -0
  122. algo_engine/exchange_profile/c_profile_dispatcher.pyx +1 -0
  123. algo_engine/includes/algo_engine/base/c_allocator_protocol.c +11608 -0
  124. algo_engine/includes/algo_engine/base/c_intern_string.c +5908 -0
  125. algo_engine/includes/algo_engine/base/c_market_data/c_candlestick.c +18996 -0
  126. algo_engine/includes/algo_engine/base/c_market_data/c_internal.c +14059 -0
  127. algo_engine/includes/algo_engine/base/c_market_data/c_market_data.c +30420 -0
  128. algo_engine/includes/algo_engine/base/c_market_data/c_market_data.h +1464 -0
  129. algo_engine/includes/algo_engine/base/c_market_data/c_market_data_buffer.c +29407 -0
  130. algo_engine/includes/algo_engine/base/c_market_data/c_market_data_buffer.h +941 -0
  131. algo_engine/includes/algo_engine/base/c_market_data/c_market_data_config.h +29 -0
  132. algo_engine/includes/algo_engine/base/c_market_data/c_tick.c +44303 -0
  133. algo_engine/includes/algo_engine/base/c_market_data/c_trade_utils.c +28702 -0
  134. algo_engine/includes/algo_engine/base/c_market_data/c_transaction.c +23558 -0
  135. algo_engine/includes/algo_engine/engine/c_event_engine.c +16325 -0
  136. algo_engine/includes/algo_engine/engine/c_market_engine.c +24258 -0
  137. algo_engine/includes/algo_engine/exchange_profile/c_ex_profile_base.c +87 -0
  138. algo_engine/includes/algo_engine/exchange_profile/c_ex_profile_base.h +1204 -0
  139. algo_engine/includes/algo_engine/exchange_profile/c_ex_profile_cn.c +968 -0
  140. algo_engine/includes/algo_engine/exchange_profile/c_ex_profile_cn.h +39 -0
  141. algo_engine/includes/algo_engine/exchange_profile/c_exchange_profile.c +52375 -0
  142. algo_engine/includes/algo_engine/exchange_profile/c_profile_cn.c +7798 -0
  143. algo_engine/includes/algo_engine/exchange_profile/c_profile_default.c +7798 -0
  144. algo_engine/includes/algo_engine/exchange_profile/c_profile_dispatcher.c +7798 -0
  145. algo_engine/monitor/__init__.py +15 -0
  146. algo_engine/monitor/advanced_data_interface.py +334 -0
  147. algo_engine/strategy/__init__.py +44 -0
  148. algo_engine/strategy/strategy_engine.py +441 -0
  149. algo_engine/utils/__init__.py +3 -0
  150. algo_engine/utils/commit_regularizer.py +49 -0
  151. algo_engine/utils/data_utils.py +296 -0
  152. pyalgoengine-0.12.3.dist-info/METADATA +142 -0
  153. pyalgoengine-0.12.3.dist-info/RECORD +156 -0
  154. pyalgoengine-0.12.3.dist-info/WHEEL +5 -0
  155. pyalgoengine-0.12.3.dist-info/licenses/LICENSE +21 -0
  156. pyalgoengine-0.12.3.dist-info/top_level.txt +2 -0
@@ -0,0 +1,635 @@
1
+ import abc
2
+ import datetime
3
+ import enum
4
+ import inspect
5
+ import logging
6
+ import operator
7
+ import warnings
8
+ from collections.abc import Sequence, Mapping, Iterable, Callable
9
+ from typing import Literal, Protocol, runtime_checkable, get_type_hints, Self
10
+
11
+ from . import LOGGER
12
+ from ..base import MarketData, DataType, MarketDataBuffer
13
+
14
+ LOGGER = LOGGER.getChild('Replay')
15
+ __all__ = ['PyDataScope', 'MarketDateCallable', 'MarketDataLoader', 'MarketDataBulkLoader', 'Replay', 'SimpleReplay', 'ProgressReplay', 'ProgressiveReplay']
16
+
17
+
18
+ class PyDataScope(enum.Flag):
19
+ SCOPE_TRANSACTION = enum.auto()
20
+ SCOPE_ORDER = enum.auto()
21
+ SCOPE_TICK = enum.auto()
22
+ SCOPE_TICK_LITE = enum.auto()
23
+
24
+ SCOPE_ALL = SCOPE_TRANSACTION | SCOPE_ORDER | SCOPE_TICK
25
+
26
+ @classmethod
27
+ def _missing_(cls, value: Literal['TickData', 'TickDataLite', 'OrderData', 'TransactionData']):
28
+ if isinstance(value, int):
29
+ return super()._missing_(value)
30
+
31
+ if isinstance(value, str):
32
+ dtypes = value.split(',')
33
+ elif isinstance(value, Iterable):
34
+ dtypes = value
35
+ else:
36
+ raise TypeError(value)
37
+
38
+ _ = PyDataScope(0)
39
+ for dtype in dtypes:
40
+ _ = _.from_str(dtype)
41
+ return _
42
+
43
+ @classmethod
44
+ def get_dtype(cls, dtype: DataType | str) -> str | Literal['TickData', 'TickDataLite', 'OrderData', 'TransactionData']:
45
+ match dtype:
46
+ case 'TickData' | 'TickDataLite' | 'OrderData' | 'TransactionData':
47
+ return str(dtype)
48
+ case 'TradeData': # handle the alias
49
+ return 'TransactionData'
50
+ case DataType.DTYPE_TICK | DataType.DTYPE_ORDER | DataType.DTYPE_TRANSACTION:
51
+ return DataType(dtype).name.removeprefix('DTYPE_').capitalize() + 'Data'
52
+ case DataType.DTYPE_TICK_LITE:
53
+ return 'Data'.join(_.capitalize() for _ in DataType(dtype).name.removeprefix('DTYPE_').split('_'))
54
+ case _:
55
+ raise ValueError(f'Invalid dtype {dtype}, expect str or int.')
56
+
57
+ def __iter__(self):
58
+ return iter(self.to_dtype())
59
+
60
+ def to_dtype(self) -> list[DataType]:
61
+ scope = list(super().__iter__())
62
+ scope_dtype = set()
63
+
64
+ for dtype in scope:
65
+
66
+ if dtype is PyDataScope.SCOPE_TRANSACTION:
67
+ scope_dtype.add(DataType.DTYPE_TRANSACTION)
68
+ elif dtype is PyDataScope.SCOPE_ORDER:
69
+ scope_dtype.add(DataType.DTYPE_ORDER)
70
+ elif dtype is PyDataScope.SCOPE_TICK_LITE:
71
+ scope_dtype.add(DataType.DTYPE_TICK_LITE)
72
+ elif dtype is PyDataScope.SCOPE_TICK:
73
+ scope_dtype.add(DataType.DTYPE_TICK)
74
+
75
+ return list(scope_dtype)
76
+
77
+ def to_int(self) -> list[int]:
78
+ return [int(_) for _ in self.to_dtype()]
79
+
80
+ def to_str(self) -> list[str]:
81
+ return [self.get_dtype(_) for _ in self.to_dtype()]
82
+
83
+ def from_str(self, dtype: Literal['TickData', 'TickDataLite', 'OrderData', 'TransactionData']) -> Self:
84
+ match dtype:
85
+ case 'TickData':
86
+ return self | self.SCOPE_TICK
87
+ case 'TickDataLite':
88
+ return self | self.SCOPE_TICK_LITE
89
+ case 'OrderData':
90
+ return self | self.SCOPE_ORDER
91
+ case 'TransactionData' | 'TradeData':
92
+ return self | self.SCOPE_TRANSACTION
93
+ case _:
94
+ raise ValueError(f'Invalid str {dtype}.')
95
+
96
+
97
+ @runtime_checkable
98
+ class MarketDateCallable(Protocol):
99
+ def __call__(self, market_date: datetime.date) -> None:
100
+ ...
101
+
102
+
103
+ @runtime_checkable
104
+ class MarketDataLoader(Protocol):
105
+ def __call__(self, market_date: datetime.date, ticker: str, dtype: str | DataType) -> Sequence[MarketData] | Mapping[float, MarketData]:
106
+ pass
107
+
108
+
109
+ @runtime_checkable
110
+ class MarketDataBulkLoader(Protocol):
111
+ def __call__(self, market_date: datetime.date, tickers: Sequence[str], dtypes: Sequence[str | DataType] | PyDataScope) -> Sequence[MarketData] | Mapping[float, MarketData] | MarketDataBuffer:
112
+ pass
113
+
114
+
115
+ def check_protocol_signature(func: Callable, protocol: type) -> bool:
116
+ if not callable(func):
117
+ raise TypeError(f"{func} is not callable")
118
+
119
+ proto_sig = inspect.signature(protocol.__call__)
120
+ func_sig = inspect.signature(func)
121
+
122
+ proto_params = list(proto_sig.parameters.values())[1:] # Skip 'self'
123
+ func_params = list(func_sig.parameters.values())
124
+ enable_keywords = False
125
+
126
+ # Check for *args (VAR_POSITIONAL) — not allowed
127
+ for p in func_params:
128
+ if p.kind == inspect.Parameter.VAR_POSITIONAL:
129
+ raise TypeError(f"{func.__name__} uses *args, which is not allowed")
130
+ elif p.kind == inspect.Parameter.VAR_KEYWORD:
131
+ enable_keywords = True
132
+
133
+ # Extract positional args (POSITIONAL_ONLY or POSITIONAL_OR_KEYWORD)
134
+ proto_arg_names = [p.name for p in proto_params if p.kind in (
135
+ inspect.Parameter.POSITIONAL_ONLY,
136
+ inspect.Parameter.POSITIONAL_OR_KEYWORD
137
+ )]
138
+
139
+ func_arg_names = [p.name for p in func_params if p.kind in (
140
+ inspect.Parameter.POSITIONAL_ONLY,
141
+ inspect.Parameter.POSITIONAL_OR_KEYWORD
142
+ )]
143
+
144
+ # Check if required positional args match (ignore **kwargs)
145
+ if not enable_keywords and sorted(proto_arg_names) != sorted(func_arg_names):
146
+ warnings.warn(
147
+ f"{func} argument names {func_arg_names} do not match protocol {proto_arg_names}",
148
+ stacklevel=2
149
+ )
150
+ return False
151
+
152
+ # Type hint comparison (warn if mismatched, but allow)
153
+ proto_hints = get_type_hints(protocol.__call__)
154
+ func_hints = get_type_hints(func)
155
+
156
+ for pname in proto_arg_names:
157
+ expected = proto_hints.get(pname)
158
+ actual = func_hints.get(pname)
159
+ if expected and actual and expected != actual:
160
+ warnings.warn(
161
+ f"Type hint mismatch for parameter '{pname}': expected {expected}, got {actual}",
162
+ stacklevel=2
163
+ )
164
+
165
+ # Optional: check return type
166
+ expected_ret = proto_hints.get("return")
167
+ actual_ret = func_hints.get("return")
168
+ if expected_ret and actual_ret and expected_ret != actual_ret:
169
+ warnings.warn(
170
+ f"Return type mismatch: expected {expected_ret}, got {actual_ret}",
171
+ stacklevel=2
172
+ )
173
+
174
+ return True
175
+
176
+
177
+ class Replay(object, metaclass=abc.ABCMeta):
178
+ # __slots__ = ('start_date', 'end_date', 'market_date', 'calendar', 'bod', 'eod', 'subscription', '_calendar', '_market_date', '_status', '_progress')
179
+
180
+ def __init__(self, start_date: datetime.date = None, end_date: datetime.date = None, market_date: datetime.date = None, calendar: Sequence[datetime.date] = None, bod: MarketDateCallable = None, eod: MarketDateCallable = None) -> None:
181
+ self.start_date = start_date or market_date or calendar[0]
182
+ self.end_date = end_date or calendar[-1]
183
+ self.market_date = market_date or start_date
184
+ self.calendar = calendar or []
185
+
186
+ self.bod = []
187
+ self.eod = []
188
+ self.subscription = {}
189
+
190
+ if bod is not None:
191
+ self.add_bod(bod)
192
+
193
+ if eod is not None:
194
+ self.add_eod(eod)
195
+
196
+ def add_bod(self, func: MarketDateCallable, priority: int = None) -> None:
197
+ if priority is None:
198
+ self.bod.append(func)
199
+ else:
200
+ self.bod.insert(priority, func)
201
+
202
+ def add_eod(self, func: MarketDateCallable, priority: int = None):
203
+ if priority is None:
204
+ self.eod.append(func)
205
+ else:
206
+ self.eod.insert(priority, func)
207
+
208
+ def add_subscription(self, ticker: str, dtype: DataType | str):
209
+ dtype = PyDataScope.get_dtype(dtype)
210
+ topic = f'{ticker}.{dtype}'
211
+
212
+ self.subscription[topic] = (ticker, dtype)
213
+
214
+ def remove_subscription(self, ticker: str, dtype: DataType | str):
215
+ dtype = PyDataScope.get_dtype(dtype)
216
+ topic = f'{ticker}.{dtype}'
217
+
218
+ try:
219
+ self.subscription.pop(topic)
220
+ except KeyError as _:
221
+ LOGGER.info(f'{topic} not in {self.subscription}')
222
+
223
+ @abc.abstractmethod
224
+ def __next__(self):
225
+ ...
226
+
227
+ @abc.abstractmethod
228
+ def __iter__(self):
229
+ ...
230
+
231
+
232
+ class SimpleReplay(Replay):
233
+ def __init__(
234
+ self,
235
+ loader: MarketDataBulkLoader | MarketDataLoader = None,
236
+ market_date: datetime.date = None,
237
+ start_date: datetime.date = None,
238
+ end_date: datetime.date = None,
239
+ calendar: Sequence[datetime.date] = None,
240
+ bod: MarketDateCallable = None,
241
+ eod: MarketDateCallable = None
242
+ ):
243
+ super().__init__(market_date=market_date, start_date=start_date, end_date=end_date, calendar=calendar, bod=bod, eod=eod)
244
+ self.loader = loader
245
+
246
+ def __iter__(self):
247
+ self._calendar = self.calendar or [self.start_date + datetime.timedelta(days=i) for i in range((self.end_date - self.start_date).days + 1)]
248
+ self._market_date = sorted(_ for _ in self._calendar if _ >= self.market_date)[0]
249
+ self._status = {market_date: 'skipped' if market_date < self.market_date else 'idle' for market_date in self._calendar}
250
+ self._idx_buffer = 0
251
+ self._idx_date = sum([1 for _ in self._calendar if _ < self.market_date])
252
+
253
+ for func in self.bod:
254
+ func(self._market_date)
255
+
256
+ self._safe_load()
257
+
258
+ return self
259
+
260
+ def __next__(self) -> MarketData:
261
+ if self._idx_buffer < self._buffer_size:
262
+ self._idx_buffer += 1
263
+ return next(self._buffer)
264
+
265
+ for func in self.eod:
266
+ func(self._market_date)
267
+
268
+ self._idx_buffer = 0
269
+ self._idx_date += 1
270
+
271
+ if self._idx_date >= len(self._calendar):
272
+ self._calendar.clear()
273
+ del self._calendar
274
+ del self._market_date
275
+ del self._status
276
+ del self._idx_buffer
277
+ del self._idx_date
278
+ del self._buffer
279
+ del self._buffer_size
280
+ raise StopIteration()
281
+
282
+ self._market_date = self._calendar[self._idx_date]
283
+
284
+ for func in self.bod:
285
+ func(self._market_date)
286
+
287
+ self._safe_load()
288
+ return self.__next__()
289
+
290
+ def __repr__(self):
291
+ return f'{self.__class__.__name__}{{id={id(self)}, from={self.start_date}, to={self.end_date}}}'
292
+
293
+ def _bulk_load_protocol(self):
294
+ LOGGER.info(f'{self} loading {self._market_date} {(', '.join(self.dtypes)) if self.dtypes else 'data'} for {len(self.tickers)} tickers...')
295
+ buffer = self.loader(market_date=self._market_date, tickers=self.tickers, dtypes=self.dtypes)
296
+ LOGGER.info(f'{self} sorting {self._market_date} data...')
297
+ buffer.sort()
298
+
299
+ if isinstance(buffer, MarketDataBuffer):
300
+ self._buffer = buffer
301
+ self._buffer_size = len(self._buffer)
302
+ elif isinstance(buffer, Sequence):
303
+ self._buffer = iter(buffer)
304
+ self._buffer_size = len(buffer)
305
+ elif isinstance(buffer, Mapping):
306
+ self._buffer = iter(buffer.values())
307
+ self._buffer_size = len(buffer)
308
+ LOGGER.info(f'{self} {self._market_date} total {self._buffer_size:,} items loaded.')
309
+
310
+ def _individual_load_protocol(self):
311
+ buffer = []
312
+ for topic, (_ticker, _dtype) in self.subscription.items():
313
+ LOGGER.info(f'{self} loading {self._market_date} {_ticker} {_dtype}...')
314
+ data = self.loader(market_date=self._market_date, ticker=_ticker, dtype=_dtype)
315
+ if isinstance(data, Mapping):
316
+ buffer.extend(list(data.values()))
317
+ elif isinstance(data, Sequence):
318
+ buffer.extend(data)
319
+ else:
320
+ raise TypeError(f'The loader {self.loader} returned {type(data)}. Expect a sequence or mapping of MarketData')
321
+ LOGGER.info(f'{self} sorting {self._market_date} data...')
322
+ buffer.sort(key=operator.attrgetter('timestamp', 'ticker', '_dtype'))
323
+ self._buffer = iter(buffer)
324
+ self._buffer_size = len(buffer)
325
+ LOGGER.info(f'{self} {self._market_date} total {self._buffer_size:,} items loaded.')
326
+
327
+ def _safe_load(self):
328
+ if self.loader is None:
329
+ assert hasattr(self, '_buffer') and isinstance(self._buffer, Iterable), f'Without assigning a data loader, the _buffer of {self.__class__.__name__} should be set in bod process.'
330
+ return None
331
+
332
+ is_bulk_loader = check_protocol_signature(self.loader, MarketDataBulkLoader)
333
+ is_individual_loader = check_protocol_signature(self.loader, MarketDataLoader)
334
+
335
+ if (is_bulk_loader and is_individual_loader) or (not is_bulk_loader and not is_individual_loader):
336
+ try:
337
+ return self._bulk_load_protocol()
338
+ except Exception as e:
339
+ LOGGER.info('Failed to load data using MarketDataBulkLoader protocol!')
340
+
341
+ try:
342
+ return self._individual_load_protocol()
343
+ except Exception as e:
344
+ LOGGER.info('Failed to load data using MarketDataLoader protocol!')
345
+ raise
346
+
347
+ if is_bulk_loader:
348
+ return self._bulk_load_protocol()
349
+
350
+ return self._individual_load_protocol()
351
+
352
+ @property
353
+ def progress(self) -> float:
354
+ if not hasattr(self, '_buffer'):
355
+ raise RuntimeError(f'{self.__class__.__name__} not started yet.')
356
+
357
+ return (self._idx_date + self._idx_buffer / self._buffer_size) / len(self._calendar)
358
+
359
+ @property
360
+ def tickers(self) -> list[str]:
361
+ tickers = set()
362
+ for _, (ticker, dtype) in self.subscription.items():
363
+ tickers.add(ticker)
364
+ return list(tickers)
365
+
366
+ @property
367
+ def dtypes(self) -> list[str]:
368
+ dtypes = set()
369
+ for _, (ticker, dtype) in self.subscription.items():
370
+ dtypes.add(dtype)
371
+ return list(dtypes)
372
+
373
+ @property
374
+ def status(self) -> dict[datetime.date, str]:
375
+ if not hasattr(self, '_status'):
376
+ raise RuntimeError(f'{self.__class__.__name__} not started yet.')
377
+
378
+ return self._status
379
+
380
+
381
+ class ProgressReplay(SimpleReplay):
382
+ def __init__(
383
+ self,
384
+ loader: MarketDataBulkLoader | MarketDataLoader = None,
385
+ market_date: datetime.date = None,
386
+ start_date: datetime.date = None,
387
+ end_date: datetime.date = None,
388
+ calendar: Sequence[datetime.date] = None,
389
+ bod: MarketDateCallable = None,
390
+ eod: MarketDateCallable = None,
391
+ **pbar_config
392
+ ):
393
+ super().__init__(
394
+ loader=loader,
395
+ market_date=market_date,
396
+ start_date=start_date,
397
+ end_date=end_date,
398
+ calendar=calendar,
399
+ bod=bod,
400
+ eod=eod
401
+ )
402
+
403
+ self.pbar_config = {
404
+ 'backend': pbar_config.pop('backend', 'tqdm'), # tqdm or native
405
+ 'config': pbar_config,
406
+ }
407
+ self._pbar = None
408
+
409
+ def _init_pbar_tqdm(self):
410
+ from tqdm.auto import tqdm
411
+ from tqdm.std import tqdm as tqdm_std
412
+ from tqdm.contrib.logging import _TqdmLoggingHandler, _get_first_found_console_logging_handler, _is_console_logging_handler
413
+
414
+ tqdm_config = {
415
+ 'total': 1,
416
+ 'unit_scale': True,
417
+ 'unit': 'percent',
418
+ 'mininterval': 0.1,
419
+ 'miniters': 0.001,
420
+ **self.pbar_config['config'],
421
+ }
422
+ self._pbar = tqdm(**tqdm_config)
423
+
424
+ self.pbar_config['loggers'] = loggers = [LOGGER.root] + [_ for _ in LOGGER.root.manager.loggerDict.values() if isinstance(_, logging.Logger) and _.handlers]
425
+ self.pbar_config['original_handlers_list'] = [logger.handlers for logger in loggers]
426
+ for logger in loggers:
427
+ tqdm_handler = _TqdmLoggingHandler(tqdm_std)
428
+ orig_handler = _get_first_found_console_logging_handler(logger.handlers)
429
+ if orig_handler is not None:
430
+ tqdm_handler.setFormatter(orig_handler.formatter)
431
+ tqdm_handler.stream = orig_handler.stream
432
+ logger.handlers = [handler for handler in logger.handlers if not _is_console_logging_handler(handler)] + [tqdm_handler]
433
+
434
+ self.add_bod(self._init_pbar_tqdm_secondary, priority=0)
435
+ self.add_eod(self._close_pbar_tqdm_secondary, priority=0)
436
+ self.add_bod(self._update_tqdm_prefix, priority=0)
437
+ self._update_pbar_progress = self._update_tqdm_progress
438
+
439
+ def _init_pbar_tqdm_secondary(self, market_date):
440
+ from tqdm.auto import tqdm
441
+
442
+ tqdm_secondary_config = {
443
+ 'total': 1,
444
+ 'unit_scale': True,
445
+ 'unit': 'percent',
446
+ 'mininterval': 0.1,
447
+ 'miniters': 0.001,
448
+ **self.pbar_config['config'],
449
+ }
450
+ self._pbar_secondary = tqdm(**tqdm_secondary_config)
451
+ prompt = f'Progress Total ({self._idx_date + 1} / {len(self._calendar)})'
452
+ prompt_secondary = f'Progress [{market_date:%Y-%m-%d}]'
453
+ prompt_length = max(len(prompt), len(prompt_secondary))
454
+ self._pbar_secondary.n = 0
455
+ self._pbar_secondary.set_description(prompt_secondary.ljust(prompt_length))
456
+ self._pbar_secondary.refresh()
457
+
458
+ def _close_pbar_tqdm_secondary(self, market_date: datetime.date):
459
+ self._pbar_secondary.n = 1
460
+ # self._pbar_secondary.refresh()
461
+ self._pbar_secondary.close()
462
+ self._pbar_secondary = None
463
+
464
+ def _init_pbar_native(self):
465
+ from ..base import Progress
466
+
467
+ progress_config = dict(
468
+ tasks=1,
469
+ tick_size=0.001,
470
+ **self.pbar_config['config'],
471
+ )
472
+
473
+ self.add_bod(self._update_native_prefix, priority=0)
474
+ self._pbar = Progress(**progress_config)
475
+ self._update_pbar_progress = self._update_native_progress
476
+
477
+ def _update_tqdm_prefix(self, market_date: datetime.date):
478
+ prompt = f'Progress Total ({self._idx_date + 1} / {len(self._calendar)})'
479
+ self._pbar.set_description(prompt)
480
+ self._pbar.refresh()
481
+
482
+ def _update_native_prefix(self, market_date: datetime.date):
483
+ self._pbar.prompt = f'Replay {market_date:%Y-%m-%d} ({self._idx_date + 1} / {len(self._calendar)}):'
484
+ self._pbar.output()
485
+
486
+ def _close_pbar_tqdm(self):
487
+ for logger, original_handlers in zip(self.pbar_config['loggers'], self.pbar_config['original_handlers_list']):
488
+ logger.handlers = original_handlers
489
+
490
+ self._pbar.n = 1
491
+ # self._pbar.refresh()
492
+ self._pbar.close()
493
+ self._pbar = None
494
+
495
+ def _close_pbar_native(self):
496
+ self._pbar.done_tasks = 1
497
+ self._pbar.output()
498
+
499
+ def _update_tqdm_progress(self):
500
+ self._pbar.n = self.progress
501
+ self._pbar.update(0)
502
+
503
+ self._pbar_secondary.n = self._idx_buffer / self._buffer_size
504
+ self._pbar_secondary.update(0)
505
+
506
+ def _update_native_progress(self):
507
+ self._pbar.done_tasks = self.progress
508
+
509
+ if (not self._pbar.tick_size) \
510
+ or self._pbar.progress >= self._pbar.tick_size + self._pbar.last_output \
511
+ or self._pbar.is_done:
512
+ self._pbar.output()
513
+
514
+ def __iter__(self):
515
+ pbar_backend = self.pbar_config['backend']
516
+
517
+ match pbar_backend:
518
+ case 'tqdm':
519
+ self._init_pbar_tqdm()
520
+ case 'native':
521
+ self._init_pbar_native()
522
+ case _:
523
+ raise NotImplementedError(f'Invalid pbar backend {pbar_backend}')
524
+
525
+ return super().__iter__()
526
+
527
+ def __next__(self) -> MarketData:
528
+ try:
529
+ result = super().__next__()
530
+ if self._pbar is not None:
531
+ self._update_pbar_progress()
532
+ return result
533
+ except StopIteration:
534
+ if self._pbar is not None:
535
+ pbar_backend = self.pbar_config['backend']
536
+ match pbar_backend:
537
+ case 'tqdm':
538
+ self._close_pbar_tqdm()
539
+ case 'native':
540
+ self._close_pbar_native()
541
+ case _:
542
+ raise NotImplementedError(f'Invalid pbar backend {pbar_backend}')
543
+ raise
544
+
545
+
546
+ class ProgressiveReplay(SimpleReplay):
547
+ """
548
+ progressively loading and replaying market data
549
+
550
+ requires arguments
551
+ loader: a data loading function. Expect loader = Callable(market_date: datetime.date, ticker: str, dtype: str| type) -> dict[any, MarketData]
552
+ start_date & end_date: the given replay period
553
+ or calendar: the given replay calendar.
554
+
555
+ accepts kwargs:
556
+ ticker / tickers: the given symbols to replay, expect a str| list[str]
557
+ dtype / dtypes: the given dtype(s) of symbol to replay, expect a str | type, list[str | type]. default = all, which is (TradeData, TickData, OrderBook)
558
+ subscription / subscribe: the given ticker-dtype pair to replay, expect a list[dict[str, str | type]]
559
+ """
560
+
561
+ def __init__(
562
+ self,
563
+ loader: MarketDataLoader,
564
+ tickers: str | Sequence[str] = None,
565
+ dtypes: str | DataType | Sequence[str] | Sequence[DataType] = None,
566
+ market_date: datetime.date = None,
567
+ start_date: datetime.date = None,
568
+ end_date: datetime.date = None,
569
+ calendar: Sequence[datetime.date] = None,
570
+ bod: MarketDateCallable = None,
571
+ eod: MarketDateCallable = None,
572
+ **progress_config
573
+ ) -> None:
574
+ warnings.warn('User ProgressReplay instead!', DeprecationWarning, stacklevel=2)
575
+ self.loader = loader
576
+ super().__init__(loader=loader, market_date=market_date, start_date=start_date, end_date=end_date, calendar=calendar, bod=bod, eod=eod)
577
+
578
+ tickers = tickers or []
579
+ dtypes = dtypes or ['TransactionData', 'TickData', 'OrderData']
580
+
581
+ if not isinstance(loader, MarketDataLoader):
582
+ raise TypeError('loader function has 3 requires args, market_date, ticker and dtype.')
583
+
584
+ if isinstance(tickers, str):
585
+ tickers = [tickers]
586
+ elif isinstance(tickers, Iterable):
587
+ tickers = list(tickers)
588
+ else:
589
+ raise TypeError(f'Invalid ticker {tickers}, expect str or list[str]')
590
+
591
+ if isinstance(dtypes, (str, int, DataType)):
592
+ dtypes = [dtypes]
593
+ elif isinstance(dtypes, Iterable):
594
+ dtypes = list(dtypes)
595
+ else:
596
+ raise TypeError(f'Invalid dtype {dtypes}, expect str or list[str]')
597
+
598
+ for ticker in tickers:
599
+ for dtype in dtypes:
600
+ self.add_subscription(ticker=ticker, dtype=dtype)
601
+
602
+ self.progress_config = dict(
603
+ tasks=1,
604
+ **progress_config
605
+ )
606
+ self._pbar = None
607
+ self.add_bod(self._update_progress_bar, priority=0)
608
+
609
+ def __iter__(self):
610
+ from ..base import Progress
611
+ self._pbar = Progress(**self.progress_config)
612
+ return super().__iter__()
613
+
614
+ def __next__(self) -> MarketData:
615
+ try:
616
+ result = super().__next__()
617
+ if self._pbar:
618
+ self._pbar.done_tasks = self.progress
619
+
620
+ if (not self._pbar.tick_size) \
621
+ or self._pbar.progress >= self._pbar.tick_size + self._pbar.last_output \
622
+ or self._pbar.is_done:
623
+ self._pbar.output()
624
+
625
+ return result
626
+ except StopIteration:
627
+ if self._pbar is not None and not self._pbar.is_done:
628
+ self.progress.done_tasks = 1
629
+ self._pbar.output()
630
+ raise
631
+
632
+ def _update_progress_bar(self, market_date: datetime.date):
633
+ if self._pbar:
634
+ self.progress.prompt = f'Replay {market_date:%Y-%m-%d} ({self._idx_date + 1} / {len(self._calendar)}):'
635
+ self._pbar.output()