diff --git a/README.md b/README.md index 508145a..98de680 100644 --- a/README.md +++ b/README.md @@ -133,8 +133,11 @@ update_history_with_config( - **`update_history`**: incremental append based on existing SQLite `MAX(time)` per symbol (and timeframe for rates); account-level deals use a separate cursor when `include_account_events=True`. - **`rates` table**: normalized storage with `symbol` and `timeframe` columns. - **Rate compatibility views**: mt5cli manages all `rate_*` views. Naming is `rate___` when a symbol has one timeframe, otherwise `rate____` (for example `rate_EURUSD__M1_1`). Stale `rate_*` views are dropped and recreated when rates change for offline tools such as mteor optimize. -- **Rate view resolution**: use `mt5cli.history.resolve_rate_view_name()` / `resolve_rate_view_names()` to map symbols and granularities to existing SQLite compatibility views without creating databases. +- **Rate view resolution**: use `resolve_rate_view_name()` / `resolve_rate_view_names()` to map symbols and granularities to existing SQLite compatibility views without creating databases. Both accept `None` (or a missing path) and return deterministic default names unless `require_existing=True`. - **Rate view loading**: use `load_rate_data()` / `load_rate_data_from_connection()` to load a SQLite rate table or view into a `DatetimeIndex` DataFrame. +- **Multi-series rate loading**: use `build_rate_targets()` to build neutral `RateTarget(symbol, timeframe)` pairs, `resolve_rate_tables()` to map them to table/view names (pass `require_existing=True` for strict resolution), and `load_rate_series_from_sqlite()` to load them into a mapping keyed by `(symbol, integer timeframe)`. The loader requires existing managed views unless `explicit_tables` is supplied, and rejects duplicate `(symbol, timeframe)` targets. +- **Multi-account latest rates**: use `collect_latest_rates_for_accounts()` with `AccountSpec` to read the latest bars for several account groups, merged into a `(symbol, integer timeframe)` mapping. +- **MT5 session helper**: use the `mt5_session()` context manager to attach to (or, when `Mt5Config.path` is set, launch) an MT5 terminal, log in, and yield a connected `Mt5CliClient` that shuts down on exit. - **SQLite export helpers**: use `export_dataframe_to_sqlite()` for append mode, optional index export, and post-write deduplication by key columns. - **Recent ticks and margins**: `recent_ticks()` and `minimum_margins()` SDK helpers (and matching CLI commands) cover common downstream read-only queries. diff --git a/docs/api/history.md b/docs/api/history.md index a9235ca..e08d137 100644 --- a/docs/api/history.md +++ b/docs/api/history.md @@ -183,3 +183,35 @@ rates = load_rate_data(Path("history.db"), view, count=1000) The loader accepts close-based OHLC rate data or tick-like bid/ask data. It validates that `time` exists, parses timestamps with pandas, and returns a DataFrame indexed by ascending `DatetimeIndex` named `time`. + +### Multi-series rate loading + +For loading many rate series at once, build neutral `RateTarget` pairs and load +them from SQLite in one call. View names are resolved via the same +compatibility-view rules, or you can pass `explicit_tables` to bypass resolution: + +```python +from pathlib import Path + +from mt5cli import build_rate_targets, load_rate_series_from_sqlite + +targets = build_rate_targets(["EURUSD", "GBPUSD"], ["M1", "H1"]) +series = load_rate_series_from_sqlite(Path("history.db"), targets, count=1000) +frame = series["EURUSD", 1] # keyed by (symbol, integer timeframe) +``` + +- `build_rate_targets()` returns `RateTarget(symbol, timeframe)` pairs in + row-major order, normalizing timeframe names such as `"M1"` to their integer + values; set `allow_missing_symbol=True` to address series solely by + `explicit_tables` (targets carry `symbol=None`). +- `resolve_rate_tables()` maps targets to table or view names and validates that + any `explicit_tables` count matches the target count. Pass + `require_existing=True` to raise `ValueError` instead of returning a + best-guess name when the database or managed view is missing. When + `explicit_tables` is provided, names are returned as-is and + `require_existing` is ignored. +- `load_rate_series_from_sqlite()` returns a mapping keyed by + `(symbol, integer timeframe)`. Unless `explicit_tables` is supplied, it + requires existing managed `rate_*` compatibility views and raises + `ValueError` when they are missing. Duplicate `(symbol, timeframe)` targets + are rejected. diff --git a/mt5cli/__init__.py b/mt5cli/__init__.py index a620bc0..14336e5 100644 --- a/mt5cli/__init__.py +++ b/mt5cli/__init__.py @@ -2,13 +2,28 @@ from importlib.metadata import version -from .history import load_rate_data, load_rate_data_from_connection +from .history import ( + RateTarget, + build_rate_targets, + build_rate_view_name, + load_rate_data, + load_rate_data_from_connection, + load_rate_series_from_sqlite, + resolve_history_datasets, + resolve_history_tick_flags, + resolve_history_timeframes, + resolve_rate_tables, + resolve_rate_view_name, + resolve_rate_view_names, +) from .sdk import ( + AccountSpec, Mt5CliClient, account_info, build_config, collect_history, collect_latest_rates, + collect_latest_rates_for_accounts, copy_rates_from, copy_rates_from_pos, copy_rates_range, @@ -20,6 +35,7 @@ from .sdk import ( latest_rates, market_book, minimum_margins, + mt5_session, mt5_summary, mt5_summary_as_df, orders, @@ -37,23 +53,35 @@ from .sdk import ( version as mt5_version, ) from .utils import ( + TICK_FLAG_MAP, + TIMEFRAME_MAP, Dataset, IfExists, detect_format, export_dataframe, export_dataframe_to_sqlite, + parse_datetime, + parse_tick_flags, + parse_timeframe, ) __version__ = version(__package__) if __package__ else None __all__ = [ + "TICK_FLAG_MAP", + "TIMEFRAME_MAP", + "AccountSpec", "Dataset", "IfExists", "Mt5CliClient", + "RateTarget", "account_info", "build_config", + "build_rate_targets", + "build_rate_view_name", "collect_history", "collect_latest_rates", + "collect_latest_rates_for_accounts", "copy_rates_from", "copy_rates_from_pos", "copy_rates_range", @@ -68,15 +96,26 @@ __all__ = [ "latest_rates", "load_rate_data", "load_rate_data_from_connection", + "load_rate_series_from_sqlite", "market_book", "minimum_margins", + "mt5_session", "mt5_summary", "mt5_summary_as_df", "mt5_version", "orders", + "parse_datetime", + "parse_tick_flags", + "parse_timeframe", "positions", "recent_history_deals", "recent_ticks", + "resolve_history_datasets", + "resolve_history_tick_flags", + "resolve_history_timeframes", + "resolve_rate_tables", + "resolve_rate_view_name", + "resolve_rate_view_names", "symbol_info", "symbol_info_tick", "symbols", diff --git a/mt5cli/history.py b/mt5cli/history.py index 51c931b..926e2d0 100644 --- a/mt5cli/history.py +++ b/mt5cli/history.py @@ -4,6 +4,7 @@ from __future__ import annotations import logging import sqlite3 +from dataclasses import dataclass from datetime import UTC, datetime from pathlib import Path from typing import TYPE_CHECKING, Literal, cast @@ -135,14 +136,17 @@ def _require_non_empty_identifier(identifier: str, kind: str) -> str: def _open_history_connection( - conn_or_path: SqliteConnOrPath, + conn_or_path: SqliteConnOrPath | None, ) -> tuple[sqlite3.Connection | None, bool]: """Open a read-only SQLite connection when given a path. Returns: - A connection and whether the caller should close it. When the path does - not exist, returns ``(None, False)`` without creating a database file. + A connection and whether the caller should close it. When ``conn_or_path`` + is None or the path does not exist, returns ``(None, False)`` without + creating a database file. """ + if conn_or_path is None: + return None, False if isinstance(conn_or_path, sqlite3.Connection): return conn_or_path, False path = Path(conn_or_path) @@ -376,7 +380,7 @@ def _resolve_rate_view_name_from_context( def resolve_rate_view_name( - conn_or_path: SqliteConnOrPath, + conn_or_path: SqliteConnOrPath | None, symbol: str, granularity: str, *, @@ -385,7 +389,9 @@ def resolve_rate_view_name( """Resolve the mt5cli-managed rate compatibility view name. Args: - conn_or_path: SQLite database path or open connection. + conn_or_path: SQLite database path or open connection. When None or a + non-existing path and ``require_existing`` is False, the deterministic + default view name is returned without creating a database file. symbol: Symbol stored in the normalized ``rates`` table. granularity: Timeframe name (for example ``M1``) or integer string. require_existing: When True, require the database and a managed view to exist. @@ -429,7 +435,7 @@ def resolve_rate_view_name( def resolve_rate_view_names( - conn_or_path: SqliteConnOrPath, + conn_or_path: SqliteConnOrPath | None, symbols: Sequence[str], granularities: Sequence[str], *, @@ -438,7 +444,9 @@ def resolve_rate_view_names( """Resolve rate compatibility view names for symbol and granularity pairs. Args: - conn_or_path: SQLite database path or open connection. + conn_or_path: SQLite database path or open connection. When None or a + non-existing path and ``require_existing`` is False, deterministic + default view names are returned without creating a database file. symbols: Symbols stored in the normalized ``rates`` table. granularities: Timeframe names (for example ``M1``) or integer strings. require_existing: When True, require the database and managed views to exist. @@ -482,6 +490,222 @@ def resolve_rate_view_names( conn.close() +@dataclass(frozen=True) +class RateTarget: + """A single rate series identified by symbol and timeframe. + + Attributes: + symbol: MT5 symbol name, or None when the rate series is addressed only + by an explicit table (for example a custom SQLite view). + timeframe: MT5 timeframe as an integer or name (for example ``M1``). + """ + + symbol: str | None + timeframe: int | str + + def __post_init__(self) -> None: + """Normalize accepted timeframe aliases to the stored integer value.""" + if not isinstance(self.timeframe, int): + object.__setattr__(self, "timeframe", parse_timeframe(self.timeframe)) + + @property + def timeframe_int(self) -> int: + """Return the timeframe as its integer MT5 value.""" + return cast("int", self.timeframe) + + +def build_rate_targets( + symbols: Sequence[str], + timeframes: Sequence[int | str], + *, + allow_missing_symbol: bool = False, +) -> list[RateTarget]: + """Build rate targets for every symbol and timeframe combination. + + Args: + symbols: MT5 symbol names. May be empty when ``allow_missing_symbol``. + timeframes: MT5 timeframes as integers or names (for example ``M1``). + allow_missing_symbol: When True and ``symbols`` is empty, build targets + with ``symbol=None`` for each timeframe instead of raising. + + Returns: + Targets in row-major order: every timeframe for the first symbol, then + every timeframe for the next symbol, and so on. + + Raises: + ValueError: If ``timeframes`` is empty, or ``symbols`` is empty and + ``allow_missing_symbol`` is False. + """ + if not timeframes: + msg = "At least one timeframe is required." + raise ValueError(msg) + if not symbols: + if not allow_missing_symbol: + msg = "At least one symbol is required." + raise ValueError(msg) + return [RateTarget(symbol=None, timeframe=tf) for tf in timeframes] + return [ + RateTarget(symbol=symbol, timeframe=tf) + for symbol in symbols + for tf in timeframes + ] + + +def resolve_rate_tables( + conn_or_path: SqliteConnOrPath | None, + targets: Sequence[RateTarget], + explicit_tables: Sequence[str] | None = None, + *, + require_existing: bool = False, +) -> list[str]: + """Resolve SQLite table or view names for rate targets. + + Args: + conn_or_path: SQLite database path or open connection. May be None when + ``explicit_tables`` is provided, or when ``require_existing`` is + False and deterministic default view names are sufficient. + targets: Rate targets to resolve. + explicit_tables: Optional explicit table or view names. When provided, + they are used as-is and must match the number of targets. + require_existing: When True, require the database and managed views to + exist for each symbol target. Ignored when ``explicit_tables`` is + provided. + + Returns: + Table or view names aligned with ``targets``. + + Raises: + ValueError: If ``targets`` is empty, ``explicit_tables`` length does not + match the target count, a target without a symbol is resolved + without an explicit table, or ``require_existing`` is True and the + database or a managed view is missing. + """ + target_list = list(targets) + if not target_list: + msg = "At least one rate target is required." + raise ValueError(msg) + if explicit_tables is not None: + tables = list(explicit_tables) + if len(tables) != len(target_list): + msg = ( + f"Expected {len(target_list)} explicit table(s) " + f"to match the targets, got {len(tables)}." + ) + raise ValueError(msg) + return tables + if any(target.symbol is None for target in target_list): + msg = ( + "Cannot resolve a rate table for a target without a symbol; " + "provide explicit_tables." + ) + raise ValueError(msg) + conn, should_close = _open_history_connection(conn_or_path) + try: + if conn is None: + if require_existing: + path = ( + conn_or_path + if isinstance(conn_or_path, (Path, str)) + else "database" + ) + msg = f"SQLite database not found: {path}" + raise ValueError(msg) + timeframe_counts = None + existing_views: set[str] = set() + else: + timeframe_counts = _load_rates_timeframe_counts(conn) + existing_views = _load_existing_rate_views(conn) + resolved: list[str] = [] + for target in target_list: + symbol = cast("str", target.symbol) + timeframe = target.timeframe_int + resolved.append( + _resolve_rate_view_name_from_context( + symbol=symbol, + timeframe=timeframe, + granularity_name=resolve_granularity_name(timeframe), + timeframe_counts=timeframe_counts, + existing_views=existing_views, + require_existing=require_existing, + ), + ) + return resolved + finally: + if should_close and conn is not None: + conn.close() + + +def load_rate_series_from_sqlite( + conn_or_path: SqliteConnOrPath, + targets: Sequence[RateTarget], + count: int, + explicit_tables: Sequence[str] | None = None, +) -> dict[tuple[str | None, int], pd.DataFrame]: + """Load multiple rate series from a SQLite database. + + Args: + conn_or_path: SQLite database path or open connection. + targets: Rate targets to load. Each ``(symbol, timeframe_int)`` pair + must be unique. + count: Number of most recent rows to load per series. + explicit_tables: Optional explicit table or view names matching targets. + When omitted, managed ``rate_*`` compatibility views must already + exist in the database. + + Returns: + Mapping keyed by ``(symbol, timeframe_int)`` to each rate DataFrame. + + Raises: + ValueError: If ``count`` is not positive, targets are empty, duplicate + ``(symbol, timeframe_int)`` pairs are present, or table resolution + fails. + """ + if count <= 0: + msg = "count must be positive." + raise ValueError(msg) + target_list = list(targets) + if not target_list: + msg = "At least one rate target is required." + raise ValueError(msg) + if explicit_tables is None and any(target.symbol is None for target in target_list): + msg = ( + "Cannot resolve a rate table for a target without a symbol; " + "provide explicit_tables." + ) + raise ValueError(msg) + seen_keys: set[tuple[str | None, int]] = set() + for target in target_list: + key = (target.symbol, target.timeframe_int) + if key in seen_keys: + symbol_repr = repr(target.symbol) + msg = f"Duplicate rate target: ({symbol_repr}, {target.timeframe_int})" + raise ValueError(msg) + seen_keys.add(key) + tables = ( + resolve_rate_tables(None, target_list, explicit_tables) + if explicit_tables is not None + else None + ) + conn, should_close = _open_existing_sqlite_database(conn_or_path) + try: + resolved_tables = tables or resolve_rate_tables( + conn, + target_list, + require_existing=True, + ) + return { + (target.symbol, target.timeframe_int): load_rate_data_from_connection( + conn, + table, + count=count, + ) + for target, table in zip(target_list, resolved_tables, strict=True) + } + finally: + if should_close: + conn.close() + + def get_table_columns(conn: sqlite3.Connection, table: str) -> set[str]: """Return existing SQLite columns for a table.""" quoted_table = quote_sqlite_identifier(table) diff --git a/mt5cli/sdk.py b/mt5cli/sdk.py index 5606c24..10641d1 100644 --- a/mt5cli/sdk.py +++ b/mt5cli/sdk.py @@ -6,7 +6,7 @@ import json import logging import sqlite3 from contextlib import contextmanager -from dataclasses import dataclass +from dataclasses import dataclass, field from datetime import UTC, datetime, timedelta from pathlib import Path from typing import TYPE_CHECKING, Self, TypeVar, cast @@ -40,11 +40,13 @@ T = TypeVar("T") logger = logging.getLogger(__name__) __all__ = [ + "AccountSpec", "Mt5CliClient", "account_info", "build_config", "collect_history", "collect_latest_rates", + "collect_latest_rates_for_accounts", "copy_rates_from", "copy_rates_from_pos", "copy_rates_range", @@ -56,6 +58,7 @@ __all__ = [ "latest_rates", "market_book", "minimum_margins", + "mt5_session", "mt5_summary", "mt5_summary_as_df", "orders", @@ -277,6 +280,26 @@ def _run_with_client( return fetch_fn(client) +@contextmanager +def mt5_session(config: Mt5Config | None = None) -> Iterator[Mt5CliClient]: + """Open an MT5 terminal session and yield a connected client. + + Launches the MetaTrader 5 terminal using ``Mt5Config.path`` (when set), + logs in, yields a connected :class:`Mt5CliClient`, and always shuts the + terminal down on exit. + + Args: + config: MT5 connection configuration. Defaults to an empty config that + attaches to a running terminal. + + Yields: + Connected ``Mt5CliClient`` bound to the session. + """ + mt5_config = config or build_config() + with _connected_client(mt5_config) as client: + yield Mt5CliClient.from_connected_client(client) + + class Mt5CliClient: """Programmatic client for read-only MetaTrader 5 data access.""" @@ -1075,6 +1098,120 @@ def collect_latest_rates( ) +@dataclass(frozen=True) +class AccountSpec: + """Connection parameters and symbols for one MT5 account group. + + Attributes: + symbols: Symbols to load latest rates for under this account. + login: Trading account login. String values are coerced to int when + non-empty. + password: Trading account password. + server: Trading server name. + path: Path to the MetaTrader5 terminal EXE file. + timeout: Connection timeout in milliseconds. + """ + + symbols: Sequence[str] + login: int | str | None = None + password: str | None = field(default=None, repr=False) + server: str | None = None + path: str | None = None + timeout: int | None = None + + +def _coerce_login(login: int | str | None) -> int | None: + """Coerce a login value to int, treating empty strings as unset. + + Returns: + Integer login, or None when unset or an empty string. + """ + if login is None or isinstance(login, int): + return login + text = login.strip() + if not text: + return None + return int(text) + + +def _build_account_config( + account: AccountSpec, + base_config: Mt5Config | None, +) -> Mt5Config: + """Build an ``Mt5Config`` for an account, falling back to ``base_config``. + + Returns: + Merged MT5 configuration for the account. + """ + login = _coerce_login(account.login) + if login is None and base_config is not None: + login = base_config.login + return build_config( + path=account.path or (base_config.path if base_config else None), + login=login, + password=account.password or (base_config.password if base_config else None), + server=account.server or (base_config.server if base_config else None), + timeout=account.timeout + if account.timeout is not None + else (base_config.timeout if base_config else None), + ) + + +def collect_latest_rates_for_accounts( + accounts: Sequence[AccountSpec], + timeframes: Sequence[int | str], + count: int, + *, + start_pos: int = 0, + base_config: Mt5Config | None = None, +) -> dict[tuple[str, int], pd.DataFrame]: + """Collect latest rates across multiple MT5 account groups. + + Each account is connected in turn, its symbols are read for every + timeframe, and the resulting frames are merged into a single mapping. + + Args: + accounts: Account groups to read. Each must define at least one symbol. + timeframes: MT5 timeframes as integers or names (for example ``M1``). + count: Number of most recent bars to read per symbol/timeframe. + start_pos: Initial bar position offset. + base_config: Optional base configuration whose fields fill any value not + set on an individual account. + + Returns: + Mapping keyed by ``(symbol, timeframe_int)``. When accounts share a + symbol/timeframe pair, the last account processed wins. + + Raises: + ValueError: If ``accounts``, ``timeframes``, or any account's symbols are + empty, or ``count`` is not positive. + """ + account_list = list(accounts) + if not account_list: + msg = "At least one account is required." + raise ValueError(msg) + if not timeframes: + msg = "At least one timeframe is required." + raise ValueError(msg) + if any(not account.symbols for account in account_list): + msg = "Each account requires at least one symbol." + raise ValueError(msg) + _require_positive(count, "count") + result: dict[tuple[str, int], pd.DataFrame] = {} + for account in account_list: + config = _build_account_config(account, base_config) + with Mt5CliClient(config=config) as client: + result.update( + client.collect_latest_rates( + account.symbols, + timeframes, + count=count, + start_pos=start_pos, + ), + ) + return result + + def copy_rates_range( symbol: str, timeframe: int | str, diff --git a/tests/test_history.py b/tests/test_history.py index b0adc30..c44ece3 100644 --- a/tests/test_history.py +++ b/tests/test_history.py @@ -10,14 +10,18 @@ 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 history from mt5cli.history import ( DEFAULT_HISTORY_TIMEFRAMES, + RateTarget, append_dataframe, augment_written_columns_from_sqlite, + build_rate_targets, build_rate_view_name, create_cash_events_view, create_history_indexes, @@ -33,6 +37,7 @@ from mt5cli.history import ( load_incremental_start_datetimes, load_rate_data, load_rate_data_from_connection, + load_rate_series_from_sqlite, parse_sqlite_timestamp, quote_sqlite_identifier, record_written_columns, @@ -40,6 +45,7 @@ from mt5cli.history import ( resolve_history_datasets, resolve_history_tick_flags, resolve_history_timeframes, + resolve_rate_tables, resolve_rate_view_name, resolve_rate_view_names, write_collected_datasets, @@ -60,6 +66,21 @@ class TestResolveRateViewName: assert resolve_rate_view_name(db_path, "EURUSD", "M1") == "rate_EURUSD__1" assert not db_path.exists() + def test_none_path_returns_default_name(self) -> None: + """Test a None connection or path returns the deterministic default.""" + assert resolve_rate_view_name(None, "EURUSD", "M1") == "rate_EURUSD__1" + assert resolve_rate_view_names(None, ["EURUSD"], ["M1", "H1"]) == [ + "rate_EURUSD__1", + "rate_EURUSD__16385", + ] + + def test_none_path_with_require_existing_raises(self) -> None: + """Test a None path under strict mode raises a clear error.""" + with pytest.raises(ValueError, match="SQLite database not found"): + resolve_rate_view_name(None, "EURUSD", "M1", require_existing=True) + with pytest.raises(ValueError, match="SQLite database not found"): + resolve_rate_view_names(None, ["EURUSD"], ["M1"], require_existing=True) + def test_no_rates_table_falls_back_to_single_timeframe_name( self, tmp_path: Path, @@ -1915,3 +1936,305 @@ class TestWriteHelpers: ) assert get_table_columns(conn, "rates") == {"time", "open"} create_history_indexes(conn, written_columns) + + +class TestRateSourceHelpers: + """Tests for generic rate-source SDK helpers.""" + + def test_rate_target_timeframe_int(self) -> None: + """Test RateTarget resolves named and integer timeframes.""" + target = RateTarget(symbol="EURUSD", timeframe="M1") + assert target.timeframe == 1 + assert target.timeframe_int == 1 + assert RateTarget(symbol="EURUSD", timeframe=16385).timeframe_int == 16385 + + def test_build_rate_targets_row_major(self) -> None: + """Test targets are built in row-major symbol/timeframe order.""" + targets = build_rate_targets(["EURUSD", "GBPUSD"], ["M1", "H1"]) + assert [(t.symbol, t.timeframe) for t in targets] == [ + ("EURUSD", 1), + ("EURUSD", 16385), + ("GBPUSD", 1), + ("GBPUSD", 16385), + ] + + def test_build_rate_targets_allows_missing_symbol(self) -> None: + """Test missing symbols produce None-symbol targets when allowed.""" + targets = build_rate_targets([], ["M1", "H1"], allow_missing_symbol=True) + assert [(t.symbol, t.timeframe) for t in targets] == [ + (None, 1), + (None, 16385), + ] + + @pytest.mark.parametrize( + ("symbols", "timeframes", "match"), + [ + (["EURUSD"], [], "At least one timeframe"), + ([], ["M1"], "At least one symbol"), + ], + ) + def test_build_rate_targets_rejects_empty( + self, + symbols: list[str], + timeframes: list[str], + match: str, + ) -> None: + """Test target building input validation.""" + with pytest.raises(ValueError, match=match): + build_rate_targets(symbols, timeframes) + + def test_resolve_rate_tables_uses_explicit_tables(self) -> None: + """Test explicit tables bypass view resolution when counts match.""" + targets = build_rate_targets([], ["M1", "H1"], allow_missing_symbol=True) + assert resolve_rate_tables(None, targets, ["t1", "t2"]) == ["t1", "t2"] + + def test_resolve_rate_tables_rejects_mismatched_explicit_count(self) -> None: + """Test explicit table count must match the number of targets.""" + targets = build_rate_targets(["EURUSD"], ["M1"]) + with pytest.raises(ValueError, match="Expected 1 explicit table"): + resolve_rate_tables(None, targets, ["t1", "t2"]) + + def test_resolve_rate_tables_rejects_empty_targets(self) -> None: + """Test resolving requires at least one target.""" + with pytest.raises(ValueError, match="At least one rate target"): + resolve_rate_tables(None, []) + + def test_resolve_rate_tables_requires_symbol_without_explicit(self) -> None: + """Test None-symbol targets require explicit tables.""" + targets = build_rate_targets([], ["M1"], allow_missing_symbol=True) + with pytest.raises(ValueError, match="without a symbol"): + resolve_rate_tables(None, targets) + + def test_resolve_rate_tables_resolves_view_names(self) -> None: + """Test symbol targets resolve to default view names without a database.""" + targets = build_rate_targets(["EURUSD"], ["M1", "H1"]) + assert resolve_rate_tables(None, targets) == [ + "rate_EURUSD__1", + "rate_EURUSD__16385", + ] + + def test_resolve_rate_tables_none_path_with_require_existing_raises(self) -> None: + """Test strict mode rejects a missing database path.""" + targets = build_rate_targets(["EURUSD"], ["M1"]) + with pytest.raises(ValueError, match="SQLite database not found"): + resolve_rate_tables(None, targets, require_existing=True) + + def test_resolve_rate_tables_missing_db_with_require_existing_raises( + self, + tmp_path: Path, + ) -> None: + """Test strict mode rejects a non-existing database path.""" + db_path = tmp_path / "missing.db" + targets = build_rate_targets(["EURUSD"], ["M1"]) + with pytest.raises(ValueError, match="SQLite database not found"): + resolve_rate_tables(db_path, targets, require_existing=True) + + def test_resolve_rate_tables_missing_view_with_require_existing_raises( + self, + tmp_path: Path, + ) -> None: + """Test strict mode rejects databases without managed rate views.""" + db_path = tmp_path / "no-views.db" + with sqlite3.connect(db_path) as conn: + conn.execute( + "CREATE TABLE rates(" + " symbol TEXT, timeframe INTEGER, time TEXT, close REAL)", + ) + conn.execute( + "INSERT INTO rates(symbol, timeframe, time, close) VALUES (?, ?, ?, ?)", + ("EURUSD", 1, "2024-01-01T00:00:00+00:00", 1.0), + ) + targets = build_rate_targets(["EURUSD"], ["M1"]) + with pytest.raises(ValueError, match="No rate compatibility view exists"): + resolve_rate_tables(db_path, targets, require_existing=True) + + def test_resolve_rate_tables_with_require_existing_resolves_views( + self, + tmp_path: Path, + ) -> None: + """Test strict mode resolves existing managed rate views.""" + db_path = tmp_path / "strict-views.db" + with sqlite3.connect(db_path) as conn: + conn.execute( + "CREATE TABLE rates(" + " symbol TEXT, timeframe INTEGER, time TEXT, close REAL)", + ) + conn.execute( + "INSERT INTO rates(symbol, timeframe, time, close) VALUES (?, ?, ?, ?)", + ("EURUSD", 1, "2024-01-01T00:00:00+00:00", 1.0), + ) + create_rate_compatibility_views(conn) + targets = build_rate_targets(["EURUSD"], ["M1"]) + assert resolve_rate_tables(db_path, targets, require_existing=True) == [ + "rate_EURUSD__1", + ] + + def test_resolve_rate_tables_batches_sqlite_metadata( + self, + tmp_path: Path, + mocker: MockerFixture, + ) -> None: + """Test resolving multiple targets loads SQLite metadata once.""" + db_path = tmp_path / "batch-rate-tables.db" + with sqlite3.connect(db_path) as conn: + conn.execute( + "CREATE TABLE rates(" + " symbol TEXT, timeframe INTEGER, time TEXT, close REAL)", + ) + conn.executemany( + "INSERT INTO rates(symbol, timeframe, time, close) VALUES (?, ?, ?, ?)", + [ + ("EURUSD", 1, "2024-01-01T00:00:00+00:00", 1.0), + ("EURUSD", 16385, "2024-01-01T01:00:00+00:00", 1.1), + ("GBPUSD", 1, "2024-01-01T00:00:00+00:00", 1.2), + ], + ) + create_rate_compatibility_views(conn) + counts_spy = mocker.spy(history, "_load_rates_timeframe_counts") + views_spy = mocker.spy(history, "_load_existing_rate_views") + + targets = build_rate_targets(["EURUSD", "GBPUSD"], ["M1", "H1"]) + assert resolve_rate_tables(db_path, targets) == [ + "rate_EURUSD__M1_1", + "rate_EURUSD__H1_16385", + "rate_GBPUSD__1", + "rate_GBPUSD__16385", + ] + assert counts_spy.call_count == 1 + assert views_spy.call_count == 1 + + def test_load_rate_series_from_sqlite(self, tmp_path: Path) -> None: + """Test loading multiple rate series keyed by symbol and timeframe.""" + db_path = tmp_path / "series.db" + with sqlite3.connect(db_path) as conn: + conn.execute( + "CREATE TABLE rates(" + " symbol TEXT, timeframe INTEGER, time TEXT, close REAL)", + ) + conn.executemany( + "INSERT INTO rates(symbol, timeframe, time, close) VALUES (?, ?, ?, ?)", + [ + ("EURUSD", 1, "2024-01-01T00:00:00+00:00", 1.0), + ("EURUSD", 1, "2024-01-01T00:01:00+00:00", 1.1), + ], + ) + create_rate_compatibility_views(conn) + targets = build_rate_targets(["EURUSD"], ["M1"]) + result = load_rate_series_from_sqlite(db_path, targets, count=2) + assert set(result) == {("EURUSD", 1)} + assert len(result["EURUSD", 1]) == 2 + + def test_load_rate_series_reuses_path_connection( + self, + tmp_path: Path, + mocker: MockerFixture, + ) -> None: + """Test loading from a path opens SQLite once for resolve and reads.""" + db_path = tmp_path / "single-open-series.db" + with sqlite3.connect(db_path) as conn: + conn.execute( + "CREATE TABLE rates(" + " symbol TEXT, timeframe INTEGER, time TEXT, close REAL)", + ) + conn.execute( + "INSERT INTO rates(symbol, timeframe, time, close) VALUES (?, ?, ?, ?)", + ("EURUSD", 1, "2024-01-01T00:00:00+00:00", 1.0), + ) + create_rate_compatibility_views(conn) + connect_spy = mocker.spy(history.sqlite3, "connect") + + result = load_rate_series_from_sqlite( + db_path, + build_rate_targets(["EURUSD"], ["M1"]), + count=1, + ) + + assert set(result) == {("EURUSD", 1)} + assert connect_spy.call_count == 1 + + def test_load_rate_series_with_explicit_tables(self, tmp_path: Path) -> None: + """Test explicit tables and None-symbol targets load series.""" + db_path = tmp_path / "explicit.db" + with sqlite3.connect(db_path) as conn: + conn.execute("CREATE TABLE custom_view(time TEXT, close REAL)") + conn.execute( + "INSERT INTO custom_view(time, close) VALUES (?, ?)", + ("2024-01-01T00:00:00+00:00", 1.0), + ) + targets = build_rate_targets([], ["M1"], allow_missing_symbol=True) + result = load_rate_series_from_sqlite( + db_path, + targets, + count=1, + explicit_tables=["custom_view"], + ) + assert set(result) == {(None, 1)} + + def test_load_rate_series_rejects_non_positive_count(self) -> None: + """Test loading requires a positive count.""" + targets = build_rate_targets(["EURUSD"], ["M1"]) + with pytest.raises(ValueError, match="count must be positive"): + load_rate_series_from_sqlite("unused.db", targets, count=0) + + def test_load_rate_series_rejects_empty_targets(self) -> None: + """Test loading requires at least one target before opening SQLite.""" + with pytest.raises(ValueError, match="At least one rate target"): + load_rate_series_from_sqlite("unused.db", [], count=1) + + def test_load_rate_series_requires_symbol_without_explicit_tables(self) -> None: + """Test None-symbol targets require explicit tables before opening SQLite.""" + targets = build_rate_targets([], ["M1"], allow_missing_symbol=True) + with pytest.raises(ValueError, match="without a symbol"): + load_rate_series_from_sqlite("unused.db", targets, count=1) + + def test_load_rate_series_requires_existing_managed_views( + self, + tmp_path: Path, + ) -> None: + """Test loading without explicit tables requires managed rate views.""" + db_path = tmp_path / "no-managed-views.db" + with sqlite3.connect(db_path) as conn: + conn.execute( + "CREATE TABLE rates(" + " symbol TEXT, timeframe INTEGER, time TEXT, close REAL)", + ) + conn.execute( + "INSERT INTO rates(symbol, timeframe, time, close) VALUES (?, ?, ?, ?)", + ("EURUSD", 1, "2024-01-01T00:00:00+00:00", 1.0), + ) + targets = build_rate_targets(["EURUSD"], ["M1"]) + with pytest.raises(ValueError, match="No rate compatibility view exists"): + load_rate_series_from_sqlite(db_path, targets, count=1) + + def test_load_rate_series_rejects_duplicate_targets(self) -> None: + """Test duplicate (symbol, timeframe) targets are rejected.""" + targets = [ + RateTarget("EURUSD", 1), + RateTarget("EURUSD", "M1"), + ] + with pytest.raises(ValueError, match=r"Duplicate rate target: \('EURUSD', 1\)"): + load_rate_series_from_sqlite("unused.db", targets, count=1) + + def test_load_rate_series_rejects_duplicate_targets_with_explicit_tables( + self, + tmp_path: Path, + ) -> None: + """Test duplicate targets are rejected even with explicit tables.""" + db_path = tmp_path / "duplicate-explicit.db" + with sqlite3.connect(db_path) as conn: + conn.execute("CREATE TABLE custom_view(time TEXT, close REAL)") + conn.execute( + "INSERT INTO custom_view(time, close) VALUES (?, ?)", + ("2024-01-01T00:00:00+00:00", 1.0), + ) + targets = [ + RateTarget("EURUSD", 1), + RateTarget("EURUSD", 1), + ] + with pytest.raises(ValueError, match=r"Duplicate rate target: \('EURUSD', 1\)"): + load_rate_series_from_sqlite( + db_path, + targets, + count=1, + explicit_tables=["custom_view", "custom_view"], + ) diff --git a/tests/test_sdk.py b/tests/test_sdk.py index 4217a19..131cc32 100644 --- a/tests/test_sdk.py +++ b/tests/test_sdk.py @@ -15,16 +15,18 @@ from pytest_mock import MockerFixture # noqa: TC002 if TYPE_CHECKING: from pathlib import Path - from pdmt5 import Mt5DataClient + from pdmt5 import Mt5Config, Mt5DataClient from mt5cli import sdk from mt5cli.history import DEFAULT_HISTORY_TIMEFRAMES from mt5cli.sdk import ( + AccountSpec, Mt5CliClient, account_info, build_config, collect_history, collect_latest_rates, + collect_latest_rates_for_accounts, copy_rates_from, copy_rates_from_pos, copy_rates_range, @@ -36,6 +38,7 @@ from mt5cli.sdk import ( latest_rates, market_book, minimum_margins, + mt5_session, mt5_summary, mt5_summary_as_df, orders, @@ -1248,3 +1251,177 @@ class TestMinimumMargins: ) client.order_calc_margin.assert_any_call(0, "EURUSD", 0.01, 1.1010) client.order_calc_margin.assert_any_call(1, "EURUSD", 0.01, 1.1000) + + +class TestMt5Session: + """Tests for the mt5_session context manager.""" + + def test_yields_connected_client_and_shuts_down( + self, + mocker: MockerFixture, + ) -> None: + """Test mt5_session connects, yields a client wrapper, and shuts down.""" + mock_client = MagicMock() + mt5_data_client = mocker.patch( + "mt5cli.sdk.Mt5DataClient", + return_value=mock_client, + ) + + with mt5_session(build_config(path="/opt/mt5/terminal64.exe")) as client: + mock_client.initialize_and_login_mt5.assert_called_once() + assert isinstance(client, Mt5CliClient) + + config = mt5_data_client.call_args.kwargs["config"] + assert config.path == "/opt/mt5/terminal64.exe" + mock_client.shutdown.assert_called_once() + + def test_default_config_attaches_to_running_terminal( + self, + mocker: MockerFixture, + ) -> None: + """Test mt5_session builds a default config when none is supplied.""" + mock_client = MagicMock() + mt5_data_client = mocker.patch( + "mt5cli.sdk.Mt5DataClient", + return_value=mock_client, + ) + + with mt5_session(): + pass + + mt5_data_client.assert_called_once() + mock_client.shutdown.assert_called_once() + + +class TestAccountSpec: + """Tests for account configuration helpers.""" + + def test_repr_omits_password(self) -> None: + """Test AccountSpec repr does not expose plaintext passwords.""" + spec = AccountSpec(symbols=["EURUSD"], login=123, password="secret") + + assert "secret" not in repr(spec) + assert "password" not in repr(spec) + + @pytest.mark.parametrize( + ("login", "expected"), + [ + (None, None), + (123, 123), + ("", None), + (" ", None), + ("456", 456), + ], + ) + def test_coerce_login( + self, + login: int | str | None, + expected: int | None, + ) -> None: + """Test login values are normalized for account configs.""" + assert sdk._coerce_login(login) == expected # type: ignore[reportPrivateUsage] + + def test_coerce_login_rejects_non_numeric_string(self) -> None: + """Test non-numeric login strings raise ValueError.""" + with pytest.raises(ValueError, match="invalid literal"): + sdk._coerce_login("abc") # type: ignore[reportPrivateUsage] + + +class TestCollectLatestRatesForAccounts: + """Tests for collect_latest_rates_for_accounts.""" + + def test_merges_results_across_accounts( + self, + mock_client: MagicMock, + mocker: MockerFixture, + ) -> None: + """Test rates are collected and merged for each account group.""" + mt5_data_client = mocker.patch( + "mt5cli.sdk.Mt5DataClient", + return_value=mock_client, + ) + accounts = [ + AccountSpec(symbols=["EURUSD"], login="123"), + AccountSpec(symbols=["GBPUSD"], login=456), + ] + + result = collect_latest_rates_for_accounts(accounts, ["M1"], count=2) + + assert set(result) == {("EURUSD", 1), ("GBPUSD", 1)} + assert mt5_data_client.call_count == 2 + assert mock_client.initialize_and_login_mt5.call_count == 2 + assert mock_client.shutdown.call_count == 2 + + def test_builds_config_from_account_and_base( + self, + mock_client: MagicMock, + mocker: MockerFixture, + ) -> None: + """Test account fields override base_config, empty login falls back.""" + configs: list[object] = [] + + def _record_config(*, config: object) -> MagicMock: + configs.append(config) + return mock_client + + mocker.patch("mt5cli.sdk.Mt5DataClient", side_effect=_record_config) + base = build_config(login=999, server="Base-Server", timeout=5000) + accounts = [ + AccountSpec(symbols=["EURUSD"], login="", server="Acct-Server"), + ] + + collect_latest_rates_for_accounts(accounts, ["M1"], count=1, base_config=base) + + assert len(configs) == 1 + config = cast("Mt5Config", configs[0]) + assert config.login == 999 + assert config.server == "Acct-Server" + assert config.timeout == 5000 + + @pytest.mark.parametrize( + ("accounts", "timeframes", "count", "match"), + [ + ([], ["M1"], 1, "At least one account"), + ([AccountSpec(symbols=["EURUSD"])], [], 1, "At least one timeframe"), + ( + [AccountSpec(symbols=[])], + ["M1"], + 1, + "Each account requires at least one symbol", + ), + ( + [AccountSpec(symbols=["EURUSD"])], + ["M1"], + 0, + "count must be positive", + ), + ], + ) + def test_rejects_invalid_inputs( + self, + accounts: list[AccountSpec], + timeframes: list[str], + count: int, + match: str, + ) -> None: + """Test input validation for account-level rate collection.""" + with pytest.raises(ValueError, match=match): + collect_latest_rates_for_accounts(accounts, timeframes, count) + + def test_rejects_empty_symbols_before_connecting( + self, + mocker: MockerFixture, + ) -> None: + """Test all account symbols are validated before any MT5 connection.""" + mt5_data_client = mocker.patch("mt5cli.sdk.Mt5DataClient") + accounts = [ + AccountSpec(symbols=["EURUSD"], login=123), + AccountSpec(symbols=[], login=456), + ] + + with pytest.raises( + ValueError, match="Each account requires at least one symbol" + ): + collect_latest_rates_for_accounts(accounts, ["M1"], count=1) + + mt5_data_client.assert_not_called()