Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 4 additions & 3 deletions backend/app/api/__init__.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
from fastapi import APIRouter
from app.api import auth, experiments, models, configs, benchmarks, factors, data, train, monitoring
from app.api import auth, experiments, models, configs, benchmarks, factors, data, train, monitoring, tasks

# Create main API router
api_router = APIRouter()
Expand All @@ -12,7 +12,8 @@
api_router.include_router(benchmarks.router, prefix="/benchmarks", tags=["benchmarks"])
api_router.include_router(factors.router, prefix="/factors", tags=["factors"])
api_router.include_router(data.router, prefix="/data", tags=["data"])
api_router.include_router(train.router, prefix="/train", tags=["train"]) # 添加训练API路由器
api_router.include_router(monitoring.router, prefix="/monitoring", tags=["monitoring"]) # 添加监控API路由器
api_router.include_router(train.router, prefix="/train", tags=["train"])
api_router.include_router(monitoring.router, prefix="/monitoring", tags=["monitoring"])
api_router.include_router(tasks.router, prefix="/tasks", tags=["tasks"])


7 changes: 5 additions & 2 deletions backend/app/api/auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -183,8 +183,11 @@ def register_user(
role=user.role
)

# Send verification email if email is provided
if user.email:
# Skip email verification if configured
if settings.skip_email_verification:
db_user.email_verified = True
elif user.email:
# Send verification email if email is provided
send_verification_email(db_user)

db.add(db_user)
Expand Down
26 changes: 26 additions & 0 deletions backend/app/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,32 @@ class Settings:

# Email verification settings
verification_token_expire_minutes = int(os.getenv("VERIFICATION_TOKEN_EXPIRE_MINUTES", "1440")) # 24 hours
skip_email_verification = os.getenv("SKIP_EMAIL_VERIFICATION", "False").lower() in ("true", "1", "t")

# CORS settings
cors_origins = os.getenv("CORS_ORIGINS", "http://localhost:3001,http://localhost:3000,http://localhost:8000,http://127.0.0.1:3001,http://127.0.0.1:3000,http://127.0.0.1:8000")

# Default production origins always included
_default_origins = [
"http://116.62.59.244",
"http://qlib.hoo.ink",
"http://ddns.hoo.ink:8000",
]

def get_cors_origins(self) -> list:
"""Parse CORS origins from settings, combining env var and defaults."""
origins = self.cors_origins
if isinstance(origins, str):
parsed = [o.strip() for o in origins.split(",") if o.strip()]
elif isinstance(origins, list):
parsed = origins
else:
parsed = []
# Merge with defaults, avoiding duplicates
for o in self._default_origins:
if o not in parsed:
parsed.append(o)
return parsed
Comment on lines +47 to +60

Copilot AI Mar 30, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

When cors_origins is already a list, parsed = origins means get_cors_origins() mutates the underlying settings list by appending defaults. That can cause surprising side effects across calls. Use a copy for list inputs (e.g., parsed = list(origins)) before appending/deduping.

Copilot uses AI. Check for mistakes.

# Create settings instance
settings = Settings()
4 changes: 2 additions & 2 deletions backend/app/services/training_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,8 +9,8 @@ class TrainingClient:
"""API client for communicating with the training server"""

def __init__(self):
self.base_url = settings.TRAINING_SERVER_URL
self.timeout = settings.TRAINING_SERVER_TIMEOUT
self.base_url = settings.training_server_url
self.timeout = settings.training_server_timeout
self.client = httpx.AsyncClient(
base_url=self.base_url,
timeout=self.timeout,
Expand Down
16 changes: 5 additions & 11 deletions backend/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -67,16 +67,8 @@ async def send_log(self, log: str, task_id: str):
# Configure CORS from environment variable if available, otherwise use production origins
from app.db.database import settings

# Get allowed origins from settings or use default production origins
allow_origins = getattr(settings, "cors_origins", [
"http://116.62.59.244", # Production IP
"http://qlib.hoo.ink", # Production domain
"http://ddns.hoo.ink:8000" # DDNS server for training
])

# Ensure allow_origins is a list
if isinstance(allow_origins, str):
allow_origins = allow_origins.split(",")
# Get allowed origins from settings
allow_origins = settings.get_cors_origins()

app.add_middleware(
CORSMiddleware,
Expand Down Expand Up @@ -207,7 +199,7 @@ async def monitor_performance(request, call_next):
return response

# Include API router without training endpoints
from app.api import auth, experiments, models, configs, benchmarks, factors, data, monitoring
from app.api import auth, experiments, models, configs, benchmarks, factors, data, monitoring, train, tasks
Comment on lines 201 to +202

Copilot AI Mar 30, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The comment says “without training endpoints” but train is now imported and later mounted. Update the comment to reflect the new behavior (or remove it) to avoid misleading future changes.

Copilot uses AI. Check for mistakes.

# Create main API router without training
main_api_router = APIRouter()
Expand All @@ -221,6 +213,8 @@ async def monitor_performance(request, call_next):
main_api_router.include_router(factors.router, prefix="/factors", tags=["factors"])
main_api_router.include_router(data.router, prefix="/data", tags=["data"])
main_api_router.include_router(monitoring.router, prefix="/monitoring", tags=["monitoring"])
main_api_router.include_router(train.router, prefix="/train", tags=["train"])
main_api_router.include_router(tasks.router, prefix="/tasks", tags=["tasks"])

# Include main API router
app.include_router(main_api_router, prefix="/api")
Expand Down
12 changes: 2 additions & 10 deletions backend/train_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -67,16 +67,8 @@ async def send_log(self, log: str, task_id: str):
# Configure CORS from environment variable if available, otherwise use production origins
from app.db.database import settings

# Get allowed origins from settings or use default production origins
allow_origins = getattr(settings, "cors_origins", [
"http://116.62.59.244", # Production IP
"http://qlib.hoo.ink", # Production domain
"http://ddns.hoo.ink:8000" # DDNS server for training
])

# Ensure allow_origins is a list
if isinstance(allow_origins, str):
allow_origins = allow_origins.split(",")
# Get allowed origins from settings
allow_origins = settings.get_cors_origins()

app.add_middleware(
CORSMiddleware,
Expand Down
259 changes: 259 additions & 0 deletions scripts/import_investment_data.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,259 @@
#!/usr/bin/env python3
"""
Import historical stock data using chenditc/investment_data.

This script downloads and converts data from the investment_data project
into qlib format for use with the qlib_t management platform.

Usage:
python scripts/import_investment_data.py [--qlib-dir ~/.qlib/qlib_data/cn_data]

Prerequisites:
pip install investment_data

Reference: https://github.com/chenditc/investment_data
"""

import argparse
import logging
import os
import subprocess
import sys

logging.basicConfig(
level=logging.INFO,
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s'
)
logger = logging.getLogger(__name__)


def install_investment_data():
"""Install investment_data package if not already installed."""
try:
import investment_data # noqa: F401
logger.info("investment_data package is already installed.")
return True
except ImportError:
logger.info("Installing investment_data package...")
result = subprocess.run(
[sys.executable, "-m", "pip", "install", "investment_data"],
capture_output=True,
text=True
)
if result.returncode != 0:
logger.error(f"Failed to install investment_data: {result.stderr}")
return False
logger.info("investment_data package installed successfully.")
return True


def download_qlib_data(qlib_dir: str):
"""
Download and convert data to qlib format using investment_data.

The investment_data package provides a CLI command to download
and convert Chinese A-share market data into qlib format.
"""
# Ensure the target directory exists
os.makedirs(qlib_dir, exist_ok=True)

logger.info(f"Downloading qlib data to: {qlib_dir}")
logger.info("This may take several minutes depending on your network speed...")

# Use investment_data CLI to download data
# The package provides: investment_data download_qlib_data --target_dir <dir>
try:
result = subprocess.run(
[
sys.executable, "-m", "investment_data",
"download_qlib_data",
"--target_dir", qlib_dir
],
capture_output=True,
text=True,
timeout=3600 # 1 hour timeout
)

if result.returncode == 0:
logger.info("Data download completed successfully.")
if result.stdout:
logger.info(f"Output: {result.stdout[-500:]}")
return True
else:
logger.warning(f"CLI method returned non-zero: {result.stderr}")
# Try alternative method
return download_qlib_data_alternative(qlib_dir)
except subprocess.TimeoutExpired:
logger.error("Data download timed out after 1 hour.")
return False
except FileNotFoundError:
logger.warning("investment_data CLI not found, trying alternative method...")
return download_qlib_data_alternative(qlib_dir)


def download_qlib_data_alternative(qlib_dir: str):
"""
Alternative method to download data using Python API.
"""
try:
logger.info("Trying alternative download method via Python API...")

# Try using the Python API directly
from investment_data import download_qlib_data as _download
_download(target_dir=qlib_dir)
logger.info("Data download completed successfully via Python API.")
return True
except ImportError:
logger.warning("Python API method not available.")
except Exception as e:
logger.warning(f"Python API method failed: {e}")

# Final fallback: use wget to download from GitHub releases
try:
logger.info("Trying to download from GitHub releases...")
import urllib.request
import zipfile
import tempfile

# chenditc/investment_data releases contain pre-built qlib data
release_url = "https://github.com/chenditc/investment_data/releases/latest"
Comment on lines +114 to +119

Copilot AI Mar 30, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

urllib.request, zipfile, tempfile, and release_url are currently unused in this fallback path (the function just prints manual instructions and returns False). Either remove these unused imports/variable or implement the intended “download release asset + extract” fallback to match the in-code comment.

Copilot uses AI. Check for mistakes.

logger.info(f"Checking latest release at: {release_url}")
logger.info("Please manually download qlib data from:")
logger.info(" https://github.com/chenditc/investment_data/releases")
logger.info(f" and extract to: {qlib_dir}")
logger.info("")
logger.info("Or run the following commands:")
logger.info(f" pip install investment_data")
logger.info(f" python -m investment_data download_qlib_data --target_dir {qlib_dir}")

return False
except Exception as e:
logger.error(f"Failed to download data: {e}")
return False


def verify_qlib_data(qlib_dir: str):
"""Verify that qlib data was downloaded correctly."""
required_dirs = ["instruments", "calendars", "features"]

logger.info(f"Verifying qlib data in: {qlib_dir}")

missing = []
for d in required_dirs:
path = os.path.join(qlib_dir, d)
if not os.path.exists(path):
missing.append(d)
else:
# Count files in directory
file_count = sum(1 for _ in os.scandir(path) if _.is_file() or _.is_dir())
logger.info(f" {d}/: {file_count} items")

if missing:
logger.warning(f"Missing directories: {missing}")
return False

# Check instruments file
instruments_dir = os.path.join(qlib_dir, "instruments")
if os.path.exists(instruments_dir):
all_txt = os.path.join(instruments_dir, "all.txt")
if os.path.exists(all_txt):
with open(all_txt, 'r') as f:
lines = f.readlines()
logger.info(f" instruments/all.txt: {len(lines)} instruments")
else:
logger.warning(" instruments/all.txt not found")

logger.info("Data verification completed.")
return True


def init_qlib_with_data(qlib_dir: str):
"""Initialize qlib with the downloaded data to verify it works."""
try:
import qlib
from qlib.config import REG_CN

logger.info(f"Initializing qlib with provider_uri: {qlib_dir}")
qlib.init(provider_uri=qlib_dir, region=REG_CN)

from qlib.data import D

# Test getting instruments
instruments = D.instruments(market="all")
logger.info(f"QLib initialized successfully. Market instruments loaded.")

# Test getting calendar
calendar = D.calendar(start_time="2020-01-01", end_time="2020-01-31")
logger.info(f"Calendar test: {len(calendar)} trading days in Jan 2020")

return True
except ImportError:
logger.warning("qlib not installed, skipping initialization test")
return True # Data might still be valid
except Exception as e:
logger.error(f"Failed to initialize qlib: {e}")
return False


def main():
parser = argparse.ArgumentParser(
description="Import historical stock data using chenditc/investment_data"
)
parser.add_argument(
"--qlib-dir",
default=os.path.expanduser("~/.qlib/qlib_data/cn_data"),
help="Target directory for qlib data (default: ~/.qlib/qlib_data/cn_data)"
)
parser.add_argument(
"--skip-download",
action="store_true",
help="Skip download and only verify existing data"
)
parser.add_argument(
"--verify-only",
action="store_true",
help="Only verify existing data, don't download or initialize"
)

args = parser.parse_args()
qlib_dir = args.qlib_dir

logger.info("=" * 60)
logger.info("QLib Historical Data Import Tool")
logger.info(f"Data source: chenditc/investment_data")
logger.info(f"Target directory: {qlib_dir}")
logger.info("=" * 60)

if args.verify_only:
success = verify_qlib_data(qlib_dir)
sys.exit(0 if success else 1)

if not args.skip_download:
# Step 1: Install investment_data package
if not install_investment_data():
logger.error("Failed to install investment_data package.")
sys.exit(1)

# Step 2: Download data
if not download_qlib_data(qlib_dir):
logger.error("Failed to download data. See instructions above for manual download.")
sys.exit(1)

# Step 3: Verify data
if not verify_qlib_data(qlib_dir):
logger.warning("Data verification found issues, but continuing...")

# Step 4: Test qlib initialization
if not init_qlib_with_data(qlib_dir):
logger.warning("QLib initialization test failed, but data may still be usable.")

logger.info("=" * 60)
logger.info("Data import completed!")
logger.info(f"Data location: {qlib_dir}")
logger.info("You can now start the backend server to use this data.")
logger.info("=" * 60)


if __name__ == "__main__":
main()
Loading
Loading