From 5b44318d5579b64c6c81137b60af3e871fde2aa3 Mon Sep 17 00:00:00 2001 From: Daichi Narushima <1938249+dceoy@users.noreply.github.com> Date: Mon, 8 Jun 2026 22:54:53 +0900 Subject: [PATCH] Add programmatic SDK and refactor mt5cli into cli, sdk, and utils (#15) * Refactor cli.py into cli and utils modules Extract constants, enums, Click parameter types, and parse/export utility functions into a new mt5cli/utils.py module, keeping the typer app, commands, and collect-history SQLite helpers in cli.py. https://claude.ai/code/session_016JwSEhPyq6phXySktQ1FGU * Address review comments * Add programmatic SDK layer for read-only MT5 data collection. Expose Mt5CliClient and collect_history through the package API while keeping CLI commands as thin adapters over the SDK. Co-authored-by: Cursor * Harden SDK connection lifecycle and scope internal helpers as private. Co-authored-by: Cursor * Export build_config in the public API and bump version to 0.4.0. Co-authored-by: Cursor * Remove duplicate scripts/ in favor of local-qa skill script. Co-authored-by: Cursor --------- Co-authored-by: Claude Co-authored-by: Cursor --- AGENTS.md | 4 +- docs/api/index.md | 45 +- docs/api/sdk.md | 3 + docs/api/utils.md | 3 + docs/index.md | 42 +- mkdocs.yml | 2 + mt5cli/__init__.py | 48 +- mt5cli/cli.py | 1032 ++++--------------------------------------- mt5cli/sdk.py | 1023 ++++++++++++++++++++++++++++++++++++++++++ mt5cli/utils.py | 408 +++++++++++++++++ pyproject.toml | 2 +- tests/test_cli.py | 328 +------------- tests/test_sdk.py | 466 +++++++++++++++++++ tests/test_utils.py | 335 ++++++++++++++ uv.lock | 2 +- 15 files changed, 2479 insertions(+), 1264 deletions(-) create mode 100644 docs/api/sdk.md create mode 100644 docs/api/utils.md create mode 100644 mt5cli/sdk.py create mode 100644 mt5cli/utils.py create mode 100644 tests/test_sdk.py create mode 100644 tests/test_utils.py diff --git a/AGENTS.md b/AGENTS.md index 8852bd9..e8e6e8a 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -29,9 +29,11 @@ uv sync - `mt5cli/`: Main package directory - `__init__.py`: Package initialization and exports (`detect_format`, `export_dataframe`) - `cli.py`: CLI application with typer-based commands for data export + - `utils.py`: Constants, enums, parameter types, parsers, and export utilities - `__main__.py`: Entry point for `python -m mt5cli` - `tests/`: Comprehensive test suite (pytest-based) - - `test_cli.py`: Tests for CLI commands, parameter types, and export functions + - `test_cli.py`: Tests for CLI commands and collect-history behavior + - `test_utils.py`: Tests for utility constants, parameter types, parsers, and export functions - `docs/`: MkDocs documentation with API reference - `docs/index.md`: Main documentation - `docs/api/`: Auto-generated API documentation for all modules diff --git a/docs/api/index.md b/docs/api/index.md index 8013438..01ecaaa 100644 --- a/docs/api/index.md +++ b/docs/api/index.md @@ -10,12 +10,22 @@ The mt5cli package consists of the following modules: Command-line interface module providing typer-based commands for exporting MetaTrader 5 data to CSV, JSON, Parquet, and SQLite3 formats. +### [Utils](utils.md) + +Utility module providing constants, enums, Click parameter types, and helper functions for parsing and exporting data. + +### [SDK](sdk.md) + +Programmatic SDK for read-only MetaTrader 5 data collection. Returns pandas DataFrames and provides `collect_history` for SQLite bulk collection. + ## Architecture Overview The package follows a simple architecture built on top of pdmt5: -1. **CLI Layer** (`cli.py`): Typer application with subcommands for each data type, custom Click parameter types for datetime/timeframe/tick flags parsing, and format detection/export utilities. -2. **Data Layer** (via `pdmt5`): Uses `Mt5DataClient` and `Mt5Config` from the pdmt5 package for all MetaTrader 5 data access. +1. **CLI Layer** (`cli.py`): Typer application with subcommands that delegate to the SDK and export results. +2. **SDK Layer** (`sdk.py`): Read-only data access functions, `Mt5CliClient`, and `collect_history` orchestration. +3. **Utils Layer** (`utils.py`): Constants, enums, custom Click parameter types, parsing helpers, and format detection/export utilities. +4. **Data Layer** (via `pdmt5`): Uses `Mt5DataClient` and `Mt5Config` from the pdmt5 package for all MetaTrader 5 data access. ## Usage Guidelines @@ -47,15 +57,38 @@ mt5cli -o data.db --table symbols symbols --group "*USD*" ## Python API ```python -from mt5cli import detect_format, export_dataframe -import pandas as pd +from datetime import UTC, datetime +from pathlib import Path + +from mt5cli import ( + Mt5CliClient, + collect_history, + copy_rates_range, + detect_format, + export_dataframe, +) + +# Fetch rates programmatically +rates = copy_rates_range( + "EURUSD", + timeframe="H1", + date_from="2024-01-01", + date_to="2024-02-01", +) # Detect output format from file extension fmt = detect_format(Path("output.parquet")) # Returns "parquet" # Export a DataFrame -df = pd.DataFrame({"symbol": ["EURUSD"], "bid": [1.1234]}) -export_dataframe(df, Path("output.csv"), "csv") +export_dataframe(rates, Path("output.csv"), "csv") + +# Collect history into SQLite +collect_history( + Path("history.db"), + symbols=["EURUSD"], + date_from=datetime(2024, 1, 1, tzinfo=UTC), + date_to=datetime(2024, 2, 1, tzinfo=UTC), +) ``` ## Examples diff --git a/docs/api/sdk.md b/docs/api/sdk.md new file mode 100644 index 0000000..a7af946 --- /dev/null +++ b/docs/api/sdk.md @@ -0,0 +1,3 @@ +# SDK Module + +::: mt5cli.sdk diff --git a/docs/api/utils.md b/docs/api/utils.md new file mode 100644 index 0000000..854ce0b --- /dev/null +++ b/docs/api/utils.md @@ -0,0 +1,3 @@ +# Utils Module + +::: mt5cli.utils diff --git a/docs/index.md b/docs/index.md index 605bff8..692cecf 100644 --- a/docs/index.md +++ b/docs/index.md @@ -20,6 +20,44 @@ mt5cli is a CLI application that exports MetaTrader 5 trading data to multiple f 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. + +```python +from datetime import UTC, datetime +from pathlib import Path + +from mt5cli import Mt5CliClient, collect_history, copy_rates_range, export_dataframe + +# One-off fetch with module-level helpers +rates = copy_rates_range( + "EURUSD", + timeframe="H1", + date_from="2024-01-01", + date_to="2024-02-01", +) +export_dataframe(rates, Path("rates.csv"), "csv") + +# Reuse one MT5 connection for multiple calls +with Mt5CliClient(login=12345, password="secret", server="Broker-Demo") as client: + account = client.account_info() + positions = client.positions() + +# Bulk SQLite collection (same behavior as the collect-history CLI command) +collect_history( + Path("history.db"), + symbols=["EURUSD", "GBPUSD"], + date_from=datetime(2024, 1, 1, tzinfo=UTC), + date_to=datetime(2024, 2, 1, tzinfo=UTC), + timeframe="M1", + flags="ALL", + with_views=True, +) +``` + +Timeframes, tick flags, and ISO 8601 date strings are accepted wherever noted in the SDK API. + ## Quick Start ```bash @@ -138,7 +176,9 @@ History orders and deals are fetched per symbol and concatenated, so the symbol Browse the API documentation for detailed module information: -- [CLI Module](api/cli.md) - CLI application with export commands and utility functions +- [CLI Module](api/cli.md) - CLI application with export commands +- [SDK Module](api/sdk.md) - Programmatic read-only data collection API +- [Utils Module](api/utils.md) - Constants, parameter types, parsers, and export utilities ## Development diff --git a/mkdocs.yml b/mkdocs.yml index 6a128b6..53e9d48 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -56,6 +56,8 @@ nav: - API Reference: - Overview: api/index.md - CLI: api/cli.md + - SDK: api/sdk.md + - Utils: api/utils.md markdown_extensions: - admonition diff --git a/mt5cli/__init__.py b/mt5cli/__init__.py index c4f174e..ee57b80 100644 --- a/mt5cli/__init__.py +++ b/mt5cli/__init__.py @@ -1,12 +1,56 @@ -"""mt5cli: Command-line tool for MetaTrader 5.""" +"""mt5cli: Command-line tool and SDK for MetaTrader 5.""" from importlib.metadata import version -from .cli import detect_format, export_dataframe +from .sdk import ( + Mt5CliClient, + account_info, + build_config, + collect_history, + copy_rates_from, + copy_rates_from_pos, + copy_rates_range, + copy_ticks_from, + copy_ticks_range, + history_deals, + history_orders, + last_error, + market_book, + orders, + positions, + symbol_info, + symbol_info_tick, + symbols, + terminal_info, +) +from .sdk import ( + version as mt5_version, +) +from .utils import detect_format, export_dataframe __version__ = version(__package__) if __package__ else None __all__ = [ + "Mt5CliClient", + "account_info", + "build_config", + "collect_history", + "copy_rates_from", + "copy_rates_from_pos", + "copy_rates_range", + "copy_ticks_from", + "copy_ticks_range", "detect_format", "export_dataframe", + "history_deals", + "history_orders", + "last_error", + "market_book", + "mt5_version", + "orders", + "positions", + "symbol_info", + "symbol_info_tick", + "symbols", + "terminal_info", ] diff --git a/mt5cli/cli.py b/mt5cli/cli.py index 24a24e1..0189a36 100644 --- a/mt5cli/cli.py +++ b/mt5cli/cli.py @@ -2,18 +2,28 @@ from __future__ import annotations -import json import logging -import sqlite3 from dataclasses import dataclass -from datetime import UTC, datetime -from enum import StrEnum -from pathlib import Path -from typing import TYPE_CHECKING, Annotated, Any, TypeGuard, cast +from datetime import datetime # noqa: TC003 +from pathlib import Path # noqa: TC003 +from typing import TYPE_CHECKING, Annotated, Any, cast -import click import typer -from pdmt5 import Mt5Config, Mt5DataClient +from pdmt5 import Mt5Config + +from . import sdk +from .utils import ( + DATETIME_TYPE, + REQUEST_TYPE, + TICK_FLAGS_TYPE, + TIMEFRAME_TYPE, + Dataset, + IfExists, + LogLevel, + OutputFormat, + detect_format, + export_dataframe, +) if TYPE_CHECKING: from collections.abc import Callable @@ -22,235 +32,6 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) -# --------------------------------------------------------------------------- -# Constants -# --------------------------------------------------------------------------- - -TIMEFRAME_MAP: dict[str, int] = { - "M1": 1, - "M2": 2, - "M3": 3, - "M4": 4, - "M5": 5, - "M6": 6, - "M10": 10, - "M12": 12, - "M15": 15, - "M20": 20, - "M30": 30, - "H1": 16385, - "H2": 16386, - "H3": 16387, - "H4": 16388, - "H6": 16390, - "H8": 16392, - "H12": 16396, - "D1": 16408, - "W1": 32769, - "MN1": 49153, -} - -TICK_FLAG_MAP: dict[str, int] = { - "ALL": 1, - "INFO": 2, - "TRADE": 4, -} - -_TRADE_DEAL_TYPES: tuple[int, int] = (0, 1) -_TRADE_DEAL_TYPES_SQL = f"({', '.join(str(value) for value in _TRADE_DEAL_TYPES)})" -_POSITIONS_VIEW_REQUIRED_COLUMNS: frozenset[str] = frozenset({ - "position_id", - "symbol", - "time", - "type", - "entry", - "volume", - "price", - "profit", -}) - -_FORMAT_EXTENSIONS: dict[str, str] = { - ".csv": "csv", - ".json": "json", - ".parquet": "parquet", - ".pq": "parquet", - ".db": "sqlite3", - ".sqlite": "sqlite3", - ".sqlite3": "sqlite3", -} - -# --------------------------------------------------------------------------- -# Enums -# --------------------------------------------------------------------------- - - -class OutputFormat(StrEnum): - """Supported output file formats.""" - - csv = "csv" - json = "json" - parquet = "parquet" - sqlite3 = "sqlite3" - - -class LogLevel(StrEnum): - """Logging verbosity levels.""" - - DEBUG = "DEBUG" - INFO = "INFO" - WARNING = "WARNING" - ERROR = "ERROR" - - -class Dataset(StrEnum): - """Datasets supported by the ``collect-history`` command.""" - - rates = "rates" - ticks = "ticks" - history_orders = "history-orders" - history_deals = "history-deals" - - -class IfExists(StrEnum): - """SQLite table conflict behavior for the ``collect-history`` command.""" - - APPEND = "append" - REPLACE = "replace" - FAIL = "fail" - - -_DATASET_TABLE_NAMES: dict[Dataset, str] = { - Dataset.rates: "rates", - Dataset.ticks: "ticks", - Dataset.history_orders: "history_orders", - Dataset.history_deals: "history_deals", -} - - -# --------------------------------------------------------------------------- -# Click parameter types -# --------------------------------------------------------------------------- - - -class _DateTimeType(click.ParamType): - """Click parameter type for ISO 8601 datetime strings.""" - - name = "DATETIME" - - def convert( - self, - value: object, - param: click.Parameter | None, - ctx: click.Context | None, - ) -> datetime: - """Convert a string value to a timezone-aware datetime. - - Args: - value: Raw value from the command line. - param: Click parameter instance. - ctx: Click context. - - Returns: - Parsed datetime. - """ - if isinstance(value, datetime): - return value - try: - return parse_datetime(str(value)) - except ValueError as exc: - self.fail(str(exc), param, ctx) - - -class _TimeframeType(click.ParamType): - """Click parameter type for MT5 timeframe values.""" - - name = "TIMEFRAME" - - def convert( - self, - value: object, - param: click.Parameter | None, - ctx: click.Context | None, - ) -> int: - """Convert a string or integer value to a timeframe integer. - - Args: - value: Raw value from the command line. - param: Click parameter instance. - ctx: Click context. - - Returns: - Integer timeframe value. - """ - if isinstance(value, int): - return value - try: - return parse_timeframe(str(value)) - except ValueError as exc: - self.fail(str(exc), param, ctx) - - -class _TickFlagsType(click.ParamType): - """Click parameter type for MT5 tick copy flags.""" - - name = "FLAGS" - - def convert( - self, - value: object, - param: click.Parameter | None, - ctx: click.Context | None, - ) -> int: - """Convert a string or integer value to a tick flags integer. - - Args: - value: Raw value from the command line. - param: Click parameter instance. - ctx: Click context. - - Returns: - Integer tick flag value. - """ - if isinstance(value, int): - return value - try: - return parse_tick_flags(str(value)) - except ValueError as exc: - self.fail(str(exc), param, ctx) - - -class _RequestType(click.ParamType): - """Click parameter type for JSON order requests.""" - - name = "REQUEST" - - def convert( - self, - value: object, - param: click.Parameter | None, - ctx: click.Context | None, - ) -> dict[str, Any]: - """Convert a raw CLI value to an order request dictionary. - - Args: - value: Raw value from the command line. - param: Click parameter instance. - ctx: Click context. - - Returns: - Parsed request dictionary. - """ - try: - return parse_request(str(value)) - except ValueError as exc: - self.fail(str(exc), param, ctx) - - -DATETIME_TYPE = _DateTimeType() -TIMEFRAME_TYPE = _TimeframeType() -TICK_FLAGS_TYPE = _TickFlagsType() -REQUEST_TYPE = _RequestType() - # --------------------------------------------------------------------------- # Export context # --------------------------------------------------------------------------- @@ -266,186 +47,6 @@ class _ExportContext: config: Mt5Config -# --------------------------------------------------------------------------- -# Public utility functions -# --------------------------------------------------------------------------- - - -def detect_format( - output_path: Path, - explicit_format: str | None = None, -) -> str: - """Detect the output format from a file extension or explicit format string. - - Args: - output_path: Path to the output file. - explicit_format: Explicitly specified format, if any. - - Returns: - The detected format string. - - Raises: - ValueError: If the format cannot be determined. - """ - if explicit_format is not None: - return explicit_format - suffix = output_path.suffix.lower() - if suffix in _FORMAT_EXTENSIONS: - return _FORMAT_EXTENSIONS[suffix] - msg = ( - f"Cannot detect format from extension '{suffix}'." - " Use --format to specify the output format." - ) - raise ValueError(msg) - - -def export_dataframe( - df: pd.DataFrame, - output_path: Path, - output_format: str, - table_name: str = "data", -) -> None: - """Export a pandas DataFrame to the specified file format. - - Args: - df: DataFrame to export. - output_path: Path to the output file. - output_format: Output format (csv, json, parquet, or sqlite3). - table_name: Table name for SQLite3 output. - - Raises: - ValueError: If the output format is not supported. - """ - if output_format == "csv": - df.to_csv(output_path, index=False) - elif output_format == "json": - df.to_json( - output_path, - orient="records", - date_format="iso", - indent=2, - ) - elif output_format == "parquet": - df.to_parquet(output_path, index=False) - elif output_format == "sqlite3": - with sqlite3.connect(output_path) as conn: - df.to_sql( # type: ignore[reportUnknownMemberType] - table_name, - conn, - if_exists="replace", - index=False, - ) - else: - msg = f"Unsupported output format: {output_format}" - raise ValueError(msg) - - -def parse_datetime(value: str) -> datetime: - """Parse an ISO 8601 datetime string to a timezone-aware datetime. - - Args: - value: ISO 8601 datetime string (e.g., '2024-01-01' or - '2024-01-01T12:00:00+00:00'). - - Returns: - Parsed datetime with UTC timezone if no timezone is specified. - - Raises: - ValueError: If the string cannot be parsed. - """ - try: - dt = datetime.fromisoformat(value) - except ValueError: - msg = f"Invalid datetime format: '{value}'. Use ISO 8601 format." - raise ValueError(msg) from None - if dt.tzinfo is None: - dt = dt.replace(tzinfo=UTC) - return dt - - -def parse_timeframe(value: str) -> int: - """Parse a timeframe string or integer value. - - Args: - value: Timeframe name (e.g., 'M1', 'H1', 'D1') or integer value. - - Returns: - Integer timeframe value. - - Raises: - ValueError: If the timeframe is invalid. - """ - upper = value.upper() - if upper in TIMEFRAME_MAP: - return TIMEFRAME_MAP[upper] - try: - return int(value) - except ValueError: - valid = ", ".join(TIMEFRAME_MAP) - msg = f"Invalid timeframe: '{value}'. Use one of: {valid}, or an integer." - raise ValueError(msg) from None - - -def parse_tick_flags(value: str) -> int: - """Parse tick flags string or integer value. - - Args: - value: Tick flag name (ALL, INFO, TRADE) or integer value. - - Returns: - Integer tick flag value. - - Raises: - ValueError: If the flag is invalid. - """ - upper = value.upper() - if upper in TICK_FLAG_MAP: - return TICK_FLAG_MAP[upper] - try: - return int(value) - except ValueError: - valid = ", ".join(TICK_FLAG_MAP) - msg = f"Invalid tick flags: '{value}'. Use one of: {valid}, or an integer." - raise ValueError(msg) from None - - -def _is_request_dict(value: object) -> TypeGuard[dict[str, Any]]: - return isinstance(value, dict) - - -def parse_request(value: str) -> dict[str, Any]: - """Parse a JSON-formatted order request string or file reference. - - Args: - value: JSON object string, or '@path' to read JSON from a file. - - Returns: - Parsed request dictionary. - - Raises: - ValueError: If the request file cannot be read or the value is not a - JSON object. - """ - if value.startswith("@"): - path = Path(value[1:]) - try: - text = path.read_text(encoding="utf-8") - except (OSError, UnicodeDecodeError) as exc: - msg = f"Failed to read JSON request file '{path}': {exc}" - raise ValueError(msg) from exc - else: - text = value - try: - parsed: object = json.loads(text) - except json.JSONDecodeError as exc: - msg = f"Invalid JSON request: {exc}" - raise ValueError(msg) from exc - if not _is_request_dict(parsed): - msg = "Order request must be a JSON object." - raise ValueError(msg) - return parsed - - # --------------------------------------------------------------------------- # Typer application # --------------------------------------------------------------------------- @@ -466,34 +67,33 @@ def _get_export_context(ctx: typer.Context) -> _ExportContext: def _execute_export( ctx: typer.Context, - fetch_fn: Callable[[Mt5DataClient], pd.DataFrame], + fetch_fn: Callable[[], pd.DataFrame], ) -> None: - """Execute the common connect-fetch-export-shutdown workflow. + """Execute the common fetch-export workflow. Args: ctx: Typer context carrying shared options. - fetch_fn: Callable that receives a connected client and returns a - DataFrame. + fetch_fn: Callable that returns a DataFrame via the SDK layer. """ export_ctx = _get_export_context(ctx) - client = Mt5DataClient(config=export_ctx.config) - client.initialize_and_login_mt5() - try: - df = fetch_fn(client) - export_dataframe( - df=df, - output_path=export_ctx.output, - output_format=export_ctx.output_format, - table_name=export_ctx.table, - ) - logger.info( - "Exported %d rows to %s (%s)", - len(df), - export_ctx.output, - export_ctx.output_format, - ) - finally: - client.shutdown() + df = fetch_fn() + export_dataframe( + df=df, + output_path=export_ctx.output, + output_format=export_ctx.output_format, + table_name=export_ctx.table, + ) + logger.info( + "Exported %d rows to %s (%s)", + len(df), + export_ctx.output, + export_ctx.output_format, + ) + + +def _sdk_client(ctx: typer.Context) -> sdk.Mt5CliClient: + export_ctx = _get_export_context(ctx) + return sdk.Mt5CliClient(config=export_ctx.config) @app.callback() @@ -593,14 +193,10 @@ def rates_from( count: Annotated[int, typer.Option(help="Number of records.")], ) -> None: """Export rates from a start date.""" + client = _sdk_client(ctx) _execute_export( ctx, - lambda c: c.copy_rates_from_as_df( - symbol=symbol, - timeframe=timeframe, - date_from=date_from, - count=count, - ), + lambda: client.copy_rates_from(symbol, timeframe, date_from, count), ) @@ -619,14 +215,10 @@ def rates_from_pos( count: Annotated[int, typer.Option(help="Number of records.")], ) -> None: """Export rates from a start position.""" + client = _sdk_client(ctx) _execute_export( ctx, - lambda c: c.copy_rates_from_pos_as_df( - symbol=symbol, - timeframe=timeframe, - start_pos=start_pos, - count=count, - ), + lambda: client.copy_rates_from_pos(symbol, timeframe, start_pos, count), ) @@ -651,14 +243,10 @@ def rates_range( ], ) -> None: """Export rates for a date range.""" + client = _sdk_client(ctx) _execute_export( ctx, - lambda c: c.copy_rates_range_as_df( - symbol=symbol, - timeframe=timeframe, - date_from=date_from, - date_to=date_to, - ), + lambda: client.copy_rates_range(symbol, timeframe, date_from, date_to), ) @@ -680,14 +268,10 @@ def ticks_from( ], ) -> None: """Export ticks from a start date.""" + client = _sdk_client(ctx) _execute_export( ctx, - lambda c: c.copy_ticks_from_as_df( - symbol=symbol, - date_from=date_from, - count=count, - flags=flags, - ), + lambda: client.copy_ticks_from(symbol, date_from, count, flags), ) @@ -709,27 +293,23 @@ def ticks_range( ], ) -> None: """Export ticks for a date range.""" + client = _sdk_client(ctx) _execute_export( ctx, - lambda c: c.copy_ticks_range_as_df( - symbol=symbol, - date_from=date_from, - date_to=date_to, - flags=flags, - ), + lambda: client.copy_ticks_range(symbol, date_from, date_to, flags), ) @app.command() def account_info(ctx: typer.Context) -> None: """Export account information.""" - _execute_export(ctx, lambda c: c.account_info_as_df()) + _execute_export(ctx, _sdk_client(ctx).account_info) @app.command() def terminal_info(ctx: typer.Context) -> None: """Export terminal information.""" - _execute_export(ctx, lambda c: c.terminal_info_as_df()) + _execute_export(ctx, _sdk_client(ctx).terminal_info) @app.command() @@ -741,10 +321,8 @@ def symbols( ] = None, ) -> None: """Export symbol list.""" - _execute_export( - ctx, - lambda c: c.symbols_get_as_df(group=group), - ) + client = _sdk_client(ctx) + _execute_export(ctx, lambda: client.symbols(group=group)) @app.command() @@ -753,10 +331,8 @@ def symbol_info( symbol: Annotated[str, typer.Option(help="Symbol name.")], ) -> None: """Export symbol details.""" - _execute_export( - ctx, - lambda c: c.symbol_info_as_df(symbol=symbol), - ) + client = _sdk_client(ctx) + _execute_export(ctx, lambda: client.symbol_info(symbol)) @app.command() @@ -767,13 +343,10 @@ def orders( ticket: Annotated[int | None, typer.Option(help="Ticket filter.")] = None, ) -> None: """Export active orders.""" + client = _sdk_client(ctx) _execute_export( ctx, - lambda c: c.orders_get_as_df( - symbol=symbol, - group=group, - ticket=ticket, - ), + lambda: client.orders(symbol=symbol, group=group, ticket=ticket), ) @@ -785,13 +358,10 @@ def positions( ticket: Annotated[int | None, typer.Option(help="Ticket filter.")] = None, ) -> None: """Export open positions.""" + client = _sdk_client(ctx) _execute_export( ctx, - lambda c: c.positions_get_as_df( - symbol=symbol, - group=group, - ticket=ticket, - ), + lambda: client.positions(symbol=symbol, group=group, ticket=ticket), ) @@ -812,9 +382,10 @@ def history_orders( position: Annotated[int | None, typer.Option(help="Position ticket.")] = None, ) -> None: """Export historical orders.""" + client = _sdk_client(ctx) _execute_export( ctx, - lambda c: c.history_orders_get_as_df( + lambda: client.history_orders( date_from=date_from, date_to=date_to, group=group, @@ -842,9 +413,10 @@ def history_deals( position: Annotated[int | None, typer.Option(help="Position ticket.")] = None, ) -> None: """Export historical deals.""" + client = _sdk_client(ctx) _execute_export( ctx, - lambda c: c.history_deals_get_as_df( + lambda: client.history_deals( date_from=date_from, date_to=date_to, group=group, @@ -858,13 +430,13 @@ def history_deals( @app.command() def version(ctx: typer.Context) -> None: """Export MetaTrader5 version information.""" - _execute_export(ctx, lambda c: c.version_as_df()) + _execute_export(ctx, _sdk_client(ctx).version) @app.command() def last_error(ctx: typer.Context) -> None: """Export the last error information.""" - _execute_export(ctx, lambda c: c.last_error_as_df()) + _execute_export(ctx, _sdk_client(ctx).last_error) @app.command() @@ -873,10 +445,8 @@ def symbol_info_tick( symbol: Annotated[str, typer.Option(help="Symbol name.")], ) -> None: """Export the last tick for a symbol.""" - _execute_export( - ctx, - lambda c: c.symbol_info_tick_as_df(symbol=symbol), - ) + client = _sdk_client(ctx) + _execute_export(ctx, lambda: client.symbol_info_tick(symbol)) @app.command() @@ -885,10 +455,8 @@ def market_book( symbol: Annotated[str, typer.Option(help="Symbol name.")], ) -> None: """Export market depth (order book) for a symbol.""" - _execute_export( - ctx, - lambda c: c.market_book_get_as_df(symbol=symbol), - ) + client = _sdk_client(ctx) + _execute_export(ctx, lambda: client.market_book(symbol)) @app.command() @@ -900,10 +468,15 @@ def order_check( ], ) -> None: """Check funds sufficiency for a trading operation.""" - _execute_export( - ctx, - lambda c: c.order_check_as_df(request=request), - ) + export_ctx = _get_export_context(ctx) + + def _fetch() -> pd.DataFrame: + return sdk._run_with_client( # noqa: SLF001 # pyright: ignore[reportPrivateUsage] + export_ctx.config, + lambda c: c.order_check_as_df(request=request), + ) + + _execute_export(ctx, _fetch) @app.command() @@ -926,404 +499,15 @@ def order_send( if not yes: msg = "Pass --yes to send a live trade request." raise typer.BadParameter(msg, param_hint="--yes") - _execute_export( - ctx, - lambda c: c.order_send_as_df(request=request), - ) + export_ctx = _get_export_context(ctx) - -def _create_cash_events_view( - conn: sqlite3.Connection, - deals_columns: set[str], -) -> bool: - """Create the cash_events SQLite view derived from history_deals. - - Args: - conn: Open SQLite connection. - deals_columns: Column names present in the history_deals table. - - Returns: - True if the view was created, False if required columns are missing. - """ - if "type" not in deals_columns: - logger.warning("Skipping cash_events view: history_deals.type is missing") - return False - conn.execute("DROP VIEW IF EXISTS cash_events") - conn.execute( - "CREATE VIEW cash_events AS" # noqa: S608 - f" SELECT * FROM history_deals WHERE type NOT IN {_TRADE_DEAL_TYPES_SQL}", - ) - return True - - -def _create_positions_reconstructed_view( - conn: sqlite3.Connection, - deals_columns: set[str], -) -> bool: - """Create the positions_reconstructed SQLite view derived from history_deals. - - The view aggregates trade deals (``type IN (0, 1)``) by ``position_id`` and - excludes positions that have no closing deal (``entry IN (1, 3)``), so - still-open positions and reversal-only fragments are filtered out. - - Open/close prices are volume-weighted averages over the corresponding - entry deals. Reversal deals (``DEAL_ENTRY_INOUT = 2``) are reported via - ``volume_reversal`` and ``reversal_count``; they do not contribute to the - open or close volume/price weights because a single reversal deal mixes a - close of the existing direction with the open of the new direction. - - Args: - conn: Open SQLite connection. - deals_columns: Column names present in the history_deals table. - - Returns: - True if the view was created, False if required columns are missing. - """ - if not _POSITIONS_VIEW_REQUIRED_COLUMNS.issubset(deals_columns): - missing = ", ".join(sorted(_POSITIONS_VIEW_REQUIRED_COLUMNS - deals_columns)) - logger.warning( - "Skipping positions_reconstructed view: history_deals missing columns: %s", - missing, - ) - return False - conn.execute("DROP VIEW IF EXISTS positions_reconstructed") - conn.execute( - "CREATE VIEW positions_reconstructed AS" # noqa: S608 - " SELECT" - " position_id," - " symbol," - " MIN(CASE WHEN entry = 0 THEN time END) AS open_time," - " MAX(CASE WHEN entry IN (1, 2, 3) THEN time END) AS close_time," - " MIN(CASE WHEN entry = 0 THEN type END) AS direction," - " SUM(CASE WHEN entry = 0 THEN volume ELSE 0 END) AS volume_open," - " SUM(CASE WHEN entry IN (1, 3) THEN volume ELSE 0 END) AS volume_close," - " SUM(CASE WHEN entry = 2 THEN volume ELSE 0 END) AS volume_reversal," - " CASE" - " WHEN SUM(CASE WHEN entry = 0 THEN volume ELSE 0 END) > 0" - " THEN SUM(CASE WHEN entry = 0 THEN price * volume ELSE 0 END)" - " / SUM(CASE WHEN entry = 0 THEN volume ELSE 0 END)" - " END AS open_price," - " CASE" - " WHEN SUM(CASE WHEN entry IN (1, 3) THEN volume ELSE 0 END) > 0" - " THEN SUM(CASE WHEN entry IN (1, 3) THEN price * volume ELSE 0 END)" - " / SUM(CASE WHEN entry IN (1, 3) THEN volume ELSE 0 END)" - " END AS close_price," - " SUM(profit) AS total_profit," - " SUM(CASE WHEN entry = 2 THEN 1 ELSE 0 END) AS reversal_count," - " COUNT(*) AS deals_count" - " FROM history_deals" - f" WHERE type IN {_TRADE_DEAL_TYPES_SQL} AND position_id != 0" - " GROUP BY position_id, symbol" - " HAVING SUM(CASE WHEN entry IN (1, 3) THEN 1 ELSE 0 END) > 0", - ) - return True - - -def _write_frame_to_sqlite( - conn: sqlite3.Connection, - frame: pd.DataFrame, - table_name: str, - if_exists: IfExists, -) -> bool: - """Write a non-empty-schema frame to SQLite. - - Args: - conn: Open SQLite connection. - frame: DataFrame to write. - table_name: Target SQLite table name. - if_exists: Table conflict behavior. - - Returns: - True if a table was written, False if the frame had no columns. - """ - if len(frame.columns) == 0: - logger.warning("Skipping %s: dataset returned no columns", table_name) - return False - frame.to_sql( # type: ignore[reportUnknownMemberType] - table_name, - conn, - if_exists=if_exists.value, - index=False, - chunksize=50_000, - method="multi", - ) - return True - - -def _create_collect_history_indexes( - conn: sqlite3.Connection, - written_columns: dict[Dataset, set[str]], -) -> None: - """Create useful indexes for collected history tables when present.""" - if {"symbol", "time"}.issubset(written_columns.get(Dataset.rates, set())): - conn.execute( - "CREATE INDEX IF NOT EXISTS idx_rates_symbol_time ON rates(symbol, time)", - ) - if {"symbol", "time"}.issubset(written_columns.get(Dataset.ticks, set())): - conn.execute( - "CREATE INDEX IF NOT EXISTS idx_ticks_symbol_time ON ticks(symbol, time)", - ) - if {"position_id", "symbol"}.issubset( - written_columns.get(Dataset.history_deals, set()) - ): - conn.execute( - "CREATE INDEX IF NOT EXISTS idx_history_deals_position_symbol" - " ON history_deals(position_id, symbol)", + def _fetch() -> pd.DataFrame: + return sdk._run_with_client( # noqa: SLF001 # pyright: ignore[reportPrivateUsage] + export_ctx.config, + lambda c: c.order_send_as_df(request=request), ) - -def _record_written_columns( - written_columns: dict[Dataset, set[str]], - dataset: Dataset, - frame: pd.DataFrame, -) -> None: - """Remember columns for datasets written during streaming collection.""" - columns = set(frame.columns) - if dataset in written_columns: - written_columns[dataset].update(columns) - else: - written_columns[dataset] = columns - - -def _write_streamed_frame( - conn: sqlite3.Connection, - frame: pd.DataFrame, - dataset: Dataset, - table_exists: bool, - if_exists: IfExists, - written_columns: dict[Dataset, set[str]], -) -> bool: - """Write one streamed dataset frame and track table state. - - Args: - conn: Open SQLite connection. - frame: DataFrame to write. - dataset: Dataset being written. - table_exists: Whether this dataset table has already been written. - if_exists: Initial table conflict behavior. - written_columns: Mutable map of columns written by dataset. - - Returns: - True if the dataset table exists after this write attempt. - """ - write_mode = IfExists.APPEND if table_exists else if_exists - if _write_frame_to_sqlite( - conn, - frame, - _DATASET_TABLE_NAMES[dataset], - write_mode, - ): - _record_written_columns(written_columns, dataset, frame) - return True - return table_exists - - -def _write_rates_dataset( - conn: sqlite3.Connection, - client: Mt5DataClient, - symbols: list[str], - timeframe: int, - date_from: datetime, - date_to: datetime, - if_exists: IfExists, - written_columns: dict[Dataset, set[str]], -) -> bool: - """Stream rates frames into SQLite. - - Args: - conn: Open SQLite connection. - client: Connected MT5 data client. - symbols: Symbols to collect. - timeframe: Rates timeframe integer. - date_from: Start date. - date_to: End date. - if_exists: Initial table conflict behavior. - written_columns: Mutable map of columns written by dataset. - - Returns: - True if the rates table was written. - """ - table_exists = False - for sym in symbols: - frame = client.copy_rates_range_as_df( - symbol=sym, - timeframe=timeframe, - date_from=date_from, - date_to=date_to, - ) - frame.insert(0, "symbol", sym) - frame.insert(1, "timeframe", timeframe) - table_exists = _write_streamed_frame( - conn, - frame, - Dataset.rates, - table_exists, - if_exists, - written_columns, - ) - return table_exists - - -def _write_ticks_dataset( - conn: sqlite3.Connection, - client: Mt5DataClient, - symbols: list[str], - flags: int, - date_from: datetime, - date_to: datetime, - if_exists: IfExists, - written_columns: dict[Dataset, set[str]], -) -> bool: - """Stream ticks frames into SQLite. - - Args: - conn: Open SQLite connection. - client: Connected MT5 data client. - symbols: Symbols to collect. - flags: Tick copy flags integer. - date_from: Start date. - date_to: End date. - if_exists: Initial table conflict behavior. - written_columns: Mutable map of columns written by dataset. - - Returns: - True if the ticks table was written. - """ - table_exists = False - for sym in symbols: - frame = client.copy_ticks_range_as_df( - symbol=sym, - date_from=date_from, - date_to=date_to, - flags=flags, - ) - frame.insert(0, "symbol", sym) - table_exists = _write_streamed_frame( - conn, - frame, - Dataset.ticks, - table_exists, - if_exists, - written_columns, - ) - return table_exists - - -def _write_history_dataset( - conn: sqlite3.Connection, - fetch: Callable[..., pd.DataFrame], - dataset: Dataset, - symbols: list[str], - date_from: datetime, - date_to: datetime, - if_exists: IfExists, - written_columns: dict[Dataset, set[str]], -) -> bool: - """Stream a history dataset into SQLite with exact symbol filtering. - - Args: - conn: Open SQLite connection. - fetch: Bound history_orders_get_as_df / history_deals_get_as_df method. - dataset: History dataset being written. - symbols: Symbols to collect. - date_from: Start date. - date_to: End date. - if_exists: Initial table conflict behavior. - written_columns: Mutable map of columns written by dataset. - - Returns: - True if the history table was written. - """ - table_exists = False - for sym in symbols: - frame = fetch(date_from=date_from, date_to=date_to, symbol=sym) - if "symbol" in frame.columns: - frame = frame[frame["symbol"] == sym] - table_exists = _write_streamed_frame( - conn, - frame, - dataset, - table_exists, - if_exists, - written_columns, - ) - return table_exists - - -def _write_collected_datasets( - conn: sqlite3.Connection, - client: Mt5DataClient, - symbols: list[str], - datasets: set[Dataset], - timeframe: int, - flags: int, - date_from: datetime, - date_to: datetime, - if_exists: IfExists, -) -> tuple[set[Dataset], dict[Dataset, set[str]]]: - """Collect selected datasets and stream each symbol frame into SQLite. - - Args: - conn: Open SQLite connection. - client: Connected MT5 data client. - symbols: Symbols to collect. - datasets: Selected datasets to write. - timeframe: Rates timeframe integer. - flags: Tick copy flags integer. - date_from: Start date. - date_to: End date. - if_exists: Initial table conflict behavior. - - Returns: - Written datasets and their columns. - """ - written_columns: dict[Dataset, set[str]] = {} - written_tables: set[Dataset] = set() - if Dataset.rates in datasets and _write_rates_dataset( - conn, - client, - symbols, - timeframe, - date_from, - date_to, - if_exists, - written_columns, - ): - written_tables.add(Dataset.rates) - if Dataset.ticks in datasets and _write_ticks_dataset( - conn, - client, - symbols, - flags, - date_from, - date_to, - if_exists, - written_columns, - ): - written_tables.add(Dataset.ticks) - if Dataset.history_orders in datasets and _write_history_dataset( - conn, - client.history_orders_get_as_df, - Dataset.history_orders, - symbols, - date_from, - date_to, - if_exists, - written_columns, - ): - written_tables.add(Dataset.history_orders) - if Dataset.history_deals in datasets and _write_history_dataset( - conn, - client.history_deals_get_as_df, - Dataset.history_deals, - symbols, - date_from, - date_to, - if_exists, - written_columns, - ): - written_tables.add(Dataset.history_deals) - return written_tables, written_columns + _execute_export(ctx, _fetch) @app.command() @@ -1409,42 +593,18 @@ def collect_history( ) raise typer.BadParameter(msg) datasets = set(dataset) if dataset else set(Dataset) - client = Mt5DataClient(config=export_ctx.config) - client.initialize_and_login_mt5() - try: - with sqlite3.connect(export_ctx.output) as conn: - conn.execute("PRAGMA journal_mode=WAL") - conn.execute("PRAGMA synchronous=NORMAL") - written_tables, written_columns = _write_collected_datasets( - conn, - client, - symbol, - datasets, - timeframe, - flags, - date_from, - date_to, - if_exists, - ) - _create_collect_history_indexes(conn, written_columns) - if with_views and Dataset.history_deals in written_tables: - _create_cash_events_view(conn, written_columns[Dataset.history_deals]) - _create_positions_reconstructed_view( - conn, - written_columns[Dataset.history_deals], - ) - elif with_views: - logger.warning( - "--with-views ignored: history_deals table was not written" - ) - logger.info( - "Collected %s for %d symbol(s) into %s", - ", ".join(sorted(ds.value for ds in datasets)), - len(symbol), - export_ctx.output, - ) - finally: - client.shutdown() + sdk.collect_history( + output=export_ctx.output, + symbols=symbol, + date_from=date_from, + date_to=date_to, + datasets=datasets, + timeframe=timeframe, + flags=flags, + if_exists=if_exists, + with_views=with_views, + config=export_ctx.config, + ) def main() -> None: diff --git a/mt5cli/sdk.py b/mt5cli/sdk.py new file mode 100644 index 0000000..28ada99 --- /dev/null +++ b/mt5cli/sdk.py @@ -0,0 +1,1023 @@ +"""Programmatic SDK for MetaTrader 5 data collection.""" + +from __future__ import annotations + +import logging +import sqlite3 +from contextlib import contextmanager +from datetime import datetime +from pathlib import Path # noqa: TC003 +from typing import TYPE_CHECKING, Self, TypeVar + +from pdmt5 import Mt5Config, Mt5DataClient + +from .utils import ( + Dataset, + IfExists, + parse_datetime, + parse_tick_flags, + parse_timeframe, +) + +if TYPE_CHECKING: + from collections.abc import Callable, Iterator + + import pandas as pd + +T = TypeVar("T") + +logger = logging.getLogger(__name__) + +__all__ = [ + "Mt5CliClient", + "account_info", + "build_config", + "collect_history", + "copy_rates_from", + "copy_rates_from_pos", + "copy_rates_range", + "copy_ticks_from", + "copy_ticks_range", + "history_deals", + "history_orders", + "last_error", + "market_book", + "orders", + "positions", + "symbol_info", + "symbol_info_tick", + "symbols", + "terminal_info", + "version", +] + +_TRADE_DEAL_TYPES: tuple[int, int] = (0, 1) +_TRADE_DEAL_TYPES_SQL = f"({', '.join(str(value) for value in _TRADE_DEAL_TYPES)})" +_POSITIONS_VIEW_REQUIRED_COLUMNS: frozenset[str] = frozenset({ + "position_id", + "symbol", + "time", + "type", + "entry", + "volume", + "price", + "profit", +}) + + +def _coerce_timeframe(timeframe: int | str) -> int: + if isinstance(timeframe, int): + return timeframe + return parse_timeframe(timeframe) + + +def _coerce_tick_flags(flags: int | str) -> int: + if isinstance(flags, int): + return flags + return parse_tick_flags(flags) + + +def _require_datetime(value: datetime | str) -> datetime: + if isinstance(value, datetime): + return value + return parse_datetime(value) + + +def _coerce_datetime(value: datetime | str | None) -> datetime | None: + if value is None or isinstance(value, datetime): + return value + return parse_datetime(value) + + +def build_config( + *, + path: str | None = None, + login: int | None = None, + password: str | None = None, + server: str | None = None, + timeout: int | None = None, +) -> Mt5Config: + """Build an ``Mt5Config`` from optional connection parameters. + + Returns: + Configured ``Mt5Config`` instance. + """ + return Mt5Config( + path=path, + login=login, + password=password, + server=server, + timeout=timeout, + ) + + +@contextmanager +def _connected_client(config: Mt5Config) -> Iterator[Mt5DataClient]: + """Initialize MT5, yield a connected client, and always shut down. + + Args: + config: MT5 connection configuration. + + Yields: + Connected ``Mt5DataClient`` instance. + """ + client = Mt5DataClient(config=config) + try: + client.initialize_and_login_mt5() + yield client + finally: + client.shutdown() + + +def _run_with_client( + config: Mt5Config, + fetch_fn: Callable[[Mt5DataClient], T], +) -> T: + """Connect, run ``fetch_fn``, and shut down safely. + + Args: + config: MT5 connection configuration. + fetch_fn: Callable receiving a connected client. + + Returns: + Value returned by ``fetch_fn``. + """ + with _connected_client(config) as client: + return fetch_fn(client) + + +class Mt5CliClient: + """Programmatic client for read-only MetaTrader 5 data access.""" + + def __init__( + self, + *, + path: str | None = None, + login: int | None = None, + password: str | None = None, + server: str | None = None, + timeout: int | None = None, + config: Mt5Config | None = None, + ) -> None: + """Initialize the SDK client. + + Args: + path: Path to MetaTrader5 terminal EXE file. + login: Trading account login. + password: Trading account password. + server: Trading server name. + timeout: Connection timeout in milliseconds. + config: Optional pre-built ``Mt5Config`` (overrides other args). + """ + self._config = config or build_config( + path=path, + login=login, + password=password, + server=server, + timeout=timeout, + ) + self._client: Mt5DataClient | None = None + + @property + def config(self) -> Mt5Config: + """Return the underlying MT5 configuration.""" + return self._config + + def __enter__(self) -> Self: + """Open a persistent MT5 connection for multiple calls. + + Returns: + This client instance. + """ + client = Mt5DataClient(config=self._config) + try: + client.initialize_and_login_mt5() + except Exception: + client.shutdown() + raise + self._client = client + return self + + def __exit__( + self, + exc_type: type[BaseException] | None, + exc: BaseException | None, + tb: object, + ) -> None: + """Shut down the persistent MT5 connection.""" + if self._client is not None: + self._client.shutdown() + self._client = None + + def _fetch(self, fetch_fn: Callable[[Mt5DataClient], pd.DataFrame]) -> pd.DataFrame: + if self._client is not None: + return fetch_fn(self._client) + return _run_with_client(self._config, fetch_fn) + + def copy_rates_from( + self, + symbol: str, + timeframe: int | str, + date_from: datetime | str, + count: int, + ) -> pd.DataFrame: + """Return rates starting from a date.""" + tf = _coerce_timeframe(timeframe) + start = _require_datetime(date_from) + return self._fetch( + lambda c: c.copy_rates_from_as_df( + symbol=symbol, + timeframe=tf, + date_from=start, + count=count, + ), + ) + + def copy_rates_from_pos( + self, + symbol: str, + timeframe: int | str, + start_pos: int, + count: int, + ) -> pd.DataFrame: + """Return rates starting from a bar position.""" + tf = _coerce_timeframe(timeframe) + return self._fetch( + lambda c: c.copy_rates_from_pos_as_df( + symbol=symbol, + timeframe=tf, + start_pos=start_pos, + count=count, + ), + ) + + def copy_rates_range( + self, + symbol: str, + timeframe: int | str, + date_from: datetime | str, + date_to: datetime | str, + ) -> pd.DataFrame: + """Return rates for a date range.""" + tf = _coerce_timeframe(timeframe) + start = _require_datetime(date_from) + end = _require_datetime(date_to) + return self._fetch( + lambda c: c.copy_rates_range_as_df( + symbol=symbol, + timeframe=tf, + date_from=start, + date_to=end, + ), + ) + + def copy_ticks_from( + self, + symbol: str, + date_from: datetime | str, + count: int, + flags: int | str, + ) -> pd.DataFrame: + """Return ticks starting from a date.""" + start = _require_datetime(date_from) + tick_flags = _coerce_tick_flags(flags) + return self._fetch( + lambda c: c.copy_ticks_from_as_df( + symbol=symbol, + date_from=start, + count=count, + flags=tick_flags, + ), + ) + + def copy_ticks_range( + self, + symbol: str, + date_from: datetime | str, + date_to: datetime | str, + flags: int | str, + ) -> pd.DataFrame: + """Return ticks for a date range.""" + start = _require_datetime(date_from) + end = _require_datetime(date_to) + tick_flags = _coerce_tick_flags(flags) + return self._fetch( + lambda c: c.copy_ticks_range_as_df( + symbol=symbol, + date_from=start, + date_to=end, + flags=tick_flags, + ), + ) + + def account_info(self) -> pd.DataFrame: + """Return account information.""" + return self._fetch(lambda c: c.account_info_as_df()) + + def terminal_info(self) -> pd.DataFrame: + """Return terminal information.""" + return self._fetch(lambda c: c.terminal_info_as_df()) + + def symbols(self, group: str | None = None) -> pd.DataFrame: + """Return the symbol list.""" + return self._fetch(lambda c: c.symbols_get_as_df(group=group)) + + def symbol_info(self, symbol: str) -> pd.DataFrame: + """Return details for one symbol.""" + return self._fetch(lambda c: c.symbol_info_as_df(symbol=symbol)) + + def orders( + self, + symbol: str | None = None, + group: str | None = None, + ticket: int | None = None, + ) -> pd.DataFrame: + """Return active orders.""" + return self._fetch( + lambda c: c.orders_get_as_df( + symbol=symbol, + group=group, + ticket=ticket, + ), + ) + + def positions( + self, + symbol: str | None = None, + group: str | None = None, + ticket: int | None = None, + ) -> pd.DataFrame: + """Return open positions.""" + return self._fetch( + lambda c: c.positions_get_as_df( + symbol=symbol, + group=group, + ticket=ticket, + ), + ) + + def history_orders( + self, + date_from: datetime | str | None = None, + date_to: datetime | str | None = None, + group: str | None = None, + symbol: str | None = None, + ticket: int | None = None, + position: int | None = None, + ) -> pd.DataFrame: + """Return historical orders.""" + start = _coerce_datetime(date_from) + end = _coerce_datetime(date_to) + return self._fetch( + lambda c: c.history_orders_get_as_df( + date_from=start, + date_to=end, + group=group, + symbol=symbol, + ticket=ticket, + position=position, + ), + ) + + def history_deals( + self, + date_from: datetime | str | None = None, + date_to: datetime | str | None = None, + group: str | None = None, + symbol: str | None = None, + ticket: int | None = None, + position: int | None = None, + ) -> pd.DataFrame: + """Return historical deals.""" + start = _coerce_datetime(date_from) + end = _coerce_datetime(date_to) + return self._fetch( + lambda c: c.history_deals_get_as_df( + date_from=start, + date_to=end, + group=group, + symbol=symbol, + ticket=ticket, + position=position, + ), + ) + + def version(self) -> pd.DataFrame: + """Return MetaTrader5 version information.""" + return self._fetch(lambda c: c.version_as_df()) + + def last_error(self) -> pd.DataFrame: + """Return the last error information.""" + return self._fetch(lambda c: c.last_error_as_df()) + + def symbol_info_tick(self, symbol: str) -> pd.DataFrame: + """Return the last tick for a symbol.""" + return self._fetch(lambda c: c.symbol_info_tick_as_df(symbol=symbol)) + + def market_book(self, symbol: str) -> pd.DataFrame: + """Return market depth for a symbol.""" + return self._fetch(lambda c: c.market_book_get_as_df(symbol=symbol)) + + +def _create_cash_events_view( + conn: sqlite3.Connection, + deals_columns: set[str], +) -> bool: + """Create the cash_events SQLite view derived from history_deals. + + Returns: + True if the view was created, False if required columns are missing. + """ + if "type" not in deals_columns: + logger.warning("Skipping cash_events view: history_deals.type is missing") + return False + conn.execute("DROP VIEW IF EXISTS cash_events") + conn.execute( + "CREATE VIEW cash_events AS" # noqa: S608 + f" SELECT * FROM history_deals WHERE type NOT IN {_TRADE_DEAL_TYPES_SQL}", + ) + return True + + +def _create_positions_reconstructed_view( + conn: sqlite3.Connection, + deals_columns: set[str], +) -> bool: + """Create the positions_reconstructed SQLite view derived from history_deals. + + Returns: + True if the view was created, False if required columns are missing. + """ + if not _POSITIONS_VIEW_REQUIRED_COLUMNS.issubset(deals_columns): + missing = ", ".join(sorted(_POSITIONS_VIEW_REQUIRED_COLUMNS - deals_columns)) + logger.warning( + "Skipping positions_reconstructed view: history_deals missing columns: %s", + missing, + ) + return False + conn.execute("DROP VIEW IF EXISTS positions_reconstructed") + conn.execute( + "CREATE VIEW positions_reconstructed AS" # noqa: S608 + " SELECT" + " position_id," + " symbol," + " MIN(CASE WHEN entry = 0 THEN time END) AS open_time," + " MAX(CASE WHEN entry IN (1, 2, 3) THEN time END) AS close_time," + " MIN(CASE WHEN entry = 0 THEN type END) AS direction," + " SUM(CASE WHEN entry = 0 THEN volume ELSE 0 END) AS volume_open," + " SUM(CASE WHEN entry IN (1, 3) THEN volume ELSE 0 END) AS volume_close," + " SUM(CASE WHEN entry = 2 THEN volume ELSE 0 END) AS volume_reversal," + " CASE" + " WHEN SUM(CASE WHEN entry = 0 THEN volume ELSE 0 END) > 0" + " THEN SUM(CASE WHEN entry = 0 THEN price * volume ELSE 0 END)" + " / SUM(CASE WHEN entry = 0 THEN volume ELSE 0 END)" + " END AS open_price," + " CASE" + " WHEN SUM(CASE WHEN entry IN (1, 3) THEN volume ELSE 0 END) > 0" + " THEN SUM(CASE WHEN entry IN (1, 3) THEN price * volume ELSE 0 END)" + " / SUM(CASE WHEN entry IN (1, 3) THEN volume ELSE 0 END)" + " END AS close_price," + " SUM(profit) AS total_profit," + " SUM(CASE WHEN entry = 2 THEN 1 ELSE 0 END) AS reversal_count," + " COUNT(*) AS deals_count" + " FROM history_deals" + f" WHERE type IN {_TRADE_DEAL_TYPES_SQL} AND position_id != 0" + " GROUP BY position_id, symbol" + " HAVING SUM(CASE WHEN entry IN (1, 3) THEN 1 ELSE 0 END) > 0", + ) + return True + + +def _write_frame_to_sqlite( + conn: sqlite3.Connection, + frame: pd.DataFrame, + table_name: str, + if_exists: IfExists, +) -> bool: + """Write a non-empty-schema frame to SQLite. + + Returns: + True if a table was written, False if the frame had no columns. + """ + if len(frame.columns) == 0: + logger.warning("Skipping %s: dataset returned no columns", table_name) + return False + frame.to_sql( # type: ignore[reportUnknownMemberType] + table_name, + conn, + if_exists=if_exists.value, + index=False, + chunksize=50_000, + method="multi", + ) + return True + + +def _create_collect_history_indexes( + conn: sqlite3.Connection, + written_columns: dict[Dataset, set[str]], +) -> None: + """Create useful indexes for collected history tables when present.""" + if {"symbol", "time"}.issubset(written_columns.get(Dataset.rates, set())): + conn.execute( + "CREATE INDEX IF NOT EXISTS idx_rates_symbol_time ON rates(symbol, time)", + ) + if {"symbol", "time"}.issubset(written_columns.get(Dataset.ticks, set())): + conn.execute( + "CREATE INDEX IF NOT EXISTS idx_ticks_symbol_time ON ticks(symbol, time)", + ) + if {"position_id", "symbol"}.issubset( + written_columns.get(Dataset.history_deals, set()) + ): + conn.execute( + "CREATE INDEX IF NOT EXISTS idx_history_deals_position_symbol" + " ON history_deals(position_id, symbol)", + ) + + +def _record_written_columns( + written_columns: dict[Dataset, set[str]], + dataset: Dataset, + frame: pd.DataFrame, +) -> None: + """Remember columns for datasets written during streaming collection.""" + columns = set(frame.columns) + if dataset in written_columns: + written_columns[dataset].update(columns) + else: + written_columns[dataset] = columns + + +def _write_streamed_frame( + conn: sqlite3.Connection, + frame: pd.DataFrame, + dataset: Dataset, + table_exists: bool, + if_exists: IfExists, + written_columns: dict[Dataset, set[str]], +) -> bool: + """Write one streamed dataset frame and track table state. + + Returns: + True if the dataset table exists after this write attempt. + """ + write_mode = IfExists.APPEND if table_exists else if_exists + if _write_frame_to_sqlite( + conn, + frame, + dataset.table_name, + write_mode, + ): + _record_written_columns(written_columns, dataset, frame) + return True + return table_exists + + +def _write_rates_dataset( + conn: sqlite3.Connection, + client: Mt5DataClient, + symbols: list[str], + timeframe: int, + date_from: datetime, + date_to: datetime, + if_exists: IfExists, + written_columns: dict[Dataset, set[str]], +) -> bool: + """Stream rates frames into SQLite. + + Returns: + True if the rates table was written. + """ + table_exists = False + for sym in symbols: + frame = client.copy_rates_range_as_df( + symbol=sym, + timeframe=timeframe, + date_from=date_from, + date_to=date_to, + ) + frame.insert(0, "symbol", sym) + frame.insert(1, "timeframe", timeframe) + table_exists = _write_streamed_frame( + conn, + frame, + Dataset.rates, + table_exists, + if_exists, + written_columns, + ) + return table_exists + + +def _write_ticks_dataset( + conn: sqlite3.Connection, + client: Mt5DataClient, + symbols: list[str], + flags: int, + date_from: datetime, + date_to: datetime, + if_exists: IfExists, + written_columns: dict[Dataset, set[str]], +) -> bool: + """Stream ticks frames into SQLite. + + Returns: + True if the ticks table was written. + """ + table_exists = False + for sym in symbols: + frame = client.copy_ticks_range_as_df( + symbol=sym, + date_from=date_from, + date_to=date_to, + flags=flags, + ) + frame.insert(0, "symbol", sym) + table_exists = _write_streamed_frame( + conn, + frame, + Dataset.ticks, + table_exists, + if_exists, + written_columns, + ) + return table_exists + + +def _write_history_dataset( + conn: sqlite3.Connection, + fetch: Callable[..., pd.DataFrame], + dataset: Dataset, + symbols: list[str], + date_from: datetime, + date_to: datetime, + if_exists: IfExists, + written_columns: dict[Dataset, set[str]], +) -> bool: + """Stream a history dataset into SQLite with exact symbol filtering. + + Returns: + True if the history table was written. + """ + table_exists = False + for sym in symbols: + frame = fetch(date_from=date_from, date_to=date_to, symbol=sym) + if "symbol" in frame.columns: + frame = frame[frame["symbol"] == sym] + table_exists = _write_streamed_frame( + conn, + frame, + dataset, + table_exists, + if_exists, + written_columns, + ) + return table_exists + + +def _write_collected_datasets( + conn: sqlite3.Connection, + client: Mt5DataClient, + symbols: list[str], + datasets: set[Dataset], + timeframe: int, + flags: int, + date_from: datetime, + date_to: datetime, + if_exists: IfExists, +) -> tuple[set[Dataset], dict[Dataset, set[str]]]: + """Collect selected datasets and stream each symbol frame into SQLite. + + Returns: + Written datasets and their columns. + """ + written_columns: dict[Dataset, set[str]] = {} + written_tables: set[Dataset] = set() + if Dataset.rates in datasets and _write_rates_dataset( + conn, + client, + symbols, + timeframe, + date_from, + date_to, + if_exists, + written_columns, + ): + written_tables.add(Dataset.rates) + if Dataset.ticks in datasets and _write_ticks_dataset( + conn, + client, + symbols, + flags, + date_from, + date_to, + if_exists, + written_columns, + ): + written_tables.add(Dataset.ticks) + if Dataset.history_orders in datasets and _write_history_dataset( + conn, + client.history_orders_get_as_df, + Dataset.history_orders, + symbols, + date_from, + date_to, + if_exists, + written_columns, + ): + written_tables.add(Dataset.history_orders) + if Dataset.history_deals in datasets and _write_history_dataset( + conn, + client.history_deals_get_as_df, + Dataset.history_deals, + symbols, + date_from, + date_to, + if_exists, + written_columns, + ): + written_tables.add(Dataset.history_deals) + return written_tables, written_columns + + +def collect_history( + output: Path, + symbols: list[str], + date_from: datetime | str, + date_to: datetime | str, + *, + datasets: set[Dataset] | None = None, + timeframe: int | str = 1, + flags: int | str = 1, + if_exists: IfExists = IfExists.FAIL, + with_views: bool = False, + config: Mt5Config | None = None, +) -> None: + """Collect historical datasets into a single SQLite database. + + Args: + output: SQLite database path. + symbols: Symbols to collect. + date_from: Start date. + date_to: End date. + datasets: Datasets to include (defaults to all). + timeframe: Rates timeframe as integer or name (e.g. ``M1``). + flags: Tick copy flags as integer or name (e.g. ``ALL``). + if_exists: Behavior when a target table already exists. + with_views: Create ``cash_events`` and ``positions_reconstructed`` views. + config: MT5 connection configuration. + """ + start = _require_datetime(date_from) + end = _require_datetime(date_to) + selected = datasets if datasets is not None else set(Dataset) + tf = _coerce_timeframe(timeframe) + tick_flags = _coerce_tick_flags(flags) + mt5_config = config or build_config() + with _connected_client(mt5_config) as client, sqlite3.connect(output) as conn: + conn.execute("PRAGMA journal_mode=WAL") + conn.execute("PRAGMA synchronous=NORMAL") + written_tables, written_columns = _write_collected_datasets( + conn, + client, + symbols, + selected, + tf, + tick_flags, + start, + end, + if_exists, + ) + _create_collect_history_indexes(conn, written_columns) + if with_views and Dataset.history_deals in written_tables: + _create_cash_events_view(conn, written_columns[Dataset.history_deals]) + _create_positions_reconstructed_view( + conn, + written_columns[Dataset.history_deals], + ) + elif with_views: + logger.warning( + "--with-views ignored: history_deals table was not written", + ) + logger.info( + "Collected %s for %d symbol(s) into %s", + ", ".join(sorted(ds.value for ds in selected)), + len(symbols), + output, + ) + + +def _make_client(*, config: Mt5Config | None = None) -> Mt5CliClient: + return Mt5CliClient(config=config) if config is not None else Mt5CliClient() + + +def copy_rates_from( + symbol: str, + timeframe: int | str, + date_from: datetime | str, + count: int, + *, + config: Mt5Config | None = None, +) -> pd.DataFrame: + """Return rates starting from a date.""" + return _make_client(config=config).copy_rates_from( + symbol, + timeframe, + date_from, + count, + ) + + +def copy_rates_from_pos( + symbol: str, + timeframe: int | str, + start_pos: int, + count: int, + *, + config: Mt5Config | None = None, +) -> pd.DataFrame: + """Return rates starting from a bar position.""" + return _make_client(config=config).copy_rates_from_pos( + symbol, + timeframe, + start_pos, + count, + ) + + +def copy_rates_range( + symbol: str, + timeframe: int | str, + date_from: datetime | str, + date_to: datetime | str, + *, + config: Mt5Config | None = None, +) -> pd.DataFrame: + """Return rates for a date range.""" + return _make_client(config=config).copy_rates_range( + symbol, + timeframe, + date_from, + date_to, + ) + + +def copy_ticks_from( + symbol: str, + date_from: datetime | str, + count: int, + flags: int | str, + *, + config: Mt5Config | None = None, +) -> pd.DataFrame: + """Return ticks starting from a date.""" + return _make_client(config=config).copy_ticks_from( + symbol, + date_from, + count, + flags, + ) + + +def copy_ticks_range( + symbol: str, + date_from: datetime | str, + date_to: datetime | str, + flags: int | str, + *, + config: Mt5Config | None = None, +) -> pd.DataFrame: + """Return ticks for a date range.""" + return _make_client(config=config).copy_ticks_range( + symbol, + date_from, + date_to, + flags, + ) + + +def account_info(*, config: Mt5Config | None = None) -> pd.DataFrame: + """Return account information.""" + return _make_client(config=config).account_info() + + +def terminal_info(*, config: Mt5Config | None = None) -> pd.DataFrame: + """Return terminal information.""" + return _make_client(config=config).terminal_info() + + +def symbols( + group: str | None = None, + *, + config: Mt5Config | None = None, +) -> pd.DataFrame: + """Return the symbol list.""" + return _make_client(config=config).symbols(group=group) + + +def symbol_info( + symbol: str, + *, + config: Mt5Config | None = None, +) -> pd.DataFrame: + """Return details for one symbol.""" + return _make_client(config=config).symbol_info(symbol) + + +def orders( + symbol: str | None = None, + group: str | None = None, + ticket: int | None = None, + *, + config: Mt5Config | None = None, +) -> pd.DataFrame: + """Return active orders.""" + return _make_client(config=config).orders( + symbol=symbol, + group=group, + ticket=ticket, + ) + + +def positions( + symbol: str | None = None, + group: str | None = None, + ticket: int | None = None, + *, + config: Mt5Config | None = None, +) -> pd.DataFrame: + """Return open positions.""" + return _make_client(config=config).positions( + symbol=symbol, + group=group, + ticket=ticket, + ) + + +def history_orders( + date_from: datetime | str | None = None, + date_to: datetime | str | None = None, + group: str | None = None, + symbol: str | None = None, + ticket: int | None = None, + position: int | None = None, + *, + config: Mt5Config | None = None, +) -> pd.DataFrame: + """Return historical orders.""" + return _make_client(config=config).history_orders( + date_from=date_from, + date_to=date_to, + group=group, + symbol=symbol, + ticket=ticket, + position=position, + ) + + +def history_deals( + date_from: datetime | str | None = None, + date_to: datetime | str | None = None, + group: str | None = None, + symbol: str | None = None, + ticket: int | None = None, + position: int | None = None, + *, + config: Mt5Config | None = None, +) -> pd.DataFrame: + """Return historical deals.""" + return _make_client(config=config).history_deals( + date_from=date_from, + date_to=date_to, + group=group, + symbol=symbol, + ticket=ticket, + position=position, + ) + + +def version(*, config: Mt5Config | None = None) -> pd.DataFrame: + """Return MetaTrader5 version information.""" + return _make_client(config=config).version() + + +def last_error(*, config: Mt5Config | None = None) -> pd.DataFrame: + """Return the last error information.""" + return _make_client(config=config).last_error() + + +def symbol_info_tick( + symbol: str, + *, + config: Mt5Config | None = None, +) -> pd.DataFrame: + """Return the last tick for a symbol.""" + return _make_client(config=config).symbol_info_tick(symbol) + + +def market_book( + symbol: str, + *, + config: Mt5Config | None = None, +) -> pd.DataFrame: + """Return market depth for a symbol.""" + return _make_client(config=config).market_book(symbol) diff --git a/mt5cli/utils.py b/mt5cli/utils.py new file mode 100644 index 0000000..e393a57 --- /dev/null +++ b/mt5cli/utils.py @@ -0,0 +1,408 @@ +"""Utility constants, types, and functions for the mt5cli package.""" + +from __future__ import annotations + +import importlib +import json +from datetime import UTC, datetime +from enum import StrEnum +from pathlib import Path +from typing import TYPE_CHECKING, Any, TypeGuard, cast + +import click + +if TYPE_CHECKING: + import pandas as pd + +# --------------------------------------------------------------------------- +# Constants +# --------------------------------------------------------------------------- + +TIMEFRAME_MAP: dict[str, int] = { + "M1": 1, + "M2": 2, + "M3": 3, + "M4": 4, + "M5": 5, + "M6": 6, + "M10": 10, + "M12": 12, + "M15": 15, + "M20": 20, + "M30": 30, + "H1": 16385, + "H2": 16386, + "H3": 16387, + "H4": 16388, + "H6": 16390, + "H8": 16392, + "H12": 16396, + "D1": 16408, + "W1": 32769, + "MN1": 49153, +} + +TICK_FLAG_MAP: dict[str, int] = { + "ALL": 1, + "INFO": 2, + "TRADE": 4, +} + +_FORMAT_EXTENSIONS: dict[str, str] = { + ".csv": "csv", + ".json": "json", + ".parquet": "parquet", + ".pq": "parquet", + ".db": "sqlite3", + ".sqlite": "sqlite3", + ".sqlite3": "sqlite3", +} + +# --------------------------------------------------------------------------- +# Enums +# --------------------------------------------------------------------------- + + +class OutputFormat(StrEnum): + """Supported output file formats.""" + + csv = "csv" + json = "json" + parquet = "parquet" + sqlite3 = "sqlite3" + + +class LogLevel(StrEnum): + """Logging verbosity levels.""" + + DEBUG = "DEBUG" + INFO = "INFO" + WARNING = "WARNING" + ERROR = "ERROR" + + +class Dataset(StrEnum): + """Datasets supported by the ``collect-history`` command.""" + + rates = "rates" + ticks = "ticks" + history_orders = "history-orders" + history_deals = "history-deals" + + @property + def table_name(self) -> str: + """Return the SQLite table name for this dataset.""" + return self.value.replace("-", "_") + + +class IfExists(StrEnum): + """SQLite table conflict behavior for the ``collect-history`` command.""" + + APPEND = "append" + REPLACE = "replace" + FAIL = "fail" + + +# --------------------------------------------------------------------------- +# Click parameter types +# --------------------------------------------------------------------------- + + +class _DateTimeType(click.ParamType): + """Click parameter type for ISO 8601 datetime strings.""" + + name = "DATETIME" + + def convert( + self, + value: object, + param: click.Parameter | None, + ctx: click.Context | None, + ) -> datetime: + """Convert a string value to a timezone-aware datetime. + + Args: + value: Raw value from the command line. + param: Click parameter instance. + ctx: Click context. + + Returns: + Parsed datetime. + """ + if isinstance(value, datetime): + return value + try: + return parse_datetime(str(value)) + except ValueError as exc: + self.fail(str(exc), param, ctx) + + +class _TimeframeType(click.ParamType): + """Click parameter type for MT5 timeframe values.""" + + name = "TIMEFRAME" + + def convert( + self, + value: object, + param: click.Parameter | None, + ctx: click.Context | None, + ) -> int: + """Convert a string or integer value to a timeframe integer. + + Args: + value: Raw value from the command line. + param: Click parameter instance. + ctx: Click context. + + Returns: + Integer timeframe value. + """ + if isinstance(value, int): + return value + try: + return parse_timeframe(str(value)) + except ValueError as exc: + self.fail(str(exc), param, ctx) + + +class _TickFlagsType(click.ParamType): + """Click parameter type for MT5 tick copy flags.""" + + name = "FLAGS" + + def convert( + self, + value: object, + param: click.Parameter | None, + ctx: click.Context | None, + ) -> int: + """Convert a string or integer value to a tick flags integer. + + Args: + value: Raw value from the command line. + param: Click parameter instance. + ctx: Click context. + + Returns: + Integer tick flag value. + """ + if isinstance(value, int): + return value + try: + return parse_tick_flags(str(value)) + except ValueError as exc: + self.fail(str(exc), param, ctx) + + +class _RequestType(click.ParamType): + """Click parameter type for JSON order requests.""" + + name = "REQUEST" + + def convert( + self, + value: object, + param: click.Parameter | None, + ctx: click.Context | None, + ) -> dict[str, Any]: + """Convert a raw CLI value to an order request dictionary. + + Args: + value: Raw value from the command line. + param: Click parameter instance. + ctx: Click context. + + Returns: + Parsed request dictionary. + """ + try: + return parse_request(str(value)) + except ValueError as exc: + self.fail(str(exc), param, ctx) + + +DATETIME_TYPE = _DateTimeType() +TIMEFRAME_TYPE = _TimeframeType() +TICK_FLAGS_TYPE = _TickFlagsType() +REQUEST_TYPE = _RequestType() + +# --------------------------------------------------------------------------- +# Public utility functions +# --------------------------------------------------------------------------- + + +def detect_format( + output_path: Path, + explicit_format: str | None = None, +) -> str: + """Detect the output format from a file extension or explicit format string. + + Args: + output_path: Path to the output file. + explicit_format: Explicitly specified format, if any. + + Returns: + The detected format string. + + Raises: + ValueError: If the format cannot be determined. + """ + if explicit_format is not None: + return explicit_format + suffix = output_path.suffix.lower() + if suffix in _FORMAT_EXTENSIONS: + return _FORMAT_EXTENSIONS[suffix] + msg = ( + f"Cannot detect format from extension '{suffix}'." + " Use --format to specify the output format." + ) + raise ValueError(msg) + + +def export_dataframe( + df: pd.DataFrame, + output_path: Path, + output_format: str, + table_name: str = "data", +) -> None: + """Export a pandas DataFrame to the specified file format. + + Args: + df: DataFrame to export. + output_path: Path to the output file. + output_format: Output format (csv, json, parquet, or sqlite3). + table_name: Table name for SQLite3 output. + + Raises: + ValueError: If the output format is not supported. + """ + if output_format == "csv": + df.to_csv(output_path, index=False) + elif output_format == "json": + df.to_json( + output_path, + orient="records", + date_format="iso", + indent=2, + ) + 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, + ) + else: + msg = f"Unsupported output format: {output_format}" + raise ValueError(msg) + + +def parse_datetime(value: str) -> datetime: + """Parse an ISO 8601 datetime string to a timezone-aware datetime. + + Args: + value: ISO 8601 datetime string (e.g., '2024-01-01' or + '2024-01-01T12:00:00+00:00'). + + Returns: + Parsed datetime with UTC timezone if no timezone is specified. + + Raises: + ValueError: If the string cannot be parsed. + """ + try: + dt = datetime.fromisoformat(value) + except ValueError: + msg = f"Invalid datetime format: '{value}'. Use ISO 8601 format." + raise ValueError(msg) from None + if dt.tzinfo is None: + dt = dt.replace(tzinfo=UTC) + return dt + + +def parse_timeframe(value: str) -> int: + """Parse a timeframe string or integer value. + + Args: + value: Timeframe name (e.g., 'M1', 'H1', 'D1') or integer value. + + Returns: + Integer timeframe value. + + Raises: + ValueError: If the timeframe is invalid. + """ + upper = value.upper() + if upper in TIMEFRAME_MAP: + return TIMEFRAME_MAP[upper] + try: + return int(value) + except ValueError: + valid = ", ".join(TIMEFRAME_MAP) + msg = f"Invalid timeframe: '{value}'. Use one of: {valid}, or an integer." + raise ValueError(msg) from None + + +def parse_tick_flags(value: str) -> int: + """Parse tick flags string or integer value. + + Args: + value: Tick flag name (ALL, INFO, TRADE) or integer value. + + Returns: + Integer tick flag value. + + Raises: + ValueError: If the flag is invalid. + """ + upper = value.upper() + if upper in TICK_FLAG_MAP: + return TICK_FLAG_MAP[upper] + try: + return int(value) + except ValueError: + valid = ", ".join(TICK_FLAG_MAP) + msg = f"Invalid tick flags: '{value}'. Use one of: {valid}, or an integer." + raise ValueError(msg) from None + + +def _is_request_dict(value: object) -> TypeGuard[dict[str, Any]]: + return isinstance(value, dict) + + +def parse_request(value: str) -> dict[str, Any]: + """Parse a JSON-formatted order request string or file reference. + + Args: + value: JSON object string, or '@path' to read JSON from a file. + + Returns: + Parsed request dictionary. + + Raises: + ValueError: If the request file cannot be read or the value is not a + JSON object. + """ + if value.startswith("@"): + path = Path(value[1:]) + try: + text = path.read_text(encoding="utf-8") + except (OSError, UnicodeDecodeError) as exc: + msg = f"Failed to read JSON request file '{path}': {exc}" + raise ValueError(msg) from exc + else: + text = value + try: + parsed: object = json.loads(text) + except json.JSONDecodeError as exc: + msg = f"Invalid JSON request: {exc}" + raise ValueError(msg) from exc + if not _is_request_dict(parsed): + msg = "Order request must be a JSON object." + raise ValueError(msg) + return parsed diff --git a/pyproject.toml b/pyproject.toml index a9d0071..c663a5a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "mt5cli" -version = "0.3.0" +version = "0.4.0" 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 70251e2..79c21ab 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -19,22 +19,11 @@ if TYPE_CHECKING: from pathlib import Path from mt5cli.cli import ( - DATETIME_TYPE, - REQUEST_TYPE, - TICK_FLAG_MAP, - TICK_FLAGS_TYPE, - TIMEFRAME_MAP, - TIMEFRAME_TYPE, _execute_export, # type: ignore[reportPrivateUsage] _ExportContext, # type: ignore[reportPrivateUsage] + _sdk_client, # type: ignore[reportPrivateUsage] app, - detect_format, - export_dataframe, main, - parse_datetime, - parse_request, - parse_tick_flags, - parse_timeframe, ) runner = CliRunner() @@ -46,299 +35,6 @@ def normalize_cli_output(output: str) -> str: return " ".join(_ANSI_ESCAPE_RE.sub("", output).split()) -# --------------------------------------------------------------------------- -# detect_format -# --------------------------------------------------------------------------- - - -class TestDetectFormat: - """Tests for detect_format.""" - - def test_explicit_format_returned(self, tmp_path: Path) -> None: - """Test that explicit format overrides extension.""" - result = detect_format(tmp_path / "data.txt", explicit_format="csv") - assert result == "csv" - - @pytest.mark.parametrize( - ("filename", "expected"), - [ - ("data.csv", "csv"), - ("data.json", "json"), - ("data.parquet", "parquet"), - ("data.pq", "parquet"), - ("data.db", "sqlite3"), - ("data.sqlite", "sqlite3"), - ("data.sqlite3", "sqlite3"), - ("DATA.CSV", "csv"), - ("DATA.JSON", "json"), - ("DATA.PARQUET", "parquet"), - ], - ) - def test_auto_detect_from_extension( - self, - tmp_path: Path, - filename: str, - expected: str, - ) -> None: - """Test format auto-detection from file extension.""" - result = detect_format(tmp_path / filename) - assert result == expected - - def test_unknown_extension_raises(self, tmp_path: Path) -> None: - """Test that unknown extension raises ValueError.""" - with pytest.raises(ValueError, match="Cannot detect format"): - detect_format(tmp_path / "data.xyz") - - -# --------------------------------------------------------------------------- -# export_dataframe -# --------------------------------------------------------------------------- - - -class TestExportDataframe: - """Tests for export_dataframe.""" - - @pytest.fixture - def sample_df(self) -> pd.DataFrame: - """Create a sample DataFrame for testing.""" - return pd.DataFrame({"a": [1, 2, 3], "b": ["x", "y", "z"]}) - - def test_export_csv(self, tmp_path: Path, sample_df: pd.DataFrame) -> None: - """Test CSV export.""" - output = tmp_path / "out.csv" - export_dataframe(sample_df, output, "csv") - result = pd.read_csv(output) - pd.testing.assert_frame_equal(result, sample_df) - - def test_export_json(self, tmp_path: Path, sample_df: pd.DataFrame) -> None: - """Test JSON export.""" - output = tmp_path / "out.json" - export_dataframe(sample_df, output, "json") - with output.open() as f: - records = json.load(f) - assert len(records) == 3 - assert records[0]["a"] == 1 - - def test_export_parquet(self, tmp_path: Path, sample_df: pd.DataFrame) -> None: - """Test Parquet export.""" - output = tmp_path / "out.parquet" - export_dataframe(sample_df, output, "parquet") - result = pd.read_parquet(output) - pd.testing.assert_frame_equal(result, sample_df) - - def test_export_sqlite3(self, tmp_path: Path, sample_df: pd.DataFrame) -> None: - """Test SQLite3 export.""" - output = tmp_path / "out.db" - export_dataframe(sample_df, output, "sqlite3", table_name="test_table") - with sqlite3.connect(output) as conn: - result = pd.read_sql( # type: ignore[reportUnknownMemberType] - "SELECT * FROM test_table", - conn, - ) - pd.testing.assert_frame_equal(result, sample_df) - - def test_unsupported_format_raises( - self, - tmp_path: Path, - sample_df: pd.DataFrame, - ) -> None: - """Test that unsupported format raises ValueError.""" - with pytest.raises(ValueError, match="Unsupported output format"): - export_dataframe(sample_df, tmp_path / "out.txt", "xml") - - -# --------------------------------------------------------------------------- -# Parse helpers -# --------------------------------------------------------------------------- - - -class TestParseDatetime: - """Tests for parse_datetime.""" - - def test_valid_date(self) -> None: - """Test parsing a date string.""" - result = parse_datetime("2024-01-15") - assert result == datetime(2024, 1, 15, tzinfo=UTC) - - def test_valid_datetime_with_tz(self) -> None: - """Test parsing a datetime with timezone.""" - result = parse_datetime("2024-01-15T12:00:00+00:00") - assert result == datetime(2024, 1, 15, 12, 0, 0, tzinfo=UTC) - - def test_invalid_format_raises(self) -> None: - """Test that invalid format raises ValueError.""" - with pytest.raises(ValueError, match="Invalid datetime"): - parse_datetime("not-a-date") - - -class TestParseTimeframe: - """Tests for parse_timeframe.""" - - @pytest.mark.parametrize( - ("value", "expected"), - [("M1", 1), ("h1", 16385), ("D1", 16408), ("MN1", 49153)], - ) - def test_named_timeframe(self, value: str, expected: int) -> None: - """Test parsing named timeframes.""" - assert parse_timeframe(value) == expected - - def test_integer_timeframe(self) -> None: - """Test parsing integer timeframe.""" - assert parse_timeframe("42") == 42 - - def test_invalid_timeframe_raises(self) -> None: - """Test that invalid timeframe raises ValueError.""" - with pytest.raises(ValueError, match="Invalid timeframe"): - parse_timeframe("INVALID") - - -class TestParseTickFlags: - """Tests for parse_tick_flags.""" - - @pytest.mark.parametrize( - ("value", "expected"), - [("ALL", 1), ("info", 2), ("TRADE", 4)], - ) - def test_named_flag(self, value: str, expected: int) -> None: - """Test parsing named tick flags.""" - assert parse_tick_flags(value) == expected - - def test_integer_flag(self) -> None: - """Test parsing integer tick flag.""" - assert parse_tick_flags("7") == 7 - - def test_invalid_flag_raises(self) -> None: - """Test that invalid flag raises ValueError.""" - with pytest.raises(ValueError, match="Invalid tick flags"): - parse_tick_flags("INVALID") - - -# --------------------------------------------------------------------------- -# parse_request -# --------------------------------------------------------------------------- - - -class TestParseRequest: - """Tests for parse_request.""" - - def test_inline_json(self) -> None: - """Test parsing an inline JSON object string.""" - result = parse_request('{"action": 1, "symbol": "EURUSD"}') - assert result == {"action": 1, "symbol": "EURUSD"} - - def test_file_reference(self, tmp_path: Path) -> None: - """Test parsing JSON from a file via the @path syntax.""" - path = tmp_path / "req.json" - path.write_text('{"action": 2}', encoding="utf-8") - result = parse_request(f"@{path}") - assert result == {"action": 2} - - def test_invalid_json_raises(self) -> None: - """Test that invalid JSON raises ValueError.""" - with pytest.raises(ValueError, match="Invalid JSON request"): - parse_request("not json") - - def test_non_object_raises(self) -> None: - """Test that a non-object JSON raises ValueError.""" - with pytest.raises(ValueError, match="must be a JSON object"): - parse_request("[1, 2, 3]") - - def test_missing_file_raises(self, tmp_path: Path) -> None: - """Test that a missing request file raises ValueError.""" - path = tmp_path / "missing.json" - with pytest.raises(ValueError, match="Failed to read JSON request file"): - parse_request(f"@{path}") - - -# --------------------------------------------------------------------------- -# Constants -# --------------------------------------------------------------------------- - - -class TestConstants: - """Tests for module constants.""" - - def test_timeframe_map_has_expected_keys(self) -> None: - """Test that TIMEFRAME_MAP contains standard timeframes.""" - for key in ("M1", "M5", "M15", "M30", "H1", "H4", "D1", "W1", "MN1"): - assert key in TIMEFRAME_MAP - - def test_tick_flag_map_has_expected_keys(self) -> None: - """Test that TICK_FLAG_MAP contains standard flags.""" - assert set(TICK_FLAG_MAP) == {"ALL", "INFO", "TRADE"} - - -# --------------------------------------------------------------------------- -# Click ParamTypes -# --------------------------------------------------------------------------- - - -class TestDateTimeType: - """Tests for _DateTimeType.""" - - def test_convert_string(self) -> None: - """Test converting a string to datetime.""" - result = DATETIME_TYPE.convert("2024-06-15", None, None) - assert result == datetime(2024, 6, 15, tzinfo=UTC) - - def test_convert_datetime_passthrough(self) -> None: - """Test that datetime values pass through unchanged.""" - dt = datetime(2024, 1, 1, tzinfo=UTC) - assert DATETIME_TYPE.convert(dt, None, None) is dt - - def test_convert_invalid(self) -> None: - """Test that invalid values raise BadParameter.""" - with pytest.raises(Exception, match="Invalid datetime"): - DATETIME_TYPE.convert("bad", None, None) - - -class TestTimeframeType: - """Tests for _TimeframeType.""" - - def test_convert_string(self) -> None: - """Test converting a string to timeframe integer.""" - assert TIMEFRAME_TYPE.convert("H1", None, None) == 16385 - - def test_convert_int_passthrough(self) -> None: - """Test that integer values pass through unchanged.""" - assert TIMEFRAME_TYPE.convert(42, None, None) == 42 - - def test_convert_invalid(self) -> None: - """Test that invalid values raise BadParameter.""" - with pytest.raises(Exception, match="Invalid timeframe"): - TIMEFRAME_TYPE.convert("bad", None, None) - - -class TestTickFlagsType: - """Tests for _TickFlagsType.""" - - def test_convert_string(self) -> None: - """Test converting a string to tick flags integer.""" - assert TICK_FLAGS_TYPE.convert("ALL", None, None) == 1 - - def test_convert_int_passthrough(self) -> None: - """Test that integer values pass through unchanged.""" - assert TICK_FLAGS_TYPE.convert(7, None, None) == 7 - - def test_convert_invalid(self) -> None: - """Test that invalid values raise BadParameter.""" - with pytest.raises(Exception, match="Invalid tick flags"): - TICK_FLAGS_TYPE.convert("bad", None, None) - - -class TestRequestType: - """Tests for _RequestType.""" - - def test_convert_string(self) -> None: - """Test converting a JSON string to a request dictionary.""" - assert REQUEST_TYPE.convert('{"action": 1}', None, None) == {"action": 1} - - def test_convert_invalid(self) -> None: - """Test that invalid values raise BadParameter.""" - with pytest.raises(Exception, match="Invalid JSON request"): - REQUEST_TYPE.convert("bad", None, None) - - # --------------------------------------------------------------------------- # _execute_export # --------------------------------------------------------------------------- @@ -355,7 +51,7 @@ class TestExecuteExport: """Test that shutdown is called even when fetch raises.""" mock_client = MagicMock() mock_client.account_info_as_df.side_effect = RuntimeError("boom") - mocker.patch("mt5cli.cli.Mt5DataClient", return_value=mock_client) + mocker.patch("mt5cli.sdk.Mt5DataClient", return_value=mock_client) ctx = MagicMock() ctx.obj = _ExportContext( output=tmp_path / "out.csv", @@ -364,7 +60,7 @@ class TestExecuteExport: config=MagicMock(), ) with pytest.raises(RuntimeError, match="boom"): - _execute_export(ctx, lambda c: c.account_info_as_df()) + _execute_export(ctx, _sdk_client(ctx).account_info) mock_client.shutdown.assert_called_once() @@ -397,7 +93,7 @@ def mock_client(mocker: MockerFixture) -> MagicMock: client.market_book_get_as_df.return_value = sample_df client.order_check_as_df.return_value = sample_df client.order_send_as_df.return_value = sample_df - mocker.patch("mt5cli.cli.Mt5DataClient", return_value=client) + mocker.patch("mt5cli.sdk.Mt5DataClient", return_value=client) return client @@ -931,7 +627,7 @@ class TestCallback: mock_client = MagicMock() mock_client.account_info_as_df.return_value = pd.DataFrame({"a": [1]}) mocker.patch( - "mt5cli.cli.Mt5DataClient", + "mt5cli.sdk.Mt5DataClient", return_value=mock_client, ) mock_config = mocker.patch("mt5cli.cli.Mt5Config") @@ -984,7 +680,7 @@ class TestCallback: {"s": ["EURUSD"]}, ) mocker.patch( - "mt5cli.cli.Mt5DataClient", + "mt5cli.sdk.Mt5DataClient", return_value=mock_client, ) output = tmp_path / "out.db" @@ -1090,7 +786,7 @@ def _build_history_client(mocker: MockerFixture) -> MagicMock: client.history_orders_get_as_df.side_effect = _orders client.history_deals_get_as_df.side_effect = _deals - mocker.patch("mt5cli.cli.Mt5DataClient", return_value=client) + mocker.patch("mt5cli.sdk.Mt5DataClient", return_value=client) return client @@ -1431,7 +1127,7 @@ class TestCollectHistory: "ticket": [3, 4], "symbol": ["EURUSD", "EURUSDm"], }) - mocker.patch("mt5cli.cli.Mt5DataClient", return_value=client) + mocker.patch("mt5cli.sdk.Mt5DataClient", return_value=client) output = tmp_path / "history.db" result = runner.invoke( app, @@ -1519,9 +1215,9 @@ class TestCollectHistory: client.copy_ticks_range_as_df.return_value = pd.DataFrame({"x": [1]}) client.history_orders_get_as_df.return_value = pd.DataFrame({"x": [1]}) client.history_deals_get_as_df.return_value = pd.DataFrame({"x": [1]}) - mocker.patch("mt5cli.cli.Mt5DataClient", return_value=client) + mocker.patch("mt5cli.sdk.Mt5DataClient", return_value=client) output = tmp_path / "history.db" - with caplog.at_level(logging.WARNING, logger="mt5cli.cli"): + with caplog.at_level(logging.WARNING, logger="mt5cli.sdk"): result = runner.invoke( app, [ @@ -1559,7 +1255,7 @@ class TestCollectHistory: client = MagicMock() client.copy_rates_range_as_df.return_value = pd.DataFrame({"time": [1]}) client.history_deals_get_as_df.return_value = pd.DataFrame() - mocker.patch("mt5cli.cli.Mt5DataClient", return_value=client) + mocker.patch("mt5cli.sdk.Mt5DataClient", return_value=client) output = tmp_path / "history.db" result = runner.invoke( app, @@ -1598,7 +1294,7 @@ class TestCollectHistory: ) -> None: """Test that --with-views warns when history_deals is not written.""" output = tmp_path / "history.db" - with caplog.at_level(logging.WARNING, logger="mt5cli.cli"): + with caplog.at_level(logging.WARNING, logger="mt5cli.sdk"): result = runner.invoke( app, [ diff --git a/tests/test_sdk.py b/tests/test_sdk.py new file mode 100644 index 0000000..3432ef1 --- /dev/null +++ b/tests/test_sdk.py @@ -0,0 +1,466 @@ +"""Tests for mt5cli.sdk module.""" + +from __future__ import annotations + +import logging +import sqlite3 +from datetime import UTC, datetime +from typing import TYPE_CHECKING +from unittest.mock import MagicMock + +import pandas as pd +import pytest +from pytest_mock import MockerFixture # noqa: TC002 + +if TYPE_CHECKING: + from pathlib import Path + +from mt5cli import sdk +from mt5cli.sdk import ( + Mt5CliClient, + account_info, + build_config, + collect_history, + copy_rates_from, + copy_rates_from_pos, + copy_rates_range, + copy_ticks_from, + copy_ticks_range, + history_deals, + history_orders, + last_error, + market_book, + orders, + positions, + symbol_info, + symbol_info_tick, + symbols, + terminal_info, + version, +) +from mt5cli.utils import Dataset + +_DEALS_FIXTURE: dict[str, list[object]] = { + "ticket": [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14], + "position_id": [100, 100, 100, 0, 200, 200, 300, 400, 400, 500, 500, 600, 600, 600], + "symbol": [ + "EURUSD", + "EURUSD", + "EURUSD", + "", + "EURUSD", + "EURUSD", + "GBPUSD", + "GBPUSD", + "GBPUSD", + "EURUSD", + "EURUSD", + "GBPUSD", + "GBPUSD", + "GBPUSD", + ], + "time": [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14], + "type": [0, 0, 1, 2, 0, 1, 0, 0, 2, 0, 1, 0, 1, 1], + "entry": [0, 0, 1, 0, 0, 1, 0, 0, 2, 0, 3, 0, 2, 1], + "volume": [1.0, 3.0, 4.0, 0.0, 2.0, 2.0, 5.0, 1.0, 1.0, 2.0, 2.0, 3.0, 1.0, 3.0], + "price": [ + 1.10, + 1.20, + 1.50, + 0.0, + 2.00, + 2.20, + 1.30, + 1.30, + 1.40, + 1.00, + 1.05, + 1.10, + 9.99, + 1.40, + ], + "profit": [0.0, 0.0, 10.0, 5.0, 0.0, 8.0, 0.0, 0.0, -1.0, 0.0, 3.0, 0.0, -2.0, 7.0], +} + + +@pytest.fixture +def mock_client(mocker: MockerFixture) -> MagicMock: + """Create and patch a mock Mt5DataClient for SDK tests.""" + client = MagicMock() + sample_df = pd.DataFrame({"col": [1]}) + client.copy_rates_from_as_df.return_value = sample_df + client.copy_rates_from_pos_as_df.return_value = sample_df + client.copy_rates_range_as_df.return_value = sample_df + client.copy_ticks_from_as_df.return_value = sample_df + client.copy_ticks_range_as_df.return_value = sample_df + client.account_info_as_df.return_value = sample_df + client.terminal_info_as_df.return_value = sample_df + client.symbols_get_as_df.return_value = sample_df + client.symbol_info_as_df.return_value = sample_df + client.orders_get_as_df.return_value = sample_df + client.positions_get_as_df.return_value = sample_df + client.history_orders_get_as_df.return_value = sample_df + client.history_deals_get_as_df.return_value = sample_df + client.version_as_df.return_value = sample_df + client.last_error_as_df.return_value = sample_df + client.symbol_info_tick_as_df.return_value = sample_df + client.market_book_get_as_df.return_value = sample_df + mocker.patch("mt5cli.sdk.Mt5DataClient", return_value=client) + return client + + +def _build_history_client(mocker: MockerFixture) -> MagicMock: + """Build a mocked Mt5DataClient with per-symbol history results.""" + client = MagicMock() + + def _rates(**kwargs: object) -> pd.DataFrame: + return pd.DataFrame({ + "time": [1], + "open": [1.0], + "symbol_arg": [kwargs.get("symbol")], + }) + + def _ticks(**kwargs: object) -> pd.DataFrame: + return pd.DataFrame({ + "time": [1], + "bid": [1.0], + "symbol_arg": [kwargs.get("symbol")], + }) + + client.copy_rates_range_as_df.side_effect = _rates + client.copy_ticks_range_as_df.side_effect = _ticks + + def _orders(**kwargs: object) -> pd.DataFrame: + return pd.DataFrame({"ticket": [10], "symbol": [kwargs.get("symbol")]}) + + def _deals(**kwargs: object) -> pd.DataFrame: + sym = kwargs.get("symbol") + df = pd.DataFrame(_DEALS_FIXTURE) + return df[df["symbol"] == sym].reset_index(drop=True) + + client.history_orders_get_as_df.side_effect = _orders + client.history_deals_get_as_df.side_effect = _deals + mocker.patch("mt5cli.sdk.Mt5DataClient", return_value=client) + return client + + +class TestConnectionLifecycle: + """Tests for MT5 connection lifecycle helpers.""" + + def test_connected_client_shuts_down(self, mocker: MockerFixture) -> None: + """Test that _connected_client always shuts down.""" + mock_client = MagicMock() + mocker.patch("mt5cli.sdk.Mt5DataClient", return_value=mock_client) + config = MagicMock() + with sdk._connected_client(config): # type: ignore[reportPrivateUsage] + mock_client.initialize_and_login_mt5.assert_called_once() + mock_client.shutdown.assert_called_once() + + def test_connected_client_shutdown_on_init_failure( + self, + mocker: MockerFixture, + ) -> None: + """Test that shutdown is called when initialize/login fails.""" + mock_client = MagicMock() + mock_client.initialize_and_login_mt5.side_effect = RuntimeError( + "login failed", + ) + mocker.patch("mt5cli.sdk.Mt5DataClient", return_value=mock_client) + with ( + pytest.raises(RuntimeError, match="login failed"), + sdk._connected_client(MagicMock()), # type: ignore[reportPrivateUsage] + ): + pass + mock_client.shutdown.assert_called_once() + + def test_run_with_client_shutdown_on_error( + self, + mocker: MockerFixture, + ) -> None: + """Test that shutdown is called even when fetch raises.""" + mock_client = MagicMock() + mock_client.account_info_as_df.side_effect = RuntimeError("boom") + mocker.patch("mt5cli.sdk.Mt5DataClient", return_value=mock_client) + with pytest.raises(RuntimeError, match="boom"): + sdk._run_with_client( # type: ignore[reportPrivateUsage] + MagicMock(), + lambda c: c.account_info_as_df(), + ) + mock_client.shutdown.assert_called_once() + + def test_client_context_manager_reuses_connection( + self, + mocker: MockerFixture, + ) -> None: + """Test that context-managed client reuses one connection.""" + mock_client = MagicMock() + mock_client.account_info_as_df.return_value = pd.DataFrame({"a": [1]}) + mock_client.terminal_info_as_df.return_value = pd.DataFrame({"b": [2]}) + mocker.patch("mt5cli.sdk.Mt5DataClient", return_value=mock_client) + with Mt5CliClient() as client: + client.account_info() + client.terminal_info() + assert client.config is not None + mock_client.initialize_and_login_mt5.assert_called_once() + mock_client.shutdown.assert_called_once() + assert mock_client.account_info_as_df.call_count == 1 + assert mock_client.terminal_info_as_df.call_count == 1 + + def test_client_context_manager_shutdown_on_init_failure( + self, + mocker: MockerFixture, + ) -> None: + """Test that shutdown is called when context manager login fails.""" + mock_client = MagicMock() + mock_client.initialize_and_login_mt5.side_effect = RuntimeError( + "login failed", + ) + mocker.patch("mt5cli.sdk.Mt5DataClient", return_value=mock_client) + client = Mt5CliClient() + with pytest.raises(RuntimeError, match="login failed"), client: + pass + mock_client.shutdown.assert_called_once() + assert client._client is None # type: ignore[reportPrivateUsage] + + def test_exit_without_enter_is_noop(self) -> None: + """Test that __exit__ without __enter__ does not fail.""" + client = Mt5CliClient() + client.__exit__(None, None, None) + + +class TestModuleFunctions: + """Tests for module-level SDK wrappers.""" + + @pytest.mark.parametrize( + ("fn", "args", "method"), + [ + ( + copy_rates_from, + ("EURUSD", "M1", "2024-01-01", 10), + "copy_rates_from_as_df", + ), + ( + copy_rates_from_pos, + ("EURUSD", "M1", 0, 10), + "copy_rates_from_pos_as_df", + ), + ( + copy_ticks_from, + ("EURUSD", "2024-01-01", 10, "ALL"), + "copy_ticks_from_as_df", + ), + ( + copy_ticks_range, + ("EURUSD", "2024-01-01", "2024-02-01", "ALL"), + "copy_ticks_range_as_df", + ), + (account_info, (), "account_info_as_df"), + (terminal_info, (), "terminal_info_as_df"), + (symbols, ("*USD*",), "symbols_get_as_df"), + (symbol_info, ("EURUSD",), "symbol_info_as_df"), + (orders, (), "orders_get_as_df"), + (positions, (), "positions_get_as_df"), + (history_orders, (), "history_orders_get_as_df"), + (history_deals, (), "history_deals_get_as_df"), + (version, (), "version_as_df"), + (last_error, (), "last_error_as_df"), + (symbol_info_tick, ("EURUSD",), "symbol_info_tick_as_df"), + (market_book, ("EURUSD",), "market_book_get_as_df"), + ], + ) + def test_module_functions_delegate( + self, + mock_client: MagicMock, + fn: object, + args: tuple[object, ...], + method: str, + ) -> None: + """Test module-level functions call the expected client methods.""" + config = build_config(login=123) + result = fn(*args, config=config) # type: ignore[operator] + assert isinstance(result, pd.DataFrame) + getattr(mock_client, method).assert_called_once() + + +class TestMt5CliClient: + """Tests for Mt5CliClient SDK methods.""" + + def test_copy_rates_range_returns_dataframe( + self, + mock_client: MagicMock, + ) -> None: + """Test that copy_rates_range returns a DataFrame.""" + df = Mt5CliClient().copy_rates_range( + "EURUSD", + "D1", + "2024-01-01", + "2024-02-01", + ) + assert isinstance(df, pd.DataFrame) + mock_client.copy_rates_range_as_df.assert_called_once_with( + symbol="EURUSD", + timeframe=16408, + date_from=datetime(2024, 1, 1, tzinfo=UTC), + date_to=datetime(2024, 2, 1, tzinfo=UTC), + ) + + def test_copy_ticks_from_parses_flags( + self, + mock_client: MagicMock, + ) -> None: + """Test that string tick flags are parsed.""" + Mt5CliClient().copy_ticks_from("EURUSD", "2024-01-01", 100, "INFO") + mock_client.copy_ticks_from_as_df.assert_called_once_with( + symbol="EURUSD", + date_from=datetime(2024, 1, 1, tzinfo=UTC), + count=100, + flags=2, + ) + + def test_history_orders_accepts_string_dates( + self, + mock_client: MagicMock, + ) -> None: + """Test that string datetime inputs are parsed.""" + Mt5CliClient().history_orders( + date_from="2024-01-01", + date_to="2024-02-01", + ) + mock_client.history_orders_get_as_df.assert_called_once_with( + date_from=datetime(2024, 1, 1, tzinfo=UTC), + date_to=datetime(2024, 2, 1, tzinfo=UTC), + group=None, + symbol=None, + ticket=None, + position=None, + ) + + def test_module_function_delegates_to_client( + self, + mock_client: MagicMock, + ) -> None: + """Test module-level copy_rates_range delegates to the client.""" + df = copy_rates_range( + "USDJPY", + "M1", + "2024-01-01", + "2024-02-01", + ) + assert isinstance(df, pd.DataFrame) + mock_client.copy_rates_range_as_df.assert_called_once() + + +class TestCollectHistory: + """Tests for collect_history SDK function.""" + + @pytest.fixture + def history_client(self, mocker: MockerFixture) -> MagicMock: + """Create a mocked Mt5DataClient with history-style DataFrames.""" + return _build_history_client(mocker) + + def test_collect_history_writes_all_tables( + self, + tmp_path: Path, + history_client: MagicMock, + ) -> None: + """Test that collect_history writes rates, ticks, and history tables.""" + output = tmp_path / "history.db" + collect_history( + output, + ["EURUSD", "GBPUSD"], + "2024-01-01", + "2024-02-01", + ) + assert history_client.copy_rates_range_as_df.call_count == 2 + assert history_client.copy_ticks_range_as_df.call_count == 2 + with sqlite3.connect(output) as conn: + tables = { + row[0] + for row in conn.execute( + "SELECT name FROM sqlite_master WHERE type='table'", + ).fetchall() + } + assert {"rates", "ticks", "history_orders", "history_deals"} <= tables + + def test_collect_history_with_views( + self, + tmp_path: Path, + history_client: MagicMock, # noqa: ARG002 + ) -> None: + """Test that with_views creates cash_events and positions views.""" + output = tmp_path / "history.db" + collect_history( + output, + ["EURUSD", "GBPUSD"], + "2024-01-01", + "2024-02-01", + with_views=True, + ) + with sqlite3.connect(output) as conn: + views = { + row[0] + for row in conn.execute( + "SELECT name FROM sqlite_master WHERE type='view'", + ).fetchall() + } + positions = { + row[0] + for row in conn.execute( + "SELECT position_id FROM positions_reconstructed", + ).fetchall() + } + assert {"cash_events", "positions_reconstructed"} <= views + assert set(positions) == {100, 200, 500, 600} + + def test_collect_history_rates_table_has_timeframe( + self, + tmp_path: Path, + history_client: MagicMock, # noqa: ARG002 + ) -> None: + """Test that the rates table carries the requested timeframe value.""" + output = tmp_path / "history.db" + collect_history( + output, + ["EURUSD"], + "2024-01-01", + "2024-02-01", + datasets={Dataset.rates}, + timeframe="H1", + ) + with sqlite3.connect(output) as conn: + rows = conn.execute( + "SELECT DISTINCT timeframe FROM rates", + ).fetchall() + assert rows == [(16385,)] + + def test_collect_history_views_skipped_when_columns_missing( + self, + tmp_path: Path, + mocker: MockerFixture, + caplog: pytest.LogCaptureFixture, + ) -> None: + """Test that views are not created when required columns are missing.""" + client = MagicMock() + client.copy_rates_range_as_df.return_value = pd.DataFrame({"x": [1]}) + client.copy_ticks_range_as_df.return_value = pd.DataFrame({"x": [1]}) + client.history_orders_get_as_df.return_value = pd.DataFrame({"x": [1]}) + client.history_deals_get_as_df.return_value = pd.DataFrame({"x": [1]}) + mocker.patch("mt5cli.sdk.Mt5DataClient", return_value=client) + output = tmp_path / "history.db" + with caplog.at_level(logging.WARNING, logger="mt5cli.sdk"): + collect_history( + output, + ["EURUSD"], + "2024-01-01", + "2024-02-01", + with_views=True, + ) + with sqlite3.connect(output) as conn: + views = { + row[0] + for row in conn.execute( + "SELECT name FROM sqlite_master WHERE type='view'", + ).fetchall() + } + assert "cash_events" not in views + assert "positions_reconstructed" not in views diff --git a/tests/test_utils.py b/tests/test_utils.py new file mode 100644 index 0000000..fa635c8 --- /dev/null +++ b/tests/test_utils.py @@ -0,0 +1,335 @@ +"""Tests for mt5cli.utils module.""" + +from __future__ import annotations + +import json +import sqlite3 +from datetime import UTC, datetime +from typing import TYPE_CHECKING + +import pandas as pd +import pytest + +if TYPE_CHECKING: + from pathlib import Path + +from mt5cli.utils import ( + DATETIME_TYPE, + REQUEST_TYPE, + TICK_FLAG_MAP, + TICK_FLAGS_TYPE, + TIMEFRAME_MAP, + TIMEFRAME_TYPE, + Dataset, + detect_format, + export_dataframe, + parse_datetime, + parse_request, + parse_tick_flags, + parse_timeframe, +) + +# --------------------------------------------------------------------------- +# detect_format +# --------------------------------------------------------------------------- + + +class TestDetectFormat: + """Tests for detect_format.""" + + def test_explicit_format_returned(self, tmp_path: Path) -> None: + """Test that explicit format overrides extension.""" + result = detect_format(tmp_path / "data.txt", explicit_format="csv") + assert result == "csv" + + @pytest.mark.parametrize( + ("filename", "expected"), + [ + ("data.csv", "csv"), + ("data.json", "json"), + ("data.parquet", "parquet"), + ("data.pq", "parquet"), + ("data.db", "sqlite3"), + ("data.sqlite", "sqlite3"), + ("data.sqlite3", "sqlite3"), + ("DATA.CSV", "csv"), + ("DATA.JSON", "json"), + ("DATA.PARQUET", "parquet"), + ], + ) + def test_auto_detect_from_extension( + self, + tmp_path: Path, + filename: str, + expected: str, + ) -> None: + """Test format auto-detection from file extension.""" + result = detect_format(tmp_path / filename) + assert result == expected + + def test_unknown_extension_raises(self, tmp_path: Path) -> None: + """Test that unknown extension raises ValueError.""" + with pytest.raises(ValueError, match="Cannot detect format"): + detect_format(tmp_path / "data.xyz") + + +# --------------------------------------------------------------------------- +# export_dataframe +# --------------------------------------------------------------------------- + + +class TestExportDataframe: + """Tests for export_dataframe.""" + + @pytest.fixture + def sample_df(self) -> pd.DataFrame: + """Create a sample DataFrame for testing.""" + return pd.DataFrame({"a": [1, 2, 3], "b": ["x", "y", "z"]}) + + def test_export_csv(self, tmp_path: Path, sample_df: pd.DataFrame) -> None: + """Test CSV export.""" + output = tmp_path / "out.csv" + export_dataframe(sample_df, output, "csv") + result = pd.read_csv(output) + pd.testing.assert_frame_equal(result, sample_df) + + def test_export_json(self, tmp_path: Path, sample_df: pd.DataFrame) -> None: + """Test JSON export.""" + output = tmp_path / "out.json" + export_dataframe(sample_df, output, "json") + with output.open() as f: + records = json.load(f) + assert len(records) == 3 + assert records[0]["a"] == 1 + + def test_export_parquet(self, tmp_path: Path, sample_df: pd.DataFrame) -> None: + """Test Parquet export.""" + output = tmp_path / "out.parquet" + export_dataframe(sample_df, output, "parquet") + result = pd.read_parquet(output) + pd.testing.assert_frame_equal(result, sample_df) + + def test_export_sqlite3(self, tmp_path: Path, sample_df: pd.DataFrame) -> None: + """Test SQLite3 export.""" + output = tmp_path / "out.db" + export_dataframe(sample_df, output, "sqlite3", table_name="test_table") + with sqlite3.connect(output) as conn: + result = pd.read_sql( # type: ignore[reportUnknownMemberType] + "SELECT * FROM test_table", + conn, + ) + pd.testing.assert_frame_equal(result, sample_df) + + def test_unsupported_format_raises( + self, + tmp_path: Path, + sample_df: pd.DataFrame, + ) -> None: + """Test that unsupported format raises ValueError.""" + with pytest.raises(ValueError, match="Unsupported output format"): + export_dataframe(sample_df, tmp_path / "out.txt", "xml") + + +# --------------------------------------------------------------------------- +# Parse helpers +# --------------------------------------------------------------------------- + + +class TestParseDatetime: + """Tests for parse_datetime.""" + + def test_valid_date(self) -> None: + """Test parsing a date string.""" + result = parse_datetime("2024-01-15") + assert result == datetime(2024, 1, 15, tzinfo=UTC) + + def test_valid_datetime_with_tz(self) -> None: + """Test parsing a datetime with timezone.""" + result = parse_datetime("2024-01-15T12:00:00+00:00") + assert result == datetime(2024, 1, 15, 12, 0, 0, tzinfo=UTC) + + def test_invalid_format_raises(self) -> None: + """Test that invalid format raises ValueError.""" + with pytest.raises(ValueError, match="Invalid datetime"): + parse_datetime("not-a-date") + + +class TestParseTimeframe: + """Tests for parse_timeframe.""" + + @pytest.mark.parametrize( + ("value", "expected"), + [("M1", 1), ("h1", 16385), ("D1", 16408), ("MN1", 49153)], + ) + def test_named_timeframe(self, value: str, expected: int) -> None: + """Test parsing named timeframes.""" + assert parse_timeframe(value) == expected + + def test_integer_timeframe(self) -> None: + """Test parsing integer timeframe.""" + assert parse_timeframe("42") == 42 + + def test_invalid_timeframe_raises(self) -> None: + """Test that invalid timeframe raises ValueError.""" + with pytest.raises(ValueError, match="Invalid timeframe"): + parse_timeframe("INVALID") + + +class TestParseTickFlags: + """Tests for parse_tick_flags.""" + + @pytest.mark.parametrize( + ("value", "expected"), + [("ALL", 1), ("info", 2), ("TRADE", 4)], + ) + def test_named_flag(self, value: str, expected: int) -> None: + """Test parsing named tick flags.""" + assert parse_tick_flags(value) == expected + + def test_integer_flag(self) -> None: + """Test parsing integer tick flag.""" + assert parse_tick_flags("7") == 7 + + def test_invalid_flag_raises(self) -> None: + """Test that invalid flag raises ValueError.""" + with pytest.raises(ValueError, match="Invalid tick flags"): + parse_tick_flags("INVALID") + + +# --------------------------------------------------------------------------- +# parse_request +# --------------------------------------------------------------------------- + + +class TestParseRequest: + """Tests for parse_request.""" + + def test_inline_json(self) -> None: + """Test parsing an inline JSON object string.""" + result = parse_request('{"action": 1, "symbol": "EURUSD"}') + assert result == {"action": 1, "symbol": "EURUSD"} + + def test_file_reference(self, tmp_path: Path) -> None: + """Test parsing JSON from a file via the @path syntax.""" + path = tmp_path / "req.json" + path.write_text('{"action": 2}', encoding="utf-8") + result = parse_request(f"@{path}") + assert result == {"action": 2} + + def test_invalid_json_raises(self) -> None: + """Test that invalid JSON raises ValueError.""" + with pytest.raises(ValueError, match="Invalid JSON request"): + parse_request("not json") + + def test_non_object_raises(self) -> None: + """Test that a non-object JSON raises ValueError.""" + with pytest.raises(ValueError, match="must be a JSON object"): + parse_request("[1, 2, 3]") + + def test_missing_file_raises(self, tmp_path: Path) -> None: + """Test that a missing request file raises ValueError.""" + path = tmp_path / "missing.json" + with pytest.raises(ValueError, match="Failed to read JSON request file"): + parse_request(f"@{path}") + + +# --------------------------------------------------------------------------- +# Constants +# --------------------------------------------------------------------------- + + +class TestConstants: + """Tests for module constants.""" + + def test_timeframe_map_has_expected_keys(self) -> None: + """Test that TIMEFRAME_MAP contains standard timeframes.""" + for key in ("M1", "M5", "M15", "M30", "H1", "H4", "D1", "W1", "MN1"): + assert key in TIMEFRAME_MAP + + def test_tick_flag_map_has_expected_keys(self) -> None: + """Test that TICK_FLAG_MAP contains standard flags.""" + assert set(TICK_FLAG_MAP) == {"ALL", "INFO", "TRADE"} + + @pytest.mark.parametrize( + ("dataset", "expected"), + [ + (Dataset.rates, "rates"), + (Dataset.ticks, "ticks"), + (Dataset.history_orders, "history_orders"), + (Dataset.history_deals, "history_deals"), + ], + ) + def test_dataset_table_name(self, dataset: Dataset, expected: str) -> None: + """Test dataset SQLite table names.""" + assert dataset.table_name == expected + + +# --------------------------------------------------------------------------- +# Click ParamTypes +# --------------------------------------------------------------------------- + + +class TestDateTimeType: + """Tests for _DateTimeType.""" + + def test_convert_string(self) -> None: + """Test converting a string to datetime.""" + result = DATETIME_TYPE.convert("2024-06-15", None, None) + assert result == datetime(2024, 6, 15, tzinfo=UTC) + + def test_convert_datetime_passthrough(self) -> None: + """Test that datetime values pass through unchanged.""" + dt = datetime(2024, 1, 1, tzinfo=UTC) + assert DATETIME_TYPE.convert(dt, None, None) is dt + + def test_convert_invalid(self) -> None: + """Test that invalid values raise BadParameter.""" + with pytest.raises(Exception, match="Invalid datetime"): + DATETIME_TYPE.convert("bad", None, None) + + +class TestTimeframeType: + """Tests for _TimeframeType.""" + + def test_convert_string(self) -> None: + """Test converting a string to timeframe integer.""" + assert TIMEFRAME_TYPE.convert("H1", None, None) == 16385 + + def test_convert_int_passthrough(self) -> None: + """Test that integer values pass through unchanged.""" + assert TIMEFRAME_TYPE.convert(42, None, None) == 42 + + def test_convert_invalid(self) -> None: + """Test that invalid values raise BadParameter.""" + with pytest.raises(Exception, match="Invalid timeframe"): + TIMEFRAME_TYPE.convert("bad", None, None) + + +class TestTickFlagsType: + """Tests for _TickFlagsType.""" + + def test_convert_string(self) -> None: + """Test converting a string to tick flags integer.""" + assert TICK_FLAGS_TYPE.convert("ALL", None, None) == 1 + + def test_convert_int_passthrough(self) -> None: + """Test that integer values pass through unchanged.""" + assert TICK_FLAGS_TYPE.convert(7, None, None) == 7 + + def test_convert_invalid(self) -> None: + """Test that invalid values raise BadParameter.""" + with pytest.raises(Exception, match="Invalid tick flags"): + TICK_FLAGS_TYPE.convert("bad", None, None) + + +class TestRequestType: + """Tests for _RequestType.""" + + def test_convert_string(self) -> None: + """Test converting a JSON string to a request dictionary.""" + assert REQUEST_TYPE.convert('{"action": 1}', None, None) == {"action": 1} + + def test_convert_invalid(self) -> None: + """Test that invalid values raise BadParameter.""" + with pytest.raises(Exception, match="Invalid JSON request"): + REQUEST_TYPE.convert("bad", None, None) diff --git a/uv.lock b/uv.lock index 4cc83de..eae6710 100644 --- a/uv.lock +++ b/uv.lock @@ -487,7 +487,7 @@ wheels = [ [[package]] name = "mt5cli" -version = "0.3.0" +version = "0.4.0" source = { editable = "." } dependencies = [ { name = "click" },