"""Tests for index downloader and read_index_daily functionality.""" import pandas as pd import os from unittest.mock import patch, MagicMock from datetime import date from pathlib import Path import pytest from sanguo_data.config import DataConfig def test_download_index_writes_parquet(tmp_path): """Test that download_index writes parquet files with correct structure.""" # Sample data that baostock would return sample_data = [ ["2024-01-02", "3495.0", "3505.0", "3490.0", "3500.0", "100000"], ["2024-01-03", "3505.0", "3515.0", "3500.0", "3510.0", "120000"], ["2024-01-04", "3515.0", "3525.0", "3510.0", "3520.0", "110000"], ] # Create a simple baostock mock class MockBaostock: class MockResult: def __init__(self, data): self.error_code = "success" self.error_msg = "success" self.data = data self.fields = ["date", "open", "high", "low", "close", "volume"] self.row_index = 0 def next(self): if self.row_index < len(self.data): row = self.data[self.row_index] self.row_index += 1 return True return False def get_row_data(self): return self.data[self.row_index - 1] def login(self): return self.MockResult([]) def logout(self): return self.MockResult([]) def query_history_k_data_plus(self, *args, **kwargs): return self.MockResult(sample_data) # Patch baostock module import sys sys.modules["baostock"] = MockBaostock() try: # Import after patching from sanguo_data.index_downloader import download_index # Download index data download_index("sh000300", 2024, 2024, str(tmp_path)) finally: # Clean up the mock del sys.modules["baostock"] # Verify parquet file was created expected_file = tmp_path / "2024" / "sh000300_daily.parquet" assert expected_file.exists(), f"Expected parquet file at {expected_file}" # Verify parquet content df_read = pd.read_parquet(expected_file) assert len(df_read) == 3 assert "close" in df_read.columns assert "date" in df_read.columns assert df_read["close"].iloc[0] == 3500.0 def test_download_index_clears_proxy(tmp_path): """Test that download_index clears proxy environment variables.""" # Set proxy variables os.environ["http_proxy"] = "http://evil:8080" os.environ["https_proxy"] = "https://evil:8080" sample_data = [["2024-01-02", "3495.0", "3505.0", "3490.0", "3500.0", "100000"]] # Create a simple baostock mock class MockBaostock: class MockResult: def __init__(self, data): self.error_code = "success" self.error_msg = "success" self.data = data self.fields = ["date", "open", "high", "low", "close", "volume"] self.row_index = 0 def next(self): if self.row_index < len(self.data): row = self.data[self.row_index] self.row_index += 1 return True return False def get_row_data(self): return self.data[self.row_index - 1] def login(self): return self.MockResult([]) def logout(self): return self.MockResult([]) def query_history_k_data_plus(self, *args, **kwargs): return self.MockResult(sample_data) # Patch baostock module import sys sys.modules["baostock"] = MockBaostock() try: # Import after patching from sanguo_data.index_downloader import download_index # Download index data download_index("sh000300", 2024, 2024, str(tmp_path)) finally: # Clean up the mock del sys.modules["baostock"] # Verify proxy variables were cleared assert "http_proxy" not in os.environ assert "https_proxy" not in os.environ def test_read_index_daily_reads_from_dbbardata(tmp_path): """read_index_daily 从 dbbardata 表读指数日线(2026-08-01 切自 vnpy DbBarData)。 建 tmp sqlite 库 + dbbardata 表插指数行,测真实 SQL 路径(非 mock)。 中证指数用 sh 前缀 → exchange=SSE(点位由 sina_index_eod 灌入)。 """ import sqlite3 from sanguo_data.datareader import read_index_daily db = tmp_path / "q.db" conn = sqlite3.connect(db) conn.execute( "CREATE TABLE dbbardata (symbol TEXT, exchange TEXT, datetime TEXT, interval TEXT, " "volume REAL, turnover REAL, open_interest REAL, open_price REAL, " "high_price REAL, low_price REAL, close_price REAL)" ) conn.executemany( "INSERT INTO dbbardata (symbol,exchange,datetime,interval,volume,turnover," "open_interest,open_price,high_price,low_price,close_price) VALUES (?,?,?,?,?,?,?,?,?,?,?)", [ ("000300", "SSE", "2024-01-02", "d", 100000, 0, 0, 3495.0, 3505.0, 3490.0, 3500.0), ("000300", "SSE", "2024-01-03", "d", 120000, 0, 0, 3505.0, 3515.0, 3500.0, 3510.0), ("000300", "SSE", "2024-01-04", "d", 110000, 0, 0, 3515.0, 3525.0, 3510.0, 3520.0), ], ) conn.commit() conn.close() cfg = DataConfig( data_paths={"vnpy_db": str(db)}, data_sources={}, validation={}, performance={}, ) result = read_index_daily("sh000300", date(2024, 1, 1), date(2024, 12, 31), cfg) assert len(result) == 3 assert list(result.columns) == ["date", "open", "high", "low", "close", "volume"] assert result["close"].iloc[0] == 3500.0 def test_read_index_daily_filters_date_range(tmp_path): """read_index_daily 用 substr(datetime,1,10) 比日期,正确过滤 start/end 范围。""" import sqlite3 from sanguo_data.datareader import read_index_daily db = tmp_path / "q.db" conn = sqlite3.connect(db) conn.execute( "CREATE TABLE dbbardata (symbol TEXT, exchange TEXT, datetime TEXT, interval TEXT, " "volume REAL, turnover REAL, open_interest REAL, open_price REAL, " "high_price REAL, low_price REAL, close_price REAL)" ) conn.executemany( "INSERT INTO dbbardata (symbol,exchange,datetime,interval,volume,turnover," "open_interest,open_price,high_price,low_price,close_price) VALUES (?,?,?,?,?,?,?,?,?,?,?)", [ ("000300", "SSE", "2024-01-02", "d", 100, 0, 0, 1, 1, 1, 3500.0), ("000300", "SSE", "2024-06-15", "d", 100, 0, 0, 1, 1, 1, 3600.0), ("000300", "SSE", "2025-01-10", "d", 100, 0, 0, 1, 1, 1, 3700.0), # 范围外 ], ) conn.commit() conn.close() cfg = DataConfig( data_paths={"vnpy_db": str(db)}, data_sources={}, validation={}, performance={}, ) result = read_index_daily("sh000300", "2024-01-01", "2024-12-31", cfg) assert len(result) == 2 # 只 2024 两行, 2025 范围外排除 assert result["close"].iloc[1] == 3600.0