95 lines
3.2 KiB
Python
95 lines
3.2 KiB
Python
import glob
|
|
import os
|
|
from typing import Dict, List, Optional
|
|
|
|
import pandas as pd
|
|
from pt_trading.fit_method import PairsTradingFitMethod
|
|
|
|
|
|
def resolve_datafiles(config: Dict, cli_datafiles: Optional[str] = None) -> List[str]:
|
|
"""
|
|
Resolve the list of data files to process.
|
|
CLI datafiles take priority over config datafiles.
|
|
Supports wildcards in config but not in CLI.
|
|
"""
|
|
if cli_datafiles:
|
|
# CLI override - comma-separated list, no wildcards
|
|
datafiles = [f.strip() for f in cli_datafiles.split(",")]
|
|
# Make paths absolute relative to data directory
|
|
data_dir = config.get("data_directory", "./data")
|
|
resolved_files = []
|
|
for df in datafiles:
|
|
if not os.path.isabs(df):
|
|
df = os.path.join(data_dir, df)
|
|
resolved_files.append(df)
|
|
return resolved_files
|
|
|
|
# Use config datafiles with wildcard support
|
|
config_datafiles = config.get("datafiles", [])
|
|
data_dir = config.get("data_directory", "./data")
|
|
resolved_files = []
|
|
|
|
for pattern in config_datafiles:
|
|
if "*" in pattern or "?" in pattern:
|
|
# Handle wildcards
|
|
if not os.path.isabs(pattern):
|
|
pattern = os.path.join(data_dir, pattern)
|
|
matched_files = glob.glob(pattern)
|
|
resolved_files.extend(matched_files)
|
|
else:
|
|
# Handle explicit file path
|
|
if not os.path.isabs(pattern):
|
|
pattern = os.path.join(data_dir, pattern)
|
|
resolved_files.append(pattern)
|
|
|
|
return sorted(list(set(resolved_files))) # Remove duplicates and sort
|
|
|
|
|
|
def create_pairs(
|
|
datafiles: List[str],
|
|
fit_method: PairsTradingFitMethod,
|
|
config: Dict,
|
|
instruments: List[Dict[str, str]],
|
|
) -> List:
|
|
from pt_trading.trading_pair import TradingPair
|
|
from tools.data_loader import load_market_data
|
|
|
|
all_indexes = range(len(instruments))
|
|
unique_index_pairs = [(i, j) for i in all_indexes for j in all_indexes if i < j]
|
|
pairs = []
|
|
|
|
# Update config to use the specified instruments
|
|
config_copy = config.copy()
|
|
config_copy["instruments"] = instruments
|
|
|
|
market_data_df = pd.DataFrame()
|
|
extra_minutes = 0
|
|
if "execution_price" in config_copy:
|
|
extra_minutes = config_copy["execution_price"]["shift"]
|
|
|
|
for datafile in datafiles:
|
|
md_df = load_market_data(
|
|
datafile=datafile,
|
|
instruments=instruments,
|
|
db_table_name=config_copy["market_data_loading"][instruments[0]["instrument_type"]]["db_table_name"],
|
|
trading_hours=config_copy["trading_hours"],
|
|
extra_minutes=extra_minutes,
|
|
)
|
|
market_data_df = pd.concat([market_data_df, md_df])
|
|
|
|
if len(set(market_data_df["symbol"])) != 2: # both symbols must be present for a pair
|
|
print(f"WARNING: insufficient data in files: {datafiles}")
|
|
return []
|
|
|
|
for a_index, b_index in unique_index_pairs:
|
|
symbol_a=instruments[a_index]["symbol"]
|
|
symbol_b=instruments[b_index]["symbol"]
|
|
pair = fit_method.create_trading_pair(
|
|
config=config_copy,
|
|
market_data=market_data_df,
|
|
symbol_a=symbol_a,
|
|
symbol_b=symbol_b,
|
|
)
|
|
pairs.append(pair)
|
|
return pairs
|