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
« prev ^ index » next coverage.py v7.13.1, created at 2026-01-28 16:27 +0200
1"""Parquet storage module.
3Efficient data storage and retrieval using Parquet format.
4"""
6from __future__ import annotations
8import os
9from pathlib import Path
10from typing import TYPE_CHECKING
12import pandas as pd
13import pyarrow as pa
14import pyarrow.parquet as pq
16from monte_neo.utils.logger import get_logger
18if TYPE_CHECKING:
19 pass
21logger = get_logger(__name__)
24class ParquetStorage:
25 """Efficient Parquet-based data storage."""
27 def __init__(self, base_dir: str | Path | None = None) -> None:
28 """Initialize storage.
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()
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)
43 def _get_path(
44 self,
45 symbol: str,
46 timeframe: str,
47 category: str = "raw",
48 ) -> Path:
49 """Generate file path for data.
51 Args:
52 symbol: Trading pair symbol.
53 timeframe: Candle interval.
54 category: Data category (raw, processed, results).
56 Returns:
57 Path to the Parquet file.
58 """
59 filename = f"{symbol}_{timeframe}.parquet"
60 return self.base_dir / category / filename
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.
72 Args:
73 df: DataFrame to save.
74 symbol: Trading pair symbol.
75 timeframe: Candle interval.
76 category: Data category.
77 compression: Compression algorithm.
79 Returns:
80 Path to saved file.
81 """
82 path = self._get_path(symbol, timeframe, category)
84 # Convert to PyArrow table for better control
85 table = pa.Table.from_pandas(df)
87 pq.write_table(
88 table,
89 path,
90 compression=compression,
91 use_dictionary=True,
92 write_statistics=True,
93 )
95 logger.info(f"Saved {len(df)} rows to {path}")
96 return path
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.
107 Args:
108 symbol: Trading pair symbol.
109 timeframe: Candle interval.
110 category: Data category.
111 columns: Optional list of columns to load.
113 Returns:
114 Loaded DataFrame.
115 """
116 path = self._get_path(symbol, timeframe, category)
118 if not path.exists():
119 raise FileNotFoundError(f"Data file not found: {path}")
121 df = pq.read_table(path, columns=columns).to_pandas()
122 logger.info(f"Loaded {len(df)} rows from {path}")
123 return df
125 def exists(
126 self,
127 symbol: str,
128 timeframe: str,
129 category: str = "raw",
130 ) -> bool:
131 """Check if data file exists.
133 Args:
134 symbol: Trading pair symbol.
135 timeframe: Candle interval.
136 category: Data category.
138 Returns:
139 True if file exists.
140 """
141 return self._get_path(symbol, timeframe, category).exists()
143 def list_files(self, category: str = "raw") -> list[dict]:
144 """List all data files in category.
146 Args:
147 category: Data category.
149 Returns:
150 List of file info dictionaries.
151 """
152 category_dir = self.base_dir / category
153 files = []
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 )
167 return sorted(files, key=lambda x: x["symbol"])
169 def delete(
170 self,
171 symbol: str,
172 timeframe: str,
173 category: str = "raw",
174 ) -> bool:
175 """Delete a data file.
177 Args:
178 symbol: Trading pair symbol.
179 timeframe: Candle interval.
180 category: Data category.
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
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.
200 Args:
201 symbol: Trading pair symbol.
202 timeframe: Candle interval.
203 category: Data category.
205 Returns:
206 Dictionary with file metadata or None.
207 """
208 path = self._get_path(symbol, timeframe, category)
210 if not path.exists():
211 return None
213 parquet_file = pq.ParquetFile(path)
214 metadata = parquet_file.metadata
215 schema = parquet_file.schema_arrow
217 # Try to get date range from statistics
218 start_date = None
219 end_date = None
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 = []
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
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)
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}")
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 }