Coverage for src / monte_neo / data / storage.py: 93%

82 statements  

« prev     ^ index     » next       coverage.py v7.13.1, created at 2026-01-28 16:27 +0200

1"""Parquet storage module. 

2 

3Efficient data storage and retrieval using Parquet format. 

4""" 

5 

6from __future__ import annotations 

7 

8import os 

9from pathlib import Path 

10from typing import TYPE_CHECKING 

11 

12import pandas as pd 

13import pyarrow as pa 

14import pyarrow.parquet as pq 

15 

16from monte_neo.utils.logger import get_logger 

17 

18if TYPE_CHECKING: 

19 pass 

20 

21logger = get_logger(__name__) 

22 

23 

24class ParquetStorage: 

25 """Efficient Parquet-based data storage.""" 

26 

27 def __init__(self, base_dir: str | Path | None = None) -> None: 

28 """Initialize storage. 

29 

30 Args: 

31 base_dir: Base directory for data storage. 

32 """ 

33 if base_dir is None: 

34 base_dir = os.getenv("MONTE_NEO_DATA_DIR", "./data") 

35 self.base_dir = Path(base_dir) 

36 self._ensure_directories() 

37 

38 def _ensure_directories(self) -> None: 

39 """Create necessary directories.""" 

40 for subdir in ["raw", "processed", "results"]: 

41 (self.base_dir / subdir).mkdir(parents=True, exist_ok=True) 

42 

43 def _get_path( 

44 self, 

45 symbol: str, 

46 timeframe: str, 

47 category: str = "raw", 

48 ) -> Path: 

49 """Generate file path for data. 

50 

51 Args: 

52 symbol: Trading pair symbol. 

53 timeframe: Candle interval. 

54 category: Data category (raw, processed, results). 

55 

56 Returns: 

57 Path to the Parquet file. 

58 """ 

59 filename = f"{symbol}_{timeframe}.parquet" 

60 return self.base_dir / category / filename 

61 

62 def save( 

63 self, 

64 df: pd.DataFrame, 

65 symbol: str, 

66 timeframe: str, 

67 category: str = "raw", 

68 compression: str = "snappy", 

69 ) -> Path: 

70 """Save DataFrame to Parquet. 

71 

72 Args: 

73 df: DataFrame to save. 

74 symbol: Trading pair symbol. 

75 timeframe: Candle interval. 

76 category: Data category. 

77 compression: Compression algorithm. 

78 

79 Returns: 

80 Path to saved file. 

81 """ 

82 path = self._get_path(symbol, timeframe, category) 

83 

84 # Convert to PyArrow table for better control 

85 table = pa.Table.from_pandas(df) 

86 

87 pq.write_table( 

88 table, 

89 path, 

90 compression=compression, 

91 use_dictionary=True, 

92 write_statistics=True, 

93 ) 

94 

95 logger.info(f"Saved {len(df)} rows to {path}") 

96 return path 

97 

98 def load( 

99 self, 

100 symbol: str, 

101 timeframe: str, 

102 category: str = "raw", 

103 columns: list[str] | None = None, 

104 ) -> pd.DataFrame: 

105 """Load DataFrame from Parquet. 

106 

107 Args: 

108 symbol: Trading pair symbol. 

109 timeframe: Candle interval. 

110 category: Data category. 

111 columns: Optional list of columns to load. 

112 

113 Returns: 

114 Loaded DataFrame. 

115 """ 

116 path = self._get_path(symbol, timeframe, category) 

117 

118 if not path.exists(): 

119 raise FileNotFoundError(f"Data file not found: {path}") 

120 

121 df = pq.read_table(path, columns=columns).to_pandas() 

122 logger.info(f"Loaded {len(df)} rows from {path}") 

123 return df 

124 

125 def exists( 

126 self, 

127 symbol: str, 

128 timeframe: str, 

129 category: str = "raw", 

130 ) -> bool: 

131 """Check if data file exists. 

132 

133 Args: 

134 symbol: Trading pair symbol. 

135 timeframe: Candle interval. 

136 category: Data category. 

137 

138 Returns: 

139 True if file exists. 

140 """ 

141 return self._get_path(symbol, timeframe, category).exists() 

142 

143 def list_files(self, category: str = "raw") -> list[dict]: 

144 """List all data files in category. 

145 

146 Args: 

147 category: Data category. 

148 

149 Returns: 

150 List of file info dictionaries. 

151 """ 

152 category_dir = self.base_dir / category 

153 files = [] 

154 

155 for path in category_dir.glob("*.parquet"): 

156 parts = path.stem.split("_") 

157 if len(parts) >= 2: 

158 files.append( 

159 { 

160 "symbol": parts[0], 

161 "timeframe": parts[1], 

162 "path": path, 

163 "size_mb": path.stat().st_size / (1024 * 1024), 

164 } 

165 ) 

166 

167 return sorted(files, key=lambda x: x["symbol"]) 

168 

169 def delete( 

170 self, 

171 symbol: str, 

172 timeframe: str, 

173 category: str = "raw", 

174 ) -> bool: 

175 """Delete a data file. 

176 

177 Args: 

178 symbol: Trading pair symbol. 

179 timeframe: Candle interval. 

180 category: Data category. 

181 

182 Returns: 

183 True if file was deleted. 

184 """ 

185 path = self._get_path(symbol, timeframe, category) 

186 if path.exists(): 

187 path.unlink() 

188 logger.info(f"Deleted {path}") 

189 return True 

190 return False 

191 

192 def get_info( 

193 self, 

194 symbol: str, 

195 timeframe: str, 

196 category: str = "raw", 

197 ) -> dict | None: 

198 """Get metadata about a data file. 

199 

200 Args: 

201 symbol: Trading pair symbol. 

202 timeframe: Candle interval. 

203 category: Data category. 

204 

205 Returns: 

206 Dictionary with file metadata or None. 

207 """ 

208 path = self._get_path(symbol, timeframe, category) 

209 

210 if not path.exists(): 

211 return None 

212 

213 parquet_file = pq.ParquetFile(path) 

214 metadata = parquet_file.metadata 

215 schema = parquet_file.schema_arrow 

216 

217 # Try to get date range from statistics 

218 start_date = None 

219 end_date = None 

220 

221 try: 

222 # Assuming timestamp is the index or first column 

223 # We check all row groups to find global min/max 

224 # This handles unsorted data too, though time data is usually sorted 

225 min_vals = [] 

226 max_vals = [] 

227 

228 # Find timestamp column index 

229 ts_col_idx = -1 

230 for i, name in enumerate(schema.names): 

231 if name == "timestamp" or name == "__index_level_0__": 

232 ts_col_idx = i 

233 break 

234 

235 if ts_col_idx >= 0: 

236 for rg in range(metadata.num_row_groups): 

237 col_meta = metadata.row_group(rg).column(ts_col_idx) 

238 if col_meta.is_stats_set: 

239 stats = col_meta.statistics 

240 if stats.has_min_max: 

241 min_vals.append(stats.min) 

242 max_vals.append(stats.max) 

243 

244 if min_vals and max_vals: 

245 start_date = min(min_vals) 

246 end_date = max(max_vals) 

247 except Exception as e: 

248 logger.debug(f"Could not extract date range from metadata: {e}") 

249 

250 return { 

251 "path": path, 

252 "rows": metadata.num_rows, 

253 "columns": [field.name for field in schema], 

254 "size_mb": path.stat().st_size / (1024 * 1024), 

255 "compression": metadata.row_group(0).column(0).compression, 

256 "start_date": start_date, 

257 "end_date": end_date, 

258 }