From b2bb2ad0a0cd87aaa6b8fa4e0a204f28a2f2f963 Mon Sep 17 00:00:00 2001 From: Daichi Narushima <1938249+dceoy@users.noreply.github.com> Date: Tue, 9 Jun 2026 03:29:03 +0900 Subject: [PATCH] Add rate view resolution and downstream SDK helpers (#18) * Add public helpers to resolve rate compatibility view names. Expose resolve_rate_view_name and resolve_rate_view_names in mt5cli.history so consumers can derive mt5cli-managed SQLite view names from stored rates metadata without reimplementing the naming rules. Co-authored-by: Cursor * Add reusable export, tick-window, and margin helpers for downstream tools. Expose SQLite append/dedup export, recent tick retrieval, and minimum margin summary through the SDK and CLI so projects like mteor can depend on mt5cli instead of duplicating MT5 data plumbing. Co-authored-by: Cursor * Bump version to 0.4.3. Co-authored-by: Cursor * Address PR review feedback for rate view resolution and SDK helpers. Harden SQLite read-only connections, tighten view discovery, improve recent_ticks fetch efficiency, default SQLite export to append, and expand tests and docs. Co-authored-by: Cursor * Fix read-only SQLite URI construction on Windows. Use Path.as_uri() so encoded file URIs work cross-platform with mode=ro. Co-authored-by: Cursor --------- Co-authored-by: Cursor --- README.md | 5 + docs/api/history.md | 35 ++++++ docs/api/index.md | 20 +++ docs/index.md | 30 ++++- mt5cli/__init__.py | 13 +- mt5cli/cli.py | 48 ++++++++ mt5cli/history.py | 225 ++++++++++++++++++++++++++++++++++ mt5cli/sdk.py | 173 +++++++++++++++++++++++++- mt5cli/utils.py | 65 ++++++++-- pyproject.toml | 2 +- tests/test_cli.py | 61 +++++++++- tests/test_history.py | 277 ++++++++++++++++++++++++++++++++++++++++++ tests/test_sdk.py | 174 +++++++++++++++++++++++++- tests/test_utils.py | 108 ++++++++++++++++ uv.lock | 2 +- 15 files changed, 1215 insertions(+), 23 deletions(-) diff --git a/README.md b/README.md index 25e4fea..35d3863 100644 --- a/README.md +++ b/README.md @@ -57,6 +57,7 @@ python -m mt5cli -o account.csv account-info | `rates-range` | Export rates for a date range | | `ticks-from` | Export ticks from a start date | | `ticks-range` | Export ticks for a date range | +| `ticks-recent` | Export ticks from a recent trailing window | | `account-info` | Export account information | | `terminal-info` | Export terminal information | | `version` | Export MetaTrader 5 version information | @@ -64,6 +65,7 @@ python -m mt5cli -o account.csv account-info | `symbols` | Export symbol list | | `symbol-info` | Export symbol details | | `symbol-info-tick` | Export the last tick for a symbol | +| `minimum-margins` | Export minimum-volume buy and sell margin requirements | | `market-book` | Export market depth (order book) | | `orders` | Export active orders | | `positions` | Export open positions | @@ -127,6 +129,9 @@ update_history_with_config( - **`update_history`**: incremental append based on existing SQLite `MAX(time)` per symbol (and timeframe for rates); account-level deals use a separate cursor when `include_account_events=True`. - **`rates` table**: normalized storage with `symbol` and `timeframe` columns. - **Rate compatibility views**: mt5cli manages all `rate_*` views. Naming is `rate___` when a symbol has one timeframe, otherwise `rate____` (for example `rate_EURUSD__M1_1`). Stale `rate_*` views are dropped and recreated when rates change for offline tools such as mteor optimize. +- **Rate view resolution**: use `mt5cli.history.resolve_rate_view_name()` / `resolve_rate_view_names()` to map symbols and granularities to existing SQLite compatibility views without creating databases. +- **SQLite export helpers**: use `export_dataframe_to_sqlite()` for append mode, optional index export, and post-write deduplication by key columns. +- **Recent ticks and margins**: `recent_ticks()` and `minimum_margins()` SDK helpers (and matching CLI commands) cover common downstream read-only queries. ## Requirements diff --git a/docs/api/history.md b/docs/api/history.md index f47a96a..97524e6 100644 --- a/docs/api/history.md +++ b/docs/api/history.md @@ -129,3 +129,38 @@ when required columns are missing. The `update_history` SDK path uses the same base tables and optional `cash_events` / `positions_reconstructed` views. It additionally maintains `rate___` compatibility views when `create_rate_views=True`. + +### Rate view resolution + +Downstream tools can resolve mt5cli-managed compatibility view names from an +existing SQLite history database without creating files or guessing legacy +naming schemes: + +```python +from pathlib import Path + +from mt5cli.history import resolve_rate_view_name, resolve_rate_view_names + +# Single symbol and granularity +view = resolve_rate_view_name(Path("history.db"), "EURUSD", "M1") + +# Batch resolution in row-major order +views = resolve_rate_view_names( + Path("history.db"), + ["EURUSD", "GBPUSD"], + ["M1", "H1"], +) +``` + +Resolution rules: + +- Returns `rate___` when a symbol stores one timeframe. +- Returns `rate____` when multiple timeframes + are stored for the same symbol. +- When multiple naming candidates apply, prefers an existing managed + `rate_*__*` view from the candidate list. +- Falls back to single-timeframe naming when the database path is missing or + `rates` metadata is unavailable. +- Pass `require_existing=True` to raise `ValueError` instead of returning a + best-guess name when the database or view is missing. +- Accepts either a SQLite path or an open `sqlite3.Connection`. diff --git a/docs/api/index.md b/docs/api/index.md index 86b17df..3b26de0 100644 --- a/docs/api/index.md +++ b/docs/api/index.md @@ -65,12 +65,18 @@ from datetime import UTC, datetime from pathlib import Path from mt5cli import ( + Dataset, + IfExists, Mt5CliClient, collect_history, copy_rates_range, detect_format, export_dataframe, + export_dataframe_to_sqlite, + minimum_margins, + recent_ticks, ) +from mt5cli.history import resolve_rate_view_name # Fetch rates programmatically rates = copy_rates_range( @@ -86,6 +92,20 @@ fmt = detect_format(Path("output.parquet")) # Returns "parquet" # Export a DataFrame export_dataframe(rates, Path("output.csv"), "csv") +# Append to SQLite with deduplication +export_dataframe_to_sqlite( + rates, + Path("history.db"), + "rates", + if_exists=IfExists.APPEND, + deduplicate_on=("symbol", "timeframe", "time"), +) + +# Resolve rate compatibility views and fetch recent ticks +view = resolve_rate_view_name(Path("history.db"), "EURUSD", "M1") +ticks = recent_ticks("EURUSD", seconds=300) +margins = minimum_margins("EURUSD") + # Collect history into SQLite collect_history( Path("history.db"), diff --git a/docs/index.md b/docs/index.md index 2c4548f..f6bb875 100644 --- a/docs/index.md +++ b/docs/index.md @@ -22,13 +22,22 @@ pip install mt5cli ## Programmatic usage / SDK usage -mt5cli can be used as a small Python SDK for read-only MetaTrader 5 data collection. SDK functions return pandas DataFrames without writing files. Use `export_dataframe` when you need to persist results. +mt5cli can be used as a small Python SDK for read-only MetaTrader 5 data collection. SDK functions return pandas DataFrames without writing files. Use `export_dataframe` or `export_dataframe_to_sqlite` when you need to persist results. ```python from datetime import UTC, datetime from pathlib import Path -from mt5cli import Mt5CliClient, collect_history, copy_rates_range, export_dataframe +from mt5cli import ( + Mt5CliClient, + collect_history, + copy_rates_range, + export_dataframe, + export_dataframe_to_sqlite, + minimum_margins, + recent_ticks, +) +from mt5cli.history import resolve_rate_view_name # One-off fetch with module-level helpers rates = copy_rates_range( @@ -39,6 +48,13 @@ rates = copy_rates_range( ) export_dataframe(rates, Path("rates.csv"), "csv") +# Resolve SQLite rate compatibility views for downstream tools +view = resolve_rate_view_name(Path("history.db"), "EURUSD", "M1") + +# Recent tick window and minimum margin summary +ticks = recent_ticks("EURUSD", seconds=300) +margins = minimum_margins("EURUSD") + # Reuse one MT5 connection for multiple calls with Mt5CliClient(login=12345, password="secret", server="Broker-Demo") as client: account = client.account_info() @@ -92,10 +108,11 @@ mt5cli --login 12345 --password mypass --server MyBroker-Demo \ ### Ticks -| Command | Description | -| ------------- | ------------------------------ | -| `ticks-from` | Export ticks from a start date | -| `ticks-range` | Export ticks for a date range | +| Command | Description | +| -------------- | ----------------------------------- | +| `ticks-from` | Export ticks from a start date | +| `ticks-range` | Export ticks for a date range | +| `ticks-recent` | Export ticks from a trailing window | ### Information @@ -108,6 +125,7 @@ mt5cli --login 12345 --password mypass --server MyBroker-Demo \ | `symbols` | Export symbol list | | `symbol-info` | Export symbol details | | `symbol-info-tick` | Export the last tick for a symbol | +| `minimum-margins` | Export minimum-volume margin summary | | `market-book` | Export market depth (order book) | ### Trading diff --git a/mt5cli/__init__.py b/mt5cli/__init__.py index f01ac06..7a30458 100644 --- a/mt5cli/__init__.py +++ b/mt5cli/__init__.py @@ -16,8 +16,10 @@ from .sdk import ( history_orders, last_error, market_book, + minimum_margins, orders, positions, + recent_ticks, symbol_info, symbol_info_tick, symbols, @@ -28,7 +30,13 @@ from .sdk import ( from .sdk import ( version as mt5_version, ) -from .utils import Dataset, IfExists, detect_format, export_dataframe +from .utils import ( + Dataset, + IfExists, + detect_format, + export_dataframe, + export_dataframe_to_sqlite, +) __version__ = version(__package__) if __package__ else None @@ -46,13 +54,16 @@ __all__ = [ "copy_ticks_range", "detect_format", "export_dataframe", + "export_dataframe_to_sqlite", "history_deals", "history_orders", "last_error", "market_book", + "minimum_margins", "mt5_version", "orders", "positions", + "recent_ticks", "symbol_info", "symbol_info_tick", "symbols", diff --git a/mt5cli/cli.py b/mt5cli/cli.py index 0189a36..4a61c6c 100644 --- a/mt5cli/cli.py +++ b/mt5cli/cli.py @@ -300,6 +300,44 @@ def ticks_range( ) +@app.command() +def ticks_recent( + ctx: typer.Context, + symbol: Annotated[str, typer.Option(help="Symbol name.")], + seconds: Annotated[ + float, + typer.Option(help="Lookback window in seconds."), + ], + date_to: Annotated[ + datetime | None, + typer.Option(click_type=DATETIME_TYPE, help="Window end date."), + ] = None, + count: Annotated[ + int, + typer.Option(help="Maximum number of ticks to return."), + ] = 10000, + flags: Annotated[ + int, + typer.Option( + click_type=TICK_FLAGS_TYPE, + help="Tick flags (ALL, INFO, TRADE, or integer).", + ), + ] = 1, +) -> None: + """Export ticks from a recent time window.""" + client = _sdk_client(ctx) + _execute_export( + ctx, + lambda: client.recent_ticks( + symbol, + seconds, + date_to=date_to, + count=count, + flags=flags, + ), + ) + + @app.command() def account_info(ctx: typer.Context) -> None: """Export account information.""" @@ -335,6 +373,16 @@ def symbol_info( _execute_export(ctx, lambda: client.symbol_info(symbol)) +@app.command() +def minimum_margins( + ctx: typer.Context, + symbol: Annotated[str, typer.Option(help="Symbol name.")], +) -> None: + """Export minimum-volume buy and sell margin requirements.""" + client = _sdk_client(ctx) + _execute_export(ctx, lambda: client.minimum_margins(symbol)) + + @app.command() def orders( ctx: typer.Context, diff --git a/mt5cli/history.py b/mt5cli/history.py index 6edba85..fbe6e87 100644 --- a/mt5cli/history.py +++ b/mt5cli/history.py @@ -5,6 +5,7 @@ from __future__ import annotations import logging import sqlite3 from datetime import UTC, datetime +from pathlib import Path from typing import TYPE_CHECKING, Literal import pandas as pd @@ -122,6 +123,230 @@ def build_rate_view_name( return f"rate_{symbol}__{granularity}_{timeframe}" +SqliteConnOrPath = sqlite3.Connection | Path | str + + +def _open_history_connection( + conn_or_path: SqliteConnOrPath, +) -> tuple[sqlite3.Connection | None, bool]: + """Open a read-only SQLite connection when given a path. + + Returns: + A connection and whether the caller should close it. When the path does + not exist, returns ``(None, False)`` without creating a database file. + """ + if isinstance(conn_or_path, sqlite3.Connection): + return conn_or_path, False + path = Path(conn_or_path) + if not path.exists(): + return None, False + conn = sqlite3.connect(f"{path.resolve().as_uri()}?mode=ro", uri=True) + return conn, True + + +def _load_rates_timeframe_counts(conn: sqlite3.Connection) -> dict[str, int] | None: + """Return distinct timeframe counts per symbol from the normalized rates table.""" + columns = get_table_columns(conn, Dataset.rates.table_name) + if not {"symbol", "timeframe"}.issubset(columns): + return None + rows = conn.execute( + "SELECT symbol, COUNT(DISTINCT timeframe) FROM rates GROUP BY symbol", + ).fetchall() + return {str(symbol): int(count) for symbol, count in rows} + + +def _load_existing_rate_views(conn: sqlite3.Connection) -> set[str]: + """Return mt5cli-managed ``rate_*__*`` compatibility view names.""" + rows = conn.execute( + "SELECT name FROM sqlite_master WHERE type = 'view' AND name GLOB 'rate_*__*'", + ).fetchall() + return {str(row[0]) for row in rows} + + +def _rate_view_name_candidates( + *, + symbol: str, + granularity: str, + granularity_count: int, + timeframe: int, +) -> list[str]: + """Return candidate view names in preference order.""" + single = build_rate_view_name( + symbol=symbol, + granularity=granularity, + granularity_count=1, + timeframe=timeframe, + ) + if granularity_count <= 1: + return [single] + multi = build_rate_view_name( + symbol=symbol, + granularity=granularity, + granularity_count=granularity_count, + timeframe=timeframe, + ) + return [multi, single] + + +def _resolve_rate_view_name_from_context( + *, + symbol: str, + timeframe: int, + granularity_name: str, + timeframe_counts: dict[str, int] | None, + existing_views: set[str], + require_existing: bool = False, +) -> str: + """Resolve one rate view name using preloaded SQLite metadata. + + Returns: + Preferred mt5cli-managed rate compatibility view name. + + Raises: + ValueError: If ``require_existing`` is True and no managed view exists. + """ + if timeframe_counts is None or symbol not in timeframe_counts: + candidates = [ + build_rate_view_name( + symbol=symbol, + granularity=granularity_name, + granularity_count=1, + timeframe=timeframe, + ), + build_rate_view_name( + symbol=symbol, + granularity=granularity_name, + granularity_count=2, + timeframe=timeframe, + ), + ] + else: + candidates = _rate_view_name_candidates( + symbol=symbol, + granularity=granularity_name, + granularity_count=timeframe_counts[symbol], + timeframe=timeframe, + ) + for candidate in candidates: + if candidate in existing_views: + return candidate + if require_existing: + msg = ( + f"No rate compatibility view exists for symbol {symbol!r} " + f"and granularity {granularity_name!r}; " + f"candidates: {', '.join(candidates)}." + ) + raise ValueError(msg) + return candidates[0] + + +def resolve_rate_view_name( + conn_or_path: SqliteConnOrPath, + symbol: str, + granularity: str, + *, + require_existing: bool = False, +) -> str: + """Resolve the mt5cli-managed rate compatibility view name. + + Args: + conn_or_path: SQLite database path or open connection. + symbol: Symbol stored in the normalized ``rates`` table. + granularity: Timeframe name (for example ``M1``) or integer string. + require_existing: When True, require the database and a managed view to exist. + + Returns: + View name such as ``rate_EURUSD__1`` or ``rate_EURUSD__M1_1``. + + Raises: + ValueError: If ``require_existing`` is True and the database or view is missing. + """ + timeframe = parse_timeframe(granularity) + granularity_name = resolve_granularity_name(timeframe) + conn, should_close = _open_history_connection(conn_or_path) + try: + if conn is None: + if require_existing: + path = ( + conn_or_path + if isinstance(conn_or_path, (Path, str)) + else "database" + ) + msg = f"SQLite database not found: {path}" + raise ValueError(msg) + return build_rate_view_name( + symbol=symbol, + granularity=granularity_name, + granularity_count=1, + timeframe=timeframe, + ) + return _resolve_rate_view_name_from_context( + symbol=symbol, + timeframe=timeframe, + granularity_name=granularity_name, + timeframe_counts=_load_rates_timeframe_counts(conn), + existing_views=_load_existing_rate_views(conn), + require_existing=require_existing, + ) + finally: + if should_close and conn is not None: + conn.close() + + +def resolve_rate_view_names( + conn_or_path: SqliteConnOrPath, + symbols: Sequence[str], + granularities: Sequence[str], + *, + require_existing: bool = False, +) -> list[str]: + """Resolve rate compatibility view names for symbol and granularity pairs. + + Args: + conn_or_path: SQLite database path or open connection. + symbols: Symbols stored in the normalized ``rates`` table. + granularities: Timeframe names (for example ``M1``) or integer strings. + require_existing: When True, require the database and managed views to exist. + + Returns: + View names in row-major order: every ``granularity`` for the first + symbol, then every granularity for the next symbol, and so on. + """ + conn, should_close = _open_history_connection(conn_or_path) + try: + if conn is None: + return [ + resolve_rate_view_name( + conn_or_path, + symbol, + granularity, + require_existing=require_existing, + ) + for symbol in symbols + for granularity in granularities + ] + timeframe_counts = _load_rates_timeframe_counts(conn) + existing_views = _load_existing_rate_views(conn) + resolved: list[str] = [] + for symbol in symbols: + for granularity in granularities: + timeframe = parse_timeframe(granularity) + resolved.append( + _resolve_rate_view_name_from_context( + symbol=symbol, + timeframe=timeframe, + granularity_name=resolve_granularity_name(timeframe), + timeframe_counts=timeframe_counts, + existing_views=existing_views, + require_existing=require_existing, + ), + ) + return resolved + finally: + if should_close and conn is not None: + conn.close() + + def get_table_columns(conn: sqlite3.Connection, table: str) -> set[str]: """Return existing SQLite columns for a table.""" rows = conn.execute(f"PRAGMA table_info({table})").fetchall() diff --git a/mt5cli/sdk.py b/mt5cli/sdk.py index 6e82903..16a3521 100644 --- a/mt5cli/sdk.py +++ b/mt5cli/sdk.py @@ -10,6 +10,7 @@ from datetime import UTC, datetime, timedelta from pathlib import Path from typing import TYPE_CHECKING, Self, TypeVar +import pandas as pd from pdmt5 import Mt5Config, Mt5DataClient from .history import ( @@ -33,8 +34,6 @@ from .utils import ( if TYPE_CHECKING: from collections.abc import Callable, Iterator, Sequence - import pandas as pd - T = TypeVar("T") logger = logging.getLogger(__name__) @@ -53,8 +52,10 @@ __all__ = [ "history_orders", "last_error", "market_book", + "minimum_margins", "orders", "positions", + "recent_ticks", "symbol_info", "symbol_info_tick", "symbols", @@ -89,6 +90,89 @@ def _coerce_datetime(value: datetime | str | None) -> datetime | None: return parse_datetime(value) +def _coerce_tick_time(value: object) -> datetime: + if isinstance(value, datetime): + return value + if isinstance(value, str): + return parse_datetime(value) + if isinstance(value, (int, float)): + return datetime.fromtimestamp(value, tz=UTC) + msg = f"Unsupported tick time value: {value!r}" + raise TypeError(msg) + + +def _filter_ticks_to_end(frame: pd.DataFrame, end: datetime) -> pd.DataFrame: + if frame.empty or "time" not in frame.columns: + return frame + times = pd.to_datetime(frame["time"], utc=True) + return frame.loc[times <= end].reset_index(drop=True) + + +def _fetch_recent_ticks( + client: Mt5DataClient, + symbol: str, + seconds: float, + date_to: datetime | None, + count: int, + flags: int, +) -> pd.DataFrame: + if date_to is not None: + end = date_to + else: + tick = client.symbol_info_tick(symbol) + end = _coerce_tick_time(tick.time) + start = end - timedelta(seconds=seconds) + if count > 0: + from_frame = _filter_ticks_to_end( + client.copy_ticks_from_as_df( + symbol=symbol, + date_from=start, + count=count, + flags=flags, + ), + end, + ) + if len(from_frame) < count: + return from_frame + frame = client.copy_ticks_range_as_df( + symbol=symbol, + date_from=start, + date_to=end, + flags=flags, + ) + if count > 0 and len(frame) > count: + return frame.tail(count).reset_index(drop=True) + return frame + + +def _fetch_minimum_margins(client: Mt5DataClient, symbol: str) -> pd.DataFrame: + sym = client.symbol_info(symbol) + account = client.account_info() + tick = client.symbol_info_tick(symbol) + volume_min = sym.volume_min + buy_margin = client.order_calc_margin( + client.mt5.ORDER_TYPE_BUY, + symbol, + volume_min, + tick.ask, + ) + sell_margin = client.order_calc_margin( + client.mt5.ORDER_TYPE_SELL, + symbol, + volume_min, + tick.bid, + ) + return pd.DataFrame([ + { + "symbol": symbol, + "account_currency": account.currency, + "volume_min": volume_min, + "buy_margin": buy_margin, + "sell_margin": sell_margin, + } + ]) + + def build_config( *, path: str | None = None, @@ -418,6 +502,57 @@ class Mt5CliClient: """Return market depth for a symbol.""" return self._fetch(lambda c: c.market_book_get_as_df(symbol=symbol)) + def recent_ticks( + self, + symbol: str, + seconds: float, + *, + date_to: datetime | str | None = None, + count: int = 10000, + flags: int | str = "ALL", + ) -> pd.DataFrame: + """Return ticks from a recent time window. + + Args: + symbol: Symbol name. + seconds: Lookback window in seconds ending at ``date_to``. + date_to: Window end time. When ``None``, uses the latest + ``symbol_info_tick().time`` rather than wall-clock now. + count: Maximum ticks to return. Values ``<= 0`` return the full + window without trimming. Positive values keep the most recent + ticks; when the window is sparse, ``copy_ticks_from`` avoids + fetching the entire range. + flags: Tick flags as ``ALL``, ``INFO``, ``TRADE``, or an integer. + + Returns: + Tick DataFrame with MT5 tick columns such as ``time``, ``bid``, + ``ask``, ``last``, and ``volume``. + """ + tick_flags = _coerce_tick_flags(flags) + end = _coerce_datetime(date_to) + return self._fetch( + lambda c: _fetch_recent_ticks( + c, + symbol, + seconds, + end, + count, + tick_flags, + ), + ) + + def minimum_margins(self, symbol: str) -> pd.DataFrame: + """Return minimum-volume buy and sell margin requirements. + + Args: + symbol: Symbol name. + + Returns: + One-row DataFrame with columns ``symbol``, ``account_currency``, + ``volume_min``, ``buy_margin``, and ``sell_margin``. + """ + return self._fetch(lambda c: _fetch_minimum_margins(c, symbol)) + def _resolve_incremental_settings( selected_datasets: set[Dataset], @@ -915,3 +1050,37 @@ def market_book( ) -> pd.DataFrame: """Return market depth for a symbol.""" return _make_client(config=config).market_book(symbol) + + +def recent_ticks( + symbol: str, + seconds: float, + *, + date_to: datetime | str | None = None, + count: int = 10000, + flags: int | str = "ALL", + config: Mt5Config | None = None, +) -> pd.DataFrame: + """Return ticks from a recent time window ending at ``date_to`` or now. + + See ``Mt5CliClient.recent_ticks`` for parameter and return details. + """ + return _make_client(config=config).recent_ticks( + symbol, + seconds, + date_to=date_to, + count=count, + flags=flags, + ) + + +def minimum_margins( + symbol: str, + *, + config: Mt5Config | None = None, +) -> pd.DataFrame: + """Return minimum-volume buy and sell margin requirements. + + See ``Mt5CliClient.minimum_margins`` for return details. + """ + return _make_client(config=config).minimum_margins(symbol) diff --git a/mt5cli/utils.py b/mt5cli/utils.py index e393a57..fb8fba1 100644 --- a/mt5cli/utils.py +++ b/mt5cli/utils.py @@ -2,16 +2,18 @@ from __future__ import annotations -import importlib import json +import sqlite3 from datetime import UTC, datetime from enum import StrEnum from pathlib import Path -from typing import TYPE_CHECKING, Any, TypeGuard, cast +from typing import TYPE_CHECKING, Any, TypeGuard import click if TYPE_CHECKING: + from collections.abc import Sequence + import pandas as pd # --------------------------------------------------------------------------- @@ -260,6 +262,50 @@ def detect_format( raise ValueError(msg) +def export_dataframe_to_sqlite( + df: pd.DataFrame, + output_path: Path, + table_name: str = "data", + *, + if_exists: IfExists = IfExists.APPEND, + index: bool = False, + index_label: str | None = None, + deduplicate_on: Sequence[str] | None = None, +) -> None: + """Write a DataFrame to SQLite with configurable append and deduplication. + + Args: + df: DataFrame to export. + output_path: SQLite database path. + table_name: Target table name. + if_exists: Conflict behavior when the table already exists. + index: Whether to write the DataFrame index as a column. + index_label: Column name for the index when ``index=True``. + deduplicate_on: Optional key columns to deduplicate after writing, + keeping the latest ``ROWID`` per key group. Deduplication scans the + full table, so repeated appends cost O(table size); index the key + columns when appending frequently. + """ + with sqlite3.connect(output_path) as conn: + df.to_sql( # type: ignore[reportUnknownMemberType] + table_name, + conn, + if_exists=if_exists.value, + index=index, + index_label=index_label, + ) + if deduplicate_on: + from .history import drop_duplicates_in_table # noqa: PLC0415 + + drop_duplicates_in_table( + conn.cursor(), + table_name, + list(deduplicate_on), + keep="last", + ) + conn.commit() + + def export_dataframe( df: pd.DataFrame, output_path: Path, @@ -289,14 +335,13 @@ def export_dataframe( elif output_format == "parquet": df.to_parquet(output_path, index=False) elif output_format == "sqlite3": - sqlite3 = cast("Any", importlib.import_module("sqlite3")) - with sqlite3.connect(output_path) as conn: - df.to_sql( # type: ignore[reportUnknownMemberType] - table_name, - conn, - if_exists="replace", - index=False, - ) + export_dataframe_to_sqlite( + df, + output_path, + table_name, + if_exists=IfExists.REPLACE, + index=False, + ) else: msg = f"Unsupported output format: {output_format}" raise ValueError(msg) diff --git a/pyproject.toml b/pyproject.toml index 0ff659d..332ba54 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "mt5cli" -version = "0.4.2" +version = "0.4.3" description = "Command-line tool for MetaTrader 5" authors = [{name = "dceoy", email = "dceoy@users.noreply.github.com"}] maintainers = [{name = "dceoy", email = "dceoy@users.noreply.github.com"}] diff --git a/tests/test_cli.py b/tests/test_cli.py index 8b91700..dae8f7b 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -6,7 +6,7 @@ import json import logging import re import sqlite3 -from datetime import UTC, datetime +from datetime import UTC, datetime, timedelta from typing import TYPE_CHECKING from unittest.mock import MagicMock @@ -316,6 +316,65 @@ class TestCommands: flags=2, ) + def test_ticks_recent( + self, + tmp_path: Path, + mock_client: MagicMock, + ) -> None: + """Test ticks-recent command.""" + output = tmp_path / "out.csv" + result = runner.invoke( + app, + [ + "-o", + str(output), + "ticks-recent", + "--symbol", + "EURUSD", + "--seconds", + "120", + "--date-to", + "2024-01-02", + "--count", + "500", + "--flags", + "ALL", + ], + ) + assert result.exit_code == 0, result.output + mock_client.copy_ticks_from_as_df.assert_called_once_with( + symbol="EURUSD", + date_from=datetime(2024, 1, 2, tzinfo=UTC) - timedelta(seconds=120), + count=500, + flags=1, + ) + mock_client.copy_ticks_range_as_df.assert_not_called() + + def test_minimum_margins( + self, + tmp_path: Path, + mock_client: MagicMock, + ) -> None: + """Test minimum-margins command.""" + sym = MagicMock(volume_min=0.01) + account = MagicMock(currency="USD") + tick = MagicMock(ask=1.1010, bid=1.1000) + mock_client.symbol_info.return_value = sym + mock_client.account_info.return_value = account + mock_client.symbol_info_tick.return_value = tick + mock_client.order_calc_margin.side_effect = [12.5, 12.4] + mock_client.mt5.ORDER_TYPE_BUY = 0 + mock_client.mt5.ORDER_TYPE_SELL = 1 + output = tmp_path / "out.csv" + result = runner.invoke( + app, + ["-o", str(output), "minimum-margins", "--symbol", "EURUSD"], + ) + assert result.exit_code == 0, result.output + mock_client.symbol_info.assert_called_once_with("EURUSD") + mock_client.order_calc_margin.assert_any_call(0, "EURUSD", 0.01, 1.1010) + mock_client.order_calc_margin.assert_any_call(1, "EURUSD", 0.01, 1.1000) + def test_orders( self, tmp_path: Path, diff --git a/tests/test_history.py b/tests/test_history.py index f631f17..ce5075f 100644 --- a/tests/test_history.py +++ b/tests/test_history.py @@ -38,6 +38,8 @@ from mt5cli.history import ( resolve_history_datasets, resolve_history_tick_flags, resolve_history_timeframes, + resolve_rate_view_name, + resolve_rate_view_names, write_collected_datasets, write_history_dataset, write_incremental_datasets, @@ -47,6 +49,281 @@ from mt5cli.history import ( from mt5cli.utils import TIMEFRAME_MAP, Dataset, IfExists +class TestResolveRateViewName: + """Tests for resolve_rate_view_name and resolve_rate_view_names.""" + + def test_missing_database_path_does_not_create_file(self, tmp_path: Path) -> None: + """Test resolving against a missing path does not create a database.""" + db_path = tmp_path / "missing.db" + assert resolve_rate_view_name(db_path, "EURUSD", "M1") == "rate_EURUSD__1" + assert not db_path.exists() + + def test_no_rates_table_falls_back_to_single_timeframe_name( + self, + tmp_path: Path, + ) -> None: + """Test databases without a rates table use single-timeframe naming.""" + db_path = tmp_path / "no-rates.db" + with sqlite3.connect(db_path) as conn: + conn.execute("CREATE TABLE ticks(symbol TEXT, time TEXT)") + assert resolve_rate_view_name(db_path, "EURUSD", "M1") == "rate_EURUSD__1" + + def test_single_timeframe_for_one_symbol(self, tmp_path: Path) -> None: + """Test one stored timeframe resolves to the short view name.""" + db_path = tmp_path / "single-timeframe.db" + with sqlite3.connect(db_path) as conn: + conn.execute( + "CREATE TABLE rates(" + " symbol TEXT, timeframe INTEGER, time TEXT, close REAL)", + ) + conn.execute( + "INSERT INTO rates(symbol, timeframe, time, close) VALUES (?, ?, ?, ?)", + ("EURUSD", 1, "2024-01-01T00:00:00+00:00", 1.0), + ) + create_rate_compatibility_views(conn) + assert resolve_rate_view_name(db_path, "EURUSD", "M1") == "rate_EURUSD__1" + + def test_multiple_timeframes_for_one_symbol(self, tmp_path: Path) -> None: + """Test multiple stored timeframes resolve to disambiguated view names.""" + db_path = tmp_path / "multi-timeframe.db" + with sqlite3.connect(db_path) as conn: + conn.execute( + "CREATE TABLE rates(" + " symbol TEXT, timeframe INTEGER, time TEXT, close REAL)", + ) + conn.executemany( + "INSERT INTO rates(symbol, timeframe, time, close) VALUES (?, ?, ?, ?)", + [ + ("EURUSD", 1, "2024-01-01T00:00:00+00:00", 1.0), + ("EURUSD", TIMEFRAME_MAP["H1"], "2024-01-01T01:00:00+00:00", 1.1), + ], + ) + create_rate_compatibility_views(conn) + assert resolve_rate_view_name(db_path, "EURUSD", "M1") == "rate_EURUSD__M1_1" + assert ( + resolve_rate_view_name(db_path, "EURUSD", "H1") == "rate_EURUSD__H1_16385" + ) + + def test_prefers_multi_name_when_both_candidate_views_exist( + self, + tmp_path: Path, + ) -> None: + """Test multi-timeframe metadata wins over stale single-timeframe views.""" + db_path = tmp_path / "stale-and-current-views.db" + with sqlite3.connect(db_path) as conn: + conn.execute( + "CREATE TABLE rates(" + " symbol TEXT, timeframe INTEGER, time TEXT, close REAL)", + ) + conn.executemany( + "INSERT INTO rates(symbol, timeframe, time, close) VALUES (?, ?, ?, ?)", + [ + ("EURUSD", 1, "2024-01-01T00:00:00+00:00", 1.0), + ("EURUSD", TIMEFRAME_MAP["H1"], "2024-01-01T01:00:00+00:00", 1.1), + ], + ) + conn.execute( + 'CREATE VIEW "rate_EURUSD__1" AS' + " SELECT time, close FROM rates" + " WHERE symbol = 'EURUSD' AND timeframe = 1", + ) + conn.execute( + 'CREATE VIEW "rate_EURUSD__M1_1" AS' + " SELECT time, close FROM rates" + " WHERE symbol = 'EURUSD' AND timeframe = 1", + ) + assert resolve_rate_view_name(db_path, "EURUSD", "M1") == "rate_EURUSD__M1_1" + + def test_prefers_existing_view_when_metadata_unavailable( + self, + tmp_path: Path, + ) -> None: + """Test an existing managed view is preferred without rates metadata.""" + db_path = tmp_path / "view-only.db" + with sqlite3.connect(db_path) as conn: + conn.execute("CREATE TABLE ticks(symbol TEXT, time TEXT)") + conn.execute('CREATE VIEW "rate_EURUSD__M1_1" AS SELECT 1 AS close') + assert resolve_rate_view_name(db_path, "EURUSD", "M1") == "rate_EURUSD__M1_1" + + def test_symbol_absent_from_rates_metadata_uses_candidate_pair( + self, + tmp_path: Path, + ) -> None: + """Test symbols missing from rates metadata still resolve known views.""" + db_path = tmp_path / "other-symbol-only.db" + with sqlite3.connect(db_path) as conn: + conn.execute( + "CREATE TABLE rates(" + " symbol TEXT, timeframe INTEGER, time TEXT, close REAL)", + ) + conn.execute( + "INSERT INTO rates(symbol, timeframe, time, close) VALUES (?, ?, ?, ?)", + ("GBPUSD", 1, "2024-01-01T00:00:00+00:00", 1.0), + ) + conn.execute( + 'CREATE VIEW "rate_EURUSD__1" AS' + " SELECT time, close FROM rates" + " WHERE symbol = 'EURUSD' AND timeframe = 1", + ) + assert resolve_rate_view_name(db_path, "EURUSD", "M1") == "rate_EURUSD__1" + + def test_ignores_non_compatibility_rate_views(self, tmp_path: Path) -> None: + """Test unrelated rate_* views without the __ separator are ignored.""" + db_path = tmp_path / "summary-view.db" + with sqlite3.connect(db_path) as conn: + conn.execute("CREATE TABLE ticks(symbol TEXT, time TEXT)") + conn.execute('CREATE VIEW "rate_summary" AS SELECT 1 AS close') + assert resolve_rate_view_name(db_path, "EURUSD", "M1") == "rate_EURUSD__1" + + def test_invalid_granularity_propagates_value_error(self, tmp_path: Path) -> None: + """Test invalid granularities raise ValueError from parse_timeframe.""" + with pytest.raises(ValueError, match="Invalid timeframe"): + resolve_rate_view_name(tmp_path / "unused.db", "EURUSD", "BAD") + with pytest.raises(ValueError, match="Invalid timeframe"): + resolve_rate_view_names(tmp_path / "unused.db", ["EURUSD"], ["BAD"]) + + def test_resolve_rate_view_names_for_multiple_pairs(self, tmp_path: Path) -> None: + """Test batch resolution returns row-major symbol/granularity pairs.""" + db_path = tmp_path / "batch-resolve.db" + with sqlite3.connect(db_path) as conn: + conn.execute( + "CREATE TABLE rates(" + " symbol TEXT, timeframe INTEGER, time TEXT, close REAL)", + ) + conn.executemany( + "INSERT INTO rates(symbol, timeframe, time, close) VALUES (?, ?, ?, ?)", + [ + ("EURUSD", 1, "2024-01-01T00:00:00+00:00", 1.0), + ("EURUSD", TIMEFRAME_MAP["H1"], "2024-01-01T01:00:00+00:00", 1.1), + ("GBPUSD", 1, "2024-01-01T00:00:00+00:00", 1.2), + ], + ) + create_rate_compatibility_views(conn) + assert resolve_rate_view_names( + db_path, + ["EURUSD", "GBPUSD"], + ["M1", "H1"], + ) == [ + "rate_EURUSD__M1_1", + "rate_EURUSD__H1_16385", + "rate_GBPUSD__1", + "rate_GBPUSD__16385", + ] + + @pytest.mark.parametrize( + "symbol", + ["EUR/USD", "US500.cash", "#US500"], + ) + def test_supports_broker_specific_symbols( + self, + tmp_path: Path, + symbol: str, + ) -> None: + """Test broker-specific symbols resolve to safely created view names.""" + db_path = tmp_path / "broker-symbol-resolve.db" + with sqlite3.connect(db_path) as conn: + conn.execute( + "CREATE TABLE rates(" + " symbol TEXT, timeframe INTEGER, time TEXT, close REAL)", + ) + conn.execute( + "INSERT INTO rates(symbol, timeframe, time, close) VALUES (?, ?, ?, ?)", + (symbol, 1, "2024-01-01T00:00:00+00:00", 1.0), + ) + create_rate_compatibility_views(conn) + assert resolve_rate_view_name(db_path, symbol, "M1") == build_rate_view_name( + symbol=symbol, + granularity="M1", + granularity_count=1, + timeframe=1, + ) + + def test_accepts_open_sqlite_connection(self, tmp_path: Path) -> None: + """Test resolver accepts an already-open SQLite connection.""" + db_path = tmp_path / "open-connection.db" + with sqlite3.connect(db_path) as conn: + conn.execute( + "CREATE TABLE rates(" + " symbol TEXT, timeframe INTEGER, time TEXT, close REAL)", + ) + conn.execute( + "INSERT INTO rates(symbol, timeframe, time, close) VALUES (?, ?, ?, ?)", + ("EURUSD", 1, "2024-01-01T00:00:00+00:00", 1.0), + ) + create_rate_compatibility_views(conn) + assert resolve_rate_view_name(conn, "EURUSD", "M1") == "rate_EURUSD__1" + + def test_require_existing_raises_when_database_missing( + self, + tmp_path: Path, + ) -> None: + """Test strict mode rejects missing database paths.""" + db_path = tmp_path / "missing.db" + with pytest.raises(ValueError, match="SQLite database not found"): + resolve_rate_view_name( + db_path, + "EURUSD", + "M1", + require_existing=True, + ) + with pytest.raises(ValueError, match="SQLite database not found"): + resolve_rate_view_names( + db_path, + ["EURUSD"], + ["M1"], + require_existing=True, + ) + + def test_require_existing_raises_when_view_missing(self, tmp_path: Path) -> None: + """Test strict mode rejects databases without matching rate views.""" + db_path = tmp_path / "no-view.db" + with sqlite3.connect(db_path) as conn: + conn.execute("CREATE TABLE ticks(symbol TEXT, time TEXT)") + with pytest.raises(ValueError, match="No rate compatibility view exists"): + resolve_rate_view_name( + db_path, + "EURUSD", + "M1", + require_existing=True, + ) + with pytest.raises(ValueError, match="No rate compatibility view exists"): + resolve_rate_view_names( + db_path, + ["EURUSD"], + ["M1"], + require_existing=True, + ) + + def test_require_existing_returns_existing_view(self, tmp_path: Path) -> None: + """Test strict mode returns a view when one exists.""" + db_path = tmp_path / "existing-view.db" + with sqlite3.connect(db_path) as conn: + conn.execute( + "CREATE TABLE rates(" + " symbol TEXT, timeframe INTEGER, time TEXT, close REAL)", + ) + conn.execute( + "INSERT INTO rates(symbol, timeframe, time, close) VALUES (?, ?, ?, ?)", + ("EURUSD", 1, "2024-01-01T00:00:00+00:00", 1.0), + ) + create_rate_compatibility_views(conn) + assert ( + resolve_rate_view_name( + db_path, + "EURUSD", + "M1", + require_existing=True, + ) + == "rate_EURUSD__1" + ) + assert resolve_rate_view_names( + db_path, + ["EURUSD"], + ["M1"], + require_existing=True, + ) == ["rate_EURUSD__1"] + + class TestQuoteSqliteIdentifier: """Tests for quote_sqlite_identifier.""" diff --git a/tests/test_sdk.py b/tests/test_sdk.py index 98e5237..3d8aa2a 100644 --- a/tests/test_sdk.py +++ b/tests/test_sdk.py @@ -4,7 +4,7 @@ from __future__ import annotations import logging import sqlite3 -from datetime import UTC, datetime +from datetime import UTC, datetime, timedelta from typing import TYPE_CHECKING from unittest.mock import MagicMock @@ -31,8 +31,10 @@ from mt5cli.sdk import ( history_orders, last_error, market_book, + minimum_margins, orders, positions, + recent_ticks, symbol_info, symbol_info_tick, symbols, @@ -808,3 +810,173 @@ class TestUpdateHistory: ) after = datetime.now(UTC) assert before <= captured["end"] <= after + + +class TestRecentTicks: + """Tests for recent_ticks helper.""" + + def test_recent_ticks_uses_explicit_date_to_window( + self, + mocker: MockerFixture, + ) -> None: + """Test recent_ticks fetches the requested trailing window.""" + client = MagicMock() + end = datetime(2024, 1, 2, 12, 0, 0, tzinfo=UTC) + client.copy_ticks_from_as_df.return_value = pd.DataFrame({ + "time": [end], + "bid": [1.0], + }) + mocker.patch("mt5cli.sdk.Mt5DataClient", return_value=client) + result = recent_ticks( + "EURUSD", + 60, + date_to=end, + count=100, + flags="INFO", + config=build_config(login=123), + ) + assert isinstance(result, pd.DataFrame) + client.copy_ticks_from_as_df.assert_called_once_with( + symbol="EURUSD", + date_from=end - timedelta(seconds=60), + count=100, + flags=2, + ) + client.copy_ticks_range_as_df.assert_not_called() + + def test_recent_ticks_uses_latest_tick_when_date_to_omitted( + self, + mocker: MockerFixture, + ) -> None: + """Test recent_ticks anchors the window on the latest tick time.""" + client = MagicMock() + tick = MagicMock() + tick.time = datetime(2024, 1, 2, 12, 0, 0, tzinfo=UTC) + client.symbol_info_tick.return_value = tick + client.copy_ticks_from_as_df.return_value = pd.DataFrame({ + "time": [1, 2], + "bid": [1.0, 1.1], + }) + client.copy_ticks_range_as_df.return_value = pd.DataFrame({ + "time": [1, 2, 3], + "bid": [1.0, 1.1, 1.2], + }) + mocker.patch("mt5cli.sdk.Mt5DataClient", return_value=client) + result = Mt5CliClient().recent_ticks("EURUSD", 30, count=2, flags="ALL") + assert len(result) == 2 + client.symbol_info_tick.assert_called_once_with("EURUSD") + client.copy_ticks_from_as_df.assert_called_once() + _, kwargs = client.copy_ticks_range_as_df.call_args + assert kwargs["symbol"] == "EURUSD" + assert kwargs["date_to"] == tick.time + assert kwargs["date_from"] == tick.time - timedelta(seconds=30) + assert kwargs["flags"] == 1 + + def test_recent_ticks_rejects_unsupported_tick_time( + self, + mocker: MockerFixture, + ) -> None: + """Test recent_ticks raises when the latest tick time is unsupported.""" + client = MagicMock() + tick = MagicMock() + tick.time = object() + client.symbol_info_tick.return_value = tick + mocker.patch("mt5cli.sdk.Mt5DataClient", return_value=client) + with pytest.raises(TypeError, match="Unsupported tick time value"): + Mt5CliClient().recent_ticks("EURUSD", 30) + + @pytest.mark.parametrize( + "tick_time", + [ + "2024-01-02T12:00:00+00:00", + 1704196800, + ], + ) + def test_recent_ticks_coerces_string_and_unix_tick_times( + self, + mocker: MockerFixture, + tick_time: str | int, + ) -> None: + """Test recent_ticks accepts string and unix tick timestamps.""" + client = MagicMock() + tick = MagicMock() + tick.time = tick_time + client.symbol_info_tick.return_value = tick + expected_end = ( + datetime(2024, 1, 2, 12, 0, 0, tzinfo=UTC) + if isinstance(tick_time, str) + else datetime.fromtimestamp(tick_time, tz=UTC) + ) + client.copy_ticks_from_as_df.return_value = pd.DataFrame({ + "time": [expected_end], + }) + mocker.patch("mt5cli.sdk.Mt5DataClient", return_value=client) + Mt5CliClient().recent_ticks("EURUSD", 30) + _, kwargs = client.copy_ticks_from_as_df.call_args + assert kwargs["date_from"] == expected_end - timedelta(seconds=30) + + def test_recent_ticks_returns_full_frame_when_count_not_positive( + self, + mocker: MockerFixture, + ) -> None: + """Test non-positive count returns the full range without trimming.""" + client = MagicMock() + end = datetime(2024, 1, 2, 12, 0, 0, tzinfo=UTC) + client.copy_ticks_range_as_df.return_value = pd.DataFrame({ + "time": [1, 2, 3], + "bid": [1.0, 1.1, 1.2], + }) + mocker.patch("mt5cli.sdk.Mt5DataClient", return_value=client) + result = recent_ticks( + "EURUSD", + 60, + date_to=end, + count=0, + config=build_config(login=123), + ) + assert len(result) == 3 + client.copy_ticks_from_as_df.assert_not_called() + client.copy_ticks_range_as_df.assert_called_once_with( + symbol="EURUSD", + date_from=end - timedelta(seconds=60), + date_to=end, + flags=1, + ) + + +class TestMinimumMargins: + """Tests for minimum_margins helper.""" + + def test_minimum_margins_shape( + self, + mocker: MockerFixture, + ) -> None: + """Test minimum_margins returns the expected summary columns.""" + client = MagicMock() + sym = MagicMock(volume_min=0.01) + account = MagicMock(currency="USD") + tick = MagicMock(ask=1.1010, bid=1.1000) + client.symbol_info.return_value = sym + client.account_info.return_value = account + client.symbol_info_tick.return_value = tick + client.order_calc_margin.side_effect = [12.5, 12.4] + client.mt5.ORDER_TYPE_BUY = 0 + client.mt5.ORDER_TYPE_SELL = 1 + mocker.patch("mt5cli.sdk.Mt5DataClient", return_value=client) + + result = minimum_margins("EURUSD", config=build_config(login=123)) + + pd.testing.assert_frame_equal( + result, + pd.DataFrame([ + { + "symbol": "EURUSD", + "account_currency": "USD", + "volume_min": 0.01, + "buy_margin": 12.5, + "sell_margin": 12.4, + } + ]), + ) + client.order_calc_margin.assert_any_call(0, "EURUSD", 0.01, 1.1010) + client.order_calc_margin.assert_any_call(1, "EURUSD", 0.01, 1.1000) diff --git a/tests/test_utils.py b/tests/test_utils.py index fa635c8..24da1b3 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -21,8 +21,10 @@ from mt5cli.utils import ( TIMEFRAME_MAP, TIMEFRAME_TYPE, Dataset, + IfExists, detect_format, export_dataframe, + export_dataframe_to_sqlite, parse_datetime, parse_request, parse_tick_flags, @@ -130,6 +132,112 @@ class TestExportDataframe: export_dataframe(sample_df, tmp_path / "out.txt", "xml") +class TestExportDataframeToSqlite: + """Tests for export_dataframe_to_sqlite.""" + + def test_append_preserves_existing_rows(self, tmp_path: Path) -> None: + """Test append mode keeps prior rows in the SQLite table.""" + output = tmp_path / "append.db" + first = pd.DataFrame({"id": [1], "value": ["a"]}) + second = pd.DataFrame({"id": [2], "value": ["b"]}) + export_dataframe_to_sqlite(first, output, "items", if_exists=IfExists.REPLACE) + export_dataframe_to_sqlite(second, output, "items", if_exists=IfExists.APPEND) + with sqlite3.connect(output) as conn: + result = pd.read_sql( # type: ignore[reportUnknownMemberType] + "SELECT id, value FROM items ORDER BY id", + conn, + ) + pd.testing.assert_frame_equal( + result, + pd.DataFrame({"id": [1, 2], "value": ["a", "b"]}), + ) + + def test_deduplicate_keeps_latest_row(self, tmp_path: Path) -> None: + """Test deduplication keeps the latest ROWID for key columns.""" + output = tmp_path / "dedup.db" + first = pd.DataFrame({ + "symbol": ["EURUSD", "EURUSD"], + "time": ["2024-01-01", "2024-01-01"], + "bid": [1.0, 1.1], + }) + second = pd.DataFrame({ + "symbol": ["EURUSD"], + "time": ["2024-01-01"], + "bid": [1.2], + }) + export_dataframe_to_sqlite( + first, + output, + "ticks", + if_exists=IfExists.REPLACE, + deduplicate_on=("symbol", "time"), + ) + export_dataframe_to_sqlite( + second, + output, + "ticks", + if_exists=IfExists.APPEND, + deduplicate_on=("symbol", "time"), + ) + with sqlite3.connect(output) as conn: + result = pd.read_sql( # type: ignore[reportUnknownMemberType] + "SELECT symbol, time, bid FROM ticks", + conn, + ) + pd.testing.assert_frame_equal( + result.reset_index(drop=True), + pd.DataFrame({ + "symbol": ["EURUSD"], + "time": ["2024-01-01"], + "bid": [1.2], + }), + ) + + def test_default_if_exists_appends_without_dropping_rows( + self, + tmp_path: Path, + ) -> None: + """Test the default append mode keeps prior rows.""" + output = tmp_path / "default-append.db" + first = pd.DataFrame({"id": [1], "value": ["a"]}) + second = pd.DataFrame({"id": [2], "value": ["b"]}) + export_dataframe_to_sqlite(first, output, "items") + export_dataframe_to_sqlite(second, output, "items") + with sqlite3.connect(output) as conn: + result = pd.read_sql( # type: ignore[reportUnknownMemberType] + "SELECT id, value FROM items ORDER BY id", + conn, + ) + pd.testing.assert_frame_equal( + result, + pd.DataFrame({"id": [1, 2], "value": ["a", "b"]}), + ) + + def test_writes_index_with_label(self, tmp_path: Path) -> None: + """Test optional index export with a custom label.""" + output = tmp_path / "index.db" + frame = pd.DataFrame( + {"value": [1.0]}, index=pd.Index(["EURUSD"], name="symbol") + ) + export_dataframe_to_sqlite( + frame, + output, + "margins", + if_exists=IfExists.REPLACE, + index=True, + index_label="symbol", + ) + with sqlite3.connect(output) as conn: + result = pd.read_sql( # type: ignore[reportUnknownMemberType] + "SELECT symbol, value FROM margins", + conn, + ) + pd.testing.assert_frame_equal( + result, + pd.DataFrame({"symbol": ["EURUSD"], "value": [1.0]}), + ) + + # --------------------------------------------------------------------------- # Parse helpers # --------------------------------------------------------------------------- diff --git a/uv.lock b/uv.lock index ce1d024..925f709 100644 --- a/uv.lock +++ b/uv.lock @@ -487,7 +487,7 @@ wheels = [ [[package]] name = "mt5cli" -version = "0.4.2" +version = "0.4.3" source = { editable = "." } dependencies = [ { name = "click" },