Compare commits
2 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 5b44318d55 | |||
| 7f70073301 |
@@ -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
|
||||
|
||||
+39
-6
@@ -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
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
# SDK Module
|
||||
|
||||
::: mt5cli.sdk
|
||||
@@ -0,0 +1,3 @@
|
||||
# Utils Module
|
||||
|
||||
::: mt5cli.utils
|
||||
+41
-1
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
+46
-2
@@ -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",
|
||||
]
|
||||
|
||||
+96
-936
File diff suppressed because it is too large
Load Diff
+1023
File diff suppressed because it is too large
Load Diff
+408
@@ -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
|
||||
+1
-5
@@ -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"}]
|
||||
@@ -48,10 +48,6 @@ dev = [
|
||||
"pymdown-extensions >= 10.21.2",
|
||||
]
|
||||
|
||||
[tool.uv.build-backend]
|
||||
source-include = ["mt5cli/**", "LICENSE"]
|
||||
source-exclude = ["tests/**"]
|
||||
|
||||
[tool.ruff]
|
||||
line-length = 88
|
||||
exclude = ["build", ".venv"]
|
||||
|
||||
+12
-316
@@ -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,
|
||||
[
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user