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 <cursoragent@cursor.com>

* Harden SDK connection lifecycle and scope internal helpers as private.

Co-authored-by: Cursor <cursoragent@cursor.com>

* Export build_config in the public API and bump version to 0.4.0.

Co-authored-by: Cursor <cursoragent@cursor.com>

* Remove duplicate scripts/ in favor of local-qa skill script.

Co-authored-by: Cursor <cursoragent@cursor.com>

---------

Co-authored-by: Claude <noreply@anthropic.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Daichi Narushima
2026-06-08 22:54:53 +09:00
committed by GitHub
parent 7f70073301
commit 5b44318d55
15 changed files with 2479 additions and 1264 deletions
+12 -316
View File
@@ -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,
[
+466
View File
@@ -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
+335
View File
@@ -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)